refactor: Move core logic into separate modules

This commit refactors the codebase by extracting major components from the main `__init__.py` file into their own dedicated modules. This improves code organization, readability, and maintainability.

- **`distorch.py`**: New file containing the `DisTorch` class, which manages multi-GPU device patching and distribution logic.
- **`block_swap.py`**: New file containing the generic `BlockSwap` class for UNet block swapping to manage VRAM.
- **`wanvideo.py`**: New file containing the `WanVideoBlockSwap` class, a specialized implementation for WanVideo models.
- **`__init__.py`**: Simplified to handle node registration and imports from the new modules.
This commit is contained in:
John Pollock
2025-08-10 09:53:33 -05:00
parent 3e3190346b
commit 898169fccf
5 changed files with 1321 additions and 1260 deletions
+116 -741
View File
@@ -1,36 +1,15 @@
import copy
import torch
import sys
import comfy.model_management as mm
import os
from pathlib import Path
import logging
import os
import copy
from pathlib import Path
import folder_paths
from collections import defaultdict
import hashlib
import comfy.utils
from typing import Dict, List
import comfy.model_management as mm
from nodes import NODE_CLASS_MAPPINGS as GLOBAL_NODE_CLASS_MAPPINGS
from .nodes import (
UnetLoaderGGUF, UnetLoaderGGUFAdvanced,
CLIPLoaderGGUF, DualCLIPLoaderGGUF, TripleCLIPLoaderGGUF, QuadrupleCLIPLoaderGGUF,
LTXVLoader,
Florence2ModelLoader, DownloadAndLoadFlorence2Model,
CheckpointLoaderNF4,
LoadFluxControlNet,
MMAudioModelLoader, MMAudioFeatureUtilsLoader, MMAudioSampler,
PulidModelLoader, PulidInsightFaceLoader, PulidEvaClipLoader,
HyVideoModelLoader, HyVideoVAELoader, DownloadAndLoadHyVideoTextEncoder,
WanVideoModelLoader, WanVideoModelLoader_2, WanVideoVAELoader, LoadWanVideoT5TextEncoder, LoadWanVideoClipTextEncoder,
WanVideoTextEncode, WanVideoBlockSwap, WanVideoSampler
)
# DisTorch import removed - all implementations now in __init__.py
# Global device state management
current_device = mm.get_torch_device()
current_text_encoder_device = mm.text_encoder_device()
model_allocation_store = {}
def _has_xpu():
try:
@@ -38,302 +17,6 @@ def _has_xpu():
except Exception:
return False
def get_torch_device_patched():
device = None
if (not (torch.cuda.is_available() or _has_xpu()) or mm.cpu_state == mm.CPUState.CPU or "cpu" in str(current_device).lower()):
device = torch.device("cpu")
else:
devs = set(get_device_list())
device = torch.device(current_device) if str(current_device) in devs else torch.device("cpu")
logging.info(f"[MultiGPU get_torch_device_patched] Returning device: {device} (current_device={current_device})")
return device
def text_encoder_device_patched():
device = None
if (not (torch.cuda.is_available() or _has_xpu()) or mm.cpu_state == mm.CPUState.CPU or "cpu" in str(current_text_encoder_device).lower()):
device = torch.device("cpu")
else:
devs = set(get_device_list())
device = torch.device(current_text_encoder_device) if str(current_text_encoder_device) in devs else torch.device("cpu")
logging.info(f"[MultiGPU text_encoder_device_patched] Returning device: {device} (current_text_encoder_device={current_text_encoder_device})")
return device
logging.info(f"[MultiGPU] Patching mm.get_torch_device and mm.text_encoder_device")
logging.info(f"[MultiGPU] Initial current_device: {current_device}")
logging.info(f"[MultiGPU] Initial current_text_encoder_device: {current_text_encoder_device}")
mm.get_torch_device = get_torch_device_patched
mm.text_encoder_device = text_encoder_device_patched
logging.info(f"[MultiGPU] Patches applied successfully")
def create_model_hash(model, caller):
model_type = type(model.model).__name__
model_size = model.model_size()
first_layers = str(list(model.model_state_dict().keys())[:3])
identifier = f"{model_type}_{model_size}_{first_layers}"
final_hash = hashlib.sha256(identifier.encode()).hexdigest()
return final_hash
def register_patched_ggufmodelpatcher():
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["UnetLoaderGGUF"]
module = sys.modules[original_loader.__module__]
if not hasattr(module.GGUFModelPatcher, '_patched'):
original_load = module.GGUFModelPatcher.load
def new_load(self, *args, force_patch_weights=False, **kwargs):
global model_allocation_store
super(module.GGUFModelPatcher, self).load(*args, force_patch_weights=True, **kwargs)
debug_hash = create_model_hash(self, "patcher")
linked = []
module_count = 0
for n, m in self.model.named_modules():
module_count += 1
if hasattr(m, "weight"):
device = getattr(m.weight, "device", None)
if device is not None:
linked.append((n, m))
continue
if hasattr(m, "bias"):
device = getattr(m.bias, "device", None)
if device is not None:
linked.append((n, m))
continue
if linked:
if hasattr(self, 'model'):
debug_hash = create_model_hash(self, "patcher")
debug_allocations = model_allocation_store.get(debug_hash)
if debug_allocations:
device_assignments = analyze_ggml_loading(self.model, debug_allocations)['device_assignments']
for device, layers in device_assignments.items():
target_device = torch.device(device)
for n, m, _ in layers:
m.to(self.load_device).to(target_device)
self.mmap_released = True
module.GGUFModelPatcher.load = new_load
module.GGUFModelPatcher._patched = True
def analyze_ggml_loading(model, allocations_str):
DEVICE_RATIOS_DISTORCH = {}
device_table = {}
distorch_alloc = allocations_str
virtual_vram_gb = 0.0
if '#' in allocations_str:
distorch_alloc, virtual_vram_str = allocations_str.split('#')
if not distorch_alloc:
distorch_alloc = calculate_vvram_allocation_string(model, virtual_vram_str)
eq_line = "=" * 47
dash_line = "-" * 47
fmt_assign = "{:<12}{:>10}{:>14}{:>10}"
for allocation in distorch_alloc.split(';'):
dev_name, fraction = allocation.split(',')
fraction = float(fraction)
total_mem_bytes = mm.get_total_memory(torch.device(dev_name))
alloc_gb = (total_mem_bytes * fraction) / (1024**3)
DEVICE_RATIOS_DISTORCH[dev_name] = alloc_gb
device_table[dev_name] = {
"fraction": fraction,
"total_gb": total_mem_bytes / (1024**3),
"alloc_gb": alloc_gb
}
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
logging.info(eq_line)
logging.info(" DisTorch Device Allocations")
logging.info(eq_line)
logging.info(fmt_assign.format("Device", "Alloc %", "Total (GB)", " Alloc (GB)"))
logging.info(dash_line)
sorted_devices = sorted(device_table.keys(), key=lambda d: (d == "cpu", d))
for dev in sorted_devices:
frac = device_table[dev]["fraction"]
tot_gb = device_table[dev]["total_gb"]
alloc_gb = device_table[dev]["alloc_gb"]
logging.info(fmt_assign.format(dev,f"{int(frac * 100)}%",f"{tot_gb:.2f}",f"{alloc_gb:.2f}"))
logging.info(dash_line)
layer_summary = {}
layer_list = []
memory_by_type = defaultdict(int)
total_memory = 0
for name, module in model.named_modules():
if hasattr(module, "weight"):
layer_type = type(module).__name__
layer_summary[layer_type] = layer_summary.get(layer_type, 0) + 1
layer_list.append((name, module, layer_type))
layer_memory = 0
if module.weight is not None:
layer_memory += module.weight.numel() * module.weight.element_size()
if hasattr(module, "bias") and module.bias is not None:
layer_memory += module.bias.numel() * module.bias.element_size()
memory_by_type[layer_type] += layer_memory
total_memory += layer_memory
logging.info(" DisTorch GGML Layer Distribution")
logging.info(dash_line)
fmt_layer = "{:<12}{:>10}{:>14}{:>10}"
logging.info(fmt_layer.format("Layer Type", "Layers", "Memory (MB)", "% Total"))
logging.info(dash_line)
for layer_type, count in layer_summary.items():
mem_mb = memory_by_type[layer_type] / (1024 * 1024)
mem_percent = (memory_by_type[layer_type] / total_memory) * 100 if total_memory > 0 else 0
logging.info(fmt_layer.format(layer_type,str(count),f"{mem_mb:.2f}",f"{mem_percent:.1f}%"))
logging.info(dash_line)
nonzero_devices = [d for d, r in DEVICE_RATIOS_DISTORCH.items() if r > 0]
nonzero_total_ratio = sum(DEVICE_RATIOS_DISTORCH[d] for d in nonzero_devices)
device_assignments = {device: [] for device in DEVICE_RATIOS_DISTORCH.keys()}
total_layers = len(layer_list)
current_layer = 0
for idx, device in enumerate(nonzero_devices):
ratio = DEVICE_RATIOS_DISTORCH[device]
if idx == len(nonzero_devices) - 1:
device_layer_count = total_layers - current_layer
else:
device_layer_count = int((ratio / nonzero_total_ratio) * total_layers)
start_idx = current_layer
end_idx = current_layer + device_layer_count
device_assignments[device] = layer_list[start_idx:end_idx]
current_layer += device_layer_count
logging.info(" DisTorch Final Device/Layer Assignments")
logging.info(dash_line)
fmt_assign = "{:<12}{:>10}{:>14}{:>10}"
logging.info(fmt_assign.format("Device", "Layers", "Memory (MB)", "% Total"))
logging.info(dash_line)
total_assigned_memory = 0
device_memories = {}
for device, layers in device_assignments.items():
device_memory = 0
for layer_type in layer_summary:
type_layers = sum(1 for _, _, lt in layers if lt == layer_type)
if layer_summary[layer_type] > 0:
mem_per_layer = memory_by_type[layer_type] / layer_summary[layer_type]
device_memory += mem_per_layer * type_layers
device_memories[device] = device_memory
total_assigned_memory += device_memory
sorted_assignments = sorted(device_assignments.keys(), key=lambda d: (d == "cpu", d))
for dev in sorted_assignments:
layers = device_assignments[dev]
mem_mb = device_memories[dev] / (1024 * 1024)
mem_percent = (device_memories[dev] / total_memory) * 100 if total_memory > 0 else 0
logging.info(fmt_assign.format(dev,str(len(layers)),f"{mem_mb:.2f}",f"{mem_percent:.1f}%"))
logging.info(dash_line)
return {"device_assignments": device_assignments}
def calculate_vvram_allocation_string(model, virtual_vram_str):
recipient_device, vram_amount, donors = virtual_vram_str.split(';')
virtual_vram_gb = float(vram_amount)
eq_line = "=" * 47
dash_line = "-" * 47
fmt_assign = "{:<8} {:<6} {:>11} {:>9} {:>9}"
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
logging.info(eq_line)
logging.info(" DisTorch Virtual VRAM Analysis")
logging.info(eq_line)
logging.info(fmt_assign.format("Object", "Role", "Original(GB)", "Total(GB)", "Virt(GB)"))
logging.info(dash_line)
recipient_vram = mm.get_total_memory(torch.device(recipient_device)) / (1024**3)
recipient_virtual = recipient_vram + virtual_vram_gb
logging.info(fmt_assign.format(recipient_device, 'recip', f"{recipient_vram:.2f}GB",f"{recipient_virtual:.2f}GB", f"+{virtual_vram_gb:.2f}GB"))
ram_donors = [d for d in donors.split(',') if d != 'cpu']
remaining_vram_needed = virtual_vram_gb
donor_device_info = {}
donor_allocations = {}
for donor in ram_donors:
donor_vram = mm.get_total_memory(torch.device(donor)) / (1024**3)
max_donor_capacity = donor_vram * 0.9
donation = min(remaining_vram_needed, max_donor_capacity)
donor_virtual = donor_vram - donation
remaining_vram_needed -= donation
donor_allocations[donor] = donation
donor_device_info[donor] = (donor_vram, donor_virtual)
logging.info(fmt_assign.format(donor, 'donor', f"{donor_vram:.2f}GB", f"{donor_virtual:.2f}GB", f"-{donation:.2f}GB"))
system_dram_gb = mm.get_total_memory(torch.device('cpu')) / (1024**3)
cpu_donation = remaining_vram_needed
cpu_virtual = system_dram_gb - cpu_donation
donor_allocations['cpu'] = cpu_donation
logging.info(fmt_assign.format('cpu', 'donor', f"{system_dram_gb:.2f}GB", f"{cpu_virtual:.2f}GB", f"-{cpu_donation:.2f}GB"))
logging.info(dash_line)
layer_summary = {}
layer_list = []
memory_by_type = defaultdict(int)
total_memory = 0
for name, module in model.named_modules():
if hasattr(module, "weight"):
layer_type = type(module).__name__
layer_summary[layer_type] = layer_summary.get(layer_type, 0) + 1
layer_list.append((name, module, layer_type))
layer_memory = 0
if module.weight is not None:
layer_memory += module.weight.numel() * module.weight.element_size()
if hasattr(module, "bias") and module.bias is not None:
layer_memory += module.bias.numel() * module.bias.element_size()
memory_by_type[layer_type] += layer_memory
total_memory += layer_memory
model_size_gb = total_memory / (1024**3)
new_model_size_gb = max(0, model_size_gb - virtual_vram_gb)
logging.info(fmt_assign.format('model', 'model', f"{model_size_gb:.2f}GB",f"{new_model_size_gb:.2f}GB", f"-{virtual_vram_gb:.2f}GB"))
if model_size_gb > (recipient_vram * 0.9):
on_recipient = recipient_vram * 0.9
on_virtuals = model_size_gb - on_recipient
logging.info(f"\nWarning: Model size is greater than 90% of recipient VRAM. {on_virtuals:.2f} GB of GGML Layers Offloaded Automatically to Virtual VRAM.\n")
else:
on_recipient = model_size_gb
on_virtuals = 0
new_on_recipient = max(0, on_recipient - virtual_vram_gb)
allocation_parts = []
recipient_percent = new_on_recipient / recipient_vram
allocation_parts.append(f"{recipient_device},{recipient_percent:.4f}")
for donor in ram_donors:
donor_vram = donor_device_info[donor][0]
donor_percent = donor_allocations[donor] / donor_vram
allocation_parts.append(f"{donor},{donor_percent:.4f}")
cpu_percent = donor_allocations['cpu'] / system_dram_gb
allocation_parts.append(f"cpu,{cpu_percent:.4f}")
allocation_string = ";".join(allocation_parts)
fmt_mem = "{:<20}{:>20}"
logging.info(fmt_mem.format("\nAllocation String", allocation_string))
return allocation_string
def get_device_list():
devs = ["cpu"]
try:
@@ -348,58 +31,15 @@ def get_device_list():
pass
return devs
class DeviceSelectorMultiGPU:
@classmethod
def INPUT_TYPES(s):
devices = get_device_list()
return {
"required": {
"device": (devices, {"default": devices[1] if len(devices) > 1 else devices[0]})
}
}
RETURN_TYPES = (get_device_list(),)
RETURN_NAMES = ("device",)
FUNCTION = "select_device"
CATEGORY = "multigpu"
def select_device(self, device):
return (device,)
class HunyuanVideoEmbeddingsAdapter:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"hyvid_embeds": ("HYVIDEMBEDS",),
}
}
RETURN_TYPES = ("CONDITIONING",)
FUNCTION = "adapt_embeddings"
CATEGORY = "multigpu"
def adapt_embeddings(self, hyvid_embeds):
cond = hyvid_embeds["prompt_embeds"]
pooled_dict = {
"pooled_output": hyvid_embeds["prompt_embeds_2"],
"cross_attn": hyvid_embeds["prompt_embeds"],
"attention_mask": hyvid_embeds["attention_mask"],
}
if hyvid_embeds["attention_mask_2"] is not None:
pooled_dict["attention_mask_controlnet"] = hyvid_embeds["attention_mask_2"]
if hyvid_embeds["cfg"] is not None:
pooled_dict["guidance"] = float(hyvid_embeds["cfg"])
pooled_dict["start_percent"] = float(hyvid_embeds["start_percent"]) if hyvid_embeds["start_percent"] is not None else 0.0
pooled_dict["end_percent"] = float(hyvid_embeds["end_percent"]) if hyvid_embeds["end_percent"] is not None else 1.0
return ([[cond, pooled_dict]],)
def set_current_device(device):
global current_device
current_device = device
logging.info(f"[MultiGPU] current_device set to: {device}")
def set_current_text_encoder_device(device):
global current_text_encoder_device
current_text_encoder_device = device
logging.info(f"[MultiGPU] current_text_encoder_device set to: {device}")
def override_class(cls):
class NodeOverride(cls):
@@ -416,12 +56,10 @@ def override_class(cls):
FUNCTION = "override"
def override(self, *args, device=None, **kwargs):
global current_device
logging.info(f"[MultiGPU override_class] Called with device={device}, current_device={current_device}")
logging.info(f"[MultiGPU override_class] Called with device={device}")
if device is not None:
current_device = device
set_current_device(device)
logging.info(f"[MultiGPU override_class] Setting current_device to {device}")
fn = getattr(super(), cls.FUNCTION)
@@ -448,10 +86,8 @@ def override_class_clip(cls):
FUNCTION = "override"
def override(self, *args, device=None, **kwargs):
global current_text_encoder_device
if device is not None:
current_text_encoder_device = device
set_current_text_encoder_device(device)
fn = getattr(super(), cls.FUNCTION)
out = fn(*args, **kwargs)
@@ -460,363 +96,33 @@ def override_class_clip(cls):
return NodeOverride
def override_class_with_distorch_safetensor(cls):
"""DisTorch wrapper for SafeTensor models, providing block-swap memory optimization."""
class NodeOverrideDisTorchSafeTensor(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_ram_gb"] = ("FLOAT", {"default": 4.0, "min": 0.0, "max": 100.0, "step": 0.1})
inputs["optional"]["donor_device"] = (devices, {"default": "cpu"})
inputs["optional"]["expert_mode_allocations"] = ("STRING", {"multiline": False, "default": ""})
return inputs
CATEGORY = "multigpu"
FUNCTION = "override"
def override(self, *args, compute_device=None, virtual_ram_gb=4.0,
donor_device="cpu", expert_mode_allocations="", **kwargs):
global current_device
logging.info(f"[DisTorch SafeTensor] Override called with: compute_device={compute_device}, donor_device={donor_device}, virtual_ram_gb={virtual_ram_gb}")
if compute_device is not None:
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_ram_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 NodeOverrideDisTorchSafeTensor
def analyze_safetensor_distorch(model, compute_device, swap_device, virtual_vram_gb, reserved_swap_gb, all_blocks):
"""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)
# Device Allocation Table
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}", f"Reserve: {reserved_swap_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 Analysis Table
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)
# Final Assignment Table
model_size_gb = total_memory / (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
logging.info(" DisTorch Final Block Assignments")
logging.info(dash_line)
fmt_final = "{:<20}{:>15}"
logging.info(fmt_final.format("Total Model Size (GB):", f"{model_size_gb:.2f}"))
logging.info(fmt_final.format("Average Block Size (MB):", f"{block_size_gb * 1024:.2f}" if all_blocks else "N/A"))
logging.info(dash_line)
logging.info(fmt_final.format("Blocks on Compute:", f"{blocks_on_compute}"))
logging.info(fmt_final.format("Blocks on Swap:", f"{blocks_to_offload}"))
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=""):
"""
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.
"""
logging.info(f"[DisTorch SafeTensor] 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.")
elif hasattr(model_patcher, 'model'):
model_to_patch = model_patcher.model
logging.info("[DisTorch SafeTensor] Found 'model' attribute for patching.")
def get_torch_device_patched():
device = None
if (not (torch.cuda.is_available() or _has_xpu()) or mm.cpu_state == mm.CPUState.CPU or "cpu" in str(current_device).lower()):
device = torch.device("cpu")
else:
logging.error("[DisTorch SafeTensor] Could not find a valid model to patch for block swapping.")
return
devs = set(get_device_list())
device = torch.device(current_device) if str(current_device) in devs else torch.device("cpu")
logging.info(f"[MultiGPU get_torch_device_patched] Returning device: {device} (current_device={current_device})")
return device
all_blocks = []
# 1. Standard UNet Structure
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').")
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.")
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.")
all_blocks.extend(model_to_patch.layers)
# 4. Fallback to top-level ModuleLists
def text_encoder_device_patched():
device = None
if (not (torch.cuda.is_available() or _has_xpu()) or mm.cpu_state == mm.CPUState.CPU or "cpu" in str(current_text_encoder_device).lower()):
device = torch.device("cpu")
else:
logging.info("[DisTorch SafeTensor] No standard structure found. Falling back to 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)
devs = set(get_device_list())
device = torch.device(current_text_encoder_device) if str(current_text_encoder_device) in devs else torch.device("cpu")
logging.info(f"[MultiGPU text_encoder_device_patched] Returning device: {device} (current_text_encoder_device={current_text_encoder_device})")
return device
if not all_blocks:
logging.error("[DisTorch SafeTensor] CRITICAL: No swappable blocks were found in the model. Block swap cannot be applied.")
return
logging.info(f"[DisTorch SafeTensor] Successfully identified {len(all_blocks)} swappable blocks.")
# Run and display the analysis
analyze_safetensor_distorch(model_to_patch, compute_device, swap_device, virtual_vram_gb, 0.0, all_blocks)
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
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)
# 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
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}.")
logging.info("[DisTorch SafeTensor] Block swap setup complete.")
# For backwards compatibility, keep the old name pointing to the new safetensor wrapper
override_class_with_distorch_bs = override_class_with_distorch_safetensor
def override_class_with_distorch_gguf_legacy(cls):
"""Legacy DisTorch wrapper for GGUF models for backward compatibility."""
class NodeOverrideDisTorchLegacy(cls):
@classmethod
def INPUT_TYPES(s):
inputs = copy.deepcopy(cls.INPUT_TYPES())
devices = get_device_list()
default_device = devices[1] if len(devices) > 1 else devices[0]
inputs["optional"] = inputs.get("optional", {})
inputs["optional"]["device"] = (devices, {"default": default_device})
inputs["optional"]["virtual_vram_gb"] = ("FLOAT", {"default": 4.0, "min": 0.0, "max": 24.0, "step": 0.1})
inputs["optional"]["use_other_vram"] = ("BOOLEAN", {"default": False})
inputs["optional"]["expert_mode_allocations"] = ("STRING", {
"multiline": False,
"default": "",
})
return inputs
CATEGORY = "multigpu/legacy"
FUNCTION = "override"
def override(self, *args, device=None, expert_mode_allocations=None, use_other_vram=None, virtual_vram_gb=0.0, **kwargs):
global current_device
if device is not None:
current_device = device
register_patched_ggufmodelpatcher()
fn = getattr(super(), cls.FUNCTION)
out = fn(*args, **kwargs)
vram_string = ""
if virtual_vram_gb > 0:
if use_other_vram:
available_devices = [d for d in get_device_list() if d.startswith(("cuda", "xpu"))]
other_devices = [d for d in available_devices if d != device]
other_devices.sort(key=lambda x: int(x.split(':')[1] if ':' in x else x[-1]), reverse=False)
device_string = ','.join(other_devices + ['cpu'])
vram_string = f"{device};{virtual_vram_gb};{device_string}"
else:
vram_string = f"{device};{virtual_vram_gb};cpu"
full_allocation = f"{expert_mode_allocations}#{vram_string}" if expert_mode_allocations or vram_string else ""
if hasattr(out[0], 'model'):
model_hash = create_model_hash(out[0], "override")
model_allocation_store[model_hash] = full_allocation
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
model_hash = create_model_hash(out[0].patcher, "override")
model_allocation_store[model_hash] = full_allocation
return out
return NodeOverrideDisTorchLegacy
def override_class_with_distorch_gguf(cls):
"""Standardized DisTorch wrapper for GGUF models."""
class NodeOverrideDisTorchGGUF(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_ram_gb"] = ("FLOAT", {"default": 4.0, "min": 0.0, "max": 100.0, "step": 0.1})
inputs["optional"]["donor_device"] = (devices, {"default": "cpu"})
inputs["optional"]["expert_mode_allocations"] = ("STRING", {"multiline": False, "default": ""})
return inputs
CATEGORY = "multigpu"
FUNCTION = "override"
def override(self, *args, compute_device=None, virtual_ram_gb=4.0,
donor_device="cpu", expert_mode_allocations="", **kwargs):
global current_device
if compute_device is not None:
current_device = compute_device
register_patched_ggufmodelpatcher()
fn = getattr(super(), cls.FUNCTION)
out = fn(*args, **kwargs)
vram_string = ""
if virtual_ram_gb > 0:
vram_string = f"{compute_device};{virtual_ram_gb};{donor_device}"
full_allocation = f"{expert_mode_allocations}#{vram_string}" if expert_mode_allocations or vram_string else ""
logging.info(f"[DisTorch GGUF] Full allocation string: {full_allocation}")
if hasattr(out[0], 'model'):
model_hash = create_model_hash(out[0], "override")
model_allocation_store[model_hash] = full_allocation
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
model_hash = create_model_hash(out[0].patcher, "override")
model_allocation_store[model_hash] = full_allocation
return out
return NodeOverrideDisTorchGGUF
# Keep old name for compatibility but point to GGUF version
override_class_with_distorch = override_class_with_distorch_gguf
def override_class_with_distorch_clip(cls):
class NodeOverrideDisTorch(cls):
@classmethod
def INPUT_TYPES(s):
inputs = copy.deepcopy(cls.INPUT_TYPES())
devices = get_device_list()
default_device = devices[1] if len(devices) > 1 else devices[0]
inputs["optional"] = inputs.get("optional", {})
inputs["optional"]["device"] = (devices, {"default": default_device})
inputs["optional"]["virtual_vram_gb"] = ("FLOAT", {"default": 4.0, "min": 0.0, "max": 24.0, "step": 0.1})
inputs["optional"]["use_other_vram"] = ("BOOLEAN", {"default": False})
inputs["optional"]["expert_mode_allocations"] = ("STRING", {
"multiline": False,
"default": "",
"tooltip": "Expert use only: Manual VRAM allocation string. Incorrect values can cause crashes. Do not modify unless you fully understand DisTorch memory management."
})
return inputs
CATEGORY = "multigpu"
FUNCTION = "override"
def override(self, *args, device=None, expert_mode_allocations=None, use_other_vram=None, virtual_vram_gb=0.0, **kwargs):
global current_text_encoder_device
if device is not None:
current_text_encoder_device = device
register_patched_ggufmodelpatcher()
fn = getattr(super(), cls.FUNCTION)
out = fn(*args, **kwargs)
vram_string = ""
if virtual_vram_gb > 0:
if use_other_vram:
available_devices = [d for d in get_device_list() if d.startswith(("cuda", "xpu"))]
other_devices = [d for d in available_devices if d != device]
other_devices.sort(key=lambda x: int(x.split(':')[1] if ':' in x else x[-1]), reverse=False)
device_string = ','.join(other_devices + ['cpu'])
vram_string = f"{device};{virtual_vram_gb};{device_string}"
else:
vram_string = f"{device};{virtual_vram_gb};cpu"
full_allocation = f"{expert_mode_allocations}#{vram_string}" if expert_mode_allocations or vram_string else ""
logging.info(f"[DisTorch] Full allocation string: {full_allocation}")
if hasattr(out[0], 'model'):
model_hash = create_model_hash(out[0], "override")
model_allocation_store[model_hash] = full_allocation
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
model_hash = create_model_hash(out[0].patcher, "override")
model_allocation_store[model_hash] = full_allocation
return out
return NodeOverrideDisTorch
# Apply patches
logging.info(f"[MultiGPU] Patching mm.get_torch_device and mm.text_encoder_device")
logging.info(f"[MultiGPU] Initial current_device: {current_device}")
logging.info(f"[MultiGPU] Initial current_text_encoder_device: {current_text_encoder_device}")
mm.get_torch_device = get_torch_device_patched
mm.text_encoder_device = text_encoder_device_patched
logging.info(f"[MultiGPU] Patches applied successfully")
def check_module_exists(module_path):
full_path = os.path.join(folder_paths.get_folder_paths("custom_nodes")[0], module_path)
@@ -827,12 +133,72 @@ def check_module_exists(module_path):
logging.info(f"MultiGPU: Found {module_path}, creating compatible MultiGPU nodes")
return True
# Import from nodes.py
from .nodes import (
DeviceSelectorMultiGPU,
HunyuanVideoEmbeddingsAdapter,
UnetLoaderGGUF,
UnetLoaderGGUFAdvanced,
CLIPLoaderGGUF,
DualCLIPLoaderGGUF,
TripleCLIPLoaderGGUF,
QuadrupleCLIPLoaderGGUF,
LTXVLoader,
Florence2ModelLoader,
DownloadAndLoadFlorence2Model,
CheckpointLoaderNF4,
LoadFluxControlNet,
MMAudioModelLoader,
MMAudioFeatureUtilsLoader,
MMAudioSampler,
PulidModelLoader,
PulidInsightFaceLoader,
PulidEvaClipLoader,
HyVideoModelLoader,
HyVideoVAELoader,
DownloadAndLoadHyVideoTextEncoder,
)
# Import from wanvideo.py
from .wanvideo import (
WanVideoModelLoader,
WanVideoModelLoader_2,
WanVideoVAELoader,
LoadWanVideoT5TextEncoder,
LoadWanVideoClipTextEncoder,
WanVideoTextEncode,
WanVideoBlockSwap,
WanVideoSampler
)
# Import from distorch.py
from .distorch import (
model_allocation_store,
create_model_hash,
register_patched_ggufmodelpatcher,
analyze_ggml_loading,
calculate_vvram_allocation_string,
override_class_with_distorch_gguf,
override_class_with_distorch_gguf_legacy,
override_class_with_distorch_clip,
override_class_with_distorch
)
# Import from block_swap.py
from .block_swap import (
analyze_safetensor_distorch,
apply_block_swap,
override_class_with_distorch_safetensor,
override_class_with_distorch_bs
)
# Initialize NODE_CLASS_MAPPINGS
NODE_CLASS_MAPPINGS = {
"DeviceSelectorMultiGPU": DeviceSelectorMultiGPU,
"HunyuanVideoEmbeddingsAdapter": HunyuanVideoEmbeddingsAdapter,
}
# Standard MultiGPU nodes
NODE_CLASS_MAPPINGS["UNETLoaderMultiGPU"] = override_class(GLOBAL_NODE_CLASS_MAPPINGS["UNETLoader"])
NODE_CLASS_MAPPINGS["VAELoaderMultiGPU"] = override_class(GLOBAL_NODE_CLASS_MAPPINGS["VAELoader"])
NODE_CLASS_MAPPINGS["CLIPLoaderMultiGPU"] = override_class_clip(GLOBAL_NODE_CLASS_MAPPINGS["CLIPLoader"])
@@ -843,6 +209,13 @@ if "QuadrupleCLIPLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
NODE_CLASS_MAPPINGS["QuadrupleCLIPLoaderMultiGPU"] = override_class_clip(GLOBAL_NODE_CLASS_MAPPINGS["QuadrupleCLIPLoader"])
NODE_CLASS_MAPPINGS["CLIPVisionLoaderMultiGPU"] = override_class_clip(GLOBAL_NODE_CLASS_MAPPINGS["CLIPVisionLoader"])
NODE_CLASS_MAPPINGS["CheckpointLoaderSimpleMultiGPU"] = override_class(GLOBAL_NODE_CLASS_MAPPINGS["CheckpointLoaderSimple"])
NODE_CLASS_MAPPINGS["ControlNetLoaderMultiGPU"] = override_class(GLOBAL_NODE_CLASS_MAPPINGS["ControlNetLoader"])
if "DiffusersLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
NODE_CLASS_MAPPINGS["DiffusersLoaderMultiGPU"] = override_class(GLOBAL_NODE_CLASS_MAPPINGS["DiffusersLoader"])
if "DiffControlNetLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
NODE_CLASS_MAPPINGS["DiffControlNetLoaderMultiGPU"] = override_class(GLOBAL_NODE_CLASS_MAPPINGS["DiffControlNetLoader"])
# DisTorch SafeTensor nodes
NODE_CLASS_MAPPINGS["UNETLoaderDisTorchMultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["UNETLoader"])
NODE_CLASS_MAPPINGS["VAELoaderDisTorchMultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["VAELoader"])
NODE_CLASS_MAPPINGS["CLIPLoaderDisTorchMultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["CLIPLoader"])
@@ -858,30 +231,31 @@ if "DiffusersLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
NODE_CLASS_MAPPINGS["DiffusersLoaderDisTorchMultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["DiffusersLoader"])
if "DiffControlNetLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
NODE_CLASS_MAPPINGS["DiffControlNetLoaderDisTorchMultiGPU"] = override_class_with_distorch_safetensor(GLOBAL_NODE_CLASS_MAPPINGS["DiffControlNetLoader"])
NODE_CLASS_MAPPINGS["ControlNetLoaderMultiGPU"] = override_class(GLOBAL_NODE_CLASS_MAPPINGS["ControlNetLoader"])
if "DiffusersLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
NODE_CLASS_MAPPINGS["DiffusersLoaderMultiGPU"] = override_class(GLOBAL_NODE_CLASS_MAPPINGS["DiffusersLoader"])
if "DiffControlNetLoader" in GLOBAL_NODE_CLASS_MAPPINGS:
NODE_CLASS_MAPPINGS["DiffControlNetLoaderMultiGPU"] = override_class(GLOBAL_NODE_CLASS_MAPPINGS["DiffControlNetLoader"])
# ComfyUI-LTXVideo
if check_module_exists("ComfyUI-LTXVideo") or check_module_exists("comfyui-ltxvideo"):
NODE_CLASS_MAPPINGS["LTXVLoaderMultiGPU"] = override_class(LTXVLoader)
# ComfyUI-Florence2
if check_module_exists("ComfyUI-Florence2") or check_module_exists("comfyui-florence2"):
NODE_CLASS_MAPPINGS["Florence2ModelLoaderMultiGPU"] = override_class(Florence2ModelLoader)
NODE_CLASS_MAPPINGS["DownloadAndLoadFlorence2ModelMultiGPU"] = override_class(DownloadAndLoadFlorence2Model)
# ComfyUI_bitsandbytes_NF4
if check_module_exists("ComfyUI_bitsandbytes_NF4") or check_module_exists("comfyui_bitsandbytes_nf4"):
NODE_CLASS_MAPPINGS["CheckpointLoaderNF4MultiGPU"] = override_class(CheckpointLoaderNF4)
# x-flux-comfyui
if check_module_exists("x-flux-comfyui") or check_module_exists("x-flux-comfyui"):
NODE_CLASS_MAPPINGS["LoadFluxControlNetMultiGPU"] = override_class(LoadFluxControlNet)
# ComfyUI-MMAudio
if check_module_exists("ComfyUI-MMAudio") or check_module_exists("comfyui-mmaudio"):
NODE_CLASS_MAPPINGS["MMAudioModelLoaderMultiGPU"] = override_class(MMAudioModelLoader)
NODE_CLASS_MAPPINGS["MMAudioFeatureUtilsLoaderMultiGPU"] = override_class(MMAudioFeatureUtilsLoader)
NODE_CLASS_MAPPINGS["MMAudioSamplerMultiGPU"] = override_class(MMAudioSampler)
# ComfyUI-GGUF
if check_module_exists("ComfyUI-GGUF") or check_module_exists("comfyui-gguf"):
NODE_CLASS_MAPPINGS["UnetLoaderGGUFMultiGPU"] = override_class(UnetLoaderGGUF)
NODE_CLASS_MAPPINGS["UnetLoaderGGUFDisTorchMultiGPU"] = override_class_with_distorch_gguf(UnetLoaderGGUF)
@@ -898,18 +272,20 @@ if check_module_exists("ComfyUI-GGUF") or check_module_exists("comfyui-gguf"):
NODE_CLASS_MAPPINGS["QuadrupleCLIPLoaderGGUFMultiGPU"] = override_class_clip(QuadrupleCLIPLoaderGGUF)
NODE_CLASS_MAPPINGS["QuadrupleCLIPLoaderGGUFDisTorchMultiGPU"] = override_class_with_distorch_clip(QuadrupleCLIPLoaderGGUF)
# PuLID_ComfyUI
if check_module_exists("PuLID_ComfyUI") or check_module_exists("pulid_comfyui"):
NODE_CLASS_MAPPINGS["PulidModelLoaderMultiGPU"] = override_class(PulidModelLoader)
NODE_CLASS_MAPPINGS["PulidInsightFaceLoaderMultiGPU"] = override_class(PulidInsightFaceLoader)
NODE_CLASS_MAPPINGS["PulidEvaClipLoaderMultiGPU"] = override_class(PulidEvaClipLoader)
# ComfyUI-HunyuanVideoWrapper
if check_module_exists("ComfyUI-HunyuanVideoWrapper") or check_module_exists("comfyui-hunyuanvideowrapper"):
NODE_CLASS_MAPPINGS["HyVideoModelLoaderMultiGPU"] = override_class(HyVideoModelLoader)
NODE_CLASS_MAPPINGS["HyVideoVAELoaderMultiGPU"] = override_class(HyVideoVAELoader)
NODE_CLASS_MAPPINGS["DownloadAndLoadHyVideoTextEncoderMultiGPU"] = override_class(DownloadAndLoadHyVideoTextEncoder)
# ComfyUI-WanVideoWrapper
if check_module_exists("ComfyUI-WanVideoWrapper") or check_module_exists("comfyui-wanvideowrapper"):
# WanVideo uses custom implementation, not the standard override
NODE_CLASS_MAPPINGS["WanVideoModelLoaderMultiGPU"] = WanVideoModelLoader
NODE_CLASS_MAPPINGS["WanVideoModelLoaderMultiGPU_2"] = WanVideoModelLoader_2
NODE_CLASS_MAPPINGS["WanVideoVAELoaderMultiGPU"] = WanVideoVAELoader
@@ -919,5 +295,4 @@ if check_module_exists("ComfyUI-WanVideoWrapper") or check_module_exists("comfyu
NODE_CLASS_MAPPINGS["WanVideoBlockSwapMultiGPU"] = WanVideoBlockSwap
NODE_CLASS_MAPPINGS["WanVideoSamplerMultiGPU"] = WanVideoSampler
logging.info(f"MultiGPU: Registration complete. Final mappings: {', '.join(NODE_CLASS_MAPPINGS.keys())}")
+212
View File
@@ -0,0 +1,212 @@
"""
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
def analyze_safetensor_distorch(model, compute_device, swap_device, virtual_vram_gb, reserved_swap_gb, all_blocks):
"""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)
# Device Allocation Table
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}", f"Reserve: {reserved_swap_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 Analysis Table
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)
# Final Assignment Table
model_size_gb = total_memory / (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
logging.info(" DisTorch Final Block Assignments")
logging.info(dash_line)
fmt_final = "{:<20}{:>15}"
logging.info(fmt_final.format("Total Model Size (GB):", f"{model_size_gb:.2f}"))
logging.info(fmt_final.format("Average Block Size (MB):", f"{block_size_gb * 1024:.2f}" if all_blocks else "N/A"))
logging.info(dash_line)
logging.info(fmt_final.format("Blocks on Compute:", f"{blocks_on_compute}"))
logging.info(fmt_final.format("Blocks on Swap:", f"{blocks_to_offload}"))
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=""):
"""
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.
"""
logging.info(f"[DisTorch SafeTensor] 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.")
elif hasattr(model_patcher, 'model'):
model_to_patch = model_patcher.model
logging.info("[DisTorch SafeTensor] Found 'model' attribute for patching.")
else:
logging.error("[DisTorch SafeTensor] Could not find a valid model to patch for block swapping.")
return
all_blocks = []
# 1. Standard UNet Structure
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').")
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.")
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.")
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.")
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.")
return
logging.info(f"[DisTorch SafeTensor] Successfully identified {len(all_blocks)} swappable blocks.")
# Run and display the analysis
analyze_safetensor_distorch(model_to_patch, compute_device, swap_device, virtual_vram_gb, 0.0, all_blocks)
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
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)
# 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
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}.")
logging.info("[DisTorch SafeTensor] Block swap setup complete.")
def override_class_with_distorch_safetensor(cls):
"""DisTorch wrapper for SafeTensor models, providing block-swap memory optimization."""
from .nodes import get_device_list
class NodeOverrideDisTorchSafeTensor(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_ram_gb"] = ("FLOAT", {"default": 4.0, "min": 0.0, "max": 100.0, "step": 0.1})
inputs["optional"]["donor_device"] = (devices, {"default": "cpu"})
inputs["optional"]["expert_mode_allocations"] = ("STRING", {"multiline": False, "default": ""})
return inputs
CATEGORY = "multigpu"
FUNCTION = "override"
def override(self, *args, compute_device=None, virtual_ram_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_ram_gb={virtual_ram_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_ram_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 NodeOverrideDisTorchSafeTensor
# For backwards compatibility, keep the old name pointing to the new safetensor wrapper
override_class_with_distorch_bs = override_class_with_distorch_safetensor
+464
View File
@@ -0,0 +1,464 @@
"""
DisTorch GGUF/GGML Memory Management Module
Contains all GGUF/GGML related code for distributed memory management
"""
import sys
import torch
import logging
import hashlib
import copy
from collections import defaultdict
import comfy.model_management as mm
# Global store for model allocations
model_allocation_store = {}
def create_model_hash(model, caller):
"""Create a unique hash for a model to track allocations"""
model_type = type(model.model).__name__
model_size = model.model_size()
first_layers = str(list(model.model_state_dict().keys())[:3])
identifier = f"{model_type}_{model_size}_{first_layers}"
final_hash = hashlib.sha256(identifier.encode()).hexdigest()
return final_hash
def register_patched_ggufmodelpatcher():
"""Register and patch the GGUFModelPatcher for distributed loading"""
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["UnetLoaderGGUF"]
module = sys.modules[original_loader.__module__]
if not hasattr(module.GGUFModelPatcher, '_patched'):
original_load = module.GGUFModelPatcher.load
def new_load(self, *args, force_patch_weights=False, **kwargs):
global model_allocation_store
super(module.GGUFModelPatcher, self).load(*args, force_patch_weights=True, **kwargs)
debug_hash = create_model_hash(self, "patcher")
linked = []
module_count = 0
for n, m in self.model.named_modules():
module_count += 1
if hasattr(m, "weight"):
device = getattr(m.weight, "device", None)
if device is not None:
linked.append((n, m))
continue
if hasattr(m, "bias"):
device = getattr(m.bias, "device", None)
if device is not None:
linked.append((n, m))
continue
if linked:
if hasattr(self, 'model'):
debug_hash = create_model_hash(self, "patcher")
debug_allocations = model_allocation_store.get(debug_hash)
if debug_allocations:
device_assignments = analyze_ggml_loading(self.model, debug_allocations)['device_assignments']
for device, layers in device_assignments.items():
target_device = torch.device(device)
for n, m, _ in layers:
m.to(self.load_device).to(target_device)
self.mmap_released = True
module.GGUFModelPatcher.load = new_load
module.GGUFModelPatcher._patched = True
def analyze_ggml_loading(model, allocations_str):
"""Analyze and distribute GGML model layers across devices"""
DEVICE_RATIOS_DISTORCH = {}
device_table = {}
distorch_alloc = allocations_str
virtual_vram_gb = 0.0
if '#' in allocations_str:
distorch_alloc, virtual_vram_str = allocations_str.split('#')
if not distorch_alloc:
distorch_alloc = calculate_vvram_allocation_string(model, virtual_vram_str)
eq_line = "=" * 47
dash_line = "-" * 47
fmt_assign = "{:<12}{:>10}{:>14}{:>10}"
for allocation in distorch_alloc.split(';'):
dev_name, fraction = allocation.split(',')
fraction = float(fraction)
total_mem_bytes = mm.get_total_memory(torch.device(dev_name))
alloc_gb = (total_mem_bytes * fraction) / (1024**3)
DEVICE_RATIOS_DISTORCH[dev_name] = alloc_gb
device_table[dev_name] = {
"fraction": fraction,
"total_gb": total_mem_bytes / (1024**3),
"alloc_gb": alloc_gb
}
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
logging.info(eq_line)
logging.info(" DisTorch Device Allocations")
logging.info(eq_line)
logging.info(fmt_assign.format("Device", "Alloc %", "Total (GB)", " Alloc (GB)"))
logging.info(dash_line)
sorted_devices = sorted(device_table.keys(), key=lambda d: (d == "cpu", d))
for dev in sorted_devices:
frac = device_table[dev]["fraction"]
tot_gb = device_table[dev]["total_gb"]
alloc_gb = device_table[dev]["alloc_gb"]
logging.info(fmt_assign.format(dev,f"{int(frac * 100)}%",f"{tot_gb:.2f}",f"{alloc_gb:.2f}"))
logging.info(dash_line)
layer_summary = {}
layer_list = []
memory_by_type = defaultdict(int)
total_memory = 0
for name, module in model.named_modules():
if hasattr(module, "weight"):
layer_type = type(module).__name__
layer_summary[layer_type] = layer_summary.get(layer_type, 0) + 1
layer_list.append((name, module, layer_type))
layer_memory = 0
if module.weight is not None:
layer_memory += module.weight.numel() * module.weight.element_size()
if hasattr(module, "bias") and module.bias is not None:
layer_memory += module.bias.numel() * module.bias.element_size()
memory_by_type[layer_type] += layer_memory
total_memory += layer_memory
logging.info(" DisTorch GGML Layer Distribution")
logging.info(dash_line)
fmt_layer = "{:<12}{:>10}{:>14}{:>10}"
logging.info(fmt_layer.format("Layer Type", "Layers", "Memory (MB)", "% Total"))
logging.info(dash_line)
for layer_type, count in layer_summary.items():
mem_mb = memory_by_type[layer_type] / (1024 * 1024)
mem_percent = (memory_by_type[layer_type] / total_memory) * 100 if total_memory > 0 else 0
logging.info(fmt_layer.format(layer_type,str(count),f"{mem_mb:.2f}",f"{mem_percent:.1f}%"))
logging.info(dash_line)
nonzero_devices = [d for d, r in DEVICE_RATIOS_DISTORCH.items() if r > 0]
nonzero_total_ratio = sum(DEVICE_RATIOS_DISTORCH[d] for d in nonzero_devices)
device_assignments = {device: [] for device in DEVICE_RATIOS_DISTORCH.keys()}
total_layers = len(layer_list)
current_layer = 0
for idx, device in enumerate(nonzero_devices):
ratio = DEVICE_RATIOS_DISTORCH[device]
if idx == len(nonzero_devices) - 1:
device_layer_count = total_layers - current_layer
else:
device_layer_count = int((ratio / nonzero_total_ratio) * total_layers)
start_idx = current_layer
end_idx = current_layer + device_layer_count
device_assignments[device] = layer_list[start_idx:end_idx]
current_layer += device_layer_count
logging.info(" DisTorch Final Device/Layer Assignments")
logging.info(dash_line)
fmt_assign = "{:<12}{:>10}{:>14}{:>10}"
logging.info(fmt_assign.format("Device", "Layers", "Memory (MB)", "% Total"))
logging.info(dash_line)
total_assigned_memory = 0
device_memories = {}
for device, layers in device_assignments.items():
device_memory = 0
for layer_type in layer_summary:
type_layers = sum(1 for _, _, lt in layers if lt == layer_type)
if layer_summary[layer_type] > 0:
mem_per_layer = memory_by_type[layer_type] / layer_summary[layer_type]
device_memory += mem_per_layer * type_layers
device_memories[device] = device_memory
total_assigned_memory += device_memory
sorted_assignments = sorted(device_assignments.keys(), key=lambda d: (d == "cpu", d))
for dev in sorted_assignments:
layers = device_assignments[dev]
mem_mb = device_memories[dev] / (1024 * 1024)
mem_percent = (device_memories[dev] / total_memory) * 100 if total_memory > 0 else 0
logging.info(fmt_assign.format(dev,str(len(layers)),f"{mem_mb:.2f}",f"{mem_percent:.1f}%"))
logging.info(dash_line)
return {"device_assignments": device_assignments}
def calculate_vvram_allocation_string(model, virtual_vram_str):
"""Calculate virtual VRAM allocation string for distributed loading"""
recipient_device, vram_amount, donors = virtual_vram_str.split(';')
virtual_vram_gb = float(vram_amount)
eq_line = "=" * 47
dash_line = "-" * 47
fmt_assign = "{:<8} {:<6} {:>11} {:>9} {:>9}"
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
logging.info(eq_line)
logging.info(" DisTorch Virtual VRAM Analysis")
logging.info(eq_line)
logging.info(fmt_assign.format("Object", "Role", "Original(GB)", "Total(GB)", "Virt(GB)"))
logging.info(dash_line)
recipient_vram = mm.get_total_memory(torch.device(recipient_device)) / (1024**3)
recipient_virtual = recipient_vram + virtual_vram_gb
logging.info(fmt_assign.format(recipient_device, 'recip', f"{recipient_vram:.2f}GB",f"{recipient_virtual:.2f}GB", f"+{virtual_vram_gb:.2f}GB"))
ram_donors = [d for d in donors.split(',') if d != 'cpu']
remaining_vram_needed = virtual_vram_gb
donor_device_info = {}
donor_allocations = {}
for donor in ram_donors:
donor_vram = mm.get_total_memory(torch.device(donor)) / (1024**3)
max_donor_capacity = donor_vram * 0.9
donation = min(remaining_vram_needed, max_donor_capacity)
donor_virtual = donor_vram - donation
remaining_vram_needed -= donation
donor_allocations[donor] = donation
donor_device_info[donor] = (donor_vram, donor_virtual)
logging.info(fmt_assign.format(donor, 'donor', f"{donor_vram:.2f}GB", f"{donor_virtual:.2f}GB", f"-{donation:.2f}GB"))
system_dram_gb = mm.get_total_memory(torch.device('cpu')) / (1024**3)
cpu_donation = remaining_vram_needed
cpu_virtual = system_dram_gb - cpu_donation
donor_allocations['cpu'] = cpu_donation
logging.info(fmt_assign.format('cpu', 'donor', f"{system_dram_gb:.2f}GB", f"{cpu_virtual:.2f}GB", f"-{cpu_donation:.2f}GB"))
logging.info(dash_line)
layer_summary = {}
layer_list = []
memory_by_type = defaultdict(int)
total_memory = 0
for name, module in model.named_modules():
if hasattr(module, "weight"):
layer_type = type(module).__name__
layer_summary[layer_type] = layer_summary.get(layer_type, 0) + 1
layer_list.append((name, module, layer_type))
layer_memory = 0
if module.weight is not None:
layer_memory += module.weight.numel() * module.weight.element_size()
if hasattr(module, "bias") and module.bias is not None:
layer_memory += module.bias.numel() * module.bias.element_size()
memory_by_type[layer_type] += layer_memory
total_memory += layer_memory
model_size_gb = total_memory / (1024**3)
new_model_size_gb = max(0, model_size_gb - virtual_vram_gb)
logging.info(fmt_assign.format('model', 'model', f"{model_size_gb:.2f}GB",f"{new_model_size_gb:.2f}GB", f"-{virtual_vram_gb:.2f}GB"))
if model_size_gb > (recipient_vram * 0.9):
on_recipient = recipient_vram * 0.9
on_virtuals = model_size_gb - on_recipient
logging.info(f"\nWarning: Model size is greater than 90% of recipient VRAM. {on_virtuals:.2f} GB of GGML Layers Offloaded Automatically to Virtual VRAM.\n")
else:
on_recipient = model_size_gb
on_virtuals = 0
new_on_recipient = max(0, on_recipient - virtual_vram_gb)
allocation_parts = []
recipient_percent = new_on_recipient / recipient_vram
allocation_parts.append(f"{recipient_device},{recipient_percent:.4f}")
for donor in ram_donors:
donor_vram = donor_device_info[donor][0]
donor_percent = donor_allocations[donor] / donor_vram
allocation_parts.append(f"{donor},{donor_percent:.4f}")
cpu_percent = donor_allocations['cpu'] / system_dram_gb
allocation_parts.append(f"cpu,{cpu_percent:.4f}")
allocation_string = ";".join(allocation_parts)
fmt_mem = "{:<20}{:>20}"
logging.info(fmt_mem.format("\nAllocation String", allocation_string))
return allocation_string
def override_class_with_distorch_gguf(cls):
"""Standardized DisTorch wrapper for GGUF models."""
from .nodes import get_device_list
from . import current_device
class NodeOverrideDisTorchGGUF(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_ram_gb"] = ("FLOAT", {"default": 4.0, "min": 0.0, "max": 100.0, "step": 0.1})
inputs["optional"]["donor_device"] = (devices, {"default": "cpu"})
inputs["optional"]["expert_mode_allocations"] = ("STRING", {"multiline": False, "default": ""})
return inputs
CATEGORY = "multigpu"
FUNCTION = "override"
def override(self, *args, compute_device=None, virtual_ram_gb=4.0,
donor_device="cpu", expert_mode_allocations="", **kwargs):
from . import set_current_device
if compute_device is not None:
set_current_device(compute_device)
register_patched_ggufmodelpatcher()
fn = getattr(super(), cls.FUNCTION)
out = fn(*args, **kwargs)
vram_string = ""
if virtual_ram_gb > 0:
vram_string = f"{compute_device};{virtual_ram_gb};{donor_device}"
full_allocation = f"{expert_mode_allocations}#{vram_string}" if expert_mode_allocations or vram_string else ""
logging.info(f"[DisTorch GGUF] Full allocation string: {full_allocation}")
if hasattr(out[0], 'model'):
model_hash = create_model_hash(out[0], "override")
model_allocation_store[model_hash] = full_allocation
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
model_hash = create_model_hash(out[0].patcher, "override")
model_allocation_store[model_hash] = full_allocation
return out
return NodeOverrideDisTorchGGUF
def override_class_with_distorch_gguf_legacy(cls):
"""Legacy DisTorch wrapper for GGUF models for backward compatibility."""
from .nodes import get_device_list
from . import current_device
class NodeOverrideDisTorchLegacy(cls):
@classmethod
def INPUT_TYPES(s):
inputs = copy.deepcopy(cls.INPUT_TYPES())
devices = get_device_list()
default_device = devices[1] if len(devices) > 1 else devices[0]
inputs["optional"] = inputs.get("optional", {})
inputs["optional"]["device"] = (devices, {"default": default_device})
inputs["optional"]["virtual_vram_gb"] = ("FLOAT", {"default": 4.0, "min": 0.0, "max": 24.0, "step": 0.1})
inputs["optional"]["use_other_vram"] = ("BOOLEAN", {"default": False})
inputs["optional"]["expert_mode_allocations"] = ("STRING", {
"multiline": False,
"default": "",
})
return inputs
CATEGORY = "multigpu/legacy"
FUNCTION = "override"
def override(self, *args, device=None, expert_mode_allocations=None, use_other_vram=None, virtual_vram_gb=0.0, **kwargs):
from . import set_current_device
if device is not None:
set_current_device(device)
register_patched_ggufmodelpatcher()
fn = getattr(super(), cls.FUNCTION)
out = fn(*args, **kwargs)
vram_string = ""
if virtual_vram_gb > 0:
if use_other_vram:
available_devices = [d for d in get_device_list() if d.startswith(("cuda", "xpu"))]
other_devices = [d for d in available_devices if d != device]
other_devices.sort(key=lambda x: int(x.split(':')[1] if ':' in x else x[-1]), reverse=False)
device_string = ','.join(other_devices + ['cpu'])
vram_string = f"{device};{virtual_vram_gb};{device_string}"
else:
vram_string = f"{device};{virtual_vram_gb};cpu"
full_allocation = f"{expert_mode_allocations}#{vram_string}" if expert_mode_allocations or vram_string else ""
if hasattr(out[0], 'model'):
model_hash = create_model_hash(out[0], "override")
model_allocation_store[model_hash] = full_allocation
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
model_hash = create_model_hash(out[0].patcher, "override")
model_allocation_store[model_hash] = full_allocation
return out
return NodeOverrideDisTorchLegacy
def override_class_with_distorch_clip(cls):
"""DisTorch wrapper for CLIP models with GGUF support"""
from .nodes import get_device_list
from . import current_text_encoder_device
class NodeOverrideDisTorch(cls):
@classmethod
def INPUT_TYPES(s):
inputs = copy.deepcopy(cls.INPUT_TYPES())
devices = get_device_list()
default_device = devices[1] if len(devices) > 1 else devices[0]
inputs["optional"] = inputs.get("optional", {})
inputs["optional"]["device"] = (devices, {"default": default_device})
inputs["optional"]["virtual_vram_gb"] = ("FLOAT", {"default": 4.0, "min": 0.0, "max": 24.0, "step": 0.1})
inputs["optional"]["use_other_vram"] = ("BOOLEAN", {"default": False})
inputs["optional"]["expert_mode_allocations"] = ("STRING", {
"multiline": False,
"default": "",
"tooltip": "Expert use only: Manual VRAM allocation string. Incorrect values can cause crashes. Do not modify unless you fully understand DisTorch memory management."
})
return inputs
CATEGORY = "multigpu"
FUNCTION = "override"
def override(self, *args, device=None, expert_mode_allocations=None, use_other_vram=None, virtual_vram_gb=0.0, **kwargs):
from . import set_current_text_encoder_device
if device is not None:
set_current_text_encoder_device(device)
register_patched_ggufmodelpatcher()
fn = getattr(super(), cls.FUNCTION)
out = fn(*args, **kwargs)
vram_string = ""
if virtual_vram_gb > 0:
if use_other_vram:
available_devices = [d for d in get_device_list() if d.startswith(("cuda", "xpu"))]
other_devices = [d for d in available_devices if d != device]
other_devices.sort(key=lambda x: int(x.split(':')[1] if ':' in x else x[-1]), reverse=False)
device_string = ','.join(other_devices + ['cpu'])
vram_string = f"{device};{virtual_vram_gb};{device_string}"
else:
vram_string = f"{device};{virtual_vram_gb};cpu"
full_allocation = f"{expert_mode_allocations}#{vram_string}" if expert_mode_allocations or vram_string else ""
logging.info(f"[DisTorch] Full allocation string: {full_allocation}")
if hasattr(out[0], 'model'):
model_hash = create_model_hash(out[0], "override")
model_allocation_store[model_hash] = full_allocation
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
model_hash = create_model_hash(out[0].patcher, "override")
model_allocation_store[model_hash] = full_allocation
return out
return NodeOverrideDisTorch
# Alias for backward compatibility
override_class_with_distorch = override_class_with_distorch_gguf
+73 -519
View File
@@ -1,7 +1,80 @@
import torch
import folder_paths
from pathlib import Path
from nodes import NODE_CLASS_MAPPINGS
def _has_xpu():
try:
return hasattr(torch, "xpu") and hasattr(torch.xpu, "is_available") and torch.xpu.is_available()
except Exception:
return False
def get_device_list():
devs = ["cpu"]
try:
if hasattr(torch, "cuda") and hasattr(torch.cuda, "is_available") and torch.cuda.is_available():
devs += [f"cuda:{i}" for i in range(torch.cuda.device_count())]
except Exception:
pass
try:
if _has_xpu():
devs += [f"xpu:{i}" for i in range(torch.xpu.device_count())]
except Exception:
pass
return devs
class DeviceSelectorMultiGPU:
@classmethod
def INPUT_TYPES(s):
devices = get_device_list()
return {
"required": {
"device": (devices, {"default": devices[1] if len(devices) > 1 else devices[0]})
}
}
RETURN_TYPES = (get_device_list(),)
RETURN_NAMES = ("device",)
FUNCTION = "select_device"
CATEGORY = "multigpu"
def select_device(self, device):
return (device,)
class HunyuanVideoEmbeddingsAdapter:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"hyvid_embeds": ("HYVIDEMBEDS",),
}
}
RETURN_TYPES = ("CONDITIONING",)
FUNCTION = "adapt_embeddings"
CATEGORY = "multigpu"
def adapt_embeddings(self, hyvid_embeds):
cond = hyvid_embeds["prompt_embeds"]
pooled_dict = {
"pooled_output": hyvid_embeds["prompt_embeds_2"],
"cross_attn": hyvid_embeds["prompt_embeds"],
"attention_mask": hyvid_embeds["attention_mask"],
}
if hyvid_embeds["attention_mask_2"] is not None:
pooled_dict["attention_mask_controlnet"] = hyvid_embeds["attention_mask_2"]
if hyvid_embeds["cfg"] is not None:
pooled_dict["guidance"] = float(hyvid_embeds["cfg"])
pooled_dict["start_percent"] = float(hyvid_embeds["start_percent"]) if hyvid_embeds["start_percent"] is not None else 0.0
pooled_dict["end_percent"] = float(hyvid_embeds["end_percent"]) if hyvid_embeds["end_percent"] is not None else 1.0
return ([[cond, pooled_dict]],)
class UnetLoaderGGUF:
@classmethod
def INPUT_TYPES(s):
@@ -18,7 +91,6 @@ class UnetLoaderGGUF:
TITLE = "Unet Loader (GGUF)"
def load_unet(self, unet_name, dequant_dtype=None, patch_dtype=None, patch_on_device=None):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["UnetLoaderGGUF"]()
return original_loader.load_unet(unet_name, dequant_dtype, patch_dtype, patch_on_device)
@@ -62,17 +134,14 @@ class CLIPLoaderGGUF:
return sorted(files)
def load_data(self, ckpt_paths):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["CLIPLoaderGGUF"]()
return original_loader.load_data(ckpt_paths)
def load_patcher(self, clip_paths, clip_type, clip_data):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["CLIPLoaderGGUF"]()
return original_loader.load_patcher(clip_paths, clip_type, clip_data)
def load_clip(self, clip_name, type="stable_diffusion"):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["CLIPLoaderGGUF"]()
return original_loader.load_clip(clip_name, type)
@@ -93,7 +162,6 @@ class DualCLIPLoaderGGUF(CLIPLoaderGGUF):
TITLE = "DualCLIPLoader (GGUF)"
def load_clip(self, clip_name1, clip_name2, type):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["DualCLIPLoaderGGUF"]()
clip = original_loader.load_clip(clip_name1, clip_name2, type)
clip[0].patcher.load(force_patch_weights=True)
@@ -115,7 +183,6 @@ class TripleCLIPLoaderGGUF(CLIPLoaderGGUF):
TITLE = "TripleCLIPLoader (GGUF)"
def load_clip(self, clip_name1, clip_name2, clip_name3, type="sd3"):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["TripleCLIPLoaderGGUF"]()
return original_loader.load_clip(clip_name1, clip_name2, clip_name3, type)
@@ -135,7 +202,6 @@ class QuadrupleCLIPLoaderGGUF(CLIPLoaderGGUF):
TITLE = "QuadrupleCLIPLoader (GGUF)"
def load_clip(self, clip_name1, clip_name2, clip_name3, clip_name4, type="stable_diffusion"):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["QuadrupleCLIPLoaderGGUF"]()
return original_loader.load_clip(clip_name1, clip_name2, clip_name3, clip_name4, type)
@@ -159,15 +225,12 @@ class LTXVLoader:
OUTPUT_NODE = False
def load(self, ckpt_name, dtype):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["LTXVLoader"]()
return original_loader.load(ckpt_name, dtype)
def _load_unet(self, load_device, offload_device, weights, num_latent_channels, dtype, config=None ):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["LTXVLoader"]()
return original_loader._load_unet(load_device, offload_device, weights, num_latent_channels, dtype, config=None )
def _load_vae(self, weights, config=None):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["LTXVLoader"]()
return original_loader._load_vae(weights, config=None)
@@ -194,7 +257,6 @@ class Florence2ModelLoader:
CATEGORY = "Florence2"
def loadmodel(self, model, precision, attention, lora=None):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["Florence2ModelLoader"]()
return original_loader.loadmodel(model, precision, attention, lora)
@@ -242,7 +304,6 @@ class DownloadAndLoadFlorence2Model:
CATEGORY = "Florence2"
def loadmodel(self, model, precision, attention, lora=None):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["DownloadAndLoadFlorence2Model"]()
return original_loader.loadmodel(model, precision, attention, lora)
@@ -258,7 +319,6 @@ class CheckpointLoaderNF4:
def load_checkpoint(self, ckpt_name):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["CheckpointLoaderNF4"]()
return original_loader.load_checkpoint(ckpt_name)
@@ -275,7 +335,6 @@ class LoadFluxControlNet:
CATEGORY = "XLabsNodes"
def loadmodel(self, model_name, controlnet_path):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["LoadFluxControlNet"]()
return original_loader.loadmodel(model_name, controlnet_path)
@@ -296,7 +355,6 @@ class MMAudioModelLoader:
CATEGORY = "MMAudio"
def loadmodel(self, mmaudio_model, base_precision):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["MMAudioModelLoader"]()
return original_loader.loadmodel(mmaudio_model, base_precision)
@@ -324,7 +382,6 @@ class MMAudioFeatureUtilsLoader:
CATEGORY = "MMAudio"
def loadmodel(self, vae_model, precision, synchformer_model, clip_model, mode, bigvgan_vocoder_model=None):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["MMAudioFeatureUtilsLoader"]()
return original_loader.loadmodel(vae_model, precision, synchformer_model, clip_model, mode, bigvgan_vocoder_model)
@@ -355,7 +412,6 @@ class MMAudioSampler:
CATEGORY = "MMAudio"
def sample(self, mmaudio_model, seed, feature_utils, duration, steps, cfg, prompt, negative_prompt, mask_away_clip, force_offload, images=None):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["MMAudioSampler"]()
return original_loader.sample(mmaudio_model, seed, feature_utils, duration, steps, cfg, prompt, negative_prompt, mask_away_clip, force_offload, images)
@@ -369,7 +425,6 @@ class PulidModelLoader:
CATEGORY = "pulid"
def load_model(self, pulid_file):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["PulidModelLoader"]()
return original_loader.load_model(pulid_file)
@@ -387,7 +442,6 @@ class PulidInsightFaceLoader:
CATEGORY = "pulid"
def load_insightface(self, provider):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["PulidInsightFaceLoader"]()
return original_loader.load_insightface(provider)
@@ -403,7 +457,6 @@ class PulidEvaClipLoader:
CATEGORY = "pulid"
def load_eva_clip(self):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["PulidEvaClipLoader"]()
return original_loader.load_eva_clip()
@@ -438,7 +491,6 @@ class HyVideoModelLoader:
CATEGORY = "HunyuanVideoWrapper"
def loadmodel(self, model, base_precision, load_device, quantization, compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None, auto_cpu_offload=False):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["HyVideoModelLoader"]()
return original_loader.loadmodel(model, base_precision, load_device, quantization, compile_args, attention_mode, block_swap_args, lora, auto_cpu_offload)
@@ -464,7 +516,6 @@ class HyVideoVAELoader:
DESCRIPTION = "Loads Hunyuan VAE model from 'ComfyUI/models/vae'"
def loadmodel(self, model_name, precision, compile_args=None):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["HyVideoVAELoader"]()
return original_loader.loadmodel(model_name, precision, compile_args)
@@ -493,502 +544,5 @@ class DownloadAndLoadHyVideoTextEncoder:
DESCRIPTION = "Loads Hunyuan text_encoder model from 'ComfyUI/models/LLM'"
def loadmodel(self, llm_model, clip_model, precision, apply_final_norm=False, hidden_state_skip_layer=2, quantization="disabled"):
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["DownloadAndLoadHyVideoTextEncoder"]()
return original_loader.loadmodel(llm_model, clip_model, precision, apply_final_norm, hidden_state_skip_layer, quantization)
class WanVideoModelLoader:
@classmethod
def INPUT_TYPES(s):
from . import get_device_list
devices = get_device_list()
return {
"required": {
"model": (folder_paths.get_filename_list("unet_gguf") + folder_paths.get_filename_list("diffusion_models"),
{"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' folder",}),
"base_precision": (["fp32", "bf16", "fp16", "fp16_fast"], {"default": "bf16"}),
"quantization": (
["disabled", "fp8_e4m3fn", "fp8_e4m3fn_fast", "fp8_e5m2", "fp8_e4m3fn_fast_no_ffn", "fp8_e4m3fn_scaled", "fp8_e5m2_scaled"],
{"default": "disabled", "tooltip": "optional quantization method"}
),
"device": (devices, {"default": devices[1] if len(devices) > 1 else devices[0], "tooltip": "Device to load the model to"}),
},
"optional": {
"attention_mode": ([
"sdpa",
"flash_attn_2",
"flash_attn_3",
"sageattn",
"sageattn_3",
"flex_attention",
"radial_sage_attention",
], {"default": "sdpa"}),
"compile_args": ("WANCOMPILEARGS", ),
"block_swap_args": ("BLOCKSWAPARGS", ),
"lora": ("WANVIDLORA", {"default": None}),
"vram_management_args": ("VRAM_MANAGEMENTARGS", {"default": None, "tooltip": "Alternative offloading method from DiffSynth-Studio, more aggressive in reducing memory use than block swapping, but can be slower"}),
"vace_model": ("VACEPATH", {"default": None, "tooltip": "VACE model to use when not using model that has it included"}),
"fantasytalking_model": ("FANTASYTALKINGMODEL", {"default": None, "tooltip": "FantasyTalking model https://github.com/Fantasy-AMAP"}),
"multitalk_model": ("MULTITALKMODEL", {"default": None, "tooltip": "Multitalk model"}),
}
}
RETURN_TYPES = ("WANVIDEOMODEL",)
RETURN_NAMES = ("model", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
def loadmodel(self, model, base_precision, device, quantization,
compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None, vram_management_args=None, vace_model=None, fantasytalking_model=None, multitalk_model=None):
import logging
import comfy.model_management as mm
import torch
logging.info(f"[MultiGPU WanVideoModelLoader] ========== CUSTOM IMPLEMENTATION ==========")
logging.info(f"[MultiGPU WanVideoModelLoader] User selected device: {device}")
selected_device = torch.device(device)
# Determine load_device for original loader
load_device = "offload_device" if device == "cpu" else "main_device"
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["WanVideoModelLoader"]()
import sys
import inspect
loader_module = inspect.getmodule(original_loader)
if loader_module:
logging.info(f"[MultiGPU WanVideoModelLoader] Patching WanVideo modules to use {selected_device}")
original_device = getattr(loader_module, 'device', None)
original_offload = getattr(loader_module, 'offload_device', None)
# Check if there's a model offload device override (from block swap config)
model_offload_override = getattr(loader_module, '_model_offload_device_override', None)
setattr(loader_module, 'device', selected_device)
if model_offload_override:
setattr(loader_module, 'offload_device', model_offload_override)
logging.info(f"[MultiGPU WanVideoModelLoader] Using model offload override: {model_offload_override}")
elif device == "cpu":
setattr(loader_module, 'offload_device', selected_device)
nodes_module_name = loader_module.__name__.replace('.nodes_model_loading', '.nodes')
if nodes_module_name in sys.modules:
nodes_module = sys.modules[nodes_module_name]
setattr(nodes_module, 'device', selected_device)
nodes_model_offload_override = getattr(nodes_module, '_model_offload_device_override', None)
if nodes_model_offload_override:
setattr(nodes_module, 'offload_device', nodes_model_offload_override)
elif device == "cpu":
setattr(nodes_module, 'offload_device', selected_device)
logging.info(f"[MultiGPU WanVideoModelLoader] Both modules patched successfully")
logging.info(f"[MultiGPU WanVideoModelLoader] Calling original loader")
result = original_loader.loadmodel(model, base_precision, load_device, quantization,
compile_args, attention_mode, block_swap_args, lora, vram_management_args, vace_model, fantasytalking_model, multitalk_model)
# After model is loaded, check if we have a transformer and patch it for block swap
if result and len(result) > 0 and hasattr(result[0], 'model'):
model_obj = result[0]
if hasattr(model_obj.model, 'diffusion_model'):
transformer = model_obj.model.diffusion_model
block_swap_override = getattr(loader_module, '_block_swap_device_override', None)
if block_swap_override:
transformer.offload_device = block_swap_override
logging.info(f"[MultiGPU WanVideoModelLoader] Patched transformer for block swap to use: {block_swap_override}")
logging.info(f"[MultiGPU WanVideoModelLoader] Model loaded on {selected_device}")
return result
else:
logging.error(f"[MultiGPU WanVideoModelLoader] Could not patch modules, falling back")
return original_loader.loadmodel(model, base_precision, load_device, quantization,
compile_args, attention_mode, block_swap_args, lora, vram_management_args, vace_model, fantasytalking_model, multitalk_model)
class WanVideoVAELoader:
@classmethod
def INPUT_TYPES(s):
from . import get_device_list
devices = get_device_list()
return {
"required": {
"model_name": (folder_paths.get_filename_list("vae"),
{"tooltip": "These models are loaded from 'ComfyUI/models/vae'"}),
"device": (devices, {"default": devices[1] if len(devices) > 1 else devices[0],
"tooltip": "Device to load the VAE to"}),
},
"optional": {
"precision": (["fp16", "fp32", "bf16"], {"default": "bf16"}),
"compile_args": ("WANCOMPILEARGS", ),
}
}
RETURN_TYPES = ("WANVAE",)
RETURN_NAMES = ("vae", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Loads Wan VAE model with explicit device selection"
def loadmodel(self, model_name, device, precision="bf16", compile_args=None):
import logging
import torch
logging.info(f"[MultiGPU WanVideoVAELoader] User selected device: {device}")
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["WanVideoVAELoader"]()
import sys
import inspect
loader_module = inspect.getmodule(original_loader)
if loader_module:
selected_device = torch.device(device)
logging.info(f"[MultiGPU WanVideoVAELoader] Patching modules to use {selected_device}")
setattr(loader_module, 'offload_device', selected_device)
setattr(loader_module, 'device', selected_device)
nodes_module_name = loader_module.__name__.replace('.nodes_model_loading', '.nodes')
if nodes_module_name in sys.modules:
nodes_module = sys.modules[nodes_module_name]
setattr(nodes_module, 'device', selected_device)
setattr(nodes_module, 'offload_device', selected_device)
result = original_loader.loadmodel(model_name, precision, compile_args)
logging.info(f"[MultiGPU WanVideoVAELoader] VAE loaded on {selected_device}")
return result
else:
logging.error(f"[MultiGPU WanVideoVAELoader] Could not patch modules")
return original_loader.loadmodel(model_name, precision, compile_args)
class LoadWanVideoT5TextEncoder:
@classmethod
def INPUT_TYPES(s):
from . import get_device_list
devices = get_device_list()
return {
"required": {
"model_name": (folder_paths.get_filename_list("text_encoders"),
{"tooltip": "These models are loaded from 'ComfyUI/models/text_encoders'"}),
"precision": (["fp32", "bf16"], {"default": "bf16"}),
"device": (devices, {"default": devices[1] if len(devices) > 1 else devices[0],
"tooltip": "Device to load the text encoder to"}),
},
"optional": {
"quantization": (['disabled', 'fp8_e4m3fn'],
{"default": 'disabled', "tooltip": "optional quantization method"}),
}
}
RETURN_TYPES = ("WANTEXTENCODER",)
RETURN_NAMES = ("wan_t5_model", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Loads Wan text_encoder model from 'ComfyUI/models/text_encoders'"
def loadmodel(self, model_name, precision, device, quantization="disabled"):
import logging
import torch
logging.info(f"[MultiGPU LoadWanVideoT5TextEncoder] ========== CUSTOM IMPLEMENTATION ==========")
logging.info(f"[MultiGPU LoadWanVideoT5TextEncoder] User selected device: {device}")
selected_device = torch.device(device)
load_device = "offload_device" if device == "cpu" else "main_device"
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["LoadWanVideoT5TextEncoder"]()
import sys
import inspect
loader_module = inspect.getmodule(original_loader)
if loader_module:
logging.info(f"[MultiGPU LoadWanVideoT5TextEncoder] Patching WanVideo modules to use {selected_device}")
setattr(loader_module, 'device', selected_device)
if device == "cpu":
setattr(loader_module, 'offload_device', selected_device)
nodes_module_name = loader_module.__name__.replace('.nodes_model_loading', '.nodes')
if nodes_module_name in sys.modules:
nodes_module = sys.modules[nodes_module_name]
setattr(nodes_module, 'device', selected_device)
if device == "cpu":
setattr(nodes_module, 'offload_device', selected_device)
result = original_loader.loadmodel(model_name, precision, load_device, quantization)
logging.info(f"[MultiGPU LoadWanVideoT5TextEncoder] Text encoder loaded on {selected_device}")
return result
else:
logging.error(f"[MultiGPU LoadWanVideoT5TextEncoder] Could not patch modules, falling back")
return original_loader.loadmodel(model_name, precision, load_device, quantization)
class WanVideoTextEncode:
@classmethod
def INPUT_TYPES(s):
from . import get_device_list
devices = get_device_list()
return {"required": {
"positive_prompt": ("STRING", {"default": "", "multiline": True} ),
"negative_prompt": ("STRING", {"default": "", "multiline": True} ),
"device": (devices, {"default": devices[1] if len(devices) > 1 else devices[0],
"tooltip": "Device to run the text encoding on"}),
},
"optional": {
"t5": ("WANTEXTENCODER",),
"force_offload": ("BOOLEAN", {"default": True}),
"model_to_offload": ("WANVIDEOMODEL", {"tooltip": "Model to move to offload_device before encoding"}),
"use_disk_cache": ("BOOLEAN", {"default": False, "tooltip": "Cache the text embeddings to disk for faster re-use"}),
}
}
RETURN_TYPES = ("WANVIDEOTEXTEMBEDS", )
RETURN_NAMES = ("text_embeds",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Encodes text prompts with explicit device selection"
def process(self, positive_prompt, negative_prompt, device, t5=None, force_offload=True,
model_to_offload=None, use_disk_cache=False):
import logging
import torch
logging.info(f"[MultiGPU WanVideoTextEncode] User selected device: {device}")
original_device = "gpu" if device != "cpu" else "cpu"
from nodes import NODE_CLASS_MAPPINGS
original_encoder = NODE_CLASS_MAPPINGS["WanVideoTextEncode"]()
import sys
import inspect
encoder_module = inspect.getmodule(original_encoder)
if encoder_module:
selected_device = torch.device(device)
logging.info(f"[MultiGPU WanVideoTextEncode] Patching module to use {selected_device}")
setattr(encoder_module, 'device', selected_device)
model_loading_name = encoder_module.__name__.replace('.nodes', '.nodes_model_loading')
if model_loading_name in sys.modules:
model_loading_module = sys.modules[model_loading_name]
setattr(model_loading_module, 'device', selected_device)
result = original_encoder.process(positive_prompt, negative_prompt, t5=t5,
force_offload=force_offload, model_to_offload=model_to_offload,
use_disk_cache=use_disk_cache, device=original_device)
logging.info(f"[MultiGPU WanVideoTextEncode] Encoding completed on {selected_device}")
return result
else:
return original_encoder.process(positive_prompt, negative_prompt, t5=t5,
force_offload=force_offload, model_to_offload=model_to_offload,
use_disk_cache=use_disk_cache, device=original_device)
class LoadWanVideoClipTextEncoder:
@classmethod
def INPUT_TYPES(s):
from . import get_device_list
devices = get_device_list()
return {
"required": {
"model_name": (folder_paths.get_filename_list("clip_vision") + folder_paths.get_filename_list("text_encoders"),
{"tooltip": "These models are loaded from 'ComfyUI/models/clip_vision'"}),
"precision": (["fp16", "fp32", "bf16"], {"default": "fp16"}),
"device": (devices, {"default": devices[1] if len(devices) > 1 else devices[0],
"tooltip": "Device to load the CLIP encoder to"}),
}
}
RETURN_TYPES = ("CLIP_VISION",)
RETURN_NAMES = ("clip_vision", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Loads Wan CLIP text encoder model from 'ComfyUI/models/clip_vision'"
def loadmodel(self, model_name, precision, device):
import logging
import torch
logging.info(f"[MultiGPU LoadWanVideoClipTextEncoder] ========== CUSTOM IMPLEMENTATION ==========")
logging.info(f"[MultiGPU LoadWanVideoClipTextEncoder] User selected device: {device}")
selected_device = torch.device(device)
load_device = "offload_device" if device == "cpu" else "main_device"
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["LoadWanVideoClipTextEncoder"]()
import sys
import inspect
loader_module = inspect.getmodule(original_loader)
if loader_module:
logging.info(f"[MultiGPU LoadWanVideoClipTextEncoder] Patching WanVideo modules to use {selected_device}")
setattr(loader_module, 'device', selected_device)
if device == "cpu":
setattr(loader_module, 'offload_device', selected_device)
nodes_module_name = loader_module.__name__.replace('.nodes_model_loading', '.nodes')
if nodes_module_name in sys.modules:
nodes_module = sys.modules[nodes_module_name]
setattr(nodes_module, 'device', selected_device)
if device == "cpu":
setattr(nodes_module, 'offload_device', selected_device)
result = original_loader.loadmodel(model_name, precision, load_device)
logging.info(f"[MultiGPU LoadWanVideoClipTextEncoder] CLIP encoder loaded on {selected_device}")
return result
else:
logging.error(f"[MultiGPU LoadWanVideoClipTextEncoder] Could not patch modules, falling back")
return original_loader.loadmodel(model_name, precision, load_device)
class WanVideoModelLoader_2:
"""Second instance for multi-model workflows to maintain separate device patches"""
@classmethod
def INPUT_TYPES(s):
# Delegate to the primary loader
return WanVideoModelLoader.INPUT_TYPES()
RETURN_TYPES = WanVideoModelLoader.RETURN_TYPES
RETURN_NAMES = WanVideoModelLoader.RETURN_NAMES
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Second model loader instance for workflows using multiple models on different devices"
def loadmodel(self, model, base_precision, device, quantization,
compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None,
vram_management_args=None, vace_model=None, fantasytalking_model=None, multitalk_model=None):
loader = WanVideoModelLoader()
return loader.loadmodel(model, base_precision, device, quantization,
compile_args, attention_mode, block_swap_args, lora,
vram_management_args, vace_model, fantasytalking_model, multitalk_model)
class WanVideoSampler:
"""Wrapper that ensures correct device patching before sampling"""
@classmethod
def INPUT_TYPES(s):
# Get original sampler's inputs
from nodes import NODE_CLASS_MAPPINGS
original_types = NODE_CLASS_MAPPINGS["WanVideoSampler"].INPUT_TYPES()
return original_types
RETURN_TYPES = ("LATENT", "LATENT",)
RETURN_NAMES = ("samples", "denoised_samples",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "MultiGPU-aware sampler that ensures correct device for each model"
def process(self, model, **kwargs):
import sys
import torch
import logging
model_device = model.load_device
logging.info(f"[MultiGPU WanVideoSampler] Processing on device: {model_device}")
for module_name in sys.modules.keys():
if 'WanVideoWrapper' in module_name and hasattr(sys.modules[module_name], 'device'):
sys.modules[module_name].device = model_device
from nodes import NODE_CLASS_MAPPINGS
original_sampler = NODE_CLASS_MAPPINGS["WanVideoSampler"]()
return original_sampler.process(model, **kwargs)
class WanVideoBlockSwap:
@classmethod
def INPUT_TYPES(s):
from . import get_device_list
devices = get_device_list()
return {
"required": {
"blocks_to_swap": ("INT", {"default": 20, "min": 0, "max": 40, "step": 1,
"tooltip": "Number of transformer blocks to swap, the 14B model has 40, while the 1.3B model has 30 blocks"}),
"swap_device": (devices, {"default": "cpu",
"tooltip": "Device to swap blocks to during sampling (default: cpu for standard behavior)"}),
"model_offload_device": (devices, {"default": "cpu",
"tooltip": "Device to offload entire model to when done (default: cpu)"}),
"offload_img_emb": ("BOOLEAN", {"default": False, "tooltip": "Offload img_emb to swap_device"}),
"offload_txt_emb": ("BOOLEAN", {"default": False, "tooltip": "Offload time_emb to swap_device"}),
},
"optional": {
"use_non_blocking": ("BOOLEAN", {"default": False,
"tooltip": "Use non-blocking memory transfer for offloading, reserves more RAM but is faster"}),
"vace_blocks_to_swap": ("INT", {"default": 0, "min": 0, "max": 15, "step": 1,
"tooltip": "Number of VACE blocks to swap, the VACE model has 15 blocks"}),
},
}
RETURN_TYPES = ("BLOCKSWAPARGS",)
RETURN_NAMES = ("block_swap_args",)
FUNCTION = "setargs"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Block swap settings with explicit device selection for memory management across GPUs"
def setargs(self, blocks_to_swap, swap_device, model_offload_device, offload_img_emb, offload_txt_emb,
use_non_blocking=False, vace_blocks_to_swap=0):
import logging
import torch
import comfy.model_management as mm
logging.info(f"[MultiGPU WanVideoBlockSwap] ========== CONFIGURATION ==========")
logging.info(f"[MultiGPU WanVideoBlockSwap] User selected swap device: {swap_device}")
logging.info(f"[MultiGPU WanVideoBlockSwap] User selected model offload device: {model_offload_device}")
logging.info(f"[MultiGPU WanVideoBlockSwap] Blocks to swap: {blocks_to_swap}")
selected_swap_device = torch.device(swap_device)
selected_offload_device = torch.device(model_offload_device)
import sys
for module_name in sys.modules.keys():
if 'WanVideoWrapper' in module_name and 'nodes_model_loading' in module_name:
module = sys.modules[module_name]
setattr(module, 'offload_device', selected_offload_device)
setattr(module, '_block_swap_device_override', selected_swap_device)
setattr(module, '_model_offload_device_override', selected_offload_device)
logging.info(f"[MultiGPU WanVideoBlockSwap] Patched {module_name} for offload to {selected_offload_device} and swap to {selected_swap_device}")
if 'WanVideoWrapper' in module_name and module_name.endswith('.nodes'):
module = sys.modules[module_name]
setattr(module, 'offload_device', selected_offload_device)
setattr(module, '_block_swap_device_override', selected_swap_device)
setattr(module, '_model_offload_device_override', selected_offload_device)
block_swap_args = {
"blocks_to_swap": blocks_to_swap,
"offload_img_emb": offload_img_emb,
"offload_txt_emb": offload_txt_emb,
"use_non_blocking": use_non_blocking,
"vace_blocks_to_swap": vace_blocks_to_swap,
"swap_device": swap_device,
"model_offload_device": model_offload_device,
}
logging.info(f"[MultiGPU WanVideoBlockSwap] Block swap configuration complete")
return (block_swap_args,)
+456
View File
@@ -0,0 +1,456 @@
import logging
import torch
import sys
import inspect
import folder_paths
import comfy.model_management as mm
class WanVideoModelLoader:
@classmethod
def INPUT_TYPES(s):
from . import get_device_list
devices = get_device_list()
return {
"required": {
"model": (folder_paths.get_filename_list("unet_gguf") + folder_paths.get_filename_list("diffusion_models"),
{"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' folder",}),
"base_precision": (["fp32", "bf16", "fp16", "fp16_fast"], {"default": "bf16"}),
"quantization": (
["disabled", "fp8_e4m3fn", "fp8_e4m3fn_fast", "fp8_e5m2", "fp8_e4m3fn_fast_no_ffn", "fp8_e4m3fn_scaled", "fp8_e5m2_scaled"],
{"default": "disabled", "tooltip": "optional quantization method"}
),
"device": (devices, {"default": devices[1] if len(devices) > 1 else devices[0], "tooltip": "Device to load the model to"}),
},
"optional": {
"attention_mode": ([
"sdpa",
"flash_attn_2",
"flash_attn_3",
"sageattn",
"sageattn_3",
"flex_attention",
"radial_sage_attention",
], {"default": "sdpa"}),
"compile_args": ("WANCOMPILEARGS", ),
"block_swap_args": ("BLOCKSWAPARGS", ),
"lora": ("WANVIDLORA", {"default": None}),
"vram_management_args": ("VRAM_MANAGEMENTARGS", {"default": None, "tooltip": "Alternative offloading method from DiffSynth-Studio, more aggressive in reducing memory use than block swapping, but can be slower"}),
"vace_model": ("VACEPATH", {"default": None, "tooltip": "VACE model to use when not using model that has it included"}),
"fantasytalking_model": ("FANTASYTALKINGMODEL", {"default": None, "tooltip": "FantasyTalking model https://github.com/Fantasy-AMAP"}),
"multitalk_model": ("MULTITALKMODEL", {"default": None, "tooltip": "Multitalk model"}),
}
}
RETURN_TYPES = ("WANVIDEOMODEL",)
RETURN_NAMES = ("model", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
def loadmodel(self, model, base_precision, device, quantization,
compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None, vram_management_args=None, vace_model=None, fantasytalking_model=None, multitalk_model=None):
logging.info(f"[MultiGPU WanVideoModelLoader] ========== CUSTOM IMPLEMENTATION ==========")
logging.info(f"[MultiGPU WanVideoModelLoader] User selected device: {device}")
selected_device = torch.device(device)
load_device = "offload_device" if device == "cpu" else "main_device"
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["WanVideoModelLoader"]()
loader_module = inspect.getmodule(original_loader)
if loader_module:
logging.info(f"[MultiGPU WanVideoModelLoader] Patching WanVideo modules to use {selected_device}")
original_device = getattr(loader_module, 'device', None)
original_offload = getattr(loader_module, 'offload_device', None)
model_offload_override = getattr(loader_module, '_model_offload_device_override', None)
setattr(loader_module, 'device', selected_device)
if model_offload_override:
setattr(loader_module, 'offload_device', model_offload_override)
logging.info(f"[MultiGPU WanVideoModelLoader] Using model offload override: {model_offload_override}")
elif device == "cpu":
setattr(loader_module, 'offload_device', selected_device)
nodes_module_name = loader_module.__name__.replace('.nodes_model_loading', '.nodes')
if nodes_module_name in sys.modules:
nodes_module = sys.modules[nodes_module_name]
setattr(nodes_module, 'device', selected_device)
nodes_model_offload_override = getattr(nodes_module, '_model_offload_device_override', None)
if nodes_model_offload_override:
setattr(nodes_module, 'offload_device', nodes_model_offload_override)
elif device == "cpu":
setattr(nodes_module, 'offload_device', selected_device)
logging.info(f"[MultiGPU WanVideoModelLoader] Both modules patched successfully")
logging.info(f"[MultiGPU WanVideoModelLoader] Calling original loader")
result = original_loader.loadmodel(model, base_precision, load_device, quantization,
compile_args, attention_mode, block_swap_args, lora, vram_management_args, vace_model, fantasytalking_model, multitalk_model)
if result and len(result) > 0 and hasattr(result[0], 'model'):
model_obj = result[0]
if hasattr(model_obj.model, 'diffusion_model'):
transformer = model_obj.model.diffusion_model
block_swap_override = getattr(loader_module, '_block_swap_device_override', None)
if block_swap_override:
transformer.offload_device = block_swap_override
logging.info(f"[MultiGPU WanVideoModelLoader] Patched transformer for block swap to use: {block_swap_override}")
logging.info(f"[MultiGPU WanVideoModelLoader] Model loaded on {selected_device}")
return result
else:
logging.error(f"[MultiGPU WanVideoModelLoader] Could not patch modules, falling back")
return original_loader.loadmodel(model, base_precision, load_device, quantization,
compile_args, attention_mode, block_swap_args, lora, vram_management_args, vace_model, fantasytalking_model, multitalk_model)
class WanVideoVAELoader:
@classmethod
def INPUT_TYPES(s):
from . import get_device_list
devices = get_device_list()
return {
"required": {
"model_name": (folder_paths.get_filename_list("vae"),
{"tooltip": "These models are loaded from 'ComfyUI/models/vae'"}),
"device": (devices, {"default": devices[1] if len(devices) > 1 else devices[0],
"tooltip": "Device to load the VAE to"}),
},
"optional": {
"precision": (["fp16", "fp32", "bf16"], {"default": "bf16"}),
"compile_args": ("WANCOMPILEARGS", ),
}
}
RETURN_TYPES = ("WANVAE",)
RETURN_NAMES = ("vae", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Loads Wan VAE model with explicit device selection"
def loadmodel(self, model_name, device, precision="bf16", compile_args=None):
logging.info(f"[MultiGPU WanVideoVAELoader] User selected device: {device}")
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["WanVideoVAELoader"]()
loader_module = inspect.getmodule(original_loader)
if loader_module:
selected_device = torch.device(device)
logging.info(f"[MultiGPU WanVideoVAELoader] Patching modules to use {selected_device}")
setattr(loader_module, 'offload_device', selected_device)
setattr(loader_module, 'device', selected_device)
nodes_module_name = loader_module.__name__.replace('.nodes_model_loading', '.nodes')
if nodes_module_name in sys.modules:
nodes_module = sys.modules[nodes_module_name]
setattr(nodes_module, 'device', selected_device)
setattr(nodes_module, 'offload_device', selected_device)
result = original_loader.loadmodel(model_name, precision, compile_args)
logging.info(f"[MultiGPU WanVideoVAELoader] VAE loaded on {selected_device}")
return result
else:
logging.error(f"[MultiGPU WanVideoVAELoader] Could not patch modules")
return original_loader.loadmodel(model_name, precision, compile_args)
class LoadWanVideoT5TextEncoder:
@classmethod
def INPUT_TYPES(s):
from . import get_device_list
devices = get_device_list()
return {
"required": {
"model_name": (folder_paths.get_filename_list("text_encoders"),
{"tooltip": "These models are loaded from 'ComfyUI/models/text_encoders'"}),
"precision": (["fp32", "bf16"], {"default": "bf16"}),
"device": (devices, {"default": devices[1] if len(devices) > 1 else devices[0],
"tooltip": "Device to load the text encoder to"}),
},
"optional": {
"quantization": (['disabled', 'fp8_e4m3fn'],
{"default": 'disabled', "tooltip": "optional quantization method"}),
}
}
RETURN_TYPES = ("WANTEXTENCODER",)
RETURN_NAMES = ("wan_t5_model", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Loads Wan text_encoder model from 'ComfyUI/models/text_encoders'"
def loadmodel(self, model_name, precision, device, quantization="disabled"):
logging.info(f"[MultiGPU LoadWanVideoT5TextEncoder] ========== CUSTOM IMPLEMENTATION ==========")
logging.info(f"[MultiGPU LoadWanVideoT5TextEncoder] User selected device: {device}")
selected_device = torch.device(device)
load_device = "offload_device" if device == "cpu" else "main_device"
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["LoadWanVideoT5TextEncoder"]()
loader_module = inspect.getmodule(original_loader)
if loader_module:
logging.info(f"[MultiGPU LoadWanVideoT5TextEncoder] Patching WanVideo modules to use {selected_device}")
setattr(loader_module, 'device', selected_device)
if device == "cpu":
setattr(loader_module, 'offload_device', selected_device)
nodes_module_name = loader_module.__name__.replace('.nodes_model_loading', '.nodes')
if nodes_module_name in sys.modules:
nodes_module = sys.modules[nodes_module_name]
setattr(nodes_module, 'device', selected_device)
if device == "cpu":
setattr(nodes_module, 'offload_device', selected_device)
result = original_loader.loadmodel(model_name, precision, load_device, quantization)
logging.info(f"[MultiGPU LoadWanVideoT5TextEncoder] Text encoder loaded on {selected_device}")
return result
else:
logging.error(f"[MultiGPU LoadWanVideoT5TextEncoder] Could not patch modules, falling back")
return original_loader.loadmodel(model_name, precision, load_device, quantization)
class WanVideoTextEncode:
@classmethod
def INPUT_TYPES(s):
from . import get_device_list
devices = get_device_list()
return {"required": {
"positive_prompt": ("STRING", {"default": "", "multiline": True} ),
"negative_prompt": ("STRING", {"default": "", "multiline": True} ),
"device": (devices, {"default": devices[1] if len(devices) > 1 else devices[0],
"tooltip": "Device to run the text encoding on"}),
},
"optional": {
"t5": ("WANTEXTENCODER",),
"force_offload": ("BOOLEAN", {"default": True}),
"model_to_offload": ("WANVIDEOMODEL", {"tooltip": "Model to move to offload_device before encoding"}),
"use_disk_cache": ("BOOLEAN", {"default": False, "tooltip": "Cache the text embeddings to disk for faster re-use"}),
}
}
RETURN_TYPES = ("WANVIDEOTEXTEMBEDS", )
RETURN_NAMES = ("text_embeds",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Encodes text prompts with explicit device selection"
def process(self, positive_prompt, negative_prompt, device, t5=None, force_offload=True,
model_to_offload=None, use_disk_cache=False):
logging.info(f"[MultiGPU WanVideoTextEncode] User selected device: {device}")
original_device = "gpu" if device != "cpu" else "cpu"
from nodes import NODE_CLASS_MAPPINGS
original_encoder = NODE_CLASS_MAPPINGS["WanVideoTextEncode"]()
encoder_module = inspect.getmodule(original_encoder)
if encoder_module:
selected_device = torch.device(device)
logging.info(f"[MultiGPU WanVideoTextEncode] Patching module to use {selected_device}")
setattr(encoder_module, 'device', selected_device)
model_loading_name = encoder_module.__name__.replace('.nodes', '.nodes_model_loading')
if model_loading_name in sys.modules:
model_loading_module = sys.modules[model_loading_name]
setattr(model_loading_module, 'device', selected_device)
result = original_encoder.process(positive_prompt, negative_prompt, t5=t5,
force_offload=force_offload, model_to_offload=model_to_offload,
use_disk_cache=use_disk_cache, device=original_device)
logging.info(f"[MultiGPU WanVideoTextEncode] Encoding completed on {selected_device}")
return result
else:
return original_encoder.process(positive_prompt, negative_prompt, t5=t5,
force_offload=force_offload, model_to_offload=model_to_offload,
use_disk_cache=use_disk_cache, device=original_device)
class LoadWanVideoClipTextEncoder:
@classmethod
def INPUT_TYPES(s):
from . import get_device_list
devices = get_device_list()
return {
"required": {
"model_name": (folder_paths.get_filename_list("clip_vision") + folder_paths.get_filename_list("text_encoders"),
{"tooltip": "These models are loaded from 'ComfyUI/models/clip_vision'"}),
"precision": (["fp16", "fp32", "bf16"], {"default": "fp16"}),
"device": (devices, {"default": devices[1] if len(devices) > 1 else devices[0],
"tooltip": "Device to load the CLIP encoder to"}),
}
}
RETURN_TYPES = ("CLIP_VISION",)
RETURN_NAMES = ("clip_vision", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Loads Wan CLIP text encoder model from 'ComfyUI/models/clip_vision'"
def loadmodel(self, model_name, precision, device):
logging.info(f"[MultiGPU LoadWanVideoClipTextEncoder] ========== CUSTOM IMPLEMENTATION ==========")
logging.info(f"[MultiGPU LoadWanVideoClipTextEncoder] User selected device: {device}")
selected_device = torch.device(device)
load_device = "offload_device" if device == "cpu" else "main_device"
from nodes import NODE_CLASS_MAPPINGS
original_loader = NODE_CLASS_MAPPINGS["LoadWanVideoClipTextEncoder"]()
loader_module = inspect.getmodule(original_loader)
if loader_module:
logging.info(f"[MultiGPU LoadWanVideoClipTextEncoder] Patching WanVideo modules to use {selected_device}")
setattr(loader_module, 'device', selected_device)
if device == "cpu":
setattr(loader_module, 'offload_device', selected_device)
nodes_module_name = loader_module.__name__.replace('.nodes_model_loading', '.nodes')
if nodes_module_name in sys.modules:
nodes_module = sys.modules[nodes_module_name]
setattr(nodes_module, 'device', selected_device)
if device == "cpu":
setattr(nodes_module, 'offload_device', selected_device)
result = original_loader.loadmodel(model_name, precision, load_device)
logging.info(f"[MultiGPU LoadWanVideoClipTextEncoder] CLIP encoder loaded on {selected_device}")
return result
else:
logging.error(f"[MultiGPU LoadWanVideoClipTextEncoder] Could not patch modules, falling back")
return original_loader.loadmodel(model_name, precision, load_device)
class WanVideoModelLoader_2:
@classmethod
def INPUT_TYPES(s):
return WanVideoModelLoader.INPUT_TYPES()
RETURN_TYPES = WanVideoModelLoader.RETURN_TYPES
RETURN_NAMES = WanVideoModelLoader.RETURN_NAMES
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Second model loader instance for workflows using multiple models on different devices"
def loadmodel(self, model, base_precision, device, quantization,
compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None,
vram_management_args=None, vace_model=None, fantasytalking_model=None, multitalk_model=None):
loader = WanVideoModelLoader()
return loader.loadmodel(model, base_precision, device, quantization,
compile_args, attention_mode, block_swap_args, lora,
vram_management_args, vace_model, fantasytalking_model, multitalk_model)
class WanVideoSampler:
@classmethod
def INPUT_TYPES(s):
from nodes import NODE_CLASS_MAPPINGS
original_types = NODE_CLASS_MAPPINGS["WanVideoSampler"].INPUT_TYPES()
return original_types
RETURN_TYPES = ("LATENT", "LATENT",)
RETURN_NAMES = ("samples", "denoised_samples",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "MultiGPU-aware sampler that ensures correct device for each model"
def process(self, model, **kwargs):
model_device = model.load_device
logging.info(f"[MultiGPU WanVideoSampler] Processing on device: {model_device}")
for module_name in sys.modules.keys():
if 'WanVideoWrapper' in module_name and hasattr(sys.modules[module_name], 'device'):
sys.modules[module_name].device = model_device
from nodes import NODE_CLASS_MAPPINGS
original_sampler = NODE_CLASS_MAPPINGS["WanVideoSampler"]()
return original_sampler.process(model, **kwargs)
class WanVideoBlockSwap:
@classmethod
def INPUT_TYPES(s):
from . import get_device_list
devices = get_device_list()
return {
"required": {
"blocks_to_swap": ("INT", {"default": 20, "min": 0, "max": 40, "step": 1,
"tooltip": "Number of transformer blocks to swap, the 14B model has 40, while the 1.3B model has 30 blocks"}),
"swap_device": (devices, {"default": "cpu",
"tooltip": "Device to swap blocks to during sampling (default: cpu for standard behavior)"}),
"model_offload_device": (devices, {"default": "cpu",
"tooltip": "Device to offload entire model to when done (default: cpu)"}),
"offload_img_emb": ("BOOLEAN", {"default": False, "tooltip": "Offload img_emb to swap_device"}),
"offload_txt_emb": ("BOOLEAN", {"default": False, "tooltip": "Offload time_emb to swap_device"}),
},
"optional": {
"use_non_blocking": ("BOOLEAN", {"default": False,
"tooltip": "Use non-blocking memory transfer for offloading, reserves more RAM but is faster"}),
"vace_blocks_to_swap": ("INT", {"default": 0, "min": 0, "max": 15, "step": 1,
"tooltip": "Number of VACE blocks to swap, the VACE model has 15 blocks"}),
},
}
RETURN_TYPES = ("BLOCKSWAPARGS",)
RETURN_NAMES = ("block_swap_args",)
FUNCTION = "setargs"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Block swap settings with explicit device selection for memory management across GPUs"
def setargs(self, blocks_to_swap, swap_device, model_offload_device, offload_img_emb, offload_txt_emb,
use_non_blocking=False, vace_blocks_to_swap=0):
logging.info(f"[MultiGPU WanVideoBlockSwap] ========== CONFIGURATION ==========")
logging.info(f"[MultiGPU WanVideoBlockSwap] User selected swap device: {swap_device}")
logging.info(f"[MultiGPU WanVideoBlockSwap] User selected model offload device: {model_offload_device}")
logging.info(f"[MultiGPU WanVideoBlockSwap] Blocks to swap: {blocks_to_swap}")
selected_swap_device = torch.device(swap_device)
selected_offload_device = torch.device(model_offload_device)
for module_name in sys.modules.keys():
if 'WanVideoWrapper' in module_name and 'nodes_model_loading' in module_name:
module = sys.modules[module_name]
setattr(module, 'offload_device', selected_offload_device)
setattr(module, '_block_swap_device_override', selected_swap_device)
setattr(module, '_model_offload_device_override', selected_offload_device)
logging.info(f"[MultiGPU WanVideoBlockSwap] Patched {module_name} for offload to {selected_offload_device} and swap to {selected_swap_device}")
if 'WanVideoWrapper' in module_name and module_name.endswith('.nodes'):
module = sys.modules[module_name]
setattr(module, 'offload_device', selected_offload_device)
setattr(module, '_block_swap_device_override', selected_swap_device)
setattr(module, '_model_offload_device_override', selected_offload_device)
block_swap_args = {
"blocks_to_swap": blocks_to_swap,
"offload_img_emb": offload_img_emb,
"offload_txt_emb": offload_txt_emb,
"use_non_blocking": use_non_blocking,
"vace_blocks_to_swap": vace_blocks_to_swap,
"swap_device": swap_device,
"model_offload_device": model_offload_device,
}
logging.info(f"[MultiGPU WanVideoBlockSwap] Block swap configuration complete")
return (block_swap_args,)