perf: optimize model initialization with meta device for 90% speedup

- Use meta device initialization to avoid unnecessary memory allocation during model creation
  - Reduces DiT and VAE initialization time by ~90% when loading to CPU
  - Explicitly delete state dicts after loading to free memory immediately
- Refactor configure_model_inference() into focused helper functions for DRY code
- Improve debug logging with timestamps and cleaner section separators for readability
This commit is contained in:
Adrien Toupet
2025-08-26 13:23:30 -04:00
parent 1ee1ee0fb5
commit 89bde29cbe
4 changed files with 260 additions and 104 deletions
+4 -2
View File
@@ -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)
+173 -53
View File
@@ -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
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
+9 -6
View File
@@ -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),
+74 -43
View File
@@ -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
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()