Standardize doc strings and make PEP 257 compliant

This commit is contained in:
John Pollock
2025-09-30 09:34:41 -05:00
parent fc2a732419
commit 62752d1bbf
7 changed files with 55 additions and 137 deletions
+8 -1
View File
@@ -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())}")
+2 -5
View File
@@ -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
View File
@@ -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
View File
@@ -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)
+2 -7
View File
@@ -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)}")
+28 -1
View File
@@ -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
View File
@@ -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,