feat: add model inspection utilities for tracking and analysis
- Extend module docstring to include inspection capabilities - Add create_model_identifier() to generate unique hashes from model type and size - Add analyze_tensor_locations() to analyze tensor device placement and memory usage - Include imports for hashlib, psutil, and comfy.model_management to support new features These utilities enable end-to-end tracking of model state and placement for better debugging and management in multi-GPU setups.
This commit is contained in:
+276
-2
@@ -1,11 +1,14 @@
|
||||
"""
|
||||
Device detection and management utilities for ComfyUI-MultiGPU.
|
||||
Single source of truth for all device enumeration and compatibility checks.
|
||||
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")
|
||||
|
||||
@@ -269,3 +272,274 @@ def soft_empty_cache_multigpu(logger):
|
||||
elif device_str.startswith("corex:"):
|
||||
torch.corex.empty_cache() # Hypothetical based on ComfyUI's ixuca support
|
||||
logger.debug("[MultiGPU_Device_Utils] Cleared cache for CoreX")
|
||||
|
||||
|
||||
# ==========================================================================================
|
||||
# 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)
|
||||
|
||||
Reference in New Issue
Block a user