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:
@@ -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
@@ -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
@@ -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"}
|
||||
|
||||
@@ -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())
|
||||
|
||||
|
||||
@@ -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]):
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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]):
|
||||
|
||||
@@ -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):
|
||||
"""
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
@@ -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
|
||||
Reference in New Issue
Block a user