Files
pollockjj-ComfyUI-MultiGPU/block_swap.py
T
John Pollock 04b5bb0a6f refactor: Enhance block swap memory analysis report
The memory analysis function, `analyze_safetensor_distorch`, has been improved to provide a more accurate and detailed report.

Instead of estimating the number of swapped blocks based on an average size, the function now receives the actual list of blocks being swapped. It generates a per-block table detailing each block's ID, type, size, and its final assignment (COMPUTE or SWAP).

This provides users with a precise breakdown of the memory offload, reflecting the actual state of the model rather than a theoretical calculation.

Additionally, the unused `log_memory_usage` helper function has been removed.
2025-08-11 19:19:58 -05:00

231 lines
9.4 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 = {}
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.")
else:
logging.warning(f"[BlockSwap] Model type '{model_type}' is not yet supported for block swapping.")
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