Paranoia before cleanup

This commit is contained in:
John Pollock
2025-10-04 12:32:37 -05:00
parent e3750fd737
commit e6d19951d7
3 changed files with 118 additions and 208 deletions
+67 -22
View File
@@ -74,33 +74,78 @@ def register_patched_safetensor_modelpatcher():
original_loaded_model_memory_required = mm.LoadedModel.model_memory_required
def patched_loaded_model_memory_required(self, device):
"""Drive unload behavior purely by unload_distorch_model flag"""
"""Truth table for memory reporting:
eject_models=0, is_distorch=0: return original
eject_models=0, is_distorch=1: return original - virtual_vram_gb_bytes
eject_models=1, is_distorch=0: mutually exclusive (shouldn't occur)
eject_models=1, is_distorch=1: return MAX memory to force eviction"""
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}")
# 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')
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}")
logger.mgpu_mm_log(f"[MEM_REPORT][{model_name}] Memory assessment requested for model on device: {device}")
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
# Not a DisTorch model - use original behavior
logger.mgpu_mm_log("[IS_DISTORCH_MODEL] Non-DisTorch model - Using original Comfy memory calculation")
# GET ORIGINAL MEMORY REQUIREMENT
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
# CHECK FOR EJECT_MODELS PROPERTY
has_eject_models = hasattr(getattr(getattr(self, 'model', None), 'model', None), '_mgpu_eject_models')
logger.mgpu_mm_log(f"[MEM_REPORT][{model_name}] Original needs: {original_gb:.2f}GB, has_eject_models={has_eject_models}")
# CHECK IF DISTORCH MODEL WITH VIRTUAL VRAM PROPERTY
is_distorch_model = hasattr(getattr(getattr(self, 'model', None), 'model', None), '_mgpu_virtual_vram_gb')
# TRUTH TABLE APPLICATION
if has_eject_models:
if not is_distorch_model:
logger.mgpu_mm_log(f"[MEM_REPORT][{model_name}] ERROR: eject_models=1 but not DisTorch (mutually exclusive)")
# eject_models=1, is_distorch=1: RETURN MAX MEMORY TO FORCE EVICTION
logger.mgpu_mm_log(f"[MEM_REPORT][{model_name}] eject_models=1, is_distorch={is_distorch_model} → FORCING EVICTION WITH MAX MEMORY")
# DISABLED: Manual ejection should happen automatically when MAX memory is returned
DISABLE_MANUAL_EJECTION = True # TODO: Remove this once auto-eviction confirmed
if not DISABLE_MANUAL_EJECTION:
logger.mgpu_mm_log(f"======= DIRECT MODEL EJECTION START[{model_name}] =======")
logger.mgpu_mm_log(f"[DIRECT_EJECTION][{model_name}] Current loaded models count: {len(mm.current_loaded_models)}")
# DIRECTLY UNLOAD ALL MODELS
models_unloaded = []
for i, lm in enumerate(mm.current_loaded_models):
model_name_to_eject = type(getattr(lm.model, 'model', lm.model)).__name__ if lm.model else 'Unknown'
logger.mgpu_mm_log(f"[DIRECT_EJECTION][{model_name}] UNLOADING MODEL {i+1}/{len(mm.current_loaded_models)}: {model_name_to_eject}")
try:
lm.model_unload(unpatch_weights=True)
models_unloaded.append(model_name_to_eject)
logger.mgpu_mm_log(f"[DIRECT_EJECTION][{model_name}] SUCCESSFULLY UNLOADED: {model_name_to_eject}")
except Exception as e:
logger.mgpu_mm_log(f"[DIRECT_EJECTION][{model_name}] ERROR unloading {model_name_to_eject}: {e}")
mm.current_loaded_models = []
logger.mgpu_mm_log(f"[DIRECT_EJECTION][{model_name}] Models unloaded: {models_unloaded}")
logger.mgpu_mm_log(f"======= DIRECT MODEL EJECTION COMPLETE[{model_name}] =======")
multigpu_memory_log("eject_models_post", "complete")
# RETURN MAX MEMORY - Should trigger auto-eviction by Comfy Core
total_device_memory = mm.get_total_memory(device)
max_gb = total_device_memory / (1024**3)
logger.mgpu_mm_log(f"[MEM_REPORT][{model_name}] Returning MAX memory ({max_gb:.2f}GB) for auto-eviction by Comfy Core")
return total_device_memory
elif is_distorch_model:
# eject_models=0, is_distorch=1: SUBTRACT VIRTUAL VRAM FROM ORIGINAL
virtual_vram_gb = getattr(getattr(self, 'model', None), 'model', None)._mgpu_virtual_vram_gb
virtual_vram_bytes = virtual_vram_gb * (1024**3)
adjusted_result = max(0, original_result - virtual_vram_bytes)
adjusted_gb = adjusted_result / (1024**3) if adjusted_result else 0
logger.mgpu_mm_log(f"[MEM_REPORT][{model_name}] eject_models=0, is_distorch=1 → adjusted {original_gb:.2f}GB - {virtual_vram_gb:.2f}GB = {adjusted_gb:.2f}GB (DisTorch allocation)")
multigpu_memory_log("distorch_allocation", "reported")
return adjusted_result
else:
# eject_models=0, is_distorch=0: RETURN ORIGINAL
logger.mgpu_mm_log(f"[MEM_REPORT][{model_name}] eject_models=0, is_distorch=0 → returning original {original_gb:.2f}GB")
multigpu_memory_log("keep_loaded_memory_check", "end")
return original_result
mm.LoadedModel.model_memory_required = patched_loaded_model_memory_required
-153
View File
@@ -11,35 +11,12 @@ 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")
# ==========================================================================================
# 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)
@@ -232,133 +209,3 @@ def force_full_system_cleanup(reason="manual", force=True):
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
# ==========================================================================================
# Core Patching: unload_all_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 with selective ejection support and comprehensive diagnostics."""
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 = []
for i, lm in enumerate(mm.current_loaded_models):
mp = lm.model # weakref call to ModelPatcher
# 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)
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
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("[DELEGATION] All models flagged to be kept - delegating to standard unload_all_models")
_mgpu_original_unload_all_models()
return
if kept_models:
logger.mgpu_mm_log(f"[SELECTIVE_UNLOAD] Proceeding with selective unload: retaining {len(kept_models)}, unloading {len(models_to_unload)}")
# 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)
except Exception as 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"[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("[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
mm.unload_all_models._mgpu_eject_distorch_patched = True
logger.info("[MultiGPU Core Patching] mm.unload_all_models patched successfully")
else:
logger.debug("[MultiGPU Core Patching] mm.unload_all_models already patched - skipping")
+51 -33
View File
@@ -37,7 +37,7 @@ def _create_distorch_safetensor_v2_override(cls, device_param_name, device_sette
inputs["optional"]["virtual_vram_gb"] = ("FLOAT", {"default": 4.0, "min": 0.0, "max": 128.0, "step": 0.1})
inputs["optional"]["donor_device"] = (devices, {"default": "cpu"})
inputs["optional"]["expert_mode_allocations"] = ("STRING", {"multiline": False, "default": ""})
inputs["optional"]["keep_loaded"] = ("BOOLEAN", {"default": True})
inputs["optional"]["eject_models"] = ("BOOLEAN", {"default": True})
return inputs
CATEGORY = "multigpu/distorch_2"
@@ -45,10 +45,10 @@ def _create_distorch_safetensor_v2_override(cls, device_param_name, device_sette
TITLE = f"{cls.TITLE if hasattr(cls, 'TITLE') else cls.__name__} (DisTorch2)"
@classmethod
def IS_CHANGED(s, *args, virtual_vram_gb=4.0, donor_device="cpu",
expert_mode_allocations="", keep_loaded=True, **kwargs):
def IS_CHANGED(s, *args, virtual_vram_gb=4.0, donor_device="cpu",
expert_mode_allocations="", eject_models=True, **kwargs):
device_value = kwargs.get(device_param_name)
settings_str = f"{device_value}{virtual_vram_gb}{donor_device}{expert_mode_allocations}{keep_loaded}"
settings_str = f"{device_value}{virtual_vram_gb}{donor_device}{expert_mode_allocations}{eject_models}"
current_hash = hashlib.sha256(settings_str.encode()).hexdigest()
if not hasattr(cls, '_last_hash'):
@@ -60,19 +60,40 @@ def _create_distorch_safetensor_v2_override(cls, device_param_name, device_sette
return current_hash
def override(self, *args, virtual_vram_gb=4.0, donor_device="cpu",
expert_mode_allocations="", keep_loaded=True, **kwargs):
expert_mode_allocations="", eject_models=True, **kwargs):
device_value = kwargs.get(device_param_name)
unload_distorch_model = not keep_loaded
import comfy.model_management as mm
if eject_models:
logger.mgpu_mm_log(f"[EJECT_MODELS_SETUP] eject_models=True - marking all loaded models for eviction, device target: {device_value}")
ejection_count = 0
for i, lm in enumerate(mm.current_loaded_models):
# Set _mgpu_unload_distorch_model=True on all models to force Comfy Core eviction
model_name = type(getattr(lm.model, 'model', lm.model)).__name__ if lm.model else 'Unknown'
if hasattr(lm.model, 'model') and lm.model.model is not None:
lm.model.model._mgpu_unload_distorch_model = True
logger.mgpu_mm_log(f"[EJECT_MARKED] Model {i}: {model_name} (id=0x{id(lm):x}) → marked for eviction")
ejection_count += 1
elif lm.model is not None:
lm.model._mgpu_unload_distorch_model = True
logger.mgpu_mm_log(f"[EJECT_MARKED] Model {i}: {model_name} (direct patcher) → marked for eviction")
ejection_count += 1
logger.mgpu_mm_log(f"[EJECT_MODELS_SETUP_COMPLETE] Marked {ejection_count} models for Comfy Core eviction during load_models_gpu")
else:
logger.mgpu_mm_log(f"[EJECT_MODELS_SETUP] eject_models=False - loading without eviction")
if device_value is not None:
device_setter_func(device_value)
# Strip MultiGPU-specific parameters before calling original function
clean_kwargs = {k: v for k, v in kwargs.items()
if k not in [device_param_name, 'virtual_vram_gb',
'donor_device', 'expert_mode_allocations',
'keep_loaded']}
# Strip MultiGPU-specific parameters before calling original function (REMOVE eject_models, keep_loaded and virtual_vram_gb since we handle them above)
clean_kwargs = {k: v for k, v in kwargs.items()
if k not in [device_param_name, 'virtual_vram_gb',
'donor_device', 'expert_mode_allocations',
'eject_models']}
if apply_device_kwarg_workaround:
clean_kwargs['device'] = 'default'
@@ -106,7 +127,7 @@ def _create_distorch_safetensor_v2_override(cls, device_param_name, device_sette
logger.debug(f"[MultiGPU DisTorch V2] Stored allocation for model {model_hash[:8]}: {full_allocation}")
logger.info(f"[MultiGPU DisTorch V2] Full allocation string: {full_allocation}")
logger.mgpu_mm_log(f"[FLAG_SET_START] Setting '_mgpu_unload_distorch_model' to: {unload_distorch_model} (keep_loaded={keep_loaded})")
logger.mgpu_mm_log(f"[MODEL_SETUP] Setting DisTorch model properties: virtual_vram_gb={virtual_vram_gb}")
if hasattr(out[0], 'model'):
mp = out[0]
@@ -115,16 +136,19 @@ def _create_distorch_safetensor_v2_override(cls, device_param_name, device_sette
inner_model_id = id(inner_model) if inner_model else None
inner_model_name = type(inner_model).__name__ if inner_model else "None"
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}")
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}")
# SET VIRTUAL VRAM PROPERTY FOR MEMORY CALCULATION
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")
inner_model._mgpu_virtual_vram_gb = virtual_vram_gb
logger.mgpu_mm_log(f"[VIRTUAL_VRAM_SET] Set _mgpu_virtual_vram_gb={virtual_vram_gb}GB on inner model (id=0x{inner_model_id:x}) for memory assessment")
# SET EJECT MODELS PROPERTY IF ENABLED
if eject_models and inner_model:
inner_model._mgpu_eject_models = True
logger.mgpu_mm_log(f"[EJECT_FLAG_SET] Set _mgpu_eject_models=True on inner model (id=0x{inner_model_id:x}) - will trigger ejection during load_models_gpu")
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
mp = out[0].patcher
mp_id = id(mp)
@@ -132,19 +156,13 @@ def _create_distorch_safetensor_v2_override(cls, device_param_name, device_sette
inner_model_id = id(inner_model) if inner_model else None
inner_model_name = type(inner_model).__name__ if inner_model else "None"
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}")
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}")
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)
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}")
# SET VIRTUAL VRAM PROPERTY FOR MEMORY CALCULATION
if inner_model:
inner_model._mgpu_virtual_vram_gb = virtual_vram_gb
logger.mgpu_mm_log(f"[VIRTUAL_VRAM_SET] Set _mgpu_virtual_vram_gb={virtual_vram_gb}GB on inner model (id=0x{inner_model_id:x}) for memory assessment")
return out