From d666205fd9b99ae7853c45e6533e7fe722fd9f3b Mon Sep 17 00:00:00 2001 From: John Pollock Date: Mon, 11 Aug 2025 14:28:20 -0500 Subject: [PATCH] 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. --- block_swap.py | 223 +++++++++++++++++++++++++++++++++++++++++--------- 1 file changed, 183 insertions(+), 40 deletions(-) diff --git a/block_swap.py b/block_swap.py index 15b2be1..4dfc319 100644 --- a/block_swap.py +++ b/block_swap.py @@ -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):