Optimized RoPE stability fix and Patch RoPE for blockswap
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user