diff --git a/__init__.py b/__init__.py index cbbd98e..10fc88e 100644 --- a/__init__.py +++ b/__init__.py @@ -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())}") \ No newline at end of file +logger.info(f"[MultiGPU] Registration complete. Final mappings: {', '.join(NODE_CLASS_MAPPINGS.keys())}") diff --git a/checkpoint_multigpu.py b/checkpoint_multigpu.py index d893841..6d0b8c9 100644 --- a/checkpoint_multigpu.py +++ b/checkpoint_multigpu.py @@ -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')) diff --git a/device_utils.py b/device_utils.py index d48962a..a86cf84 100644 --- a/device_utils.py +++ b/device_utils.py @@ -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=|cpu=/|=/|..." - - 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) diff --git a/distorch_2.py b/distorch_2.py index e6e62ca..250a07e 100644 --- a/distorch_2.py +++ b/distorch_2.py @@ -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) diff --git a/model_management_mgpu.py b/model_management_mgpu.py index 6a26149..9a890ef 100644 --- a/model_management_mgpu.py +++ b/model_management_mgpu.py @@ -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)}") diff --git a/nodes.py b/nodes.py index fac43f8..a5ad9b4 100644 --- a/nodes.py +++ b/nodes.py @@ -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 \ No newline at end of file + return out diff --git a/wrappers.py b/wrappers.py index 0da36ff..483adc9 100644 --- a/wrappers.py +++ b/wrappers.py @@ -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,