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:
+116
-741
@@ -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
@@ -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
@@ -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
|
||||
@@ -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
@@ -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,)
|
||||
Reference in New Issue
Block a user