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