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:
+183
-40
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user