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.
231 lines
9.4 KiB
Python
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
|