Merge pull request #390 from AInVFX/main

v2.5.19: new logo, remove dead flash-attn wrapper, graceful DLL fallback, improved VRAM tracking
This commit is contained in:
Adrien Toupet
2025-12-10 01:55:47 -05:00
committed by GitHub
11 changed files with 251 additions and 362 deletions
+10
View File
@@ -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)
+7 -18
View File
@@ -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")
+1 -1
View File
@@ -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"}
+2 -1
View File
@@ -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())
+2 -1
View File
@@ -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]):
+2 -1
View File
@@ -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,
+2 -1
View File
@@ -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]):
+58 -216
View File
@@ -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):
"""
+51 -40
View File
@@ -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):
+1 -1
View File
@@ -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
+115 -82
View File
@@ -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