Replace try-catch wrapped comfyui_memory_load calls with streamlined multigpu_memory_log function. Removes exception handling overhead and uses consistent config hash identifiers for UNet, VAE, and CLIP model loading phases.
752 lines
30 KiB
Python
752 lines
30 KiB
Python
"""
|
|
Device detection, management, and inspection utilities for ComfyUI-MultiGPU.
|
|
Single source of truth for all device enumeration, compatibility checks, and state inspection.
|
|
Handles all device types supported by ComfyUI core.
|
|
"""
|
|
|
|
import torch
|
|
import logging
|
|
import hashlib
|
|
import psutil
|
|
import comfy.model_management as mm
|
|
|
|
logger = logging.getLogger("MultiGPU")
|
|
|
|
# Module-level cache for device list (populated once on first call)
|
|
_DEVICE_LIST_CACHE = None
|
|
|
|
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)
|
|
try:
|
|
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)")
|
|
except Exception as e:
|
|
logger.debug(f"[MultiGPU_Device_Utils] CUDA detection failed: {e}")
|
|
|
|
# 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
|
|
try:
|
|
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)")
|
|
except Exception as e:
|
|
logger.debug(f"[MultiGPU_Device_Utils] XPU detection failed: {e}")
|
|
|
|
# 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 Exception as e:
|
|
logger.debug(f"[MultiGPU_Device_Utils] NPU detection failed: {e}")
|
|
|
|
# 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 Exception as e:
|
|
logger.debug(f"[MultiGPU_Device_Utils] MLU detection failed: {e}")
|
|
|
|
# MPS device (Apple Metal - single device only)
|
|
try:
|
|
if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
|
|
devs.append("mps")
|
|
logger.debug("[MultiGPU_Device_Utils] Found MPS device")
|
|
except Exception as e:
|
|
logger.debug(f"[MultiGPU_Device_Utils] MPS detection failed: {e}")
|
|
|
|
# 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 Exception as e:
|
|
logger.debug(f"[MultiGPU_Device_Utils] DirectML detection failed: {e}")
|
|
|
|
# IXUCA/CoreX devices (special accelerator)
|
|
try:
|
|
if hasattr(torch, "corex"):
|
|
# CoreX typically exposes single device, but check if there's a count method
|
|
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 Exception as e:
|
|
logger.debug(f"[MultiGPU_Device_Utils] CoreX detection failed: {e}")
|
|
|
|
# Cache the result for future calls
|
|
_DEVICE_LIST_CACHE = devs
|
|
|
|
# Log only once when initially populated
|
|
logger.info(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
|
|
try:
|
|
if torch.cuda.is_available():
|
|
return True
|
|
except:
|
|
pass
|
|
|
|
# Check XPU (Intel GPU)
|
|
try:
|
|
if hasattr(torch, "xpu") and torch.xpu.is_available():
|
|
return True
|
|
except:
|
|
pass
|
|
|
|
# Check NPU (Ascend)
|
|
try:
|
|
import torch_npu
|
|
if hasattr(torch, "npu") and torch.npu.is_available():
|
|
return True
|
|
except:
|
|
pass
|
|
|
|
# Check MLU (Cambricon)
|
|
try:
|
|
import torch_mlu
|
|
if hasattr(torch, "mlu") and torch.mlu.is_available():
|
|
return True
|
|
except:
|
|
pass
|
|
|
|
# Check MPS (Apple Metal)
|
|
try:
|
|
if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
|
|
return True
|
|
except:
|
|
pass
|
|
|
|
# Check DirectML
|
|
try:
|
|
import torch_directml
|
|
if torch_directml.device_count() > 0:
|
|
return True
|
|
except:
|
|
pass
|
|
|
|
# Check CoreX/IXUCA
|
|
try:
|
|
if hasattr(torch, "corex"):
|
|
return True
|
|
except:
|
|
pass
|
|
|
|
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
|
|
|
|
|
|
def soft_empty_cache_multigpu():
|
|
"""
|
|
Replicate ComfyUI's cache clearing but for ALL devices in MultiGPU.
|
|
MultiGPU adaptation of ComfyUI's soft_empty_cache() functionality.
|
|
Uses context managers to ensure the calling thread's device context is restored.
|
|
"""
|
|
import gc
|
|
|
|
logger.info("[MultiGPU_Device_Utils] soft_empty_cache_multigpu: starting GC and multi-device cache clear")
|
|
# Record pre-GC snapshot for general system view
|
|
multigpu_memory_log("general", "pre-soft-empty")
|
|
|
|
# Python GC (same as all implementations)
|
|
gc.collect()
|
|
logger.info("[MultiGPU_Device_Utils] soft_empty_cache_multigpu: garbage collection complete")
|
|
|
|
# Clear cache for ALL devices (not just ComfyUI's single device)
|
|
all_devices = get_device_list()
|
|
logger.info(f"[MultiGPU_Device_Utils] 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])
|
|
# Use context manager for safe switching and automatic restoration
|
|
logger.info(f"[MultiGPU_Device_Utils] Clearing CUDA cache on {device_str} (idx={device_idx})")
|
|
with torch.cuda.device(device_idx):
|
|
torch.cuda.empty_cache()
|
|
if hasattr(torch.cuda, "ipc_collect"):
|
|
torch.cuda.ipc_collect() # ComfyUI's CUDA optimization
|
|
logger.info(f"[MultiGPU_Device_Utils] Cleared CUDA cache (and IPC if available) on {device_str}")
|
|
|
|
elif device_str == "mps":
|
|
if hasattr(torch, "mps") and hasattr(torch.mps, "empty_cache"):
|
|
logger.info("[MultiGPU_Device_Utils] Clearing MPS cache")
|
|
torch.mps.empty_cache()
|
|
logger.info("[MultiGPU_Device_Utils] Cleared MPS cache")
|
|
|
|
elif device_str.startswith("xpu:"):
|
|
if hasattr(torch, "xpu") and hasattr(torch.xpu, "empty_cache"):
|
|
logger.info(f"[MultiGPU_Device_Utils] Clearing XPU cache on {device_str}")
|
|
torch.xpu.empty_cache()
|
|
logger.info(f"[MultiGPU_Device_Utils] Cleared XPU cache on {device_str}")
|
|
|
|
elif device_str.startswith("npu:"):
|
|
if hasattr(torch, "npu") and hasattr(torch.npu, "empty_cache"):
|
|
logger.info(f"[MultiGPU_Device_Utils] Clearing NPU cache on {device_str}")
|
|
torch.npu.empty_cache()
|
|
logger.info(f"[MultiGPU_Device_Utils] Cleared NPU cache on {device_str}")
|
|
|
|
elif device_str.startswith("mlu:"):
|
|
if hasattr(torch, "mlu") and hasattr(torch.mlu, "empty_cache"):
|
|
logger.info(f"[MultiGPU_Device_Utils] Clearing MLU cache on {device_str}")
|
|
torch.mlu.empty_cache()
|
|
logger.info(f"[MultiGPU_Device_Utils] Cleared MLU cache on {device_str}")
|
|
|
|
elif device_str.startswith("corex:"):
|
|
if hasattr(torch, "corex") and hasattr(torch.corex, "empty_cache"):
|
|
logger.info(f"[MultiGPU_Device_Utils] Clearing CoreX cache on {device_str}")
|
|
torch.corex.empty_cache()
|
|
logger.info(f"[MultiGPU_Device_Utils] Cleared CoreX cache on {device_str}")
|
|
|
|
# Record post-GC snapshot for general system view
|
|
multigpu_memory_log("general", "post-soft-empty")
|
|
|
|
|
|
|
|
def _bytes_to_gib(b: int) -> float:
|
|
"""Convert bytes to GiB as a float."""
|
|
return float(b) / (1024.0 ** 3)
|
|
|
|
|
|
def comfyui_memory_load(tag: str) -> str:
|
|
"""
|
|
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 = _bytes_to_gib(vm.used)
|
|
cpu_total_gib = _bytes_to_gib(vm.total)
|
|
|
|
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 = _bytes_to_gib(used)
|
|
total_gib = _bytes_to_gib(total or 0)
|
|
if total_gib > 0:
|
|
segments.append(f"{dev_str}={used_gib:.2f}/{total_gib:.2f}")
|
|
|
|
return "|".join(segments)
|
|
|
|
|
|
# ==========================================================================================
|
|
# Delta-capable memory logging (identifier + tag) with timestamped series
|
|
# ==========================================================================================
|
|
|
|
from datetime import datetime, timezone
|
|
|
|
# Stores the last snapshot per identifier: identifier -> (last_tag, snapshot_map)
|
|
# snapshot_map: device_str -> (used_bytes, total_bytes)
|
|
_MEM_SNAPSHOT_LAST = {}
|
|
|
|
# Full chronological series per identifier: identifier -> list[(timestamp, tag, snapshot_map)]
|
|
_MEM_SNAPSHOT_SERIES = {}
|
|
|
|
|
|
def _capture_memory_snapshot() -> dict[str, tuple[int, int]]:
|
|
"""
|
|
Capture an absolute memory snapshot for CPU and all non-CPU devices.
|
|
Values are returned in bytes (used, total) for each device string key.
|
|
"""
|
|
snapshot: dict[str, tuple[int, int]] = {}
|
|
|
|
# CPU
|
|
vm = psutil.virtual_memory()
|
|
snapshot["cpu"] = (vm.used, vm.total)
|
|
|
|
# Non-CPU devices
|
|
devices = [d for d in get_device_list() if d != "cpu"]
|
|
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)
|
|
system_free = free_info[0] if isinstance(free_info, tuple) else free_info
|
|
used = max(0, (total or 0) - (system_free or 0))
|
|
snapshot[dev_str] = (used, total or 0)
|
|
|
|
return snapshot
|
|
|
|
|
|
def _format_delta_gib(delta_bytes: int) -> str:
|
|
"""Format a signed GiB delta with two decimals."""
|
|
gib = _bytes_to_gib(abs(delta_bytes))
|
|
sign = "+" if delta_bytes >= 0 else "-"
|
|
return f"{sign}{gib:.2f}"
|
|
|
|
|
|
def memory_print_summary(log: logging.Logger = logger):
|
|
"""
|
|
Print the entire run as absolute actuals with timestamps for each identifier.
|
|
One line per recorded snapshot in insertion order.
|
|
Format:
|
|
YYYY-MM-DDTHH:MM:SS.mmmZ identifier tag | cpu=U/T | cuda:0=U/T | ...
|
|
(GiB values, two decimals)
|
|
"""
|
|
from . import logger as mgpu_logger
|
|
|
|
# Stable identifier order for readability
|
|
for identifier in sorted(_MEM_SNAPSHOT_SERIES.keys()):
|
|
series = _MEM_SNAPSHOT_SERIES[identifier]
|
|
if not series:
|
|
continue
|
|
mgpu_logger.memory(f"=== memory summary: {identifier} ===")
|
|
for ts, tag, snap in series:
|
|
# Build device list (cpu first, then sorted devices)
|
|
parts = []
|
|
# CPU
|
|
cpu_used, cpu_total = snap.get("cpu", (0, 0))
|
|
parts.append(f"cpu={_bytes_to_gib(cpu_used):.2f}/{_bytes_to_gib(cpu_total):.2f}")
|
|
# Non-CPU (sorted)
|
|
devs = sorted([k for k in snap.keys() if k != "cpu"])
|
|
for dev in devs:
|
|
used, total = snap[dev]
|
|
parts.append(f"{dev}={_bytes_to_gib(used):.2f}/{_bytes_to_gib(total):.2f}")
|
|
ts_str = ts.strftime("%Y-%m-%dT%H:%M:%S.%f")[:-3] + "Z"
|
|
mgpu_logger.memory(f"{ts_str} {identifier} {tag} | " + " | ".join(parts))
|
|
|
|
|
|
def multigpu_memory_log(identifier: str, tag: str, log: logging.Logger = logger):
|
|
"""
|
|
Record a timestamped memory snapshot for the given identifier and tag.
|
|
- INFO: per-device deltas vs. previous snapshot for the same identifier (GiB, signed, no totals).
|
|
- DEBUG: absolute snapshot string via comfyui_memory_load(tag) prefixed by identifier.
|
|
- Special identifier: 'print_summary' will dump the entire series as actuals with timestamps.
|
|
"""
|
|
from . import logger as mgpu_logger
|
|
|
|
if identifier == "print_summary":
|
|
memory_print_summary(log=log)
|
|
return
|
|
|
|
# Capture current snapshot and timestamp
|
|
ts = datetime.now(timezone.utc)
|
|
curr = _capture_memory_snapshot()
|
|
|
|
# Append to full series
|
|
series = _MEM_SNAPSHOT_SERIES.get(identifier)
|
|
if series is None:
|
|
series = []
|
|
_MEM_SNAPSHOT_SERIES[identifier] = series
|
|
series.append((ts, tag, curr))
|
|
|
|
# Compute and log delta vs last
|
|
if identifier in _MEM_SNAPSHOT_LAST:
|
|
prev_tag, prev = _MEM_SNAPSHOT_LAST[identifier]
|
|
# Union of device keys
|
|
keys = set(prev.keys()) | set(curr.keys())
|
|
# Stable order: cpu first, then sorted devices
|
|
ordered = ["cpu"] + sorted([k for k in keys if k != "cpu"])
|
|
parts = []
|
|
for k in ordered:
|
|
p_used, _p_tot = prev.get(k, (0, curr.get(k, (0, 0))[1]))
|
|
c_used, _c_tot = curr.get(k, (0, prev.get(k, (0, 0))[1]))
|
|
delta = c_used - p_used
|
|
parts.append(f"{k}={_format_delta_gib(delta)}")
|
|
mgpu_logger.memory(f"{identifier} {tag} - {prev_tag}: " + " | ".join(parts))
|
|
else:
|
|
# Baseline vs zero
|
|
keys = set(curr.keys())
|
|
ordered = ["cpu"] + sorted([k for k in keys if k != "cpu"])
|
|
parts = []
|
|
for k in ordered:
|
|
c_used, _c_tot = curr.get(k, (0, 0))
|
|
parts.append(f"{k}=+{_bytes_to_gib(c_used):.2f}")
|
|
mgpu_logger.memory(f"{identifier} {tag} - <baseline>: " + " | ".join(parts))
|
|
|
|
# DEBUG absolute
|
|
mgpu_logger.memory(f"{identifier}, {comfyui_memory_load(tag)}")
|
|
|
|
# Update last snapshot
|
|
_MEM_SNAPSHOT_LAST[identifier] = (tag, curr)
|
|
|
|
|
|
# ==========================================================================================
|
|
# Model Management Inspection Utilities (End-to-End Tracking)
|
|
# ==========================================================================================
|
|
|
|
def create_model_identifier(model_patcher):
|
|
"""Creates a concise, unique identifier for a model patcher based on type and size."""
|
|
if not model_patcher or not model_patcher.model:
|
|
return "N/A (Detached)"
|
|
|
|
model = model_patcher.model
|
|
model_type = type(model).__name__
|
|
|
|
# Try the fast path first (using size calculated by ModelPatcher)
|
|
try:
|
|
model_size = model_patcher.model_size()
|
|
except Exception:
|
|
model_size = 0
|
|
|
|
# If the fast path fails or returns 0, perform a safe deep inspection
|
|
if model_size == 0:
|
|
try:
|
|
# Safely inspect parameters without triggering hooks/loads
|
|
with model_patcher.use_ejected(skip_and_inject_on_exit_only=True):
|
|
# We must iterate parameters() AND buffers() as both consume memory
|
|
params = list(model.parameters()) + list(model.buffers())
|
|
# Use data_ptr to handle potential weight tying/shared tensors correctly
|
|
seen_tensors = set()
|
|
for p in params:
|
|
if p.data_ptr() not in seen_tensors:
|
|
model_size += p.numel() * p.element_size()
|
|
seen_tensors.add(p.data_ptr())
|
|
except Exception as e:
|
|
logger.debug(f"[MultiGPU_Inspection] Error during safe size calculation for identifier: {e}")
|
|
return f"{model_type} (ID_Err)"
|
|
|
|
# Create a hash based on type and calculated size
|
|
identifier = f"{model_type}_{model_size}"
|
|
model_hash = hashlib.sha256(identifier.encode()).hexdigest()
|
|
return f"{model_type} ({model_hash[:8]})"
|
|
|
|
|
|
def analyze_tensor_locations(model_patcher):
|
|
"""
|
|
Analyzes the physical device placement of model tensors (parameters and buffers).
|
|
This provides the Ground Truth location of the data, handling shared weights correctly.
|
|
"""
|
|
device_summary = {}
|
|
seen_tensors = set()
|
|
total_memory = 0
|
|
|
|
if not model_patcher or not model_patcher.model:
|
|
return {"error": "Model not available"}, 0
|
|
|
|
model = model_patcher.model
|
|
|
|
# Crucial: Use the ejector to ensure we can access the model weights safely
|
|
# without interfering with injections, hooks, or triggering unintended loads (like in standard LowVRAM mode).
|
|
try:
|
|
with model_patcher.use_ejected(skip_and_inject_on_exit_only=True):
|
|
# Helper to process tensors (parameters or buffers)
|
|
def process_tensor(tensor):
|
|
nonlocal total_memory
|
|
# Use data_ptr() for unique identification of the underlying memory
|
|
if tensor.data_ptr() in seen_tensors:
|
|
return
|
|
seen_tensors.add(tensor.data_ptr())
|
|
|
|
if tensor.numel() > 0:
|
|
tensor_mem = tensor.numel() * tensor.element_size()
|
|
total_memory += tensor_mem
|
|
|
|
if hasattr(tensor, 'device'):
|
|
device = str(tensor.device)
|
|
else:
|
|
# Handle cases like NF4 quantization or other custom tensors
|
|
device = "Unknown/Managed"
|
|
|
|
if device not in device_summary:
|
|
device_summary[device] = {'tensors': 0, 'memory': 0}
|
|
|
|
device_summary[device]['tensors'] += 1
|
|
device_summary[device]['memory'] += tensor_mem
|
|
|
|
# Iterate over all parameters (weights, biases)
|
|
for param in model.parameters():
|
|
process_tensor(param)
|
|
|
|
# Iterate over all buffers (like batch norm running stats)
|
|
for buffer in model.buffers():
|
|
process_tensor(buffer)
|
|
|
|
except Exception as e:
|
|
logger.error(f"[MultiGPU_Inspection] Error during tensor location analysis: {e}")
|
|
return {"error": str(e)}, 0
|
|
|
|
return device_summary, total_memory
|
|
|
|
|
|
def inspect_model_management_state(context_description=""):
|
|
"""
|
|
Provides a detailed, structured overview of the current state of ComfyUI's model management,
|
|
including memory usage across all devices and the status, location, and patching of all loaded models.
|
|
|
|
Call this function anywhere in the code to get an immediate snapshot of the system state.
|
|
"""
|
|
|
|
# Ensure logger configuration (handles calls before full MultiGPU init if needed)
|
|
if not logger.handlers:
|
|
handler = logging.StreamHandler()
|
|
formatter = logging.Formatter('%(message)s')
|
|
handler.setFormatter(formatter)
|
|
logger.addHandler(handler)
|
|
# Default to INFO if log level isn't set by main __init__.py
|
|
if logger.level == logging.NOTSET:
|
|
logger.setLevel(logging.INFO)
|
|
|
|
# We inspect the state without forcing GC or cache clearing, which might alter the state we want to observe.
|
|
|
|
logger.info("\n" + "=" * 100)
|
|
logger.info(f" INSPECTION: ComfyUI Model Management State [Context: {context_description}]")
|
|
logger.info("=" * 100)
|
|
|
|
# 1. Device Memory Overview
|
|
# Provides context on available resources across the system.
|
|
logger.info("--- [1] System Device Memory Overview (GB) ---")
|
|
# Sys Free: Memory available to the OS. Torch Alloc: Memory reserved by PyTorch (Active + Cache).
|
|
fmt_mem = "{:<12} | {:>10} | {:>10} | {:>10} | {:>15}"
|
|
logger.info(fmt_mem.format("Device", "Total", "Sys Free", "Used", "Torch Alloc"))
|
|
logger.info("-" * 70)
|
|
|
|
all_devices = get_device_list()
|
|
# Sort devices for consistent display (CPU last)
|
|
sorted_devices = sorted(all_devices, key=lambda d: (d == 'cpu', d))
|
|
|
|
for dev_str in sorted_devices:
|
|
try:
|
|
device = torch.device(dev_str)
|
|
|
|
if dev_str == "cpu":
|
|
vm = psutil.virtual_memory()
|
|
mem_total, mem_free_sys, mem_used = vm.total, vm.available, vm.used
|
|
torch_alloc = 0 # Difficult to track accurately for CPU globally
|
|
else:
|
|
# Use ComfyUI's management functions which account for different backends (CUDA, XPU, etc.)
|
|
mem_total = mm.get_total_memory(device)
|
|
|
|
# get_free_memory returns (system_free, torch_cache_free)
|
|
free_info = mm.get_free_memory(device, torch_free_too=True)
|
|
if isinstance(free_info, tuple):
|
|
mem_free_sys = free_info[0]
|
|
else:
|
|
mem_free_sys = free_info # Fallback for backends that return single value (like MPS)
|
|
|
|
mem_used = mem_total - mem_free_sys
|
|
|
|
# Determine Torch Allocation (Reserved memory) - Specific checks for known backends
|
|
torch_alloc = 0
|
|
if device.type == 'cuda' and hasattr(torch.cuda, 'memory_stats'):
|
|
stats = torch.cuda.memory_stats(device)
|
|
torch_alloc = stats.get('reserved_bytes.all.current', 0)
|
|
elif device.type == 'xpu' and hasattr(torch, 'xpu') and hasattr(torch.xpu, 'memory_stats'):
|
|
stats = torch.xpu.memory_stats(device)
|
|
torch_alloc = stats.get('reserved_bytes.all.current', 0)
|
|
elif device.type == 'npu' and hasattr(torch, 'npu') and hasattr(torch.npu, 'memory_stats'):
|
|
stats = torch.npu.memory_stats(device)
|
|
torch_alloc = stats.get('reserved_bytes.all.current', 0)
|
|
elif device.type == 'mlu' and hasattr(torch, 'mlu') and hasattr(torch.mlu, 'memory_stats'):
|
|
stats = torch.mlu.memory_stats(device)
|
|
torch_alloc = stats.get('reserved_bytes.all.current', 0)
|
|
# MPS, DirectML, CoreX do not always expose detailed reserved memory stats easily.
|
|
|
|
logger.info(fmt_mem.format(
|
|
dev_str,
|
|
f"{mem_total / (1024**3):.2f}",
|
|
f"{mem_free_sys / (1024**3):.2f}",
|
|
f"{mem_used / (1024**3):.2f}",
|
|
f"{torch_alloc / (1024**3):.2f}"
|
|
))
|
|
except Exception as e:
|
|
logger.debug(f"Could not retrieve memory stats for {dev_str}: {e}")
|
|
|
|
logger.info("-" * 70)
|
|
|
|
# 2. Loaded Models Inspection (Logical and Physical View)
|
|
# mm.current_loaded_models holds the list of models ComfyUI is managing.
|
|
loaded_models = mm.current_loaded_models
|
|
logger.info(f"\n--- [2] Loaded Models Inspection (Count: {len(loaded_models)}) ---")
|
|
|
|
if not loaded_models:
|
|
logger.info("No models currently managed by comfy.model_management.")
|
|
logger.info("=" * 100)
|
|
return
|
|
|
|
for i, lm in enumerate(loaded_models):
|
|
logger.info(f"\nModel {i+1}/{len(loaded_models)}:")
|
|
|
|
# Check lifecycle status
|
|
mp = lm.model # weakref call to ModelPatcher
|
|
if mp is None:
|
|
# ModelPatcher is gone. Check if the underlying model is still alive (potential leak)
|
|
if lm.is_dead() and lm.real_model() is not None:
|
|
logger.warning(f" [!] Status: LEAK DETECTED (Patcher GC'd, but underlying model {lm.real_model().__class__.__name__} persists)")
|
|
else:
|
|
logger.info(f" Status: Cleaned Up (Patcher and Model GC'd)")
|
|
continue
|
|
|
|
model_id = create_model_identifier(mp)
|
|
logger.info(f" Identifier: {model_id}")
|
|
logger.info(f" Status: {'Active (In Use)' if lm.currently_used else 'Idle (Cache)'}")
|
|
|
|
# A. Logical View (What ComfyUI intends/tracks)
|
|
logger.info(" [A] Logical View (ComfyUI Tracking):")
|
|
|
|
# Devices: Target (Compute) vs Offload (Storage)
|
|
logger.info(f" Devices: Target={lm.device} | Offload={mp.offload_device} | Current (Model.device)={mp.current_loaded_device()}")
|
|
|
|
# Memory Footprint
|
|
mem_total = lm.model_memory()
|
|
mem_loaded = lm.model_loaded_memory()
|
|
mem_offloaded = lm.model_offloaded_memory()
|
|
logger.info(f" Memory (MB): Total={mem_total/(1024**2):.2f} | Loaded (on Target)={mem_loaded/(1024**2):.2f} | Offloaded={mem_offloaded/(1024**2):.2f}")
|
|
|
|
# Management Mode (LowVRAM/DisTorch)
|
|
# model_lowvram indicates if ComfyUI is managing this model partially
|
|
is_lowvram = getattr(mp.model, 'model_lowvram', False)
|
|
lowvram_patches_pending = mp.lowvram_patch_counter()
|
|
logger.info(f" Mode: {'Partial Load (LowVRAM/DisTorch)' if is_lowvram else 'Full Load'}")
|
|
if is_lowvram:
|
|
# This indicates how many weights are being managed by the partial loading system
|
|
logger.info(f" Weights Managed by LowVRAM/DisTorch System: {lowvram_patches_pending}")
|
|
|
|
# Patching (LoRAs, etc.) - Tracking Attach/Detach
|
|
num_weight_patches = len(mp.patches)
|
|
# Check the UUID applied to the actual weights vs the UUID defined in the patcher
|
|
current_weight_uuid = getattr(mp.model, 'current_weight_patches_uuid', None)
|
|
weights_synced = (mp.patches_uuid == current_weight_uuid) and (current_weight_uuid is not None)
|
|
|
|
if num_weight_patches > 0:
|
|
status = 'Applied & Synced' if weights_synced else 'Pending/Mismatch (Re-patch needed)'
|
|
logger.info(f" Patches: {num_weight_patches} weight patches defined | Status: {status}")
|
|
logger.info(f" UUIDs: Defined={str(mp.patches_uuid)[:8]}... | Applied={str(current_weight_uuid)[:8] if current_weight_uuid else 'None'}...")
|
|
|
|
# B. Physical View (Ground Truth Tensor Locations)
|
|
logger.info(" [B] Physical View (Ground Truth Tensor Locations):")
|
|
device_summary, calculated_total_mem = analyze_tensor_locations(mp)
|
|
|
|
if "error" in device_summary:
|
|
logger.error(f" Analysis Error: {device_summary['error']}")
|
|
continue
|
|
|
|
if not device_summary:
|
|
logger.info(" No tensors found (e.g., fully offloaded CLIP or utility object).")
|
|
else:
|
|
# Sort devices (CPU last)
|
|
sorted_devices = sorted(device_summary.keys(), key=lambda d: (d.startswith("cpu"), d))
|
|
fmt_loc = " {:<15} | Tensors: {:>6} | Memory (MB): {:>10.2f} | Percent: {:>6.1f}%"
|
|
for device in sorted_devices:
|
|
data = device_summary[device]
|
|
percent = (data['memory'] / calculated_total_mem) * 100 if calculated_total_mem > 0 else 0
|
|
logger.info(fmt_loc.format(device, data['tensors'], data['memory']/(1024**2), percent))
|
|
|
|
# Verification Check
|
|
if abs(calculated_total_mem - mem_total) > (1024*1024): # Allow 1MB difference
|
|
logger.warning(f" [!] Verification WARNING: Physical memory ({calculated_total_mem/(1024**2):.2f}MB) differs from logical memory ({mem_total/(1024**2):.2f}MB).")
|
|
|
|
logger.info("-" * 100)
|
|
|
|
logger.info("End of Inspection")
|
|
logger.info("=" * 100)
|