feat: implement comprehensive memory management and OOM prevention

- Add ModelPatcher lifecycle tracking with weakref-based cleanup
- Implement reference cycle fixes in LoadedModel to prevent memory leaks
- Add memory threshold monitoring and automatic cleanup triggers
- Enable multigpu memory logging for debugging (MGPU_MM_LOG=True)
- Add OOM handling with graceful cleanup and recovery mechanisms
- Import additional memory utilities for cache management and malloc trimming
This commit is contained in:
John Pollock
2025-09-23 22:52:20 -05:00
parent cd7a536645
commit 7b319544e0
4 changed files with 508 additions and 19 deletions
+163 -11
View File
@@ -1,12 +1,24 @@
import torch
import logging
import weakref
import os
import copy
from pathlib import Path
import folder_paths
import comfy.model_management as mm
import comfy.model_patcher
from nodes import NODE_CLASS_MAPPINGS as GLOBAL_NODE_CLASS_MAPPINGS
from .device_utils import get_device_list, is_accelerator_available, soft_empty_cache_multigpu
from .device_utils import (
get_device_list,
is_accelerator_available,
soft_empty_cache_multigpu,
trigger_executor_cache_reset,
check_cpu_memory_threshold,
multigpu_memory_log,
prune_distorch_stores,
try_malloc_trim,
track_modelpatcher,
)
# --- DisTorch V2 Logging Configuration ---
# Set to "E" for Engineering (DEBUG) or "P" for Production (INFO)
@@ -24,7 +36,7 @@ if not logger.handlers:
logger.addHandler(handler)
logger.setLevel(log_level)
MGPU_MM_LOG = False
MGPU_MM_LOG = True
def mgpu_mm_log_method(self, msg):
if MGPU_MM_LOG:
@@ -148,6 +160,111 @@ logger.debug(f"[MultiGPU DEBUG] Initial current_text_encoder_device: {current_te
mm.get_torch_device = get_torch_device_patched
mm.text_encoder_device = text_encoder_device_patched
# ==========================================================================================
# Core Patching: ModelPatcher Lifecycle Tracking (__init__)
# ==========================================================================================
logger.info("[MultiGPU Core Patching] Applying ModelPatcher lifecycle tracking patch (__init__).")
if not hasattr(comfy.model_patcher.ModelPatcher, '_mgpu_lifecycle_patched'):
try:
_mgpu_original_modelpatcher_init = comfy.model_patcher.ModelPatcher.__init__
def _mgpu_patched_modelpatcher_init(self, *args, **kwargs):
_mgpu_original_modelpatcher_init(self, *args, **kwargs)
# Track all ModelPatcher instances at construction time
try:
track_modelpatcher(self)
except Exception:
pass
comfy.model_patcher.ModelPatcher.__init__ = _mgpu_patched_modelpatcher_init
comfy.model_patcher.ModelPatcher._mgpu_lifecycle_patched = True
logger.info("[MultiGPU Core Patching] ModelPatcher.__init__ patched for lifecycle tracking.")
except Exception as e:
logger.error(f"[MultiGPU Core Patching] FAILED to patch ModelPatcher.__init__: {e}")
# ==========================================================================================
# Core Patching: Fix Potential Reference Cycles in LoadedModel
# ==========================================================================================
if hasattr(mm, 'LoadedModel') and hasattr(mm.LoadedModel, '_set_model'):
logger.info("[MultiGPU Core Patching] Patching mm.LoadedModel._set_model and _switch_parent to reduce reference cycles.")
_mgpu_original_set_model = mm.LoadedModel._set_model
def _mgpu_patched_set_model(self, model):
patcher_id = id(model)
# Ensure attributes exist
if not hasattr(self, '_model'):
self._model = None
if not hasattr(self, '_parent_model'):
self._parent_model = None
if not hasattr(self, '_patcher_finalizer'):
self._patcher_finalizer = None
# Reset refs
self._model = weakref.ref(model)
self._parent_model = None
# Detach any previous finalizer
if self._patcher_finalizer is not None:
try:
self._patcher_finalizer.detach()
except Exception:
pass
self._patcher_finalizer = None
# If clone, set parent and attach a weakref-based finalizer
parent = getattr(model, 'parent', None)
if parent is not None:
self._parent_model = weakref.ref(parent)
self_weak = weakref.ref(self)
def _mgpu_finalize_clone():
s = self_weak()
if s is not None and hasattr(s, '_switch_parent'):
logger.mgpu_mm_log(f"[MultiGPU_LoadedModel_Patch] Clone Patcher {patcher_id} GC'd. Switching LoadedModel to parent.")
s._switch_parent()
else:
logger.mgpu_mm_log(f"[MultiGPU_LoadedModel_Patch] Clone Patcher {patcher_id} GC'd. LoadedModel already gone or missing _switch_parent.")
try:
self._patcher_finalizer = weakref.finalize(model, _mgpu_finalize_clone)
except Exception:
self._patcher_finalizer = None
else:
logger.mgpu_mm_log(f"[MultiGPU_LoadedModel_Patch] Set base model Patcher {patcher_id}.")
mm.LoadedModel._set_model = _mgpu_patched_set_model
# Patch _switch_parent to clear references explicitly
if hasattr(mm.LoadedModel, '_switch_parent'):
_mgpu_original_switch_parent = mm.LoadedModel._switch_parent
def _mgpu_patched_switch_parent(self):
_mgpu_original_switch_parent(self)
# Clear parent and detach finalizer to avoid cycles
if hasattr(self, '_parent_model'):
self._parent_model = None
if hasattr(self, '_patcher_finalizer') and self._patcher_finalizer is not None:
try:
self._patcher_finalizer.detach()
except Exception:
pass
self._patcher_finalizer = None
mm.LoadedModel._switch_parent = _mgpu_patched_switch_parent
else:
# Fallback if core ever changes
def _mgpu_fallback_switch_parent(self):
if hasattr(self, '_parent_model') and self._parent_model is not None:
parent_model = self._parent_model()
if parent_model is not None:
self._set_model(parent_model)
self._parent_model = None
mm.LoadedModel._switch_parent = _mgpu_fallback_switch_parent
else:
logger.warning("[MultiGPU Core Patching] mm.LoadedModel not found or missing _set_model; skip cycle patch.")
def check_module_exists(module_path):
full_path = os.path.join(folder_paths.get_folder_paths("custom_nodes")[0], module_path)
logger.debug(f"[MultiGPU] Checking for module at {full_path}")
@@ -181,6 +298,7 @@ from .nodes import (
HyVideoModelLoader,
HyVideoVAELoader,
DownloadAndLoadHyVideoTextEncoder,
FullCleanupMultiGPU,
)
# Import from wanvideo.py
@@ -221,32 +339,57 @@ from .distorch_2 import (
override_class_with_distorch_safetensor_v2_clip_no_device
)
logger.info("[MultiGPU Core Patching] Patching mm.soft_empty_cache for DisTorch2 Multi-Device Allocation/Clearing")
logger.info("[MultiGPU Core Patching] Patching mm.soft_empty_cache for Comprehensive Memory Management (VRAM + CPU + Store Pruning)")
original_soft_empty_cache = mm.soft_empty_cache
def soft_empty_cache_distorch2_patched(force=False):
"""
Patched mm.soft_empty_cache. If DisTorch2 models are active, clear cache on ALL devices.
Otherwise, execute original ComfyUI behavior.
Patched mm.soft_empty_cache.
- Prunes DisTorch store bookkeeping to avoid stale references
- Manages VRAM: if DisTorch2 models are active, clear allocator caches on all devices;
otherwise delegate to original mm.soft_empty_cache.
- Manages CPU RAM: adaptive threshold-based PromptExecutor cache reset;
and force-triggered reset when explicitly requested (mirrors ComfyUI 'Free memory' button).
"""
multigpu_memory_log("patched_soft_empty", f"start:force={force}")
# Prune DisTorch stores before any clearing to drop stale references
try:
prune_distorch_stores()
except Exception:
pass
is_distorch_active = False
# Check if any loaded model is managed by DisTorch2 using the allocation store
# Detect DisTorch2-managed models
for lm in mm.current_loaded_models:
mp = lm.model # weakref call to ModelPatcher
if mp is not None:
model_hash = create_safetensor_model_hash(mp, "cache_patch_check")
if model_hash in safetensor_allocation_store and safetensor_allocation_store[model_hash]:
if model_hash in safetensor_allocation_store and safetensor_allocation_store.get(model_hash):
is_distorch_active = True
break
# Phase 2: adaptive CPU memory management
check_cpu_memory_threshold()
# VRAM allocator management
if is_distorch_active:
logger.mgpu_mm_log("DisTorch2 active: clearing caches on all devices")
logger.mgpu_mm_log("DisTorch2 active: clearing allocator caches on all devices (VRAM)")
soft_empty_cache_multigpu()
else:
logger.mgpu_mm_log("DisTorch2 not active: delegating to original mm.soft_empty_cache")
logger.mgpu_mm_log("DisTorch2 not active: delegating allocator cache clear (VRAM) to original mm.soft_empty_cache")
original_soft_empty_cache(force)
# Attempt to return CPU heap to OS on legacy path as well
try:
try_malloc_trim()
except Exception:
pass
# Phase 1/3: forced executor reset mirrors ComfyUI 'Free memory' semantics
if force:
logger.mgpu_mm_log("Force flag active: triggering executor cache reset (CPU)")
trigger_executor_cache_reset(reason="forced_soft_empty", force=True)
multigpu_memory_log("patched_soft_empty", "end")
mm.soft_empty_cache = soft_empty_cache_distorch2_patched
@@ -263,6 +406,7 @@ if hasattr(mm, 'load_models_gpu') and not hasattr(mm.load_models_gpu, "_distorch
Proactively unload large models that are not needed when loading a large DisTorch2 model.
This frees both compute and donor device memory ahead of ComfyUI's compute-only check.
"""
multigpu_memory_log("patched_load_models_gpu", "start")
# Validate models argument loudly
if not isinstance(models, (list, tuple, set)):
logger.error("[MultiGPU Core Patching] CRITICAL: mm.load_models_gpu 'models' is not a list/tuple/set. Bypassing proactive patch.")
@@ -390,6 +534,8 @@ if hasattr(mm, 'load_models_gpu') and not hasattr(mm.load_models_gpu, "_distorch
if incoming_distorch_nonzero:
logger.mgpu_mm_log("Non-Zero incoming DisTorch2 model detected. Initiating proactive unload.")
# Proactively clear PromptExecutor caches ahead of major DisTorch2 load (Phase 1)
trigger_executor_cache_reset(reason="proactive_distorch_load", force=False)
if not hasattr(mm, 'current_loaded_models'):
raise AttributeError("comfy.model_management is missing 'current_loaded_models'. Proactive unload check failed.")
@@ -478,8 +624,11 @@ if hasattr(mm, 'load_models_gpu') and not hasattr(mm.load_models_gpu, "_distorch
elif incoming_is_distorch:
logger.mgpu_mm_log("Incoming DisTorch2 model requires 0.00GB; skipping proactive unload")
# Continue with original behavior
return original_load_models_gpu(models, memory_required, force_patch_weights, minimum_memory_required, force_full_load)
# Memory Logging
multigpu_memory_log("patched_load_models_gpu", "pre-original-call")
result = original_load_models_gpu(models, memory_required, force_patch_weights, minimum_memory_required, force_full_load)
multigpu_memory_log("patched_load_models_gpu", "post-original-call")
return result
# Mark and apply the patch
patched_load_models_gpu._distorch2_proactive_patched = True
@@ -650,4 +799,7 @@ for item in registration_data:
logger.info(dash_line)
# Register maintenance node
NODE_CLASS_MAPPINGS["FullCleanupMultiGPU"] = FullCleanupMultiGPU
logger.info(f"[MultiGPU] Registration complete. Final mappings: {', '.join(NODE_CLASS_MAPPINGS.keys())}")
+304 -2
View File
@@ -9,12 +9,175 @@ 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 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")
# Module-level cache for device list (populated once on first call)
_DEVICE_LIST_CACHE = None
# ==========================================================================================
# Executor Cache Management and CPU Monitoring (Phases 1, 2, 3)
# ==========================================================================================
# 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.
@@ -246,9 +409,16 @@ def soft_empty_cache_multigpu():
# Record pre-GC snapshot for general system view
multigpu_memory_log("general", "pre-soft-empty")
# Python GC (same as all implementations)
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)
all_devices = get_device_list()
@@ -263,41 +433,53 @@ def soft_empty_cache_multigpu():
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
logger.mgpu_mm_log(f"Cleared CUDA cache (and IPC if available) on {device_str}")
multigpu_memory_log("general", f"post-empty:{device_str}")
elif device_str == "mps":
if hasattr(torch, "mps") and hasattr(torch.mps, "empty_cache"):
logger.mgpu_mm_log("Clearing MPS cache")
multigpu_memory_log("general", f"pre-empty:{device_str}")
torch.mps.empty_cache()
logger.mgpu_mm_log("Cleared MPS cache")
multigpu_memory_log("general", f"post-empty:{device_str}")
elif device_str.startswith("xpu:"):
if hasattr(torch, "xpu") and hasattr(torch.xpu, "empty_cache"):
logger.mgpu_mm_log(f"Clearing XPU cache on {device_str}")
multigpu_memory_log("general", f"pre-empty:{device_str}")
torch.xpu.empty_cache()
logger.mgpu_mm_log(f"Cleared XPU cache on {device_str}")
multigpu_memory_log("general", f"post-empty:{device_str}")
elif device_str.startswith("npu:"):
if hasattr(torch, "npu") and hasattr(torch.npu, "empty_cache"):
logger.mgpu_mm_log(f"Clearing NPU cache on {device_str}")
multigpu_memory_log("general", f"pre-empty:{device_str}")
torch.npu.empty_cache()
logger.mgpu_mm_log(f"Cleared NPU cache on {device_str}")
multigpu_memory_log("general", f"post-empty:{device_str}")
elif device_str.startswith("mlu:"):
if hasattr(torch, "mlu") and hasattr(torch.mlu, "empty_cache"):
logger.mgpu_mm_log(f"Clearing MLU cache on {device_str}")
multigpu_memory_log("general", f"pre-empty:{device_str}")
torch.mlu.empty_cache()
logger.mgpu_mm_log(f"Cleared MLU cache on {device_str}")
multigpu_memory_log("general", f"post-empty:{device_str}")
elif device_str.startswith("corex:"):
if hasattr(torch, "corex") and hasattr(torch.corex, "empty_cache"):
logger.mgpu_mm_log(f"Clearing CoreX cache on {device_str}")
multigpu_memory_log("general", f"pre-empty:{device_str}")
torch.corex.empty_cache()
logger.mgpu_mm_log(f"Cleared CoreX cache on {device_str}")
multigpu_memory_log("general", f"post-empty:{device_str}")
# Record post-GC snapshot for general system view
multigpu_memory_log("general", "post-soft-empty")
@@ -353,7 +535,6 @@ def comfyui_memory_load(tag: str) -> str:
# 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)
@@ -477,6 +658,127 @@ def multigpu_memory_log(identifier: str, tag: str, log: logging.Logger = logger)
_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)
# ==========================================================================================
+12 -5
View File
@@ -16,8 +16,6 @@ import inspect
from collections import defaultdict
import comfy.model_management as mm
import comfy.model_patcher
from . import current_device
from .device_utils import get_device_list, soft_empty_cache_multigpu, multigpu_memory_log
safetensor_allocation_store = {}
safetensor_settings_store = {}
@@ -61,6 +59,7 @@ 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")
@@ -75,6 +74,12 @@ def register_patched_safetensor_modelpatcher():
del self._distorch_block_assignments
return result
# Track active DisTorch2 ModelPatcher lifecycle for leak diagnostics
try:
track_modelpatcher(self)
except Exception:
pass
if not hasattr(self.model, 'current_weight_patches_uuid'):
self.model.current_weight_patches_uuid = None
@@ -176,6 +181,7 @@ 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
@@ -383,6 +389,7 @@ 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
@@ -794,11 +801,11 @@ def calculate_safetensor_vvram_allocation(model_patcher, virtual_vram_str):
def override_class_with_distorch_safetensor_v2(cls):
"""DisTorch 2.0 wrapper for safetensor models"""
from . import current_device
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]
@@ -891,11 +898,11 @@ def override_class_with_distorch_safetensor_v2(cls):
def override_class_with_distorch_safetensor_v2_clip(cls):
"""DisTorch 2.0 wrapper for safetensor CLIP models"""
from . import current_device
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]
@@ -989,11 +996,11 @@ def override_class_with_distorch_safetensor_v2_clip(cls):
def override_class_with_distorch_safetensor_v2_clip_no_device(cls):
"""DisTorch 2.0 wrapper for safetensor CLIP models"""
from . import current_device
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]
+29 -1
View File
@@ -2,7 +2,7 @@ import torch
import folder_paths
from pathlib import Path
from nodes import NODE_CLASS_MAPPINGS
from .device_utils import get_device_list
from .device_utils import get_device_list, force_full_system_cleanup
class DeviceSelectorMultiGPU:
@classmethod
@@ -527,3 +527,31 @@ class DownloadAndLoadHyVideoTextEncoder:
def loadmodel(self, llm_model, clip_model, precision, apply_final_norm=False, hidden_state_skip_layer=2, quantization="disabled"):
original_loader = NODE_CLASS_MAPPINGS["DownloadAndLoadHyVideoTextEncoder"]()
return original_loader.loadmodel(llm_model, clip_model, precision, apply_final_norm, hidden_state_skip_layer, quantization)
class FullCleanupMultiGPU:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"image": ("IMAGE",),
"reason": ("STRING", {"default": "inline_node", "multiline": False}),
},
"optional": {
"force": ("BOOLEAN", {"default": True}),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image",)
FUNCTION = "cleanup"
CATEGORY = "multigpu/maintenance"
TITLE = "Full System Cleanup (MultiGPU)"
def cleanup(self, image, reason, force=True):
"""
Trigger the full system cleanup to match ComfyUI's 'Free model and node cache'.
Passthroughs the input image unchanged; summary is logged via MultiGPU logger.
"""
_ = force_full_system_cleanup(reason=reason, force=force)
return (image,)