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:
+22
-18
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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")
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user