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:
+163
-11
@@ -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
@@ -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
@@ -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]
|
||||
|
||||
@@ -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,)
|
||||
|
||||
Reference in New Issue
Block a user