From fb2b6c7d4c875fc2c14a344533103cc641fb885e Mon Sep 17 00:00:00 2001 From: Adrien Toupet Date: Fri, 10 Oct 2025 17:48:12 -0400 Subject: [PATCH] Fix BlockSwap not applying on cached models and rename offload_io_components to swap_io_components - Fixed critical bug where BlockSwap configuration changes were not applied to cached models - Apply BlockSwap immediately in _handle_blockswap_config instead of deferring to materialization phase - Renamed offload_io_components to swap_io_components across entire codebase for consistency - Removed unused _pending_blockswap_config attribute --- inference_cli.py | 6 +++--- src/core/model_manager.py | 17 ++++++++++------- src/interfaces/comfyui_node.py | 14 +++++++------- src/optimization/blockswap.py | 24 ++++++++++++------------ src/optimization/memory_manager.py | 6 +++--- 5 files changed, 35 insertions(+), 32 deletions(-) diff --git a/inference_cli.py b/inference_cli.py index 8bc4170..47ee48b 100644 --- a/inference_cli.py +++ b/inference_cli.py @@ -455,7 +455,7 @@ def _gpu_processing(frames_tensor: torch.Tensor, device_list: List[str], "block_swap_config": { 'blocks_to_swap': args.blocks_to_swap, 'use_none_blocking': args.use_none_blocking, - 'offload_io_components': args.offload_io_components, + 'swap_io_components': args.swap_io_components, 'cache_model': False, }, "vae_encode_tiling_enabled": args.vae_encode_tiling_enabled, @@ -595,8 +595,8 @@ def parse_arguments() -> argparse.Namespace: help="Temporal overlap for processing (default: 0, no temporal overlap)") parser.add_argument("--prepend_frames", type=int, default=0, help="Number of frames to prepend to the video (default: 0). This can help with artifacts at the start of the video and are removed after processing") - parser.add_argument("--offload_io_components", action="store_true", - help="Offload IO components to CPU for VRAM optimization") + parser.add_argument("--swap_io_components", action="store_true", + help="Swap IO components to CPU for VRAM optimization") parser.add_argument("--vae_encode_tiling_enabled", action="store_true", help="Enable VAE encode tiling for improved VRAM usage") parser.add_argument("--vae_encode_tile_size", action=OneOrTwoValues, nargs='+', default=(512, 512), diff --git a/src/core/model_manager.py b/src/core/model_manager.py index 6774c49..aab5d78 100644 --- a/src/core/model_manager.py +++ b/src/core/model_manager.py @@ -132,31 +132,34 @@ def _handle_blockswap_config(cached_runner: VideoDiffusionInfer, if block_swap_config and block_swap_config.get("blocks_to_swap", 0) > 0: desired_config = ( block_swap_config.get("blocks_to_swap"), - block_swap_config.get("offload_io_components", False) + block_swap_config.get("swap_io_components", False) ) current_config = None if hasattr(cached_runner, "_block_swap_config"): current_config = ( cached_runner._block_swap_config.get("blocks_swapped"), - cached_runner._block_swap_config.get("offload_io_components", False) + cached_runner._block_swap_config.get("swap_io_components", False) ) # Apply changes if needed if desired_config != current_config: - fmt_curr = "disabled" if current_config is None else f"blocks={current_config[0]}, offload={current_config[1]}" - fmt_new = "disabled" if desired_config is None else f"blocks={desired_config[0]}, offload={desired_config[1]}" + fmt_curr = "disabled" if current_config is None else f"blocks={current_config[0]}, I/O={current_config[1]}" + fmt_new = "disabled" if desired_config is None else f"blocks={desired_config[0]}, I/O={desired_config[1]}" debug.log(f"BlockSwap config changed: {fmt_curr} → {fmt_new}", category="blockswap", force=True) cleanup_blockswap(cached_runner, keep_state_for_cache=False) - cached_runner._pending_blockswap_config = block_swap_config if desired_config else None if desired_config: - debug.log("BlockSwap application deferred to DiT phase", category="blockswap") + # Apply BlockSwap immediately since model is already materialized + apply_block_swap_to_dit(cached_runner, block_swap_config, debug) + else: + # Clear any pending config + cached_runner._dit_block_swap_config = None elif desired_config and hasattr(cached_runner, "_blockswap_active") and not cached_runner._blockswap_active: cached_runner._blockswap_active = True - debug.log(f"BlockSwap reactivated: blocks={desired_config[0]}, offload={desired_config[1]}", + debug.log(f"BlockSwap reactivated: blocks={desired_config[0]}, I/O={desired_config[1]}", category="blockswap", force=True) diff --git a/src/interfaces/comfyui_node.py b/src/interfaces/comfyui_node.py index 09e4982..6ae8e14 100644 --- a/src/interfaces/comfyui_node.py +++ b/src/interfaces/comfyui_node.py @@ -153,7 +153,7 @@ class SeedVR2: - model: Model filename - device: Target device - blocks_to_swap: BlockSwap configuration - - offload_io_components: I/O component offload flag + - swap_io_components: I/O component swap flag - cache_in_ram: Model caching flag - torch_compile_args: Optional compilation settings vae: VAE model configuration from SeedVR2LoadVAEModel node containing: @@ -193,7 +193,7 @@ class SeedVR2: dit_model = dit["model"] dit_device = dit["device"] blocks_to_swap = dit.get("blocks_to_swap", 0) - offload_io_components = dit.get("offload_io_components", False) + swap_io_components = dit.get("swap_io_components", False) cache_model_dit = dit.get("cache_in_ram", False) dit_torch_compile_args = dit.get("torch_compile_args", None) @@ -214,7 +214,7 @@ class SeedVR2: if blocks_to_swap > 0: block_swap_config = { "blocks_to_swap": blocks_to_swap, - "offload_io_components": offload_io_components, + "swap_io_components": swap_io_components, } # Fixed parameters (could be exposed in future) @@ -624,7 +624,7 @@ class SeedVR2LoadDiTModel: "step": 1, "tooltip": "BlockSwap: Number of transformer blocks to offload (0=disabled, 16=balanced, 32=max savings for 3b model, 36=max savings for 7b model)" }), - "offload_io_components": ("BOOLEAN", { + "swap_io_components": ("BOOLEAN", { "default": False, "tooltip": "Offload embeddings/IO layers to CPU for additional VRAM savings" }), @@ -644,7 +644,7 @@ class SeedVR2LoadDiTModel: DESCRIPTION = "Configure DiT model loading and memory optimization settings" def create_config(self, model: str, device: str, blocks_to_swap: int = 0, - offload_io_components: bool = False, cache_in_ram: bool = False, + swap_io_components: bool = False, cache_in_ram: bool = False, torch_compile_args: Optional[Dict[str, Any]] = None) -> Tuple[Dict[str, Any]]: """ Create DiT model configuration dictionary @@ -653,7 +653,7 @@ class SeedVR2LoadDiTModel: model: DiT model filename (e.g., "seedvr2_ema_3b_fp16.safetensors") device: Target device for DiT model (e.g., "cuda:0", "cpu") blocks_to_swap: Number of transformer blocks to offload for BlockSwap (0=disabled) - offload_io_components: Whether to offload input/output layers to CPU + swap_io_components: Whether to offload input/output layers to CPU cache_in_ram: Whether to keep model in RAM between runs torch_compile_args: Optional torch.compile configuration from settings node @@ -664,7 +664,7 @@ class SeedVR2LoadDiTModel: "model": model, "device": device, "blocks_to_swap": blocks_to_swap, - "offload_io_components": offload_io_components, + "swap_io_components": swap_io_components, "cache_in_ram": cache_in_ram, "torch_compile_args": torch_compile_args, } diff --git a/src/optimization/blockswap.py b/src/optimization/blockswap.py index 035ebf2..377d4cb 100644 --- a/src/optimization/blockswap.py +++ b/src/optimization/blockswap.py @@ -72,7 +72,7 @@ def apply_block_swap_to_dit(runner, block_swap_config: Dict[str, Any], debug) -> runner: VideoDiffusionInfer instance containing the model block_swap_config: Configuration dictionary with keys: - blocks_to_swap: Number of blocks to swap (from the start) - - offload_io_components: Whether to offload I/O components + - swap_io_components: Whether to offload I/O components - enable_debug: Whether to enable debug logging """ if not block_swap_config: @@ -103,7 +103,7 @@ def apply_block_swap_to_dit(runner, block_swap_config: Dict[str, Any], debug) -> blocks_to_swap = block_swap_config.get("blocks_to_swap", 0) if blocks_to_swap > 0: configs.append(f"{blocks_to_swap} blocks") - if block_swap_config.get("offload_io_components", False): + if block_swap_config.get("swap_io_components", False): configs.append("I/O components") debug.log(f"BlockSwap configured: {', '.join(configs)}", category="blockswap", force=True) debug.log("BlockSwap will swap blocks to GPU during inference", category="info", force=True) @@ -125,16 +125,16 @@ def apply_block_swap_to_dit(runner, block_swap_config: Dict[str, Any], debug) -> debug.log(f"Configuring: {blocks_to_swap}/{total_blocks} blocks for swapping", category="blockswap") # Configure I/O components - offload_io_components = block_swap_config.get("offload_io_components", False) + swap_io_components = block_swap_config.get("swap_io_components", False) io_components_offloaded = _configure_io_components(model, device, offload_device, - offload_io_components, debug) + swap_io_components, debug) # Configure block placement and memory tracking memory_stats = _configure_blocks(model, device, offload_device, debug) memory_stats['io_components'] = io_components_offloaded # Log memory summary - _log_memory_summary(memory_stats, offload_device, device, offload_io_components, + _log_memory_summary(memory_stats, offload_device, device, swap_io_components, debug) # Wrap block forward methods for dynamic swapping @@ -151,7 +151,7 @@ def apply_block_swap_to_dit(runner, block_swap_config: Dict[str, Any], debug) -> # Store configuration for debugging and cleanup runner._block_swap_config = { "blocks_swapped": blocks_to_swap, - "offload_io_components": offload_io_components, + "swap_io_components": swap_io_components, "total_blocks": total_blocks, "offload_device": offload_device, "main_device": device, @@ -168,22 +168,22 @@ def apply_block_swap_to_dit(runner, block_swap_config: Dict[str, Any], debug) -> def _configure_io_components(model, device: str, offload_device: str, - offload_io_components: bool, debug) -> List[str]: + swap_io_components: bool, debug) -> List[str]: """Configure I/O component placement and wrapping.""" io_components_offloaded = [] # Process non-block parameters for name, param in model.named_parameters(): if "block" not in name: - target_device = offload_device if offload_io_components else device + target_device = offload_device if swap_io_components else device param.data = param.data.to(target_device, non_blocking=False) - status = "(offloaded)" if offload_io_components else "" + status = "(offloaded)" if swap_io_components else "" debug.log(f" {name} → {target_device} {status}", category="blockswap") # Handle I/O modules with dynamic swapping for name, module in model.named_children(): if name != "blocks": - if offload_io_components: + if swap_io_components: module.to(offload_device) _wrap_io_forward(module, name, model, debug) io_components_offloaded.append(name) @@ -227,7 +227,7 @@ def _configure_blocks(model, device: str, offload_device: str, def _log_memory_summary(memory_stats: Dict[str, float], offload_device: str, - device: str, offload_io_components: bool, + device: str, swap_io_components: bool, debug) -> None: """Log memory usage summary.""" debug.log("BlockSwap memory configuration:", category="blockswap") @@ -241,7 +241,7 @@ def _log_memory_summary(memory_stats: Dict[str, float], offload_device: str, total_memory = memory_stats['offload_memory'] + memory_stats['main_memory'] debug.log(f" Total transformer blocks memory: {total_memory:.2f}MB", category="blockswap") - if offload_io_components and memory_stats.get('io_components'): + if swap_io_components and memory_stats.get('io_components'): debug.log(f" I/O components offloaded: {', '.join(memory_stats['io_components'])}", category="blockswap") diff --git a/src/optimization/memory_manager.py b/src/optimization/memory_manager.py index 5357eb5..aac6b05 100644 --- a/src/optimization/memory_manager.py +++ b/src/optimization/memory_manager.py @@ -606,7 +606,7 @@ def _handle_blockswap_model_movement(runner: Any, model: torch.nn.Module, block.to("cpu") # Handle I/O components - if not runner._block_swap_config.get("offload_io_components", False): + if not runner._block_swap_config.get("swap_io_components", False): # I/O components should be on GPU if not offloaded for name, module in model.named_children(): if name != "blocks": @@ -827,7 +827,7 @@ def cleanup_dit(runner: Any, debug: Optional[Any], keep_model_in_ram: bool = Fal # 5. Clear DiT-related components and temporary attributes dit_components = [ 'sampler', 'sampling_timesteps', 'schedule', - '_dit_checkpoint', '_dit_block_swap_config', '_pending_blockswap_config' + '_dit_checkpoint', '_dit_block_swap_config' ] for component in dit_components: if hasattr(runner, component): @@ -925,7 +925,7 @@ def complete_cleanup(runner: Any, debug: Optional[Any], keep_dit_in_ram: bool = # 4. Clear all temporary loading/configuration attributes temp_attributes = [ - '_dit_checkpoint', '_dit_block_swap_config', '_pending_blockswap_config', + '_dit_checkpoint', '_dit_block_swap_config', '_vae_checkpoint', '_vae_dtype_override' ] for attr in temp_attributes: