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