From 5fa2383b6af66bf2e1155ce07581ed751db9bbfe Mon Sep 17 00:00:00 2001 From: Adrien Toupet Date: Wed, 23 Jul 2025 12:10:42 -0400 Subject: [PATCH] Added logging for artifact debugging --- src/core/model_manager.py | 55 ++++++++++++ src/models/dit/nablocks/mmsr_block.py | 12 +++ src/optimization/blockswap.py | 120 +++++++++++++++++++++++++- src/optimization/compatibility.py | 28 +++++- 4 files changed, 212 insertions(+), 3 deletions(-) diff --git a/src/core/model_manager.py b/src/core/model_manager.py index a2181f1..c96f3a5 100644 --- a/src/core/model_manager.py +++ b/src/core/model_manager.py @@ -205,6 +205,61 @@ def configure_runner(model, base_cache_dir, preserve_vram=False, debug=False, bl # Apply BlockSwap if configured if blockswap_active: apply_block_swap_to_dit(runner, block_swap_config) + else: + # === ADD THIS LOGGING CODE HERE === + print("[DEBUG] Running WITHOUT blockswap - direct model execution") + + # Check if model has mixed precision + try: + # Get the actual model (handle FP8CompatibleDiT wrapper) + actual_model = runner.dit + if hasattr(actual_model, 'dit_model'): + actual_model = actual_model.dit_model + + if hasattr(actual_model, 'blocks'): + blocks = actual_model.blocks + block_dtypes = [] + + # Analyze each block's dtype + for i, block in enumerate(blocks): + try: + block_dtype = next(block.parameters()).dtype + block_dtypes.append((i, str(block_dtype))) + + # Special attention to last block + if i == len(blocks) - 1: + print(f"[DEBUG] Block {i} (LAST): dtype={block_dtype}") + if block_dtype == torch.float16: + print(f"[DEBUG] ⚠️ Last block is FP16 - potential mixed precision boundary") + + # Check all parameters in last block + param_dtypes = set() + for name, param in block.named_parameters(): + param_dtypes.add(str(param.dtype)) + if len(param_dtypes) > 1: + print(f"[DEBUG] ⚠️ Last block has mixed dtypes: {param_dtypes}") + except Exception as e: + print(f"[DEBUG] Could not analyze block {i}: {e}") + + # Print dtype distribution + dtype_counts = {} + for idx, dtype_str in block_dtypes: + dtype_counts[dtype_str] = dtype_counts.get(dtype_str, 0) + 1 + + print(f"[DEBUG] Block dtype distribution: {dtype_counts}") + print(f"[DEBUG] Total blocks: {len(blocks)}") + + # Check for FP8/FP16 boundary + if len(block_dtypes) > 1: + last_dtype = block_dtypes[-1][1] + second_last_dtype = block_dtypes[-2][1] if len(block_dtypes) > 1 else None + + if 'float8' in str(second_last_dtype) and 'float16' in str(last_dtype): + print(f"[DEBUG] ⚠️ MIXED PRECISION DETECTED: Blocks 0-{len(blocks)-2} are {second_last_dtype}, Block {len(blocks)-1} is {last_dtype}") + print(f"[DEBUG] This mixed precision boundary may cause artifacts without blockswap!") + + except Exception as e: + print(f"[DEBUG] Could not analyze model structure: {e}") #clear_vram_cache() return runner diff --git a/src/models/dit/nablocks/mmsr_block.py b/src/models/dit/nablocks/mmsr_block.py index 48e5482..80ee857 100644 --- a/src/models/dit/nablocks/mmsr_block.py +++ b/src/models/dit/nablocks/mmsr_block.py @@ -222,6 +222,18 @@ class NaMMSRTransformerBlock(MMWindowTransformerBlock): torch.LongTensor, torch.LongTensor, ]: + # === ADD DEBUG LOGGING FOR NON-BLOCKSWAP CASE === + if hasattr(self, '_block_idx') and not hasattr(self, '_blockswap_wrapped'): + # This means we're not using blockswap + is_last_block = hasattr(self, 'is_last_layer') and self.is_last_layer + if is_last_block: + print(f"[NON-BLOCKSWAP] Last block forward: vid.dtype={vid.dtype}, device={vid.device}") + # Check parameter dtypes + param_dtypes = set() + for param in self.parameters(): + param_dtypes.add(str(param.dtype)) + print(f"[NON-BLOCKSWAP] Last block param dtypes: {param_dtypes}") + hid_len = MMArg( cache("vid_len", lambda: vid_shape.prod(-1)), cache("txt_len", lambda: txt_shape.prod(-1)), diff --git a/src/optimization/blockswap.py b/src/optimization/blockswap.py index 480231d..a0f4ac2 100644 --- a/src/optimization/blockswap.py +++ b/src/optimization/blockswap.py @@ -351,23 +351,138 @@ def _wrap_block_forward(block: torch.nn.Module, block_idx: int, model: torch.nn. if hasattr(model, 'blocks_to_swap') and self._block_idx <= model.blocks_to_swap: t_start = time.time() if debugger and debugger.enabled else None + # === ENHANCED DEBUG LOGGING === + # Check if this is the last block + try: + total_blocks = len(model.blocks) if hasattr(model, 'blocks') else 0 + is_last_block = self._block_idx == total_blocks - 1 + except: + is_last_block = False + + # Log input tensor info - FIXED to handle actual inputs + if debugger and debugger.enabled: + input_info = [] + + # Log all tensor arguments + for i, arg in enumerate(args): + if isinstance(arg, torch.Tensor): + input_info.append(f"arg[{i}]: dtype={arg.dtype}, device={arg.device}, shape={arg.shape}") + + # For last block, log sample values and statistics + if is_last_block and i < 2: # Only for vid and txt tensors + with torch.no_grad(): + try: + # Get sample values + sample_vals = arg.flatten()[:5].tolist() + debugger.log(f"[Block {self._block_idx}] Input[{i}] sample values: {sample_vals}") + + # Get statistics + arg_float = arg.float() + mean_val = arg_float.mean().item() + std_val = arg_float.std().item() + min_val = arg_float.min().item() + max_val = arg_float.max().item() + debugger.log(f"[Block {self._block_idx}] Input[{i}] stats: mean={mean_val:.6f}, std={std_val:.6f}, min={min_val:.6f}, max={max_val:.6f}") + + # Check for anomalies + if abs(mean_val) > 10 or std_val > 100 or abs(min_val) > 1000 or abs(max_val) > 1000: + debugger.log(f"[Block {self._block_idx}] ⚠️ Input[{i}] has abnormal values!") + except Exception as e: + debugger.log(f"[Block {self._block_idx}] Could not compute input[{i}] statistics: {e}") + + # Also log kwargs if any contain tensors + for key, value in kwargs.items(): + if isinstance(value, torch.Tensor): + input_info.append(f"kwarg[{key}]: dtype={value.dtype}, device={value.device}, shape={value.shape}") + + if input_info: + debugger.log(f"[Block {self._block_idx}] {'LAST BLOCK (FP16)' if is_last_block else 'FP8'} Pre-swap inputs: {input_info}") + else: + debugger.log(f"[Block {self._block_idx}] {'LAST BLOCK (FP16)' if is_last_block else 'FP8'} Pre-swap: No tensor inputs found in {len(args)} args") + + # Log block parameter dtypes + param_dtypes = set() + for param in self.parameters(): + param_dtypes.add(str(param.dtype)) + debugger.log(f"[Block {self._block_idx}] Block param dtypes: {param_dtypes}") + # Only move to GPU if necessary current_device = next(self.parameters()).device target_device = torch.device(model.main_device) if current_device != target_device: + # Log before .to() operation + if debugger and debugger.enabled: + debugger.log(f"[Block {self._block_idx}] Moving from {current_device} to {target_device}") + self.to(model.main_device, non_blocking=model.use_non_blocking) + # Log after .to() operation + if debugger and debugger.enabled: + # Check if dtypes changed + new_param_dtypes = set() + for param in self.parameters(): + new_param_dtypes.add(str(param.dtype)) + if new_param_dtypes != param_dtypes: + debugger.log(f"[Block {self._block_idx}] WARNING: DTYPE CHANGE after .to(): {param_dtypes} -> {new_param_dtypes}") + # Synchronize if needed if hasattr(model, 'use_non_blocking') and not model.use_non_blocking: torch.cuda.synchronize() # Execute forward pass with OOM protection - output = original_forward(*args, **kwargs) + try: + output = original_forward(*args, **kwargs) + except Exception as e: + debugger.log(f"[Block {self._block_idx}] ERROR in forward: {e}") + raise + + # === ENHANCED OUTPUT LOGGING === + if debugger and debugger.enabled: + # Handle different output types (tuple or single tensor) + output_tensors = [] + if isinstance(output, tuple): + for i, out in enumerate(output): + if isinstance(out, torch.Tensor): + output_tensors.append((i, out)) + elif isinstance(output, torch.Tensor): + output_tensors.append((0, output)) + + for idx, out_tensor in output_tensors: + debugger.log(f"[Block {self._block_idx}] Output[{idx}]: dtype={out_tensor.dtype}, device={out_tensor.device}, shape={out_tensor.shape}") + + # Detailed analysis for last block + if is_last_block: + with torch.no_grad(): + has_nan = torch.isnan(out_tensor).any().item() + has_inf = torch.isinf(out_tensor).any().item() + + if has_nan or has_inf: + debugger.log(f"[Block {self._block_idx}] ⚠️ WARNING: Output has NaN={has_nan}, Inf={has_inf}") + + # Compute statistics + try: + out_float = out_tensor.float() + mean_val = out_float.mean().item() + std_val = out_float.std().item() + min_val = out_float.min().item() + max_val = out_float.max().item() + + debugger.log(f"[Block {self._block_idx}] Output stats: mean={mean_val:.6f}, std={std_val:.6f}, min={min_val:.6f}, max={max_val:.6f}") + + # Check for artifact indicators + if abs(mean_val) > 10 or std_val > 100 or abs(min_val) > 1000 or abs(max_val) > 1000: + debugger.log(f"[Block {self._block_idx}] ⚠️ ABNORMAL VALUES DETECTED - potential artifacts!") + except: + debugger.log(f"[Block {self._block_idx}] Could not compute output statistics") # Move back to offload device self.to(model.offload_device, non_blocking=model.use_non_blocking) + # === LOG POST-OFFLOAD STATE === + if debugger and debugger.enabled: + debugger.log(f"[Block {self._block_idx}] Post-offload complete") + # Log timing if debugger is available if debugger and t_start is not None: debugger.log_swap_time( @@ -381,6 +496,9 @@ def _wrap_block_forward(block: torch.nn.Module, block_idx: int, model: torch.nn. if torch.cuda.memory_allocated() > torch.cuda.get_device_properties(0).total_memory * 0.9: mm.soft_empty_cache() else: + # === LOG NON-SWAPPED BLOCKS === + if debugger and debugger.enabled and self._block_idx % 10 == 0: # Log every 10th non-swapped block + debugger.log(f"[Block {self._block_idx}] Running without swap (block > blocks_to_swap)") output = original_forward(*args, **kwargs) return output diff --git a/src/optimization/compatibility.py b/src/optimization/compatibility.py index 08e5c70..5e382dd 100644 --- a/src/optimization/compatibility.py +++ b/src/optimization/compatibility.py @@ -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