diff --git a/README.md b/README.md index d177f37..5096b8e 100644 --- a/README.md +++ b/README.md @@ -418,12 +418,11 @@ Configure the DiT (Diffusion Transformer) model for video upscaling. - `sdpa`: PyTorch scaled_dot_product_attention (default, stable, always available) - `flash_attn`: Flash Attention 2 (faster on supported hardware, requires flash-attn package) -- **allow_vram_overflow**: Allow VRAM to overflow to system RAM - - `False` (default): Strict VRAM limit - prevents silent swap but OOMs if exceeded - - `True`: Allow overflow - prevents OOM but may cause severe slowdown when physical VRAM exceeded - - Last resort when other memory optimizations are insufficient - - Requires ComfyUI restart to change setting - - No effect on Apple Silicon (unified memory architecture) +- **allow_vram_overflow**: Windows only - allow VRAM to overflow to system RAM + - `False` (default): Strict VRAM limit - faster when within limits + - `True`: Allow overflow - prevents OOM but causes severe slowdown + - Last resort when other optimizations are insufficient + - Requires ComfyUI restart to change - **torch_compile_args**: Connect to SeedVR2 Torch Compile Settings node for 20-40% speedup diff --git a/inference_cli.py b/inference_cli.py index 342231e..338f467 100644 --- a/inference_cli.py +++ b/inference_cli.py @@ -1336,8 +1336,8 @@ Examples: blockswap_group.add_argument("--swap_io_components", action="store_true", help="Offload DiT I/O layers for extra VRAM savings. Requires --dit_offload_device") blockswap_group.add_argument("--allow_vram_overflow", action="store_true", - help="Allow VRAM overflow to system RAM. Prevents OOM but may cause severe slowdown. " - "Last resort when other memory optimizations are insufficient. No effect on Apple Silicon (unified memory).") + help="Windows only: Allow VRAM overflow to system RAM. Prevents OOM but causes severe slowdown. " + "Last resort when other optimizations are insufficient.") # VAE Tiling vae_group = parser.add_argument_group('VAE tiling (for high resolution upscale)') diff --git a/src/interfaces/dit_model_loader.py b/src/interfaces/dit_model_loader.py index 76963d3..e26c970 100644 --- a/src/interfaces/dit_model_loader.py +++ b/src/interfaces/dit_model_loader.py @@ -116,12 +116,12 @@ class SeedVR2LoadDiTModel(io.ComfyNode): default=False, optional=True, tooltip=( - "Allow VRAM to overflow to system RAM when physical VRAM is exceeded.\n" - "• False (default): Strict VRAM limit - OOM if exceeded (faster when within limits)\n" - "• True: Allow overflow to RAM - prevents OOM but may cause severe slowdown\n" + "Windows only: Allow VRAM to overflow to system RAM.\n" + "• False (default): Strict VRAM limit - faster when within limits\n" + "• True: Allow overflow - prevents OOM but may cause severe slowdown\n" "\n" - "Last resort when other memory optimizations are insufficient.\n" - "Requires ComfyUI restart to change. No effect on Apple Silicon (unified memory)." + "Last resort when other optimizations are insufficient.\n" + "Requires ComfyUI restart to change." ) ), io.Custom("TORCH_COMPILE_ARGS").Input("torch_compile_args", diff --git a/src/optimization/memory_manager.py b/src/optimization/memory_manager.py index b8150a6..a61d5ef 100644 --- a/src/optimization/memory_manager.py +++ b/src/optimization/memory_manager.py @@ -12,7 +12,7 @@ import time import psutil import platform from typing import Tuple, Dict, Any, Optional, List, Union - + def _device_str(device: Union[torch.device, str]) -> str: """Normalized uppercase device string for comparison and logging. MPS variants → 'MPS'.""" @@ -45,27 +45,6 @@ def get_gpu_backend() -> str: return 'cpu' -def get_memory_architecture() -> str: - """Get memory architecture type for swap/overflow detection. - - This combines GPU backend with OS platform to determine how - GPU memory overflow is handled: - - Returns: - 'unified': macOS unified memory (MPS) - GPU/CPU share memory pool - 'discrete_paged': Windows WDDM - GPU memory can page to system RAM - 'discrete_strict': Linux - No automatic GPU paging, OOM on overflow - 'cpu_only': No GPU backend available - """ - if is_mps_available(): - return 'unified' - if is_cuda_available(): - if platform.system() == 'Windows': - return 'discrete_paged' - return 'discrete_strict' - return 'cpu_only' - - def get_device_list(include_none: bool = False, include_cpu: bool = False) -> List[str]: """ Get list of available compute devices for SeedVR2 @@ -215,7 +194,7 @@ def was_vram_limit_change_attempted() -> bool: return _vram_limit_change_attempted -def get_vram_usage(device: Optional[torch.device] = None, debug: Optional['Debug'] = None) -> Tuple[float, float, float]: +def get_vram_usage(device: Optional[torch.device] = None, debug: Optional['Debug'] = None) -> Tuple[float, float, float, float]: """ Get current VRAM usage metrics for monitoring. Used for tracking memory consumption during processing. @@ -225,8 +204,8 @@ def get_vram_usage(device: Optional[torch.device] = None, debug: Optional['Debug debug: Optional debug instance for logging Returns: - tuple: (allocated_gb, reserved_gb, max_reserved_gb) - Returns (0, 0, 0) if no GPU available + tuple: (allocated_gb, reserved_gb, peak_allocated_gb, peak_reserved_gb) + Returns (0, 0, 0, 0) if no GPU available """ try: if is_cuda_available(): @@ -236,18 +215,19 @@ def get_vram_usage(device: Optional[torch.device] = None, debug: Optional['Debug device = torch.device(device) allocated = torch.cuda.memory_allocated(device) / (1024**3) reserved = torch.cuda.memory_reserved(device) / (1024**3) - max_reserved = torch.cuda.max_memory_reserved(device) / (1024**3) - return allocated, reserved, max_reserved + peak_allocated = torch.cuda.max_memory_allocated(device) / (1024**3) + peak_reserved = torch.cuda.max_memory_reserved(device) / (1024**3) + return allocated, reserved, peak_allocated, peak_reserved elif is_mps_available(): # MPS doesn't support per-device queries - uses global memory tracking allocated = torch.mps.current_allocated_memory() / (1024**3) reserved = torch.mps.driver_allocated_memory() / (1024**3) - max_allocated = allocated # MPS doesn't track peak separately - return allocated, reserved, max_allocated + # MPS doesn't track peak separately + return allocated, reserved, allocated, reserved except Exception as e: if debug: debug.log(f"Failed to get VRAM usage: {e}", level="WARNING", category="memory", force=True) - return 0.0, 0.0, 0.0 + return 0.0, 0.0, 0.0, 0.0 def get_ram_usage(debug: Optional['Debug'] = None) -> Tuple[float, float, float, float]: diff --git a/src/utils/debug.py b/src/utils/debug.py index 10f3c2c..30eb23a 100644 --- a/src/utils/debug.py +++ b/src/utils/debug.py @@ -10,6 +10,7 @@ import torch import gc from typing import Optional, List, Dict, Any, Union from datetime import datetime +import platform from ..optimization.memory_manager import ( get_vram_usage, get_basic_vram_info, @@ -18,36 +19,26 @@ from ..optimization.memory_manager import ( is_vram_overflow_allowed, was_vram_limit_change_attempted, is_mps_available, - is_cuda_available, - get_memory_architecture + is_cuda_available ) from ..utils.constants import __version__ -def _format_peak_with_swap(peak_gb: float, total_vram_gb: float, arch: str = None) -> str: - """Format peak memory with architecture-aware overflow reporting. +def _format_peak_with_swap(peak_gb: float, total_vram_gb: float) -> str: + """Format peak memory, showing overflow breakdown on Windows. Args: peak_gb: Peak reserved memory from PyTorch total_vram_gb: Physical GPU VRAM capacity - arch: Memory architecture from get_memory_architecture(), or None to auto-detect """ if total_vram_gb <= 0: return f"{peak_gb:.2f}GB" overflow_gb = peak_gb - total_vram_gb - if overflow_gb <= 0: + if overflow_gb <= 0 or platform.system() != 'Windows': return f"{peak_gb:.2f}GB" - if arch is None: - arch = get_memory_architecture() - - if arch == 'discrete_paged': - return f"{peak_gb:.2f}GB ({total_vram_gb:.0f}GB GPU + {overflow_gb:.2f}GB system RAM)" - elif arch == 'discrete_strict': - return f"{peak_gb:.2f}GB (exceeded {total_vram_gb:.0f}GB by {overflow_gb:.2f}GB)" - # unified or cpu_only - no swap concept - return f"{peak_gb:.2f}GB" + return f"{peak_gb:.2f}GB ({total_vram_gb:.0f}GB GPU + {overflow_gb:.2f}GB system RAM)" class Debug: @@ -108,7 +99,8 @@ class Debug: self.vram_history: List[float] = [] self.active_timer_stack: List[str] = [] self.timer_namespace: str = "" - self.phase_vram_peaks: Dict[str, float] = {} + self.phase_vram_peaks_alloc: Dict[str, float] = {} + self.phase_vram_peaks_rsv: Dict[str, float] = {} self.phase_ram_peaks: Dict[str, float] = {} @torch._dynamo.disable # Skip tracing to avoid datetime.now() warnings @@ -252,30 +244,19 @@ class Debug: self.log(f"{cuda_line} | ComfyUI: {comfy_str}" if comfy_str else cuda_line, category="info") def _print_vram_overflow_status(self) -> bool: - """Print VRAM overflow status - warnings always shown, info only in debug mode. - - Returns: - True if a forced warning was printed, False otherwise. - """ - is_mps = is_mps_available() - force = False + """Print VRAM overflow status (Windows only). Returns True if warning was printed.""" + if platform.system() != 'Windows': + return False if was_vram_limit_change_attempted(): self.log("allow_vram_overflow setting changed - restart ComfyUI to apply", level="WARNING", category="memory", force=True) - force = True + return True elif is_vram_overflow_allowed(): - if is_mps: - self.log("allow_vram_overflow: enabled (no effect on Apple Silicon unified memory)", category="info") - else: - self.log("allow_vram_overflow: enabled - may cause severe slowdown if physical VRAM exceeded", level="WARNING", category="memory", force=True) - force = True + self.log("allow_vram_overflow: enabled - may cause severe slowdown if physical VRAM exceeded", level="WARNING", category="memory", force=True) + return True else: - if is_mps: - self.log("allow_vram_overflow: disabled (no effect on Apple Silicon unified memory)", category="info") - else: - self.log("allow_vram_overflow: disabled (recommended for best performance)", category="success") - - return force + self.log("allow_vram_overflow: disabled (recommended)", category="success") + return False def print_footer(self) -> None: """Print the footer with links - always displayed""" @@ -446,19 +427,13 @@ class Debug: if show_diff and self.memory_checkpoints: self._log_memory_diff(current_metrics=memory_info, force=force) - # Architecture-aware overflow warnings - arch = memory_info.get('arch', 'cpu_only') + # Overflow warning (Windows only - WDDM can page to system RAM) overflow = memory_info.get('vram_overflow', 0.0) - if overflow > 0 and not is_vram_overflow_allowed(): - if arch == 'discrete_paged': - self.log(f"VRAM overflow: {overflow:.2f}GB paged to system RAM - severe slowdown expected. " - "Consider optimizing (e.g., reduce resolution, batch size, enable BlockSwap, VAE tiling...).", - level="WARNING", category="memory", force=True) - elif arch == 'discrete_strict': - self.log(f"VRAM exceeded physical limit by {overflow:.2f}GB - OOM risk. " - "Consider optimizing (e.g., reduce resolution, batch size, enable BlockSwap, VAE tiling...).", - level="WARNING", category="memory", force=True) + if overflow > 0 and platform.system() == 'Windows' and not is_vram_overflow_allowed(): + self.log(f"VRAM overflow: {overflow:.2f}GB paged to system RAM - severe slowdown expected. " + "Consider optimizing (e.g., reduce resolution, batch size, enable BlockSwap, VAE tiling...).", + level="WARNING", category="memory", force=True) # Log detailed analysis if requested if detailed_tensors and tensor_stats.get('details'): @@ -469,10 +444,15 @@ class Debug: # Update phase peaks if we're in an active phase if self.current_phase: - if memory_info['vram_peak_since_last'] > 0: - self.phase_vram_peaks[self.current_phase] = max( - self.phase_vram_peaks.get(self.current_phase, 0), - memory_info['vram_peak_since_last'] + if memory_info['vram_peak_alloc'] > 0: + self.phase_vram_peaks_alloc[self.current_phase] = max( + self.phase_vram_peaks_alloc.get(self.current_phase, 0), + memory_info['vram_peak_alloc'] + ) + if memory_info['vram_peak_rsv'] > 0: + self.phase_vram_peaks_rsv[self.current_phase] = max( + self.phase_vram_peaks_rsv.get(self.current_phase, 0), + memory_info['vram_peak_rsv'] ) if memory_info['ram_process'] > 0: self.phase_ram_peaks[self.current_phase] = max( @@ -484,17 +464,18 @@ class Debug: reset_vram_peak(device=None, debug=self) def _collect_memory_metrics(self) -> Dict[str, Any]: - """Collect current memory metrics with architecture-aware reporting.""" - arch = get_memory_architecture() + """Collect current memory metrics.""" + is_mps = is_mps_available() + has_gpu = is_mps or is_cuda_available() metrics = { 'vram_allocated': 0.0, 'vram_reserved': 0.0, 'vram_free': 0.0, 'vram_total': 0.0, - 'vram_peak_since_last': 0.0, + 'vram_peak_alloc': 0.0, + 'vram_peak_rsv': 0.0, 'vram_overflow': 0.0, - 'arch': arch, 'ram_process': 0.0, 'ram_available': 0.0, 'ram_total': 0.0, @@ -503,28 +484,26 @@ class Debug: 'summary_ram': "" } - if arch == 'cpu_only': - pass # No GPU metrics - else: - metrics['vram_allocated'], metrics['vram_reserved'], metrics['vram_peak_since_last'] = get_vram_usage(device=None, debug=self) + if has_gpu: + metrics['vram_allocated'], metrics['vram_reserved'], metrics['vram_peak_alloc'], metrics['vram_peak_rsv'] = get_vram_usage(device=None, debug=self) vram_info = get_basic_vram_info(device=None) if "error" not in vram_info and vram_info["total_gb"] > 0: metrics['vram_free'] = vram_info["free_gb"] metrics['vram_total'] = vram_info["total_gb"] + metrics['vram_overflow'] = max(0.0, metrics['vram_peak_rsv'] - metrics['vram_total']) - # Calculate overflow: reserved beyond physical VRAM - metrics['vram_overflow'] = max(0.0, metrics['vram_peak_since_last'] - metrics['vram_total']) - - backend = "Unified Memory" if arch == 'unified' else "VRAM" - peak_str = _format_peak_with_swap(metrics['vram_peak_since_last'], metrics['vram_total'], arch) + backend = "Unified Memory" if is_mps else "VRAM" + peak_alloc_str = _format_peak_with_swap(metrics['vram_peak_alloc'], metrics['vram_total']) metrics['summary_vram'] = ( f" [{backend}] {metrics['vram_allocated']:.2f}GB allocated / " f"{metrics['vram_reserved']:.2f}GB reserved / " - f"Peak: {peak_str} / " + f"Peak: {peak_alloc_str} / " f"{metrics['vram_free']:.2f}GB free / " f"{metrics['vram_total']:.2f}GB total" ) + + self.vram_history.append(metrics['vram_reserved']) # RAM metrics metrics['ram_process'], metrics['ram_available'], metrics['ram_total'], metrics['ram_others'] = get_ram_usage(debug=self) @@ -537,10 +516,6 @@ class Debug: f"{metrics['ram_total']:.2f}GB total" ) - # Track reserved (matches nvidia-smi) for pressure history - if arch != 'cpu_only': - self.vram_history.append(metrics['vram_reserved']) - return metrics def _collect_tensor_stats(self, detailed: bool = False) -> Dict[str, Any]: @@ -667,8 +642,8 @@ class Debug: self.log(f"Memory changes: {', '.join(diffs)}", category="memory", force=force, indent_level=1) def log_peak_memory_summary(self, force: bool = True) -> None: - """Display peak memory usage across all phases (VRAM and RAM combined)""" - if not self.phase_vram_peaks and not self.phase_ram_peaks: + """Display peak memory usage across all phases.""" + if not self.phase_vram_peaks_alloc and not self.phase_ram_peaks: return phase_names = { @@ -678,11 +653,11 @@ class Debug: 'phase4': 'Post-processing' } - arch = get_memory_architecture() + is_mps = is_mps_available() - # Get total VRAM for overflow detection + # Get total VRAM for overflow formatting (Windows only) total_vram_gb = 0.0 - if arch not in ('unified', 'cpu_only'): + if not is_mps: vram_info = get_basic_vram_info(device=None) if "error" not in vram_info: total_vram_gb = vram_info["total_gb"] @@ -691,25 +666,29 @@ class Debug: self.log("────────────────────────", category="none", force=force) self.log("Peak memory by phase:", category="memory", force=force) - all_phases = sorted(set(self.phase_vram_peaks.keys()) | set(self.phase_ram_peaks.keys())) + all_phases = sorted(set(self.phase_vram_peaks_alloc.keys()) | set(self.phase_ram_peaks.keys())) for phase_key in all_phases: phase_num = phase_key[-1] phase_name = phase_names.get(phase_key, phase_key) - vram = self.phase_vram_peaks.get(phase_key, 0) + alloc = self.phase_vram_peaks_alloc.get(phase_key, 0) + rsv = self.phase_vram_peaks_rsv.get(phase_key, 0) ram = self.phase_ram_peaks.get(phase_key, 0) - if arch == 'unified': - self.log(f"Phase {phase_num} ({phase_name}): {vram:.2f}GB", category="memory", indent_level=1, force=force) + if is_mps: + self.log(f"{phase_num}. {phase_name}: {alloc:.2f}GB", category="memory", indent_level=1, force=force) else: - self.log(f"Phase {phase_num} ({phase_name}): {_format_peak_with_swap(vram, total_vram_gb, arch)} | RAM {ram:.2f}GB", category="memory", indent_level=1, force=force) + rsv_str = _format_peak_with_swap(rsv, total_vram_gb) + self.log(f"{phase_num}. {phase_name}: VRAM {alloc:.2f}GB allocated, {rsv_str} reserved | RAM {ram:.2f}GB", category="memory", indent_level=1, force=force) - overall_vram = max(self.phase_vram_peaks.values()) if self.phase_vram_peaks else 0 + overall_alloc = max(self.phase_vram_peaks_alloc.values()) if self.phase_vram_peaks_alloc else 0 + overall_rsv = max(self.phase_vram_peaks_rsv.values()) if self.phase_vram_peaks_rsv else 0 overall_ram = max(self.phase_ram_peaks.values()) if self.phase_ram_peaks else 0 - if arch == 'unified': - self.log(f"Overall peak: {overall_vram:.2f}GB", category="memory", force=force) + if is_mps: + self.log(f"Overall peak: {overall_alloc:.2f}GB", category="memory", force=force) else: - self.log(f"Overall peak: {_format_peak_with_swap(overall_vram, total_vram_gb, arch)} | RAM {overall_ram:.2f}GB", category="memory", force=force) + overall_rsv_str = _format_peak_with_swap(overall_rsv, total_vram_gb) + self.log(f"Overall peak: VRAM {overall_alloc:.2f}GB allocated, {overall_rsv_str} reserved | RAM {overall_ram:.2f}GB", category="memory", force=force) @torch._dynamo.disable # Skip tracing to avoid time.time() warnings def _store_checkpoint(self, label: str, metrics: Dict[str, Any]) -> None: @@ -819,6 +798,7 @@ class Debug: self.timer_durations.clear() self.timer_messages.clear() self.active_timer_stack.clear() - self.phase_vram_peaks.clear() + self.phase_vram_peaks_alloc.clear() + self.phase_vram_peaks_rsv.clear() self.phase_ram_peaks.clear() self.current_phase = None \ No newline at end of file