Standardize doc strings and make PEP 257 compliant
This commit is contained in:
+8
-1
@@ -35,11 +35,13 @@ if not logger.handlers:
|
||||
logger.setLevel(log_level)
|
||||
|
||||
def mgpu_mm_log_method(self, msg):
|
||||
"""Add MultiGPU model management logging method to logger instance."""
|
||||
if MGPU_MM_LOG:
|
||||
self.info(f"[MultiGPU Model Management] {msg}")
|
||||
logger.mgpu_mm_log = mgpu_mm_log_method.__get__(logger, type(logger))
|
||||
|
||||
def check_module_exists(module_path):
|
||||
"""Check if a custom node module exists in ComfyUI custom_nodes directory."""
|
||||
full_path = os.path.join(folder_paths.get_folder_paths("custom_nodes")[0], module_path)
|
||||
logger.debug(f"[MultiGPU] Checking for module at {full_path}")
|
||||
if not os.path.exists(full_path):
|
||||
@@ -52,16 +54,19 @@ current_device = mm.get_torch_device()
|
||||
current_text_encoder_device = mm.text_encoder_device()
|
||||
|
||||
def set_current_device(device):
|
||||
"""Set the current device context for MultiGPU operations."""
|
||||
global current_device
|
||||
current_device = device
|
||||
logger.debug(f"[MultiGPU Initialization] current_device set to: {device}")
|
||||
|
||||
def set_current_text_encoder_device(device):
|
||||
"""Set the current text encoder device context for CLIP models."""
|
||||
global current_text_encoder_device
|
||||
current_text_encoder_device = device
|
||||
logger.debug(f"[MultiGPU Initialization] current_text_encoder_device set to: {device}")
|
||||
|
||||
def get_torch_device_patched():
|
||||
"""Return MultiGPU-aware device selection for patched mm.get_torch_device."""
|
||||
device = None
|
||||
if (not is_accelerator_available() or mm.cpu_state == mm.CPUState.CPU or "cpu" in str(current_device).lower()):
|
||||
device = torch.device("cpu")
|
||||
@@ -72,6 +77,7 @@ def get_torch_device_patched():
|
||||
return device
|
||||
|
||||
def text_encoder_device_patched():
|
||||
"""Return MultiGPU-aware text encoder device for patched mm.text_encoder_device."""
|
||||
device = None
|
||||
if (not is_accelerator_available() or mm.cpu_state == mm.CPUState.CPU or "cpu" in str(current_text_encoder_device).lower()):
|
||||
device = torch.device("cpu")
|
||||
@@ -191,6 +197,7 @@ logger.info(dash_line)
|
||||
registration_data = []
|
||||
|
||||
def register_and_count(module_names, node_map):
|
||||
"""Register MultiGPU node wrappers for detected custom node modules."""
|
||||
found = False
|
||||
for name in module_names:
|
||||
if check_module_exists(name):
|
||||
@@ -281,4 +288,4 @@ for item in registration_data:
|
||||
logger.info(fmt_reg.format(item['name'], item['found'], str(item['count'])))
|
||||
logger.info(dash_line)
|
||||
|
||||
logger.info(f"[MultiGPU] Registration complete. Final mappings: {', '.join(NODE_CLASS_MAPPINGS.keys())}")
|
||||
logger.info(f"[MultiGPU] Registration complete. Final mappings: {', '.join(NODE_CLASS_MAPPINGS.keys())}")
|
||||
|
||||
@@ -19,10 +19,7 @@ checkpoint_distorch_config = {}
|
||||
original_load_state_dict_guess_config = None
|
||||
|
||||
def patch_load_state_dict_guess_config():
|
||||
"""
|
||||
Monkey patch the load_state_dict_guess_config function to replace its logic
|
||||
with a MultiGPU-aware implementation.
|
||||
"""
|
||||
"""Monkey patch comfy.sd.load_state_dict_guess_config with MultiGPU-aware checkpoint loading."""
|
||||
global original_load_state_dict_guess_config
|
||||
|
||||
if original_load_state_dict_guess_config is not None:
|
||||
@@ -36,7 +33,7 @@ def patch_load_state_dict_guess_config():
|
||||
def patched_load_state_dict_guess_config(sd, output_vae=True, output_clip=True, output_clipvision=False,
|
||||
embedding_directory=None, output_model=True, model_options={},
|
||||
te_model_options={}, metadata=None):
|
||||
|
||||
"""Patched checkpoint loader with MultiGPU and DisTorch2 device placement support."""
|
||||
from . import set_current_device, set_current_text_encoder_device, current_device, current_text_encoder_device
|
||||
|
||||
sd_size = sum(p.numel() for p in sd.values() if hasattr(p, 'numel'))
|
||||
|
||||
+11
-94
@@ -1,9 +1,3 @@
|
||||
"""
|
||||
Device detection, management, and inspection utilities for ComfyUI-MultiGPU.
|
||||
Single source of truth for all device enumeration, compatibility checks, and VRAM management.
|
||||
Handles all device types supported by ComfyUI core.
|
||||
"""
|
||||
|
||||
import torch
|
||||
import logging
|
||||
import hashlib
|
||||
@@ -13,13 +7,8 @@ import gc
|
||||
|
||||
logger = logging.getLogger("MultiGPU")
|
||||
|
||||
# Module-level cache for device list (populated once on first call)
|
||||
_DEVICE_LIST_CACHE = None
|
||||
|
||||
# ==========================================================================================
|
||||
# Device Detection and Management
|
||||
# ==========================================================================================
|
||||
|
||||
def get_device_list():
|
||||
"""
|
||||
Enumerate ALL physically available devices that can store torch tensors.
|
||||
@@ -38,25 +27,19 @@ def get_device_list():
|
||||
"""
|
||||
global _DEVICE_LIST_CACHE
|
||||
|
||||
# Return cached result if already populated
|
||||
if _DEVICE_LIST_CACHE is not None:
|
||||
return _DEVICE_LIST_CACHE
|
||||
|
||||
# First time - do the actual detection
|
||||
devs = []
|
||||
|
||||
# CPU is always physically present and can store tensors
|
||||
devs.append("cpu")
|
||||
|
||||
# CUDA devices (NVIDIA GPUs)
|
||||
if hasattr(torch, "cuda") and hasattr(torch.cuda, "is_available") and torch.cuda.is_available():
|
||||
device_count = torch.cuda.device_count()
|
||||
devs += [f"cuda:{i}" for i in range(device_count)]
|
||||
logger.debug(f"[MultiGPU_Device_Utils] Found {device_count} CUDA device(s)")
|
||||
|
||||
# XPU devices (Intel GPUs)
|
||||
try:
|
||||
# Try to import intel extension first (may be required for XPU support)
|
||||
import intel_extension_for_pytorch as ipex
|
||||
except ImportError:
|
||||
pass
|
||||
@@ -66,7 +49,6 @@ def get_device_list():
|
||||
devs += [f"xpu:{i}" for i in range(device_count)]
|
||||
logger.debug(f"[MultiGPU_Device_Utils] Found {device_count} XPU device(s)")
|
||||
|
||||
# NPU devices (Ascend NPUs from Huawei)
|
||||
try:
|
||||
import torch_npu
|
||||
if hasattr(torch, "npu") and hasattr(torch.npu, "is_available") and torch.npu.is_available():
|
||||
@@ -76,7 +58,6 @@ def get_device_list():
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
# MLU devices (Cambricon MLUs)
|
||||
try:
|
||||
import torch_mlu
|
||||
if hasattr(torch, "mlu") and hasattr(torch.mlu, "is_available") and torch.mlu.is_available():
|
||||
@@ -86,12 +67,10 @@ def get_device_list():
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
# MPS device (Apple Metal - single device only)
|
||||
if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
|
||||
devs.append("mps")
|
||||
logger.debug("[MultiGPU_Device_Utils] Found MPS device")
|
||||
|
||||
# DirectML devices (Windows DirectML for AMD/Intel/NVIDIA)
|
||||
try:
|
||||
import torch_directml
|
||||
adapter_count = torch_directml.device_count()
|
||||
@@ -101,7 +80,6 @@ def get_device_list():
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
# IXUCA/CoreX devices (special accelerator)
|
||||
try:
|
||||
if hasattr(torch, "corex"):
|
||||
if hasattr(torch.corex, "device_count"):
|
||||
@@ -114,115 +92,69 @@ def get_device_list():
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
# Cache the result for future calls
|
||||
_DEVICE_LIST_CACHE = devs
|
||||
|
||||
# Log only once when initially populated
|
||||
logger.debug(f"[MultiGPU_Device_Utils] Device list initialized: {devs}")
|
||||
|
||||
return devs
|
||||
|
||||
def is_accelerator_available():
|
||||
"""
|
||||
Check if any accelerator device is available.
|
||||
Used by patched functions to determine CPU fallback.
|
||||
|
||||
Returns True if any GPU/accelerator is available, False otherwise.
|
||||
"""
|
||||
# Check CUDA
|
||||
"""Check if any GPU or accelerator device is available including CUDA, XPU, NPU, MLU, MPS, DirectML, or CoreX."""
|
||||
if hasattr(torch, "cuda") and torch.cuda.is_available():
|
||||
return True
|
||||
|
||||
# Check XPU (Intel GPU)
|
||||
if hasattr(torch, "xpu") and hasattr(torch.xpu, "is_available") and torch.xpu.is_available():
|
||||
return True
|
||||
|
||||
# Check NPU (Ascend)
|
||||
try:
|
||||
import torch_npu
|
||||
if hasattr(torch, "npu") and hasattr(torch.npu, "is_available") and torch.npu.is_available():
|
||||
return True
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
# Check MLU (Cambricon)
|
||||
|
||||
try:
|
||||
import torch_mlu
|
||||
if hasattr(torch, "mlu") and hasattr(torch.mlu, "is_available") and torch.mlu.is_available():
|
||||
return True
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
# Check MPS (Apple Metal)
|
||||
|
||||
if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
|
||||
return True
|
||||
|
||||
# Check DirectML
|
||||
|
||||
try:
|
||||
import torch_directml
|
||||
if torch_directml.device_count() > 0:
|
||||
return True
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
# Check CoreX/IXUCA
|
||||
|
||||
if hasattr(torch, "corex"):
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
def is_device_compatible(device_string):
|
||||
"""
|
||||
Check if a device string represents a valid, available device.
|
||||
|
||||
Args:
|
||||
device_string: Device identifier like "cuda:0", "cpu", "xpu:1", etc.
|
||||
|
||||
Returns:
|
||||
True if the device is available, False otherwise.
|
||||
"""
|
||||
"""Check if a device string represents a valid available device."""
|
||||
available_devices = get_device_list()
|
||||
return device_string in available_devices
|
||||
|
||||
def get_device_type(device_string):
|
||||
"""
|
||||
Extract the device type from a device string.
|
||||
|
||||
Args:
|
||||
device_string: Device identifier like "cuda:0", "cpu", "xpu:1", etc.
|
||||
|
||||
Returns:
|
||||
Device type string (e.g., "cuda", "cpu", "xpu", "npu", "mlu", "mps", "directml", "corex")
|
||||
"""
|
||||
"""Extract device type from device string (e.g. 'cuda' from 'cuda:0')."""
|
||||
if ":" in device_string:
|
||||
return device_string.split(":")[0]
|
||||
return device_string
|
||||
|
||||
def parse_device_string(device_string):
|
||||
"""
|
||||
Parse a device string into type and index.
|
||||
|
||||
Args:
|
||||
device_string: Device identifier like "cuda:0", "cpu", "xpu:1", etc.
|
||||
|
||||
Returns:
|
||||
Tuple of (device_type, device_index) where index is None for non-indexed devices
|
||||
"""
|
||||
"""Parse device string into (device_type, device_index) tuple."""
|
||||
if ":" in device_string:
|
||||
parts = device_string.split(":")
|
||||
return parts[0], int(parts[1])
|
||||
return device_string, None
|
||||
|
||||
# ==========================================================================================
|
||||
# VRAM Management (Multi-device cache clearing)
|
||||
# ==========================================================================================
|
||||
|
||||
def soft_empty_cache_multigpu():
|
||||
"""
|
||||
Replicate ComfyUI's cache clearing but for ALL devices in MultiGPU.
|
||||
Uses context managers to ensure the calling thread's device context is restored.
|
||||
"""
|
||||
# Import model management functions
|
||||
"""Clear allocator caches across all devices using context managers to preserve calling thread device context."""
|
||||
from .model_management_mgpu import multigpu_memory_log
|
||||
|
||||
logger.mgpu_mm_log("soft_empty_cache_multigpu: starting GC and multi-device cache clear")
|
||||
@@ -301,14 +233,7 @@ logger.info("[MultiGPU Core Patching] Patching mm.soft_empty_cache for Comprehen
|
||||
original_soft_empty_cache = mm.soft_empty_cache
|
||||
|
||||
def soft_empty_cache_distorch2_patched(force=False):
|
||||
"""
|
||||
Patched mm.soft_empty_cache.
|
||||
- Prunes DisTorch store bookkeeping to avoid stale references
|
||||
- Manages VRAM: if DisTorch2 models are active, clear allocator caches on all devices;
|
||||
otherwise delegate to original mm.soft_empty_cache.
|
||||
- Manages CPU RAM: adaptive threshold-based PromptExecutor cache reset;
|
||||
and force-triggered reset when explicitly requested (mirrors ComfyUI 'Free memory' button).
|
||||
"""
|
||||
"""Patched mm.soft_empty_cache managing VRAM across all devices, CPU RAM with adaptive thresholding, and DisTorch store pruning."""
|
||||
from .model_management_mgpu import multigpu_memory_log, check_cpu_memory_threshold, trigger_executor_cache_reset
|
||||
from .distorch_2 import safetensor_allocation_store, create_safetensor_model_hash
|
||||
|
||||
@@ -364,15 +289,7 @@ mm.soft_empty_cache = soft_empty_cache_distorch2_patched
|
||||
# ==========================================================================================
|
||||
|
||||
def comfyui_memory_load(tag):
|
||||
"""
|
||||
Returns a single-line, pipe-delimited snapshot of system and device memory usage.
|
||||
|
||||
Format: "tag=<TAG>|cpu=<used_GiB>/<total_GiB>|<device>=<used_GiB>/<total_GiB>|..."
|
||||
- CPU values represent system RAM via psutil.
|
||||
- Device values represent VRAM via comfy.model_management across all non-CPU devices.
|
||||
- Device identifiers use the torch device string from get_device_list() (e.g., 'cuda:0', 'xpu:0', 'mps').
|
||||
- Values are in GiB with 2 decimals.
|
||||
"""
|
||||
"""Return single-line pipe-delimited snapshot of system and device memory usage in GiB."""
|
||||
# CPU RAM
|
||||
vm = psutil.virtual_memory()
|
||||
cpu_used_gib = vm.used / (1024.0 ** 3)
|
||||
|
||||
+3
-17
@@ -52,7 +52,6 @@ def create_safetensor_model_hash(model, caller):
|
||||
logger.debug(f"[MultiGPU DisTorch V2] Created hash for {caller}: {final_hash[:8]}...")
|
||||
return final_hash
|
||||
|
||||
|
||||
def register_patched_safetensor_modelpatcher():
|
||||
"""Register and patch the ModelPatcher for distributed safetensor loading"""
|
||||
from comfy.model_patcher import wipe_lowvram_weight, move_weight_functions
|
||||
@@ -216,12 +215,8 @@ def register_patched_safetensor_modelpatcher():
|
||||
comfy.model_patcher.ModelPatcher._distorch_patched = True
|
||||
logger.info("[MultiGPU Core Patching] Successfully patched ModelPatcher.partially_load")
|
||||
|
||||
|
||||
def _extract_clip_head_blocks(raw_block_list, compute_device):
|
||||
"""
|
||||
Helper: Identify and pre-assign CLIP head blocks to compute device.
|
||||
Returns (head_blocks, distributable_blocks, block_assignments, head_memory)
|
||||
"""
|
||||
"""Identify and pre-assign CLIP head blocks to compute device returning head_blocks, distributable_blocks, block_assignments, and head_memory."""
|
||||
head_keywords = ['embed', 'wte', 'wpe', 'token_embedding', 'position_embedding']
|
||||
head_blocks = []
|
||||
distributable_blocks = []
|
||||
@@ -238,7 +233,6 @@ def _extract_clip_head_blocks(raw_block_list, compute_device):
|
||||
|
||||
return head_blocks, distributable_blocks, block_assignments, head_memory
|
||||
|
||||
|
||||
def analyze_safetensor_loading(model_patcher, allocations_string, is_clip=False):
|
||||
"""
|
||||
Analyze and distribute safetensor model blocks across devices.
|
||||
@@ -463,7 +457,6 @@ def analyze_safetensor_loading(model_patcher, allocations_string, is_clip=False)
|
||||
"block_assignments": block_assignments
|
||||
}
|
||||
|
||||
|
||||
def parse_memory_string(mem_str):
|
||||
"""Parses a memory string (e.g., '4.0g', '512M') and returns bytes."""
|
||||
mem_str = mem_str.strip().lower()
|
||||
@@ -484,11 +477,7 @@ def parse_memory_string(mem_str):
|
||||
return val
|
||||
|
||||
def calculate_fraction_from_byte_expert_string(model_patcher, byte_str):
|
||||
"""
|
||||
Converts a user-provided byte string (e.g., "cuda:1,4gb;cpu,*") into a
|
||||
fractional VRAM allocation string that the main assignment logic can use.
|
||||
This function strictly respects device order and byte quotas.
|
||||
"""
|
||||
"""Convert byte allocation string (e.g. 'cuda:1,4gb;cpu,*') to fractional VRAM allocation string respecting device order and byte quotas."""
|
||||
raw_block_list = model_patcher._load_list()
|
||||
total_model_memory = sum(module_size for module_size, _, _, _ in raw_block_list)
|
||||
remaining_model_bytes = total_model_memory
|
||||
@@ -547,10 +536,7 @@ def calculate_fraction_from_byte_expert_string(model_patcher, byte_str):
|
||||
return allocations_string
|
||||
|
||||
def calculate_fraction_from_ratio_expert_string(model_patcher, ratio_str):
|
||||
"""
|
||||
Converts a user-provided ratio string (which describes how to split the MODEL)
|
||||
into a fraction string (which describes the fraction of DEVICE VRAM to use).
|
||||
"""
|
||||
"""Convert ratio allocation string (e.g. 'cuda:0,25%;cpu,75%') describing model split to fractional VRAM allocation string."""
|
||||
raw_block_list = model_patcher._load_list()
|
||||
total_model_memory = sum(module_size for module_size, _, _, _ in raw_block_list)
|
||||
|
||||
|
||||
@@ -206,10 +206,7 @@ def check_cpu_memory_threshold(threshold_percent=CPU_MEMORY_THRESHOLD_PERCENT):
|
||||
multigpu_memory_log("cpu_monitor", f"skip_hysteresis:{current_usage:.1f}pct")
|
||||
|
||||
def force_full_system_cleanup(reason="manual", force=True):
|
||||
"""
|
||||
Mirror ComfyUI-Manager 'Free model and node cache' by setting both flags:
|
||||
unload_models=True and free_memory=True
|
||||
"""
|
||||
"""Mirror ComfyUI-Manager 'Free model and node cache' by setting unload_models=True and free_memory=True flags."""
|
||||
vm = psutil.virtual_memory()
|
||||
pre_cpu = vm.used
|
||||
pre_models = len(mm.current_loaded_models)
|
||||
@@ -246,9 +243,7 @@ if not hasattr(mm.unload_all_models, '_mgpu_eject_distorch_patched'):
|
||||
_mgpu_original_unload_all_models = mm.unload_all_models
|
||||
|
||||
def _mgpu_patched_unload_all_models():
|
||||
"""
|
||||
Patched mm.unload_all_models with comprehensive diagnostics and fixed path alignment.
|
||||
"""
|
||||
"""Patched mm.unload_all_models with selective ejection support and comprehensive diagnostics."""
|
||||
|
||||
logger.mgpu_mm_log(f"[UNLOAD_START] Patched unload_all_models called - initial model count: {len(mm.current_loaded_models)}")
|
||||
|
||||
|
||||
@@ -21,6 +21,7 @@ class DeviceSelectorMultiGPU:
|
||||
CATEGORY = "multigpu"
|
||||
|
||||
def select_device(self, device):
|
||||
"""Select target device from available device list."""
|
||||
return (device,)
|
||||
|
||||
|
||||
@@ -38,6 +39,7 @@ class HunyuanVideoEmbeddingsAdapter:
|
||||
CATEGORY = "multigpu"
|
||||
|
||||
def adapt_embeddings(self, hyvid_embeds):
|
||||
"""Adapt HunyuanVideo embeddings to standard ComfyUI conditioning format."""
|
||||
cond = hyvid_embeds["prompt_embeds"]
|
||||
|
||||
pooled_dict = {
|
||||
@@ -73,6 +75,7 @@ class UnetLoaderGGUF:
|
||||
TITLE = "Unet Loader (GGUF)"
|
||||
|
||||
def load_unet(self, unet_name, dequant_dtype=None, patch_dtype=None, patch_on_device=None):
|
||||
"""Load GGUF format UNet model."""
|
||||
original_loader = NODE_CLASS_MAPPINGS["UnetLoaderGGUF"]()
|
||||
return original_loader.load_unet(unet_name, dequant_dtype, patch_dtype, patch_on_device)
|
||||
|
||||
@@ -110,20 +113,24 @@ class CLIPLoaderGGUF:
|
||||
|
||||
@classmethod
|
||||
def get_filename_list(s):
|
||||
"""Get combined list of CLIP and CLIP_GGUF model files."""
|
||||
files = []
|
||||
files += folder_paths.get_filename_list("clip")
|
||||
files += folder_paths.get_filename_list("clip_gguf")
|
||||
return sorted(files)
|
||||
|
||||
def load_data(self, ckpt_paths):
|
||||
"""Load CLIP model data from checkpoint paths."""
|
||||
original_loader = NODE_CLASS_MAPPINGS["CLIPLoaderGGUF"]()
|
||||
return original_loader.load_data(ckpt_paths)
|
||||
|
||||
def load_patcher(self, clip_paths, clip_type, clip_data):
|
||||
"""Create ModelPatcher for CLIP model."""
|
||||
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", device=None):
|
||||
"""Load CLIP model from GGUF or standard format."""
|
||||
original_loader = NODE_CLASS_MAPPINGS["CLIPLoaderGGUF"]()
|
||||
return original_loader.load_clip(clip_name, type)
|
||||
|
||||
@@ -144,6 +151,7 @@ class DualCLIPLoaderGGUF(CLIPLoaderGGUF):
|
||||
TITLE = "DualCLIPLoader (GGUF)"
|
||||
|
||||
def load_clip(self, clip_name1, clip_name2, type, device=None):
|
||||
"""Load dual CLIP model configuration."""
|
||||
original_loader = NODE_CLASS_MAPPINGS["DualCLIPLoaderGGUF"]()
|
||||
clip = original_loader.load_clip(clip_name1, clip_name2, type)
|
||||
clip[0].patcher.load(force_patch_weights=True)
|
||||
@@ -165,6 +173,7 @@ class TripleCLIPLoaderGGUF(CLIPLoaderGGUF):
|
||||
TITLE = "TripleCLIPLoader (GGUF)"
|
||||
|
||||
def load_clip(self, clip_name1, clip_name2, clip_name3, type="sd3"):
|
||||
"""Load triple CLIP model configuration for SD3."""
|
||||
original_loader = NODE_CLASS_MAPPINGS["TripleCLIPLoaderGGUF"]()
|
||||
return original_loader.load_clip(clip_name1, clip_name2, clip_name3, type)
|
||||
|
||||
@@ -184,6 +193,7 @@ class QuadrupleCLIPLoaderGGUF(CLIPLoaderGGUF):
|
||||
TITLE = "QuadrupleCLIPLoader (GGUF)"
|
||||
|
||||
def load_clip(self, clip_name1, clip_name2, clip_name3, clip_name4, type="stable_diffusion"):
|
||||
"""Load quadruple CLIP model configuration."""
|
||||
original_loader = NODE_CLASS_MAPPINGS["QuadrupleCLIPLoaderGGUF"]()
|
||||
return original_loader.load_clip(clip_name1, clip_name2, clip_name3, clip_name4, type)
|
||||
|
||||
@@ -207,12 +217,15 @@ class LTXVLoader:
|
||||
OUTPUT_NODE = False
|
||||
|
||||
def load(self, ckpt_name, dtype):
|
||||
"""Load LTXV model and VAE with specified precision."""
|
||||
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 ):
|
||||
"""Load LTXV UNet with device-specific configuration."""
|
||||
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):
|
||||
"""Load LTXV VAE from weights."""
|
||||
original_loader = NODE_CLASS_MAPPINGS["LTXVLoader"]()
|
||||
return original_loader._load_vae(weights, config=None)
|
||||
|
||||
@@ -239,6 +252,7 @@ class Florence2ModelLoader:
|
||||
CATEGORY = "Florence2"
|
||||
|
||||
def loadmodel(self, model, precision, attention, lora=None):
|
||||
"""Load Florence2 vision model with specified precision and attention mode."""
|
||||
original_loader = NODE_CLASS_MAPPINGS["Florence2ModelLoader"]()
|
||||
return original_loader.loadmodel(model, precision, attention, lora)
|
||||
|
||||
@@ -286,6 +300,7 @@ class DownloadAndLoadFlorence2Model:
|
||||
CATEGORY = "Florence2"
|
||||
|
||||
def loadmodel(self, model, precision, attention, lora=None):
|
||||
"""Download and load Florence2 model from HuggingFace."""
|
||||
original_loader = NODE_CLASS_MAPPINGS["DownloadAndLoadFlorence2Model"]()
|
||||
return original_loader.loadmodel(model, precision, attention, lora)
|
||||
|
||||
@@ -301,6 +316,7 @@ class CheckpointLoaderNF4:
|
||||
|
||||
|
||||
def load_checkpoint(self, ckpt_name):
|
||||
"""Load checkpoint in NF4 quantized format."""
|
||||
original_loader = NODE_CLASS_MAPPINGS["CheckpointLoaderNF4"]()
|
||||
return original_loader.load_checkpoint(ckpt_name)
|
||||
|
||||
@@ -317,6 +333,7 @@ class LoadFluxControlNet:
|
||||
CATEGORY = "XLabsNodes"
|
||||
|
||||
def loadmodel(self, model_name, controlnet_path):
|
||||
"""Load Flux ControlNet model."""
|
||||
original_loader = NODE_CLASS_MAPPINGS["LoadFluxControlNet"]()
|
||||
return original_loader.loadmodel(model_name, controlnet_path)
|
||||
|
||||
@@ -337,6 +354,7 @@ class MMAudioModelLoader:
|
||||
CATEGORY = "MMAudio"
|
||||
|
||||
def loadmodel(self, mmaudio_model, base_precision):
|
||||
"""Load MMAudio model with specified precision."""
|
||||
original_loader = NODE_CLASS_MAPPINGS["MMAudioModelLoader"]()
|
||||
return original_loader.loadmodel(mmaudio_model, base_precision)
|
||||
|
||||
@@ -364,6 +382,7 @@ class MMAudioFeatureUtilsLoader:
|
||||
CATEGORY = "MMAudio"
|
||||
|
||||
def loadmodel(self, vae_model, precision, synchformer_model, clip_model, mode, bigvgan_vocoder_model=None):
|
||||
"""Load MMAudio feature extraction utilities including VAE, Synchformer, and CLIP."""
|
||||
original_loader = NODE_CLASS_MAPPINGS["MMAudioFeatureUtilsLoader"]()
|
||||
return original_loader.loadmodel(vae_model, precision, synchformer_model, clip_model, mode, bigvgan_vocoder_model)
|
||||
|
||||
@@ -394,6 +413,7 @@ 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):
|
||||
"""Sample audio from MMAudio model with conditioning."""
|
||||
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)
|
||||
|
||||
@@ -407,6 +427,7 @@ class PulidModelLoader:
|
||||
CATEGORY = "pulid"
|
||||
|
||||
def load_model(self, pulid_file):
|
||||
"""Load PuLID identity preservation model."""
|
||||
original_loader = NODE_CLASS_MAPPINGS["PulidModelLoader"]()
|
||||
return original_loader.load_model(pulid_file)
|
||||
|
||||
@@ -424,6 +445,7 @@ class PulidInsightFaceLoader:
|
||||
CATEGORY = "pulid"
|
||||
|
||||
def load_insightface(self, provider):
|
||||
"""Load InsightFace face analysis model for PuLID."""
|
||||
original_loader = NODE_CLASS_MAPPINGS["PulidInsightFaceLoader"]()
|
||||
return original_loader.load_insightface(provider)
|
||||
|
||||
@@ -439,6 +461,7 @@ class PulidEvaClipLoader:
|
||||
CATEGORY = "pulid"
|
||||
|
||||
def load_eva_clip(self):
|
||||
"""Load EVA CLIP model for PuLID."""
|
||||
original_loader = NODE_CLASS_MAPPINGS["PulidEvaClipLoader"]()
|
||||
return original_loader.load_eva_clip()
|
||||
|
||||
@@ -473,6 +496,7 @@ 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):
|
||||
"""Load HunyuanVideo model with specified precision and quantization."""
|
||||
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)
|
||||
|
||||
@@ -498,6 +522,7 @@ class HyVideoVAELoader:
|
||||
DESCRIPTION = "Loads Hunyuan VAE model from 'ComfyUI/models/vae'"
|
||||
|
||||
def loadmodel(self, model_name, precision, compile_args=None):
|
||||
"""Load HunyuanVideo VAE model."""
|
||||
original_loader = NODE_CLASS_MAPPINGS["HyVideoVAELoader"]()
|
||||
return original_loader.loadmodel(model_name, precision, compile_args)
|
||||
|
||||
@@ -526,6 +551,7 @@ 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"):
|
||||
"""Download and load HunyuanVideo text encoder from HuggingFace."""
|
||||
original_loader = NODE_CLASS_MAPPINGS["DownloadAndLoadHyVideoTextEncoder"]()
|
||||
return original_loader.loadmodel(llm_model, clip_model, precision, apply_final_norm, hidden_state_skip_layer, quantization)
|
||||
|
||||
@@ -542,6 +568,7 @@ class UNetLoaderLP:
|
||||
TITLE = "UNet Loader (LP)"
|
||||
|
||||
def load_unet(self, unet_name):
|
||||
"""Load UNet with low-precision LoRA flag for CPU storage optimization."""
|
||||
original_loader = NODE_CLASS_MAPPINGS["UNETLoader"]()
|
||||
out = original_loader.load_unet(unet_name)
|
||||
|
||||
@@ -551,4 +578,4 @@ class UNetLoaderLP:
|
||||
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
|
||||
out[0].patcher.model._distorch_high_precision_loras = False
|
||||
|
||||
return out
|
||||
return out
|
||||
|
||||
+1
-12
@@ -16,18 +16,7 @@ logger = logging.getLogger("MultiGPU")
|
||||
# ============================================================================
|
||||
|
||||
def _create_distorch_safetensor_v2_override(cls, device_param_name, device_setter_func, apply_device_kwarg_workaround):
|
||||
"""
|
||||
Internal factory function - creates DisTorch 2.0 override class with parameterized behavior.
|
||||
|
||||
Args:
|
||||
cls: The base class to override
|
||||
device_param_name: Parameter name ("compute_device" or "device")
|
||||
device_setter_func: Function to call for device setting
|
||||
apply_device_kwarg_workaround: If True, sets kwargs['device'] = 'default' for ComfyUI compatibility
|
||||
|
||||
Returns:
|
||||
Override class with specified behavior
|
||||
"""
|
||||
"""Internal factory function creating DisTorch2 override class with parameterized device selection behavior."""
|
||||
from .distorch_2 import (
|
||||
register_patched_safetensor_modelpatcher,
|
||||
safetensor_allocation_store,
|
||||
|
||||
Reference in New Issue
Block a user