diff --git a/src/__init__.py b/src/__init__.py index 5626235..e0b8cc0 100644 --- a/src/__init__.py +++ b/src/__init__.py @@ -70,8 +70,6 @@ if MODULES_AVAILABLE['performance']: if MODULES_AVAILABLE['compatibility']: from src.optimization.compatibility import ( FP8CompatibleDiT, - apply_fp8_compatibility_hooks, - remove_compatibility_hooks, ) @@ -126,7 +124,7 @@ __all__ = [ 'validate_video_format', 'ensure_4n_plus_1_format', 'calculate_padding_requirements', 'apply_wavelet_reconstruction', 'temporal_consistency_check', # Compatibility - 'FP8CompatibleDiT', 'apply_fp8_compatibility_hooks', 'remove_compatibility_hooks', + 'FP8CompatibleDiT', # Core Model & Generation & Infer 'configure_runner', 'load_quantized_state_dict', 'configure_dit_model_inference', 'configure_vae_model_inference', diff --git a/src/core/infer.py b/src/core/infer.py index 3d6d2dd..129b451 100644 --- a/src/core/infer.py +++ b/src/core/infer.py @@ -306,23 +306,15 @@ class VideoDiffusionInfer(): if cfg_scale is None: cfg_scale = self.config.diffusion.cfg.scale - # πŸš€ OPTIMISATION: DΓ©tecter le dtype du modΓ¨le pour performance optimale - model_dtype = next(self.dit.parameters()).dtype + # πŸš€ OPTIMISATION: Use BFloat16 autocast for all models + # - FP8 models: BFloat16 required for arithmetic operations + # - FP16 models: BFloat16 provides better numerical stability and prevents black frames + # - BFloat16 models: Already optimal + target_dtype = torch.bfloat16 + if self.debug: - print(f"🎯 model_dtype: {model_dtype}") - # Adapter les dtypes selon le modΓ¨le - if model_dtype in (torch.float8_e4m3fn, torch.float8_e5m2): - # FP8 natif: utiliser BFloat16 pour les calculs intermΓ©diaires (compatible) - target_dtype = torch.float16 - #print(f"πŸš€ FP8 model detected: using BFloat16 for intermediate calculations") - elif model_dtype == torch.float16: - target_dtype = torch.bfloat16 - #print(f"🎯 FP16 model: using FP16 pipeline") - else: - target_dtype = torch.bfloat16 - #print(f"🎯 BFloat16 model: using BFloat16 pipeline") - if self.debug: - print(f"🎯 target_dtype: {target_dtype}") + model_dtype = next(self.dit.parameters()).dtype + print(f"🎯 Model dtype: {model_dtype}, using {target_dtype} for autocast") # Text embeddings. assert type(texts_pos[0]) is type(texts_neg[0]) if isinstance(texts_pos[0], str): diff --git a/src/models/dit/modulation.py b/src/models/dit/modulation.py index 47c9f11..b1db4d9 100644 --- a/src/models/dit/modulation.py +++ b/src/models/dit/modulation.py @@ -86,11 +86,24 @@ class AdaSingle(nn.Module): getattr(self, f"{layer}_scale"), getattr(self, f"{layer}_gate"), ) + + # Handle potential FP8 parameters - convert to computation dtype + if hasattr(torch, 'float8_e4m3fn'): + fp8_types = (torch.float8_e4m3fn, torch.float8_e5m2) + + # Convert FP8 parameters to BFloat16 for arithmetic operations + if shiftB.dtype in fp8_types: + shiftB = shiftB.to(torch.bfloat16) + if scaleB.dtype in fp8_types: + scaleB = scaleB.to(torch.bfloat16) + if gateB.dtype in fp8_types: + gateB = gateB.to(torch.bfloat16) if mode == "in": return hid.mul_(scaleA + scaleB).add_(shiftA + shiftB) if mode == "out": return hid.mul_(gateA + gateB) + raise NotImplementedError def extra_repr(self) -> str: diff --git a/src/models/dit/normalization.py b/src/models/dit/normalization.py index e03d396..6f074e4 100644 --- a/src/models/dit/normalization.py +++ b/src/models/dit/normalization.py @@ -86,6 +86,13 @@ class CustomRMSNorm(nn.Module): normalized = input / rms if self.elementwise_affine: + # Convert FP8 weight to BFloat16 for arithmetic operations + if hasattr(torch, 'float8_e4m3fn'): + fp8_types = (torch.float8_e4m3fn, torch.float8_e5m2) + if self.weight.dtype in fp8_types: + weight = self.weight.to(torch.bfloat16) + return normalized * weight + return normalized * self.weight return normalized diff --git a/src/models/dit_v2/modulation.py b/src/models/dit_v2/modulation.py index 9a99025..ba74c6e 100644 --- a/src/models/dit_v2/modulation.py +++ b/src/models/dit_v2/modulation.py @@ -92,25 +92,26 @@ class AdaSingle(nn.Module): getattr(self, f"{layer}_gate", None), ) - # πŸš€ FP8 COMPATIBILITY: Convert parameters to match embedding dtype - # This prevents "Promotion for Float8 Types is not supported" errors - target_dtype = shiftA.dtype - + # Handle potential FP8 parameters - convert to computation dtype + if hasattr(torch, 'float8_e4m3fn'): + fp8_types = (torch.float8_e4m3fn, torch.float8_e5m2) + + # Convert FP8 parameters to BFloat16 for arithmetic operations + if shiftB is not None and shiftB.dtype in fp8_types: + shiftB = shiftB.to(torch.bfloat16) + if scaleB is not None and scaleB.dtype in fp8_types: + scaleB = scaleB.to(torch.bfloat16) + if gateB is not None and gateB.dtype in fp8_types: + gateB = gateB.to(torch.bfloat16) + if mode == "in": - # Convert parameters to match embedding dtype for FP8 compatibility - if scaleB is not None and scaleB.dtype != target_dtype: - scaleB = scaleB.to(target_dtype) - if shiftB is not None and shiftB.dtype != target_dtype: - shiftB = shiftB.to(target_dtype) - return hid.mul_(scaleA + scaleB).add_(shiftA + shiftB) - if mode == "out": - # Convert gate parameter to match embedding dtype for FP8 compatibility - if gateB is not None and gateB.dtype != target_dtype: - gateB = gateB.to(target_dtype) - - return hid.mul_(gateA + gateB) + if gateB is not None: + return hid.mul_(gateA + gateB) + else: + # If no gate parameter, just use the embedding gate + return hid.mul_(gateA) raise NotImplementedError diff --git a/src/models/dit_v2/normalization.py b/src/models/dit_v2/normalization.py index 980854f..49a733a 100644 --- a/src/models/dit_v2/normalization.py +++ b/src/models/dit_v2/normalization.py @@ -97,12 +97,14 @@ class CustomRMSNorm(nn.Module): normalized = input / rms if self.elementwise_affine: - # πŸš€ FP8 COMPATIBILITY: Convert weight to match normalized dtype - # This prevents "Promotion for Float8 Types is not supported" errors - weight = self.weight - if weight.dtype != normalized.dtype: - weight = weight.to(normalized.dtype) - return normalized * weight + # Convert FP8 weight to BFloat16 for arithmetic operations + if hasattr(torch, 'float8_e4m3fn'): + fp8_types = (torch.float8_e4m3fn, torch.float8_e5m2) + if self.weight.dtype in fp8_types: + weight = self.weight.to(torch.bfloat16) + return normalized * weight + + return normalized * self.weight return normalized diff --git a/src/optimization/__init__.py b/src/optimization/__init__.py index 55fca1f..faf8afe 100644 --- a/src/optimization/__init__.py +++ b/src/optimization/__init__.py @@ -23,8 +23,6 @@ from .performance import ( # Compatibility functions and classes from .compatibility import ( FP8CompatibleDiT, - apply_fp8_compatibility_hooks, - remove_compatibility_hooks, ) __all__ = [ @@ -43,7 +41,5 @@ __all__ = [ # Compatibility "FP8CompatibleDiT", - "apply_fp8_compatibility_hooks", - "remove_compatibility_hooks", ] ''' \ No newline at end of file diff --git a/src/optimization/compatibility.py b/src/optimization/compatibility.py index 14699c4..08e5c70 100644 --- a/src/optimization/compatibility.py +++ b/src/optimization/compatibility.py @@ -5,7 +5,6 @@ Contains FP8/FP16 compatibility layers and wrappers for different model architec Extracted from: seedvr2.py (lines 1045-1630) """ -import time import torch from typing import List, Tuple, Union, Any, Optional @@ -15,7 +14,7 @@ class FP8CompatibleDiT(torch.nn.Module): Wrapper for DiT models with automatic compatibility management + advanced optimizations - FP8: Keeps native FP8 parameters, converts inputs/outputs - FP16: Uses native FP16 - - RoPE: ALWAYS forced to BFloat16 for maximum compatibility + - RoPE: Converted from FP8 to BFloat16 only when detected as FP8 - Flash Attention: Automatic optimization of attention layers """ @@ -27,26 +26,14 @@ class FP8CompatibleDiT(torch.nn.Module): self.is_fp16_model = self.model_dtype == torch.float16 # Only convert if not already done (e.g., when reusing cached weights) - if not skip_conversion: - # Detect model type - is_nadit_7b = self._is_nadit_model() # NaDiT 7B (dit/nadit) - is_nadit_v2_3b = self._is_nadit_v2_model() # NaDiT v2 3B (dit_v2/nadit) + if not skip_conversion and self.is_fp8_model: + # Only FP8 models need RoPE frequency conversion + # FP16 and BFloat16 models work as-is without conversion + model_variant = "7B" if self._is_nadit_model() else "3B" if self._is_nadit_v2_model() else "Unknown" + print(f"🎯 Detected NaDiT {model_variant} FP8 - Converting RoPE freqs for FP8 compatibility") + self._convert_rope_freqs() - if is_nadit_7b: - # 🎯 CRITICAL FIX: ALL NaDiT 7B models (FP8 AND FP16) require BFloat16 conversion - # 7B architecture has dtype compatibility issues regardless of storage format - if self.is_fp8_model: - print("🎯 Detected NaDiT 7B FP8 - Converting all parameters to BFloat16") - self._force_nadit_bfloat16() - else: - print("🎯 Detected NaDiT 7B FP16") - - - elif self.is_fp8_model and is_nadit_v2_3b: - # For NaDiT v2 3B FP8: Convert ALL model to BFloat16 - print("🎯 Detected NaDiT v2 3B FP8 - Converting all parameters to BFloat16") - self._force_nadit_bfloat16() - + # πŸš€ FLASH ATTENTION OPTIMIZATION (Phase 2) self._apply_flash_attention_optimization() @@ -59,80 +46,27 @@ class FP8CompatibleDiT(torch.nn.Module): def _is_nadit_model(self) -> bool: """Detect if this is a NaDiT model (7B) with precise logic""" - # 🎯 PRIMARY METHOD: Check emb_scale attribute (specific to 7B) - # This is the most reliable criterion to distinguish 7B vs 3B - if hasattr(self.dit_model, 'emb_scale'): - return True - - # 🎯 SECONDARY METHOD: Check module path for NaDiT 7B (dit/nadit, not dit_v2) + # Check module path for dit (not dit_v2) model_module = str(self.dit_model.__class__.__module__).lower() - if 'dit.nadit' in model_module and 'dit_v2' not in model_module: - return True - - return False - + return 'dit.nadit' in model_module and 'dit_v2' not in model_module + def _is_nadit_v2_model(self) -> bool: """Detect if this is a NaDiT v2 model (3B) with precise logic""" - # 🎯 PRIMARY METHOD: Check module path for NaDiT v2 (dit_v2/nadit) + # Check module path for dit_v2 model_module = str(self.dit_model.__class__.__module__).lower() - if 'dit_v2' in model_module: - return True + return 'dit_v2' in model_module - # 🎯 SECONDARY METHOD: Check specific 3B structure - # NaDiT v2 3B has vid_in, txt_in, emb_in but NO emb_scale - if (hasattr(self.dit_model, 'vid_in') and - hasattr(self.dit_model, 'txt_in') and - hasattr(self.dit_model, 'emb_in') and - not hasattr(self.dit_model, 'emb_scale')): # Absence of emb_scale = 3B - return True - - return False - - def _force_rope_bfloat16(self) -> None: - """🎯 Force ALL RoPE modules to BFloat16 for maximum compatibility""" - rope_count = 0 - for name, module in self.dit_model.named_modules(): - # Identify RoPE modules by name or type - if any(keyword in name.lower() for keyword in ['rope', 'rotary', 'embedding']): - # Convert all parameters of this module to BFloat16 - for param_name, param in module.named_parameters(): - if param.dtype != torch.bfloat16: - param.data = param.data.to(torch.bfloat16) - rope_count += 1 - - # Also convert buffers (non-trainable parameters) - for buffer_name, buffer in module.named_buffers(): - if buffer.dtype != torch.bfloat16: - buffer.data = buffer.data.to(torch.bfloat16) - rope_count += 1 - - def _force_nadit_bfloat16(self) -> None: - """🎯 Force ALL NaDiT parameters to BFloat16 to avoid promotion errors""" - print("πŸ”§ Converting ALL NaDiT parameters to BFloat16 for type compatibility...") - t = time.time() - converted_count = 0 - original_dtype = None - - # Convert ALL parameters to BFloat16 (FP8, FP16, etc.) - for name, param in self.dit_model.named_parameters(): - if original_dtype is None: - original_dtype = param.dtype - if param.dtype != torch.bfloat16: - param.data = param.data.to(torch.bfloat16) - converted_count += 1 - - # Also convert buffers - for name, buffer in self.dit_model.named_buffers(): - if buffer.dtype != torch.bfloat16: - buffer.data = buffer.data.to(torch.bfloat16) - converted_count += 1 - - print(f" βœ… Converted {converted_count} parameters/buffers from {original_dtype} to BFloat16") - - # Update detected dtype - self.model_dtype = torch.bfloat16 - self.is_fp8_model = False # Model is no longer FP8 after conversion - + def _convert_rope_freqs(self) -> None: + """Convert RoPE frequency buffers for FP8 compatibility""" + converted = 0 + for module in self.dit_model.modules(): + if 'RotaryEmbedding' in type(module).__name__: + if hasattr(module, 'rope') and hasattr(module.rope, 'freqs'): + if module.rope.freqs.dtype in (torch.float8_e4m3fn, torch.float8_e5m2): + module.rope.freqs.data = module.rope.freqs.to(torch.bfloat16) + converted += 1 + print(f" βœ… Converted {converted} RoPE frequency buffers") + def _apply_flash_attention_optimization(self) -> None: """πŸš€ FLASH ATTENTION OPTIMIZATION - 30-50% speedup of attention layers""" attention_layers_optimized = 0 @@ -330,81 +264,46 @@ class FP8CompatibleDiT(torch.nn.Module): return module._original_forward(x, *args, **kwargs) def forward(self, *args, **kwargs): - """Forward pass with intelligent type management according to architecture""" - is_nadit_7b = self._is_nadit_model() - is_nadit_v2_3b = self._is_nadit_v2_model() + """Forward pass with minimal dtype conversion overhead - # Input conversion according to architecture - if is_nadit_7b or is_nadit_v2_3b: - # For NaDiT models (7B and v2 3B): Everything to BFloat16 + Conversion strategy: + - FP16 models: Keep everything in FP16 (no conversion needed) + - FP8 models: Convert FP8 tensors to BFloat16 (required for arithmetic) + - BFloat16 models: No conversion needed + """ + + # Only convert if we have an FP8 model for arithmetic operations + if self.is_fp8_model: + fp8_dtypes = (torch.float8_e4m3fn, torch.float8_e5m2) + + # Convert args converted_args = [] for arg in args: - if isinstance(arg, torch.Tensor): - if arg.dtype in (torch.float32, torch.float8_e4m3fn, torch.float8_e5m2): - converted_args.append(arg.to(torch.bfloat16)) - else: - converted_args.append(arg) + if isinstance(arg, torch.Tensor) and arg.dtype in fp8_dtypes: + converted_args.append(arg.to(torch.bfloat16)) else: converted_args.append(arg) + # Convert kwargs converted_kwargs = {} for key, value in kwargs.items(): - if isinstance(value, torch.Tensor): - if value.dtype in (torch.float32, torch.float8_e4m3fn, torch.float8_e5m2): - converted_kwargs[key] = value.to(torch.bfloat16) - else: - converted_kwargs[key] = value + if isinstance(value, torch.Tensor) and value.dtype in fp8_dtypes: + converted_kwargs[key] = value.to(torch.bfloat16) else: converted_kwargs[key] = value args = tuple(converted_args) kwargs = converted_kwargs - else: - # For standard models: Conversion according to model dtype - if self.is_fp8_model: - # Convert FP8 β†’ BFloat16 for calculations - converted_args = [] - for arg in args: - if isinstance(arg, torch.Tensor) and arg.dtype in (torch.float8_e4m3fn, torch.float8_e5m2): - converted_args.append(arg.to(torch.bfloat16)) - else: - converted_args.append(arg) - - converted_kwargs = {} - for key, value in kwargs.items(): - if isinstance(value, torch.Tensor) and value.dtype in (torch.float8_e4m3fn, torch.float8_e5m2): - converted_kwargs[key] = value.to(torch.bfloat16) - else: - converted_kwargs[key] = value - - args = tuple(converted_args) - kwargs = converted_kwargs - elif self.is_fp16_model: - # Convert Float32 β†’ FP16 for FP16 models - converted_args = [] - for arg in args: - if isinstance(arg, torch.Tensor) and arg.dtype == torch.float32: - converted_args.append(arg.to(torch.float16)) - else: - converted_args.append(arg) - - converted_kwargs = {} - for key, value in kwargs.items(): - if isinstance(value, torch.Tensor) and value.dtype == torch.float32: - converted_kwargs[key] = value.to(torch.float16) - else: - converted_kwargs[key] = value - - args = tuple(converted_args) - kwargs = converted_kwargs + # Execute forward pass try: return self.dit_model(*args, **kwargs) except Exception as e: - print(f"❌ Error in forward pass: {e}") - print(f" Model type: NaDiT 7B={is_nadit_7b}, NaDiT v2 3B={is_nadit_v2_3b}") - print(f" Args dtypes: {[arg.dtype if isinstance(arg, torch.Tensor) else type(arg) for arg in args]}") - print(f" Kwargs dtypes: {[(k, v.dtype if isinstance(v, torch.Tensor) else type(v)) for k, v in kwargs.items()]}") + print(f"❌ Forward pass error: {e}") + if self.is_fp8_model: + print(f" FP8 model - converted FP8 tensors to BFloat16") + else: + print(f" {self.model_dtype} model - no conversion applied") raise def __getattr__(self, name): @@ -421,65 +320,4 @@ class FP8CompatibleDiT(torch.nn.Module): if hasattr(self, 'dit_model'): setattr(self.dit_model, name, value) else: - super().__setattr__(name, value) - - -def apply_fp8_compatibility_hooks(model: torch.nn.Module) -> List[Tuple[str, Any]]: - """ - Hook system to intercept problematic FP8 modules - Alternative if the wrapper is not sufficient. - - Args: - model: Model to apply hooks to - - Returns: - List of (module_name, hook) tuples for cleanup - """ - def create_fp8_safe_hook(original_dtype: torch.dtype): - def hook_fn(module, input, output): - # Convert FP8 output β†’ BFloat16 if necessary for compatibility - if isinstance(output, torch.Tensor) and output.dtype in (torch.float8_e4m3fn, torch.float8_e5m2): - # Temporarily keep in BFloat16 to avoid downstream errors - return output.to(torch.bfloat16) - elif isinstance(output, (tuple, list)): - # Handle multiple outputs - converted_output = [] - for item in output: - if isinstance(item, torch.Tensor) and item.dtype in (torch.float8_e4m3fn, torch.float8_e5m2): - converted_output.append(item.to(torch.bfloat16)) - else: - converted_output.append(item) - return type(output)(converted_output) - return output - return hook_fn - - # Apply hooks to critical modules - problematic_modules = [] - for name, module in model.named_modules(): - # Identify RoPE and attention modules that cause FP8 problems - if any(keyword in name.lower() for keyword in ['rope', 'rotary', 'attention', 'mmattn']): - if hasattr(module, 'register_forward_hook'): - hook = module.register_forward_hook(create_fp8_safe_hook(torch.float8_e4m3fn)) - problematic_modules.append((name, hook)) - - print(f"πŸ”§ Applied FP8 compatibility hooks to {len(problematic_modules)} modules") - return problematic_modules - - -def remove_compatibility_hooks(hooks: List[Tuple[str, Any]]) -> None: - """ - Remove previously applied compatibility hooks - - Args: - hooks: List of (module_name, hook) tuples from apply_fp8_compatibility_hooks - """ - removed_count = 0 - for name, hook in hooks: - try: - hook.remove() - removed_count += 1 - except Exception as e: - print(f"⚠️ Failed to remove hook from {name}: {e}") - - print(f"🧹 Removed {removed_count}/{len(hooks)} compatibility hooks") - + super().__setattr__(name, value) \ No newline at end of file