Files
xmarre-ComfyUI-GPU-Resident…/cleanup.py
coderabbitai[bot] 61ed3e16d4 📝 Add docstrings to codex/external-fallback-free-memory
Docstrings generation was requested by @xmarre.

The following files were modified:

* `cleanup.py`
* `patches.py`
2026-04-16 01:47:19 +00:00

405 lines
17 KiB
Python

from __future__ import annotations
import os
from contextlib import contextmanager
from typing import Any
import comfy.model_management as model_management
import torch
from .external_residency import EXTERNAL_REGISTRY, ensure_external_integrations_installed, external_trim_enabled
from .residency import REGISTRY
_ADAPTIVE_HEADROOM_RATIO = 0.125
_ADAPTIVE_HEADROOM_FLOOR_BYTES = 256 * 1024 * 1024
_ADAPTIVE_HEADROOM_CEIL_BYTES = 1024 * 1024 * 1024
def _safe_free_memory(device) -> int:
return int(model_management.get_free_memory(device))
def _normalize_trim_device(device: str | torch.device | None):
if device is None or isinstance(device, torch.device):
return device
try:
return torch.device(device)
except Exception:
return device
def _safe_is_dead(loaded) -> bool:
try:
return loaded.is_dead()
except Exception:
return False
def _device_matches(device_a, device_b) -> bool:
normalized_a = _normalize_trim_device(device_a)
normalized_b = _normalize_trim_device(device_b)
if normalized_a is None or normalized_b is None:
return False
if isinstance(normalized_a, torch.device) and isinstance(normalized_b, torch.device):
if normalized_a.type != normalized_b.type:
return False
if normalized_a.type == "cuda":
index_a = 0 if normalized_a.index is None else normalized_a.index
index_b = 0 if normalized_b.index is None else normalized_b.index
return index_a == index_b
return True
return str(normalized_a) == str(normalized_b)
def _should_force_cpu_offload(model: Any, *, active_device=None, force: bool = False) -> bool:
if force:
return True
active_device = _normalize_trim_device(active_device)
if not isinstance(active_device, torch.device) or active_device.type != "cuda":
return False
return _device_matches(active_device, getattr(model, "offload_device", None))
@contextmanager
def _temporary_offload_device(model: Any, device):
if model is None or device is None or not hasattr(model, "offload_device"):
yield
return
original_device = model.offload_device
model.offload_device = device
try:
yield
finally:
model.offload_device = original_device
def unload_loaded_model(
loaded,
*,
active_device: str | torch.device | None = None,
force_offload_to_cpu: bool = False,
unpatch_weights: bool = True,
) -> bool:
# This helper is intentionally scoped to full unloads owned by this plugin.
model = getattr(loaded, "model", None)
if model is None:
return False
active_device = _normalize_trim_device(active_device if active_device is not None else getattr(loaded, "device", None))
force_cpu_offload = _should_force_cpu_offload(
model,
active_device=active_device,
force=force_offload_to_cpu,
)
unload_target = torch.device("cpu") if force_cpu_offload else None
with _temporary_offload_device(model, unload_target):
return loaded.model_unload(None, unpatch_weights=unpatch_weights)
def _sort_key_for_candidate(entry, *, sticky_respected: bool) -> tuple[int, int, float]:
if not sticky_respected:
return (0, 0, getattr(entry, "last_touched", 0.0) if entry is not None else 0.0)
priority = int(getattr(entry, "priority", 0)) if entry is not None else 0
last_touched = float(getattr(entry, "last_touched", 0.0)) if entry is not None else 0.0
return (1, priority, last_touched)
def _should_keep_loaded_model(model: Any, keep_models: tuple[Any, ...]) -> bool:
if not keep_models:
return False
for keep in keep_models:
if keep is None:
continue
if model is keep:
return True
is_clone = getattr(model, "is_clone", None)
if callable(is_clone):
try:
if is_clone(keep):
return True
except Exception:
pass
return False
def _trim_candidates(
*,
device,
respect_sticky: bool,
sticky_floor_priority: int,
keep_models: tuple[Any, ...],
include_external: bool | None = None,
) -> list[tuple[Any, Any, bool, bool]]:
"""
Collects and returns eviction/unload candidates from in-memory models and, optionally, external integrations.
Filters currently loaded models by the given device, excludes dead or missing models, and skips models listed in `keep_models`. If `respect_sticky` is true, entries whose registry metadata mark them as sticky and whose priority meets or exceeds `sticky_floor_priority` are flagged so they are treated as higher-priority to keep. When `include_external` is true (or when `include_external` is None and external trimming is enabled), candidates from the external registry are included.
Parameters:
device: Device filter for candidates; if not None only candidates matching this device are considered.
respect_sticky (bool): Whether to respect sticky registry entries when computing candidate priority.
sticky_floor_priority (int): Minimum priority value for a registry entry to be considered sticky.
keep_models (tuple[Any, ...]): Objects that must not be selected as candidates.
include_external (bool | None): If True include external-registry candidates; if False exclude them; if None defer to runtime external_trim_enabled().
Returns:
list[tuple[Any, Any, bool, bool]]: A sorted list of tuples (candidate_obj, registry_entry_or_None, sticky_respected, is_external_candidate).
- candidate_obj: The loaded model object (internal) or the external object.
- registry_entry_or_None: Registry metadata for the candidate, or None if unavailable.
- sticky_respected: `True` when the candidate is marked sticky and meets `sticky_floor_priority`.
- is_external_candidate: `True` for candidates originating from the external registry.
"""
candidates: list[tuple[Any, Any, bool, bool]] = []
ensure_external_integrations_installed()
for loaded in list(model_management.current_loaded_models):
if device is not None and loaded.device != device:
continue
if _safe_is_dead(loaded):
continue
model = getattr(loaded, "model", None)
if model is None:
continue
if _should_keep_loaded_model(model, keep_models):
continue
entry = REGISTRY.entry_for_object(model)
sticky_respected = (
respect_sticky
and entry is not None
and bool(getattr(entry, "sticky", False))
and int(getattr(entry, "priority", 0)) >= int(sticky_floor_priority)
)
candidates.append((loaded, entry, sticky_respected, False))
should_include_external = external_trim_enabled() if include_external is None else bool(include_external)
if should_include_external:
for external_obj, entry, sticky_respected in EXTERNAL_REGISTRY.candidates(
device=device,
respect_sticky=respect_sticky,
sticky_floor_priority=sticky_floor_priority,
keep_models=keep_models,
):
candidates.append((external_obj, entry, sticky_respected, True))
candidates.sort(key=lambda item: _sort_key_for_candidate(item[1], sticky_respected=item[2]))
return candidates
def trim_resident_vram(
*,
device: str | torch.device | None = None,
target_free_vram_bytes: int,
respect_sticky: bool,
sticky_floor_priority: int,
allow_partial_unload: bool,
keep_models: tuple[Any, ...] = (),
include_external: bool | None = None,
) -> dict[str, Any]:
"""
Trim resident models (and optional external integration objects) until the requested amount of free VRAM is available or a stopping condition occurs.
Tries to free GPU memory on `device` by evicting or unloading loaded models (and, optionally, external candidates). Respects sticky/priority hints, can perform partial unloads of pinned RAM when supported, and records each attempted action in the returned report.
Parameters:
device (str | torch.device | None): Target device to free (defaults to model_management.get_torch_device()).
target_free_vram_bytes (int): Desired amount of free VRAM, in bytes.
respect_sticky (bool): If true, prefer protecting entries marked as sticky with sufficient priority.
sticky_floor_priority (int): Minimum priority required for a sticky entry to be respected.
allow_partial_unload (bool): If true, allow partial unloads and attempts to free pinned host RAM before full eviction.
keep_models (tuple[Any, ...]): Sequence of model objects that must not be unloaded (exact matches or recognized clones).
include_external (bool | None): If None, use external_trim_enabled() at runtime; otherwise force inclusion/exclusion of external candidates.
Returns:
dict[str, Any]: Report of the trimming operation containing:
- status: "met_target", "partial", or "error".
- stopped_reason: reason the loop stopped (e.g., "target_met", "no_candidates", "no_progress", "error").
- target_met (bool): whether the target free VRAM was reached.
- device (str): string form of the device used.
- target_free_vram_bytes (int), free_before_bytes (int), free_after_bytes (int).
- freed_vram_bytes (int): total freed VRAM during this call.
- respect_sticky (bool), sticky_floor_priority (int), allow_partial_unload (bool).
- external_trim_enabled (bool): computed flag indicating whether external candidates were considered.
- actions (list[dict]): ordered per-candidate action records; each entry includes metadata such as
entry_id, basename, tracked, external_candidate, sticky_respected, priority,
need_before_bytes, loaded_before_bytes, freed_pinned_ram_bytes,
mode (e.g., "full_unload", "partial_unload", "external_evict", "error"),
loaded_after_bytes, freed_vram_bytes, free_after_bytes, and any warnings/errors.
"""
ensure_external_integrations_installed()
cleanup_models_gc = getattr(model_management, "cleanup_models_gc", None)
if callable(cleanup_models_gc):
cleanup_models_gc()
REGISTRY.refresh_runtime_state()
EXTERNAL_REGISTRY.refresh_runtime_state()
device = _normalize_trim_device(device)
if device is None:
device = model_management.get_torch_device()
free_before = _safe_free_memory(device)
actions: list[dict[str, Any]] = []
soft_empty_cache = getattr(model_management, "soft_empty_cache", None)
stopped_reason = "target_met"
while True:
free_now = _safe_free_memory(device)
need = int(target_free_vram_bytes) - free_now
if need <= 0:
stopped_reason = "target_met"
break
candidates = _trim_candidates(
device=device,
respect_sticky=respect_sticky,
sticky_floor_priority=sticky_floor_priority,
keep_models=keep_models,
include_external=include_external,
)
if not candidates:
stopped_reason = "no_candidates"
break
candidate, entry, sticky_respected, is_external_candidate = candidates[0]
model = candidate if is_external_candidate else candidate.model
loaded_before = int(getattr(entry, "loaded_bytes", 0)) if is_external_candidate else int(candidate.model_loaded_memory())
action = {
"entry_id": getattr(entry, "entry_id", None),
"basename": None if entry is None else os.path.basename(getattr(entry, "source_path", "") or ""),
"tracked": entry is not None,
"external_candidate": bool(is_external_candidate),
"sticky_respected": sticky_respected,
"priority": None if entry is None else int(getattr(entry, "priority", 0)),
"need_before_bytes": need,
"loaded_before_bytes": loaded_before,
"freed_pinned_ram_bytes": 0,
}
if (not is_external_candidate and allow_partial_unload and hasattr(model, "pinned_memory_size") and hasattr(model, "partially_unload_ram")):
try:
pinned_memory = int(model.pinned_memory_size())
if pinned_memory > 0:
pinned_budget = min(pinned_memory, max(need, 0))
model.partially_unload_ram(pinned_budget)
action["freed_pinned_ram_bytes"] = pinned_budget
except Exception as exc:
action["pinned_ram_warning"] = str(exc)
if is_external_candidate:
try:
fully_unloaded = EXTERNAL_REGISTRY.evict(entry)
action["mode"] = "external_evict"
except Exception as exc:
action["mode"] = "error"
action["error"] = str(exc)
actions.append(action)
stopped_reason = "error"
break
else:
try:
fully_unloaded = candidate.model_unload(need if allow_partial_unload else None)
action["mode"] = "full_unload" if fully_unloaded else "partial_unload"
except Exception as exc:
if allow_partial_unload:
try:
fully_unloaded = candidate.model_unload(None)
action["mode"] = "full_unload_fallback"
action["partial_unload_warning"] = str(exc)
except Exception as fallback_exc:
action["mode"] = "error"
action["error"] = str(fallback_exc)
actions.append(action)
stopped_reason = "error"
break
else:
action["mode"] = "error"
action["error"] = str(exc)
actions.append(action)
stopped_reason = "error"
break
if fully_unloaded and not is_external_candidate:
try:
model_management.current_loaded_models.remove(candidate)
except ValueError:
pass
if callable(soft_empty_cache):
soft_empty_cache()
REGISTRY.refresh_runtime_state()
EXTERNAL_REGISTRY.refresh_runtime_state()
free_after = _safe_free_memory(device)
if is_external_candidate:
action["loaded_after_bytes"] = 0 if fully_unloaded else int(getattr(entry, "loaded_bytes", 0))
else:
action["loaded_after_bytes"] = 0 if fully_unloaded else int(candidate.model_loaded_memory())
action["freed_vram_bytes"] = max(0, free_after - free_now)
action["free_after_bytes"] = free_after
actions.append(action)
if action["freed_vram_bytes"] <= 0 and not fully_unloaded:
stopped_reason = "no_progress"
break
free_after = _safe_free_memory(device)
target_met = free_after >= int(target_free_vram_bytes)
if target_met:
stopped_reason = "target_met"
return {
"status": "met_target" if target_met else ("error" if stopped_reason == "error" else "partial"),
"stopped_reason": stopped_reason,
"target_met": target_met,
"device": str(device),
"target_free_vram_bytes": int(target_free_vram_bytes),
"free_before_bytes": free_before,
"free_after_bytes": free_after,
"freed_vram_bytes": max(0, free_after - free_before),
"respect_sticky": bool(respect_sticky),
"sticky_floor_priority": int(sticky_floor_priority),
"allow_partial_unload": bool(allow_partial_unload),
"external_trim_enabled": bool(external_trim_enabled() if include_external is None else include_external),
"actions": actions,
}
def adaptive_headroom_bytes(required_bytes: int) -> int:
required = max(0, int(required_bytes))
if required == 0:
return 0
return min(
_ADAPTIVE_HEADROOM_CEIL_BYTES,
max(_ADAPTIVE_HEADROOM_FLOOR_BYTES, int(required * _ADAPTIVE_HEADROOM_RATIO)),
)
def trim_resident_vram_for_load(
*,
required_bytes: int,
reason: str,
device: str | torch.device | None = None,
respect_sticky: bool = True,
sticky_floor_priority: int = 0,
allow_partial_unload: bool = True,
keep_models: tuple[Any, ...] = (),
) -> dict[str, Any]:
estimated_load_bytes = max(0, int(required_bytes))
headroom_bytes = adaptive_headroom_bytes(estimated_load_bytes)
target_free_vram_bytes = estimated_load_bytes + headroom_bytes
report = trim_resident_vram(
device=device,
target_free_vram_bytes=target_free_vram_bytes,
respect_sticky=respect_sticky,
sticky_floor_priority=sticky_floor_priority,
allow_partial_unload=allow_partial_unload,
keep_models=keep_models,
)
report["trim_strategy"] = "adaptive_load_request"
report["trim_reason"] = str(reason)
report["estimated_load_bytes"] = estimated_load_bytes
report["adaptive_headroom_bytes"] = headroom_bytes
report["kept_loaded_models"] = len([model for model in keep_models if model is not None])
return report