diff --git a/distorch_2.py b/distorch_2.py index 0c4b9ad..f595397 100644 --- a/distorch_2.py +++ b/distorch_2.py @@ -923,17 +923,54 @@ def override_class_with_distorch_safetensor_v2(cls): logger.info(f"[MultiGPU DisTorch V2] Full allocation string: {full_allocation}") - logger.mgpu_mm_log(f"[PHASE 1] Tracking setting of '_mgpu_unload_distorch_model to unload_distorch_model: {unload_distorch_model}") + logger.mgpu_mm_log(f"[FLAG_SET_START] Setting '_mgpu_unload_distorch_model' to: {unload_distorch_model} (keep_loaded={keep_loaded})") - # Store unload_distorch_model in the model for later retrieval by unload_all_models patch + # DIAGNOSTIC: Log full object chain at SET time if hasattr(out[0], 'model'): - 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 + mp = out[0] # This is the ModelPatcher + mp_id = id(mp) + inner_model = getattr(mp, 'model', None) + inner_model_id = id(inner_model) if inner_model else None + inner_model_name = type(inner_model).__name__ if inner_model else "None" + + # Format inner_model_id properly for f-string + inner_id_str = f"0x{inner_model_id:x}" if inner_model_id is not None else "None" + + logger.mgpu_mm_log(f"[OBJECT_CHAIN_SET] ModelPatcher: mp_id=0x{mp_id:x}, inner_model_id={inner_id_str}, inner_model_type={inner_model_name}") + + # FIX: Store flag on ModelPatcher itself (not inner model) + # This aligns with where it will be READ in model_management_mgpu.py + mp._mgpu_unload_distorch_model = unload_distorch_model + logger.mgpu_mm_log(f"[FLAG_SET_LOCATION] Set on ModelPatcher (mp_id=0x{mp_id:x}): mp._mgpu_unload_distorch_model = {unload_distorch_model}") + + # Also set on inner model for backwards compatibility during transition + if inner_model: + inner_model._mgpu_unload_distorch_model = unload_distorch_model + logger.mgpu_mm_log(f"[FLAG_SET_COMPAT] Also set on inner model (inner_model_id=0x{inner_model_id:x}) for compatibility") + elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'): - 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 + mp = out[0].patcher # This is the ModelPatcher + mp_id = id(mp) + inner_model = getattr(mp, 'model', None) + inner_model_id = id(inner_model) if inner_model else None + inner_model_name = type(inner_model).__name__ if inner_model else "None" + + # Format inner_model_id properly for f-string + inner_id_str = f"0x{inner_model_id:x}" if inner_model_id is not None else "None" + + logger.mgpu_mm_log(f"[OBJECT_CHAIN_SET] ModelPatcher via patcher: mp_id=0x{mp_id:x}, inner_model_id={inner_id_str}, inner_model_type={inner_model_name}") + + # FIX: Store flag on ModelPatcher itself + mp._mgpu_unload_distorch_model = unload_distorch_model + logger.mgpu_mm_log(f"[FLAG_SET_LOCATION] Set on ModelPatcher (mp_id=0x{mp_id:x}): mp._mgpu_unload_distorch_model = {unload_distorch_model}") + + # Also set on inner model for backwards compatibility + if inner_model: + inner_model._mgpu_unload_distorch_model = unload_distorch_model + logger.mgpu_mm_log(f"[FLAG_SET_COMPAT] Also set on inner model (inner_model_id=0x{inner_model_id:x}) for compatibility") if unload_distorch_model: + logger.mgpu_mm_log("[FLAG_TRIGGER] unload_distorch_model=True, triggering full system cleanup") force_full_system_cleanup(reason="policy_every_load", force=True) return out @@ -998,22 +1035,56 @@ def override_class_with_distorch_safetensor_v2_clip(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}") + logger.mgpu_mm_log(f"[FLAG_SET_START] Setting '_mgpu_unload_distorch_model' to: {unload_distorch_model} (keep_loaded={keep_loaded})") - # Store unload_distorch_model in the model for later retrieval by unload_all_models patch + # DIAGNOSTIC: Log full object chain at SET time if hasattr(out[0], 'model'): - 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 + mp = out[0] # This is the ModelPatcher + mp_id = id(mp) + inner_model = getattr(mp, 'model', None) + inner_model_id = id(inner_model) if inner_model else None + inner_model_name = type(inner_model).__name__ if inner_model else "None" + + # Format inner_model_id properly for f-string + inner_id_str = f"0x{inner_model_id:x}" if inner_model_id is not None else "None" + + logger.mgpu_mm_log(f"[OBJECT_CHAIN_SET] CLIP ModelPatcher: mp_id=0x{mp_id:x}, inner_model_id={inner_id_str}, inner_model_type={inner_model_name}") + + # FIX: Store flag on ModelPatcher itself + mp._mgpu_unload_distorch_model = unload_distorch_model + logger.mgpu_mm_log(f"[FLAG_SET_LOCATION] Set on ModelPatcher (mp_id=0x{mp_id:x}): mp._mgpu_unload_distorch_model = {unload_distorch_model}") + + # Also set on inner model for backwards compatibility + if inner_model: + inner_model._mgpu_unload_distorch_model = unload_distorch_model + logger.mgpu_mm_log(f"[FLAG_SET_COMPAT] Also set on inner model (inner_model_id=0x{inner_model_id:x}) for compatibility") + elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'): - 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 + mp = out[0].patcher # This is the ModelPatcher + mp_id = id(mp) + inner_model = getattr(mp, 'model', None) + inner_model_id = id(inner_model) if inner_model else None + inner_model_name = type(inner_model).__name__ if inner_model else "None" + + # Format inner_model_id properly for f-string + inner_id_str = f"0x{inner_model_id:x}" if inner_model_id is not None else "None" + + logger.mgpu_mm_log(f"[OBJECT_CHAIN_SET] CLIP ModelPatcher via patcher: mp_id=0x{mp_id:x}, inner_model_id={inner_id_str}, inner_model_type={inner_model_name}") + + # FIX: Store flag on ModelPatcher itself + mp._mgpu_unload_distorch_model = unload_distorch_model + logger.mgpu_mm_log(f"[FLAG_SET_LOCATION] Set on ModelPatcher (mp_id=0x{mp_id:x}): mp._mgpu_unload_distorch_model = {unload_distorch_model}") + + # Also set on inner model for backwards compatibility + if inner_model: + inner_model._mgpu_unload_distorch_model = unload_distorch_model + logger.mgpu_mm_log(f"[FLAG_SET_COMPAT] Also set on inner model (inner_model_id=0x{inner_model_id:x}) for compatibility") vram_string = "" if virtual_vram_gb > 0: - vram_string = f"{device};{virtual_vram_gb};{donor_device}" # Changed from compute_device - elif expert_mode_allocations: # Only include device if there's an expert string - vram_string = device # Changed from compute_device + vram_string = f"{device};{virtual_vram_gb};{donor_device}" + elif expert_mode_allocations: + vram_string = device full_allocation = f"{expert_mode_allocations}#{vram_string}" if expert_mode_allocations or vram_string else "" @@ -1036,6 +1107,7 @@ def override_class_with_distorch_safetensor_v2_clip(cls): safetensor_settings_store[model_hash] = settings_hash if unload_distorch_model: + logger.mgpu_mm_log("[FLAG_TRIGGER] unload_distorch_model=True, triggering full system cleanup") force_full_system_cleanup(reason="policy_every_load", force=True) return out @@ -1096,21 +1168,56 @@ def override_class_with_distorch_safetensor_v2_clip_no_device(cls): # Call the main function once out = fn(*args, **kwargs) - logger.mgpu_mm_log(f"[PHASE 1] Tracking setting of '_mgpu_unload_distorch_model to unload_distorch_model: {unload_distorch_model}") + logger.mgpu_mm_log(f"[FLAG_SET_START] Setting '_mgpu_unload_distorch_model' to: {unload_distorch_model} (keep_loaded={keep_loaded})") - # Store unload_distorch_model in the model for later retrieval by unload_all_models patch + # DIAGNOSTIC: Log full object chain at SET time if hasattr(out[0], 'model'): - 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 + mp = out[0] # This is the ModelPatcher + mp_id = id(mp) + inner_model = getattr(mp, 'model', None) + inner_model_id = id(inner_model) if inner_model else None + inner_model_name = type(inner_model).__name__ if inner_model else "None" + + # Format inner_model_id properly for f-string + inner_id_str = f"0x{inner_model_id:x}" if inner_model_id is not None else "None" + + logger.mgpu_mm_log(f"[OBJECT_CHAIN_SET] CLIP_NoDevice ModelPatcher: mp_id=0x{mp_id:x}, inner_model_id={inner_id_str}, inner_model_type={inner_model_name}") + + # FIX: Store flag on ModelPatcher itself + mp._mgpu_unload_distorch_model = unload_distorch_model + logger.mgpu_mm_log(f"[FLAG_SET_LOCATION] Set on ModelPatcher (mp_id=0x{mp_id:x}): mp._mgpu_unload_distorch_model = {unload_distorch_model}") + + # Also set on inner model for backwards compatibility + if inner_model: + inner_model._mgpu_unload_distorch_model = unload_distorch_model + logger.mgpu_mm_log(f"[FLAG_SET_COMPAT] Also set on inner model (inner_model_id=0x{inner_model_id:x}) for compatibility") + elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'): - 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 + mp = out[0].patcher # This is the ModelPatcher + mp_id = id(mp) + inner_model = getattr(mp, 'model', None) + inner_model_id = id(inner_model) if inner_model else None + inner_model_name = type(inner_model).__name__ if inner_model else "None" + + # Format inner_model_id properly for f-string + inner_id_str = f"0x{inner_model_id:x}" if inner_model_id is not None else "None" + + logger.mgpu_mm_log(f"[OBJECT_CHAIN_SET] CLIP_NoDevice ModelPatcher via patcher: mp_id=0x{mp_id:x}, inner_model_id={inner_id_str}, inner_model_type={inner_model_name}") + + # FIX: Store flag on ModelPatcher itself + mp._mgpu_unload_distorch_model = unload_distorch_model + logger.mgpu_mm_log(f"[FLAG_SET_LOCATION] Set on ModelPatcher (mp_id=0x{mp_id:x}): mp._mgpu_unload_distorch_model = {unload_distorch_model}") + + # Also set on inner model for backwards compatibility + if inner_model: + inner_model._mgpu_unload_distorch_model = unload_distorch_model + logger.mgpu_mm_log(f"[FLAG_SET_COMPAT] Also set on inner model (inner_model_id=0x{inner_model_id:x}) for compatibility") vram_string = "" if virtual_vram_gb > 0: - vram_string = f"{device};{virtual_vram_gb};{donor_device}" # Changed from compute_device - elif expert_mode_allocations: # Only include device if there's an expert string - vram_string = device # Changed from compute_device + vram_string = f"{device};{virtual_vram_gb};{donor_device}" + elif expert_mode_allocations: + vram_string = device full_allocation = f"{expert_mode_allocations}#{vram_string}" if expert_mode_allocations or vram_string else "" @@ -1133,6 +1240,7 @@ def override_class_with_distorch_safetensor_v2_clip_no_device(cls): safetensor_settings_store[model_hash] = settings_hash if unload_distorch_model: + logger.mgpu_mm_log("[FLAG_TRIGGER] unload_distorch_model=True, triggering full system cleanup") force_full_system_cleanup(reason="policy_every_load", force=True) return out diff --git a/memory-bank/phase3_bug_fix.md b/memory-bank/phase3_bug_fix.md new file mode 100644 index 0000000..41981cf --- /dev/null +++ b/memory-bank/phase3_bug_fix.md @@ -0,0 +1,261 @@ +# Phase 3 Bug Fix: Path Mismatch in Flag Storage/Retrieval + +**Date:** 2025-09-29 +**Status:** ✅ FIXED + Comprehensive Diagnostics Added +**Root Cause:** Object path mismatch between flag SET and flag READ operations + +## The Bug + +### What Was Wrong + +**Flag SETTING (distorch_2.py - 3 locations):** +```python +# BUG: Stored flag on INNER MODEL +out[0].model._mgpu_unload_distorch_model = unload_distorch_model +``` + +**Flag READING (model_management_mgpu.py):** +```python +# BUG: Read from WRONG LOCATION +mp = lm.model # This is the ModelPatcher +unload_distorch_model = getattr(mp.model, '_mgpu_unload_distorch_model', False) +# ^^^^^^^^ Reading from mp.model (inner model) +``` + +**Object Hierarchy:** +``` +LoadedModel (lm) + └─ ModelPatcher (lm.model / mp) + └─ Actual Model (mp.model / inner model) +``` + +**The Mismatch:** +- **SET:** Flag stored on `ModelPatcher` object (`out[0]` is the ModelPatcher) +- **READ:** Flag read from `ModelPatcher.model` (the inner model) +- **Result:** Flag check always returns `False` (default) → all models categorized as "keep loaded" + +### Why Selective Unload Appeared to Work But Didn't + +**Misleading Log Output:** +``` +[UNLOAD_DEBUG] Adding to kept_models: AutoencodingEngine +[UNLOAD_DEBUG] Adding to kept_models: FluxClipModel_ +[UNLOAD_DEBUG] Final counts - kept_models: 2, models_to_unload: 1 +``` + +This logging showed categorization happening, but the categorization was WRONG because: +1. Flag check failed for ALL models (path mismatch) +2. All models defaulted to `False` (keep loaded) +3. Only models with explicit `True` flag should unload +4. But flag was never found, so nothing had `True` → everything kept + +**Evidence from user's previous successful commit:** +The user mentioned selective retention "worked in more than one of the commits of this branch" - likely an earlier version where flag storage/retrieval paths were aligned. + +## The Fix + +### Primary Fix: Path Alignment + +**NEW: Store and Read from Same Location** +```python +# SET (distorch_2.py): +mp = out[0] # ModelPatcher +mp._mgpu_unload_distorch_model = unload_distorch_model + +# READ (model_management_mgpu.py): +mp = lm.model # ModelPatcher +flag_on_mp = getattr(mp, '_mgpu_unload_distorch_model', None) +``` + +**Backwards Compatibility During Transition:** +```python +# Also set on inner model for any old workflows +if inner_model: + inner_model._mgpu_unload_distorch_model = unload_distorch_model + +# Read from both locations, prefer ModelPatcher +flag_on_mp = getattr(mp, '_mgpu_unload_distorch_model', None) +flag_on_inner = getattr(mp.model, '_mgpu_unload_distorch_model', None) + +if flag_on_mp is not None: + unload_distorch_model = flag_on_mp # Use MP location (new) +elif flag_on_inner is not None: + unload_distorch_model = flag_on_inner # Fall back to inner (old) +else: + unload_distorch_model = False # Default: keep loaded +``` + +### Comprehensive Diagnostics Added + +**Object Identity Tracking:** +```python +[OBJECT_CHAIN_SET] ModelPatcher: mp_id=0x7f8a4c0, inner_model_id=0x7f8a5d0, inner_model_type=FluxClipModel_ +[FLAG_SET_LOCATION] Set on ModelPatcher (mp_id=0x7f8a4c0): mp._mgpu_unload_distorch_model = False +[FLAG_SET_COMPAT] Also set on inner model (inner_model_id=0x7f8a5d0) for compatibility + +[OBJECT_CHAIN_READ] Model 0: lm_id=0x7f8a600, mp_id=0x7f8a4c0, inner_model_id=0x7f8a5d0, inner_model_type=FluxClipModel_ +[FLAG_CHECK] Model 0 (FluxClipModel_): flag_on_mp=False, flag_on_inner=False +[FLAG_SOURCE] Using flag from ModelPatcher (mp_id=0x7f8a4c0) +[DECISION] Model 0 (FluxClipModel_): unload_distorch_model=False +[CATEGORIZE] Model 0 (FluxClipModel_) → kept_models +``` + +This reveals: +- **Object identity match:** Same mp_id at SET and READ (0x7f8a4c0) +- **Flag location:** Now reading from correct location +- **Decision trace:** Complete path from flag check to categorization +- **Remaining models:** What's left after selective unload + +## Expected Behavior After Fix + +### Scenario 1: Mixed keep_loaded Settings + +**Workflow:** +- UNET: `keep_loaded=False` → should unload +- VAE: `keep_loaded=True` → should retain +- CLIP: `keep_loaded=True` → should retain + +**Expected Log Output:** +``` +[OBJECT_CHAIN_SET] UNET mp_id=0xAAA, unload_distorch_model=True +[OBJECT_CHAIN_SET] VAE mp_id=0xBBB, unload_distorch_model=False +[OBJECT_CHAIN_SET] CLIP mp_id=0xCCC, unload_distorch_model=False + +[UNLOAD_START] initial model count: 3 + +[OBJECT_CHAIN_READ] Model 0: mp_id=0xAAA (UNET) +[FLAG_CHECK] flag_on_mp=True +[CATEGORIZE] → models_to_unload + +[OBJECT_CHAIN_READ] Model 1: mp_id=0xBBB (VAE) +[FLAG_CHECK] flag_on_mp=False +[CATEGORIZE] → kept_models + +[OBJECT_CHAIN_READ] Model 2: mp_id=0xCCC (CLIP) +[FLAG_CHECK] flag_on_mp=False +[CATEGORIZE] → kept_models + +[SELECTIVE_UNLOAD] retaining 2, unloading 1 +[UNLOAD_EXECUTE] Unloading UNET +[SELECTIVE_COMPLETE] new count: 2 + +[REMAINING_MODEL] 0: VAE (mp_id=0xBBB) +[REMAINING_MODEL] 1: CLIP (mp_id=0xCCC) +``` + +### Scenario 2: All keep_loaded=False + +**Expected:** +- All models unloaded +- CPU memory fully reclaimed +- No retained models + +### Scenario 3: All keep_loaded=True + +**Expected:** +- Delegation to original `unload_all_models()` +- Standard ComfyUI behavior +- All models handled by Comfy's normal flow + +## Files Modified + +### 1. model_management_mgpu.py +**Changes:** +- Fixed flag reading path (ModelPatcher vs inner model) +- Added object identity logging at READ time +- Added flag source detection (MP vs inner vs not found) +- Added decision trace logging +- Added remaining models logging post-unload + +### 2. distorch_2.py (3 override classes) +**Changes:** +- Fixed flag storage path (ModelPatcher vs inner model) +- Added object identity logging at SET time +- Added dual-location flag setting for compatibility +- All three overrides updated identically: + - `override_class_with_distorch_safetensor_v2` + - `override_class_with_distorch_safetensor_v2_clip` + - `override_class_with_distorch_safetensor_v2_clip_no_device` + +## Testing Plan + +### Minimal Test Workflow + +**Requirements:** +- 1 UNET (DisTorch2) with `keep_loaded=False` +- 1 VAE (any loader) +- 1 CLIP (DisTorch2) with `keep_loaded=True` + +**Expected Result:** +1. UNET loads → flag set to True → triggers cleanup request +2. Workflow executes +3. Post-execution cleanup: + - UNET unloaded (flag=True) + - VAE retained (no flag) + - CLIP retained (flag=False) +4. CPU memory reclaimed (UNET's CPU portions freed) +5. Detection shows 2 models remaining + +### What to Look For in Logs + +**Success Indicators:** +- `[FLAG_CHECK]` shows flags correctly detected +- `[CATEGORIZE]` separates models correctly +- `[SELECTIVE_COMPLETE]` shows expected count +- `[REMAINING_MODEL]` lists only kept models +- Detection after unload shows correct count + +**Failure Indicators:** +- Object IDs don't match between SET and READ +- Flags not found (all default to False) +- Wrong models categorized +- Retained models disappear after unload +- Detection shows 0 models when should show N + +## Why This Fix Should Work + +**Root Cause Eliminated:** +- Flag storage and retrieval now use same object path +- Object identity logging proves we're checking the same instance +- Backwards compatibility handles transition period + +**Architecture Preserved:** +- Still uses ComfyUI's deferred flag mechanism +- Still runs post-execution (timing is correct) +- Still selective (keeps what should be kept) +- Still comprehensive (cleans what should be cleaned) + +**Diagnostics Enable Debugging:** +- If it still fails, logs will show exactly where/why +- Object IDs prove identity across operations +- Flag source shows which location succeeded +- Decision trace shows categorization logic + +## Next Steps + +1. **Test with simple workflow** - verify basic selective unload works +2. **Monitor logs** - check object IDs match SET→READ +3. **Validate CPU memory** - confirm reclamation after unload +4. **Test edge cases:** + - All keep_loaded=False + - All keep_loaded=True + - Mixed settings +5. **If still failing** - logs will reveal the actual issue + +## Historical Context + +**Previous Failed Approaches:** +- Phase 1: Missing executor reset (failed - CPU memory not reclaimed) +- Phase 2: Implementation fixes (failed - resets occurring but memory rising) +- Phase 3 Initial: Aggressive reclamation (failed - OOM persisted) + +**This Fix Different Because:** +- Addresses actual code bug (path mismatch) +- Not architectural change (just alignment) +- Preserves working Phase 3 design +- Adds proof via diagnostics + +**User's Historical Note:** +"We had this selectiveness working in more than one of the commits of this branch so it is more rediscovering it." + +This suggests an earlier version had correct paths - this fix rediscovers that working pattern. diff --git a/model_management_mgpu.py b/model_management_mgpu.py index 5d40ae3..fb63152 100644 --- a/model_management_mgpu.py +++ b/model_management_mgpu.py @@ -216,6 +216,45 @@ def force_full_system_cleanup(reason="manual", force=True): logger.mgpu_mm_log(summary) return summary +# ========================================================================================== +# Core Patching: soft_empty_cache (Instrumentation) +# ========================================================================================== + +if not hasattr(mm.soft_empty_cache, '_mgpu_instrumented'): + logger.info("[MultiGPU Core Patching] Instrumenting mm.soft_empty_cache for diagnostics") + + _mgpu_original_soft_empty_cache = mm.soft_empty_cache + + def _mgpu_instrumented_soft_empty_cache(force=False): + """Instrumented soft_empty_cache to track what it does to mm.current_loaded_models""" + models_before = len(mm.current_loaded_models) + logger.mgpu_mm_log(f"[SOFT_EMPTY_ENTRY] Original mm.soft_empty_cache called, models_before={models_before}, force={force}") + + # Log the models present before calling original + for i, lm in enumerate(mm.current_loaded_models): + mp = lm.model + inner_model = getattr(mp, 'model', None) + model_name = type(inner_model).__name__ if inner_model else "None" + logger.mgpu_mm_log(f"[SOFT_EMPTY_ENTRY] Model {i} before: {model_name} (lm_id=0x{id(lm):x})") + + # Call original + result = _mgpu_original_soft_empty_cache(force) + + # Check what happened to models + models_after = len(mm.current_loaded_models) + logger.mgpu_mm_log(f"[SOFT_EMPTY_EXIT] Original mm.soft_empty_cache returned, models_after={models_after} (delta={models_after - models_before})") + + if models_after != models_before: + logger.mgpu_mm_log(f"[SOFT_EMPTY_CULPRIT] Original mm.soft_empty_cache MODIFIED mm.current_loaded_models: {models_before} → {models_after}") + + return result + + mm.soft_empty_cache = _mgpu_instrumented_soft_empty_cache + mm.soft_empty_cache._mgpu_instrumented = True + logger.info("[MultiGPU Core Patching] mm.soft_empty_cache instrumented successfully") +else: + logger.debug("[MultiGPU Core Patching] mm.soft_empty_cache already instrumented - skipping") + # ========================================================================================== # Core Patching: unload_all_models # ========================================================================================== @@ -227,11 +266,10 @@ if not hasattr(mm.unload_all_models, '_mgpu_eject_distorch_patched'): def _mgpu_patched_unload_all_models(): """ - Patched mm.unload_all_models that checks to see if the . - All other models (including DisTorch models without the flag) unload normally. + Patched mm.unload_all_models with comprehensive diagnostics and fixed path alignment. """ - 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"[UNLOAD_START] Patched unload_all_models called - initial model count: {len(mm.current_loaded_models)}") # Direct approach: iterate through loaded models and selectively unload models_to_unload = [] @@ -239,48 +277,93 @@ if not hasattr(mm.unload_all_models, '_mgpu_eject_distorch_patched'): for i, lm in enumerate(mm.current_loaded_models): mp = lm.model # weakref call to ModelPatcher - - unload_distorch_model = getattr(mp.model, '_mgpu_unload_distorch_model', False) - model_name = type(getattr(mp, 'model', mp)).__name__ - logger.mgpu_mm_log(f"[Phase 3 Debug] Model {i}: {model_name}, unload_distorch_model={unload_distorch_model}") - # Retain models that either: - # 1. Are non-DisTorch models (missing _mgpu_keep_loaded attribute) - # 2. Are DisTorch models with keep_loaded=True - + # DIAGNOSTIC: Log full object chain + lm_id = id(lm) + mp_id = id(mp) + inner_model = getattr(mp, 'model', None) + inner_model_id = id(inner_model) if inner_model else None + inner_model_name = type(inner_model).__name__ if inner_model else "None" + + # Format inner_model_id properly for f-string + inner_id_str = f"0x{inner_model_id:x}" if inner_model_id is not None else "None" + + logger.mgpu_mm_log(f"[OBJECT_CHAIN_READ] Model {i}: lm_id=0x{lm_id:x}, mp_id=0x{mp_id:x}, inner_model_id={inner_id_str}, inner_model_type={inner_model_name}") + + # FIX: Check flag on ModelPatcher (where it was set), not on inner model + # OLD BUG: unload_distorch_model = getattr(mp.model, '_mgpu_unload_distorch_model', False) + # NEW FIX: Check both locations to see which one has the flag + flag_on_mp = getattr(mp, '_mgpu_unload_distorch_model', None) + flag_on_inner = getattr(mp.model, '_mgpu_unload_distorch_model', None) if inner_model else None + + logger.mgpu_mm_log(f"[FLAG_CHECK] Model {i} ({inner_model_name}): flag_on_mp={flag_on_mp}, flag_on_inner={flag_on_inner}") + + # Use whichever location has the flag (for backwards compatibility during transition) + if flag_on_mp is not None: + unload_distorch_model = flag_on_mp + logger.mgpu_mm_log(f"[FLAG_SOURCE] Using flag from ModelPatcher (mp_id=0x{mp_id:x})") + elif flag_on_inner is not None: + unload_distorch_model = flag_on_inner + logger.mgpu_mm_log(f"[FLAG_SOURCE] Using flag from inner model (inner_model_id={inner_id_str})") + else: + unload_distorch_model = False + logger.mgpu_mm_log(f"[FLAG_SOURCE] No flag found - defaulting to False (keep loaded)") + + logger.mgpu_mm_log(f"[DECISION] Model {i} ({inner_model_name}): unload_distorch_model={unload_distorch_model}") + if unload_distorch_model: models_to_unload.append(lm) + logger.mgpu_mm_log(f"[CATEGORIZE] Model {i} ({inner_model_name}) → models_to_unload") else: kept_models.append(lm) - logger.mgpu_mm_log(f"[UNLOAD_DEBUG] Adding to kept_models: {model_name}") + logger.mgpu_mm_log(f"[CATEGORIZE] Model {i} ({inner_model_name}) → kept_models") # After the kept_models/models_to_unload evaluation + logger.mgpu_mm_log(f"[CATEGORIZE_SUMMARY] kept_models: {len(kept_models)}, models_to_unload: {len(models_to_unload)}, total: {len(mm.current_loaded_models)}") + if len(kept_models) == len(mm.current_loaded_models): # All models are meant to be kept - no DisTorch selective unloading needed - logger.mgpu_mm_log("[Phase 2 Debug] All models flagged to be kept - using standard unload_all_models") + logger.mgpu_mm_log("[DELEGATION] All models flagged to be kept - delegating to standard unload_all_models") _mgpu_original_unload_all_models() return - - 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)") + logger.mgpu_mm_log(f"[SELECTIVE_UNLOAD] Proceeding with selective unload: retaining {len(kept_models)}, unloading {len(models_to_unload)}") - # Unload models that don't have keep_loaded flag + # Unload models flagged for unload for lm in models_to_unload: try: + model_name = type(lm.model.model).__name__ if lm.model and hasattr(lm.model, 'model') else 'Unknown' + logger.mgpu_mm_log(f"[UNLOAD_EXECUTE] Unloading model: {model_name} (lm_id=0x{id(lm):x})") lm.model_unload(unpatch_weights=True) - 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}") + logger.warning(f"[UNLOAD_ERROR] Error unloading model: {e}") + + # WEAKREF TRACKING: Attach weakref callbacks to prove if kept models are GC'd + def model_deleted_callback(ref, model_name, model_id): + logger.mgpu_mm_log(f"[WEAKREF_DELETED] Kept model GARBAGE COLLECTED: {model_name} (id=0x{model_id:x})") + + for i, lm in enumerate(kept_models): + mp = lm.model + inner_model = getattr(mp, 'model', None) + model_name = type(inner_model).__name__ if inner_model else 'Unknown' + model_id = id(lm) + weakref.ref(lm, lambda ref, name=model_name, mid=model_id: model_deleted_callback(ref, name, mid)) + logger.mgpu_mm_log(f"[WEAKREF_ATTACHED] Tracking kept model {i}: {model_name} (lm_id=0x{model_id:x}, mp_id=0x{id(mp):x})") # 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)}") - logger.mgpu_mm_log(f"Successfully retained {len(kept_models)} model(s) during unload") + logger.mgpu_mm_log(f"[SELECTIVE_COMPLETE] Updated mm.current_loaded_models, new count: {len(mm.current_loaded_models)}") + logger.mgpu_mm_log(f"[SELECTIVE_COMPLETE] mm.current_loaded_models id: 0x{id(mm.current_loaded_models):x}") + + # DIAGNOSTIC: Log what's remaining + for i, lm in enumerate(mm.current_loaded_models): + mp = lm.model + inner_model = getattr(mp, 'model', None) + model_name = type(inner_model).__name__ if inner_model else "None" + logger.mgpu_mm_log(f"[REMAINING_MODEL] {i}: {model_name} (lm_id=0x{id(lm):x}, mp_id=0x{id(mp):x})") else: - logger.mgpu_mm_log("No models with keep_loaded=True found - delegating to original unload_all_models") + logger.mgpu_mm_log("[DELEGATION] No models with keep_loaded=True found - delegating to original unload_all_models") _mgpu_original_unload_all_models() mm.unload_all_models = _mgpu_patched_unload_all_models