This commit represents a significant architectural refactoring to improve code organization, eliminate redundancy, and fix a critical bug in wrapper functions. Net reduction of 531 lines while improving maintainability and fixing functionality. ## wrappers.py (NEW FILE: +531 lines) - Created dedicated module for ALL node wrapper/override functions - Consolidated 10 wrapper types from 3 different files into single location: * DisTorch V2 SafeTensor wrappers (factory + 3 implementations) * DisTorch V1 legacy wrappers (4 GGUF/CLIP wrappers, rewritten to call V2 backend) * Standard MultiGPU wrappers (3 device selection wrappers) - CRITICAL FIX: All wrappers now strip MultiGPU-specific parameters before calling original ComfyUI functions (fixes CheckpointLoaderSimple TypeError) - Improved architecture: clear separation between wrapper UI and backend logic ## distorch.py (DELETED: -529 lines) - Removed entire legacy DisTorch V1 file - All V1 wrapper functions moved to wrappers.py and rewritten to call V2 backend - Backend allocation functions no longer needed (V2 backend handles all cases) - Eliminates code duplication and maintenance burden ## distorch_2.py (-409 lines) - Removed duplicate _create_distorch_safetensor_v2_override factory function (was incorrectly present in both distorch_2.py and wrappers.py) - Removed 3 wrapper export functions (moved to wrappers.py) - File now contains ONLY backend logic: * register_patched_safetensor_modelpatcher() * analyze_safetensor_loading() and analyze_safetensor_loading_clip() * calculate_safetensor_vvram_allocation() * Allocation stores and model hash functions - Added clear documentation comment about wrapper migration ## __init__.py (-230 lines) - Removed 3 local wrapper function definitions (moved to wrappers.py) - Removed soft_empty_cache_distorch2_patched (moved to device_utils.py) - Removed all distorch.py imports (file deleted) - Added imports from new wrappers.py module (10 wrapper functions) - Updated imports from distorch_2.py (backend functions only, no wrappers) - Improved architecture: __init__.py now focused on initialization and registration ## device_utils.py (+68 lines) - Moved soft_empty_cache_distorch2_patched() from __init__.py - Added comprehensive memory management patch in architecturally correct location - Patch includes: * DisTorch2 detection and multi-device VRAM management * Adaptive CPU memory threshold checking * Force flag support for executor cache reset (Manager parity) - Applied patch at module level: mm.soft_empty_cache = soft_empty_cache_distorch2_patched - Behavior preserved: patch still executes when device_utils is imported by __init__.py ## nodes.py (-30 lines) - Removed unused wrapper function imports - Cleaned up import statements to reflect new architecture ## Impact Summary - Improved architecture: Clear separation between wrappers (UI) and backend (logic) - Eliminated distorch.py: Reduced from 3 files to 2 (wrappers.py + distorch_2.py) - Net code reduction: 531 lines removed while adding functionality - Better maintainability: Single source of truth for all wrapper functions - Preserved behavior: All patches execute correctly, no functional changes ## Breaking Changes None - this is a pure refactor with no API or behavioral changes.
404 lines
16 KiB
Python
404 lines
16 KiB
Python
"""
|
|
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
|
|
import psutil
|
|
import comfy.model_management as mm
|
|
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.
|
|
This includes all device types supported by ComfyUI core.
|
|
Results are cached after first call since devices don't change during runtime.
|
|
|
|
Returns a comprehensive list of all available devices across all types:
|
|
- CPU (always available)
|
|
- CUDA devices (NVIDIA GPUs)
|
|
- XPU devices (Intel GPUs)
|
|
- NPU devices (Ascend NPUs from Huawei)
|
|
- MLU devices (Cambricon MLUs)
|
|
- MPS device (Apple Metal)
|
|
- DirectML devices (Windows DirectML)
|
|
- CoreX/IXUCA devices
|
|
"""
|
|
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
|
|
|
|
if hasattr(torch, "xpu") and hasattr(torch.xpu, "is_available") and torch.xpu.is_available():
|
|
device_count = torch.xpu.device_count()
|
|
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():
|
|
device_count = torch.npu.device_count()
|
|
devs += [f"npu:{i}" for i in range(device_count)]
|
|
logger.debug(f"[MultiGPU_Device_Utils] Found {device_count} NPU device(s)")
|
|
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():
|
|
device_count = torch.mlu.device_count()
|
|
devs += [f"mlu:{i}" for i in range(device_count)]
|
|
logger.debug(f"[MultiGPU_Device_Utils] Found {device_count} MLU device(s)")
|
|
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()
|
|
if adapter_count > 0:
|
|
devs += [f"directml:{i}" for i in range(adapter_count)]
|
|
logger.debug(f"[MultiGPU_Device_Utils] Found {adapter_count} DirectML adapter(s)")
|
|
except ImportError:
|
|
pass
|
|
|
|
# IXUCA/CoreX devices (special accelerator)
|
|
try:
|
|
if hasattr(torch, "corex"):
|
|
if hasattr(torch.corex, "device_count"):
|
|
device_count = torch.corex.device_count()
|
|
devs += [f"corex:{i}" for i in range(device_count)]
|
|
logger.debug(f"[MultiGPU_Device_Utils] Found {device_count} CoreX device(s)")
|
|
else:
|
|
devs.append("corex:0")
|
|
logger.debug("[MultiGPU_Device_Utils] Found CoreX device")
|
|
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
|
|
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.
|
|
"""
|
|
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")
|
|
"""
|
|
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
|
|
"""
|
|
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
|
|
from .model_management_mgpu import multigpu_memory_log
|
|
|
|
logger.mgpu_mm_log("soft_empty_cache_multigpu: starting GC and multi-device cache clear")
|
|
|
|
gc.collect()
|
|
|
|
# Clear cache for ALL devices (not just ComfyUI's single device)
|
|
all_devices = get_device_list()
|
|
logger.mgpu_mm_log(f"soft_empty_cache_multigpu: devices to clear = {all_devices}")
|
|
|
|
# Check global availability first to avoid unnecessary iteration if backend is missing
|
|
is_cuda_available = hasattr(torch, "cuda") and hasattr(torch.cuda, "is_available") and torch.cuda.is_available()
|
|
|
|
for device_str in all_devices:
|
|
if device_str.startswith("cuda:"):
|
|
if is_cuda_available:
|
|
device_idx = int(device_str.split(":")[1])
|
|
logger.mgpu_mm_log(f"Clearing CUDA cache on {device_str} (idx={device_idx})")
|
|
multigpu_memory_log("general", f"pre-empty:{device_str}")
|
|
with torch.cuda.device(device_idx):
|
|
torch.cuda.empty_cache()
|
|
if hasattr(torch.cuda, "ipc_collect"):
|
|
torch.cuda.ipc_collect()
|
|
logger.mgpu_mm_log(f"Cleared CUDA cache (and IPC if available) on {device_str}")
|
|
multigpu_memory_log("general", f"post-empty:{device_str}")
|
|
|
|
elif device_str == "mps":
|
|
if hasattr(torch, "mps") and hasattr(torch.mps, "empty_cache"):
|
|
logger.mgpu_mm_log("Clearing MPS cache")
|
|
multigpu_memory_log("general", f"pre-empty:{device_str}")
|
|
torch.mps.empty_cache()
|
|
logger.mgpu_mm_log("Cleared MPS cache")
|
|
multigpu_memory_log("general", f"post-empty:{device_str}")
|
|
|
|
elif device_str.startswith("xpu:"):
|
|
if hasattr(torch, "xpu") and hasattr(torch.xpu, "empty_cache"):
|
|
logger.mgpu_mm_log(f"Clearing XPU cache on {device_str}")
|
|
multigpu_memory_log("general", f"pre-empty:{device_str}")
|
|
torch.xpu.empty_cache()
|
|
logger.mgpu_mm_log(f"Cleared XPU cache on {device_str}")
|
|
multigpu_memory_log("general", f"post-empty:{device_str}")
|
|
|
|
elif device_str.startswith("npu:"):
|
|
if hasattr(torch, "npu") and hasattr(torch.npu, "empty_cache"):
|
|
logger.mgpu_mm_log(f"Clearing NPU cache on {device_str}")
|
|
multigpu_memory_log("general", f"pre-empty:{device_str}")
|
|
torch.npu.empty_cache()
|
|
logger.mgpu_mm_log(f"Cleared NPU cache on {device_str}")
|
|
multigpu_memory_log("general", f"post-empty:{device_str}")
|
|
|
|
elif device_str.startswith("mlu:"):
|
|
if hasattr(torch, "mlu") and hasattr(torch.mlu, "empty_cache"):
|
|
logger.mgpu_mm_log(f"Clearing MLU cache on {device_str}")
|
|
multigpu_memory_log("general", f"pre-empty:{device_str}")
|
|
torch.mlu.empty_cache()
|
|
logger.mgpu_mm_log(f"Cleared MLU cache on {device_str}")
|
|
multigpu_memory_log("general", f"post-empty:{device_str}")
|
|
|
|
elif device_str.startswith("corex:"):
|
|
if hasattr(torch, "corex") and hasattr(torch.corex, "empty_cache"):
|
|
logger.mgpu_mm_log(f"Clearing CoreX cache on {device_str}")
|
|
multigpu_memory_log("general", f"pre-empty:{device_str}")
|
|
torch.corex.empty_cache()
|
|
logger.mgpu_mm_log(f"Cleared CoreX cache on {device_str}")
|
|
multigpu_memory_log("general", f"post-empty:{device_str}")
|
|
|
|
multigpu_memory_log("general", "post-soft-empty")
|
|
|
|
|
|
# ==========================================================================================
|
|
# Comprehensive Memory Management (VRAM + CPU + Store Pruning)
|
|
# ==========================================================================================
|
|
|
|
logger.info("[MultiGPU Core Patching] Patching mm.soft_empty_cache for Comprehensive Memory Management (VRAM + CPU + Store Pruning)")
|
|
|
|
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).
|
|
"""
|
|
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
|
|
|
|
multigpu_memory_log("patched_soft_empty", f"start:force={force}")
|
|
is_distorch_active = False
|
|
|
|
# Detect DisTorch2-managed models
|
|
logger.mgpu_mm_log(f"[DETECT_DEBUG] Checking DisTorch2 active status - loaded models: {len(mm.current_loaded_models)}, store entries: {len(safetensor_allocation_store)}")
|
|
|
|
for i, lm in enumerate(mm.current_loaded_models):
|
|
mp = lm.model # weakref call to ModelPatcher
|
|
if mp is not None:
|
|
try:
|
|
model_hash = create_safetensor_model_hash(mp, "cache_patch_check")
|
|
in_store = model_hash in safetensor_allocation_store
|
|
alloc_value = safetensor_allocation_store.get(model_hash, "")
|
|
model_name = type(getattr(mp, 'model', mp)).__name__
|
|
unload_distorch_model = getattr(getattr(mp, 'model', None), '_mgpu_unload_distorch_model', False)
|
|
|
|
logger.mgpu_mm_log(f"[DETECT_DEBUG] Model {i}: {model_name}, hash={model_hash[:8]}, in_store={in_store}, alloc_value='{alloc_value}', unload_distorch_model={unload_distorch_model}")
|
|
|
|
if in_store and alloc_value:
|
|
is_distorch_active = True
|
|
logger.mgpu_mm_log(f"[DETECT_DEBUG] DisTorch2 ACTIVE detected on model: {model_name}")
|
|
break
|
|
except Exception as e:
|
|
logger.mgpu_mm_log(f"[DETECT_DEBUG] Model {i}: Error during detection - {e}")
|
|
|
|
logger.mgpu_mm_log(f"[DETECT_DEBUG] Final DisTorch2 active status: {is_distorch_active}")
|
|
|
|
# Phase 2: adaptive CPU memory management
|
|
check_cpu_memory_threshold()
|
|
|
|
# VRAM allocator management
|
|
if is_distorch_active:
|
|
logger.mgpu_mm_log("DisTorch2 active: clearing allocator caches on all devices (VRAM)")
|
|
soft_empty_cache_multigpu()
|
|
else:
|
|
logger.mgpu_mm_log("DisTorch2 not active: delegating allocator cache clear (VRAM) to original mm.soft_empty_cache")
|
|
original_soft_empty_cache(force)
|
|
# Optional: return CPU heap to OS (not part of Comfy Core)
|
|
|
|
# Phase 1/3: forced executor reset mirrors ComfyUI 'Free memory' semantics
|
|
if force:
|
|
logger.mgpu_mm_log("Force flag active: triggering executor cache reset (CPU)")
|
|
trigger_executor_cache_reset(reason="forced_soft_empty", force=True)
|
|
multigpu_memory_log("patched_soft_empty", "end")
|
|
|
|
mm.soft_empty_cache = soft_empty_cache_distorch2_patched
|
|
|
|
# ==========================================================================================
|
|
# Memory Inspection Utilities
|
|
# ==========================================================================================
|
|
|
|
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.
|
|
"""
|
|
# CPU RAM
|
|
vm = psutil.virtual_memory()
|
|
cpu_used_gib = vm.used / (1024.0 ** 3)
|
|
cpu_total_gib = vm.total / (1024.0 ** 3)
|
|
|
|
segments = [f"tag={tag}", f"cpu={cpu_used_gib:.2f}/{cpu_total_gib:.2f}"]
|
|
|
|
# Enumerate non-CPU devices
|
|
devices = [d for d in get_device_list() if d != "cpu"]
|
|
|
|
# Append per-device VRAM used/total
|
|
for dev_str in devices:
|
|
device = torch.device(dev_str)
|
|
total = mm.get_total_memory(device)
|
|
free_info = mm.get_free_memory(device, torch_free_too=True)
|
|
# free_info may be a tuple (system_free, torch_cache_free) or a single value
|
|
if isinstance(free_info, tuple):
|
|
system_free = free_info[0]
|
|
else:
|
|
system_free = free_info
|
|
used = max(0, (total or 0) - (system_free or 0))
|
|
|
|
used_gib = used / (1024.0 ** 3)
|
|
total_gib = (total or 0) / (1024.0 ** 3)
|
|
if total_gib > 0:
|
|
segments.append(f"{dev_str}={used_gib:.2f}/{total_gib:.2f}")
|
|
|
|
return "|".join(segments)
|