commiting this steaming pile of hot garbage for future dissection to see if I want any organs from this terminally ill branch

This commit is contained in:
John Pollock
2025-09-25 14:36:13 -05:00
parent bd672479fa
commit fda5d6ed00
6 changed files with 4446 additions and 161 deletions
+22 -18
View File
@@ -313,6 +313,7 @@ from .nodes import (
HyVideoModelLoader,
HyVideoVAELoader,
DownloadAndLoadHyVideoTextEncoder,
UNetLoaderLP,
FullCleanupMultiGPU,
)
@@ -376,13 +377,28 @@ def soft_empty_cache_distorch2_patched(force=False):
is_distorch_active = False
# Detect DisTorch2-managed models
for lm in mm.current_loaded_models:
logger.mgpu_mm_log(f"[DETECT_DEBUG] Checking DisTorch2 active status - loaded models: {len(mm.current_loaded_models)}, store entries: {len(safetensor_allocation_store)}")
for i, lm in enumerate(mm.current_loaded_models):
mp = lm.model # weakref call to ModelPatcher
if mp is not None:
model_hash = create_safetensor_model_hash(mp, "cache_patch_check")
if model_hash in safetensor_allocation_store and safetensor_allocation_store.get(model_hash):
is_distorch_active = True
break
try:
model_hash = create_safetensor_model_hash(mp, "cache_patch_check")
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)
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}")
if in_store and alloc_value:
is_distorch_active = True
logger.mgpu_mm_log(f"[DETECT_DEBUG] DisTorch2 ACTIVE detected on model: {model_name}")
break
except Exception as e:
logger.mgpu_mm_log(f"[DETECT_DEBUG] Model {i}: Error during detection - {e}")
logger.mgpu_mm_log(f"[DETECT_DEBUG] Final DisTorch2 active status: {is_distorch_active}")
# Phase 2: adaptive CPU memory management
check_cpu_memory_threshold()
@@ -645,19 +661,6 @@ if hasattr(mm, 'load_models_gpu') and not hasattr(mm.load_models_gpu, "_distorch
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")
# Cleanup policy triggers (flags-only, Manager semantics)
if MGPU_CLEANUP_POLICY in ("threshold", "every_load+threshold", "threshold+every_load"):
try:
check_cpu_memory_threshold(threshold_percent=MGPU_CPU_RESET_THRESHOLD * 100.0)
except Exception:
pass
if MGPU_CLEANUP_POLICY in ("every_load", "every_load+threshold", "threshold+every_load"):
try:
# flags-only; prompt worker performs unload/reset/gc
force_full_system_cleanup(reason="policy_every_load", force=False)
except Exception:
pass
return result
# Mark and apply the patch
@@ -681,6 +684,7 @@ NODE_CLASS_MAPPINGS = {
"HunyuanVideoEmbeddingsAdapter": HunyuanVideoEmbeddingsAdapter,
"CheckpointLoaderAdvancedMultiGPU": CheckpointLoaderAdvancedMultiGPU,
"CheckpointLoaderAdvancedDisTorch2MultiGPU": CheckpointLoaderAdvancedDisTorch2MultiGPU,
"UNetLoaderLP": UNetLoaderLP,
}
# Standard MultiGPU nodes
+72 -111
View File
@@ -67,8 +67,11 @@ def register_patched_safetensor_modelpatcher():
multigpu_memory_log(f"safetensor:{debug_hash[:8]}", "pre-load")
allocations = safetensor_allocation_store.get(debug_hash)
if not hasattr(self.model, '_distorch_high_precision_loras') or not allocations:
# Set default precision flag before checking
if not hasattr(self.model, '_distorch_high_precision_loras'):
self.model._distorch_high_precision_loras = True
if not allocations:
result = original_partially_load(self, device_to, extra_memory, force_patch_weights)
multigpu_memory_log(f"safetensor:{debug_hash[:8]}", "post-load")
if hasattr(self, '_distorch_block_assignments'):
@@ -103,7 +106,7 @@ def register_patched_safetensor_modelpatcher():
device_assignments = analyze_safetensor_loading(self, allocations)
model_original_dtype = comfy.utils.weight_dtype(self.model.state_dict())
high_precision_loras = self.model._distorch_high_precision_loras
high_precision_loras = getattr(self.model, "_distorch_high_precision_loras", True)
loading = self._load_list()
loading.sort(reverse=True)
for module_size, module_name, module_object, params in loading:
@@ -813,7 +816,7 @@ def override_class_with_distorch_safetensor_v2(cls):
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"]["high_precision_loras"] = ("BOOLEAN", {"default": True})
inputs["optional"]["keep_loaded"] = ("BOOLEAN", {"default": True})
return inputs
CATEGORY = "multigpu/distorch_2"
@@ -822,13 +825,13 @@ def override_class_with_distorch_safetensor_v2(cls):
@classmethod
def IS_CHANGED(s, *args, compute_device=None, virtual_vram_gb=4.0,
donor_device="cpu", expert_mode_allocations="", high_precision_loras=True, **kwargs):
donor_device="cpu", expert_mode_allocations="", keep_loaded=True, **kwargs):
# Create a hash of our specific settings
settings_str = f"{compute_device}{virtual_vram_gb}{donor_device}{expert_mode_allocations}{high_precision_loras}"
settings_str = f"{compute_device}{virtual_vram_gb}{donor_device}{expert_mode_allocations}{keep_loaded}"
return hashlib.sha256(settings_str.encode()).hexdigest()
def override(self, *args, compute_device=None, virtual_vram_gb=4.0,
donor_device="cpu", expert_mode_allocations="", high_precision_loras=True, **kwargs):
donor_device="cpu", expert_mode_allocations="", keep_loaded=True, **kwargs):
from . import set_current_device
if compute_device is not None:
@@ -837,39 +840,7 @@ def override_class_with_distorch_safetensor_v2(cls):
# Register our patched ModelPatcher
register_patched_safetensor_modelpatcher()
# Call original function
fn = getattr(super(), cls.FUNCTION)
# --- Check if we need to unload the model due to settings change ---
# This logic is a bit redundant with IS_CHANGED, but provides clear logging
settings_str = f"{compute_device}{virtual_vram_gb}{donor_device}{expert_mode_allocations}"
settings_hash = hashlib.sha256(settings_str.encode()).hexdigest()
# Temporarily load to get hash without applying our patch
temp_out = fn(*args, **kwargs)
model_to_check = None
if hasattr(temp_out[0], 'model'):
model_to_check = temp_out[0]
elif hasattr(temp_out[0], 'patcher') and hasattr(temp_out[0].patcher, 'model'):
model_to_check = temp_out[0].patcher
if model_to_check:
model_hash = create_safetensor_model_hash(model_to_check, "override_check")
last_settings_hash = safetensor_settings_store.get(model_hash)
if last_settings_hash != settings_hash:
logger.mgpu_mm_log(f"Settings changed for model {model_hash[:8]}. Previous settings hash: {last_settings_hash}, New settings hash: {settings_hash}. Forcing reload.")
else:
logger.mgpu_mm_log(f"Settings unchanged for model {model_hash[:8]}. Using cached model.")
out = fn(*args, **kwargs)
# Store high_precision_loras in the model for later retrieval
if hasattr(out[0], 'model'):
out[0].model._distorch_high_precision_loras = high_precision_loras
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
out[0].patcher.model._distorch_high_precision_loras = high_precision_loras
# Build allocation string
vram_string = ""
if virtual_vram_gb > 0:
vram_string = f"{compute_device};{virtual_vram_gb};{donor_device}"
@@ -877,17 +848,35 @@ def override_class_with_distorch_safetensor_v2(cls):
vram_string = compute_device
full_allocation = f"{expert_mode_allocations}#{vram_string}" if expert_mode_allocations or vram_string else ""
fn = getattr(super(), cls.FUNCTION)
# Load the model and get hash, then store allocation for future runs
out = fn(*args, **kwargs)
model_to_check = None
if hasattr(out[0], 'model'):
model_to_check = out[0]
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
model_to_check = out[0].patcher
if model_to_check:
model_hash = create_safetensor_model_hash(model_to_check, "override_store")
settings_str = f"{compute_device}{virtual_vram_gb}{donor_device}{expert_mode_allocations}"
settings_hash = hashlib.sha256(settings_str.encode()).hexdigest()
# Store allocation for next run - this enables DisTorch for subsequent loads
safetensor_allocation_store[model_hash] = full_allocation
safetensor_settings_store[model_hash] = settings_hash
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}")
# Store keep_loaded in the model for later retrieval by unload_all_models patch
if hasattr(out[0], 'model'):
model_hash = create_safetensor_model_hash(out[0], "override")
safetensor_allocation_store[model_hash] = full_allocation
safetensor_settings_store[model_hash] = settings_hash
out[0].model._mgpu_keep_loaded = keep_loaded
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
model_hash = create_safetensor_model_hash(out[0].patcher, "override")
safetensor_allocation_store[model_hash] = full_allocation
safetensor_settings_store[model_hash] = settings_hash
out[0].patcher.model._mgpu_keep_loaded = keep_loaded
return out
@@ -909,7 +898,7 @@ def override_class_with_distorch_safetensor_v2_clip(cls):
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"]["high_precision_loras"] = ("BOOLEAN", {"default": True})
inputs["optional"]["keep_loaded"] = ("BOOLEAN", {"default": True})
return inputs
CATEGORY = "multigpu/distorch_2"
@@ -918,13 +907,13 @@ def override_class_with_distorch_safetensor_v2_clip(cls):
@classmethod
def IS_CHANGED(s, *args, device=None, virtual_vram_gb=4.0, # Changed from compute_device
donor_device="cpu", expert_mode_allocations="", high_precision_loras=True, **kwargs):
donor_device="cpu", expert_mode_allocations="", keep_loaded=True, **kwargs):
# Create a hash of our specific settings
settings_str = f"{device}{virtual_vram_gb}{donor_device}{expert_mode_allocations}{high_precision_loras}" # Changed from compute_device
settings_str = f"{device}{virtual_vram_gb}{donor_device}{expert_mode_allocations}{keep_loaded}" # Changed from compute_device
return hashlib.sha256(settings_str.encode()).hexdigest()
def override(self, *args, device=None, virtual_vram_gb=4.0, # Changed from compute_device
donor_device="cpu", expert_mode_allocations="", high_precision_loras=True, **kwargs):
donor_device="cpu", expert_mode_allocations="", keep_loaded=True, **kwargs):
from . import set_current_text_encoder_device # Use text encoder device setter
if device is not None:
@@ -938,35 +927,14 @@ def override_class_with_distorch_safetensor_v2_clip(cls):
# Call original function
fn = getattr(super(), cls.FUNCTION)
# --- Check if we need to unload the model due to settings change ---
# This logic is a bit redundant with IS_CHANGED, but provides clear logging
settings_str = f"{device}{virtual_vram_gb}{donor_device}{expert_mode_allocations}" # Changed from compute_device
settings_hash = hashlib.sha256(settings_str.encode()).hexdigest()
# Temporarily load to get hash without applying our patch
temp_out = fn(*args, **kwargs)
model_to_check = None
if hasattr(temp_out[0], 'model'):
model_to_check = temp_out[0]
elif hasattr(temp_out[0], 'patcher') and hasattr(temp_out[0].patcher, 'model'):
model_to_check = temp_out[0].patcher
if model_to_check:
model_hash = create_safetensor_model_hash(model_to_check, "override_check")
last_settings_hash = safetensor_settings_store.get(model_hash)
if last_settings_hash != settings_hash:
logger.mgpu_mm_log(f"Settings changed for model {model_hash[:8]}. Previous settings hash: {last_settings_hash}, New settings hash: {settings_hash}. Forcing reload.")
else:
logger.mgpu_mm_log(f"Settings unchanged for model {model_hash[:8]}. Using cached model.")
# Call the main function once
out = fn(*args, **kwargs)
# Store high_precision_loras in the model for later retrieval
# Store keep_loaded in the model for later retrieval by unload_all_models patch
if hasattr(out[0], 'model'):
out[0].model._distorch_high_precision_loras = high_precision_loras
out[0].model._mgpu_keep_loaded = keep_loaded
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
out[0].patcher.model._distorch_high_precision_loras = high_precision_loras
out[0].patcher.model._mgpu_keep_loaded = keep_loaded
vram_string = ""
if virtual_vram_gb > 0:
@@ -978,12 +946,19 @@ def override_class_with_distorch_safetensor_v2_clip(cls):
logger.info(f"[MultiGPU DisTorch V2] Full allocation string: {full_allocation}")
# Store allocation AFTER loading for next time
model_to_check = None
if hasattr(out[0], 'model'):
model_hash = create_safetensor_model_hash(out[0], "override")
safetensor_allocation_store[model_hash] = full_allocation
safetensor_settings_store[model_hash] = settings_hash
model_to_check = out[0]
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
model_hash = create_safetensor_model_hash(out[0].patcher, "override")
model_to_check = out[0].patcher
if model_to_check:
model_hash = create_safetensor_model_hash(model_to_check, "override_store")
settings_str = f"{device}{virtual_vram_gb}{donor_device}{expert_mode_allocations}"
settings_hash = hashlib.sha256(settings_str.encode()).hexdigest()
# Store allocation for next time
safetensor_allocation_store[model_hash] = full_allocation
safetensor_settings_store[model_hash] = settings_hash
@@ -1006,7 +981,7 @@ def override_class_with_distorch_safetensor_v2_clip_no_device(cls):
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"]["high_precision_loras"] = ("BOOLEAN", {"default": True})
inputs["optional"]["keep_loaded"] = ("BOOLEAN", {"default": True})
return inputs
CATEGORY = "multigpu/distorch_2"
@@ -1015,13 +990,13 @@ def override_class_with_distorch_safetensor_v2_clip_no_device(cls):
@classmethod
def IS_CHANGED(s, *args, device=None, virtual_vram_gb=4.0, # Changed from compute_device
donor_device="cpu", expert_mode_allocations="", high_precision_loras=True, **kwargs):
donor_device="cpu", expert_mode_allocations="", keep_loaded=True, **kwargs):
# Create a hash of our specific settings
settings_str = f"{device}{virtual_vram_gb}{donor_device}{expert_mode_allocations}{high_precision_loras}" # Changed from compute_device
settings_str = f"{device}{virtual_vram_gb}{donor_device}{expert_mode_allocations}{keep_loaded}" # Changed from compute_device
return hashlib.sha256(settings_str.encode()).hexdigest()
def override(self, *args, device=None, virtual_vram_gb=4.0, # Changed from compute_device
donor_device="cpu", expert_mode_allocations="", high_precision_loras=True, **kwargs):
donor_device="cpu", expert_mode_allocations="", keep_loaded=True, **kwargs):
from . import set_current_text_encoder_device # Use text encoder device setter
if device is not None:
@@ -1033,35 +1008,14 @@ def override_class_with_distorch_safetensor_v2_clip_no_device(cls):
# Call original function
fn = getattr(super(), cls.FUNCTION)
# --- Check if we need to unload the model due to settings change ---
# This logic is a bit redundant with IS_CHANGED, but provides clear logging
settings_str = f"{device}{virtual_vram_gb}{donor_device}{expert_mode_allocations}" # Changed from compute_device
settings_hash = hashlib.sha256(settings_str.encode()).hexdigest()
# Temporarily load to get hash without applying our patch
temp_out = fn(*args, **kwargs)
model_to_check = None
if hasattr(temp_out[0], 'model'):
model_to_check = temp_out[0]
elif hasattr(temp_out[0], 'patcher') and hasattr(temp_out[0].patcher, 'model'):
model_to_check = temp_out[0].patcher
if model_to_check:
model_hash = create_safetensor_model_hash(model_to_check, "override_check")
last_settings_hash = safetensor_settings_store.get(model_hash)
if last_settings_hash != settings_hash:
logger.info(f"[MultiGPU DisTorch V2] Settings changed for model {model_hash[:8]}. Previous settings hash: {last_settings_hash}, New settings hash: {settings_hash}. Forcing reload.")
else:
logger.info(f"[MultiGPU DisTorch V2] Settings unchanged for model {model_hash[:8]}. Using cached model.")
# Call the main function once
out = fn(*args, **kwargs)
# Store high_precision_loras in the model for later retrieval
# Store keep_loaded in the model for later retrieval by unload_all_models patch
if hasattr(out[0], 'model'):
out[0].model._distorch_high_precision_loras = high_precision_loras
out[0].model._mgpu_keep_loaded = keep_loaded
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
out[0].patcher.model._distorch_high_precision_loras = high_precision_loras
out[0].patcher.model._mgpu_keep_loaded = keep_loaded
vram_string = ""
if virtual_vram_gb > 0:
@@ -1073,12 +1027,19 @@ def override_class_with_distorch_safetensor_v2_clip_no_device(cls):
logger.info(f"[MultiGPU DisTorch V2] Full allocation string: {full_allocation}")
# Store allocation AFTER loading for next time
model_to_check = None
if hasattr(out[0], 'model'):
model_hash = create_safetensor_model_hash(out[0], "override")
safetensor_allocation_store[model_hash] = full_allocation
safetensor_settings_store[model_hash] = settings_hash
model_to_check = out[0]
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
model_hash = create_safetensor_model_hash(out[0].patcher, "override")
model_to_check = out[0].patcher
if model_to_check:
model_hash = create_safetensor_model_hash(model_to_check, "override_store")
settings_str = f"{device}{virtual_vram_gb}{donor_device}{expert_mode_allocations}"
settings_hash = hashlib.sha256(settings_str.encode()).hexdigest()
# Store allocation for next time
safetensor_allocation_store[model_hash] = full_allocation
safetensor_settings_store[model_hash] = settings_hash
+146
View File
@@ -0,0 +1,146 @@
# Code References (Definitive): ComfyUI Manager “Free model and node cache”
Purpose
- Provide an end-to-end, fully verified lineage of the ComfyUI Manager “Free model and node cache” button through to the exact consumption of flags in ComfyUI core, with exact file paths and code excerpts captured from the current snapshot in this workspace.
End‑to‑End Flow (Current Snapshot)
1) UI Button (Manager) → 2) JS helper free_models(...) → 3) POST /free (Comfy core) → 4) main.py prompt_worker thread polls flags and performs:
- unload_models: comfy.model_management.unload_all_models()
- free_memory: PromptExecutor.reset()
- Additionally triggers GC and comfy.model_management.soft_empty_cache()
A) Frontend UI trigger (ComfyUI Manager)
- File: ../ComfyUI-Manager/js/comfyui-manager.js
- Location: app.registerExtension({ name: "Comfy.ManagerMenu", ... }) → setup() → ComfyButtonGroup
```js
new(await import("../../scripts/ui/components/button.js")).ComfyButton({
icon: "vacuum-outline",
action: () => {
free_models();
},
tooltip: "Unload Models"
}).element,
new(await import("../../scripts/ui/components/button.js")).ComfyButton({
icon: "vacuum",
action: () => {
free_models(true);
},
tooltip: "Free model and node cache"
}).element,
```
Semantics:
- “Unload Models” → free_models() (models only)
- “Free model and node cache” → free_models(true) (models + execution cache)
B) Frontend request construction (ComfyUI Manager)
- File: ../ComfyUI-Manager/js/common.js
- Function: export async function free_models(free_execution_cache)
```js
export async function free_models(free_execution_cache) {
try {
let mode = "";
if (free_execution_cache) {
mode = '{"unload_models": true, "free_memory": true}';
} else {
mode = '{"unload_models": true}';
}
console.log(`[ManagerFreePath] POST /free payload: ${mode}`);
let res = await api.fetchApi(`/free`, {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
body: mode
});
console.log(`[ManagerFreePath] /free status: ${res.status}`);
if (res.status == 200) {
if (free_execution_cache) {
showToast("'Models' and 'Execution Cache' have been cleared.", 3000);
} else {
showToast("Models' have been unloaded.", 3000);
}
} else {
showToast('Unloading of models failed. Installed ComfyUI may be an outdated version.', 5000);
}
} catch (error) {
console.error('[ManagerFreePath] /free error:', error);
showToast('An error occurred while trying to unload models.', 5000);
}
}
```
Semantics:
- free_models(true) → POST /free with {"unload_models": true, "free_memory": true}
- free_models() → POST /free with {"unload_models": true}
C) Core server endpoint (flags are set on the queue)
- File: ../../server.py
- Route: @routes.post("/free")
```py
@routes.post("/free")
async def post_free(request):
json_data = await request.json()
unload_models = json_data.get("unload_models", False)
free_memory = json_data.get("free_memory", False)
if unload_models:
self.prompt_queue.set_flag("unload_models", unload_models)
if free_memory:
self.prompt_queue.set_flag("free_memory", free_memory)
return web.Response(status=200)
```
Semantics:
- The HTTP endpoint itself does not unload/reset; instead it sets flags on PromptServer.prompt_queue for the background worker to consume.
D) Flag consumption and execution (definitive mechanism)
- File: ../../main.py
- Function: prompt_worker(q, server_instance)
- Excerpt (poll and handle flags, then clean up):
```py
flags = q.get_flags()
free_memory = flags.get("free_memory", False)
if flags.get("unload_models", free_memory):
comfy.model_management.unload_all_models()
need_gc = True
last_gc_collect = 0
if free_memory:
e.reset()
need_gc = True
last_gc_collect = 0
if need_gc:
current_time = time.perf_counter()
if (current_time - last_gc_collect) > gc_collect_interval:
gc.collect()
comfy.model_management.soft_empty_cache()
last_gc_collect = current_time
need_gc = False
hook_breaker_ac10a0.restore_functions()
```
Context:
- e is a PromptExecutor (created earlier in prompt_worker): `e = execution.PromptExecutor(server_instance, ...)`
- The worker thread is started in start_comfyui():
```py
threading.Thread(target=prompt_worker, daemon=True, args=(prompt_server.prompt_queue, prompt_server,)).start()
```
Interpretation (What the Manager button actually does)
- “Free model and node cache” sets unload_models: true and free_memory: true via POST /free.
- The background prompt_worker then:
- Calls comfy.model_management.unload_all_models()
- Calls e.reset() on the PromptExecutor to drop execution caches
- Performs gc.collect() and comfy.model_management.soft_empty_cache()
- This matches the “benchmark button” behavior required for CPU memory reclamation (models fully unloaded + executor reset + allocator/cache cleanup).
Implications for MultiGPU P1 (force_full_system_cleanup)
- To 100% replicate the benchmark button behavior from within MultiGPU code paths:
- Call comfy.model_management.unload_all_models()
- Trigger PromptExecutor.reset() on the active executor
- Follow up with gc.collect() and comfy.model_management.soft_empty_cache()
- Or, trigger the core behavior indirectly by POST /free with both flags set, relying on ComfyUI’s running prompt worker.
Verification Status
- All file paths and snippets above were extracted from this workspace:
- Manager JS files under ../ComfyUI-Manager/js/
- ComfyUI server and main under ../../server.py and ../../main.py
- Consumption site conclusively identified in ../../main.py prompt_worker via q.get_flags → unload_all_models + PromptExecutor.reset
File diff suppressed because it is too large Load Diff
+140 -32
View File
@@ -19,6 +19,34 @@ 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)
# ==========================================================================================
@@ -63,27 +91,44 @@ def prune_distorch_stores():
active_hashes_v2 = set()
active_hashes_v1 = set()
for lm in mm.current_loaded_models:
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:
active_hashes_v2.add(create_safetensor_model_hash(mp, "prune_check_v2"))
active_hashes_v1.add(create_model_hash(mp, "prune_check_v1"))
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.info(f"[MultiGPU_Memory_Management] Pruning {len(stale_v1)} DisTorch V1 entries")
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
# 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.info(f"[MultiGPU_Memory_Management] Pruning {len(stale_v2)} V2 {name} entries")
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")
# ==========================================================================================
@@ -117,7 +162,7 @@ def _capture_memory_snapshot():
return snapshot
def multigpu_memory_log(identifier, tag):
"""Record timestamped memory snapshot with delta logging"""
"""Record timestamped memory snapshot with clean aligned logging"""
if identifier == "print_summary":
for id_key in sorted(_MEM_SNAPSHOT_SERIES.keys()):
series = _MEM_SNAPSHOT_SERIES[id_key]
@@ -125,12 +170,13 @@ def multigpu_memory_log(identifier, tag):
for ts, tag_name, snap in series:
parts = []
cpu_used, cpu_total = snap.get("cpu", (0, 0))
parts.append(f"cpu={cpu_used/(1024**3):.2f}/{cpu_total/(1024**3):.2f}")
parts.append(f"cpu|{cpu_used/(1024**3):.2f}")
for dev in sorted([k for k in snap.keys() if k != "cpu"]):
used, total = snap[dev]
parts.append(f"{dev}={used/(1024**3):.2f}/{total/(1024**3):.2f}")
parts.append(f"{dev}|{used/(1024**3):.2f}")
ts_str = ts.strftime("%Y-%m-%dT%H:%M:%S.%f")[:-3] + "Z"
logger.mgpu_mm_log(f"{ts_str} {id_key} {tag_name} | " + " | ".join(parts))
tag_padded = f"{id_key}_{tag_name}".ljust(35)
logger.mgpu_mm_log(f"{ts_str} {tag_padded} {' '.join(parts)}")
return
ts = datetime.now(timezone.utc)
@@ -141,28 +187,20 @@ def multigpu_memory_log(identifier, tag):
_MEM_SNAPSHOT_SERIES[identifier] = []
_MEM_SNAPSHOT_SERIES[identifier].append((ts, tag, curr))
# Compute delta
if identifier in _MEM_SNAPSHOT_LAST:
prev_tag, prev = _MEM_SNAPSHOT_LAST[identifier]
keys = set(prev.keys()) | set(curr.keys())
ordered = ["cpu"] + sorted([k for k in keys if k != "cpu"])
parts = []
for k in ordered:
p_used, _ = prev.get(k, (0, 0))
c_used, _ = curr.get(k, (0, 0))
delta = c_used - p_used
sign = "+" if delta >= 0 else "-"
parts.append(f"{k}={sign}{abs(delta)/(1024**3):.2f}")
logger.mgpu_mm_log(f"{identifier} {tag} - {prev_tag}: " + " | ".join(parts))
else:
# Baseline
ordered = ["cpu"] + sorted([k for k in curr.keys() if k != "cpu"])
parts = []
for k in ordered:
c_used, _ = curr.get(k, (0, 0))
parts.append(f"{k}=+{c_used/(1024**3):.2f}")
logger.mgpu_mm_log(f"{identifier} {tag} - <baseline>: " + " | ".join(parts))
# Clean aligned format: timestamp + padded tag + memory values
ts_str = ts.strftime("%Y-%m-%dT%H:%M:%S.%f")[:-3] + "Z"
tag_padded = f"{identifier}_{tag}".ljust(35)
parts = []
cpu_used, _ = curr.get("cpu", (0, 0))
parts.append(f"cpu|{cpu_used/(1024**3):.2f}")
for dev in sorted([k for k in curr.keys() if k != "cpu"]):
used, _ = curr[dev]
parts.append(f"{dev}|{used/(1024**3):.2f}")
logger.mgpu_mm_log(f"{ts_str} {tag_padded} {' '.join(parts)}")
_MEM_SNAPSHOT_LAST[identifier] = (tag, curr)
def clear_memory_snapshot_history():
@@ -331,3 +369,73 @@ 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 with keep_loaded retention
# ==========================================================================================
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")
_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.
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)}")
# 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
if mp is not None and hasattr(mp, 'model'):
# Check if this is a DisTorch model with keep_loaded flag
keep_loaded = getattr(mp.model, '_mgpu_keep_loaded', False)
model_name = type(getattr(mp, 'model', mp)).__name__
logger.mgpu_mm_log(f"[UNLOAD_DEBUG] Model {i}: {model_name}, keep_loaded={keep_loaded}")
if keep_loaded:
kept_models.append(lm)
logger.mgpu_mm_log(f"[UNLOAD_DEBUG] Adding to kept_models: {model_name}")
# GC ANCHOR TEST: Prevent premature GC of clone patchers
add_retention_anchor(mp, "keep_loaded_test")
else:
models_to_unload.append(lm)
else:
logger.mgpu_mm_log(f"[UNLOAD_DEBUG] Model {i}: ModelPatcher is None or missing model attribute")
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:
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}")
# 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")
else:
logger.mgpu_mm_log("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_keep_loaded_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")
+24
View File
@@ -530,6 +530,30 @@ class DownloadAndLoadHyVideoTextEncoder:
return original_loader.loadmodel(llm_model, clip_model, precision, apply_final_norm, hidden_state_skip_layer, quantization)
class UNetLoaderLP:
"""UNet Loader (Low Precision) - sets LoRA precision to False for CPU storage optimization"""
@classmethod
def INPUT_TYPES(s):
return {"required": { "unet_name": (folder_paths.get_filename_list("unet"), ),
}}
RETURN_TYPES = ("MODEL",)
FUNCTION = "load_unet"
CATEGORY = "loaders"
TITLE = "UNet Loader (LP)"
def load_unet(self, unet_name):
original_loader = NODE_CLASS_MAPPINGS["UNETLoader"]()
out = original_loader.load_unet(unet_name)
# Set the low-precision LoRA flag on the loaded model
if hasattr(out[0], 'model'):
out[0].model._distorch_high_precision_loras = False
elif hasattr(out[0], 'patcher') and hasattr(out[0].patcher, 'model'):
out[0].patcher.model._distorch_high_precision_loras = False
return out
class FullCleanupMultiGPU:
@classmethod
def INPUT_TYPES(s):