From c4ae5e9e08e3987302d6f2d887881df76d27b34b Mon Sep 17 00:00:00 2001 From: John Pollock Date: Mon, 29 Sep 2025 03:53:21 -0500 Subject: [PATCH] extensive clean-up, WIP --- __init__.py | 390 +------------------------------ device_utils.py | 10 +- distorch_2.py | 137 ++++++----- memory-bank/cpu_leak_fix_plan.md | 20 ++ memory-bank/systemPatterns.md | 3 - model_management_mgpu.py | 197 +++------------- 6 files changed, 140 insertions(+), 617 deletions(-) create mode 100644 memory-bank/cpu_leak_fix_plan.md diff --git a/__init__.py b/__init__.py index 44dfc29..6551ed0 100644 --- a/__init__.py +++ b/__init__.py @@ -1,3 +1,5 @@ +DISTORCH2_UNLOAD_MODEL = False + import torch import logging import weakref @@ -17,17 +19,14 @@ from .model_management_mgpu import ( trigger_executor_cache_reset, check_cpu_memory_threshold, multigpu_memory_log, - prune_distorch_stores, - try_malloc_trim, - track_modelpatcher, force_full_system_cleanup, ) -# --- DisTorch V2 Logging Configuration --- + +MGPU_MM_LOG = True + # Set to "E" for Engineering (DEBUG) or "P" for Production (INFO) LOG_LEVEL = "P" - -# Configure logger logger = logging.getLogger("MultiGPU") logger.propagate = False @@ -39,25 +38,12 @@ if not logger.handlers: logger.addHandler(handler) logger.setLevel(log_level) -# --- MultiGPU Cleanup Policy Configuration --- -# Policy: off | threshold | every_load | every_load+threshold (alias threshold+every_load) -MGPU_CLEANUP_POLICY = os.getenv("MULTIGPU_CLEANUP_POLICY", "off").lower() -try: - MGPU_CPU_RESET_THRESHOLD = float(os.getenv("MULTIGPU_CPU_RESET_THRESHOLD", "0.85")) -except Exception: - MGPU_CPU_RESET_THRESHOLD = 0.85 -# Malloc trim (not part of Comfy Core): on | off -MGPU_MALLOC_TRIM = os.getenv("MULTIGPU_MALLOC_TRIM", "on").lower() - -logger.info(f"[MultiGPU Config] cleanup_policy={MGPU_CLEANUP_POLICY}, cpu_reset_threshold={MGPU_CPU_RESET_THRESHOLD:.2f}, malloc_trim={MGPU_MALLOC_TRIM}") - -MGPU_MM_LOG = True - def mgpu_mm_log_method(self, msg): if MGPU_MM_LOG: self.info(f"[MultiGPU Model Management] {msg}") logger.mgpu_mm_log = mgpu_mm_log_method.__get__(logger, type(logger)) +logger.mgpu_mm_log(f"[PHASE 1] DISTORCH2_UNLOAD_MODEL={DISTORCH2_UNLOAD_MODEL}") # Global device state management current_device = mm.get_torch_device() @@ -175,111 +161,6 @@ logger.debug(f"[MultiGPU DEBUG] Initial current_text_encoder_device: {current_te mm.get_torch_device = get_torch_device_patched mm.text_encoder_device = text_encoder_device_patched - -# ========================================================================================== -# Core Patching: ModelPatcher Lifecycle Tracking (__init__) -# ========================================================================================== -logger.info("[MultiGPU Core Patching] Applying ModelPatcher lifecycle tracking patch (__init__).") -if not hasattr(comfy.model_patcher.ModelPatcher, '_mgpu_lifecycle_patched'): - try: - _mgpu_original_modelpatcher_init = comfy.model_patcher.ModelPatcher.__init__ - - def _mgpu_patched_modelpatcher_init(self, *args, **kwargs): - _mgpu_original_modelpatcher_init(self, *args, **kwargs) - # Track all ModelPatcher instances at construction time - try: - track_modelpatcher(self) - except Exception: - pass - - comfy.model_patcher.ModelPatcher.__init__ = _mgpu_patched_modelpatcher_init - comfy.model_patcher.ModelPatcher._mgpu_lifecycle_patched = True - logger.info("[MultiGPU Core Patching] ModelPatcher.__init__ patched for lifecycle tracking.") - except Exception as e: - logger.error(f"[MultiGPU Core Patching] FAILED to patch ModelPatcher.__init__: {e}") - -# ========================================================================================== -# Core Patching: Fix Potential Reference Cycles in LoadedModel -# ========================================================================================== -if hasattr(mm, 'LoadedModel') and hasattr(mm.LoadedModel, '_set_model'): - logger.info("[MultiGPU Core Patching] Patching mm.LoadedModel._set_model and _switch_parent to reduce reference cycles.") - - _mgpu_original_set_model = mm.LoadedModel._set_model - - def _mgpu_patched_set_model(self, model): - patcher_id = id(model) - # Ensure attributes exist - if not hasattr(self, '_model'): - self._model = None - if not hasattr(self, '_parent_model'): - self._parent_model = None - if not hasattr(self, '_patcher_finalizer'): - self._patcher_finalizer = None - - # Reset refs - self._model = weakref.ref(model) - self._parent_model = None - - # Detach any previous finalizer - if self._patcher_finalizer is not None: - try: - self._patcher_finalizer.detach() - except Exception: - pass - self._patcher_finalizer = None - - # If clone, set parent and attach a weakref-based finalizer - parent = getattr(model, 'parent', None) - if parent is not None: - self._parent_model = weakref.ref(parent) - self_weak = weakref.ref(self) - - def _mgpu_finalize_clone(): - s = self_weak() - if s is not None and hasattr(s, '_switch_parent'): - logger.mgpu_mm_log(f"[MultiGPU_LoadedModel_Patch] Clone Patcher {patcher_id} GC'd. Switching LoadedModel to parent.") - s._switch_parent() - else: - logger.mgpu_mm_log(f"[MultiGPU_LoadedModel_Patch] Clone Patcher {patcher_id} GC'd. LoadedModel already gone or missing _switch_parent.") - - try: - self._patcher_finalizer = weakref.finalize(model, _mgpu_finalize_clone) - except Exception: - self._patcher_finalizer = None - else: - logger.mgpu_mm_log(f"[MultiGPU_LoadedModel_Patch] Set base model Patcher {patcher_id}.") - - mm.LoadedModel._set_model = _mgpu_patched_set_model - - # Patch _switch_parent to clear references explicitly - if hasattr(mm.LoadedModel, '_switch_parent'): - _mgpu_original_switch_parent = mm.LoadedModel._switch_parent - - def _mgpu_patched_switch_parent(self): - _mgpu_original_switch_parent(self) - # Clear parent and detach finalizer to avoid cycles - if hasattr(self, '_parent_model'): - self._parent_model = None - if hasattr(self, '_patcher_finalizer') and self._patcher_finalizer is not None: - try: - self._patcher_finalizer.detach() - except Exception: - pass - self._patcher_finalizer = None - - mm.LoadedModel._switch_parent = _mgpu_patched_switch_parent - else: - # Fallback if core ever changes - def _mgpu_fallback_switch_parent(self): - if hasattr(self, '_parent_model') and self._parent_model is not None: - parent_model = self._parent_model() - if parent_model is not None: - self._set_model(parent_model) - self._parent_model = None - mm.LoadedModel._switch_parent = _mgpu_fallback_switch_parent -else: - logger.warning("[MultiGPU Core Patching] mm.LoadedModel not found or missing _set_model; skip cycle patch.") - def check_module_exists(module_path): full_path = os.path.join(folder_paths.get_folder_paths("custom_nodes")[0], module_path) logger.debug(f"[MultiGPU] Checking for module at {full_path}") @@ -369,11 +250,6 @@ def soft_empty_cache_distorch2_patched(force=False): and force-triggered reset when explicitly requested (mirrors ComfyUI 'Free memory' button). """ multigpu_memory_log("patched_soft_empty", f"start:force={force}") - # Prune DisTorch stores before any clearing to drop stale references - try: - prune_distorch_stores() - except Exception: - pass is_distorch_active = False # Detect DisTorch2-managed models @@ -387,9 +263,9 @@ def soft_empty_cache_distorch2_patched(force=False): in_store = model_hash in safetensor_allocation_store alloc_value = safetensor_allocation_store.get(model_hash, "") model_name = type(getattr(mp, 'model', mp)).__name__ - keep_loaded = getattr(getattr(mp, 'model', None), '_mgpu_keep_loaded', False) + unload_distorch_model = getattr(getattr(mp, 'model', None), '_mgpu_unload_distorch_model', False) - logger.mgpu_mm_log(f"[DETECT_DEBUG] Model {i}: {model_name}, hash={model_hash[:8]}, in_store={in_store}, alloc_value='{alloc_value}', keep_loaded={keep_loaded}") + logger.mgpu_mm_log(f"[DETECT_DEBUG] Model {i}: {model_name}, hash={model_hash[:8]}, in_store={in_store}, alloc_value='{alloc_value}', unload_distorch_model={unload_distorch_model}") if in_store and alloc_value: is_distorch_active = True @@ -411,11 +287,6 @@ def soft_empty_cache_distorch2_patched(force=False): logger.mgpu_mm_log("DisTorch2 not active: delegating allocator cache clear (VRAM) to original mm.soft_empty_cache") original_soft_empty_cache(force) # Optional: return CPU heap to OS (not part of Comfy Core) - if MGPU_MALLOC_TRIM != "off": - try: - try_malloc_trim() - except Exception: - pass # Phase 1/3: forced executor reset mirrors ComfyUI 'Free memory' semantics if force: @@ -427,251 +298,6 @@ mm.soft_empty_cache = soft_empty_cache_distorch2_patched LARGE_MODEL_THRESHOLD = 2 * (1024**3) # 2 GB threshold for "large" models -# Patch only once (handles reloads) -if hasattr(mm, 'load_models_gpu') and not hasattr(mm.load_models_gpu, "_distorch2_proactive_patched"): - logger.info("[MultiGPU Core Patching] Patching mm.load_models_gpu for DisTorch2 proactive unloading") - - original_load_models_gpu = mm.load_models_gpu - - def patched_load_models_gpu(models, memory_required=0, force_patch_weights=False, minimum_memory_required=None, force_full_load=False): - """ - Proactively unload large models that are not needed when loading a large DisTorch2 model. - This frees both compute and donor device memory ahead of ComfyUI's compute-only check. - """ - multigpu_memory_log("patched_load_models_gpu", "start") - # Validate models argument loudly - if not isinstance(models, (list, tuple, set)): - logger.error("[MultiGPU Core Patching] CRITICAL: mm.load_models_gpu 'models' is not a list/tuple/set. Bypassing proactive patch.") - return original_load_models_gpu(models, memory_required, force_patch_weights, minimum_memory_required, force_full_load) - - # Detect incoming DisTorch2 request - incoming_is_distorch = False - incoming_distorch_nonzero = False - incoming_patchers = set() - incoming_loaded_names = [] - incoming_allowed_devices = None - incoming_compute_device = None - incoming_required_bytes = 0 - incoming_compute_planned_bytes = 0 - - for lm in models: - # Identify ModelPatcher (prefer direct; fall back to .patcher) - if hasattr(lm, "load_device"): - patcher = lm - elif hasattr(lm, "patcher"): - patcher = lm.patcher - else: - patcher = None - - model_for_hash = patcher if patcher is not None else getattr(lm, "model", lm) - - if patcher is not None: - incoming_patchers.add(patcher) - - # Determine required memory directly from ModelPatcher (no wrapper; no side effects) - device_str = str(patcher.load_device) - if patcher.current_loaded_device() == patcher.load_device: - required_bytes = patcher.model_size() - patcher.loaded_size() - else: - required_bytes = patcher.model_size() - - else: - device_str = "n/a" - required_bytes = 0 - - # Check DisTorch2 management via allocation store (unchanged trigger) - model_hash = create_safetensor_model_hash(model_for_hash, "load_patch_check") - if model_hash in safetensor_allocation_store and safetensor_allocation_store.get(model_hash): - incoming_is_distorch = True - if required_bytes > 0: - incoming_distorch_nonzero = True - if incoming_allowed_devices is None: - # Derive compute/donor devices from allocation string - alloc_str = safetensor_allocation_store.get(model_hash, "") - allowed = set() - if alloc_str: - parts = alloc_str.split("#", 1) - if len(parts) == 2 and parts[1]: - vram = parts[1] - segs = vram.split(";") - # compute device - if len(segs) >= 1 and segs[0]: - allowed.add(segs[0].strip()) - # donors list (comma-separated) - if len(segs) >= 3 and segs[2]: - for d in segs[2].split(","): - d = d.strip() - if d: - allowed.add(d) - else: - # Expert fraction string: "dev,fraction;dev2,fraction2;..." - for token in alloc_str.split(";"): - if "," in token: - dev, frac = token.split(",", 1) - fs = frac.strip() - numlike = fs.replace(".", "", 1).isdigit() - if numlike and float(fs) > 0.0: - allowed.add(dev.strip()) - if not allowed: - allowed = {str(patcher.load_device), "cpu"} - incoming_allowed_devices = allowed - # Determine compute device and planned bytes from allocation string - alloc = safetensor_allocation_store.get(model_hash, "") - if "#" in alloc: - vram = alloc.split("#", 1)[1] - segs = vram.split(";") - if len(segs) >= 2 and segs[0]: - incoming_compute_device = segs[0].strip() - try: - vvram_gb = float(segs[1]) - incoming_compute_planned_bytes = int(vvram_gb * (1024**3)) - except Exception: - incoming_compute_planned_bytes = 0 - else: - # Expert fractions: "dev,fraction;dev2,fraction2;..." - tokens = [t for t in alloc.split(";") if "," in t] - frac_map = {} - for t in tokens: - dev, frac = t.split(",", 1) - try: - frac_val = float(frac.strip()) - except Exception: - continue - frac_map[dev.strip()] = frac_val - if frac_map: - ld = str(patcher.load_device) - # Prefer the explicit load_device if present and > 0 - target_dev = ld if (ld in frac_map and frac_map[ld] > 0.0) else None - if target_dev is None: - # Otherwise pick highest positive fraction - target_dev = max((d for d,v in frac_map.items() if v > 0.0), key=lambda d: frac_map[d], default=None) - if target_dev is not None: - incoming_compute_device = target_dev - total = mm.get_total_memory(torch.device(target_dev)) - incoming_compute_planned_bytes = int(frac_map[target_dev] * (total or 0)) - if incoming_compute_device is None: - incoming_compute_device = str(patcher.load_device) - if incoming_compute_planned_bytes <= 0: - incoming_compute_planned_bytes = required_bytes - - # Log informational context with required bytes and device - try: - model_name = type(getattr(model_for_hash, "model", model_for_hash)).__name__ - except Exception: - model_name = "UnknownModel" - incoming_loaded_names.append(f"{model_name}:{required_bytes/(1024**3):.2f}GB req on {device_str}") - - if incoming_loaded_names: - logger.mgpu_mm_log(f"Incoming models summary: {', '.join(incoming_loaded_names)}") - - if incoming_distorch_nonzero: - logger.mgpu_mm_log("Non-Zero incoming DisTorch2 model detected. Initiating proactive unload.") - # Proactively clear PromptExecutor caches ahead of major DisTorch2 load (Phase 1) - trigger_executor_cache_reset(reason="proactive_distorch_load", force=False) - if not hasattr(mm, 'current_loaded_models'): - raise AttributeError("comfy.model_management is missing 'current_loaded_models'. Proactive unload check failed.") - - needed_patchers = incoming_patchers - # Need-based free on compute device only (scale-aware; core-aligned) - dev_str = incoming_compute_device or (next(iter(incoming_allowed_devices)) if incoming_allowed_devices else None) - freed_bytes = 0 - to_unload_indices = [] - unload_summaries = [] - if dev_str is not None: - dev_obj = torch.device(dev_str) - free_now = mm.get_free_memory(dev_obj) - try: - free_now_val = free_now[0] if isinstance(free_now, tuple) else free_now - except Exception: - free_now_val = free_now - # Use core-aligned immediate needs: planned vs. memory_required vs. minimum_memory_required - effective_needed = max(incoming_compute_planned_bytes or 0, memory_required or 0, minimum_memory_required or 0) - need_bytes = max(0, effective_needed - (free_now_val or 0)) - logger.mgpu_mm_log(f"Need calc on {dev_str}: effective_needed={effective_needed/(1024**3):.2f}GB, free_now={((free_now_val or 0)/(1024**3)):.2f}GB, need_bytes={need_bytes/(1024**3):.2f}GB") - if need_bytes > 0: - logger.mgpu_mm_log(f"Need-based unload on {dev_str}: need ~{need_bytes/(1024**3):.2f}GB") - # Build candidates on this device only, excluding needed patchers - candidates = [] - for idx, lm_cur in enumerate(mm.current_loaded_models): - mp_cur = getattr(lm_cur, 'model', None) - if mp_cur is None or mp_cur in needed_patchers: - continue - if str(getattr(lm_cur, "device", "")) != dev_str: - continue - size_cur = 0 - if hasattr(lm_cur, 'model_memory'): - try: - size_cur = lm_cur.model_memory() - except Exception: - size_cur = 0 - if size_cur <= 0 and hasattr(mp_cur, 'model_size'): - size_cur = mp_cur.model_size() - candidates.append((size_cur, idx, lm_cur, mp_cur)) - # Sort by size descending - candidates.sort(key=lambda x: x[0], reverse=True) - for size_cur, idx, lm_cur, mp_cur in candidates: - model_name = type(getattr(mp_cur, 'model', mp_cur)).__name__ - logger.mgpu_mm_log(f"Unloading model on {dev_str}: {model_name} (~{size_cur/(1024**3):.2f}GB)") - success = False - if hasattr(lm_cur, 'model_unload'): - success = lm_cur.model_unload(memory_to_free=None, unpatch_weights=True) - if success: - to_unload_indices.append(idx) - unload_summaries.append(f"{model_name}:{size_cur/(1024**3):.2f}GB") - freed_bytes += size_cur - if freed_bytes >= need_bytes: - break - - # Remove from management list and clear caches - unloaded_count = 0 - for idx in sorted(to_unload_indices, reverse=True): - mm.current_loaded_models.pop(idx) - unloaded_count += 1 - - if unloaded_count > 0: - logger.mgpu_mm_log(f"Proactively unloaded {unloaded_count} large model(s): {', '.join(unload_summaries)}") - logger.mgpu_mm_log("Performing multi-device cache clear after proactive unload") - # Force multi-device cache clear via patched soft_empty_cache (which detects DisTorch2) - mm.soft_empty_cache(force=True) - else: - # Lineage-aligned cache clear when no unloads happened: apply core 25% rule, per DisTorch devices - if incoming_allowed_devices is not None and mm.vram_state != mm.VRAMState.HIGH_VRAM: - triggered = [] - for dev_str in incoming_allowed_devices: - try: - dev_obj = torch.device(dev_str) - except Exception: - continue - free_total, free_torch = mm.get_free_memory(dev_obj, torch_free_too=True) - # free_total: system free; free_torch: torch reserved-but-free - if free_torch > free_total * 0.25: - triggered.append(dev_str) - if triggered: - logger.mgpu_mm_log(f"No unloads; 25% torch-cache rule triggered on: {', '.join(triggered)}. Calling soft_empty_cache()") - mm.soft_empty_cache(force=True) - else: - logger.mgpu_mm_log("No unloads; 25% torch-cache rule not met on DisTorch devices; skipping cache clear") - else: - logger.mgpu_mm_log("No unload candidates matched criteria and either HIGH_VRAM or no DisTorch devices; skipping cache clear") - elif incoming_is_distorch: - logger.mgpu_mm_log("Incoming DisTorch2 model requires 0.00GB; skipping proactive unload") - - # Memory Logging - multigpu_memory_log("patched_load_models_gpu", "pre-original-call") - result = original_load_models_gpu(models, memory_required, force_patch_weights, minimum_memory_required, force_full_load) - multigpu_memory_log("patched_load_models_gpu", "post-original-call") - - return result - - # Mark and apply the patch - patched_load_models_gpu._distorch2_proactive_patched = True - mm.load_models_gpu = patched_load_models_gpu -else: - if not hasattr(mm, 'load_models_gpu'): - raise AttributeError("comfy.model_management is missing 'load_models_gpu'. Core patching failed.") - else: - logger.debug("[MultiGPU Core Patching] mm.load_models_gpu already patched; skipping") - # Import advanced checkpoint loaders from .checkpoint_multigpu import ( CheckpointLoaderAdvancedMultiGPU, diff --git a/device_utils.py b/device_utils.py index fd60006..96ce073 100644 --- a/device_utils.py +++ b/device_utils.py @@ -223,19 +223,11 @@ def soft_empty_cache_multigpu(): Uses context managers to ensure the calling thread's device context is restored. """ # Import model management functions - from .model_management_mgpu import multigpu_memory_log, log_tracked_modelpatchers_status, try_malloc_trim + from .model_management_mgpu import multigpu_memory_log logger.mgpu_mm_log("soft_empty_cache_multigpu: starting GC and multi-device cache clear") - multigpu_memory_log("general", "pre-soft-empty") - multigpu_memory_log("general", "pre-gc") - log_tracked_modelpatchers_status(tag="pre-gc") gc.collect() - log_tracked_modelpatchers_status(tag="post-gc") - multigpu_memory_log("general", "post-gc") - logger.mgpu_mm_log("soft_empty_cache_multigpu: garbage collection complete") - - try_malloc_trim() # Clear cache for ALL devices (not just ComfyUI's single device) all_devices = get_device_list() diff --git a/distorch_2.py b/distorch_2.py index 34cf013..7a5c3f8 100644 --- a/distorch_2.py +++ b/distorch_2.py @@ -17,7 +17,8 @@ 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 +from .model_management_mgpu import multigpu_memory_log + safetensor_allocation_store = {} safetensor_settings_store = {} @@ -58,8 +59,7 @@ def register_patched_safetensor_modelpatcher(): # Patch ComfyUI's ModelPatcher if not hasattr(comfy.model_patcher.ModelPatcher, '_distorch_patched'): - # Patch LoadedModel.model_memory_required to drive behavior purely by keep_loaded flag - # This ensures precise control over unload behavior without further core patching + # Patch LoadedModel.model_memory_required to drive behavior purely by Phase 2 = unload_distorch_model flag from comfy.model_management import current_loaded_models original_loaded_model_memory_required = None @@ -70,48 +70,40 @@ def register_patched_safetensor_modelpatcher(): if original_loaded_model_memory_required is None: # Global patch of LoadedModel class if available - try: - import comfy.model_management as mm - if hasattr(mm, 'LoadedModel'): - original_loaded_model_memory_required = mm.LoadedModel.model_memory_required + import comfy.model_management as mm - def patched_loaded_model_memory_required(self, device): - """Drive unload behavior purely by keep_loaded flag""" - multigpu_memory_log("keep_loaded_memory_check", "start") - logger.mgpu_mm_log(f"[KEEP_LOADED_DEBUG] Memory assessment requested for model on device: {device}") + original_loaded_model_memory_required = mm.LoadedModel.model_memory_required - # Check if this is a DisTorch model with keep_loaded flag - keep_loaded = getattr(getattr(self, 'model', None), '_mgpu_keep_loaded', None) + def patched_loaded_model_memory_required(self, device): + """Drive unload behavior purely by unload_distorch_model flag""" + multigpu_memory_log("unload_distorch_model_memory_check", "start") + logger.mgpu_mm_log(f"[IS_DISTORCH_MODEL] Memory assessment requested for model on device: {device}") - if keep_loaded is not None: - # This is a DisTorch model - log the decision - model_name = type(getattr(self, 'model', mp)).__name__ if getattr(self, 'model', None) else "Unknown" - logger.mgpu_mm_log(f"[KEEP_LOADED_DEBUG] DisTorch model: {model_name}, keep_loaded={keep_loaded}") + # Check if this is a DisTorch model with unload_distorch_model flag + is_distorch_model = hasattr(getattr(getattr(self, 'model', None), 'model', None), '_mgpu_unload_distorch_model') - if keep_loaded is True: - logger.mgpu_mm_log("[KEEP_LOADED_DEBUG] keep_loaded=True - Reporting 0 bytes (prevents eviction)") - multigpu_memory_log("keep_loaded_memory_check", "prevents_eviction") - return 0 - elif keep_loaded is False: - # keep_loaded=False: return full device memory to guarantee eviction - total_device_memory = mm.get_total_memory(device) - memory_gb = total_device_memory / (1024**3) - logger.mgpu_mm_log(f"[KEEP_LOADED_DEBUG] keep_loaded=False - Reporting MAX memory ({memory_gb:.2f}GB) to force complete eviction") - multigpu_memory_log("keep_loaded_memory_check", f"forces_eviction:{memory_gb:.2f}gb") - return total_device_memory + model_name = type(getattr(getattr(self, 'model', None), 'model', None)).__name__ if getattr(getattr(self, 'model', None), 'model', None) else "Unknown" + logger.mgpu_mm_log(f"[IS_DISTORCH_MODEL] DisTorch model: {model_name}, is_distorch_model={is_distorch_model}") - # Not a DisTorch model - use original behavior - logger.mgpu_mm_log("[KEEP_LOADED_DEBUG] Non-DisTorch model - Using original Comfy memory calculation") - original_result = original_loaded_model_memory_required(self, device) - original_gb = original_result / (1024**3) if original_result else 0 - logger.mgpu_mm_log(f"[KEEP_LOADED_DEBUG] Original calculation returned: {original_gb:.2f}GB") - multigpu_memory_log("keep_loaded_memory_check", "end") - return original_result + if is_distorch_model: + if self.model.model._mgpu_unload_distorch_model: + total_device_memory = mm.get_total_memory(device) + memory_gb = total_device_memory / (1024**3) + logger.mgpu_mm_log(f"[IS_DISTORCH_MODEL] _mgpu_unload_distorch_model=True - Reporting MAX memory ({memory_gb:.2f}GB) to force complete eviction") + return total_device_memory + else: + logger.mgpu_mm_log("[IS_DISTORCH_MODEL] _mgpu_unload_distorch_model=False - Reporting 0 bytes (prevents eviction)") + return 0 - mm.LoadedModel.model_memory_required = patched_loaded_model_memory_required + # Not a DisTorch model - use original behavior + logger.mgpu_mm_log("[IS_DISTORCH_MODEL] Non-DisTorch model - Using original Comfy memory calculation") + original_result = original_loaded_model_memory_required(self, device) + original_gb = original_result / (1024**3) if original_result else 0 + logger.mgpu_mm_log(f"[IS_DISTORCH_MODEL] Original calculation returned: {original_gb:.2f}GB") + multigpu_memory_log("keep_loaded_memory_check", "end") + return original_result - except (ImportError, AttributeError): - logging.warning("[MultiGPU DisTorch] Could not patch LoadedModel.model_memory_required - unload behavior may be inconsistent") + mm.LoadedModel.model_memory_required = patched_loaded_model_memory_required original_partially_load = comfy.model_patcher.ModelPatcher.partially_load @@ -134,12 +126,6 @@ def register_patched_safetensor_modelpatcher(): del self._distorch_block_assignments return result - # Track active DisTorch2 ModelPatcher lifecycle for leak diagnostics - try: - track_modelpatcher(self) - except Exception: - pass - if not hasattr(self.model, 'current_weight_patches_uuid'): self.model.current_weight_patches_uuid = None @@ -889,6 +875,17 @@ def override_class_with_distorch_safetensor_v2(cls): def override(self, *args, compute_device=None, virtual_vram_gb=4.0, donor_device="cpu", expert_mode_allocations="", keep_loaded=True, **kwargs): + from . import DISTORCH2_UNLOAD_MODEL + + logger.mgpu_mm_log(f"[PHASE 1] DISTORCH2_UNLOAD_MODEL INITIAL SETTING ={DISTORCH2_UNLOAD_MODEL}") + + unload_distorch_model = not keep_loaded + + if unload_distorch_model: + logger.mgpu_mm_log("[PHASE 1] DisTorch2 with keep_loaded=False. Setting DISTORCH2_UNLOAD_MODEL=True") + DISTORCH2_UNLOAD_MODEL = True + logger.mgpu_mm_log(f"[PHASE 1] DISTORCH2_UNLOAD_MODEL UPDATED SETTING ={DISTORCH2_UNLOAD_MODEL}") + from . import set_current_device if compute_device is not None: set_current_device(compute_device) @@ -928,11 +925,15 @@ def override_class_with_distorch_safetensor_v2(cls): logger.info(f"[MultiGPU DisTorch V2] Full allocation string: {full_allocation}") - # Store keep_loaded in the model for later retrieval by unload_all_models patch + logger.mgpu_mm_log(f"[PHASE 1] Tracking setting of '_mgpu_unload_distorch_model to unload_distorch_model: {unload_distorch_model}") + + # Store unload_distorch_model in the model for later retrieval by unload_all_models patch if hasattr(out[0], 'model'): - out[0].model._mgpu_keep_loaded = keep_loaded + logger.mgpu_mm_log(f"[PHASE 1] model {out[0].model.__class__.__name__}'_mgpu_unload_distorch_model set to: {unload_distorch_model}") + out[0].model._mgpu_unload_distorch_model = unload_distorch_model elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'): - out[0].patcher.model._mgpu_keep_loaded = keep_loaded + logger.mgpu_mm_log(f"[PHASE 1] model {out[0].patcher.model.__class__.__name__}'_mgpu_unload_distorch_model set to: {unload_distorch_model}") + out[0].patcher.model._mgpu_unload_distorch_model = unload_distorch_model return out @@ -971,6 +972,17 @@ def override_class_with_distorch_safetensor_v2_clip(cls): def override(self, *args, device=None, virtual_vram_gb=4.0, # Changed from compute_device donor_device="cpu", expert_mode_allocations="", keep_loaded=True, **kwargs): + from . import DISTORCH2_UNLOAD_MODEL + + logger.mgpu_mm_log(f"[PHASE 1] DISTORCH2_UNLOAD_MODEL INITIAL SETTING ={DISTORCH2_UNLOAD_MODEL}") + + unload_distorch_model = not keep_loaded + + if unload_distorch_model: + logger.mgpu_mm_log("[PHASE 1] DisTorch2 with keep_loaded=False. Setting DISTORCH2_UNLOAD_MODEL=True") + DISTORCH2_UNLOAD_MODEL = True + logger.mgpu_mm_log(f"[PHASE 1] DISTORCH2_UNLOAD_MODEL UPDATED SETTING ={DISTORCH2_UNLOAD_MODEL}") + from . import set_current_text_encoder_device # Use text encoder device setter if device is not None: set_current_text_encoder_device(device) @@ -987,10 +999,15 @@ def override_class_with_distorch_safetensor_v2_clip(cls): out = fn(*args, **kwargs) # Store keep_loaded in the model for later retrieval by unload_all_models patch + logger.mgpu_mm_log(f"[PHASE 1] Tracking setting of '_mgpu_unload_distorch_model to unload_distorch_model: {unload_distorch_model}") + + # Store unload_distorch_model in the model for later retrieval by unload_all_models patch if hasattr(out[0], 'model'): - out[0].model._mgpu_keep_loaded = keep_loaded + logger.mgpu_mm_log(f"[PHASE 1] model {out[0].model.__class__.__name__}'_mgpu_unload_distorch_model set to: {unload_distorch_model}") + out[0].model._mgpu_unload_distorch_model = unload_distorch_model elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'): - out[0].patcher.model._mgpu_keep_loaded = keep_loaded + logger.mgpu_mm_log(f"[PHASE 1] model {out[0].patcher.model.__class__.__name__}'_mgpu_unload_distorch_model set to: {unload_distorch_model}") + out[0].patcher.model._mgpu_unload_distorch_model = unload_distorch_model vram_string = "" if virtual_vram_gb > 0: @@ -1054,6 +1071,18 @@ def override_class_with_distorch_safetensor_v2_clip_no_device(cls): def override(self, *args, device=None, virtual_vram_gb=4.0, # Changed from compute_device donor_device="cpu", expert_mode_allocations="", keep_loaded=True, **kwargs): + from . import DISTORCH2_UNLOAD_MODEL + + logger.mgpu_mm_log(f"[PHASE 1] DISTORCH2_UNLOAD_MODEL INITIAL SETTING ={DISTORCH2_UNLOAD_MODEL}") + + unload_distorch_model = not keep_loaded + + if unload_distorch_model: + logger.mgpu_mm_log("[PHASE 1] DisTorch2 with keep_loaded=False. Setting DISTORCH2_UNLOAD_MODEL=True") + DISTORCH2_UNLOAD_MODEL = True + logger.mgpu_mm_log(f"[PHASE 1] DISTORCH2_UNLOAD_MODEL UPDATED SETTING ={DISTORCH2_UNLOAD_MODEL}") + + from . import set_current_text_encoder_device # Use text encoder device setter if device is not None: set_current_text_encoder_device(device) @@ -1067,11 +1096,15 @@ def override_class_with_distorch_safetensor_v2_clip_no_device(cls): # Call the main function once out = fn(*args, **kwargs) - # Store keep_loaded in the model for later retrieval by unload_all_models patch + logger.mgpu_mm_log(f"[PHASE 1] Tracking setting of '_mgpu_unload_distorch_model to unload_distorch_model: {unload_distorch_model}") + + # Store unload_distorch_model in the model for later retrieval by unload_all_models patch if hasattr(out[0], 'model'): - out[0].model._mgpu_keep_loaded = keep_loaded + logger.mgpu_mm_log(f"[PHASE 1] model {out[0].model.__class__.__name__}'_mgpu_unload_distorch_model set to: {unload_distorch_model}") + out[0].model._mgpu_unload_distorch_model = unload_distorch_model elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'): - out[0].patcher.model._mgpu_keep_loaded = keep_loaded + logger.mgpu_mm_log(f"[PHASE 1] model {out[0].patcher.model.__class__.__name__}'_mgpu_unload_distorch_model set to: {unload_distorch_model}") + out[0].patcher.model._mgpu_unload_distorch_model = unload_distorch_model vram_string = "" if virtual_vram_gb > 0: diff --git a/memory-bank/cpu_leak_fix_plan.md b/memory-bank/cpu_leak_fix_plan.md new file mode 100644 index 0000000..72e0898 --- /dev/null +++ b/memory-bank/cpu_leak_fix_plan.md @@ -0,0 +1,20 @@ +No. It is clear that you do not given multiple failed implementations past this point. So, lets do this in phases. + +Phase 1: Implement DISTORCH2_UNLOAD_MODEL Global correctly. It should be set to True when it sees a keep_loaded=false and should be reset at the end of our patched unload_all_models. No other code changes. Document with device snapshot and memory datalog each time a new operation is done - so when it is set and unset so it can been seen in the datalog. + +Phase 2: In Distorch_2.py, implement `_mgpu_unload` flag to any DisTorch model when keep_loaded=false and at the same time as setting DISTORCH2_UNLOAD_MODEL=True. In our patched unload_all_models() we create a simple evaluatioon loop with my pseudocode: + +if hasattr(getattr(model, 'model', None), '_mgpu_unload'): + multigpu_memory_log(model_hash, "_mgpu_unload=true") + logger.mgpu_mm_log(f"[THREE_FLAG_DEBUG] model has `_mpgu_unload` flag") +else: + logger.mgpu_mm_log(f"[THREE_FLAG_DEBUG] model does not have _mpgu_unload flag") + +At the end of the loop no matter what calls it, DISTORCH2_UNLOAD_MODEL = FALSE with an appropriate log: + logger.mgpu_mm_log("Setting DISTORCH2_UNLOAD_MODEL=False") + +Phase 3: Replace existing faulty retention or ejection logic with the loop from Phase 2: + +1. At the beginning of our patched unload_all_models, check DISTORCH2_UNLOAD_MODEL + If FALSE: run _original_unload_all_models() + IF TRUE: Using the loop from Phase 2, apply only the unload_all_models routine to the models with `_mpgu_unload` flag set, else do nothing to other models, exactly like Else loop from Phase 2. diff --git a/memory-bank/systemPatterns.md b/memory-bank/systemPatterns.md index 632fdee..096c676 100644 --- a/memory-bank/systemPatterns.md +++ b/memory-bank/systemPatterns.md @@ -341,7 +341,6 @@ def benchmark_allocation_performance(model, hardware_config, allocation_configs) - 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 @@ -386,8 +385,6 @@ def benchmark_allocation_performance(model, hardware_config, allocation_configs) - `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 index 1697aa7..6e77dc4 100644 --- a/model_management_mgpu.py +++ b/model_management_mgpu.py @@ -17,36 +17,10 @@ import ctypes import comfy.model_patcher from collections import defaultdict + + logger = logging.getLogger("MultiGPU") -# ========================================================================================== -# GC Anchor System for Model Retention Testing -# ========================================================================================== - -# Global anchor set to prevent GC of models with keep_loaded=True -_MGPU_RETENTION_ANCHORS = set() - -def add_retention_anchor(model_patcher, reason="keep_loaded"): - """Add a model patcher to the GC anchor set to prevent premature garbage collection""" - if model_patcher is not None: - _MGPU_RETENTION_ANCHORS.add(model_patcher) - model_name = type(getattr(model_patcher, 'model', model_patcher)).__name__ - logger.mgpu_mm_log(f"[GC_ANCHOR] Added retention anchor for {model_name}, reason: {reason}, total anchors: {len(_MGPU_RETENTION_ANCHORS)}") - -def remove_retention_anchor(model_patcher, reason="cleanup"): - """Remove a model patcher from the GC anchor set""" - if model_patcher is not None and model_patcher in _MGPU_RETENTION_ANCHORS: - _MGPU_RETENTION_ANCHORS.discard(model_patcher) - model_name = type(getattr(model_patcher, 'model', model_patcher)).__name__ - logger.mgpu_mm_log(f"[GC_ANCHOR] Removed retention anchor for {model_name}, reason: {reason}, total anchors: {len(_MGPU_RETENTION_ANCHORS)}") - -def clear_all_retention_anchors(reason="manual_clear"): - """Clear all retention anchors""" - count = len(_MGPU_RETENTION_ANCHORS) - _MGPU_RETENTION_ANCHORS.clear() - logger.mgpu_mm_log(f"[GC_ANCHOR] Cleared all {count} retention anchors, reason: {reason}") - - # ========================================================================================== # Model Analysis and Store Management (DisTorch V1 & V2) # ========================================================================================== @@ -85,52 +59,6 @@ def create_model_hash(model, caller): 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() - - logger.mgpu_mm_log(f"[PRUNE_DEBUG] Starting prune - current_loaded_models count: {len(mm.current_loaded_models)}") - - for i, lm in enumerate(mm.current_loaded_models): - mp = lm.model - if mp is not None: - try: - hash_v2 = create_safetensor_model_hash(mp, "prune_check_v2") - hash_v1 = create_model_hash(mp, "prune_check_v1") - active_hashes_v2.add(hash_v2) - active_hashes_v1.add(hash_v1) - - model_name = type(getattr(mp, 'model', mp)).__name__ - keep_loaded = getattr(getattr(mp, 'model', None), '_mgpu_keep_loaded', False) - has_v2_alloc = hash_v2 in safetensor_allocation_store - logger.mgpu_mm_log(f"[PRUNE_DEBUG] Model {i}: {model_name}, keep_loaded={keep_loaded}, hash={hash_v2[:8]}, has_v2_allocation={has_v2_alloc}") - except Exception as e: - logger.mgpu_mm_log(f"[PRUNE_DEBUG] Model {i}: Error getting hash - {e}") - - logger.mgpu_mm_log(f"[PRUNE_DEBUG] Active hashes V2: {len(active_hashes_v2)}, Store has: {len(safetensor_allocation_store)}") - - # V1 pruning - stale_v1 = set(model_allocation_store.keys()) - active_hashes_v1 - if stale_v1: - logger.mgpu_mm_log(f"[MultiGPU_Memory_Management] Pruning {len(stale_v1)} DisTorch V1 entries") - for k in stale_v1: - del model_allocation_store[k] - - # V2 pruning with diagnostics - for store, name in ((safetensor_allocation_store, "allocation"), (safetensor_settings_store, "settings")): - stale_v2 = set(store.keys()) - active_hashes_v2 - if stale_v2: - logger.mgpu_mm_log(f"[MultiGPU_Memory_Management] Would prune {len(stale_v2)} V2 {name} entries: {[h[:8] for h in list(stale_v2)[:5]]}") - for k in stale_v2: - del store[k] - else: - logger.mgpu_mm_log(f"[PRUNE_DEBUG] No stale {name} entries to prune") - - logger.mgpu_mm_log(f"[PRUNE_DEBUG] After pruning - V2 allocation store has: {len(safetensor_allocation_store)} entries") - multigpu_memory_log("distorch_prune", "end") - # ========================================================================================== # Memory Logging Infrastructure # ========================================================================================== @@ -203,64 +131,6 @@ def multigpu_memory_log(identifier, tag): _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 @@ -270,25 +140,6 @@ 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""" @@ -306,17 +157,12 @@ def trigger_executor_cache_reset(reason="policy", force=False): 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): @@ -370,23 +216,30 @@ def force_full_system_cleanup(reason="manual", force=True): logger.mgpu_mm_log(summary) return summary - # ========================================================================================== -# Core Patching: unload_all_models with keep_loaded retention +# Core Patching: unload_all_models # ========================================================================================== -if hasattr(mm, 'unload_all_models') and not hasattr(mm.unload_all_models, '_mgpu_keep_loaded_patched'): - logger.info("[MultiGPU Core Patching] Patching mm.unload_all_models to respect keep_loaded flag for DisTorch models") +if not hasattr(mm.unload_all_models, '_mgpu_eject_distorch_patched'): + logger.info("[MultiGPU Core Patching] Patching mm.unload_all_models for DisTorch2 ejection support") _mgpu_original_unload_all_models = mm.unload_all_models def _mgpu_patched_unload_all_models(): """ - Patched mm.unload_all_models that preserves DisTorch models with _mgpu_keep_loaded=True. + Patched mm.unload_all_models that checks to see if the . All other models (including DisTorch models without the flag) unload normally. """ - logger.mgpu_mm_log(f"[UNLOAD_DEBUG] Patched unload_all_models called - initial model count: {len(mm.current_loaded_models)}") - + from . import DISTORCH2_UNLOAD_MODEL + + logger.mgpu_mm_log(f"[Phase 2 Debug] Patched unload_all_models called - initial model count: {len(mm.current_loaded_models)}") + logger.mgpu_mm_log(f"[Phase 2 Debug] DISTORCH2_UNLOAD_MODEL={DISTORCH2_UNLOAD_MODEL}") + + if DISTORCH2_UNLOAD_MODEL == False: + logger.mgpu_mm_log("[Phase 2 Debug] Standard unload_all_models() called from Comfy Core") + _mgpu_original_unload_all_models() + return + # Direct approach: iterate through loaded models and selectively unload models_to_unload = [] kept_models = [] @@ -415,10 +268,10 @@ if hasattr(mm, 'unload_all_models') and not hasattr(mm.unload_all_models, '_mgpu models_to_unload.append(lm) logger.mgpu_mm_log(f"[UNLOAD_DEBUG] Final counts - kept_models: {len(kept_models)}, models_to_unload: {len(models_to_unload)}") - + if kept_models: logger.mgpu_mm_log(f"Found {len(kept_models)} model(s) to retain, unloading {len(models_to_unload)} model(s)") - + # Unload models that don't have keep_loaded flag for lm in models_to_unload: try: @@ -426,7 +279,7 @@ if hasattr(mm, 'unload_all_models') and not hasattr(mm.unload_all_models, '_mgpu logger.debug(f"Unloaded model: {type(lm.model.model).__name__ if lm.model else 'Unknown'}") except Exception as e: logger.warning(f"Error unloading model: {e}") - + # Remove unloaded models from current_loaded_models mm.current_loaded_models = kept_models logger.mgpu_mm_log(f"[UNLOAD_DEBUG] Updated mm.current_loaded_models, new count: {len(mm.current_loaded_models)}") @@ -434,12 +287,14 @@ if hasattr(mm, 'unload_all_models') and not hasattr(mm.unload_all_models, '_mgpu else: logger.mgpu_mm_log("No models with keep_loaded=True found - delegating to original unload_all_models") _mgpu_original_unload_all_models() + + # Phase 1: Reset DISTORCH2_UNLOAD_MODEL flag at end of unload (REGARDLESS) + logger.mgpu_mm_log("[PHASE1_DEBUG] Setting DISTORCH2_UNLOAD_MODEL=False at end of unload") + multigpu_memory_log("distorch_flag", "reset_false") + DISTORCH2_UNLOAD_MODEL = False mm.unload_all_models = _mgpu_patched_unload_all_models - mm.unload_all_models._mgpu_keep_loaded_patched = True + mm.unload_all_models._mgpu_eject_distorch_patched = True logger.info("[MultiGPU Core Patching] mm.unload_all_models patched successfully") else: - if not hasattr(mm, 'unload_all_models'): - logger.warning("[MultiGPU Core Patching] mm.unload_all_models not found - cannot patch keep_loaded retention") - else: - logger.debug("[MultiGPU Core Patching] mm.unload_all_models already patched for keep_loaded - skipping") + logger.debug("[MultiGPU Core Patching] mm.unload_all_models already patched - skipping")