fix(distorch): Add GC anchor protection for selective model retention

Problem: Models correctly categorized as "keep loaded" during selective
unload were disappearing before the next cleanup cycle. After reassigning
mm.current_loaded_models = kept_models, Python's garbage collector would
clear the models because the list was their only remaining strong reference.

Solution: Implement GC anchor system using a global set to hold strong
references to ModelPatcher objects that must survive cleanup cycles.

Changes:
- Add _MGPU_RETENTION_ANCHORS global set and helper functions
- Add early delegation check: if no DisTorch models want unload, clear
  anchors and delegate to original unload_all_models
- Add retention anchor when categorizing kept models
- Clear anchors before delegating to allow normal cleanup

Result: Self-contained, reversible protection mechanism. Models with
keep_loaded=True survive automatic cleanup but can be cleared with
explicit "Clear All Models" button. Tested on both keep_loaded=True
and keep_loaded=False scenarios.

Refs: memory-bank/distorch_selective_unload_solution.md
This commit is contained in:
John Pollock
2025-09-29 16:01:24 -05:00
parent bde51c6236
commit 07b429f3f9
+35
View File
@@ -21,6 +21,26 @@ from collections import defaultdict
logger = logging.getLogger("MultiGPU")
# ==========================================================================================
# GC Anchor System for Model Retention
# ==========================================================================================
# Global anchor set to prevent GC of models during selective unload
_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 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)
# ==========================================================================================
@@ -271,6 +291,20 @@ if not hasattr(mm.unload_all_models, '_mgpu_eject_distorch_patched'):
logger.mgpu_mm_log(f"[UNLOAD_START] Patched unload_all_models called - initial model count: {len(mm.current_loaded_models)}")
# Check if there are any DisTorch models that want to be unloaded
has_distorch_to_unload = any(
(hasattr(lm.model, '_mgpu_unload_distorch_model') and lm.model._mgpu_unload_distorch_model) or
(hasattr(getattr(lm.model, 'model', None), '_mgpu_unload_distorch_model') and lm.model.model._mgpu_unload_distorch_model)
for lm in mm.current_loaded_models
if lm.model is not None
)
if not has_distorch_to_unload:
logger.mgpu_mm_log("No DisTorch models requesting unload - clearing anchors and delegating to original unload_all_models")
clear_all_retention_anchors(reason="no_selective_unload_needed")
_mgpu_original_unload_all_models()
return
# Direct approach: iterate through loaded models and selectively unload
models_to_unload = []
kept_models = []
@@ -316,6 +350,7 @@ if not hasattr(mm.unload_all_models, '_mgpu_eject_distorch_patched'):
logger.mgpu_mm_log(f"[CATEGORIZE] Model {i} ({inner_model_name}) → models_to_unload")
else:
kept_models.append(lm)
add_retention_anchor(mp, "keep_loaded_protection")
logger.mgpu_mm_log(f"[CATEGORIZE] Model {i} ({inner_model_name}) → kept_models")
# After the kept_models/models_to_unload evaluation