This commit introduces block swapping functionality for Qwen models, enabling them to run on systems with limited VRAM by offloading layers to a swap device (e.g., CPU RAM). Key changes: - A new `QwenBlockSwapManager` class is implemented to handle the patching of Qwen transformer blocks. - The `apply_block_swap` function is extended to detect Qwen models and apply the swapping logic to their `transformer_blocks`. - A model signature for Qwen is added to `model_sig.py` to correctly identify the swappable modules. - A new diagnostic function, `log_unsupported_model_analysis`, is added to log the structure of unsupported models, aiding future development.
375 lines
16 KiB
Python
375 lines
16 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 = {}
|
|
|
|
|
|
class QwenBlockSwapManager:
|
|
"""
|
|
Manages block-swapping for Qwen models.
|
|
"""
|
|
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"[QwenBlockSwapManager] 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"[QwenBlockSwapManager] 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.")
|
|
|
|
elif model_type == "QWEN":
|
|
manager = QwenBlockSwapManager(model_patcher)
|
|
|
|
model_to_patch = model_patcher.model.diffusion_model
|
|
|
|
if not hasattr(model_to_patch, 'transformer_blocks'):
|
|
logging.error("[BlockSwap] CRITICAL: Could not find 'transformer_blocks' in Qwen model. Please analyze model structure.")
|
|
log_unsupported_model_analysis(model_patcher)
|
|
return
|
|
|
|
all_blocks = model_to_patch.transformer_blocks
|
|
|
|
if not all_blocks:
|
|
logging.error("[BlockSwap] CRITICAL: No swappable blocks found for QWEN 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 for QWEN model.")
|
|
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] QWEN block swap setup complete.")
|
|
|
|
else:
|
|
logging.warning(f"[BlockSwap] Model type '{model_type}' is not yet supported. Logging model structure for analysis.")
|
|
log_unsupported_model_analysis(model_patcher)
|
|
|
|
|
|
def log_unsupported_model_analysis(model_patcher):
|
|
"""
|
|
Logs the structure of an unsupported model for development purposes.
|
|
This is a diagnostic tool and does not modify the model.
|
|
"""
|
|
logging.info("========================================================================")
|
|
logging.info(" INTERNAL MODEL ANALYZER (UNSUPPORTED MODEL DETECTED)")
|
|
logging.info("========================================================================")
|
|
|
|
if not hasattr(model_patcher, 'model'):
|
|
logging.error("[ModelAnalyzer] Model patcher does not contain a 'model' attribute.")
|
|
return
|
|
|
|
model = model_patcher.model
|
|
logging.info(f"[ModelAnalyzer] Root Model Type: {type(model).__name__}")
|
|
|
|
if not hasattr(model, 'diffusion_model'):
|
|
logging.warning("[ModelAnalyzer] Model does not have a 'diffusion_model' attribute. Dumping root model attributes.")
|
|
_recursive_log_attrs(model, "model")
|
|
else:
|
|
diffusion_model = model.diffusion_model
|
|
logging.info(f"[ModelAnalyzer] Diffusion Model Type: {type(diffusion_model).__name__}")
|
|
_recursive_log_attrs(diffusion_model, "diffusion_model")
|
|
|
|
logging.info("========================================================================")
|
|
logging.info(" MODEL ANALYSIS COMPLETE")
|
|
logging.info("========================================================================")
|
|
|
|
def _recursive_log_attrs(module, path, seen_modules=None):
|
|
"""Helper function to recursively log model attributes."""
|
|
if seen_modules is None:
|
|
seen_modules = set()
|
|
|
|
if id(module) in seen_modules:
|
|
return
|
|
seen_modules.add(id(module))
|
|
|
|
logging.info(f"--- Analyzing path: '{path}' (Type: {type(module).__name__}) ---")
|
|
|
|
# Log named children first
|
|
children_found = False
|
|
for name, submodule in module.named_children():
|
|
children_found = True
|
|
new_path = f"{path}.{name}" if path else name
|
|
|
|
# Heuristic check for potential block lists
|
|
if isinstance(submodule, (torch.nn.ModuleList, list)) and submodule and all(isinstance(x, torch.nn.Module) for x in submodule):
|
|
logging.info(f" > [POTENTIAL BLOCK LIST] '{new_path}' | Type: {type(submodule).__name__}, Length: {len(submodule)}")
|
|
# Also inspect the first block in the list for more detail
|
|
if len(submodule) > 0:
|
|
_recursive_log_attrs(submodule[0], f"{new_path}[0]", seen_modules)
|
|
else:
|
|
logging.info(f" - Child: '{new_path}' | Type: {type(submodule).__name__}")
|
|
# Recurse into non-list children
|
|
_recursive_log_attrs(submodule, new_path, seen_modules)
|
|
|
|
if not children_found:
|
|
logging.info(" No named children found at this level.")
|
|
|
|
|
|
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
|