refactor: eliminate circular import by separating model management functions

- Create model_management_mgpu.py for centralized model lifecycle tracking
- Move memory management functions from device_utils.py to new module:
  * multigpu_memory_log, track_modelpatcher, trigger_executor_cache_reset
  * check_cpu_memory_threshold, prune_distorch_stores, try_malloc_trim
  * force_full_system_cleanup
- Update imports across codebase (distorch_2.py, distorch.py, __init__.py,
  nodes.py, checkpoint_multigpu.py)
- Resolves device_utils.py ↔ distorch_2.py circular dependency
- Follows established clean coding patterns with fail-fast error handling

Addresses critical CPU memory leak investigation infrastructure by ensuring
proper module separation for comprehensive memory management utilities.
This commit is contained in:
John Pollock
2025-09-24 17:38:15 -05:00
parent ff6efb4217
commit bd672479fa
9 changed files with 496 additions and 767 deletions
+33
View File
@@ -101,3 +101,36 @@ When working on this project, always reference the Memory Bank for context and m
4. **Targeted reference cleanup** - Performance optimization
This represents the current **highest priority technical debt** requiring resolution.
## Module Architecture Rules
### Module Boundary Principles
- **Single Responsibility**: Each module should have ONE clear purpose
- **Dependency Direction**: Dependencies should flow in ONE direction only
- **Import Hierarchy**: Lower-level modules (device_utils) should NOT import from higher-level modules (distorch_2)
### Module Hierarchy (Dependency Order)
1. `device_utils.py` - **BASE**: Device detection, VRAM management only
2. `model_management_mgpu.py` - **CORE**: Model lifecycle, memory logging, cleanup functions
3. `distorch_2.py`, `distorch.py` - **FEATURES**: DisTorch distribution logic
4. `nodes.py`, `checkpoint_multigpu.py` - **UI**: Node implementations
5. `__init__.py` - **ASSEMBLY**: Final integration and registration
### Mandatory Architecture Checks
**BEFORE adding ANY import statement:**
1. **Check Direction**: Does this create upward dependency? (FORBIDDEN)
2. **Check Purpose**: Does the function belong in this module per Single Responsibility?
3. **Check Cycles**: Run `python -c "import sys; sys.path.append('.'); import <module>"` to detect circular imports
### Function Placement Rules
- **device_utils.py**: ONLY device detection, VRAM cache management
- **model_management_mgpu.py**: Model tracking, memory logging, cleanup utilities
- **Feature modules**: Import from CORE/BASE only, never each other
- **UI modules**: Import from any lower level, implement user interfaces only
### Violation Detection
If import fails with "circular import" or "cannot import name":
1. STOP immediately - do not work around
2. Identify which module boundary was violated
3. Move misplaced function to correct architectural layer
4. Update ALL imports consistently
+2
View File
@@ -12,6 +12,8 @@ from .device_utils import (
get_device_list,
is_accelerator_available,
soft_empty_cache_multigpu,
)
from .model_management_mgpu import (
trigger_executor_cache_reset,
check_cpu_memory_threshold,
multigpu_memory_log,
+2 -1
View File
@@ -12,7 +12,8 @@ import comfy.model_management as mm
import comfy.model_detection
import comfy.clip_vision
from comfy.sd import VAE, CLIP
from .device_utils import get_device_list, soft_empty_cache_multigpu, multigpu_memory_log
from .device_utils import get_device_list, soft_empty_cache_multigpu
from .model_management_mgpu import multigpu_memory_log
from .distorch_2 import safetensor_allocation_store, safetensor_settings_store, create_safetensor_model_hash, register_patched_safetensor_modelpatcher
logger = logging.getLogger("MultiGPU")
+51 -758
View File
@@ -1,6 +1,6 @@
"""
Device detection, management, and inspection utilities for ComfyUI-MultiGPU.
Single source of truth for all device enumeration, compatibility checks, and state inspection.
Single source of truth for all device enumeration, compatibility checks, and VRAM management.
Handles all device types supported by ComfyUI core.
"""
@@ -10,30 +10,6 @@ import hashlib
import psutil
import comfy.model_management as mm
import gc
from datetime import datetime, timezone
import server
import weakref
import platform
import ctypes
import sys
import comfy.model_patcher
# DisTorch stores for pruning/diagnostics
from .distorch_2 import (
safetensor_allocation_store,
safetensor_settings_store,
create_safetensor_model_hash,
)
# Optional DisTorch v1 store support
try:
from .distorch import (
model_allocation_store,
create_model_hash,
)
except Exception:
model_allocation_store = {}
create_model_hash = None
logger = logging.getLogger("MultiGPU")
@@ -41,143 +17,9 @@ logger = logging.getLogger("MultiGPU")
_DEVICE_LIST_CACHE = None
# ==========================================================================================
# Executor Cache Management and CPU Monitoring (Phases 1, 2, 3)
# Device Detection and Management
# ==========================================================================================
# Configuration for CPU Monitoring (Phase 2)
CPU_MEMORY_THRESHOLD_PERCENT = 85.0
# Hysteresis: Only trigger again if usage increased by this amount since the last reset.
CPU_RESET_HYSTERESIS_PERCENT = 5.0
_last_cpu_usage_at_reset = 0.0
def clear_memory_snapshot_history():
"""Clears the stored memory snapshot history. (Phase 3)"""
# Logging integration
multigpu_memory_log("mem_mgmt", "pre-history-clear")
# Snapshot globals exist in this module; operate safely in case of reload
if '_MEM_SNAPSHOT_LAST' in globals():
globals()['_MEM_SNAPSHOT_LAST'].clear()
if '_MEM_SNAPSHOT_SERIES' in globals():
globals()['_MEM_SNAPSHOT_SERIES'].clear()
logger.debug("[MultiGPU_Memory_Management] Memory snapshot history cleared.")
# Logging integration
multigpu_memory_log("mem_mgmt", "post-history-clear")
def trigger_executor_cache_reset(reason="policy", force=False):
"""
(Phase 1/2 Core) Triggers PromptExecutor.reset() by setting the 'free_memory' flag.
Releases CPU-side references held by execution caches.
"""
global _last_cpu_usage_at_reset
# Ensure PromptServer singleton is available
if server.PromptServer.instance is None:
logger.debug("[MultiGPU_Memory_Management] PromptServer instance not yet initialized.")
return
prompt_server = server.PromptServer.instance
# Stability guard: Avoid during active execution unless forced
if prompt_server.prompt_queue.currently_running and not force:
logger.debug(f"[MultiGPU_Memory_Management] Skipping Executor Cache Reset during active prompt execution (Reason: {reason}).")
return
multigpu_memory_log("executor_reset", f"pre-trigger ({reason})")
logger.info(f"[MultiGPU_Memory_Management] Triggering PromptExecutor cache reset (e.reset()). Reason: {reason}")
# Diagnostics and store pruning prior to reset
analyze_cpu_memory_leaks(force=force)
prune_distorch_stores()
# Phase 3: Clear internal snapshot history as the context is resetting
clear_memory_snapshot_history()
# Set the flag on the prompt queue (ComfyUI core mechanism)
prompt_server.prompt_queue.set_flag("free_memory", True)
logger.debug("[MultiGPU_Memory_Management] 'free_memory' flag set.")
# Update usage baseline for hysteresis
vm = psutil.virtual_memory()
_last_cpu_usage_at_reset = vm.percent
# Attempt to return freed memory to OS
try_malloc_trim()
multigpu_memory_log("executor_reset", f"post-trigger ({reason})")
def _cpu_used_bytes():
try:
vm = psutil.virtual_memory()
return vm.used
except Exception:
return 0
def force_full_system_cleanup(reason="manual", force=True):
"""
Mirror ComfyUI-Manager 'Free model and node cache' semantics:
- Only set unload_models=True and free_memory=True flags on the PromptQueue
- The prompt worker (main.py) performs unload/reset/GC
"""
pre_cpu = _cpu_used_bytes()
pre_models = len(getattr(mm, "current_loaded_models", []))
multigpu_memory_log("full_cleanup", f"start:{reason}")
logger.mgpu_mm_log(f"[ManagerMatch] Requesting flags-only cleanup (reason={reason}) | pre_models={pre_models}, cpu_used_gib={pre_cpu/(1024**3):.2f}")
try:
if server.PromptServer.instance is not None:
pq = server.PromptServer.instance.prompt_queue
# Respect currently_running unless forced
if (not pq.currently_running) or force:
pq.set_flag("unload_models", True)
pq.set_flag("free_memory", True)
logger.mgpu_mm_log("[ManagerMatch] Flags set: unload_models=True, free_memory=True")
else:
logger.mgpu_mm_log("[ManagerMatch] Skipped setting flags due to active execution and force=False")
except Exception as e:
logger.mgpu_mm_log(f"[ManagerMatch] Failed to set flags: {e}")
post_cpu = _cpu_used_bytes()
post_models = len(getattr(mm, "current_loaded_models", []))
delta_cpu_mb = (post_cpu - pre_cpu) / (1024**2)
multigpu_memory_log("full_cleanup", f"requested:{reason}")
summary = (
f"[ManagerMatch] Flags-only cleanup requested (reason={reason}) | "
f"models {pre_models}->{post_models} (no immediate unload), cpu_delta_mb={delta_cpu_mb:.2f}"
)
logger.mgpu_mm_log(summary)
return summary
def check_cpu_memory_threshold(threshold_percent=CPU_MEMORY_THRESHOLD_PERCENT):
"""
(Phase 2) Checks CPU memory usage and triggers a reset if threshold is exceeded (with hysteresis).
"""
# Ensure PromptServer singleton is available
if server.PromptServer.instance is None:
return
# Stability/optimization: Do not trigger during active execution
if server.PromptServer.instance.prompt_queue.currently_running:
return
vm = psutil.virtual_memory()
current_usage = vm.percent
if current_usage > threshold_percent:
# Hysteresis gating
if current_usage > (_last_cpu_usage_at_reset + CPU_RESET_HYSTERESIS_PERCENT):
logger.warning(f"[MultiGPU_Memory_Monitor] CPU usage ({current_usage:.1f}%) exceeds threshold ({threshold_percent:.1f}%) and hysteresis.")
multigpu_memory_log("cpu_monitor", f"trigger:{current_usage:.1f}pct")
trigger_executor_cache_reset(reason="cpu_threshold_exceeded", force=False)
else:
logger.debug(f"[MultiGPU_Memory_Monitor] CPU usage high ({current_usage:.1f}%) but within hysteresis range. Skipping reset.")
multigpu_memory_log("cpu_monitor", f"skip_hysteresis:{current_usage:.1f}pct")
def get_device_list():
"""
Enumerate ALL physically available devices that can store torch tensors.
@@ -207,13 +49,10 @@ def get_device_list():
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}")
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:
@@ -221,13 +60,11 @@ def get_device_list():
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}")
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:
@@ -236,8 +73,8 @@ def get_device_list():
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}")
except ImportError:
pass
# MLU devices (Cambricon MLUs)
try:
@@ -246,16 +83,13 @@ def get_device_list():
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}")
except ImportError:
pass
# 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}")
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:
@@ -264,13 +98,12 @@ def get_device_list():
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}")
except ImportError:
pass
# 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)]
@@ -278,8 +111,8 @@ def get_device_list():
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}")
except ImportError:
pass
# Cache the result for future calls
_DEVICE_LIST_CACHE = devs
@@ -289,7 +122,6 @@ def get_device_list():
return devs
def is_accelerator_available():
"""
Check if any accelerator device is available.
@@ -298,60 +130,47 @@ def is_accelerator_available():
Returns True if any GPU/accelerator is available, False otherwise.
"""
# Check CUDA
try:
if torch.cuda.is_available():
return True
except:
pass
if hasattr(torch, "cuda") and torch.cuda.is_available():
return True
# Check XPU (Intel GPU)
try:
if hasattr(torch, "xpu") and torch.xpu.is_available():
return True
except:
pass
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 torch.npu.is_available():
if hasattr(torch, "npu") and hasattr(torch.npu, "is_available") and torch.npu.is_available():
return True
except:
except ImportError:
pass
# Check MLU (Cambricon)
try:
import torch_mlu
if hasattr(torch, "mlu") and torch.mlu.is_available():
if hasattr(torch, "mlu") and hasattr(torch.mlu, "is_available") and torch.mlu.is_available():
return True
except:
except ImportError:
pass
# Check MPS (Apple Metal)
try:
if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
return True
except:
pass
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:
except ImportError:
pass
# Check CoreX/IXUCA
try:
if hasattr(torch, "corex"):
return True
except:
pass
if hasattr(torch, "corex"):
return True
return False
def is_device_compatible(device_string):
"""
Check if a device string represents a valid, available device.
@@ -365,7 +184,6 @@ def is_device_compatible(device_string):
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.
@@ -380,7 +198,6 @@ def get_device_type(device_string):
return device_string.split(":")[0]
return device_string
def parse_device_string(device_string):
"""
Parse a device string into type and index.
@@ -396,28 +213,28 @@ def parse_device_string(device_string):
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.
MultiGPU adaptation of ComfyUI's soft_empty_cache() functionality.
Uses context managers to ensure the calling thread's device context is restored.
"""
import gc
# Import model management functions
from .model_management_mgpu import multigpu_memory_log, log_tracked_modelpatchers_status, try_malloc_trim
logger.mgpu_mm_log("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")
multigpu_memory_log("general", "pre-gc")
# Lifecycle status before GC
log_tracked_modelpatchers_status(tag="pre-gc")
gc.collect()
# Lifecycle status after GC
log_tracked_modelpatchers_status(tag="post-gc")
multigpu_memory_log("general", "post-gc")
logger.mgpu_mm_log("soft_empty_cache_multigpu: garbage collection complete")
# Attempt to release freed heap memory to OS
try_malloc_trim()
# Clear cache for ALL devices (not just ComfyUI's single device)
@@ -431,13 +248,12 @@ def soft_empty_cache_multigpu():
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.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() # ComfyUI's CUDA optimization
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}")
@@ -481,17 +297,13 @@ def soft_empty_cache_multigpu():
logger.mgpu_mm_log(f"Cleared CoreX cache on {device_str}")
multigpu_memory_log("general", f"post-empty:{device_str}")
# Record post-GC snapshot for general system view
multigpu_memory_log("general", "post-soft-empty")
# ==========================================================================================
# Memory Inspection Utilities
# ==========================================================================================
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:
def comfyui_memory_load(tag):
"""
Returns a single-line, pipe-delimited snapshot of system and device memory usage.
@@ -503,8 +315,8 @@ def comfyui_memory_load(tag: str) -> str:
"""
# CPU RAM
vm = psutil.virtual_memory()
cpu_used_gib = _bytes_to_gib(vm.used)
cpu_total_gib = _bytes_to_gib(vm.total)
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}"]
@@ -523,528 +335,9 @@ def comfyui_memory_load(tag: str) -> str:
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)
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)
# ==========================================================================================
# Delta-capable memory logging (identifier + tag) with timestamped series
# ==========================================================================================
# 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
# Stable identifier order for readability
for identifier in sorted(_MEM_SNAPSHOT_SERIES.keys()):
series = _MEM_SNAPSHOT_SERIES[identifier]
if not series:
continue
logger.mgpu_mm_log(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"
logger.mgpu_mm_log(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)}")
logger.mgpu_mm_log(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}")
logger.mgpu_mm_log(f"{identifier} {tag} - <baseline>: " + " | ".join(parts))
# Update last snapshot
_MEM_SNAPSHOT_LAST[identifier] = (tag, curr)
# ==========================================================================================
# Lifecycle Tracking and Leak Analysis Utilities
# ==========================================================================================
# Track ModelPatcher lifecycle to correlate with CPU RAM trends
if '_MGPU_TRACKED_MODELPATCHERS' not in globals():
_MGPU_TRACKED_MODELPATCHERS = weakref.WeakSet()
def track_modelpatcher(model_patcher):
"""Registers a ModelPatcher instance for lifecycle tracking."""
try:
if isinstance(model_patcher, comfy.model_patcher.ModelPatcher):
if model_patcher not in _MGPU_TRACKED_MODELPATCHERS:
_MGPU_TRACKED_MODELPATCHERS.add(model_patcher)
logger.debug(f"[MultiGPU_Lifecycle] Tracking ModelPatcher {id(model_patcher)} (tracked={len(_MGPU_TRACKED_MODELPATCHERS)})")
except Exception as e:
logger.debug(f"[MultiGPU_Lifecycle] track_modelpatcher error: {e}")
def log_tracked_modelpatchers_status(tag="checkpoint"):
"""Logs count and estimated CPU RAM for tracked ModelPatchers."""
alive_count = len(_MGPU_TRACKED_MODELPATCHERS)
total_cpu_memory_mb = 0.0
for patcher in list(_MGPU_TRACKED_MODELPATCHERS):
try:
if hasattr(patcher, "model") and patcher.model is not None:
for param in patcher.model.parameters():
if getattr(param, "device", torch.device("cpu")).type == "cpu":
total_cpu_memory_mb += (param.nelement() * param.element_size()) / (1024.0 * 1024.0)
except Exception:
continue
logger.warning(f"[MultiGPU_Lifecycle] [{tag}] Tracked ModelPatchers={alive_count}, approx CPU RAM={total_cpu_memory_mb:.2f} MB")
def analyze_cpu_memory_leaks(force=False):
"""Diagnostic: scan referrers of tracked ModelPatchers when memory is high."""
try:
vm = psutil.virtual_memory()
patchers = list(_MGPU_TRACKED_MODELPATCHERS)
if not force and len(patchers) <= 5 and vm.percent <= 80.0:
logger.debug(f"[MultiGPU_Leak_Analyzer] Skipping analysis. Patcher count ({len(patchers)}) and memory usage ({vm.percent:.1f}%) normal.")
return
logger.warning(f"[MultiGPU_Leak_Analyzer] High pressure (patchers={len(patchers)}, cpu_mem={vm.percent:.1f}%). Inspecting up to 5 referrer sets.")
for i, patcher in enumerate(patchers[:5]):
try:
referrers = gc.get_referrers(patcher)
logger.warning(f"[MultiGPU_Leak_Analyzer] Patcher #{i} id={id(patcher)} referrers={len(referrers)}")
for j, ref in enumerate(referrers[:10]):
rtype = type(ref).__name__
rmod = getattr(type(ref), "__module__", "unknown")
if isinstance(ref, dict):
logger.warning(f" Ref {j}: dict(len={len(ref)}) mod={rmod}")
elif isinstance(ref, list):
logger.warning(f" Ref {j}: list(len={len(ref)}) mod={rmod}")
else:
logger.warning(f" Ref {j}: {rtype} mod={rmod}")
except Exception:
logger.warning("[MultiGPU_Leak_Analyzer] Failed to inspect referrers for a patcher.")
except Exception as e:
logger.debug(f"[MultiGPU_Leak_Analyzer] analyze error: {e}")
def try_malloc_trim():
"""Attempt to return freed heap memory to OS (Linux/glibc)."""
try:
if platform.system() == "Linux":
libc = ctypes.CDLL("libc.so.6")
if hasattr(libc, "malloc_trim"):
logger.info("[MultiGPU_Memory_Management] malloc_trim(0) begin")
multigpu_memory_log("mem_mgmt", "pre-malloc-trim")
res = libc.malloc_trim(0)
multigpu_memory_log("mem_mgmt", "post-malloc-trim")
if res == 1:
logger.info("[MultiGPU_Memory_Management] malloc_trim(0) released memory")
else:
logger.debug("[MultiGPU_Memory_Management] malloc_trim(0) no release")
except Exception as e:
logger.debug(f"[MultiGPU_Memory_Management] malloc_trim error: {e}")
def prune_distorch_stores():
"""Prune stale allocation/settings entries not tied to active models."""
try:
multigpu_memory_log("distorch_prune", "start")
active_hashes_v2 = set()
active_hashes_v1 = set()
for lm in getattr(mm, "current_loaded_models", []):
mp = getattr(lm, "model", None)
if mp is not None:
try:
h2 = create_safetensor_model_hash(mp, "prune_check_v2")
active_hashes_v2.add(h2)
except Exception:
pass
if create_model_hash is not None:
try:
h1 = create_model_hash(mp, "prune_check_v1")
active_hashes_v1.add(h1)
except Exception:
pass
# V1
if isinstance(model_allocation_store, dict) and active_hashes_v1:
stale = set(model_allocation_store.keys()) - active_hashes_v1
if stale:
logger.info(f"[MultiGPU_Memory_Management] Pruning {len(stale)} DisTorch V1 entries")
for k in stale:
model_allocation_store.pop(k, None)
# V2
for store, name in ((safetensor_allocation_store, "allocation"), (safetensor_settings_store, "settings")):
try:
if isinstance(store, dict):
stale2 = set(store.keys()) - active_hashes_v2
if stale2:
logger.info(f"[MultiGPU_Memory_Management] Pruning {len(stale2)} V2 {name} entries")
for k in stale2:
store.pop(k, None)
except Exception:
pass
multigpu_memory_log("distorch_prune", "end")
except Exception as e:
logger.debug(f"[MultiGPU_Memory_Management] prune_distorch_stores error: {e}")
# ==========================================================================================
# 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)
+2 -1
View File
@@ -12,7 +12,8 @@ logger = logging.getLogger("MultiGPU")
import copy
from collections import defaultdict
import comfy.model_management as mm
from .device_utils import get_device_list, soft_empty_cache_multigpu, multigpu_memory_log
from .device_utils import get_device_list, soft_empty_cache_multigpu
from .model_management_mgpu import multigpu_memory_log
# Global store for model allocations
model_allocation_store = {}
+2 -6
View File
@@ -16,6 +16,8 @@ import inspect
from collections import defaultdict
import comfy.model_management as mm
import comfy.model_patcher
from .device_utils import get_device_list, soft_empty_cache_multigpu
from .model_management_mgpu import multigpu_memory_log, track_modelpatcher
safetensor_allocation_store = {}
safetensor_settings_store = {}
@@ -59,7 +61,6 @@ def register_patched_safetensor_modelpatcher():
def new_partially_load(self, device_to, extra_memory=0, full_load=False, force_patch_weights=False, **kwargs):
"""Override to use our static device assignments"""
from .device_utils import multigpu_memory_log, track_modelpatcher
global safetensor_allocation_store
debug_hash = create_safetensor_model_hash(self, "partial_load")
@@ -181,7 +182,6 @@ def analyze_safetensor_loading(model_patcher, allocations_string):
Analyze and distribute safetensor model blocks across devices
Target for refactor back into one function once stability for CLIP is established.
"""
from .device_utils import get_device_list
DEVICE_RATIOS_DISTORCH = {}
device_table = {}
distorch_alloc = allocations_string
@@ -389,7 +389,6 @@ def analyze_safetensor_loading_clip(model_patcher, allocations_string):
All other logic and UX (logging, etc.) is identical to the original.
Target for refactor once stability for CLIP is established.
"""
from .device_utils import get_device_list
DEVICE_RATIOS_DISTORCH = {}
device_table = {}
distorch_alloc = allocations_string
@@ -805,7 +804,6 @@ def override_class_with_distorch_safetensor_v2(cls):
class NodeOverrideDisTorchSafetensorV2(cls):
@classmethod
def INPUT_TYPES(s):
from .device_utils import get_device_list
inputs = copy.deepcopy(cls.INPUT_TYPES())
devices = get_device_list()
compute_device = devices[1] if len(devices) > 1 else devices[0]
@@ -902,7 +900,6 @@ def override_class_with_distorch_safetensor_v2_clip(cls):
class NodeOverrideDisTorchSafetensorV2Clip(cls):
@classmethod
def INPUT_TYPES(s):
from .device_utils import get_device_list
inputs = copy.deepcopy(cls.INPUT_TYPES())
devices = get_device_list()
default_device = devices[1] if len(devices) > 1 else devices[0]
@@ -1000,7 +997,6 @@ def override_class_with_distorch_safetensor_v2_clip_no_device(cls):
class NodeOverrideDisTorchSafetensorV2ClipNoDevice(cls):
@classmethod
def INPUT_TYPES(s):
from .device_utils import get_device_list
inputs = copy.deepcopy(cls.INPUT_TYPES())
devices = get_device_list()
default_device = devices[1] if len(devices) > 1 else devices[0]
+69
View File
@@ -322,3 +322,72 @@ def benchmark_allocation_performance(model, hardware_config, allocation_configs)
performance_ratio = distributed_time / baseline_time
assert performance_ratio < expected_slowdown_threshold(hardware_config)
```
## Module Architecture (Post-Refactoring)
### Core Module Separation
**Problem Solved**: Eliminated circular import `device_utils.py` ↔ `distorch_2.py`
**Solution**: Created `model_management_mgpu.py` as central model lifecycle hub
### Module Responsibilities
**device_utils.py** (Base Layer):
- Device enumeration and detection
- VRAM cache management (`soft_empty_cache_multigpu`)
- Pure hardware abstraction - NO model tracking
**model_management_mgpu.py** (Core Layer):
- Model lifecycle tracking (`track_modelpatcher`)
- Memory logging (`multigpu_memory_log`)
- System cleanup (`force_full_system_cleanup`, `trigger_executor_cache_reset`)
- Store pruning (`prune_distorch_stores`)
**distorch_2.py/distorch.py** (Feature Layer):
- DisTorch distribution algorithms
- SafeTensor/GGUF specific logic
- Imports FROM core/base layers ONLY
### Import Flow Architecture
```
┌─────────────────┐
│ __init__.py │ ← Assembly Layer
└─────────────────┘
↑
┌─────────────────┐
│ UI Layer │ ← nodes.py, checkpoint_multigpu.py
│ (User Interface)│
└─────────────────┘
↑
┌─────────────────┐
│ Feature Layer │ ← distorch_2.py, distorch.py
│ (DisTorch Logic)│
└─────────────────┘
↑
┌─────────────────┐
│ Core Layer │ ← model_management_mgpu.py
│ (Model Lifecycle)│
└─────────────────┘
↑
┌─────────────────┐
│ Base Layer │ ← device_utils.py
│ (Hardware) │
└─────────────────┘
```
### Architectural Validation
**Rule**: Dependencies only flow UPWARD. Violations create circular imports.
**Prevention**: Before any import, ask "Does this violate the layer hierarchy?"
### Function Migration Record
**Moved from device_utils.py to model_management_mgpu.py:**
- `multigpu_memory_log` - Memory state logging
- `track_modelpatcher` - ModelPatcher lifecycle tracking
- `trigger_executor_cache_reset` - CPU memory management
- `check_cpu_memory_threshold` - Adaptive cleanup triggers
- `prune_distorch_stores` - Store cleanup utilities
- `try_malloc_trim` - System memory reclamation
- `force_full_system_cleanup` - Full system reset
**Rationale**: These functions manage model lifecycle and memory state, not hardware detection. Separation prevents circular dependencies while maintaining clean responsibilities.
+333
View File
@@ -0,0 +1,333 @@
"""
Model Management Extensions for MultiGPU
Extends ComfyUI's model management with multi-device capabilities and lifecycle tracking.
"""
import torch
import logging
import hashlib
import psutil
import comfy.model_management as mm
import gc
from datetime import datetime, timezone
import server
import weakref
import platform
import ctypes
import comfy.model_patcher
from collections import defaultdict
logger = logging.getLogger("MultiGPU")
# ==========================================================================================
# Model Analysis and Store Management (DisTorch V1 & V2)
# ==========================================================================================
# DisTorch V2 SafeTensor stores
safetensor_allocation_store = {}
safetensor_settings_store = {}
# DisTorch V1 GGUF stores (backwards compatibility)
model_allocation_store = {}
def create_safetensor_model_hash(model, caller):
"""Create a unique hash for a safetensor model to track allocations"""
if hasattr(model, 'model'):
actual_model = model.model
model_type = type(actual_model).__name__
model_size = model.model_size() if hasattr(model, 'model_size') else sum(p.numel() * p.element_size() for p in actual_model.parameters())
first_layers = str(list(model.model_state_dict().keys() if hasattr(model, 'model_state_dict') else actual_model.state_dict().keys())[:3])
else:
model_type = type(model).__name__
model_size = sum(p.numel() * p.element_size() for p in model.parameters())
first_layers = str(list(model.state_dict().keys())[:3])
identifier = f"{model_type}_{model_size}_{first_layers}"
final_hash = hashlib.sha256(identifier.encode()).hexdigest()
logger.debug(f"[MultiGPU DisTorch V2] Created hash for {caller}: {final_hash[:8]}...")
return final_hash
def create_model_hash(model, caller):
"""Create a unique hash for a GGUF model to track allocations (DisTorch V1)"""
model_type = type(model.model).__name__
model_size = model.model_size()
first_layers = str(list(model.model_state_dict().keys())[:3])
identifier = f"{model_type}_{model_size}_{first_layers}"
final_hash = hashlib.sha256(identifier.encode()).hexdigest()
logger.debug(f"[MultiGPU_DisTorch_HASH] Created hash for {caller}: {final_hash[:8]}...")
return final_hash
def prune_distorch_stores():
"""Prune stale allocation/settings entries not tied to active models."""
multigpu_memory_log("distorch_prune", "start")
active_hashes_v2 = set()
active_hashes_v1 = set()
for lm in mm.current_loaded_models:
mp = lm.model
if mp is not None:
active_hashes_v2.add(create_safetensor_model_hash(mp, "prune_check_v2"))
active_hashes_v1.add(create_model_hash(mp, "prune_check_v1"))
# V1 pruning
stale_v1 = set(model_allocation_store.keys()) - active_hashes_v1
if stale_v1:
logger.info(f"[MultiGPU_Memory_Management] Pruning {len(stale_v1)} DisTorch V1 entries")
for k in stale_v1:
del model_allocation_store[k]
# V2 pruning
for store, name in ((safetensor_allocation_store, "allocation"), (safetensor_settings_store, "settings")):
stale_v2 = set(store.keys()) - active_hashes_v2
if stale_v2:
logger.info(f"[MultiGPU_Memory_Management] Pruning {len(stale_v2)} V2 {name} entries")
for k in stale_v2:
del store[k]
multigpu_memory_log("distorch_prune", "end")
# ==========================================================================================
# Memory Logging Infrastructure
# ==========================================================================================
_MEM_SNAPSHOT_LAST = {}
_MEM_SNAPSHOT_SERIES = {}
def _capture_memory_snapshot():
"""Capture memory snapshot for CPU and all devices"""
# Import here to avoid circular dependency
from .device_utils import get_device_list
snapshot = {}
# CPU
vm = psutil.virtual_memory()
snapshot["cpu"] = (vm.used, vm.total)
# GPU 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 - system_free)
snapshot[dev_str] = (used, total)
return snapshot
def multigpu_memory_log(identifier, tag):
"""Record timestamped memory snapshot with delta logging"""
if identifier == "print_summary":
for id_key in sorted(_MEM_SNAPSHOT_SERIES.keys()):
series = _MEM_SNAPSHOT_SERIES[id_key]
logger.mgpu_mm_log(f"=== memory summary: {id_key} ===")
for ts, tag_name, snap in series:
parts = []
cpu_used, cpu_total = snap.get("cpu", (0, 0))
parts.append(f"cpu={cpu_used/(1024**3):.2f}/{cpu_total/(1024**3):.2f}")
for dev in sorted([k for k in snap.keys() if k != "cpu"]):
used, total = snap[dev]
parts.append(f"{dev}={used/(1024**3):.2f}/{total/(1024**3):.2f}")
ts_str = ts.strftime("%Y-%m-%dT%H:%M:%S.%f")[:-3] + "Z"
logger.mgpu_mm_log(f"{ts_str} {id_key} {tag_name} | " + " | ".join(parts))
return
ts = datetime.now(timezone.utc)
curr = _capture_memory_snapshot()
# Store in series
if identifier not in _MEM_SNAPSHOT_SERIES:
_MEM_SNAPSHOT_SERIES[identifier] = []
_MEM_SNAPSHOT_SERIES[identifier].append((ts, tag, curr))
# Compute delta
if identifier in _MEM_SNAPSHOT_LAST:
prev_tag, prev = _MEM_SNAPSHOT_LAST[identifier]
keys = set(prev.keys()) | set(curr.keys())
ordered = ["cpu"] + sorted([k for k in keys if k != "cpu"])
parts = []
for k in ordered:
p_used, _ = prev.get(k, (0, 0))
c_used, _ = curr.get(k, (0, 0))
delta = c_used - p_used
sign = "+" if delta >= 0 else "-"
parts.append(f"{k}={sign}{abs(delta)/(1024**3):.2f}")
logger.mgpu_mm_log(f"{identifier} {tag} - {prev_tag}: " + " | ".join(parts))
else:
# Baseline
ordered = ["cpu"] + sorted([k for k in curr.keys() if k != "cpu"])
parts = []
for k in ordered:
c_used, _ = curr.get(k, (0, 0))
parts.append(f"{k}=+{c_used/(1024**3):.2f}")
logger.mgpu_mm_log(f"{identifier} {tag} - <baseline>: " + " | ".join(parts))
_MEM_SNAPSHOT_LAST[identifier] = (tag, curr)
def clear_memory_snapshot_history():
"""Clear stored memory snapshot history"""
multigpu_memory_log("mem_mgmt", "pre-history-clear")
_MEM_SNAPSHOT_LAST.clear()
_MEM_SNAPSHOT_SERIES.clear()
logger.debug("[MultiGPU_Memory_Management] Memory snapshot history cleared")
multigpu_memory_log("mem_mgmt", "post-history-clear")
# ==========================================================================================
# ModelPatcher Lifecycle Tracking
# ==========================================================================================
_MGPU_TRACKED_MODELPATCHERS = weakref.WeakSet()
def track_modelpatcher(model_patcher):
"""Register ModelPatcher for lifecycle tracking"""
if isinstance(model_patcher, comfy.model_patcher.ModelPatcher):
if model_patcher not in _MGPU_TRACKED_MODELPATCHERS:
_MGPU_TRACKED_MODELPATCHERS.add(model_patcher)
logger.debug(f"[MultiGPU_Lifecycle] Tracking ModelPatcher {id(model_patcher)} (tracked={len(_MGPU_TRACKED_MODELPATCHERS)})")
def log_tracked_modelpatchers_status(tag="checkpoint"):
"""Log count and estimated CPU RAM for tracked ModelPatchers"""
alive_count = len(_MGPU_TRACKED_MODELPATCHERS)
total_cpu_memory_mb = 0.0
for patcher in list(_MGPU_TRACKED_MODELPATCHERS):
if hasattr(patcher, "model") and patcher.model is not None:
for param in patcher.model.parameters():
if getattr(param, "device", torch.device("cpu")).type == "cpu":
total_cpu_memory_mb += (param.nelement() * param.element_size()) / (1024.0 * 1024.0)
logger.warning(f"[MultiGPU_Lifecycle] [{tag}] Tracked ModelPatchers={alive_count}, approx CPU RAM={total_cpu_memory_mb:.2f} MB")
def analyze_cpu_memory_leaks():
"""Diagnostic: scan referrers of tracked ModelPatchers when memory is high"""
vm = psutil.virtual_memory()
patchers = list(_MGPU_TRACKED_MODELPATCHERS)
if len(patchers) <= 5 and vm.percent <= 80.0:
logger.debug(f"[MultiGPU_Leak_Analyzer] Skipping analysis. Normal conditions: patchers={len(patchers)}, memory={vm.percent:.1f}%")
return
logger.warning(f"[MultiGPU_Leak_Analyzer] High pressure detected: patchers={len(patchers)}, cpu_mem={vm.percent:.1f}%. Analyzing referrers.")
for i, patcher in enumerate(patchers[:5]):
referrers = gc.get_referrers(patcher)
logger.warning(f"[MultiGPU_Leak_Analyzer] Patcher #{i} id={id(patcher)} referrers={len(referrers)}")
for j, ref in enumerate(referrers[:10]):
rtype = type(ref).__name__
rmod = getattr(type(ref), "__module__", "unknown")
if isinstance(ref, dict):
logger.warning(f" Ref {j}: dict(len={len(ref)}) mod={rmod}")
elif isinstance(ref, list):
logger.warning(f" Ref {j}: list(len={len(ref)}) mod={rmod}")
else:
logger.warning(f" Ref {j}: {rtype} mod={rmod}")
# ==========================================================================================
# Memory Management and Cleanup
# ==========================================================================================
CPU_MEMORY_THRESHOLD_PERCENT = 85.0
CPU_RESET_HYSTERESIS_PERCENT = 5.0
_last_cpu_usage_at_reset = 0.0
def try_malloc_trim():
"""Return freed heap memory to OS (Linux/glibc)"""
if platform.system() != "Linux":
return
libc = ctypes.CDLL("libc.so.6")
if not hasattr(libc, "malloc_trim"):
return
logger.info("[MultiGPU_Memory_Management] malloc_trim(0) begin")
multigpu_memory_log("mem_mgmt", "pre-malloc-trim")
result = libc.malloc_trim(0)
multigpu_memory_log("mem_mgmt", "post-malloc-trim")
if result == 1:
logger.info("[MultiGPU_Memory_Management] malloc_trim(0) released memory")
else:
logger.debug("[MultiGPU_Memory_Management] malloc_trim(0) no release")
def trigger_executor_cache_reset(reason="policy", force=False):
"""Trigger PromptExecutor.reset() by setting 'free_memory' flag"""
global _last_cpu_usage_at_reset
prompt_server = server.PromptServer.instance
if prompt_server is None:
logger.debug("[MultiGPU_Memory_Management] PromptServer not initialized")
return
if prompt_server.prompt_queue.currently_running and not force:
logger.debug(f"[MultiGPU_Memory_Management] Skipping reset during execution (reason: {reason})")
return
multigpu_memory_log("executor_reset", f"pre-trigger ({reason})")
logger.info(f"[MultiGPU_Memory_Management] Triggering PromptExecutor cache reset. Reason: {reason}")
analyze_cpu_memory_leaks()
prune_distorch_stores()
clear_memory_snapshot_history()
prompt_server.prompt_queue.set_flag("free_memory", True)
logger.debug("[MultiGPU_Memory_Management] 'free_memory' flag set")
vm = psutil.virtual_memory()
_last_cpu_usage_at_reset = vm.percent
try_malloc_trim()
multigpu_memory_log("executor_reset", f"post-trigger ({reason})")
def check_cpu_memory_threshold(threshold_percent=CPU_MEMORY_THRESHOLD_PERCENT):
"""Check CPU memory and trigger reset if threshold exceeded"""
if server.PromptServer.instance is None:
return
if server.PromptServer.instance.prompt_queue.currently_running:
return
vm = psutil.virtual_memory()
current_usage = vm.percent
if current_usage > threshold_percent:
if current_usage > (_last_cpu_usage_at_reset + CPU_RESET_HYSTERESIS_PERCENT):
logger.warning(f"[MultiGPU_Memory_Monitor] CPU usage ({current_usage:.1f}%) exceeds threshold ({threshold_percent:.1f}%)")
multigpu_memory_log("cpu_monitor", f"trigger:{current_usage:.1f}pct")
trigger_executor_cache_reset(reason="cpu_threshold_exceeded", force=False)
else:
logger.debug(f"[MultiGPU_Memory_Monitor] CPU usage high ({current_usage:.1f}%) but within hysteresis")
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
"""
vm = psutil.virtual_memory()
pre_cpu = vm.used
pre_models = len(mm.current_loaded_models)
multigpu_memory_log("full_cleanup", f"start:{reason}")
logger.mgpu_mm_log(f"[ManagerMatch] Requesting cleanup (reason={reason}) | pre_models={pre_models}, cpu_used_gib={pre_cpu/(1024**3):.2f}")
if server.PromptServer.instance is not None:
pq = server.PromptServer.instance.prompt_queue
if (not pq.currently_running) or force:
pq.set_flag("unload_models", True)
pq.set_flag("free_memory", True)
logger.mgpu_mm_log("[ManagerMatch] Flags set: unload_models=True, free_memory=True")
else:
logger.mgpu_mm_log("[ManagerMatch] Skipped - execution active and force=False")
vm = psutil.virtual_memory()
post_cpu = vm.used
post_models = len(mm.current_loaded_models)
delta_cpu_mb = (post_cpu - pre_cpu) / (1024**2)
multigpu_memory_log("full_cleanup", f"requested:{reason}")
summary = f"[ManagerMatch] Cleanup requested (reason={reason}) | models {pre_models}->{post_models}, cpu_delta_mb={delta_cpu_mb:.2f}"
logger.mgpu_mm_log(summary)
return summary
+2 -1
View File
@@ -2,7 +2,8 @@ import torch
import folder_paths
from pathlib import Path
from nodes import NODE_CLASS_MAPPINGS
from .device_utils import get_device_list, force_full_system_cleanup
from .device_utils import get_device_list
from .model_management_mgpu import force_full_system_cleanup
class DeviceSelectorMultiGPU:
@classmethod