Added logging for artifact debugging

This commit is contained in:
Adrien Toupet
2025-07-23 12:10:42 -04:00
parent e30fce3a35
commit 5fa2383b6a
4 changed files with 212 additions and 3 deletions
+26 -2
View File
@@ -24,6 +24,7 @@ class FP8CompatibleDiT(torch.nn.Module):
self.model_dtype = self._detect_model_dtype()
self.is_fp8_model = self.model_dtype in (torch.float8_e4m3fn, torch.float8_e5m2)
self.is_fp16_model = self.model_dtype == torch.float16
self._forward_count = 0
# Only convert if not already done (e.g., when reusing cached weights)
if not skip_conversion and self.is_fp8_model:
@@ -264,7 +265,8 @@ class FP8CompatibleDiT(torch.nn.Module):
return module._original_forward(x, *args, **kwargs)
def forward(self, *args, **kwargs):
"""Forward pass with minimal dtype conversion overhead
"""
Forward pass with minimal dtype conversion overhead
Conversion strategy:
- FP16 models: Keep everything in FP16 (no conversion needed)
@@ -272,15 +274,35 @@ class FP8CompatibleDiT(torch.nn.Module):
- BFloat16 models: No conversion needed
"""
# Increment forward counter
self._forward_count += 1
# === ADD BOUNDARY LOGGING ===
if self.is_fp8_model and hasattr(self, 'dit_model') and hasattr(self.dit_model, 'blocks'):
# Check if last block is FP16 (mixed precision)
try:
last_block = self.dit_model.blocks[-1]
last_block_dtype = next(last_block.parameters()).dtype
if last_block_dtype == torch.float16:
# Log the boundary transition
if self._forward_count % 100 == 0: # Log every 100 forward passes
print(f"[FP8CompatibleDiT] Mixed precision detected: FP8 model with FP16 last block")
except Exception as e:
# Silently handle if blocks structure is different
pass
# 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:
for i, arg in enumerate(args):
if isinstance(arg, torch.Tensor) and arg.dtype in fp8_dtypes:
converted_args.append(arg.to(torch.bfloat16))
# === LOG CONVERSIONS ===
if self._forward_count % 100 == 0:
print(f"[FP8CompatibleDiT] Converting arg[{i}] from {arg.dtype} to bfloat16")
else:
converted_args.append(arg)
@@ -289,6 +311,8 @@ class FP8CompatibleDiT(torch.nn.Module):
for key, value in kwargs.items():
if isinstance(value, torch.Tensor) and value.dtype in fp8_dtypes:
converted_kwargs[key] = value.to(torch.bfloat16)
if self._forward_count % 100 == 0:
print(f"[FP8CompatibleDiT] Converting kwarg[{key}] from {value.dtype} to bfloat16")
else:
converted_kwargs[key] = value