diff --git a/__init__.py b/__init__.py index 289d9b3..d54dfa7 100644 --- a/__init__.py +++ b/__init__.py @@ -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())}") diff --git a/block_swap.py b/block_swap.py new file mode 100644 index 0000000..f53884d --- /dev/null +++ b/block_swap.py @@ -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 diff --git a/distorch.py b/distorch.py new file mode 100644 index 0000000..5d499f8 --- /dev/null +++ b/distorch.py @@ -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 diff --git a/nodes.py b/nodes.py index 733f295..8703a4e 100644 --- a/nodes.py +++ b/nodes.py @@ -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,) diff --git a/wanvideo.py b/wanvideo.py new file mode 100644 index 0000000..83b80fd --- /dev/null +++ b/wanvideo.py @@ -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,)