extensive clean-up, WIP
This commit is contained in:
+8
-382
@@ -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,
|
||||
|
||||
+1
-9
@@ -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()
|
||||
|
||||
+85
-52
@@ -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:
|
||||
|
||||
@@ -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.
|
||||
@@ -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.
|
||||
|
||||
+26
-171
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user