diff --git a/src/optimization/blockswap.py b/src/optimization/blockswap.py index f9f99d9..d8dde1d 100644 --- a/src/optimization/blockswap.py +++ b/src/optimization/blockswap.py @@ -466,73 +466,84 @@ def _wrap_io_forward(module: torch.nn.Module, module_name: str, model: torch.nn. def _patch_rope_for_blockswap(model, debugger: BlockSwapDebugger) -> None: """ - Enhance RoPE modules to handle device mismatches when using BlockSwap. - Integrates with existing stability wrapper from FP8CompatibleDiT if present. + Patch RoPE modules to handle device mismatches gracefully. + + This complements the stability wrapper from compatibility.py by adding + device-aware error handling. Only handles device/memory errors, letting + other exceptions bubble up to the stability wrapper if present. """ rope_patches = [] for name, module in model.named_modules(): if "rope" in name.lower() and hasattr(module, "get_axial_freqs"): - # Get the true original method - original_method = getattr(module, '_original_get_axial_freqs', module.get_axial_freqs) - - # Skip if already has our unified wrapper - if hasattr(module, '_rope_wrapped') and module._rope_wrapped == 'unified': + # Skip if already wrapped by blockswap + if hasattr(module, '_blockswap_wrapped') and module._blockswap_wrapped: continue - # Store original if not already stored - if not hasattr(module, '_original_get_axial_freqs'): - module._original_get_axial_freqs = original_method - - # Store reference to current method (might be stability wrapper) + # Get current method (might be stability-wrapped) current_method = module.get_axial_freqs - def device_aware_rope_wrapper(self, *args, **kwargs): - try: - # Try the current method (original or stability-wrapped) - return current_method(*args, **kwargs) - except (RuntimeError, KeyError) as e: - error_msg = str(e).lower() - - # Handle device-related errors - if any(x in error_msg for x in ["device", "memory", "allocation"]): - debugger.log(f"RoPE device issue for {name}: {e}") - - # Get current device - current_device = next(self.parameters()).device if list(self.parameters()) else torch.device("cuda") - - # Try with stability fix on current device - try: - return call_rope_with_stability(self._original_get_axial_freqs, *args, **kwargs) - except: - # Fallback to CPU computation - debugger.log(f"RoPE fallback to CPU for {name}") + # Create device-aware wrapper with proper closure handling + def make_device_aware_wrapper(module_name, current_fn): + def device_aware_rope_wrapper(self, *args, **kwargs): + try: + # Try current method (original or stability-wrapped) + return current_fn(*args, **kwargs) + except (RuntimeError, torch.cuda.OutOfMemoryError) as e: + error_msg = str(e).lower() + # Only handle device/memory specific errors + if any(x in error_msg for x in ["device", "memory", "allocation"]): + debugger.log(f"RoPE device issue for {module_name}: {e}") + + # Get current device from parameters + current_device = next(self.parameters()).device if list(self.parameters()) else torch.device("cuda") + + # Try clearing cache first (non-invasive fix) + if hasattr(current_fn, 'cache_clear'): + current_fn.cache_clear() + try: + return current_fn(*args, **kwargs) + except: + pass + + # Fallback to CPU computation with stability + debugger.log(f"RoPE fallback to CPU for {module_name}") self.cpu() try: - result = call_rope_with_stability(self._original_get_axial_freqs, *args, **kwargs) + # Use call_rope_with_stability for CPU computation + # This ensures cache is cleared and autocast disabled + original_fn = getattr(self, '_original_get_axial_freqs', current_fn) + result = call_rope_with_stability(original_fn, *args, **kwargs) - # Restore device + # Move module back to original device self.to(current_device) - # Move result to correct device + # Move result to appropriate device if it's a tensor if hasattr(result, 'to'): target_device = args[0].device if len(args) > 0 and hasattr(args[0], 'device') else current_device return result.to(target_device) return result except Exception as cpu_error: - self.to(current_device) # Always restore device + # Always restore device even on error + self.to(current_device) raise cpu_error - else: - raise - except Exception: - # For non-device errors, apply stability fix - return call_rope_with_stability(self._original_get_axial_freqs, *args, **kwargs) + else: + # Not a device error, let it bubble up + raise + + return device_aware_rope_wrapper - # Mark as unified wrapper and bind - module._rope_wrapped = 'unified' - module.get_axial_freqs = types.MethodType(device_aware_rope_wrapper, module) + # Apply wrapper + module.get_axial_freqs = types.MethodType( + make_device_aware_wrapper(name, current_method), + module + ) + module._blockswap_wrapped = True + + # Store for cleanup (use original or previously stored) + original_method = getattr(module, '_original_get_axial_freqs', current_method) rope_patches.append((module, original_method)) if rope_patches: diff --git a/src/optimization/compatibility.py b/src/optimization/compatibility.py index fec710e..1670795 100644 --- a/src/optimization/compatibility.py +++ b/src/optimization/compatibility.py @@ -15,8 +15,7 @@ def call_rope_with_stability(method, *args, **kwargs): Call RoPE method with stability fixes: 1. Clear cache if available 2. Disable autocast to prevent numerical issues - - This is the core fix that prevents artifacts in FP8/mixed precision models. + This prevents artifacts in FP8/mixed precision models. """ if hasattr(method, 'cache_clear'): method.cache_clear() @@ -44,13 +43,11 @@ class FP8CompatibleDiT(torch.nn.Module): # Only convert if not already done (e.g., when reusing cached weights) 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() - # Apply RoPE stabilization to ALL models for numerical stability - # This prevents artifacts in FP8, mixed precision, and edge cases + # Apply RoPE stabilization for numerical stability self._stabilize_rope_computations() # 🚀 FLASH ATTENTION OPTIMIZATION (Phase 2) @@ -88,11 +85,16 @@ class FP8CompatibleDiT(torch.nn.Module): def _stabilize_rope_computations(self): """ - Stabilize RoPE computations to prevent artifacts. + Add error handling to RoPE computations to prevent artifacts. - Disables autocast during RoPE frequency calculations to prevent numerical - instability that can cause artifacts in FP8, mixed precision, and even - standard models under certain conditions. + Wraps the get_axial_freqs method of RoPE modules with a try-except handler. + During normal operation, uses the original cached method for performance. + Only on exceptions (e.g., numerical instability, NaN propagation) does it + intervene by clearing the cache and retrying the computation through + call_rope_with_stability. + + This prevents artifacts in FP8, mixed precision, and edge cases while + maintaining optimal performance for normal operations. """ if not hasattr(self.dit_model, 'blocks'): return @@ -114,16 +116,14 @@ class FP8CompatibleDiT(torch.nn.Module): module._rope_wrapped = 'stability' module._original_get_axial_freqs = original_method - # Error handler that prevents NaN propagation by disabling autocast - def make_stable_rope(orig): - def stable_rope_computation(self, *args, **kwargs): - try: - return orig(*args, **kwargs) - except Exception: - return call_rope_with_stability(orig, *args, **kwargs) - return stable_rope_computation + # Error handler that prevents NaN propagation + def stable_rope_computation(self, *args, **kwargs): + try: + return original_method(*args, **kwargs) + except Exception: + return call_rope_with_stability(original_method, *args, **kwargs) - module.get_axial_freqs = types.MethodType(make_stable_rope(original_method), module) + module.get_axial_freqs = types.MethodType(stable_rope_computation, module) rope_count += 1 if rope_count > 0: