diff --git a/.clinerules b/.clinerules index 0f77d4e..696f5ea 100644 --- a/.clinerules +++ b/.clinerules @@ -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 "` 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 diff --git a/__init__.py b/__init__.py index 54d8f36..3cfd08e 100644 --- a/__init__.py +++ b/__init__.py @@ -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, diff --git a/checkpoint_multigpu.py b/checkpoint_multigpu.py index 42b069b..6a30a9b 100644 --- a/checkpoint_multigpu.py +++ b/checkpoint_multigpu.py @@ -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") diff --git a/device_utils.py b/device_utils.py index b9989e3..fd60006 100644 --- a/device_utils.py +++ b/device_utils.py @@ -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} - : " + " | ".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) diff --git a/distorch.py b/distorch.py index ab81489..113aba9 100644 --- a/distorch.py +++ b/distorch.py @@ -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 = {} diff --git a/distorch_2.py b/distorch_2.py index a3fa23f..cd66bef 100644 --- a/distorch_2.py +++ b/distorch_2.py @@ -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] diff --git a/memory-bank/systemPatterns.md b/memory-bank/systemPatterns.md index 6517c55..632fdee 100644 --- a/memory-bank/systemPatterns.md +++ b/memory-bank/systemPatterns.md @@ -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. diff --git a/model_management_mgpu.py b/model_management_mgpu.py new file mode 100644 index 0000000..7a58aeb --- /dev/null +++ b/model_management_mgpu.py @@ -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} - : " + " | ".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 diff --git a/nodes.py b/nodes.py index fd452e8..a6f6a7d 100644 --- a/nodes.py +++ b/nodes.py @@ -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