diff --git a/README.md b/README.md index 5096b8e..e6d2d65 100644 --- a/README.md +++ b/README.md @@ -418,12 +418,6 @@ Configure the DiT (Diffusion Transformer) model for video upscaling. - `sdpa`: PyTorch scaled_dot_product_attention (default, stable, always available) - `flash_attn`: Flash Attention 2 (faster on supported hardware, requires flash-attn package) -- **allow_vram_overflow**: Windows only - allow VRAM to overflow to system RAM - - `False` (default): Strict VRAM limit - faster when within limits - - `True`: Allow overflow - prevents OOM but causes severe slowdown - - Last resort when other optimizations are insufficient - - Requires ComfyUI restart to change - - **torch_compile_args**: Connect to SeedVR2 Torch Compile Settings node for 20-40% speedup **BlockSwap Explained:** diff --git a/inference_cli.py b/inference_cli.py index 338f467..a794993 100644 --- a/inference_cli.py +++ b/inference_cli.py @@ -79,7 +79,6 @@ else: # Pre-parse arguments that must be handled before torch import _pre_parser = argparse.ArgumentParser(add_help=False) _pre_parser.add_argument("--cuda_device", type=str, default=None) - _pre_parser.add_argument("--allow_vram_overflow", action="store_true") _pre_args, _ = _pre_parser.parse_known_args() if _pre_args.cuda_device is not None: @@ -128,13 +127,9 @@ from src.core.generation_phases import ( postprocess_all_batches ) from src.utils.debug import Debug -from src.optimization.memory_manager import clear_memory, configure_vram_limit, get_gpu_backend, is_cuda_available +from src.optimization.memory_manager import clear_memory, get_gpu_backend, is_cuda_available debug = Debug(enabled=False) # Will be enabled via --debug CLI flag -# Configure VRAM limit (must be before any CUDA allocations) -if platform.system() != "Darwin": - configure_vram_limit(allow_overflow=_pre_args.allow_vram_overflow) - # ============================================================================= # Device Management Helpers # ============================================================================= @@ -1335,9 +1330,6 @@ Examples: "Requires --dit_offload_device. Default: 0 (disabled)") blockswap_group.add_argument("--swap_io_components", action="store_true", help="Offload DiT I/O layers for extra VRAM savings. Requires --dit_offload_device") - blockswap_group.add_argument("--allow_vram_overflow", action="store_true", - help="Windows only: Allow VRAM overflow to system RAM. Prevents OOM but causes severe slowdown. " - "Last resort when other optimizations are insufficient.") # VAE Tiling vae_group = parser.add_argument_group('VAE tiling (for high resolution upscale)') diff --git a/src/interfaces/dit_model_loader.py b/src/interfaces/dit_model_loader.py index e26c970..1064571 100644 --- a/src/interfaces/dit_model_loader.py +++ b/src/interfaces/dit_model_loader.py @@ -7,7 +7,7 @@ from comfy_api.latest import io from comfy_execution.utils import get_executing_context from typing import Dict, Any, Tuple from ..utils.model_registry import get_available_dit_models, DEFAULT_DIT -from ..optimization.memory_manager import get_device_list, configure_vram_limit +from ..optimization.memory_manager import get_device_list class SeedVR2LoadDiTModel(io.ComfyNode): @@ -112,18 +112,6 @@ class SeedVR2LoadDiTModel(io.ComfyNode): "Flash Attention provides speedup through optimized CUDA kernels on compatible GPUs." ) ), - io.Boolean.Input("allow_vram_overflow", - default=False, - optional=True, - tooltip=( - "Windows only: Allow VRAM to overflow to system RAM.\n" - "• False (default): Strict VRAM limit - faster when within limits\n" - "• True: Allow overflow - prevents OOM but may cause severe slowdown\n" - "\n" - "Last resort when other optimizations are insufficient.\n" - "Requires ComfyUI restart to change." - ) - ), io.Custom("TORCH_COMPILE_ARGS").Input("torch_compile_args", optional=True, tooltip=( @@ -143,7 +131,6 @@ class SeedVR2LoadDiTModel(io.ComfyNode): def execute(cls, model: str, device: str, offload_device: str = "none", cache_model: bool = False, blocks_to_swap: int = 0, swap_io_components: bool = False, attention_mode: str = "sdpa", - allow_vram_overflow: bool = False, torch_compile_args: Dict[str, Any] = None) -> io.NodeOutput: """ Create DiT model configuration for SeedVR2 main node @@ -156,7 +143,6 @@ class SeedVR2LoadDiTModel(io.ComfyNode): blocks_to_swap: Number of transformer blocks to swap (requires offload_device != device) swap_io_components: Whether to offload I/O components (requires offload_device != device) attention_mode: Attention computation backend ('sdpa' or 'flash_attn') - allow_vram_overflow: Allow VRAM overflow to system RAM (prevents OOM but slower) torch_compile_args: Optional torch.compile configuration from settings node Returns: @@ -182,9 +168,6 @@ class SeedVR2LoadDiTModel(io.ComfyNode): "(e.g., 'cpu' or another device). Set cache_model=False if you don't want to cache the model." ) - # Configure VRAM limit enforcement (once per session, first call wins) - configure_vram_limit(allow_overflow=allow_vram_overflow) - config = { "model": model, "device": device, diff --git a/src/optimization/memory_manager.py b/src/optimization/memory_manager.py index a61d5ef..fb2fc75 100644 --- a/src/optimization/memory_manager.py +++ b/src/optimization/memory_manager.py @@ -138,62 +138,6 @@ else: print(f"⚠️ Memory check failed: {vram_info['error']} - No available backend!") -# VRAM overflow configuration state -_vram_overflow_allowed: bool = True -_vram_limit_configured: bool = False -_vram_limit_change_attempted: bool = False - - -def configure_vram_limit(allow_overflow: bool = False) -> bool: - """ - Configure VRAM limit enforcement. Call early before heavy CUDA usage. - - Args: - allow_overflow: If True, allow VRAM overflow to system RAM (prevents OOM but may be slow). - If False (default), enforce strict physical VRAM limit. - - Returns: - True if configuration applied successfully, False otherwise - - Note: - Can only be configured once per session. Restart required to change. - """ - global _vram_overflow_allowed, _vram_limit_configured, _vram_limit_change_attempted - - # Already configured this session - track if user tried to change - if _vram_limit_configured: - if _vram_overflow_allowed != allow_overflow: - _vram_limit_change_attempted = True - return _vram_overflow_allowed == allow_overflow - - _vram_limit_configured = True - _vram_overflow_allowed = allow_overflow - - if allow_overflow: - return True - - if not is_cuda_available(): - return True - - try: - for i in range(torch.cuda.device_count()): - torch.cuda.set_per_process_memory_fraction(1.0, i) - return True - except RuntimeError: - _vram_overflow_allowed = True - return False - - -def is_vram_overflow_allowed() -> bool: - """Check if VRAM overflow to system RAM is allowed.""" - return _vram_overflow_allowed - - -def was_vram_limit_change_attempted() -> bool: - """Check if user tried to change VRAM limit setting after initial configuration.""" - return _vram_limit_change_attempted - - def get_vram_usage(device: Optional[torch.device] = None, debug: Optional['Debug'] = None) -> Tuple[float, float, float, float]: """ Get current VRAM usage metrics for monitoring. diff --git a/src/utils/debug.py b/src/utils/debug.py index 30eb23a..3619084 100644 --- a/src/utils/debug.py +++ b/src/utils/debug.py @@ -15,30 +15,28 @@ from ..optimization.memory_manager import ( get_vram_usage, get_basic_vram_info, get_ram_usage, - reset_vram_peak, - is_vram_overflow_allowed, - was_vram_limit_change_attempted, + reset_vram_peak, is_mps_available, is_cuda_available ) from ..utils.constants import __version__ -def _format_peak_with_swap(peak_gb: float, total_vram_gb: float) -> str: - """Format peak memory, showing overflow breakdown on Windows. +def _format_peak_with_overflow(peak_gb: float, total_vram_gb: float) -> str: + """Format peak reserved memory, showing overflow breakdown on Windows. Args: peak_gb: Peak reserved memory from PyTorch total_vram_gb: Physical GPU VRAM capacity """ if total_vram_gb <= 0: - return f"{peak_gb:.2f}GB" + return f"{peak_gb:.2f}GB reserved" overflow_gb = peak_gb - total_vram_gb if overflow_gb <= 0 or platform.system() != 'Windows': - return f"{peak_gb:.2f}GB" + return f"{peak_gb:.2f}GB reserved" - return f"{peak_gb:.2f}GB ({total_vram_gb:.0f}GB GPU + {overflow_gb:.2f}GB system RAM)" + return f"{peak_gb:.2f}GB reserved ({total_vram_gb:.0f}GB GPU + {overflow_gb:.2f}GB overflow)" class Debug: @@ -176,11 +174,6 @@ class Debug: # Environment info - only in debug mode if self.enabled: self._print_environment_info(cli) - - # VRAM overflow status - warnings always shown - vram_warning_shown = self._print_vram_overflow_status() - - self.log("", category="none", force=vram_warning_shown) def _print_environment_info(self, cli: bool = False) -> None: """Print concise environment info for bug reports - zero cost when debug disabled""" @@ -242,21 +235,7 @@ class Debug: self.log(f"Python: {py_ver} | PyTorch: {torch_ver} | Flash Attn: {flash_str} | Triton: {triton_str}", category="info") cuda_line = f"CUDA: {cuda_ver} | cuDNN: {cudnn_ver}" self.log(f"{cuda_line} | ComfyUI: {comfy_str}" if comfy_str else cuda_line, category="info") - - def _print_vram_overflow_status(self) -> bool: - """Print VRAM overflow status (Windows only). Returns True if warning was printed.""" - if platform.system() != 'Windows': - return False - - if was_vram_limit_change_attempted(): - self.log("allow_vram_overflow setting changed - restart ComfyUI to apply", level="WARNING", category="memory", force=True) - return True - elif is_vram_overflow_allowed(): - self.log("allow_vram_overflow: enabled - may cause severe slowdown if physical VRAM exceeded", level="WARNING", category="memory", force=True) - return True - else: - self.log("allow_vram_overflow: disabled (recommended)", category="success") - return False + self.log("", category="none") def print_footer(self) -> None: """Print the footer with links - always displayed""" @@ -430,7 +409,7 @@ class Debug: # Overflow warning (Windows only - WDDM can page to system RAM) overflow = memory_info.get('vram_overflow', 0.0) - if overflow > 0 and platform.system() == 'Windows' and not is_vram_overflow_allowed(): + if overflow > 0 and platform.system() == 'Windows': self.log(f"VRAM overflow: {overflow:.2f}GB paged to system RAM - severe slowdown expected. " "Consider optimizing (e.g., reduce resolution, batch size, enable BlockSwap, VAE tiling...).", level="WARNING", category="memory", force=True) @@ -494,11 +473,10 @@ class Debug: metrics['vram_overflow'] = max(0.0, metrics['vram_peak_rsv'] - metrics['vram_total']) backend = "Unified Memory" if is_mps else "VRAM" - peak_alloc_str = _format_peak_with_swap(metrics['vram_peak_alloc'], metrics['vram_total']) metrics['summary_vram'] = ( f" [{backend}] {metrics['vram_allocated']:.2f}GB allocated / " f"{metrics['vram_reserved']:.2f}GB reserved / " - f"Peak: {peak_alloc_str} / " + f"Peak: {metrics['vram_peak_alloc']:.2f}GB / " f"{metrics['vram_free']:.2f}GB free / " f"{metrics['vram_total']:.2f}GB total" ) @@ -677,8 +655,8 @@ class Debug: if is_mps: self.log(f"{phase_num}. {phase_name}: {alloc:.2f}GB", category="memory", indent_level=1, force=force) else: - rsv_str = _format_peak_with_swap(rsv, total_vram_gb) - self.log(f"{phase_num}. {phase_name}: VRAM {alloc:.2f}GB allocated, {rsv_str} reserved | RAM {ram:.2f}GB", category="memory", indent_level=1, force=force) + rsv_str = _format_peak_with_overflow(rsv, total_vram_gb) + self.log(f"{phase_num}. {phase_name}: VRAM {alloc:.2f}GB allocated, {rsv_str} | RAM {ram:.2f}GB", category="memory", indent_level=1, force=force) overall_alloc = max(self.phase_vram_peaks_alloc.values()) if self.phase_vram_peaks_alloc else 0 overall_rsv = max(self.phase_vram_peaks_rsv.values()) if self.phase_vram_peaks_rsv else 0 @@ -687,8 +665,8 @@ class Debug: if is_mps: self.log(f"Overall peak: {overall_alloc:.2f}GB", category="memory", force=force) else: - overall_rsv_str = _format_peak_with_swap(overall_rsv, total_vram_gb) - self.log(f"Overall peak: VRAM {overall_alloc:.2f}GB allocated, {overall_rsv_str} reserved | RAM {overall_ram:.2f}GB", category="memory", force=force) + overall_rsv_str = _format_peak_with_overflow(overall_rsv, total_vram_gb) + self.log(f"Overall peak: VRAM {overall_alloc:.2f}GB allocated, {overall_rsv_str} | RAM {overall_ram:.2f}GB", category="memory", force=force) @torch._dynamo.disable # Skip tracing to avoid time.time() warnings def _store_checkpoint(self, label: str, metrics: Dict[str, Any]) -> None: