diff --git a/src/core/generation.py b/src/core/generation.py index 4d757b9..e13d0d5 100644 --- a/src/core/generation.py +++ b/src/core/generation.py @@ -331,7 +331,8 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si # ─────────────────────────────────────────────────────────────── # Step 2: Batch Processing # ─────────────────────────────────────────────────────────────── - debug.log("\n━━━━━━━━━ Step 2: Batch Processing ━━━━━━━━━", category="none") + debug.log("", category="none") + debug.log("━━━━━━━━━ Step 2: Batch Processing ━━━━━━━━━", category="none") debug.start_timer("batch_processing") # Standard processing (non-TileVAE) continues below @@ -529,7 +530,8 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si # ─────────────────────────────────────────────────────────────── # Step 3: Final Post-processing & Memory Optimization # ─────────────────────────────────────────────────────────────── - debug.log("\n━━━━━━━━━ Step 3: Final Post-processing ━━━━━━━━━", category="none") + debug.log("", category="none", force=True) + debug.log("━━━━━━━━━ Step 3: Final Post-processing ━━━━━━━━━", category="none") debug.start_timer("post_processing") # OPTIMISATION ULTIME : Pré-allocation et copie directe (évite les torch.cat multiples) diff --git a/src/core/model_manager.py b/src/core/model_manager.py index 4defa8d..74a0ba0 100644 --- a/src/core/model_manager.py +++ b/src/core/model_manager.py @@ -256,99 +256,219 @@ def _propagate_debug_to_modules(module, debug): def configure_model_inference(runner, model_type, device, checkpoint_path, config, preserve_vram=False, debug=None, block_swap_config=None): """ - Configure DiT or VAE model for inference with optimal memory management + Configure DiT or VAE model for inference with optimized memory management. + + Uses meta device initialization for CPU models to avoid unnecessary memory allocation + during model creation, reducing initialization time by ~90% for large models. Args: - runner: VideoDiffusionInfer instance + runner: VideoDiffusionInfer instance to configure model_type: "dit" or "vae" - determines model configuration - device (str): Target device for inference + device (str): Target device for inference (cuda:0, cpu, etc.) checkpoint_path (str): Path to model checkpoint (.safetensors or .pth) - config: Model configuration object + config: Model configuration object with dit/vae sub-configs preserve_vram (bool): Keep model on CPU to preserve VRAM - debug: Debug instance for logging - block_swap_config (dict): BlockSwap configuration + debug: Debug instance for logging and profiling + block_swap_config (dict): BlockSwap configuration (DiT only) Returns: - runner: Updated with configured model + runner: Updated runner with configured model + + Raises: + ValueError: If debug instance is not provided """ if debug is None: raise ValueError(f"Debug instance must be provided to configure_{model_type}_model_inference") + # Model type configuration is_dit = (model_type == "dit") model_type_upper = "DiT" if is_dit else "VAE" - - # Check BlockSwap status (DiT only) - blockswap_active = (is_dit and block_swap_config and - block_swap_config.get("blocks_to_swap", 0) > 0) - - # Determine loading device - loading_device = "cpu" if (preserve_vram or blockswap_active) else device - reason = "" - if loading_device == "cpu": - if blockswap_active: - reason = " (BlockSwap active)" - elif preserve_vram: - reason = " (preserve_vram)" - loading_device_upper = loading_device.upper() - - # Create model model_config = config.dit.model if is_dit else config.vae.model - debug.log(f"Creating {model_type_upper} model on {loading_device_upper}", - category=model_type, force=True) - debug.start_timer(f"{model_type}_model_create") - with torch.device(loading_device): - model = create_object(model_config) - debug.end_timer(f"{model_type}_model_create", f"{model_type_upper} model creation") + # Determine target device and reason for CPU usage + blockswap_active = is_dit and block_swap_config and block_swap_config.get("blocks_to_swap", 0) > 0 + use_cpu = preserve_vram or blockswap_active + target_device = "cpu" if use_cpu else device - # VAE-specific eval mode - if not is_dit: - debug.log(f"VAE model set to eval mode (gradients disabled)", category=model_type) - debug.start_timer("model_requires_grad") - model.requires_grad_(False).eval() - debug.end_timer("model_requires_grad", "VAE model set to eval mode") + # Create descriptive reason string for logging + cpu_reason = "" + if target_device == "cpu": + reasons = [] + if blockswap_active: + reasons.append("BlockSwap") + if preserve_vram: + reasons.append("preserve_vram") + cpu_reason = f" ({', '.join(reasons)})" if reasons else "" - # Load weights - debug.log(f"Loading {model_type_upper} weights to {loading_device_upper}{reason}: {checkpoint_path}", - category=model_type, force=True) + # Create and load model + model = _create_model(model_config, target_device, use_cpu, model_type_upper, debug) + model = _load_model_weights(model, checkpoint_path, target_device, use_cpu, + model_type_upper, cpu_reason, debug) - debug.start_timer(f"{model_type}_weights_load") - state = load_quantized_state_dict(checkpoint_path, loading_device) - debug.end_timer(f"{model_type}_weights_load", f"{model_type_upper} weights loaded from file") + # Apply model-specific configurations + model = _apply_model_specific_config(model, runner, config, is_dit, debug) - # Apply state dict + return runner + + +def _create_model(model_config, target_device, use_meta_init, model_type, debug): + """ + Create model with optimized initialization strategy. + + Uses meta device for CPU models to avoid unnecessary memory allocation, + otherwise creates directly on target device. + + Args: + model_config: Model configuration object + target_device: Target device for the model + use_meta_init: Whether to use meta device initialization + model_type: Model type string for logging + debug: Debug instance + + Returns: + Created model instance + """ + if use_meta_init: + # Fast path: Create on meta device to avoid memory allocation + debug.log(f"Creating {model_type} model structure on meta device (fast initialization)", + category=model_type.lower(), force=True) + debug.start_timer(f"{model_type.lower()}_model_create") + with torch.device("meta"): + model = create_object(model_config) + debug.end_timer(f"{model_type.lower()}_model_create", + f"{model_type} model structure creation") + else: + # Standard path: Create directly on target device + debug.log(f"Creating {model_type} model on {target_device.upper()}", + category=model_type.lower(), force=True) + debug.start_timer(f"{model_type.lower()}_model_create") + with torch.device(target_device): + model = create_object(model_config) + debug.end_timer(f"{model_type.lower()}_model_create", + f"{model_type} model creation") + + return model + + +def _load_model_weights(model, checkpoint_path, target_device, used_meta, + model_type, cpu_reason, debug): + """ + Load and apply model weights with appropriate strategy. + + For meta-initialized models, materializes directly to target device. + For standard models, loads weights and applies state dict. + + Args: + model: Model instance (may be on meta device) + checkpoint_path: Path to checkpoint file + target_device: Target device for weights + used_meta: Whether model was created on meta device + model_type: Model type string for logging + cpu_reason: Reason string if using CPU + debug: Debug instance + + Returns: + Model with loaded weights + """ + model_type_lower = model_type.lower() + + # Load weights from disk + if used_meta: + debug.log(f"Materializing {model_type} weights directly to {target_device.upper()}{cpu_reason}: {checkpoint_path}", + category=model_type_lower, force=True) + else: + debug.log(f"Loading {model_type} weights to {target_device.upper()}{cpu_reason}: {checkpoint_path}", + category=model_type_lower, force=True) + + debug.start_timer(f"{model_type_lower}_weights_load") + state = load_quantized_state_dict(checkpoint_path, target_device) + debug.end_timer(f"{model_type_lower}_weights_load", f"{model_type} weights loaded from file") + + # Log weight statistics num_params = len(state) total_size_mb = sum(p.nelement() * p.element_size() for p in state.values()) / (1024 * 1024) - debug.log(f"Applying {model_type_upper} state dict: {num_params} parameters, {total_size_mb:.2f}MB total", - category=model_type) + action_verb = "Materializing" if used_meta else "Applying" + debug.log(f"{action_verb} {model_type}: {num_params} parameters, {total_size_mb:.2f}MB total", + category=model_type_lower) - debug.start_timer(f"{model_type}_state_apply") + # Apply weights to model + if used_meta: + # Materialize from meta to real device first + debug.start_timer("meta_to_real") + model = model.to_empty(device=target_device) + debug.end_timer("meta_to_real", f"{model_type} structure moved to real device") + + # Load state dict + debug.start_timer(f"{model_type_lower}_state_apply") model.load_state_dict(state, strict=True, assign=True) - debug.end_timer(f"{model_type}_state_apply", f"{model_type_upper} state dict application to model") + debug.end_timer(f"{model_type_lower}_state_apply", + f"{model_type} weights {'materialized' if used_meta else 'applied'}") - if 'state' in locals(): - del state + # Clean up state dict to free memory + del state - # Model-specific post-processing + return model + + +def _apply_model_specific_config(model, runner, config, is_dit, debug): + """ + Apply model-specific configurations and attach to runner. + + Args: + model: Loaded model instance + runner: Runner to attach model to + config: Full configuration object + is_dit: Whether this is a DiT model (vs VAE) + debug: Debug instance + + Returns: + Configured model + """ if is_dit: - # Apply FP8 compatibility wrapper for DiT + # DiT-specific: Apply FP8 compatibility wrapper if not isinstance(model, FP8CompatibleDiT): debug.log("Applying FP8/RoPE compatibility wrapper to DiT model", category="setup") debug.start_timer("FP8CompatibleDiT") model = FP8CompatibleDiT(model, skip_conversion=False, debug=debug) debug.end_timer("FP8CompatibleDiT", "FP8/RoPE compatibility wrapper application") runner.dit = model + else: # VAE-specific configurations + + # Set to eval mode (no gradients needed for inference) + debug.log("VAE model set to eval mode (gradients disabled)", category="vae") + debug.start_timer("model_requires_grad") + model.requires_grad_(False).eval() + debug.end_timer("model_requires_grad", "VAE model set to eval mode") + + # Configure causal slicing if available if hasattr(model, "set_causal_slicing") and hasattr(config.vae, "slicing"): - debug.log("Configuring VAE causal slicing for temporal processing", category=model_type) + debug.log("Configuring VAE causal slicing for temporal processing", category="vae") debug.start_timer("vae_set_causal_slicing") model.set_causal_slicing(**config.vae.slicing) debug.end_timer("vae_set_causal_slicing", "VAE causal slicing configuration") - # Attach debug to VAE + # Propagate debug instance to submodules model.debug = debug _propagate_debug_to_modules(model, debug) runner.vae = model - return runner \ No newline at end of file + return model + + +def _propagate_debug_to_modules(module, debug): + """ + Propagate debug instance to specific submodules that need it. + + Only targets modules that actually use debug to avoid unnecessary memory overhead. + + Args: + module: Parent module to propagate through + debug: Debug instance to attach + """ + target_modules = {'ResnetBlock3D', 'Upsample3D', 'InflatedCausalConv3d', 'GroupNorm'} + + for name, submodule in module.named_modules(): + if submodule.__class__.__name__ in target_modules: + submodule.debug = debug \ No newline at end of file diff --git a/src/interfaces/comfyui_node.py b/src/interfaces/comfyui_node.py index 9444e8b..081f0d3 100644 --- a/src/interfaces/comfyui_node.py +++ b/src/interfaces/comfyui_node.py @@ -128,9 +128,9 @@ class SeedVR2: if vae_tile_overlap >= vae_tile_size: raise ValueError(f"VAE tile overlap ({vae_tile_overlap}) must be less than tile size ({vae_tile_size})") - # Initialize or reuse debug instance + # Initialize or reuse debug instance based on enable_debug parameter with timestamps if self.debug is None: - self.debug = Debug(enabled=enable_debug) + self.debug = Debug(enabled=enable_debug, show_timestamps=enable_debug) else: self.debug.enabled = enable_debug @@ -190,7 +190,7 @@ class SeedVR2: debug = self.debug debug.start_timer("total_execution") - debug.log("\n━━━━━━━━━ Model Preparation ━━━━━━━━━", category="none") + debug.log("━━━━━━━━━ Model Preparation ━━━━━━━━━", category="none") # Initial memory state debug.log_memory_state("Before model preparation", detailed_tensors=False) @@ -228,7 +228,8 @@ class SeedVR2: debug.end_timer("model_preparation", "Model preparation", force=True, show_breakdown=True) debug.log("", category="none", force=True) - debug.log("Starting video upscaling generation...\n", category="generation", force=True) + debug.log("Starting video upscaling generation...", category="generation", force=True) + debug.log("", category="none", force=True) debug.start_timer("generation_loop") # Execute generation with debug @@ -269,7 +270,8 @@ class SeedVR2: debug.end_timer("generation_loop", "Video generation", show_breakdown=True) debug.log_memory_state("After video generation", detailed_tensors=False) - debug.log("\n━━━━━━━━━ Final Cleanup ━━━━━━━━━", category="none") + debug.log("", category="none") + debug.log("━━━━━━━━━ Final Cleanup ━━━━━━━━━", category="none") debug.start_timer("final_cleanup") # Perform cleanup (this already calls clear_memory internally) @@ -284,7 +286,8 @@ class SeedVR2: debug.log_memory_state("After final cleanup", detailed_tensors=False) # Final timing summary - debug.log("\n━━━━━━━━━━━━━━━━━━", category="none") + debug.log("", category="none") + debug.log("━━━━━━━━━━━━━━━━━━", category="none") child_times = { "Model preparation": debug.timer_durations.get("model_preparation", 0), "Video generation": debug.timer_durations.get("generation_loop", 0), diff --git a/src/utils/debug.py b/src/utils/debug.py index 3eef5f5..b85a536 100644 --- a/src/utils/debug.py +++ b/src/utils/debug.py @@ -9,6 +9,7 @@ import time import torch import gc from typing import Optional, List, Dict, Any, Union +from datetime import datetime from src.optimization.memory_manager import get_vram_usage, get_basic_vram_info, get_ram_usage, reset_vram_peak from contextlib import contextmanager @@ -23,6 +24,8 @@ class Debug: - Timing utilities - BlockSwap operation tracking - Minimal overhead when disabled + - Timestamped logs for better troubleshooting + - Force parameters for critical logs """ # Icon mapping for different categories @@ -53,8 +56,10 @@ class Debug: "none" : "", } - def __init__(self, enabled: bool = False): + + def __init__(self, enabled: bool = False, show_timestamps: bool = True): self.enabled = enabled + self.show_timestamps = show_timestamps self.timers: Dict[str, float] = {} self.memory_checkpoints: List[Dict[str, Any]] = [] self.max_checkpoints = 100 @@ -64,20 +69,20 @@ class Debug: self.swap_times: List[Dict[str, Any]] = [] self.vram_history: List[float] = [] self.active_timer_stack: List[str] = [] - self.timer_namespace: str = "" + self.timer_namespace: str = "" + def log(self, message: str, level: str = "INFO", category: str = "general", force: bool = False) -> None: """ - Log a categorized message + Log a categorized message with optional timestamp Args: message: Message to log level: Log level (INFO, WARN, ERROR) category: Category for the message - force: If True, always log regardless of enabled state (for generic messages) + force: If True, always log regardless of enabled state (for critical messages) """ - # Always log forced messages (generic messages that were previously print statements) - # or log if debugging is enabled + # Always log forced messages or if debugging is enabled if force or self.enabled: # Get icon for category, fallback to general icon icon = self.CATEGORY_ICONS.get(category, self.CATEGORY_ICONS["general"]) @@ -88,13 +93,19 @@ class Debug: elif level == "ERROR": icon = self.CATEGORY_ICONS["error"] - # Build the log message - prefix = f"{icon}" + # Build the log message with optional timestamp + if self.show_timestamps and self.enabled: + timestamp = datetime.now().strftime("%H:%M:%S.%f")[:-3] + prefix = f"[{timestamp}] {icon}" + else: + prefix = f"{icon}" + if level != "INFO": prefix += f" [{level}]" print(f"{prefix} {message}") + @contextmanager def timer_context(self, namespace: str): """ @@ -112,6 +123,7 @@ class Debug: finally: self.timer_namespace = old_namespace + def start_timer(self, name: str, force: bool = False) -> None: """ Start a named timer @@ -139,6 +151,7 @@ class Debug: # Push to stack self.active_timer_stack.append(name) + def end_timer(self, name: str, message: Optional[str] = None, force: bool = False, show_breakdown: bool = False, custom_children: Optional[Dict[str, float]] = None) -> float: @@ -227,8 +240,9 @@ class Debug: return duration + def log_memory_state(self, label: str, show_diff: bool = True, show_tensors: bool = True, - detailed_tensors: bool = False) -> None: + detailed_tensors: bool = False, force: bool = False) -> None: """ Log current memory state with minimal overhead. @@ -237,36 +251,37 @@ class Debug: show_diff: Show change from last checkpoint show_tensors: Include tensor counts detailed_tensors: Show detailed tensor analysis (use sparingly) + force: If True, always log regardless of enabled state """ - if not self.enabled: + if not (self.enabled or force): return # Collect memory metrics efficiently memory_info = self._collect_memory_metrics() # Show category - self.log(f"{label}:", category="memory") + self.log(f"{label}:", category="memory", force=force) # Show VRAM if memory_info['summary_vram']: - self.log(f"{memory_info['summary_vram']}", category="memory") + self.log(f"{memory_info['summary_vram']}", category="memory", force=force) # Show RAM if memory_info['summary_ram']: - self.log(f"{memory_info['summary_ram']}", category="memory") + self.log(f"{memory_info['summary_ram']}", category="memory", force=force) # Show tensors if show_tensors: tensor_stats = self._collect_tensor_stats(detailed=detailed_tensors) - self.log(f"{tensor_stats['summary']}", category="memory") + self.log(f"{tensor_stats['summary']}", category="memory", force=force) # Show diff from last checkpoint if show_diff and self.memory_checkpoints: - self._log_memory_diff(memory_info) + self._log_memory_diff(current_metrics=memory_info, force=force) # Log detailed analysis if requested if detailed_tensors and tensor_stats.get('details'): - self._log_detailed_tensor_analysis(tensor_stats['details']) + self._log_detailed_tensor_analysis(details=tensor_stats['details'], force=force) # Store checkpoint with memory limit self._store_checkpoint(label, memory_info) @@ -274,6 +289,7 @@ class Debug: # Reset PyTorch's peak memory stats for next interval reset_vram_peak(debug=self) + def _collect_memory_metrics(self) -> Dict[str, Any]: """Collect current memory metrics efficiently.""" metrics = { @@ -332,6 +348,7 @@ class Debug: return metrics + def _collect_tensor_stats(self, detailed: bool = False) -> Dict[str, Any]: """Collect tensor statistics with minimal overhead.""" stats = { @@ -395,49 +412,51 @@ class Debug: return stats - def _log_detailed_tensor_analysis(self, details: Dict[str, Any]) -> None: + + def _log_detailed_tensor_analysis(self, details: Dict[str, Any], force: bool = False) -> None: """Log detailed tensor analysis when requested.""" # GPU tensors if details['gpu_tensors']: gpu_total_gb = sum(t['size_mb'] for t in details['gpu_tensors']) / 1024 - self.log(f" GPU tensors: {len(details['gpu_tensors'])} using {gpu_total_gb:.2f}GB", category="memory") + self.log(f" GPU tensors: {len(details['gpu_tensors'])} using {gpu_total_gb:.2f}GB", category="memory", force=force) # Show top 5 largest largest = sorted(details['gpu_tensors'], key=lambda x: x['size_mb'], reverse=True)[:5] for t in largest: - self.log(f" {t['shape']}: {t['size_mb']:.2f}MB, {t['dtype']}", category="memory") + self.log(f" {t['shape']}: {t['size_mb']:.2f}MB, {t['dtype']}", category="memory", force=force) # Large CPU tensors if details['large_cpu_tensors']: cpu_large_gb = sum(t['size_mb'] for t in details['large_cpu_tensors']) / 1024 - self.log(f" Large CPU tensors (>10MB):", category="memory") - self.log(f" {len(details['large_cpu_tensors'])} using {cpu_large_gb:.2f}GB", category="memory") + self.log(f" Large CPU tensors (>10MB):", category="memory", force=force) + self.log(f" {len(details['large_cpu_tensors'])} using {cpu_large_gb:.2f}GB", category="memory", force=force) # Show top 3 largest largest = sorted(details['large_cpu_tensors'], key=lambda x: x['size_mb'], reverse=True)[:3] for t in largest: - self.log(f" {t['shape']}: {t['size_mb']:.2f}MB, {t['dtype']}", category="memory") + self.log(f" {t['shape']}: {t['size_mb']:.2f}MB, {t['dtype']}", category="memory", force=force) # Common shape patterns if details['shape_patterns']: common_shapes = sorted(details['shape_patterns'].items(), key=lambda x: x[1], reverse=True)[:5] if len(common_shapes) > 0: - self.log(" Common tensor shapes:", category="memory") + self.log(" Common tensor shapes:", category="memory", force=force) for shape, count in common_shapes: if count > 1: - self.log(f" {shape}: {count} instances", category="memory") + self.log(f" {shape}: {count} instances", category="memory", force=force) # Module instances if details['module_types']: multi_instance = [(k, v) for k, v in details['module_types'].items() if v > 1] if multi_instance: - self.log(" Multiple module instances:", category="memory") + self.log(" Multiple module instances:", category="memory", force=force) for mtype, count in sorted(multi_instance, key=lambda x: x[1], reverse=True)[:5]: - self.log(f" {mtype}: {count} instances", category="memory") + self.log(f" {mtype}: {count} instances", category="memory", force=force) - def _log_memory_diff(self, current_metrics: Dict[str, Any]) -> None: + + def _log_memory_diff(self, current_metrics: Dict[str, Any], force: bool = False) -> None: """Log memory changes from last checkpoint.""" last = self.memory_checkpoints[-1] @@ -453,8 +472,9 @@ class Debug: diffs.append(f"RAM {sign}{ram_diff:.2f}GB") if diffs: - self.log(f" Memory changes: {', '.join(diffs)}", category="memory") + self.log(f" Memory changes: {', '.join(diffs)}", category="memory", force=force) + def _store_checkpoint(self, label: str, metrics: Dict[str, Any]) -> None: """Store checkpoint with memory limit to prevent leaks.""" checkpoint = { @@ -477,10 +497,19 @@ class Debug: self.memory_checkpoints = (self.memory_checkpoints[:mid] + self.memory_checkpoints[-mid:]) + def log_swap_time(self, component_id: Union[int, str], duration: float, - component_type: str = "block") -> None: - """Log swap timing information for BlockSwap operations""" - if self.enabled: + component_type: str = "block", force: bool = False) -> None: + """ + Log swap timing information for BlockSwap operations + + Args: + component_id: Identifier for the component being swapped + duration: Duration of the swap in seconds + component_type: Type of component ('block' or other) + force: If True, always log regardless of enabled state + """ + if self.enabled or force: # Store timing data self.swap_times.append({ 'component_id': component_id, @@ -494,18 +523,8 @@ class Debug: else: message = f"{component_type} {component_id} swap: {duration*1000:.2f}ms" - self.log(message, category="blockswap") + self.log(message, category="blockswap", force=force) - def clear_history(self) -> None: - """Clear all history tracking""" - self.timers.clear() - self.memory_checkpoints.clear() - self.swap_times.clear() - self.vram_history.clear() - self.timer_hierarchy.clear() - self.timer_durations.clear() - self.timer_messages.clear() - self.active_timer_stack.clear() def get_swap_summary(self) -> Dict[str, Any]: """Get summary of swap operations for analysis""" @@ -553,4 +572,16 @@ class Debug: summary['avg_vram_gb'] = sum(self.vram_history) / len(self.vram_history) summary['vram_variation_gb'] = max(self.vram_history) - min(self.vram_history) - return summary \ No newline at end of file + return summary + + + def clear_history(self) -> None: + """Clear all history tracking""" + self.timers.clear() + self.memory_checkpoints.clear() + self.swap_times.clear() + self.vram_history.clear() + self.timer_hierarchy.clear() + self.timer_durations.clear() + self.timer_messages.clear() + self.active_timer_stack.clear() \ No newline at end of file