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
This commit is contained in:
Adrien Toupet
2025-10-10 17:48:12 -04:00
parent 9ee244d71b
commit fb2b6c7d4c
5 changed files with 35 additions and 32 deletions
+3 -3
View File
@@ -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),
+10 -7
View File
@@ -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)
+7 -7
View File
@@ -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,
}
+12 -12
View File
@@ -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")
+3 -3
View File
@@ -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: