diff --git a/README.md b/README.md index de87611..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)* @@ -872,6 +881,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..a794993 100644 --- a/inference_cli.py +++ b/inference_cli.py @@ -76,7 +76,7 @@ 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_args, _ = _pre_parser.parse_known_args() @@ -127,24 +127,13 @@ 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, get_gpu_backend, is_cuda_available debug = Debug(enabled=False) # Will be enabled via --debug CLI flag - # ============================================================================= # 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. @@ -160,7 +149,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": @@ -777,7 +766,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) @@ -1466,7 +1455,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.", @@ -1474,7 +1463,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.", @@ -1487,7 +1476,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/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/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/compatibility.py b/src/optimization/compatibility.py index d154dd0..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 @@ -282,11 +333,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 +455,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): """ diff --git a/src/optimization/memory_manager.py b/src/optimization/memory_manager.py index 592e37e..fb2fc75 100644 --- a/src/optimization/memory_manager.py +++ b/src/optimization/memory_manager.py @@ -10,8 +10,9 @@ import gc import sys 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'.""" @@ -19,6 +20,31 @@ 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_device_list(include_none: bool = False, include_cpu: bool = False) -> List[str]: """ Get list of available compute devices for SeedVR2 @@ -37,14 +63,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 +92,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 +106,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,29 +132,13 @@ 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!") -def _enforce_vram_limit() -> None: - """ - 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. - """ - if not torch.cuda.is_available(): - return - try: - for i in range(torch.cuda.device_count()): - torch.cuda.set_per_process_memory_fraction(1.0, i) - except Exception: - pass - -_enforce_vram_limit() - - -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. @@ -138,29 +148,30 @@ 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 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) 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 - elif hasattr(torch.backends, 'mps') and torch.backends.mps.is_available(): + 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]: @@ -251,17 +262,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 @@ -284,10 +295,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: @@ -324,7 +335,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 @@ -401,7 +412,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/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 diff --git a/src/utils/debug.py b/src/utils/debug.py index e819120..3619084 100644 --- a/src/utils/debug.py +++ b/src/utils/debug.py @@ -10,16 +10,33 @@ 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 +import platform +from ..optimization.memory_manager import ( + get_vram_usage, + get_basic_vram_info, + get_ram_usage, + 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 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)" - return f"{peak_gb:.2f}GB" +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 reserved" + + overflow_gb = peak_gb - total_vram_gb + if overflow_gb <= 0 or platform.system() != 'Windows': + return f"{peak_gb:.2f}GB reserved" + + return f"{peak_gb:.2f}GB reserved ({total_vram_gb:.0f}GB GPU + {overflow_gb:.2f}GB overflow)" class Debug: @@ -80,7 +97,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 @@ -125,26 +143,33 @@ 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) + 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: @@ -174,7 +199,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)" @@ -182,7 +207,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: @@ -381,9 +406,12 @@ 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) - 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...).", + # 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': + 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 @@ -395,10 +423,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( @@ -410,13 +443,18 @@ 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.""" + 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, 'ram_process': 0.0, 'ram_available': 0.0, 'ram_total': 0.0, @@ -425,46 +463,36 @@ 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 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: + 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']) - 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'] = "" + backend = "Unified Memory" if is_mps else "VRAM" + metrics['summary_vram'] = ( + f" [{backend}] {metrics['vram_allocated']:.2f}GB allocated / " + f"{metrics['vram_reserved']:.2f}GB reserved / " + f"Peak: {metrics['vram_peak_alloc']:.2f}GB / " + f"{metrics['vram_free']:.2f}GB free / " + f"{metrics['vram_total']:.2f}GB total" + ) + + self.vram_history.append(metrics['vram_reserved']) - # 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'] = "" - - # 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']) + 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" + ) return metrics @@ -592,8 +620,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 = { @@ -603,9 +631,9 @@ class Debug: 'phase4': 'Post-processing' } - is_mps = hasattr(torch.backends, 'mps') and torch.backends.mps.is_available() and not torch.cuda.is_available() + is_mps = is_mps_available() - # Get total VRAM for swap detection (reuse existing function) + # Get total VRAM for overflow formatting (Windows only) total_vram_gb = 0.0 if not is_mps: vram_info = get_basic_vram_info(device=None) @@ -616,25 +644,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 is_mps: - self.log(f" Phase {phase_num} ({phase_name}): {vram:.2f}GB", category="memory", force=force) + 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)} | RAM {ram:.2f}GB", category="memory", 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 + overall_ram = max(self.phase_ram_peaks.values()) if self.phase_ram_peaks else 0 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) + self.log(f"Overall peak: {overall_alloc:.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) + 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: @@ -744,6 +776,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