From e65e7fa41826fe2ec32702f141a7328a8ec84a51 Mon Sep 17 00:00:00 2001 From: Adrien Toupet Date: Tue, 9 Dec 2025 12:29:15 -0500 Subject: [PATCH 1/8] Remove dead flash attention wrapper from FP8CompatibleDiT The wrapper methods (_apply_flash_attention_optimization and related) matched NaDiT attention modules by name but required qkv or q_proj+k_proj+v_proj attributes to optimize. NaDiT uses proj_qkv instead, so the optimization path was never taken - always falling back to original forward. FlashAttentionVarlen already handles flash_attn vs sdpa switching via its attention_mode attribute, making this wrapper redundant. Removes ~200 lines of dead code. --- src/optimization/compatibility.py | 209 ------------------------------ 1 file changed, 209 deletions(-) diff --git a/src/optimization/compatibility.py b/src/optimization/compatibility.py index d154dd0..c63b007 100644 --- a/src/optimization/compatibility.py +++ b/src/optimization/compatibility.py @@ -282,11 +282,6 @@ class FP8CompatibleDiT(torch.nn.Module): self.debug.start_timer("_stabilize_rope_computations") self._stabilize_rope_computations() self.debug.end_timer("_stabilize_rope_computations", "RoPE stabilization") - - # ๐Ÿš€ FLASH ATTENTION OPTIMIZATION (Phase 2) - self.debug.start_timer("_apply_flash_attention_optimization") - self._apply_flash_attention_optimization() - self.debug.end_timer("_apply_flash_attention_optimization", "Flash Attention application") def _detect_model_dtype(self) -> torch.dtype: """Detect main model dtype""" @@ -409,210 +404,6 @@ class FP8CompatibleDiT(torch.nn.Module): if rope_count > 0: self.debug.log(f"Stabilized {rope_count} RoPE modules", category="success") - - def _apply_flash_attention_optimization(self) -> None: - """๐Ÿš€ FLASH ATTENTION OPTIMIZATION - 30-50% speedup of attention layers""" - attention_layers_optimized = 0 - flash_attention_available = self._check_flash_attention_support() - - for name, module in self.dit_model.named_modules(): - # Identify all attention layers - if self._is_attention_layer(name, module): - # Apply optimization based on availability - if self._optimize_attention_layer(name, module, flash_attention_available): - attention_layers_optimized += 1 - - if not flash_attention_available: - self.debug.log("Flash Attention not available, using PyTorch SDPA as fallback", category="info", force=True) - - def _check_flash_attention_support(self) -> bool: - """Check if Flash Attention is available""" - # Check PyTorch SDPA (includes Flash Attention on H100/A100) - if hasattr(torch.nn.functional, 'scaled_dot_product_attention'): - return True - - # Check flash-attn package (uses module-level check from top of file) - return FLASH_ATTN_AVAILABLE - - def _is_attention_layer(self, name: str, module: torch.nn.Module) -> bool: - """Identify if a module is an attention layer""" - attention_keywords = [ - 'attention', 'attn', 'self_attn', 'cross_attn', 'mhattn', 'multihead', - 'transformer_block', 'dit_block' - ] - - # Check by name - if any(keyword in name.lower() for keyword in attention_keywords): - return True - - # Check by module type - module_type = type(module).__name__.lower() - if any(keyword in module_type for keyword in attention_keywords): - return True - - # Check by attributes (modules with q, k, v projections) - if hasattr(module, 'q_proj') or hasattr(module, 'qkv') or hasattr(module, 'to_q'): - return True - - return False - - def _optimize_attention_layer(self, name: str, module: torch.nn.Module, flash_attention_available: bool) -> bool: - """Optimize a specific attention layer""" - try: - # Save original forward method - if not hasattr(module, '_original_forward'): - module._original_forward = module.forward - - # Create new optimized forward method - if flash_attention_available: - optimized_forward = self._create_flash_attention_forward(module, name) - else: - optimized_forward = self._create_sdpa_forward(module, name) - - # Replace forward method - module.forward = optimized_forward - return True - - except Exception as e: - self.debug.log(f"Failed to optimize attention layer '{name}': {e}", level="WARNING", category="dit", force=True) - return False - - def _create_flash_attention_forward(self, module: torch.nn.Module, layer_name: str): - """Create optimized forward with Flash Attention""" - original_forward = module._original_forward - - def flash_attention_forward(*args, **kwargs): - try: - # Try to use Flash Attention via SDPA - return self._sdpa_attention_forward(original_forward, module, *args, **kwargs) - except Exception as e: - # Fallback to original implementation - self.debug.log(f"Flash Attention failed for {layer_name}, using original: {e}", level="WARNING", category="dit", force=True) - return original_forward(*args, **kwargs) - - return flash_attention_forward - - def _create_sdpa_forward(self, module: torch.nn.Module, layer_name: str): - """Create optimized forward with PyTorch SDPA""" - original_forward = module._original_forward - - def sdpa_forward(*args, **kwargs): - try: - return self._sdpa_attention_forward(original_forward, module, *args, **kwargs) - except Exception as e: - # Fallback to original implementation - return original_forward(*args, **kwargs) - - return sdpa_forward - - def _sdpa_attention_forward(self, original_forward, module: torch.nn.Module, *args, **kwargs): - """Optimized forward pass using SDPA (Scaled Dot Product Attention)""" - # Detect if we can intercept and optimize this layer - if len(args) >= 1 and isinstance(args[0], torch.Tensor): - input_tensor = args[0] - - # Check dimensions to ensure it's standard attention - if len(input_tensor.shape) >= 3: # [batch, seq_len, hidden_dim] or similar - try: - return self._optimized_attention_computation(module, input_tensor, *args[1:], **kwargs) - except: - pass - - # Fallback to original implementation - return original_forward(*args, **kwargs) - - def _optimized_attention_computation(self, module: torch.nn.Module, input_tensor: torch.Tensor, *args, **kwargs): - """Optimized attention computation with SDPA""" - # Try to detect standard attention format - batch_size, seq_len = input_tensor.shape[:2] - - # Check if module has standard Q, K, V projections - if hasattr(module, 'qkv') or (hasattr(module, 'q_proj') and hasattr(module, 'k_proj') and hasattr(module, 'v_proj')): - return self._compute_sdpa_attention(module, input_tensor, *args, **kwargs) - - # If no standard format detected, use original - return module._original_forward(input_tensor, *args, **kwargs) - - def _compute_sdpa_attention(self, module: torch.nn.Module, x: torch.Tensor, *args, **kwargs): - """Optimized SDPA computation for standard attention modules""" - try: - # Case 1: Module with combined QKV projection - if hasattr(module, 'qkv'): - qkv = module.qkv(x) - # Reshape to separate Q, K, V - batch_size, seq_len, _ = qkv.shape - qkv = qkv.reshape(batch_size, seq_len, 3, -1) - q, k, v = qkv.unbind(dim=2) - - # Case 2: Separate Q, K, V projections - elif hasattr(module, 'q_proj') and hasattr(module, 'k_proj') and hasattr(module, 'v_proj'): - q = module.q_proj(x) - k = module.k_proj(x) - v = module.v_proj(x) - else: - # Unsupported format, use original - return module._original_forward(x, *args, **kwargs) - - # Detect number of heads - head_dim = getattr(module, 'head_dim', None) - num_heads = getattr(module, 'num_heads', None) - - if head_dim is None or num_heads is None: - # Try to guess from dimensions - hidden_dim = q.shape[-1] - if hasattr(module, 'num_heads'): - num_heads = module.num_heads - head_dim = hidden_dim // num_heads - else: - # Reasonable defaults - head_dim = 64 - num_heads = hidden_dim // head_dim - - # Reshape for multi-head attention - batch_size, seq_len = q.shape[:2] - q = q.view(batch_size, seq_len, num_heads, head_dim).transpose(1, 2) - k = k.view(batch_size, seq_len, num_heads, head_dim).transpose(1, 2) - v = v.view(batch_size, seq_len, num_heads, head_dim).transpose(1, 2) - - if hasattr(torch.backends, 'mps') and torch.backends.mps.is_available(): - attn_output = torch.nn.functional.scaled_dot_product_attention( - q, k, v, - dropout_p=0.0, - is_causal=False - ) - else: - # Use optimized SDPA - PyTorch 2.3+ API with CUDNN support, fallback for older versions - if hasattr(torch.nn.attention, 'sdpa_kernel'): - ctx = torch.nn.attention.sdpa_kernel([ - torch.nn.attention.SDPBackend.FLASH_ATTENTION, - torch.nn.attention.SDPBackend.EFFICIENT_ATTENTION, - torch.nn.attention.SDPBackend.CUDNN_ATTENTION, - torch.nn.attention.SDPBackend.MATH]) - else: - ctx = torch.backends.cuda.sdp_kernel(enable_flash=True, enable_math=True, enable_mem_efficient=True) - - with ctx: - attn_output = torch.nn.functional.scaled_dot_product_attention( - q, k, v, - dropout_p=0.0, - is_causal=False - ) - - # Reshape back - attn_output = attn_output.transpose(1, 2).contiguous().view( - batch_size, seq_len, num_heads * head_dim - ) - - # Output projection if it exists - if hasattr(module, 'out_proj') or hasattr(module, 'o_proj'): - proj = getattr(module, 'out_proj', None) or getattr(module, 'o_proj', None) - attn_output = proj(attn_output) - - return attn_output - - except Exception as e: - # In case of error, use original implementation - return module._original_forward(x, *args, **kwargs) def forward(self, *args, **kwargs): """ From 30bc9240435d9838af8a351ce0a7e574c4970068 Mon Sep 17 00:00:00 2001 From: Adrien Toupet Date: Tue, 9 Dec 2025 14:07:46 -0500 Subject: [PATCH 2/8] Update header logo design (thanks @naxci1, closes #378) --- src/utils/debug.py | 42 ++++++++++++++++++++++++------------------ 1 file changed, 24 insertions(+), 18 deletions(-) diff --git a/src/utils/debug.py b/src/utils/debug.py index e819120..5247645 100644 --- a/src/utils/debug.py +++ b/src/utils/debug.py @@ -125,26 +125,32 @@ class Debug: def print_header(self, cli: bool = False) -> None: """Print the header with banner - always displayed""" - # Intro logo - self.log("", category="none", force=True) - self.log(" โ•”โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•—", category="none", force=True) - self.log(" โ•‘ โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ โ–ˆโ–ˆ โ–ˆโ–ˆ โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ โ•‘", category="none", force=True) - self.log(" โ•‘ โ–ˆโ–ˆ โ–ˆโ–ˆ โ–ˆโ–ˆ โ–ˆโ–ˆ โ–ˆโ–ˆ โ–ˆโ–ˆ โ–ˆโ–ˆ โ–ˆโ–ˆ โ–ˆโ–ˆ โ–ˆโ–ˆ โ•‘", category="none", force=True) - self.log(" โ•‘ โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ โ–ˆโ–ˆ โ–ˆโ–ˆ โ–ˆโ–ˆ โ–ˆโ–ˆ โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ โ•‘", category="none", force=True) - self.log(" โ•‘ โ–ˆโ–ˆ โ–ˆโ–ˆ โ–ˆโ–ˆ โ–ˆโ–ˆ โ–ˆโ–ˆ โ–ˆโ–ˆ โ–ˆโ–ˆ โ–ˆโ–ˆ โ–ˆโ–ˆ โ–ˆโ–ˆ โ•‘", category="none", force=True) - self.log(" โ•‘ โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ โ–ˆโ–ˆโ–ˆโ–ˆ โ–ˆโ–ˆ โ–ˆโ–ˆ โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ โ•‘", category="none", force=True) + # Temporarily disable timestamps for clean header display + original_timestamps = self.show_timestamps + self.show_timestamps = False - # Version number with dynamic padding to maintain visual alignment with any version length - version_text = f"v{__version__}" - prefix = " ๐Ÿ’ป CLI mode ยท " if cli else " " - suffix = "ยฉ ByteDance Seed ยท NumZ ยท AInVFX " - emoji_compensation = 1 if cli else 0 - padding_width = 59 - len(prefix) - len(version_text) - len(suffix) - 2 - emoji_compensation - padding = " " * max(1, padding_width) - self.log(f" โ•‘{prefix}{version_text}{padding} {suffix}โ•‘", category="none", force=True) - - self.log(" โ•šโ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•โ•", category="none", force=True) + # ASCII art logo self.log("", category="none", force=True) + self.log("โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•—โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•—โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•—โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•— โ–ˆโ–ˆโ•— โ–ˆโ–ˆโ•—โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•— โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•— โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•—", category="none", force=True, indent_level=1) + self.log("โ–ˆโ–ˆโ•”โ•โ•โ•โ•โ•โ–ˆโ–ˆโ•”โ•โ•โ•โ•โ•โ–ˆโ–ˆโ•”โ•โ•โ•โ•โ•โ–ˆโ–ˆโ•”โ•โ•โ–ˆโ–ˆโ•—โ–ˆโ–ˆโ•‘ โ–ˆโ–ˆโ•‘โ–ˆโ–ˆโ•”โ•โ•โ–ˆโ–ˆโ•— โ•šโ•โ•โ•โ•โ–ˆโ–ˆโ•— โ–ˆโ–ˆโ•”โ•โ•โ•โ•โ•", category="none", force=True, indent_level=1) + self.log("โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•—โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•— โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•— โ–ˆโ–ˆโ•‘ โ–ˆโ–ˆโ•‘โ–ˆโ–ˆโ•‘ โ–ˆโ–ˆโ•‘โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•”โ• โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•”โ• โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•—", category="none", force=True, indent_level=1) + self.log("โ•šโ•โ•โ•โ•โ–ˆโ–ˆโ•‘โ–ˆโ–ˆโ•”โ•โ•โ• โ–ˆโ–ˆโ•”โ•โ•โ• โ–ˆโ–ˆโ•‘ โ–ˆโ–ˆโ•‘โ•šโ–ˆโ–ˆโ•— โ–ˆโ–ˆโ•”โ•โ–ˆโ–ˆโ•”โ•โ•โ–ˆโ–ˆโ•— โ–ˆโ–ˆโ•”โ•โ•โ•โ• โ•šโ•โ•โ•โ•โ–ˆโ–ˆโ•‘", category="none", force=True, indent_level=1) + self.log("โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•‘โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•—โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•—โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•”โ• โ•šโ–ˆโ–ˆโ–ˆโ–ˆโ•”โ• โ–ˆโ–ˆโ•‘ โ–ˆโ–ˆโ•‘ โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•— โ–ˆโ–ˆโ•— โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•‘", category="none", force=True, indent_level=1) + self.log("โ•šโ•โ•โ•โ•โ•โ•โ•โ•šโ•โ•โ•โ•โ•โ•โ•โ•šโ•โ•โ•โ•โ•โ•โ•โ•šโ•โ•โ•โ•โ•โ• โ•šโ•โ•โ•โ• โ•šโ•โ• โ•šโ•โ• โ•šโ•โ•โ•โ•โ•โ•โ• โ•šโ•โ• โ•šโ•โ•โ•โ•โ•โ•โ•", category="none", force=True, indent_level=1) + # Version and credits - left/right aligned to logo width + version_text = f"v{__version__}" + cli_indicator = "๐Ÿ’ป CLI ยท " if cli else "" + left_part = f"{cli_indicator}{version_text}" + right_part = "ยฉ ByteDance Seed ยท NumZ ยท AInVFX" + logo_width = 75 + emoji_compensation = 1 if cli else 0 + padding = logo_width - len(left_part) - len(right_part) - emoji_compensation + self.log(f"{left_part}{' ' * max(1, padding)}{right_part}", category="none", force=True, indent_level=1) + self.log("โ”" * logo_width, category="none", force=True, indent_level=1) + self.log("", category="none", force=True) + + # Restore timestamps setting + self.show_timestamps = original_timestamps # Environment info - only in debug mode if self.enabled: From 77a00f651aa1affef954ddd78529f6900b27910f Mon Sep 17 00:00:00 2001 From: Adrien Toupet Date: Tue, 9 Dec 2025 17:12:10 -0500 Subject: [PATCH 3/8] Fix: OOM regression from 2.5.14 strict VRAM limit (#367) Add allow_vram_overflow option (default: False) to make strict VRAM limit configurable. The 2.5.14 change 'Enforce physical VRAM limit' prevented PyTorch from overflowing to system RAM, causing OOM on workflows that previously worked. - Add allow_vram_overflow parameter to DiT Model Loader node - Add --allow_vram_overflow CLI flag - Show warning when enabled, track mid-session changes - Suppress swap detection warning when user explicitly allows overflow Note: Enabling overflow is a last resort - performance degrades severely when physical VRAM is exceeded. Optimizing settings (BlockSwap, VAE tiling, batch size, resolution, model size...) is always recommended. --- README.md | 8 +++++ inference_cli.py | 11 ++++-- src/interfaces/dit_model_loader.py | 19 ++++++++++- src/optimization/memory_manager.py | 54 ++++++++++++++++++++++++++---- src/utils/debug.py | 49 ++++++++++++++++++++++++--- 5 files changed, 126 insertions(+), 15 deletions(-) diff --git a/README.md b/README.md index de87611..d177f37 100644 --- a/README.md +++ b/README.md @@ -418,6 +418,13 @@ Configure the DiT (Diffusion Transformer) model for video upscaling. - `sdpa`: PyTorch scaled_dot_product_attention (default, stable, always available) - `flash_attn`: Flash Attention 2 (faster on supported hardware, requires flash-attn package) +- **allow_vram_overflow**: Allow VRAM to overflow to system RAM + - `False` (default): Strict VRAM limit - prevents silent swap but OOMs if exceeded + - `True`: Allow overflow - prevents OOM but may cause severe slowdown when physical VRAM exceeded + - Last resort when other memory optimizations are insufficient + - Requires ComfyUI restart to change setting + - No effect on Apple Silicon (unified memory architecture) + - **torch_compile_args**: Connect to SeedVR2 Torch Compile Settings node for 20-40% speedup **BlockSwap Explained:** @@ -872,6 +879,7 @@ python inference_cli.py media_folder/ \ - `--tile_debug`: Visualize tiles: 'false' (default), 'encode', or 'decode' **Performance Optimization:** +- `--allow_vram_overflow`: Allow VRAM overflow to system RAM. Prevents OOM but may cause severe slowdown - `--attention_mode`: Attention backend: 'sdpa' (default, stable) or 'flash_attn' (faster, requires package) - `--compile_dit`: Enable torch.compile for DiT model (20-40% speedup, requires PyTorch 2.0+ and Triton) - `--compile_vae`: Enable torch.compile for VAE model (15-25% speedup, requires PyTorch 2.0+ and Triton) diff --git a/inference_cli.py b/inference_cli.py index 6b50aea..ab86fed 100644 --- a/inference_cli.py +++ b/inference_cli.py @@ -76,9 +76,10 @@ if platform.system() == "Darwin": else: os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "backend:cudaMallocAsync") - # Pre-parse CUDA device argument for validation and environment setup + # Pre-parse arguments that must be handled before torch import _pre_parser = argparse.ArgumentParser(add_help=False) _pre_parser.add_argument("--cuda_device", type=str, default=None) + _pre_parser.add_argument("--allow_vram_overflow", action="store_true") _pre_args, _ = _pre_parser.parse_known_args() if _pre_args.cuda_device is not None: @@ -127,9 +128,12 @@ from src.core.generation_phases import ( postprocess_all_batches ) from src.utils.debug import Debug -from src.optimization.memory_manager import clear_memory +from src.optimization.memory_manager import clear_memory, configure_vram_limit debug = Debug(enabled=False) # Will be enabled via --debug CLI flag +# Configure VRAM limit (must be before any CUDA allocations) +if platform.system() != "Darwin": + configure_vram_limit(allow_overflow=_pre_args.allow_vram_overflow) # ============================================================================= # Device Management Helpers @@ -1341,6 +1345,9 @@ Examples: "Requires --dit_offload_device. Default: 0 (disabled)") blockswap_group.add_argument("--swap_io_components", action="store_true", help="Offload DiT I/O layers for extra VRAM savings. Requires --dit_offload_device") + blockswap_group.add_argument("--allow_vram_overflow", action="store_true", + help="Allow VRAM overflow to system RAM. Prevents OOM but may cause severe slowdown. " + "Last resort when other memory optimizations are insufficient. No effect on Apple Silicon (unified memory).") # VAE Tiling vae_group = parser.add_argument_group('VAE tiling (for high resolution upscale)') diff --git a/src/interfaces/dit_model_loader.py b/src/interfaces/dit_model_loader.py index 1064571..76963d3 100644 --- a/src/interfaces/dit_model_loader.py +++ b/src/interfaces/dit_model_loader.py @@ -7,7 +7,7 @@ from comfy_api.latest import io from comfy_execution.utils import get_executing_context from typing import Dict, Any, Tuple from ..utils.model_registry import get_available_dit_models, DEFAULT_DIT -from ..optimization.memory_manager import get_device_list +from ..optimization.memory_manager import get_device_list, configure_vram_limit class SeedVR2LoadDiTModel(io.ComfyNode): @@ -112,6 +112,18 @@ class SeedVR2LoadDiTModel(io.ComfyNode): "Flash Attention provides speedup through optimized CUDA kernels on compatible GPUs." ) ), + io.Boolean.Input("allow_vram_overflow", + default=False, + optional=True, + tooltip=( + "Allow VRAM to overflow to system RAM when physical VRAM is exceeded.\n" + "โ€ข False (default): Strict VRAM limit - OOM if exceeded (faster when within limits)\n" + "โ€ข True: Allow overflow to RAM - prevents OOM but may cause severe slowdown\n" + "\n" + "Last resort when other memory optimizations are insufficient.\n" + "Requires ComfyUI restart to change. No effect on Apple Silicon (unified memory)." + ) + ), io.Custom("TORCH_COMPILE_ARGS").Input("torch_compile_args", optional=True, tooltip=( @@ -131,6 +143,7 @@ class SeedVR2LoadDiTModel(io.ComfyNode): def execute(cls, model: str, device: str, offload_device: str = "none", cache_model: bool = False, blocks_to_swap: int = 0, swap_io_components: bool = False, attention_mode: str = "sdpa", + allow_vram_overflow: bool = False, torch_compile_args: Dict[str, Any] = None) -> io.NodeOutput: """ Create DiT model configuration for SeedVR2 main node @@ -143,6 +156,7 @@ class SeedVR2LoadDiTModel(io.ComfyNode): blocks_to_swap: Number of transformer blocks to swap (requires offload_device != device) swap_io_components: Whether to offload I/O components (requires offload_device != device) attention_mode: Attention computation backend ('sdpa' or 'flash_attn') + allow_vram_overflow: Allow VRAM overflow to system RAM (prevents OOM but slower) torch_compile_args: Optional torch.compile configuration from settings node Returns: @@ -168,6 +182,9 @@ class SeedVR2LoadDiTModel(io.ComfyNode): "(e.g., 'cpu' or another device). Set cache_model=False if you don't want to cache the model." ) + # Configure VRAM limit enforcement (once per session, first call wins) + configure_vram_limit(allow_overflow=allow_vram_overflow) + config = { "model": model, "device": device, diff --git a/src/optimization/memory_manager.py b/src/optimization/memory_manager.py index 592e37e..229f5bb 100644 --- a/src/optimization/memory_manager.py +++ b/src/optimization/memory_manager.py @@ -112,20 +112,60 @@ else: print(f"โš ๏ธ Memory check failed: {vram_info['error']} - No available backend!") -def _enforce_vram_limit() -> None: +# VRAM overflow configuration state +_vram_overflow_allowed: bool = True +_vram_limit_configured: bool = False +_vram_limit_change_attempted: bool = False + + +def configure_vram_limit(allow_overflow: bool = False) -> bool: """ - Enforce VRAM limit to physical capacity to prevent silent swap to system RAM. - Called once at module load. No-op on MPS or unsupported platforms. + Configure VRAM limit enforcement. Call early before heavy CUDA usage. + + Args: + allow_overflow: If True, allow VRAM overflow to system RAM (prevents OOM but may be slow). + If False (default), enforce strict physical VRAM limit. + + Returns: + True if configuration applied successfully, False otherwise + + Note: + Can only be configured once per session. Restart required to change. """ + global _vram_overflow_allowed, _vram_limit_configured, _vram_limit_change_attempted + + # Already configured this session - track if user tried to change + if _vram_limit_configured: + if _vram_overflow_allowed != allow_overflow: + _vram_limit_change_attempted = True + return _vram_overflow_allowed == allow_overflow + + _vram_limit_configured = True + _vram_overflow_allowed = allow_overflow + + if allow_overflow: + return True + if not torch.cuda.is_available(): - return + return True + try: for i in range(torch.cuda.device_count()): torch.cuda.set_per_process_memory_fraction(1.0, i) - except Exception: - pass + return True + except RuntimeError: + _vram_overflow_allowed = True + return False -_enforce_vram_limit() + +def is_vram_overflow_allowed() -> bool: + """Check if VRAM overflow to system RAM is allowed.""" + return _vram_overflow_allowed + + +def was_vram_limit_change_attempted() -> bool: + """Check if user tried to change VRAM limit setting after initial configuration.""" + return _vram_limit_change_attempted def get_vram_usage(device: Optional[torch.device] = None, debug: Optional['Debug'] = None) -> Tuple[float, float, float]: diff --git a/src/utils/debug.py b/src/utils/debug.py index 5247645..bf6acd0 100644 --- a/src/utils/debug.py +++ b/src/utils/debug.py @@ -10,7 +10,14 @@ import torch import gc from typing import Optional, List, Dict, Any, Union from datetime import datetime -from ..optimization.memory_manager import get_vram_usage, get_basic_vram_info, get_ram_usage, reset_vram_peak +from ..optimization.memory_manager import ( + get_vram_usage, + get_basic_vram_info, + get_ram_usage, + reset_vram_peak, + is_vram_overflow_allowed, + was_vram_limit_change_attempted +) from ..utils.constants import __version__ @@ -131,6 +138,7 @@ class Debug: # ASCII art logo self.log("", category="none", force=True) + self.log("", category="none", force=True) self.log("โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•—โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•—โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•—โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•— โ–ˆโ–ˆโ•— โ–ˆโ–ˆโ•—โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•— โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•— โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•—", category="none", force=True, indent_level=1) self.log("โ–ˆโ–ˆโ•”โ•โ•โ•โ•โ•โ–ˆโ–ˆโ•”โ•โ•โ•โ•โ•โ–ˆโ–ˆโ•”โ•โ•โ•โ•โ•โ–ˆโ–ˆโ•”โ•โ•โ–ˆโ–ˆโ•—โ–ˆโ–ˆโ•‘ โ–ˆโ–ˆโ•‘โ–ˆโ–ˆโ•”โ•โ•โ–ˆโ–ˆโ•— โ•šโ•โ•โ•โ•โ–ˆโ–ˆโ•— โ–ˆโ–ˆโ•”โ•โ•โ•โ•โ•", category="none", force=True, indent_level=1) self.log("โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•—โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•— โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•— โ–ˆโ–ˆโ•‘ โ–ˆโ–ˆโ•‘โ–ˆโ–ˆโ•‘ โ–ˆโ–ˆโ•‘โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•”โ• โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•”โ• โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ•—", category="none", force=True, indent_level=1) @@ -155,6 +163,11 @@ class Debug: # Environment info - only in debug mode if self.enabled: self._print_environment_info(cli) + + # VRAM overflow status - warnings always shown + vram_warning_shown = self._print_vram_overflow_status() + + self.log("", category="none", force=vram_warning_shown) def _print_environment_info(self, cli: bool = False) -> None: """Print concise environment info for bug reports - zero cost when debug disabled""" @@ -216,7 +229,32 @@ class Debug: self.log(f"Python: {py_ver} | PyTorch: {torch_ver} | Flash Attn: {flash_str} | Triton: {triton_str}", category="info") cuda_line = f"CUDA: {cuda_ver} | cuDNN: {cudnn_ver}" self.log(f"{cuda_line} | ComfyUI: {comfy_str}" if comfy_str else cuda_line, category="info") - self.log("", category="none") + + def _print_vram_overflow_status(self) -> bool: + """Print VRAM overflow status - warnings always shown, info only in debug mode. + + Returns: + True if a forced warning was printed, False otherwise. + """ + is_mps = hasattr(torch.backends, 'mps') and torch.backends.mps.is_available() + force = False + + if was_vram_limit_change_attempted(): + self.log("allow_vram_overflow setting changed - restart ComfyUI to apply", level="WARNING", category="memory", force=True) + force = True + elif is_vram_overflow_allowed(): + if is_mps: + self.log("allow_vram_overflow: enabled (no effect on Apple Silicon unified memory)", category="info") + else: + self.log("allow_vram_overflow: enabled - may cause severe slowdown if physical VRAM exceeded", level="WARNING", category="memory", force=True) + force = True + else: + if is_mps: + self.log("allow_vram_overflow: disabled (no effect on Apple Silicon unified memory)", category="info") + else: + self.log("allow_vram_overflow: disabled (recommended for best performance)", category="success") + + return force def print_footer(self) -> None: """Print the footer with links - always displayed""" @@ -387,10 +425,11 @@ class Debug: if show_diff and self.memory_checkpoints: self._log_memory_diff(current_metrics=memory_info, force=force) - # Warn if swap detected (peak > physical VRAM) + # Warn if swap detected (peak > physical VRAM), unless user explicitly allowed overflow if memory_info['vram_total'] > 0 and memory_info['vram_peak_since_last'] > memory_info['vram_total']: - self.log("VRAM swap detected - severe slowdown expected. Consider optimizing (e.g., reduce resolution, batch_size, enable BlockSwap, VAE tiling...).", - level="WARNING", category="memory", force=True) + if not is_vram_overflow_allowed(): + self.log("VRAM swap detected - severe slowdown expected. Consider optimizing (e.g., reduce resolution, batch_size, enable BlockSwap, VAE tiling...).", + level="WARNING", category="memory", force=True) # Log detailed analysis if requested if detailed_tensors and tensor_stats.get('details'): From 5c60716c4797196667c6ff81433ded65965be1e7 Mon Sep 17 00:00:00 2001 From: Adrien Toupet Date: Tue, 9 Dec 2025 21:06:12 -0500 Subject: [PATCH 4/8] Refactor: centralize backend detection, fix architecture-aware VRAM overflow reporting --- inference_cli.py | 22 +--- src/common/distributed/basic.py | 3 +- src/data/image/transforms/area_resize.py | 3 +- src/data/image/transforms/na_resize.py | 3 +- src/data/image/transforms/side_resize.py | 3 +- src/optimization/memory_manager.py | 79 ++++++++++--- src/utils/debug.py | 142 ++++++++++++++--------- 7 files changed, 163 insertions(+), 92 deletions(-) diff --git a/inference_cli.py b/inference_cli.py index ab86fed..342231e 100644 --- a/inference_cli.py +++ b/inference_cli.py @@ -128,7 +128,7 @@ from src.core.generation_phases import ( postprocess_all_batches ) from src.utils.debug import Debug -from src.optimization.memory_manager import clear_memory, configure_vram_limit +from src.optimization.memory_manager import clear_memory, configure_vram_limit, get_gpu_backend, is_cuda_available debug = Debug(enabled=False) # Will be enabled via --debug CLI flag # Configure VRAM limit (must be before any CUDA allocations) @@ -139,16 +139,6 @@ if platform.system() != "Darwin": # Device Management Helpers # ============================================================================= -def _get_platform_type() -> str: - """Determine the platform device type (cuda/mps/cpu).""" - if platform.system() == "Darwin": - return "mps" - elif torch.cuda.is_available(): - return "cuda" - else: - return "cpu" - - def _device_id_to_name(device_id: str, platform_type: str = None) -> str: """ Convert device ID to full device name. @@ -164,7 +154,7 @@ def _device_id_to_name(device_id: str, platform_type: str = None) -> str: return device_id if platform_type is None: - platform_type = _get_platform_type() + platform_type = get_gpu_backend() # MPS typically doesn't use indices if platform_type == "mps": @@ -781,7 +771,7 @@ def _process_frames_core( Upscaled frames tensor [T', H', W', C], Float32, range [0,1] """ # Determine platform and convert device IDs to full names - platform_type = _get_platform_type() + platform_type = get_gpu_backend() inference_device = _device_id_to_name(device_id, platform_type) # Parse offload devices (with caching defaults) @@ -1473,7 +1463,7 @@ def main() -> None: # Inform about caching defaults if args.cache_dit and args.dit_offload_device == "none": - offload_target = "system memory (CPU)" if _get_platform_type() != "mps" else "unified memory" + offload_target = "system memory (CPU)" if get_gpu_backend() != "mps" else "unified memory" debug.log( f"DiT caching enabled: Using default {offload_target} for offload. " "Set --dit_offload_device explicitly to use a different device.", @@ -1481,7 +1471,7 @@ def main() -> None: ) if args.cache_vae and args.vae_offload_device == "none": - offload_target = "system memory (CPU)" if _get_platform_type() != "mps" else "unified memory" + offload_target = "system memory (CPU)" if get_gpu_backend() != "mps" else "unified memory" debug.log( f"VAE caching enabled: Using default {offload_target} for offload. " "Set --vae_offload_device explicitly to use a different device.", @@ -1494,7 +1484,7 @@ def main() -> None: else: # Show actual CUDA device visibility debug.log(f"CUDA_VISIBLE_DEVICES: {os.environ.get('CUDA_VISIBLE_DEVICES', 'Not set (all)')}", category="device") - if torch.cuda.is_available(): + if is_cuda_available(): debug.log(f"torch.cuda.device_count(): {torch.cuda.device_count()}", category="device") debug.log(f"Using device index 0 inside script (mapped to selected GPU)", category="device") diff --git a/src/common/distributed/basic.py b/src/common/distributed/basic.py index 4615967..d92610e 100644 --- a/src/common/distributed/basic.py +++ b/src/common/distributed/basic.py @@ -21,6 +21,7 @@ from datetime import timedelta import torch import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel +from ...optimization.memory_manager import is_mps_available def get_global_rank() -> int: """ @@ -47,7 +48,7 @@ def get_device() -> torch.device: """ Get current rank device. """ - if hasattr(torch.backends, 'mps') and torch.backends.mps.is_available(): + if is_mps_available(): return torch.device("mps") return torch.device("cuda", get_local_rank()) diff --git a/src/data/image/transforms/area_resize.py b/src/data/image/transforms/area_resize.py index 5873b85..fc025da 100644 --- a/src/data/image/transforms/area_resize.py +++ b/src/data/image/transforms/area_resize.py @@ -19,6 +19,7 @@ import torch from PIL import Image from torchvision.transforms import functional as TVF from torchvision.transforms.functional import InterpolationMode +from ....optimization.memory_manager import is_mps_available class AreaResize: @@ -31,7 +32,7 @@ class AreaResize: self.max_area = max_area self.downsample_only = downsample_only self.interpolation = interpolation - if hasattr(torch.backends, 'mps') and torch.backends.mps.is_available(): + if is_mps_available(): self.interpolation = InterpolationMode.BILINEAR def __call__(self, image: Union[torch.Tensor, Image.Image]): diff --git a/src/data/image/transforms/na_resize.py b/src/data/image/transforms/na_resize.py index 61a186a..e1111c7 100644 --- a/src/data/image/transforms/na_resize.py +++ b/src/data/image/transforms/na_resize.py @@ -18,6 +18,7 @@ from torchvision.transforms import CenterCrop, Compose, InterpolationMode, Resiz from .area_resize import AreaResize from .side_resize import SideResize +from ....optimization.memory_manager import is_mps_available def NaResize( resolution: int, @@ -26,7 +27,7 @@ def NaResize( max_resolution: int = 0, interpolation: InterpolationMode = InterpolationMode.BICUBIC, ): - Interpolation = InterpolationMode.BILINEAR if (hasattr(torch.backends, 'mps') and torch.backends.mps.is_available()) else interpolation + Interpolation = InterpolationMode.BILINEAR if is_mps_available() else interpolation if mode == "area": return AreaResize( max_area=resolution**2, diff --git a/src/data/image/transforms/side_resize.py b/src/data/image/transforms/side_resize.py index 6d5273f..01362ae 100644 --- a/src/data/image/transforms/side_resize.py +++ b/src/data/image/transforms/side_resize.py @@ -17,6 +17,7 @@ import torch from PIL import Image from torchvision.transforms import InterpolationMode from torchvision.transforms import functional as TVF +from ....optimization.memory_manager import is_mps_available class SideResize: def __init__( @@ -30,7 +31,7 @@ class SideResize: self.max_size = max_size self.downsample_only = downsample_only self.interpolation = interpolation - if hasattr(torch.backends, 'mps') and torch.backends.mps.is_available(): + if is_mps_available(): self.interpolation = InterpolationMode.BILINEAR def __call__(self, image: Union[torch.Tensor, Image.Image]): diff --git a/src/optimization/memory_manager.py b/src/optimization/memory_manager.py index 229f5bb..b8150a6 100644 --- a/src/optimization/memory_manager.py +++ b/src/optimization/memory_manager.py @@ -10,6 +10,7 @@ import gc import sys import time import psutil +import platform from typing import Tuple, Dict, Any, Optional, List, Union @@ -19,6 +20,52 @@ def _device_str(device: Union[torch.device, str]) -> str: return 'MPS' if s.startswith('MPS') else s +def is_mps_available() -> bool: + """Check if MPS (Apple Metal) backend is available.""" + return hasattr(torch.backends, 'mps') and torch.backends.mps.is_available() + + +def is_cuda_available() -> bool: + """Check if CUDA backend is available.""" + return torch.cuda.is_available() + + +def get_gpu_backend() -> str: + """Get the active GPU backend type. + + Returns: + 'cuda': NVIDIA CUDA + 'mps': Apple Metal Performance Shaders + 'cpu': No GPU backend available + """ + if is_cuda_available(): + return 'cuda' + if is_mps_available(): + return 'mps' + return 'cpu' + + +def get_memory_architecture() -> str: + """Get memory architecture type for swap/overflow detection. + + This combines GPU backend with OS platform to determine how + GPU memory overflow is handled: + + Returns: + 'unified': macOS unified memory (MPS) - GPU/CPU share memory pool + 'discrete_paged': Windows WDDM - GPU memory can page to system RAM + 'discrete_strict': Linux - No automatic GPU paging, OOM on overflow + 'cpu_only': No GPU backend available + """ + if is_mps_available(): + return 'unified' + if is_cuda_available(): + if platform.system() == 'Windows': + return 'discrete_paged' + return 'discrete_strict' + return 'cpu_only' + + def get_device_list(include_none: bool = False, include_cpu: bool = False) -> List[str]: """ Get list of available compute devices for SeedVR2 @@ -37,14 +84,14 @@ def get_device_list(include_none: bool = False, include_cpu: bool = False) -> Li has_mps = False try: - if hasattr(torch, "cuda") and hasattr(torch.cuda, "is_available") and torch.cuda.is_available(): + if is_cuda_available(): devs += [f"cuda:{i}" for i in range(torch.cuda.device_count())] has_cuda = True except Exception: pass try: - if hasattr(torch.backends, "mps") and torch.backends.mps.is_available(): + if is_mps_available(): devs.append("mps") # MPS doesn't use device indices has_mps = True except Exception: @@ -66,7 +113,7 @@ def get_device_list(include_none: bool = False, include_cpu: bool = False) -> Li result.extend(devs) return result if result else [] - + def get_basic_vram_info(device: Optional[torch.device] = None) -> Dict[str, Any]: """ @@ -80,13 +127,13 @@ def get_basic_vram_info(device: Optional[torch.device] = None) -> Dict[str, Any] dict: {"free_gb": float, "total_gb": float} or {"error": str} """ try: - if torch.cuda.is_available(): + if is_cuda_available(): if device is None: device = torch.device("cuda:0") elif not isinstance(device, torch.device): device = torch.device(device) free_memory, total_memory = torch.cuda.mem_get_info(device) - elif hasattr(torch.backends, 'mps') and torch.backends.mps.is_available(): + elif is_mps_available(): # MPS doesn't support per-device queries or mem_get_info # Use system memory as proxy mem = psutil.virtual_memory() @@ -106,7 +153,7 @@ def get_basic_vram_info(device: Optional[torch.device] = None) -> Dict[str, Any] # Initial VRAM check at module load vram_info = get_basic_vram_info(device=None) if "error" not in vram_info: - backend = "MPS" if (hasattr(torch.backends, 'mps') and torch.backends.mps.is_available()) else "CUDA" + backend = "MPS" if is_mps_available() else "CUDA" print(f"๐Ÿ“Š Initial {backend} memory: {vram_info['free_gb']:.2f}GB free / {vram_info['total_gb']:.2f}GB total") else: print(f"โš ๏ธ Memory check failed: {vram_info['error']} - No available backend!") @@ -146,7 +193,7 @@ def configure_vram_limit(allow_overflow: bool = False) -> bool: if allow_overflow: return True - if not torch.cuda.is_available(): + if not is_cuda_available(): return True try: @@ -182,7 +229,7 @@ def get_vram_usage(device: Optional[torch.device] = None, debug: Optional['Debug Returns (0, 0, 0) if no GPU available """ try: - if torch.cuda.is_available(): + if is_cuda_available(): if device is None: device = torch.device("cuda:0") elif not isinstance(device, torch.device): @@ -191,7 +238,7 @@ def get_vram_usage(device: Optional[torch.device] = None, debug: Optional['Debug reserved = torch.cuda.memory_reserved(device) / (1024**3) max_reserved = torch.cuda.max_memory_reserved(device) / (1024**3) return allocated, reserved, max_reserved - elif hasattr(torch.backends, 'mps') and torch.backends.mps.is_available(): + elif is_mps_available(): # MPS doesn't support per-device queries - uses global memory tracking allocated = torch.mps.current_allocated_memory() / (1024**3) reserved = torch.mps.driver_allocated_memory() / (1024**3) @@ -291,17 +338,17 @@ def clear_memory(debug: Optional['Debug'] = None, deep: bool = False, force: boo # Use existing function for memory info mem_info = get_basic_vram_info(device=None) - if "error" not in mem_info: + if "error" not in mem_info and mem_info["total_gb"] > 0: # Check VRAM/MPS memory pressure (5% free threshold) free_ratio = mem_info["free_gb"] / mem_info["total_gb"] if free_ratio < 0.05: should_clear = True if debug: - backend = "MPS" if (hasattr(torch.backends, 'mps') and torch.backends.mps.is_available()) else "VRAM" + backend = "Unified Memory" if is_mps_available() else "VRAM" debug.log(f"{backend} pressure: {mem_info['free_gb']:.2f}GB free of {mem_info['total_gb']:.2f}GB", category="memory") # For non-MPS systems, also check system RAM separately - if not should_clear and not (hasattr(torch.backends, 'mps') and torch.backends.mps.is_available()): + if not should_clear and not is_mps_available(): mem = psutil.virtual_memory() if mem.available < mem.total * 0.05: should_clear = True @@ -324,10 +371,10 @@ def clear_memory(debug: Optional['Debug'] = None, deep: bool = False, force: boo if debug: debug.start_timer(gpu_timer) - if torch.cuda.is_available(): + if is_cuda_available(): torch.cuda.empty_cache() torch.cuda.ipc_collect() - elif hasattr(torch.backends, 'mps') and torch.backends.mps.is_available(): + elif is_mps_available(): torch.mps.empty_cache() if debug: @@ -364,7 +411,7 @@ def clear_memory(debug: Optional['Debug'] = None, deep: bool = False, force: boo handle = _os_memory_lib.GetCurrentProcess() _os_memory_lib.SetProcessWorkingSetSize(handle, -1, -1) - elif hasattr(torch.backends, 'mps') and torch.backends.mps.is_available(): + elif is_mps_available(): # macOS with MPS import ctypes # Import only when needed import ctypes.util @@ -441,7 +488,7 @@ def reset_vram_peak(device: Optional[torch.device] = None, debug: Optional['Debu if debug and debug.enabled: debug.log("Resetting VRAM peak memory statistics", category="memory") try: - if torch.cuda.is_available(): + if is_cuda_available(): if device is None: device = torch.device("cuda:0") elif not isinstance(device, torch.device): diff --git a/src/utils/debug.py b/src/utils/debug.py index bf6acd0..10f3c2c 100644 --- a/src/utils/debug.py +++ b/src/utils/debug.py @@ -16,16 +16,37 @@ from ..optimization.memory_manager import ( get_ram_usage, reset_vram_peak, is_vram_overflow_allowed, - was_vram_limit_change_attempted + was_vram_limit_change_attempted, + is_mps_available, + is_cuda_available, + get_memory_architecture ) from ..utils.constants import __version__ -def _format_peak_with_swap(peak_gb: float, total_vram_gb: float) -> str: - """Format peak memory, showing swap breakdown if overflow occurred.""" - if total_vram_gb > 0 and peak_gb > total_vram_gb: - swap_gb = peak_gb - total_vram_gb - return f"{peak_gb:.2f}GB ({total_vram_gb:.0f}GB GPU + {swap_gb:.2f}GB swap)" +def _format_peak_with_swap(peak_gb: float, total_vram_gb: float, arch: str = None) -> str: + """Format peak memory with architecture-aware overflow reporting. + + Args: + peak_gb: Peak reserved memory from PyTorch + total_vram_gb: Physical GPU VRAM capacity + arch: Memory architecture from get_memory_architecture(), or None to auto-detect + """ + if total_vram_gb <= 0: + return f"{peak_gb:.2f}GB" + + overflow_gb = peak_gb - total_vram_gb + if overflow_gb <= 0: + return f"{peak_gb:.2f}GB" + + if arch is None: + arch = get_memory_architecture() + + if arch == 'discrete_paged': + return f"{peak_gb:.2f}GB ({total_vram_gb:.0f}GB GPU + {overflow_gb:.2f}GB system RAM)" + elif arch == 'discrete_strict': + return f"{peak_gb:.2f}GB (exceeded {total_vram_gb:.0f}GB by {overflow_gb:.2f}GB)" + # unified or cpu_only - no swap concept return f"{peak_gb:.2f}GB" @@ -193,7 +214,7 @@ class Debug: cuda_ver = getattr(torch.version, 'cuda', None) or "N/A" # GPU - if torch.cuda.is_available(): + if is_cuda_available(): try: props = torch.cuda.get_device_properties(0) gpu_str = f"{props.name} ({round(props.total_memory / (1024**3))}GB)" @@ -201,7 +222,7 @@ class Debug: except Exception: gpu_str = "CUDA" cudnn_ver = "N/A" - elif getattr(getattr(torch, 'mps', None), 'is_available', lambda: False)(): + elif is_mps_available(): gpu_str = "Apple Silicon (MPS)" cudnn_ver = "N/A" else: @@ -236,7 +257,7 @@ class Debug: Returns: True if a forced warning was printed, False otherwise. """ - is_mps = hasattr(torch.backends, 'mps') and torch.backends.mps.is_available() + is_mps = is_mps_available() force = False if was_vram_limit_change_attempted(): @@ -425,10 +446,18 @@ class Debug: if show_diff and self.memory_checkpoints: self._log_memory_diff(current_metrics=memory_info, force=force) - # Warn if swap detected (peak > physical VRAM), unless user explicitly allowed overflow - if memory_info['vram_total'] > 0 and memory_info['vram_peak_since_last'] > memory_info['vram_total']: - if not is_vram_overflow_allowed(): - self.log("VRAM swap detected - severe slowdown expected. Consider optimizing (e.g., reduce resolution, batch_size, enable BlockSwap, VAE tiling...).", + # Architecture-aware overflow warnings + arch = memory_info.get('arch', 'cpu_only') + overflow = memory_info.get('vram_overflow', 0.0) + + if overflow > 0 and not is_vram_overflow_allowed(): + if arch == 'discrete_paged': + self.log(f"VRAM overflow: {overflow:.2f}GB paged to system RAM - severe slowdown expected. " + "Consider optimizing (e.g., reduce resolution, batch size, enable BlockSwap, VAE tiling...).", + level="WARNING", category="memory", force=True) + elif arch == 'discrete_strict': + self.log(f"VRAM exceeded physical limit by {overflow:.2f}GB - OOM risk. " + "Consider optimizing (e.g., reduce resolution, batch size, enable BlockSwap, VAE tiling...).", level="WARNING", category="memory", force=True) # Log detailed analysis if requested @@ -455,13 +484,17 @@ class Debug: reset_vram_peak(device=None, debug=self) def _collect_memory_metrics(self) -> Dict[str, Any]: - """Collect current memory metrics efficiently.""" + """Collect current memory metrics with architecture-aware reporting.""" + arch = get_memory_architecture() + metrics = { 'vram_allocated': 0.0, 'vram_reserved': 0.0, 'vram_free': 0.0, 'vram_total': 0.0, 'vram_peak_since_last': 0.0, + 'vram_overflow': 0.0, + 'arch': arch, 'ram_process': 0.0, 'ram_available': 0.0, 'ram_total': 0.0, @@ -470,46 +503,43 @@ class Debug: 'summary_ram': "" } - # VRAM metrics - if torch.cuda.is_available() or (hasattr(torch.backends, 'mps') and torch.backends.mps.is_available()): - metrics['vram_allocated'], metrics['vram_reserved'], current_global_peak = get_vram_usage(device=None, debug=self) - - # Calculate peak since last log_memory_state - # This captures the actual peak that occurred between calls - metrics['vram_peak_since_last'] = current_global_peak - + if arch == 'cpu_only': + pass # No GPU metrics + else: + metrics['vram_allocated'], metrics['vram_reserved'], metrics['vram_peak_since_last'] = get_vram_usage(device=None, debug=self) vram_info = get_basic_vram_info(device=None) - if "error" not in vram_info: + if "error" not in vram_info and vram_info["total_gb"] > 0: metrics['vram_free'] = vram_info["free_gb"] metrics['vram_total'] = vram_info["total_gb"] - backend = "MPS" if (hasattr(torch.backends, 'mps') and torch.backends.mps.is_available()) else "VRAM" - peak_str = _format_peak_with_swap(metrics['vram_peak_since_last'], metrics['vram_total']) - metrics['summary_vram'] = (f" [{backend}] {metrics['vram_allocated']:.2f}GB allocated / " - f"{metrics['vram_reserved']:.2f}GB reserved / " - f"Peak: {peak_str} / " - f"{metrics['vram_free']:.2f}GB free / " - f"{metrics['vram_total']:.2f}GB total") - else: - metrics['summary_vram'] = "" - else: - metrics['summary_vram'] = "" + # Calculate overflow: reserved beyond physical VRAM + metrics['vram_overflow'] = max(0.0, metrics['vram_peak_since_last'] - metrics['vram_total']) + + backend = "Unified Memory" if arch == 'unified' else "VRAM" + peak_str = _format_peak_with_swap(metrics['vram_peak_since_last'], metrics['vram_total'], arch) + metrics['summary_vram'] = ( + f" [{backend}] {metrics['vram_allocated']:.2f}GB allocated / " + f"{metrics['vram_reserved']:.2f}GB reserved / " + f"Peak: {peak_str} / " + f"{metrics['vram_free']:.2f}GB free / " + f"{metrics['vram_total']:.2f}GB total" + ) - # RAM metrics using new function + # RAM metrics metrics['ram_process'], metrics['ram_available'], metrics['ram_total'], metrics['ram_others'] = get_ram_usage(debug=self) if metrics['ram_total'] > 0: - metrics['summary_ram'] = (f" [RAM] {metrics['ram_process']:.2f}GB process / " - f"{metrics['ram_others']:.2f}GB others / " - f"{metrics['ram_available']:.2f}GB free / " - f"{metrics['ram_total']:.2f}GB total") - else: - metrics['summary_ram'] = "" + metrics['summary_ram'] = ( + f" [RAM] {metrics['ram_process']:.2f}GB process / " + f"{metrics['ram_others']:.2f}GB others / " + f"{metrics['ram_available']:.2f}GB free / " + f"{metrics['ram_total']:.2f}GB total" + ) - # Update VRAM history for tracking - if torch.cuda.is_available() or (hasattr(torch.backends, 'mps') and torch.backends.mps.is_available()): - self.vram_history.append(metrics['vram_allocated']) + # Track reserved (matches nvidia-smi) for pressure history + if arch != 'cpu_only': + self.vram_history.append(metrics['vram_reserved']) return metrics @@ -648,11 +678,11 @@ class Debug: 'phase4': 'Post-processing' } - is_mps = hasattr(torch.backends, 'mps') and torch.backends.mps.is_available() and not torch.cuda.is_available() + arch = get_memory_architecture() - # Get total VRAM for swap detection (reuse existing function) + # Get total VRAM for overflow detection total_vram_gb = 0.0 - if not is_mps: + if arch not in ('unified', 'cpu_only'): vram_info = get_basic_vram_info(device=None) if "error" not in vram_info: total_vram_gb = vram_info["total_gb"] @@ -668,18 +698,18 @@ class Debug: vram = self.phase_vram_peaks.get(phase_key, 0) ram = self.phase_ram_peaks.get(phase_key, 0) - if is_mps: - self.log(f" Phase {phase_num} ({phase_name}): {vram:.2f}GB", category="memory", force=force) + if arch == 'unified': + self.log(f"Phase {phase_num} ({phase_name}): {vram:.2f}GB", category="memory", indent_level=1, force=force) else: - self.log(f" Phase {phase_num} ({phase_name}): {_format_peak_with_swap(vram, total_vram_gb)} | RAM {ram:.2f}GB", category="memory", force=force) + self.log(f"Phase {phase_num} ({phase_name}): {_format_peak_with_swap(vram, total_vram_gb, arch)} | RAM {ram:.2f}GB", category="memory", indent_level=1, force=force) - if is_mps: - overall = max(self.phase_vram_peaks.values()) if self.phase_vram_peaks else 0 - self.log(f"Overall Peak: {overall:.2f}GB", category="memory", force=force) + overall_vram = max(self.phase_vram_peaks.values()) if self.phase_vram_peaks else 0 + overall_ram = max(self.phase_ram_peaks.values()) if self.phase_ram_peaks else 0 + + if arch == 'unified': + self.log(f"Overall peak: {overall_vram:.2f}GB", category="memory", force=force) else: - overall_vram = max(self.phase_vram_peaks.values()) if self.phase_vram_peaks else 0 - overall_ram = max(self.phase_ram_peaks.values()) if self.phase_ram_peaks else 0 - self.log(f"Overall peak: {_format_peak_with_swap(overall_vram, total_vram_gb)} | RAM {overall_ram:.2f}GB", category="memory", force=force) + self.log(f"Overall peak: {_format_peak_with_swap(overall_vram, total_vram_gb, arch)} | RAM {overall_ram:.2f}GB", category="memory", force=force) @torch._dynamo.disable # Skip tracing to avoid time.time() warnings def _store_checkpoint(self, label: str, metrics: Dict[str, Any]) -> None: From 7cbf02556122ba077438d3b273ea5aada418dc6d Mon Sep 17 00:00:00 2001 From: Adrien Toupet Date: Tue, 9 Dec 2025 23:51:51 -0500 Subject: [PATCH 5/8] Fix VRAM peak tracking: separate allocated vs reserved, Windows-only overflow - Track both peak_allocated (tensor usage) and peak_reserved (cache pool) per phase - peak_allocated resets properly between phases via reset_peak_memory_stats() - Overflow detection/warnings now Windows-only (WDDM paging behavior) - Remove get_memory_architecture() - replaced with simple is_mps + platform checks - Phase summary shows: VRAM XGB allocated, YGB reserved | RAM ZGB - Simplify MPS path (unified memory has no overflow concept) --- README.md | 11 +-- inference_cli.py | 4 +- src/interfaces/dit_model_loader.py | 10 +- src/optimization/memory_manager.py | 40 ++------ src/utils/debug.py | 146 +++++++++++++---------------- 5 files changed, 85 insertions(+), 126 deletions(-) diff --git a/README.md b/README.md index d177f37..5096b8e 100644 --- a/README.md +++ b/README.md @@ -418,12 +418,11 @@ Configure the DiT (Diffusion Transformer) model for video upscaling. - `sdpa`: PyTorch scaled_dot_product_attention (default, stable, always available) - `flash_attn`: Flash Attention 2 (faster on supported hardware, requires flash-attn package) -- **allow_vram_overflow**: Allow VRAM to overflow to system RAM - - `False` (default): Strict VRAM limit - prevents silent swap but OOMs if exceeded - - `True`: Allow overflow - prevents OOM but may cause severe slowdown when physical VRAM exceeded - - Last resort when other memory optimizations are insufficient - - Requires ComfyUI restart to change setting - - No effect on Apple Silicon (unified memory architecture) +- **allow_vram_overflow**: Windows only - allow VRAM to overflow to system RAM + - `False` (default): Strict VRAM limit - faster when within limits + - `True`: Allow overflow - prevents OOM but causes severe slowdown + - Last resort when other optimizations are insufficient + - Requires ComfyUI restart to change - **torch_compile_args**: Connect to SeedVR2 Torch Compile Settings node for 20-40% speedup diff --git a/inference_cli.py b/inference_cli.py index 342231e..338f467 100644 --- a/inference_cli.py +++ b/inference_cli.py @@ -1336,8 +1336,8 @@ Examples: blockswap_group.add_argument("--swap_io_components", action="store_true", help="Offload DiT I/O layers for extra VRAM savings. Requires --dit_offload_device") blockswap_group.add_argument("--allow_vram_overflow", action="store_true", - help="Allow VRAM overflow to system RAM. Prevents OOM but may cause severe slowdown. " - "Last resort when other memory optimizations are insufficient. No effect on Apple Silicon (unified memory).") + help="Windows only: Allow VRAM overflow to system RAM. Prevents OOM but causes severe slowdown. " + "Last resort when other optimizations are insufficient.") # VAE Tiling vae_group = parser.add_argument_group('VAE tiling (for high resolution upscale)') diff --git a/src/interfaces/dit_model_loader.py b/src/interfaces/dit_model_loader.py index 76963d3..e26c970 100644 --- a/src/interfaces/dit_model_loader.py +++ b/src/interfaces/dit_model_loader.py @@ -116,12 +116,12 @@ class SeedVR2LoadDiTModel(io.ComfyNode): default=False, optional=True, tooltip=( - "Allow VRAM to overflow to system RAM when physical VRAM is exceeded.\n" - "โ€ข False (default): Strict VRAM limit - OOM if exceeded (faster when within limits)\n" - "โ€ข True: Allow overflow to RAM - prevents OOM but may cause severe slowdown\n" + "Windows only: Allow VRAM to overflow to system RAM.\n" + "โ€ข False (default): Strict VRAM limit - faster when within limits\n" + "โ€ข True: Allow overflow - prevents OOM but may cause severe slowdown\n" "\n" - "Last resort when other memory optimizations are insufficient.\n" - "Requires ComfyUI restart to change. No effect on Apple Silicon (unified memory)." + "Last resort when other optimizations are insufficient.\n" + "Requires ComfyUI restart to change." ) ), io.Custom("TORCH_COMPILE_ARGS").Input("torch_compile_args", diff --git a/src/optimization/memory_manager.py b/src/optimization/memory_manager.py index b8150a6..a61d5ef 100644 --- a/src/optimization/memory_manager.py +++ b/src/optimization/memory_manager.py @@ -12,7 +12,7 @@ import time import psutil import platform from typing import Tuple, Dict, Any, Optional, List, Union - + def _device_str(device: Union[torch.device, str]) -> str: """Normalized uppercase device string for comparison and logging. MPS variants โ†’ 'MPS'.""" @@ -45,27 +45,6 @@ def get_gpu_backend() -> str: return 'cpu' -def get_memory_architecture() -> str: - """Get memory architecture type for swap/overflow detection. - - This combines GPU backend with OS platform to determine how - GPU memory overflow is handled: - - Returns: - 'unified': macOS unified memory (MPS) - GPU/CPU share memory pool - 'discrete_paged': Windows WDDM - GPU memory can page to system RAM - 'discrete_strict': Linux - No automatic GPU paging, OOM on overflow - 'cpu_only': No GPU backend available - """ - if is_mps_available(): - return 'unified' - if is_cuda_available(): - if platform.system() == 'Windows': - return 'discrete_paged' - return 'discrete_strict' - return 'cpu_only' - - def get_device_list(include_none: bool = False, include_cpu: bool = False) -> List[str]: """ Get list of available compute devices for SeedVR2 @@ -215,7 +194,7 @@ def was_vram_limit_change_attempted() -> bool: return _vram_limit_change_attempted -def get_vram_usage(device: Optional[torch.device] = None, debug: Optional['Debug'] = None) -> Tuple[float, float, float]: +def get_vram_usage(device: Optional[torch.device] = None, debug: Optional['Debug'] = None) -> Tuple[float, float, float, float]: """ Get current VRAM usage metrics for monitoring. Used for tracking memory consumption during processing. @@ -225,8 +204,8 @@ def get_vram_usage(device: Optional[torch.device] = None, debug: Optional['Debug debug: Optional debug instance for logging Returns: - tuple: (allocated_gb, reserved_gb, max_reserved_gb) - Returns (0, 0, 0) if no GPU available + tuple: (allocated_gb, reserved_gb, peak_allocated_gb, peak_reserved_gb) + Returns (0, 0, 0, 0) if no GPU available """ try: if is_cuda_available(): @@ -236,18 +215,19 @@ def get_vram_usage(device: Optional[torch.device] = None, debug: Optional['Debug device = torch.device(device) allocated = torch.cuda.memory_allocated(device) / (1024**3) reserved = torch.cuda.memory_reserved(device) / (1024**3) - max_reserved = torch.cuda.max_memory_reserved(device) / (1024**3) - return allocated, reserved, max_reserved + peak_allocated = torch.cuda.max_memory_allocated(device) / (1024**3) + peak_reserved = torch.cuda.max_memory_reserved(device) / (1024**3) + return allocated, reserved, peak_allocated, peak_reserved elif is_mps_available(): # MPS doesn't support per-device queries - uses global memory tracking allocated = torch.mps.current_allocated_memory() / (1024**3) reserved = torch.mps.driver_allocated_memory() / (1024**3) - max_allocated = allocated # MPS doesn't track peak separately - return allocated, reserved, max_allocated + # MPS doesn't track peak separately + return allocated, reserved, allocated, reserved except Exception as e: if debug: debug.log(f"Failed to get VRAM usage: {e}", level="WARNING", category="memory", force=True) - return 0.0, 0.0, 0.0 + return 0.0, 0.0, 0.0, 0.0 def get_ram_usage(debug: Optional['Debug'] = None) -> Tuple[float, float, float, float]: diff --git a/src/utils/debug.py b/src/utils/debug.py index 10f3c2c..30eb23a 100644 --- a/src/utils/debug.py +++ b/src/utils/debug.py @@ -10,6 +10,7 @@ import torch import gc from typing import Optional, List, Dict, Any, Union from datetime import datetime +import platform from ..optimization.memory_manager import ( get_vram_usage, get_basic_vram_info, @@ -18,36 +19,26 @@ from ..optimization.memory_manager import ( is_vram_overflow_allowed, was_vram_limit_change_attempted, is_mps_available, - is_cuda_available, - get_memory_architecture + is_cuda_available ) from ..utils.constants import __version__ -def _format_peak_with_swap(peak_gb: float, total_vram_gb: float, arch: str = None) -> str: - """Format peak memory with architecture-aware overflow reporting. +def _format_peak_with_swap(peak_gb: float, total_vram_gb: float) -> str: + """Format peak memory, showing overflow breakdown on Windows. Args: peak_gb: Peak reserved memory from PyTorch total_vram_gb: Physical GPU VRAM capacity - arch: Memory architecture from get_memory_architecture(), or None to auto-detect """ if total_vram_gb <= 0: return f"{peak_gb:.2f}GB" overflow_gb = peak_gb - total_vram_gb - if overflow_gb <= 0: + if overflow_gb <= 0 or platform.system() != 'Windows': return f"{peak_gb:.2f}GB" - if arch is None: - arch = get_memory_architecture() - - if arch == 'discrete_paged': - return f"{peak_gb:.2f}GB ({total_vram_gb:.0f}GB GPU + {overflow_gb:.2f}GB system RAM)" - elif arch == 'discrete_strict': - return f"{peak_gb:.2f}GB (exceeded {total_vram_gb:.0f}GB by {overflow_gb:.2f}GB)" - # unified or cpu_only - no swap concept - return f"{peak_gb:.2f}GB" + return f"{peak_gb:.2f}GB ({total_vram_gb:.0f}GB GPU + {overflow_gb:.2f}GB system RAM)" class Debug: @@ -108,7 +99,8 @@ class Debug: self.vram_history: List[float] = [] self.active_timer_stack: List[str] = [] self.timer_namespace: str = "" - self.phase_vram_peaks: Dict[str, float] = {} + self.phase_vram_peaks_alloc: Dict[str, float] = {} + self.phase_vram_peaks_rsv: Dict[str, float] = {} self.phase_ram_peaks: Dict[str, float] = {} @torch._dynamo.disable # Skip tracing to avoid datetime.now() warnings @@ -252,30 +244,19 @@ class Debug: self.log(f"{cuda_line} | ComfyUI: {comfy_str}" if comfy_str else cuda_line, category="info") def _print_vram_overflow_status(self) -> bool: - """Print VRAM overflow status - warnings always shown, info only in debug mode. - - Returns: - True if a forced warning was printed, False otherwise. - """ - is_mps = is_mps_available() - force = False + """Print VRAM overflow status (Windows only). Returns True if warning was printed.""" + if platform.system() != 'Windows': + return False if was_vram_limit_change_attempted(): self.log("allow_vram_overflow setting changed - restart ComfyUI to apply", level="WARNING", category="memory", force=True) - force = True + return True elif is_vram_overflow_allowed(): - if is_mps: - self.log("allow_vram_overflow: enabled (no effect on Apple Silicon unified memory)", category="info") - else: - self.log("allow_vram_overflow: enabled - may cause severe slowdown if physical VRAM exceeded", level="WARNING", category="memory", force=True) - force = True + self.log("allow_vram_overflow: enabled - may cause severe slowdown if physical VRAM exceeded", level="WARNING", category="memory", force=True) + return True else: - if is_mps: - self.log("allow_vram_overflow: disabled (no effect on Apple Silicon unified memory)", category="info") - else: - self.log("allow_vram_overflow: disabled (recommended for best performance)", category="success") - - return force + self.log("allow_vram_overflow: disabled (recommended)", category="success") + return False def print_footer(self) -> None: """Print the footer with links - always displayed""" @@ -446,19 +427,13 @@ class Debug: if show_diff and self.memory_checkpoints: self._log_memory_diff(current_metrics=memory_info, force=force) - # Architecture-aware overflow warnings - arch = memory_info.get('arch', 'cpu_only') + # Overflow warning (Windows only - WDDM can page to system RAM) overflow = memory_info.get('vram_overflow', 0.0) - if overflow > 0 and not is_vram_overflow_allowed(): - if arch == 'discrete_paged': - self.log(f"VRAM overflow: {overflow:.2f}GB paged to system RAM - severe slowdown expected. " - "Consider optimizing (e.g., reduce resolution, batch size, enable BlockSwap, VAE tiling...).", - level="WARNING", category="memory", force=True) - elif arch == 'discrete_strict': - self.log(f"VRAM exceeded physical limit by {overflow:.2f}GB - OOM risk. " - "Consider optimizing (e.g., reduce resolution, batch size, enable BlockSwap, VAE tiling...).", - level="WARNING", category="memory", force=True) + if overflow > 0 and platform.system() == 'Windows' and not is_vram_overflow_allowed(): + self.log(f"VRAM overflow: {overflow:.2f}GB paged to system RAM - severe slowdown expected. " + "Consider optimizing (e.g., reduce resolution, batch size, enable BlockSwap, VAE tiling...).", + level="WARNING", category="memory", force=True) # Log detailed analysis if requested if detailed_tensors and tensor_stats.get('details'): @@ -469,10 +444,15 @@ class Debug: # Update phase peaks if we're in an active phase if self.current_phase: - if memory_info['vram_peak_since_last'] > 0: - self.phase_vram_peaks[self.current_phase] = max( - self.phase_vram_peaks.get(self.current_phase, 0), - memory_info['vram_peak_since_last'] + if memory_info['vram_peak_alloc'] > 0: + self.phase_vram_peaks_alloc[self.current_phase] = max( + self.phase_vram_peaks_alloc.get(self.current_phase, 0), + memory_info['vram_peak_alloc'] + ) + if memory_info['vram_peak_rsv'] > 0: + self.phase_vram_peaks_rsv[self.current_phase] = max( + self.phase_vram_peaks_rsv.get(self.current_phase, 0), + memory_info['vram_peak_rsv'] ) if memory_info['ram_process'] > 0: self.phase_ram_peaks[self.current_phase] = max( @@ -484,17 +464,18 @@ class Debug: reset_vram_peak(device=None, debug=self) def _collect_memory_metrics(self) -> Dict[str, Any]: - """Collect current memory metrics with architecture-aware reporting.""" - arch = get_memory_architecture() + """Collect current memory metrics.""" + is_mps = is_mps_available() + has_gpu = is_mps or is_cuda_available() metrics = { 'vram_allocated': 0.0, 'vram_reserved': 0.0, 'vram_free': 0.0, 'vram_total': 0.0, - 'vram_peak_since_last': 0.0, + 'vram_peak_alloc': 0.0, + 'vram_peak_rsv': 0.0, 'vram_overflow': 0.0, - 'arch': arch, 'ram_process': 0.0, 'ram_available': 0.0, 'ram_total': 0.0, @@ -503,28 +484,26 @@ class Debug: 'summary_ram': "" } - if arch == 'cpu_only': - pass # No GPU metrics - else: - metrics['vram_allocated'], metrics['vram_reserved'], metrics['vram_peak_since_last'] = get_vram_usage(device=None, debug=self) + if has_gpu: + metrics['vram_allocated'], metrics['vram_reserved'], metrics['vram_peak_alloc'], metrics['vram_peak_rsv'] = get_vram_usage(device=None, debug=self) vram_info = get_basic_vram_info(device=None) if "error" not in vram_info and vram_info["total_gb"] > 0: metrics['vram_free'] = vram_info["free_gb"] metrics['vram_total'] = vram_info["total_gb"] + metrics['vram_overflow'] = max(0.0, metrics['vram_peak_rsv'] - metrics['vram_total']) - # Calculate overflow: reserved beyond physical VRAM - metrics['vram_overflow'] = max(0.0, metrics['vram_peak_since_last'] - metrics['vram_total']) - - backend = "Unified Memory" if arch == 'unified' else "VRAM" - peak_str = _format_peak_with_swap(metrics['vram_peak_since_last'], metrics['vram_total'], arch) + backend = "Unified Memory" if is_mps else "VRAM" + peak_alloc_str = _format_peak_with_swap(metrics['vram_peak_alloc'], metrics['vram_total']) metrics['summary_vram'] = ( f" [{backend}] {metrics['vram_allocated']:.2f}GB allocated / " f"{metrics['vram_reserved']:.2f}GB reserved / " - f"Peak: {peak_str} / " + f"Peak: {peak_alloc_str} / " f"{metrics['vram_free']:.2f}GB free / " f"{metrics['vram_total']:.2f}GB total" ) + + self.vram_history.append(metrics['vram_reserved']) # RAM metrics metrics['ram_process'], metrics['ram_available'], metrics['ram_total'], metrics['ram_others'] = get_ram_usage(debug=self) @@ -537,10 +516,6 @@ class Debug: f"{metrics['ram_total']:.2f}GB total" ) - # Track reserved (matches nvidia-smi) for pressure history - if arch != 'cpu_only': - self.vram_history.append(metrics['vram_reserved']) - return metrics def _collect_tensor_stats(self, detailed: bool = False) -> Dict[str, Any]: @@ -667,8 +642,8 @@ class Debug: self.log(f"Memory changes: {', '.join(diffs)}", category="memory", force=force, indent_level=1) def log_peak_memory_summary(self, force: bool = True) -> None: - """Display peak memory usage across all phases (VRAM and RAM combined)""" - if not self.phase_vram_peaks and not self.phase_ram_peaks: + """Display peak memory usage across all phases.""" + if not self.phase_vram_peaks_alloc and not self.phase_ram_peaks: return phase_names = { @@ -678,11 +653,11 @@ class Debug: 'phase4': 'Post-processing' } - arch = get_memory_architecture() + is_mps = is_mps_available() - # Get total VRAM for overflow detection + # Get total VRAM for overflow formatting (Windows only) total_vram_gb = 0.0 - if arch not in ('unified', 'cpu_only'): + if not is_mps: vram_info = get_basic_vram_info(device=None) if "error" not in vram_info: total_vram_gb = vram_info["total_gb"] @@ -691,25 +666,29 @@ class Debug: self.log("โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€", category="none", force=force) self.log("Peak memory by phase:", category="memory", force=force) - all_phases = sorted(set(self.phase_vram_peaks.keys()) | set(self.phase_ram_peaks.keys())) + all_phases = sorted(set(self.phase_vram_peaks_alloc.keys()) | set(self.phase_ram_peaks.keys())) for phase_key in all_phases: phase_num = phase_key[-1] phase_name = phase_names.get(phase_key, phase_key) - vram = self.phase_vram_peaks.get(phase_key, 0) + alloc = self.phase_vram_peaks_alloc.get(phase_key, 0) + rsv = self.phase_vram_peaks_rsv.get(phase_key, 0) ram = self.phase_ram_peaks.get(phase_key, 0) - if arch == 'unified': - self.log(f"Phase {phase_num} ({phase_name}): {vram:.2f}GB", category="memory", indent_level=1, force=force) + if is_mps: + self.log(f"{phase_num}. {phase_name}: {alloc:.2f}GB", category="memory", indent_level=1, force=force) else: - self.log(f"Phase {phase_num} ({phase_name}): {_format_peak_with_swap(vram, total_vram_gb, arch)} | RAM {ram:.2f}GB", category="memory", indent_level=1, force=force) + rsv_str = _format_peak_with_swap(rsv, total_vram_gb) + self.log(f"{phase_num}. {phase_name}: VRAM {alloc:.2f}GB allocated, {rsv_str} reserved | RAM {ram:.2f}GB", category="memory", indent_level=1, force=force) - overall_vram = max(self.phase_vram_peaks.values()) if self.phase_vram_peaks else 0 + overall_alloc = max(self.phase_vram_peaks_alloc.values()) if self.phase_vram_peaks_alloc else 0 + overall_rsv = max(self.phase_vram_peaks_rsv.values()) if self.phase_vram_peaks_rsv else 0 overall_ram = max(self.phase_ram_peaks.values()) if self.phase_ram_peaks else 0 - if arch == 'unified': - self.log(f"Overall peak: {overall_vram:.2f}GB", category="memory", force=force) + if is_mps: + self.log(f"Overall peak: {overall_alloc:.2f}GB", category="memory", force=force) else: - self.log(f"Overall peak: {_format_peak_with_swap(overall_vram, total_vram_gb, arch)} | RAM {overall_ram:.2f}GB", category="memory", force=force) + overall_rsv_str = _format_peak_with_swap(overall_rsv, total_vram_gb) + self.log(f"Overall peak: VRAM {overall_alloc:.2f}GB allocated, {overall_rsv_str} reserved | RAM {overall_ram:.2f}GB", category="memory", force=force) @torch._dynamo.disable # Skip tracing to avoid time.time() warnings def _store_checkpoint(self, label: str, metrics: Dict[str, Any]) -> None: @@ -819,6 +798,7 @@ class Debug: self.timer_durations.clear() self.timer_messages.clear() self.active_timer_stack.clear() - self.phase_vram_peaks.clear() + self.phase_vram_peaks_alloc.clear() + self.phase_vram_peaks_rsv.clear() self.phase_ram_peaks.clear() self.current_phase = None \ No newline at end of file From c010deeea1435dd1f2b990a2b6216cb32e8b8056 Mon Sep 17 00:00:00 2001 From: Adrien Toupet Date: Wed, 10 Dec 2025 00:56:27 -0500 Subject: [PATCH 6/8] Remove ineffective allow_vram_overflow setting - PyTorch's set_per_process_memory_fraction cannot prevent WDDM paging on Windows - Keep overflow detection and warning when VRAM exceeds physical limit - Simplify peak memory formatting - Remove setting from CLI, ComfyUI node, and memory_manager --- README.md | 6 ---- inference_cli.py | 10 +----- src/interfaces/dit_model_loader.py | 19 +--------- src/optimization/memory_manager.py | 56 ------------------------------ src/utils/debug.py | 48 +++++++------------------ 5 files changed, 15 insertions(+), 124 deletions(-) diff --git a/README.md b/README.md index 5096b8e..e6d2d65 100644 --- a/README.md +++ b/README.md @@ -418,12 +418,6 @@ Configure the DiT (Diffusion Transformer) model for video upscaling. - `sdpa`: PyTorch scaled_dot_product_attention (default, stable, always available) - `flash_attn`: Flash Attention 2 (faster on supported hardware, requires flash-attn package) -- **allow_vram_overflow**: Windows only - allow VRAM to overflow to system RAM - - `False` (default): Strict VRAM limit - faster when within limits - - `True`: Allow overflow - prevents OOM but causes severe slowdown - - Last resort when other optimizations are insufficient - - Requires ComfyUI restart to change - - **torch_compile_args**: Connect to SeedVR2 Torch Compile Settings node for 20-40% speedup **BlockSwap Explained:** diff --git a/inference_cli.py b/inference_cli.py index 338f467..a794993 100644 --- a/inference_cli.py +++ b/inference_cli.py @@ -79,7 +79,6 @@ else: # Pre-parse arguments that must be handled before torch import _pre_parser = argparse.ArgumentParser(add_help=False) _pre_parser.add_argument("--cuda_device", type=str, default=None) - _pre_parser.add_argument("--allow_vram_overflow", action="store_true") _pre_args, _ = _pre_parser.parse_known_args() if _pre_args.cuda_device is not None: @@ -128,13 +127,9 @@ from src.core.generation_phases import ( postprocess_all_batches ) from src.utils.debug import Debug -from src.optimization.memory_manager import clear_memory, configure_vram_limit, get_gpu_backend, is_cuda_available +from src.optimization.memory_manager import clear_memory, get_gpu_backend, is_cuda_available debug = Debug(enabled=False) # Will be enabled via --debug CLI flag -# Configure VRAM limit (must be before any CUDA allocations) -if platform.system() != "Darwin": - configure_vram_limit(allow_overflow=_pre_args.allow_vram_overflow) - # ============================================================================= # Device Management Helpers # ============================================================================= @@ -1335,9 +1330,6 @@ Examples: "Requires --dit_offload_device. Default: 0 (disabled)") blockswap_group.add_argument("--swap_io_components", action="store_true", help="Offload DiT I/O layers for extra VRAM savings. Requires --dit_offload_device") - blockswap_group.add_argument("--allow_vram_overflow", action="store_true", - help="Windows only: Allow VRAM overflow to system RAM. Prevents OOM but causes severe slowdown. " - "Last resort when other optimizations are insufficient.") # VAE Tiling vae_group = parser.add_argument_group('VAE tiling (for high resolution upscale)') diff --git a/src/interfaces/dit_model_loader.py b/src/interfaces/dit_model_loader.py index e26c970..1064571 100644 --- a/src/interfaces/dit_model_loader.py +++ b/src/interfaces/dit_model_loader.py @@ -7,7 +7,7 @@ from comfy_api.latest import io from comfy_execution.utils import get_executing_context from typing import Dict, Any, Tuple from ..utils.model_registry import get_available_dit_models, DEFAULT_DIT -from ..optimization.memory_manager import get_device_list, configure_vram_limit +from ..optimization.memory_manager import get_device_list class SeedVR2LoadDiTModel(io.ComfyNode): @@ -112,18 +112,6 @@ class SeedVR2LoadDiTModel(io.ComfyNode): "Flash Attention provides speedup through optimized CUDA kernels on compatible GPUs." ) ), - io.Boolean.Input("allow_vram_overflow", - default=False, - optional=True, - tooltip=( - "Windows only: Allow VRAM to overflow to system RAM.\n" - "โ€ข False (default): Strict VRAM limit - faster when within limits\n" - "โ€ข True: Allow overflow - prevents OOM but may cause severe slowdown\n" - "\n" - "Last resort when other optimizations are insufficient.\n" - "Requires ComfyUI restart to change." - ) - ), io.Custom("TORCH_COMPILE_ARGS").Input("torch_compile_args", optional=True, tooltip=( @@ -143,7 +131,6 @@ class SeedVR2LoadDiTModel(io.ComfyNode): def execute(cls, model: str, device: str, offload_device: str = "none", cache_model: bool = False, blocks_to_swap: int = 0, swap_io_components: bool = False, attention_mode: str = "sdpa", - allow_vram_overflow: bool = False, torch_compile_args: Dict[str, Any] = None) -> io.NodeOutput: """ Create DiT model configuration for SeedVR2 main node @@ -156,7 +143,6 @@ class SeedVR2LoadDiTModel(io.ComfyNode): blocks_to_swap: Number of transformer blocks to swap (requires offload_device != device) swap_io_components: Whether to offload I/O components (requires offload_device != device) attention_mode: Attention computation backend ('sdpa' or 'flash_attn') - allow_vram_overflow: Allow VRAM overflow to system RAM (prevents OOM but slower) torch_compile_args: Optional torch.compile configuration from settings node Returns: @@ -182,9 +168,6 @@ class SeedVR2LoadDiTModel(io.ComfyNode): "(e.g., 'cpu' or another device). Set cache_model=False if you don't want to cache the model." ) - # Configure VRAM limit enforcement (once per session, first call wins) - configure_vram_limit(allow_overflow=allow_vram_overflow) - config = { "model": model, "device": device, diff --git a/src/optimization/memory_manager.py b/src/optimization/memory_manager.py index a61d5ef..fb2fc75 100644 --- a/src/optimization/memory_manager.py +++ b/src/optimization/memory_manager.py @@ -138,62 +138,6 @@ else: print(f"โš ๏ธ Memory check failed: {vram_info['error']} - No available backend!") -# VRAM overflow configuration state -_vram_overflow_allowed: bool = True -_vram_limit_configured: bool = False -_vram_limit_change_attempted: bool = False - - -def configure_vram_limit(allow_overflow: bool = False) -> bool: - """ - Configure VRAM limit enforcement. Call early before heavy CUDA usage. - - Args: - allow_overflow: If True, allow VRAM overflow to system RAM (prevents OOM but may be slow). - If False (default), enforce strict physical VRAM limit. - - Returns: - True if configuration applied successfully, False otherwise - - Note: - Can only be configured once per session. Restart required to change. - """ - global _vram_overflow_allowed, _vram_limit_configured, _vram_limit_change_attempted - - # Already configured this session - track if user tried to change - if _vram_limit_configured: - if _vram_overflow_allowed != allow_overflow: - _vram_limit_change_attempted = True - return _vram_overflow_allowed == allow_overflow - - _vram_limit_configured = True - _vram_overflow_allowed = allow_overflow - - if allow_overflow: - return True - - if not is_cuda_available(): - return True - - try: - for i in range(torch.cuda.device_count()): - torch.cuda.set_per_process_memory_fraction(1.0, i) - return True - except RuntimeError: - _vram_overflow_allowed = True - return False - - -def is_vram_overflow_allowed() -> bool: - """Check if VRAM overflow to system RAM is allowed.""" - return _vram_overflow_allowed - - -def was_vram_limit_change_attempted() -> bool: - """Check if user tried to change VRAM limit setting after initial configuration.""" - return _vram_limit_change_attempted - - def get_vram_usage(device: Optional[torch.device] = None, debug: Optional['Debug'] = None) -> Tuple[float, float, float, float]: """ Get current VRAM usage metrics for monitoring. diff --git a/src/utils/debug.py b/src/utils/debug.py index 30eb23a..3619084 100644 --- a/src/utils/debug.py +++ b/src/utils/debug.py @@ -15,30 +15,28 @@ from ..optimization.memory_manager import ( get_vram_usage, get_basic_vram_info, get_ram_usage, - reset_vram_peak, - is_vram_overflow_allowed, - was_vram_limit_change_attempted, + reset_vram_peak, is_mps_available, is_cuda_available ) from ..utils.constants import __version__ -def _format_peak_with_swap(peak_gb: float, total_vram_gb: float) -> str: - """Format peak memory, showing overflow breakdown on Windows. +def _format_peak_with_overflow(peak_gb: float, total_vram_gb: float) -> str: + """Format peak reserved memory, showing overflow breakdown on Windows. Args: peak_gb: Peak reserved memory from PyTorch total_vram_gb: Physical GPU VRAM capacity """ if total_vram_gb <= 0: - return f"{peak_gb:.2f}GB" + return f"{peak_gb:.2f}GB reserved" overflow_gb = peak_gb - total_vram_gb if overflow_gb <= 0 or platform.system() != 'Windows': - return f"{peak_gb:.2f}GB" + return f"{peak_gb:.2f}GB reserved" - return f"{peak_gb:.2f}GB ({total_vram_gb:.0f}GB GPU + {overflow_gb:.2f}GB system RAM)" + return f"{peak_gb:.2f}GB reserved ({total_vram_gb:.0f}GB GPU + {overflow_gb:.2f}GB overflow)" class Debug: @@ -176,11 +174,6 @@ class Debug: # Environment info - only in debug mode if self.enabled: self._print_environment_info(cli) - - # VRAM overflow status - warnings always shown - vram_warning_shown = self._print_vram_overflow_status() - - self.log("", category="none", force=vram_warning_shown) def _print_environment_info(self, cli: bool = False) -> None: """Print concise environment info for bug reports - zero cost when debug disabled""" @@ -242,21 +235,7 @@ class Debug: self.log(f"Python: {py_ver} | PyTorch: {torch_ver} | Flash Attn: {flash_str} | Triton: {triton_str}", category="info") cuda_line = f"CUDA: {cuda_ver} | cuDNN: {cudnn_ver}" self.log(f"{cuda_line} | ComfyUI: {comfy_str}" if comfy_str else cuda_line, category="info") - - def _print_vram_overflow_status(self) -> bool: - """Print VRAM overflow status (Windows only). Returns True if warning was printed.""" - if platform.system() != 'Windows': - return False - - if was_vram_limit_change_attempted(): - self.log("allow_vram_overflow setting changed - restart ComfyUI to apply", level="WARNING", category="memory", force=True) - return True - elif is_vram_overflow_allowed(): - self.log("allow_vram_overflow: enabled - may cause severe slowdown if physical VRAM exceeded", level="WARNING", category="memory", force=True) - return True - else: - self.log("allow_vram_overflow: disabled (recommended)", category="success") - return False + self.log("", category="none") def print_footer(self) -> None: """Print the footer with links - always displayed""" @@ -430,7 +409,7 @@ class Debug: # Overflow warning (Windows only - WDDM can page to system RAM) overflow = memory_info.get('vram_overflow', 0.0) - if overflow > 0 and platform.system() == 'Windows' and not is_vram_overflow_allowed(): + if overflow > 0 and platform.system() == 'Windows': self.log(f"VRAM overflow: {overflow:.2f}GB paged to system RAM - severe slowdown expected. " "Consider optimizing (e.g., reduce resolution, batch size, enable BlockSwap, VAE tiling...).", level="WARNING", category="memory", force=True) @@ -494,11 +473,10 @@ class Debug: metrics['vram_overflow'] = max(0.0, metrics['vram_peak_rsv'] - metrics['vram_total']) backend = "Unified Memory" if is_mps else "VRAM" - peak_alloc_str = _format_peak_with_swap(metrics['vram_peak_alloc'], metrics['vram_total']) metrics['summary_vram'] = ( f" [{backend}] {metrics['vram_allocated']:.2f}GB allocated / " f"{metrics['vram_reserved']:.2f}GB reserved / " - f"Peak: {peak_alloc_str} / " + f"Peak: {metrics['vram_peak_alloc']:.2f}GB / " f"{metrics['vram_free']:.2f}GB free / " f"{metrics['vram_total']:.2f}GB total" ) @@ -677,8 +655,8 @@ class Debug: if is_mps: self.log(f"{phase_num}. {phase_name}: {alloc:.2f}GB", category="memory", indent_level=1, force=force) else: - rsv_str = _format_peak_with_swap(rsv, total_vram_gb) - self.log(f"{phase_num}. {phase_name}: VRAM {alloc:.2f}GB allocated, {rsv_str} reserved | RAM {ram:.2f}GB", category="memory", indent_level=1, force=force) + rsv_str = _format_peak_with_overflow(rsv, total_vram_gb) + self.log(f"{phase_num}. {phase_name}: VRAM {alloc:.2f}GB allocated, {rsv_str} | RAM {ram:.2f}GB", category="memory", indent_level=1, force=force) overall_alloc = max(self.phase_vram_peaks_alloc.values()) if self.phase_vram_peaks_alloc else 0 overall_rsv = max(self.phase_vram_peaks_rsv.values()) if self.phase_vram_peaks_rsv else 0 @@ -687,8 +665,8 @@ class Debug: if is_mps: self.log(f"Overall peak: {overall_alloc:.2f}GB", category="memory", force=force) else: - overall_rsv_str = _format_peak_with_swap(overall_rsv, total_vram_gb) - self.log(f"Overall peak: VRAM {overall_alloc:.2f}GB allocated, {overall_rsv_str} reserved | RAM {overall_ram:.2f}GB", category="memory", force=force) + overall_rsv_str = _format_peak_with_overflow(overall_rsv, total_vram_gb) + self.log(f"Overall peak: VRAM {overall_alloc:.2f}GB allocated, {overall_rsv_str} | RAM {overall_ram:.2f}GB", category="memory", force=force) @torch._dynamo.disable # Skip tracing to avoid time.time() warnings def _store_checkpoint(self, label: str, metrics: Dict[str, Any]) -> None: From 610668156343a92b352f823cf1c4b764a58456a3 Mon Sep 17 00:00:00 2001 From: Adrien Toupet Date: Wed, 10 Dec 2025 01:44:49 -0500 Subject: [PATCH 7/8] Fix graceful fallback from flash-attn #376 Add compatibility shims for corrupted flash_attn/xformers DLLs. Force-verify flash_attn_2_cuda at startup; fall back to SDPA if unavailable. --- src/optimization/compatibility.py | 65 +++++++++++++++++++++++++++---- 1 file changed, 58 insertions(+), 7 deletions(-) diff --git a/src/optimization/compatibility.py b/src/optimization/compatibility.py index c63b007..696efc9 100644 --- a/src/optimization/compatibility.py +++ b/src/optimization/compatibility.py @@ -5,9 +5,10 @@ Contains FP8/FP16 compatibility layers and wrappers for different model architec Extracted from: seedvr2.py (lines 1045-1630) """ -# Triton compatibility shim for bitsandbytes 0.45+ with triton 3.0+ -# Must be called before any diffusers import +# Compatibility shims - Must run before any torch/diffusers import import sys +import types + def ensure_triton_compat(): """Create minimal triton.ops stubs only if missing, to allow bitsandbytes import.""" @@ -20,8 +21,6 @@ def ensure_triton_compat(): except (ImportError, ModuleNotFoundError, AttributeError): pass - import types - if 'triton.ops' not in sys.modules: sys.modules['triton.ops'] = types.ModuleType('triton.ops') @@ -32,12 +31,62 @@ def ensure_triton_compat(): sys.modules['triton.ops'].matmul_perf_model = matmul_perf sys.modules['triton.ops.matmul_perf_model'] = matmul_perf -# Run immediately on import + +def ensure_flash_attn_safe(): + """ + Pre-test flash_attn package; stub if DLL is broken. + Prevents diffusers from crashing when flash_attn has broken DLLs. + """ + if 'flash_attn' in sys.modules: + return # Already loaded + + try: + import flash_attn + except (ImportError, OSError): + # DLL broken or not installed - create stub with proper __spec__ + import importlib.machinery + + stub = types.ModuleType('flash_attn') + stub.__spec__ = importlib.machinery.ModuleSpec('flash_attn', None) + stub.__file__ = None + stub.__path__ = [] + stub.__loader__ = None + # Provide attributes that diffusers/transformers import + stub.flash_attn_func = None + stub.flash_attn_varlen_func = None + sys.modules['flash_attn'] = stub + + +def ensure_xformers_flash_compat(): + """ + Pre-test xformers._C_flashattention; stub if DLL is broken. + Prevents xformers.ops.fmha.flash from crashing on import. + """ + if 'xformers._C_flashattention' in sys.modules: + return # Already loaded + + try: + from xformers import _C_flashattention # noqa: F401 + except (ImportError, OSError): + # DLL broken or not installed - create stub that fails gracefully + class _FailingStub(types.ModuleType): + """Stub that lets xformers gracefully disable its flash backend.""" + def __getattr__(self, name): + # Dunder attributes: raise AttributeError (normal Python behavior) + if name.startswith('__') and name.endswith('__'): + raise AttributeError(name) + # xformers functional attributes: raise ImportError so xformers catches it + raise ImportError("_C_flashattention unavailable") + sys.modules['xformers._C_flashattention'] = _FailingStub('xformers._C_flashattention') + + +# Run all shims immediately on import, before torch/diffusers ensure_triton_compat() +ensure_flash_attn_safe() +ensure_xformers_flash_compat() import torch -import types import os @@ -45,8 +94,10 @@ import os # 1. Flash Attention - speedup for attention operations try: from flash_attn import flash_attn_varlen_func + # Force load the CUDA extension to verify it's not corrupted + import flash_attn_2_cuda # noqa: F401 FLASH_ATTN_AVAILABLE = True -except ImportError: +except (ImportError, AttributeError, OSError): flash_attn_varlen_func = None FLASH_ATTN_AVAILABLE = False From 118c9fcbe7f26b1f1502999cdc0b633ac4bfaa3d Mon Sep 17 00:00:00 2001 From: Adrien Toupet Date: Wed, 10 Dec 2025 01:48:57 -0500 Subject: [PATCH 8/8] Release v2.5.19: new logo, remove dead flash-attn wrapper, graceful DLL fallback, improved VRAM tracking, revert VRAM limit --- README.md | 9 +++++++++ pyproject.toml | 2 +- src/utils/constants.py | 2 +- 3 files changed, 11 insertions(+), 2 deletions(-) diff --git a/README.md b/README.md index e6d2d65..ec659a6 100644 --- a/README.md +++ b/README.md @@ -36,6 +36,15 @@ We're actively working on improvements and new features. To stay informed: ## ๐Ÿš€ Updates +**2025.12.10 - Version 2.5.19** + +- **๐ŸŽจ New header logo design** - Refreshed ASCII art banner *(thanks [@naxci1](https://github.com/naxci1))* +- **๐Ÿงน Remove dead flash attention wrapper** - Removed legacy code from FP8CompatibleDiT; FlashAttentionVarlen already handles backend switching via its `attention_mode` attribute +- **๐Ÿ›ก๏ธ Fix graceful fallback from flash-attn** - Add compatibility shims for corrupted flash_attn/xformers DLLs, preventing startup crashes when CUDA extensions are broken +- **๐Ÿ“Š Improved VRAM tracking** - Separate allocated vs reserved memory tracking, Windows-only overflow detection (WDDM paging behavior) +- **โ™ป๏ธ Centralize backend detection** - Unified `is_mps_available()`, `is_cuda_available()`, `get_gpu_backend()` helpers across codebase +- **๐Ÿ”„ Revert 2.5.14 VRAM limit enforcement** - Removed `set_per_process_memory_fraction` call; Overflow detection and warnings remain. + **2025.12.09 - Version 2.5.18** - **๐Ÿš€ CLI: Streaming mode for long videos** - New `--chunk_size` flag processes videos in memory-bounded chunks, enabling arbitrarily long videos without RAM limits. Works with model caching (`--cache_dit`/`--cache_vae`) for chunk-to-chunk reuse *(inspired by [disk02](https://github.com/disk02) PR contribution)* diff --git a/pyproject.toml b/pyproject.toml index dacd4ea..365f488 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "seedvr2_videoupscaler" description = "SeedVR2 official ComfyUI integration: ByteDance-Seed's one-step diffusion-based video/image upscaling with memory-efficient inference" -version = "2.5.18" +version = "2.5.19" authors = [ {name = "numz"}, {name = "adrientoupet"} diff --git a/src/utils/constants.py b/src/utils/constants.py index 057b61d..521bb8d 100644 --- a/src/utils/constants.py +++ b/src/utils/constants.py @@ -4,7 +4,7 @@ Only includes constants actually used in the codebase """ # Version information -__version__ = "2.5.18" +__version__ = "2.5.19" import os import warnings