Refactor: Overhaul BlockSwap with hook-based manager

This commit completely rewrites the block swapping implementation for improved stability, correctness, and code structure.

Key changes:
- Replaces the fragile monkey-patching of the `forward` method with the standard PyTorch `register_forward_pre_hook`.
- Introduces a `BlockSwapManager` class to encapsulate all swapping logic, separating it from the ComfyUI node.
- Implements a "Sequential Swapping" strategy: the previously active block is offloaded before the current block is loaded, ensuring only one block is on the active device at a time.
- Adds a `cleanup` method to properly remove hooks after execution, preventing state leakage between runs.
- Fixes a critical bug where the hook signature was incorrect.
- Adds a memory logging utility for easier debugging.
This commit is contained in:
John Pollock
2025-08-11 14:28:20 -05:00
parent 62718ea12f
commit d666205fd9
+183 -40
View File
@@ -8,6 +8,151 @@ import logging
import copy
from collections import defaultdict
import comfy.model_management as mm
import torch.nn as nn
def log_memory_usage(device, stage=""):
"""Logs the memory usage of a given device."""
if not isinstance(device, torch.device):
device = torch.device(device)
if device.type == 'cuda':
stats = torch.cuda.memory_stats(device)
total_mem = mm.get_total_memory(device)
allocated = stats['allocated_bytes.all.current']
reserved = stats['reserved_bytes.all.current']
logging.info(
f"[MemLog] {stage} - {device}: "
f"Allocated: {allocated / 1024**2:.2f}MB, "
f"Reserved: {reserved / 1024**2:.2f}MB, "
f"Total: {total_mem / 1024**3:.2f}GB"
)
elif device.type == 'cpu':
# Basic CPU memory logging (less detailed than CUDA)
# This requires psutil, which might not be a dependency.
# For now, we'll just log that it's a CPU.
logging.info(f"[MemLog] {stage} - {device}: CPU memory logging is not as detailed.")
class BlockSwapManager:
"""
Manages block-swapping memory optimization using PyTorch hooks.
"""
def __init__(self, swap_device='cpu'):
self.swap_device = torch.device(swap_device)
# Determine the execution device (e.g., GPU)
self.active_device = mm.get_torch_device()
self.active_block = None
# Use a set of IDs for fast lookup of managed blocks
self.managed_block_ids = set()
self.hooks = []
def move_block(self, block, device):
"""Moves a block to the specified device."""
# Avoid moving if the target device is 'meta'
if torch.device(device).type != 'meta':
try:
block.to(device)
except Exception as e:
print(f"[BlockSwap] Warning: Failed to move block {type(block).__name__} to {device}: {e}")
def _get_block_device(self, block):
"""Robustly determines the current device of a block."""
try:
# Check the device of the first parameter found in the block
param = next(block.parameters(), None)
if param is not None:
return param.device
except Exception:
pass
return None
# CRITICAL FIX: The hook signature must accept (module, args).
def before_block_execution(self, block, args):
"""
Hook function called before a block's execution.
Implements Sequential Swapping (WanVideoWrapper style).
"""
block_id = id(block)
if block_id not in self.managed_block_ids:
return
# 1. Handle Offloading (if the active block is changing)
# CRITICAL FIX: This logic must execute regardless of the current block's device.
if self.active_block != block:
# Offload the previous block if it exists and is managed by us
if self.active_block is not None and id(self.active_block) in self.managed_block_ids:
# print(f"[BlockSwap] Offloading previous block to {self.swap_device}")
self.move_block(self.active_block, self.swap_device)
# 2. Handle Loading (only if needed)
current_device = self._get_block_device(block)
if current_device is None:
# Block has no parameters, skip loading
pass
elif current_device != self.active_device:
# print(f"[BlockSwap] Loading current block to {self.active_device}")
self.move_block(block, self.active_device)
# 3. Update the tracker
self.active_block = block
def apply_swap_optimization(self, swappable_blocks):
"""
Applies the block-swapping optimization using PyTorch forward hooks.
"""
if not swappable_blocks:
return
# print(f"[BlockSwap] Applying optimization to {len(swappable_blocks)} blocks.")
for block in swappable_blocks:
if not isinstance(block, nn.Module) or id(block) in self.managed_block_ids:
continue
# Clean up potential previous manual patches
if hasattr(block, 'original_forward'):
try:
block.forward = block.original_forward
del block.original_forward
except Exception:
pass
block_id = id(block)
self.managed_block_ids.add(block_id)
# Use register_forward_pre_hook for robustness.
try:
# CRITICAL FIX: Register the method directly, now that its signature is correct.
hook = block.register_forward_pre_hook(
self.before_block_execution
)
self.hooks.append(hook)
except Exception as e:
print(f"[BlockSwap] Warning: Failed to register hook for block {type(block).__name__}: {e}")
self.managed_block_ids.remove(block_id)
continue
# Move to CPU initially (if it has parameters)
if self._get_block_device(block) is not None:
self.move_block(block, self.swap_device)
def cleanup(self):
"""Removes hooks and restores the model state."""
# Remove hooks
for hook in self.hooks:
hook.remove()
self.hooks = []
# We rely on the surrounding environment (ComfyUI) to manage the overall
# model placement after sampling, but we ensure the last active block is returned to the GPU if needed.
if self.active_block is not None and self._get_block_device(self.active_block) != self.active_device:
self.move_block(self.active_block, self.active_device)
self.managed_block_ids = set()
self.active_block = None
def analyze_safetensor_distorch(model, compute_device, swap_device, virtual_vram_gb, all_blocks):
@@ -77,83 +222,81 @@ def analyze_safetensor_distorch(model, compute_device, swap_device, virtual_vram
def apply_block_swap(model_patcher, compute_device="cuda:0", swap_device="cpu",
virtual_vram_gb=4.0, expert_mode_allocations=""):
"""
Applies WanVideo-style block swapping by patching the forward method of individual model blocks.
This allows for offloading parts of the model to a swap device to conserve VRAM.
Applies block swapping using a manager and PyTorch hooks for robustness.
"""
logging.info(f"[DisTorch SafeTensor] Initializing block swap: compute_device={compute_device}, swap_device={swap_device}")
logging.info(f"[BlockSwap] Initializing block swap: compute_device={compute_device}, swap_device={swap_device}")
model_to_patch = None
if hasattr(model_patcher, 'model') and hasattr(model_patcher.model, 'diffusion_model'):
model_to_patch = model_patcher.model.diffusion_model
logging.info("[DisTorch SafeTensor] Found 'diffusion_model' attribute for patching.")
logging.info("[BlockSwap] Found 'diffusion_model' for patching.")
elif hasattr(model_patcher, 'model'):
model_to_patch = model_patcher.model
logging.info("[DisTorch SafeTensor] Found 'model' attribute for patching.")
logging.info("[BlockSwap] Found 'model' for patching.")
else:
logging.error("[DisTorch SafeTensor] Could not find a valid model to patch for block swapping.")
logging.error("[BlockSwap] Could not find a valid model to patch.")
return
all_blocks = []
# 1. Standard UNet Structure
# Block identification logic (remains the same)
if hasattr(model_to_patch, 'input_blocks') and hasattr(model_to_patch, 'middle_block') and hasattr(model_to_patch, 'output_blocks'):
logging.info("[DisTorch SafeTensor] Found standard UNet structure ('input_blocks', 'middle_block', 'output_blocks').")
logging.info("[BlockSwap] Found standard UNet structure.")
all_blocks.extend(model_to_patch.input_blocks)
if isinstance(model_to_patch.middle_block, torch.nn.Module):
all_blocks.append(model_to_patch.middle_block)
all_blocks.extend(model_to_patch.output_blocks)
# 2. Simple 'blocks' attribute
elif hasattr(model_to_patch, 'blocks') and isinstance(model_to_patch.blocks, torch.nn.ModuleList):
logging.info("[DisTorch SafeTensor] Found 'blocks' attribute of type ModuleList.")
logging.info("[BlockSwap] Found 'blocks' attribute.")
all_blocks.extend(model_to_patch.blocks)
# 3. Simple 'layers' attribute
elif hasattr(model_to_patch, 'layers') and isinstance(model_to_patch.layers, torch.nn.ModuleList):
logging.info("[DisTorch SafeTensor] Found 'layers' attribute of type ModuleList.")
logging.info("[BlockSwap] Found 'layers' attribute.")
all_blocks.extend(model_to_patch.layers)
# 4. Fallback to top-level ModuleLists
else:
logging.info("[DisTorch SafeTensor] No standard structure found. Falling back to searching for top-level ModuleLists.")
logging.info("[BlockSwap] No standard structure found. Searching for top-level ModuleLists.")
for child in model_to_patch.children():
if isinstance(child, torch.nn.ModuleList):
logging.info(f"[DisTorch SafeTensor] Found top-level ModuleList with {len(child)} modules. Adding them as blocks.")
all_blocks.extend(child)
if not all_blocks:
logging.error("[DisTorch SafeTensor] CRITICAL: No swappable blocks were found in the model. Block swap cannot be applied.")
logging.error("[BlockSwap] CRITICAL: No swappable blocks found.")
return
logging.info(f"[DisTorch SafeTensor] Successfully identified {len(all_blocks)} swappable blocks.")
logging.info(f"[BlockSwap] Identified {len(all_blocks)} swappable blocks.")
# Run and display the analysis
# Run analysis before making changes
analyze_safetensor_distorch(model_to_patch, compute_device, swap_device, virtual_vram_gb, all_blocks)
# Log initial memory state
log_memory_usage(compute_device, "Before Swap")
log_memory_usage(swap_device, "Before Swap")
model_size_gb = sum(p.numel() * p.element_size() for p in model_to_patch.parameters()) / (1024**3)
block_size_gb = model_size_gb / len(all_blocks) if all_blocks else 0
blocks_to_offload = int(virtual_vram_gb / block_size_gb) if block_size_gb > 0 else 0
blocks_on_compute = len(all_blocks) - blocks_to_offload
blocks_to_offload_count = int(virtual_vram_gb / block_size_gb) if block_size_gb > 0 else 0
# The blocks at the end of the list are swapped
blocks_to_swap = all_blocks[-blocks_to_offload_count:] if blocks_to_offload_count > 0 else []
for i, block in enumerate(all_blocks):
# Determine target device for this block
target_device = compute_device if i < blocks_on_compute else swap_device
block.to(target_device)
if not blocks_to_swap:
logging.warning("[BlockSwap] No blocks designated for swapping based on virtual_vram_gb. Skipping hook setup.")
return
# Patch the forward method only if the block is on the swap device
if target_device == swap_device:
original_forward = block.forward
def create_patched_forward(original_f, b, block_index, cd, sd):
def patched_forward(*args, **kwargs):
logging.info(f"[DisTorch SafeTensor] Swapping block {block_index} to {cd} for computation.")
b.to(cd, non_blocking=True)
result = original_f(*args, **kwargs)
logging.info(f"[DisTorch SafeTensor] Swapping block {block_index} back to {sd}.")
b.to(sd, non_blocking=True)
return result
return patched_forward
# Instantiate and apply the manager
manager = BlockSwapManager(swap_device=swap_device)
manager.apply_swap_optimization(blocks_to_swap)
block.forward = create_patched_forward(original_forward, block, i, torch.device(compute_device), torch.device(swap_device))
logging.info(f"[DisTorch SafeTensor] Patched forward method for block {i} on {swap_device}.")
# Store the manager on the model_patcher for lifecycle management (e.g., cleanup)
if not hasattr(model_patcher, 'block_swap_managers'):
model_patcher.block_swap_managers = []
model_patcher.block_swap_managers.append(manager)
logging.info(f"[BlockSwap] Moved {len(blocks_to_swap)} blocks to {swap_device} and applied hooks.")
logging.info("[DisTorch SafeTensor] Block swap setup complete.")
# Log memory state after moving blocks
log_memory_usage(compute_device, "After Swap")
log_memory_usage(swap_device, "After Swap")
logging.info("[BlockSwap] Block swap setup complete.")
def override_class_with_distorch_safetensor(cls):