Optimize FP8 handling to preserve native FP8 weights for memory efficiency while converting only tensors during computation
This commit is contained in:
+1
-3
@@ -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',
|
||||
|
||||
+8
-16
@@ -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):
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
'''
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user