Files
pollockjj-ComfyUI-MultiGPU/block_swap.py
T
John Pollock 235cd267bf feat(swap): Add shell-based block swapping for WanVideo models
This commit introduces a new block swapping mechanism specifically for WanVideo models to enable running them on GPUs with limited VRAM.

A new `WanVideoBlockSwapManager` is implemented which uses a pre-allocation or "shell" strategy. Instead of moving entire blocks between CPU and GPU, this approach:
1.  Pre-allocates a single "shell" block on the GPU, sized to match the largest block in the model.
2.  Offloads designated model blocks to the CPU.
3.  Patches the `forward` method of these offloaded blocks.
4.  During inference, the patched method copies the weights (`state_dict`) from the CPU block into the GPU shell just before execution.

This method avoids the overhead of allocating and deallocating GPU memory for each block, reducing memory fragmentation and potentially improving performance and or corruption copying potentially modified blocks back to the swap space.
2025-08-11 23:41:51 -05:00

469 lines
20 KiB
Python

"""
Block Swap Module for SafeTensor Models
Contains all SafeTensor DisTorch code for block-swap memory optimization
"""
import torch
import logging
import copy
from collections import defaultdict
import comfy.model_management as mm
import torch.nn as nn
from .model_sig import get_model_type
class FluxBlockSwapManager:
"""
Manages block-swapping for FLUX models using the original, fast `forward` patching method.
"""
def __init__(self, model_patcher):
self.model_patcher = model_patcher
self.patched_blocks = {}
def apply_patch(self, compute_device, swap_device, blocks_to_swap):
logging.info(f"[FluxBlockSwapManager] Applying forward patch to {len(blocks_to_swap)} blocks.")
for i, block in enumerate(blocks_to_swap):
block.to(swap_device)
original_forward = block.forward
def create_patched_forward(original_f, b, block_index, cd, sd):
def patched_forward(*args, **kwargs):
b.to(cd, non_blocking=True)
result = original_f(*args, **kwargs)
b.to(sd, non_blocking=True)
return result
return patched_forward
block.forward = create_patched_forward(original_forward, block, i, torch.device(compute_device), torch.device(swap_device))
self.patched_blocks[block] = original_forward
def cleanup(self):
logging.info(f"[FluxBlockSwapManager] Cleaning up {len(self.patched_blocks)} patched blocks.")
for block, original_forward in self.patched_blocks.items():
block.forward = original_forward
self.patched_blocks = {}
class QwenBlockSwapManager:
"""
Manages block-swapping for Qwen models.
"""
def __init__(self, model_patcher):
self.model_patcher = model_patcher
self.patched_blocks = {}
def apply_patch(self, compute_device, swap_device, blocks_to_swap):
logging.info(f"[QwenBlockSwapManager] Applying forward patch to {len(blocks_to_swap)} blocks.")
for i, block in enumerate(blocks_to_swap):
block.to(swap_device)
original_forward = block.forward
def create_patched_forward(original_f, b, block_index, cd, sd):
def patched_forward(*args, **kwargs):
b.to(cd, non_blocking=True)
result = original_f(*args, **kwargs)
b.to(sd, non_blocking=True)
return result
return patched_forward
block.forward = create_patched_forward(original_forward, block, i, torch.device(compute_device), torch.device(swap_device))
self.patched_blocks[block] = original_forward
def cleanup(self):
logging.info(f"[QwenBlockSwapManager] Cleaning up {len(self.patched_blocks)} patched blocks.")
for block, original_forward in self.patched_blocks.items():
block.forward = original_forward
self.patched_blocks = {}
class WanVideoBlockSwapManager:
"""
Manages block-swapping for WanVideo models using a pre-allocated GPU shell block.
"""
def __init__(self, model_patcher, gpu_shell_block):
self.model_patcher = model_patcher
self.gpu_shell_block = gpu_shell_block
self.patched_blocks = {}
def apply_patch(self, compute_device, swap_device, blocks_to_swap):
logging.info(f"[WanVideoBlockSwapManager] Applying state_dict patch to {len(blocks_to_swap)} blocks.")
for i, block in enumerate(blocks_to_swap):
block.to(swap_device) # Ensure the source block is on the swap device
original_forward = block.forward
def create_patched_forward(cpu_block, gpu_shell):
def patched_forward(*args, **kwargs):
logging.info(f"[DEBUG WANVIDEO SWAP] Loading state_dict from CPU block into GPU shell.")
gpu_shell.load_state_dict(cpu_block.state_dict())
logging.info(f"[DEBUG WANVIDEO SWAP] Executing forward pass on GPU shell.")
result = gpu_shell.forward(*args, **kwargs)
return result
return patched_forward
block.forward = create_patched_forward(block, self.gpu_shell_block)
self.patched_blocks[block] = original_forward
def cleanup(self):
logging.info(f"[WanVideoBlockSwapManager] Cleaning up {len(self.patched_blocks)} patched blocks.")
for block, original_forward in self.patched_blocks.items():
block.forward = original_forward
self.patched_blocks = {}
def analyze_safetensor_distorch(model, compute_device, swap_device, virtual_vram_gb, all_blocks, blocks_to_swap):
"""Provides a detailed analysis of the block swap configuration, mimicking the GGUF DisTorch style."""
eq_line = "=" * 60
dash_line = "-" * 60
logging.info(eq_line)
logging.info(" DisTorch SafeTensor Memory Analysis")
logging.info(eq_line)
fmt_assign = "{:<12}{:>15}{:>15}{:>15}"
logging.info(fmt_assign.format("Device", "Role", "Total Mem (GB)", "Config (GB)"))
logging.info(dash_line)
compute_total_gb = mm.get_total_memory(torch.device(compute_device)) / (1024**3)
swap_total_gb = mm.get_total_memory(torch.device(swap_device)) / (1024**3)
logging.info(fmt_assign.format(compute_device, "Compute", f"{compute_total_gb:.2f}", ""))
logging.info(fmt_assign.format(swap_device, "Swap", f"{swap_total_gb:.2f}", f"Offload: {virtual_vram_gb:.2f}"))
logging.info(dash_line)
block_summary = defaultdict(lambda: {'count': 0, 'memory': 0})
total_memory = 0
for block in all_blocks:
block_type = type(block).__name__
block_memory = sum(p.numel() * p.element_size() for p in block.parameters())
block_summary[block_type]['count'] += 1
block_summary[block_type]['memory'] += block_memory
total_memory += block_memory
logging.info(" DisTorch SafeTensor Block Analysis")
logging.info(dash_line)
fmt_layer = "{:<20}{:>10}{:>15}{:>12}"
logging.info(fmt_layer.format("Block Type", "Count", "Memory (MB)", "% Total"))
logging.info(dash_line)
sorted_blocks = sorted(block_summary.items(), key=lambda x: x[1]['memory'], reverse=True)
for block_type, data in sorted_blocks:
mem_mb = data['memory'] / (1024 * 1024)
mem_percent = (data['memory'] / total_memory) * 100 if total_memory > 0 else 0
logging.info(fmt_layer.format(block_type, str(data['count']), f"{mem_mb:.2f}", f"{mem_percent:.1f}%"))
logging.info(dash_line)
logging.info(" DisTorch Final Block Assignments")
logging.info(dash_line)
fmt_final = "{:<5} {:<25} {:>15} {:>15}"
logging.info(fmt_final.format("ID", "Block Type", "Size (MB)", "Assignment"))
logging.info(dash_line)
total_swapped_size_mb = 0
swapped_block_ids = {id(b) for b in blocks_to_swap}
for i, block in enumerate(all_blocks):
block_type = type(block).__name__
size_mb = sum(p.numel() * p.element_size() for p in block.parameters()) / (1024**2)
assignment = "SWAP" if id(block) in swapped_block_ids else "COMPUTE"
if assignment == "SWAP":
total_swapped_size_mb += size_mb
logging.info(fmt_final.format(i, block_type, f"{size_mb:.2f}", assignment))
logging.info(dash_line)
logging.info(f"Total Blocks Swapped: {len(blocks_to_swap)} of {len(all_blocks)}")
logging.info(f"Total VRAM Offloaded: {total_swapped_size_mb / 1024:.2f} GB")
logging.info(eq_line)
def apply_block_swap(model_patcher, compute_device="cuda:0", swap_device="cpu",
virtual_vram_gb=4.0, expert_mode_allocations=""):
"""
Identifies the model type and applies the appropriate block swapping strategy.
"""
model_type = get_model_type(model_patcher)
logging.info(f"[BlockSwap] Detected model type: {model_type}")
if model_type == "FLUX":
manager = FluxBlockSwapManager(model_patcher)
model_to_patch = model_patcher.model.diffusion_model
all_blocks = []
if hasattr(model_to_patch, 'double_blocks'):
all_blocks.extend(model_to_patch.double_blocks)
if hasattr(model_to_patch, 'single_blocks'):
all_blocks.extend(model_to_patch.single_blocks)
if not all_blocks:
logging.error("[BlockSwap] CRITICAL: No swappable blocks found for FLUX model.")
return
model_size_gb = sum(p.numel() * p.element_size() for p in model_to_patch.parameters()) / (1024**3)
if virtual_vram_gb > model_size_gb:
logging.warning(f"[BlockSwap] virtual_vram_gb ({virtual_vram_gb:.2f} GB) is larger than the model size ({model_size_gb:.2f} GB). Truncating to model size.")
virtual_vram_gb = model_size_gb
vram_target_bytes = virtual_vram_gb * (1024**3)
current_swap_size = 0
blocks_to_swap = []
for block in reversed(all_blocks):
if current_swap_size < vram_target_bytes:
block_size = sum(p.numel() * p.element_size() for p in block.parameters())
blocks_to_swap.append(block)
current_swap_size += block_size
else:
break
blocks_to_swap.reverse()
analyze_safetensor_distorch(model_to_patch, compute_device, swap_device, virtual_vram_gb, all_blocks, blocks_to_swap)
if not blocks_to_swap:
logging.warning("[BlockSwap] No blocks designated for swapping.")
return
manager.apply_patch(compute_device, swap_device, blocks_to_swap)
if not hasattr(model_patcher, 'block_swap_managers'):
model_patcher.block_swap_managers = []
model_patcher.block_swap_managers.append(manager)
logging.info("[BlockSwap] FLUX block swap setup complete.")
elif model_type == "QWEN":
manager = QwenBlockSwapManager(model_patcher)
model_to_patch = model_patcher.model.diffusion_model
if not hasattr(model_to_patch, 'transformer_blocks'):
logging.error("[BlockSwap] CRITICAL: Could not find 'transformer_blocks' in Qwen model. Please analyze model structure.")
log_unsupported_model_analysis(model_patcher)
return
all_blocks = model_to_patch.transformer_blocks
if not all_blocks:
logging.error("[BlockSwap] CRITICAL: No swappable blocks found for QWEN model.")
return
model_size_gb = sum(p.numel() * p.element_size() for p in model_to_patch.parameters()) / (1024**3)
if virtual_vram_gb > model_size_gb:
logging.warning(f"[BlockSwap] virtual_vram_gb ({virtual_vram_gb:.2f} GB) is larger than the model size ({model_size_gb:.2f} GB). Truncating to model size.")
virtual_vram_gb = model_size_gb
vram_target_bytes = virtual_vram_gb * (1024**3)
current_swap_size = 0
blocks_to_swap = []
for block in reversed(all_blocks):
if current_swap_size < vram_target_bytes:
block_size = sum(p.numel() * p.element_size() for p in block.parameters())
blocks_to_swap.append(block)
current_swap_size += block_size
else:
break
blocks_to_swap.reverse()
analyze_safetensor_distorch(model_to_patch, compute_device, swap_device, virtual_vram_gb, all_blocks, blocks_to_swap)
if not blocks_to_swap:
logging.warning("[BlockSwap] No blocks designated for swapping for QWEN model.")
return
manager.apply_patch(compute_device, swap_device, blocks_to_swap)
if not hasattr(model_patcher, 'block_swap_managers'):
model_patcher.block_swap_managers = []
model_patcher.block_swap_managers.append(manager)
logging.info("[BlockSwap] QWEN block swap setup complete.")
elif model_type == "WANVIDEO":
model_to_patch = model_patcher.model.diffusion_model
if not hasattr(model_to_patch, 'blocks'):
logging.error("[BlockSwap] CRITICAL: Could not find 'blocks' in WanVideo model. Please analyze model structure.")
log_unsupported_model_analysis(model_patcher)
return
all_blocks = model_to_patch.blocks
if not all_blocks:
logging.error("[BlockSwap] CRITICAL: No swappable blocks found for WanVideo model.")
return
# --- Pre-allocation Strategy ---
# 1. Find the largest block to create a shell
largest_block = max(all_blocks, key=lambda b: sum(p.numel() * p.element_size() for p in b.parameters()))
gpu_shell_block = copy.deepcopy(largest_block).to(compute_device)
shell_size_mb = sum(p.numel() * p.element_size() for p in gpu_shell_block.parameters()) / (1024**2)
logging.info(f"[BlockSwap] Created GPU shell block for WanVideo on {compute_device}, size: {shell_size_mb:.2f} MB")
manager = WanVideoBlockSwapManager(model_patcher, gpu_shell_block)
# --- End Pre-allocation ---
model_size_gb = sum(p.numel() * p.element_size() for p in model_to_patch.parameters()) / (1024**3)
if virtual_vram_gb > model_size_gb:
logging.warning(f"[BlockSwap] virtual_vram_gb ({virtual_vram_gb:.2f} GB) is larger than the model size ({model_size_gb:.2f} GB). Truncating to model size.")
virtual_vram_gb = model_size_gb
vram_target_bytes = virtual_vram_gb * (1024**3)
current_swap_size = 0
blocks_to_swap = []
# We still need to identify which blocks to swap (i.e., which ones will use the shell)
for block in reversed(all_blocks):
if current_swap_size < vram_target_bytes:
block_size = sum(p.numel() * p.element_size() for p in block.parameters())
blocks_to_swap.append(block)
current_swap_size += block_size
else:
# The rest of the blocks will remain on the compute device and not be patched
block.to(compute_device)
blocks_to_swap.reverse()
analyze_safetensor_distorch(model_to_patch, compute_device, swap_device, virtual_vram_gb, all_blocks, blocks_to_swap)
if not blocks_to_swap:
logging.warning("[BlockSwap] No blocks designated for swapping for WanVideo model.")
return
manager.apply_patch(compute_device, swap_device, blocks_to_swap)
if not hasattr(model_patcher, 'block_swap_managers'):
model_patcher.block_swap_managers = []
model_patcher.block_swap_managers.append(manager)
logging.info("[BlockSwap] WanVideo block swap setup complete using pre-allocation strategy.")
else:
logging.warning(f"[BlockSwap] Model type '{model_type}' is not yet supported. Logging model structure for analysis.")
log_unsupported_model_analysis(model_patcher)
def log_unsupported_model_analysis(model_patcher):
"""
Logs the structure of an unsupported model for development purposes.
This is a diagnostic tool and does not modify the model.
"""
logging.info("========================================================================")
logging.info(" INTERNAL MODEL ANALYZER (UNSUPPORTED MODEL DETECTED)")
logging.info("========================================================================")
if not hasattr(model_patcher, 'model'):
logging.error("[ModelAnalyzer] Model patcher does not contain a 'model' attribute.")
return
model = model_patcher.model
logging.info(f"[ModelAnalyzer] Root Model Type: {type(model).__name__}")
if not hasattr(model, 'diffusion_model'):
logging.warning("[ModelAnalyzer] Model does not have a 'diffusion_model' attribute. Dumping root model attributes.")
_recursive_log_attrs(model, "model")
else:
diffusion_model = model.diffusion_model
logging.info(f"[ModelAnalyzer] Diffusion Model Type: {type(diffusion_model).__name__}")
_recursive_log_attrs(diffusion_model, "diffusion_model")
logging.info("========================================================================")
logging.info(" MODEL ANALYSIS COMPLETE")
logging.info("========================================================================")
def _recursive_log_attrs(module, path, seen_modules=None):
"""Helper function to recursively log model attributes."""
if seen_modules is None:
seen_modules = set()
if id(module) in seen_modules:
return
seen_modules.add(id(module))
logging.info(f"--- Analyzing path: '{path}' (Type: {type(module).__name__}) ---")
# Log named children first
children_found = False
for name, submodule in module.named_children():
children_found = True
new_path = f"{path}.{name}" if path else name
# Heuristic check for potential block lists
if isinstance(submodule, (torch.nn.ModuleList, list)) and submodule and all(isinstance(x, torch.nn.Module) for x in submodule):
logging.info(f" > [POTENTIAL BLOCK LIST] '{new_path}' | Type: {type(submodule).__name__}, Length: {len(submodule)}")
# Also inspect the first block in the list for more detail
if len(submodule) > 0:
_recursive_log_attrs(submodule[0], f"{new_path}[0]", seen_modules)
else:
logging.info(f" - Child: '{new_path}' | Type: {type(submodule).__name__}")
# Recurse into non-list children
_recursive_log_attrs(submodule, new_path, seen_modules)
if not children_found:
logging.info(" No named children found at this level.")
def override_class_with_distorch_safetensor(cls):
"""DisTorch 2.0 wrapper for SafeTensor models, providing block-swap memory optimization."""
from .nodes import get_device_list
class NodeOverrideDisTorchSafeTensorv2(cls):
@classmethod
def INPUT_TYPES(s):
inputs = copy.deepcopy(cls.INPUT_TYPES())
devices = get_device_list()
compute_device = devices[1] if len(devices) > 1 else devices[0]
inputs["optional"] = inputs.get("optional", {})
inputs["optional"]["compute_device"] = (devices, {"default": compute_device})
inputs["optional"]["virtual_vram_gb"] = ("FLOAT", {"default": 4.0, "min": 0.0, "max": 128.0, "step": 0.1})
inputs["optional"]["donor_device"] = (devices, {"default": "cpu"})
inputs["optional"]["expert_mode_allocations"] = ("STRING", {"multiline": False, "default": ""})
return inputs
CATEGORY = "multigpu/distorch_2"
FUNCTION = "override"
def override(self, *args, compute_device=None, virtual_vram_gb=4.0,
donor_device="cpu", expert_mode_allocations="", **kwargs):
from . import set_current_device
logging.info(f"[DisTorch SafeTensor] Override called with: compute_device={compute_device}, donor_device={donor_device}, virtual_vram_gb={virtual_vram_gb}")
if compute_device is not None:
set_current_device(compute_device)
fn = getattr(super(), cls.FUNCTION)
out = fn(*args, **kwargs)
model = out[0]
if hasattr(model, 'model'):
logging.info("[DisTorch SafeTensor] Model has 'model' attribute, applying block swap.")
apply_block_swap(
model,
compute_device=compute_device,
swap_device=donor_device,
virtual_vram_gb=virtual_vram_gb,
expert_mode_allocations=expert_mode_allocations
)
else:
logging.warning("[DisTorch SafeTensor] Loaded object does not have a 'model' attribute, skipping block swap.")
return out
return NodeOverrideDisTorchSafeTensorv2
# For backwards compatibility, keep the old name pointing to the new safetensor wrapper
override_class_with_distorch_bs = override_class_with_distorch_safetensor