Optimized RoPE stability fix and Patch RoPE for blockswap

This commit is contained in:
Adrien Toupet
2025-07-24 09:32:18 -04:00
parent b551e32ae1
commit bc8f0aa6d0
2 changed files with 73 additions and 62 deletions
+55 -44
View File
@@ -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:
+18 -18
View File
@@ -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: