Compare commits
18
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
cfb707f082 | ||
|
|
19e69a8ecc | ||
|
|
f7b8fe5cc9 | ||
|
|
2643a1903b | ||
|
|
d11d8e0a25 | ||
|
|
df9128e9ee | ||
|
|
9f4b27eb23 | ||
|
|
b94fd0b3f3 | ||
|
|
1b80a6a82b | ||
|
|
6d4bd6c7c4 | ||
|
|
d533beb275 | ||
|
|
24ab29880a | ||
|
|
791a48afd3 | ||
|
|
1253b88361 | ||
|
|
ca8118de03 | ||
|
|
752def584f | ||
|
|
871c144ac9 | ||
|
|
ca0d55ce04 |
@@ -16,7 +16,7 @@ Stock ComfyUI makes separate decisions for:
|
||||
|
||||
Those are not the same thing.
|
||||
|
||||
This repo targets the second problem directly for `.safetensors` by steering eligible loads toward direct GPU ingest, then targets the first problem by overriding offload policy and by teaching `free_memory()` to respect sticky entries until the VRAM budget is actually exceeded.
|
||||
This repo targets the second problem directly for `.safetensors` by steering eligible loads toward direct GPU ingest, narrowing resident diffusion-model loads down to UNet tensors only, then targets the first problem by overriding offload policy and by teaching `free_memory()` to protect high-priority sticky entries until the VRAM budget says otherwise.
|
||||
|
||||
## What is included
|
||||
|
||||
@@ -46,6 +46,9 @@ Current patch surface:
|
||||
|
||||
- **Diffusion Model Loader Resident**
|
||||
- **Checkpoint Loader Resident**
|
||||
- **Checkpoint Model Loader Resident**
|
||||
- **Checkpoint Clip Loader Resident**
|
||||
- **Checkpoint VAE Loader Resident**
|
||||
- **Diffusion Model Selector Resident**
|
||||
|
||||
`Diffusion Model Loader Resident` mirrors the relevant KJ diffusion-loader feature surface:
|
||||
@@ -57,6 +60,8 @@ Current patch surface:
|
||||
- fp16 accumulation toggle
|
||||
- optional extra-state-dict merge
|
||||
|
||||
On `.safetensors`, the resident diffusion-model path now reads only the detected UNet keys and merges only matching keys from any extra state dict. Repeated resident loads also reuse a live equivalent object when the source path and loader-relevant options still match.
|
||||
|
||||
### Residency nodes
|
||||
|
||||
- **Set Global Residency Policy**
|
||||
@@ -72,8 +77,8 @@ The startup patcher exposes four policies:
|
||||
|
||||
- `legacy` — leave ingest/offload behavior close to stock ComfyUI.
|
||||
- `balanced` — keep the registry and diagnostics, but do not aggressively steer ingest to GPU.
|
||||
- `prefer_gpu` — prefer GPU ingest and GPU offload devices, but do not auto-pin tracked objects.
|
||||
- `sticky_gpu` — prefer GPU ingest, prefer GPU offload devices, and auto-mark tracked loader outputs sticky.
|
||||
- `prefer_gpu` — prefer GPU ingest, keep UNet/ControlNet/CLIP on the faster side of the device policy, but do not auto-pin tracked objects.
|
||||
- `sticky_gpu` — prefer GPU ingest, auto-pin the highest-value tracked outputs, and let lower-priority sticky entries yield first when VRAM pressure rises.
|
||||
|
||||
Default selection order:
|
||||
|
||||
@@ -88,11 +93,11 @@ Default selection order:
|
||||
|
||||
This repo is optimized around `.safetensors`.
|
||||
|
||||
Direct GPU ingest is attempted for `.safetensors` loads. If the direct path fails, the patcher falls back to CPU read + GPU copy and records that fallback in the registry.
|
||||
Direct GPU ingest is attempted for `.safetensors` loads. Resident diffusion-model loads take a header-first selective path and fetch only the detected UNet tensors instead of loading the whole file and filtering later. If the direct path fails, the patcher falls back to CPU read + GPU copy and records that fallback in the registry.
|
||||
|
||||
### `.ckpt` / `.pt` remain CPU-first under PyTorch
|
||||
|
||||
Those formats still go through `torch.load()`. The repo tracks that path and can still keep the resulting model hot in VRAM, but it does **not** claim true direct-to-GPU checkpoint ingest for pickle-based formats.
|
||||
Those formats still go through `torch.load()`. The repo tracks that path and can still keep the resulting model hot in VRAM, but it does **not** claim true direct-to-GPU checkpoint ingest for pickle-based formats. Resident checkpoint nodes now warn about this when a GPU-resident policy is active so the compatibility path is not mistaken for the fast path.
|
||||
|
||||
Use the included conversion helper to migrate hot models to `.safetensors`.
|
||||
|
||||
@@ -133,7 +138,17 @@ Recommended on a large VRAM machine:
|
||||
|
||||
Use **Checkpoint Loader Resident**.
|
||||
|
||||
That tracks and binds the resulting diffusion model, CLIP, and VAE independently so they appear in the registry snapshot.
|
||||
That tracks and binds the resulting diffusion model, CLIP, and VAE independently so they appear in the registry snapshot. If an equivalent live model, CLIP, or VAE already exists, the loader reuses it instead of rebuilding it.
|
||||
|
||||
### For staged checkpoint loads
|
||||
|
||||
Use:
|
||||
|
||||
- **Checkpoint Model Loader Resident** when the workflow only needs the diffusion model
|
||||
- **Checkpoint Clip Loader Resident** when the workflow only needs the text encoder
|
||||
- **Checkpoint VAE Loader Resident** when the workflow only needs the VAE
|
||||
|
||||
The model-only checkpoint node takes the same selective safetensors UNet fast path as the diffusion-model loader. CLIP-only and VAE-only nodes still use ComfyUI's checkpoint construction logic, but they avoid materializing the other outputs and can reuse an already-live equivalent object.
|
||||
|
||||
### For manual residency control
|
||||
|
||||
|
||||
+172
@@ -0,0 +1,172 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
import comfy.model_management as model_management
|
||||
|
||||
from .residency import REGISTRY
|
||||
|
||||
|
||||
def _safe_free_memory(device) -> int:
|
||||
return int(model_management.get_free_memory(device))
|
||||
|
||||
|
||||
def _safe_is_dead(loaded) -> bool:
|
||||
try:
|
||||
return loaded.is_dead()
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
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 _trim_candidates(*, device, respect_sticky: bool, sticky_floor_priority: int) -> list[tuple[Any, Any, bool]]:
|
||||
candidates: list[tuple[Any, Any, bool]] = []
|
||||
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
|
||||
|
||||
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))
|
||||
|
||||
candidates.sort(key=lambda item: _sort_key_for_candidate(item[1], sticky_respected=item[2]))
|
||||
return candidates
|
||||
|
||||
|
||||
def trim_resident_vram(
|
||||
*,
|
||||
target_free_vram_bytes: int,
|
||||
respect_sticky: bool,
|
||||
sticky_floor_priority: int,
|
||||
allow_partial_unload: bool,
|
||||
) -> dict[str, Any]:
|
||||
cleanup_models_gc = getattr(model_management, "cleanup_models_gc", None)
|
||||
if callable(cleanup_models_gc):
|
||||
cleanup_models_gc()
|
||||
|
||||
REGISTRY.refresh_runtime_state()
|
||||
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,
|
||||
)
|
||||
if not candidates:
|
||||
stopped_reason = "no_candidates"
|
||||
break
|
||||
|
||||
loaded, entry, sticky_respected = candidates[0]
|
||||
model = loaded.model
|
||||
loaded_before = int(loaded.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,
|
||||
"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 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)
|
||||
|
||||
try:
|
||||
fully_unloaded = loaded.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 = loaded.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:
|
||||
try:
|
||||
model_management.current_loaded_models.remove(loaded)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
if callable(soft_empty_cache):
|
||||
soft_empty_cache()
|
||||
|
||||
REGISTRY.refresh_runtime_state()
|
||||
free_after = _safe_free_memory(device)
|
||||
action["loaded_after_bytes"] = 0 if fully_unloaded else int(loaded.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),
|
||||
"actions": actions,
|
||||
}
|
||||
+682
-47
@@ -1,7 +1,9 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
import folder_paths
|
||||
@@ -12,7 +14,8 @@ import comfy.utils
|
||||
from comfy.cli_args import PerformanceFeature, args
|
||||
from comfy.ldm.modules.attention import attention_pytorch, wrap_attn
|
||||
|
||||
from .residency import KIND_CHECKPOINT, KIND_CLIP, KIND_MODEL, KIND_VAE, REGISTRY
|
||||
from .patches import checkpoint_component_info_from_header, infer_unet_prefix_from_keys, load_safetensors_state_dict
|
||||
from .residency import KIND_CHECKPOINT, KIND_CLIP, KIND_MODEL, KIND_VAE, POLICIES, REGISTRY
|
||||
|
||||
|
||||
_LOG = logging.getLogger(__name__)
|
||||
@@ -36,6 +39,13 @@ DTYPE_MAP = {
|
||||
"fp32": torch.float32,
|
||||
}
|
||||
|
||||
UNET_PREFIX_CANDIDATES = (
|
||||
"model.diffusion_model.",
|
||||
"model.model.",
|
||||
"net.",
|
||||
"model.",
|
||||
)
|
||||
|
||||
|
||||
def _set_cublas_linear(enabled: bool) -> None:
|
||||
if enabled:
|
||||
@@ -244,6 +254,504 @@ def _build_model_options(weight_dtype: str) -> dict[str, Any]:
|
||||
return model_options
|
||||
|
||||
|
||||
def _normalize_optional_string(value: str | None) -> str | None:
|
||||
if value is None:
|
||||
return None
|
||||
normalized = str(value).strip()
|
||||
return normalized or None
|
||||
|
||||
|
||||
def _effective_policy_name(policy_override: str | None) -> str:
|
||||
return REGISTRY.get_policy() if policy_override is None else policy_override
|
||||
|
||||
|
||||
def _make_loader_key(loader_name: str, **payload: Any) -> str:
|
||||
normalized_payload = {"loader": loader_name}
|
||||
normalized_payload.update(payload)
|
||||
return json.dumps(normalized_payload, sort_keys=True, separators=(",", ":"))
|
||||
|
||||
|
||||
def _normalize_optional_policy(value: str | None) -> str | None:
|
||||
normalized = _normalize_optional_string(value)
|
||||
return None if normalized is None else normalized.lower()
|
||||
|
||||
|
||||
def _resolve_loader_policy_and_extra_state_dict(
|
||||
*,
|
||||
loader_name: str,
|
||||
policy_override: str | None,
|
||||
extra_state_dict: str | None,
|
||||
) -> tuple[str | None, str | None]:
|
||||
normalized_policy = _normalize_optional_policy(policy_override)
|
||||
if normalized_policy is not None and normalized_policy not in POLICIES:
|
||||
raise ValueError(
|
||||
f"{loader_name}: unsupported policy_override {normalized_policy!r}. "
|
||||
f"Expected one of: {', '.join(POLICIES)}."
|
||||
)
|
||||
|
||||
normalized_extra = _normalize_optional_string(extra_state_dict)
|
||||
if normalized_extra is None:
|
||||
return normalized_policy, None
|
||||
|
||||
if normalized_policy is None and normalized_extra.lower() in POLICIES and not os.path.exists(normalized_extra):
|
||||
_LOG.warning(
|
||||
"GPU Resident Loader: interpreting legacy extra_state_dict value %r as policy_override for %s",
|
||||
normalized_extra,
|
||||
loader_name,
|
||||
)
|
||||
return normalized_extra.lower(), None
|
||||
|
||||
if not os.path.isfile(normalized_extra):
|
||||
raise FileNotFoundError(
|
||||
f"{loader_name}: extra_state_dict must point to an existing state-dict file, got {normalized_extra!r}. "
|
||||
f"If you meant to pass a residency policy, connect that STRING to policy_override instead."
|
||||
)
|
||||
|
||||
return normalized_policy, normalized_extra
|
||||
|
||||
|
||||
def _is_safetensors_path(path: str) -> bool:
|
||||
lowered = path.lower()
|
||||
return lowered.endswith(".safetensors") or lowered.endswith(".sft")
|
||||
|
||||
|
||||
def _warn_pickle_checkpoint_gpu_compatibility(loader_name: str, path: str) -> None:
|
||||
if _is_safetensors_path(path):
|
||||
return
|
||||
if not REGISTRY.wants_gpu_ingest(KIND_MODEL):
|
||||
return
|
||||
_LOG.warning(
|
||||
"%s: %s is using the compatibility path through torch.load() on CPU before tensor-by-tensor copies to GPU. "
|
||||
"Convert hot checkpoints to safetensors with scripts/convert_checkpoint_to_safetensors.py for the real fast path.",
|
||||
loader_name,
|
||||
path,
|
||||
)
|
||||
|
||||
|
||||
def _selected_unet_key_map_from_header(
|
||||
path: str,
|
||||
*,
|
||||
known_unet_keys: set[str] | None = None,
|
||||
) -> tuple[dict[str, str], str | None]:
|
||||
from safetensors import safe_open
|
||||
|
||||
with safe_open(path, framework="pt", device="cpu") as handle:
|
||||
all_keys = list(handle.keys())
|
||||
|
||||
prefix = infer_unet_prefix_from_keys(all_keys)
|
||||
if known_unet_keys is None:
|
||||
selected = {key: key[len(prefix):] for key in all_keys if key.startswith(prefix)}
|
||||
return (selected or {key: key for key in all_keys}), prefix
|
||||
|
||||
prefixes_to_try: list[str] = []
|
||||
for candidate in (prefix, *UNET_PREFIX_CANDIDATES):
|
||||
if candidate and candidate not in prefixes_to_try:
|
||||
prefixes_to_try.append(candidate)
|
||||
|
||||
selected: dict[str, str] = {}
|
||||
for key in all_keys:
|
||||
if key in known_unet_keys:
|
||||
selected[key] = key
|
||||
continue
|
||||
for candidate in prefixes_to_try:
|
||||
if key.startswith(candidate):
|
||||
stripped = key[len(candidate):]
|
||||
if stripped in known_unet_keys:
|
||||
selected[key] = stripped
|
||||
break
|
||||
return selected, prefix
|
||||
|
||||
|
||||
def _load_matching_extra_unet_state_dict(
|
||||
extra_state_dict_path: str,
|
||||
*,
|
||||
requested_device: torch.device,
|
||||
known_unet_keys: set[str],
|
||||
) -> dict[str, Any]:
|
||||
if _is_safetensors_path(extra_state_dict_path):
|
||||
selected_keys, _ = _selected_unet_key_map_from_header(
|
||||
extra_state_dict_path,
|
||||
known_unet_keys=known_unet_keys,
|
||||
)
|
||||
extra_sd, _, _, _ = load_safetensors_state_dict(
|
||||
extra_state_dict_path,
|
||||
requested_device,
|
||||
selected_keys=selected_keys,
|
||||
)
|
||||
return extra_sd
|
||||
|
||||
extra_sd = comfy.utils.load_torch_file(extra_state_dict_path)
|
||||
return _extract_unet_state_dict(extra_sd, known_unet_keys=known_unet_keys)
|
||||
|
||||
|
||||
def _extract_unet_state_dict(
|
||||
sd: dict[str, Any],
|
||||
*,
|
||||
diffusion_model_prefix: str | None = None,
|
||||
known_unet_keys: set[str] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
if known_unet_keys is None:
|
||||
if diffusion_model_prefix is None:
|
||||
diffusion_model_prefix = comfy.sd.model_detection.unet_prefix_from_state_dict(sd)
|
||||
if diffusion_model_prefix:
|
||||
prefix_len = len(diffusion_model_prefix)
|
||||
extracted = {key[prefix_len:]: value for key, value in sd.items() if key.startswith(diffusion_model_prefix)}
|
||||
if extracted:
|
||||
return extracted
|
||||
return dict(sd)
|
||||
|
||||
prefixes_to_try: list[str] = []
|
||||
|
||||
def add_prefix(prefix: str | None) -> None:
|
||||
if prefix and prefix not in prefixes_to_try:
|
||||
prefixes_to_try.append(prefix)
|
||||
|
||||
add_prefix(diffusion_model_prefix)
|
||||
add_prefix(comfy.sd.model_detection.unet_prefix_from_state_dict(sd))
|
||||
for prefix in UNET_PREFIX_CANDIDATES:
|
||||
add_prefix(prefix)
|
||||
|
||||
extracted: dict[str, Any] = {}
|
||||
for key, value in sd.items():
|
||||
if key in known_unet_keys:
|
||||
extracted[key] = value
|
||||
continue
|
||||
for prefix in prefixes_to_try:
|
||||
if not key.startswith(prefix):
|
||||
continue
|
||||
stripped = key[len(prefix):]
|
||||
if stripped in known_unet_keys:
|
||||
extracted[stripped] = value
|
||||
break
|
||||
return extracted
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _apply_policy_override(policy_override: str | None):
|
||||
if policy_override is None:
|
||||
yield
|
||||
return
|
||||
|
||||
previous_policy = REGISTRY.get_policy()
|
||||
if previous_policy == policy_override:
|
||||
yield
|
||||
return
|
||||
|
||||
REGISTRY.set_policy(policy_override)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
REGISTRY.set_policy(previous_policy)
|
||||
|
||||
|
||||
def _bind_model_for_reuse(model, *, source_path: str, note: str, loader_key: str) -> None:
|
||||
REGISTRY.bind_object(model, source_path=source_path, kind=KIND_MODEL, note=note, loader_key=loader_key)
|
||||
|
||||
|
||||
def _bind_clip_for_reuse(clip, *, source_path: str, note: str, loader_key: str) -> None:
|
||||
if clip is not None and getattr(clip, "patcher", None) is not None:
|
||||
REGISTRY.bind_object(
|
||||
clip.patcher,
|
||||
source_path=source_path,
|
||||
kind=KIND_CLIP,
|
||||
note=note,
|
||||
loader_key=loader_key,
|
||||
reusable_obj=clip,
|
||||
)
|
||||
|
||||
|
||||
def _bind_vae_for_reuse(vae, *, source_path: str, note: str, loader_key: str) -> None:
|
||||
if vae is not None and getattr(vae, "patcher", None) is not None:
|
||||
REGISTRY.bind_object(
|
||||
vae.patcher,
|
||||
source_path=source_path,
|
||||
kind=KIND_VAE,
|
||||
note=note,
|
||||
loader_key=loader_key,
|
||||
reusable_obj=vae,
|
||||
)
|
||||
|
||||
|
||||
def _load_resident_diffusion_model(
|
||||
*,
|
||||
loader_name: str,
|
||||
cache_scope: str,
|
||||
source_path: str,
|
||||
note: str,
|
||||
weight_dtype: str,
|
||||
compute_dtype: str,
|
||||
patch_cublaslinear: bool,
|
||||
sage_attention: str,
|
||||
enable_fp16_accumulation: bool,
|
||||
extra_state_dict: str | None = None,
|
||||
policy_override: str | None = None,
|
||||
) -> Any:
|
||||
model_options = _build_model_options(weight_dtype)
|
||||
effective_policy = _effective_policy_name(policy_override)
|
||||
loader_key = _make_loader_key(
|
||||
cache_scope,
|
||||
component="model",
|
||||
weight_dtype=weight_dtype,
|
||||
compute_dtype=compute_dtype,
|
||||
patch_cublaslinear=patch_cublaslinear,
|
||||
sage_attention=sage_attention,
|
||||
enable_fp16_accumulation=enable_fp16_accumulation,
|
||||
extra_state_dict=extra_state_dict,
|
||||
policy=effective_policy,
|
||||
)
|
||||
reused_model = REGISTRY.lookup_live_object(kind=KIND_MODEL, source_path=source_path, loader_key=loader_key)
|
||||
if reused_model is not None:
|
||||
_LOG.info("%s: reusing live model for %s", loader_name, source_path)
|
||||
return reused_model
|
||||
|
||||
_warn_pickle_checkpoint_gpu_compatibility(loader_name, source_path)
|
||||
explicit_device = REGISTRY.explicit_load_device(kind=KIND_MODEL, source_path=source_path)
|
||||
|
||||
with _temporary_backend_flags(
|
||||
cublas=patch_cublaslinear,
|
||||
fp16_accumulation=enable_fp16_accumulation,
|
||||
):
|
||||
with REGISTRY.load_context(
|
||||
kind=KIND_MODEL,
|
||||
source_path=source_path,
|
||||
explicit_device=explicit_device,
|
||||
cache_key=loader_key,
|
||||
):
|
||||
sd, metadata = comfy.utils.load_torch_file(source_path, return_metadata=True)
|
||||
if not _is_safetensors_path(source_path):
|
||||
sd = _extract_unet_state_dict(sd)
|
||||
if extra_state_dict:
|
||||
requested_device = explicit_device if explicit_device is not None else torch.device("cpu")
|
||||
sd.update(
|
||||
_load_matching_extra_unet_state_dict(
|
||||
extra_state_dict,
|
||||
requested_device=requested_device,
|
||||
known_unet_keys=set(sd),
|
||||
)
|
||||
)
|
||||
|
||||
model = comfy.sd.load_diffusion_model_state_dict(sd, model_options=model_options, metadata=metadata)
|
||||
_apply_model_postload_options(model, compute_dtype=compute_dtype, sage_attention=sage_attention)
|
||||
|
||||
_bind_model_for_reuse(model, source_path=source_path, note=note, loader_key=loader_key)
|
||||
return model
|
||||
|
||||
|
||||
def _checkpoint_model_loader_key(
|
||||
*,
|
||||
weight_dtype: str,
|
||||
compute_dtype: str,
|
||||
patch_cublaslinear: bool,
|
||||
sage_attention: str,
|
||||
enable_fp16_accumulation: bool,
|
||||
policy_override: str | None,
|
||||
) -> str:
|
||||
return _make_loader_key(
|
||||
"checkpoint_model",
|
||||
component="model",
|
||||
weight_dtype=weight_dtype,
|
||||
compute_dtype=compute_dtype,
|
||||
patch_cublaslinear=patch_cublaslinear,
|
||||
sage_attention=sage_attention,
|
||||
enable_fp16_accumulation=enable_fp16_accumulation,
|
||||
extra_state_dict=None,
|
||||
policy=_effective_policy_name(policy_override),
|
||||
)
|
||||
|
||||
|
||||
def _checkpoint_component_loader_key(component: str, policy_override: str | None) -> str:
|
||||
return _make_loader_key(f"checkpoint_{component}", component=component, policy=_effective_policy_name(policy_override))
|
||||
|
||||
|
||||
def _load_checkpoint_clip_only(
|
||||
*,
|
||||
ckpt_path: str,
|
||||
policy_override: str | None,
|
||||
loader_name: str,
|
||||
):
|
||||
loader_key = _checkpoint_component_loader_key("clip", policy_override)
|
||||
reused_clip = REGISTRY.lookup_live_object(kind=KIND_CLIP, source_path=ckpt_path, loader_key=loader_key)
|
||||
if reused_clip is not None:
|
||||
_LOG.info("%s: reusing live CLIP for %s", loader_name, ckpt_path)
|
||||
return reused_clip
|
||||
|
||||
_warn_pickle_checkpoint_gpu_compatibility(loader_name, ckpt_path)
|
||||
is_safetensors = _is_safetensors_path(ckpt_path)
|
||||
header_info = checkpoint_component_info_from_header(ckpt_path) if is_safetensors else None
|
||||
model_config = None if header_info is None else header_info.get("model_config")
|
||||
explicit_device = REGISTRY.explicit_load_device(kind=KIND_CLIP, source_path=ckpt_path)
|
||||
with REGISTRY.load_context(kind=KIND_CLIP, source_path=ckpt_path, explicit_device=explicit_device, cache_key=loader_key):
|
||||
if is_safetensors and model_config is None:
|
||||
requested_device = explicit_device if explicit_device is not None else torch.device("cpu")
|
||||
sd, metadata, _, _ = load_safetensors_state_dict(
|
||||
ckpt_path,
|
||||
requested_device,
|
||||
return_metadata=True,
|
||||
selected_keys=None,
|
||||
)
|
||||
else:
|
||||
sd, metadata = comfy.utils.load_torch_file(ckpt_path, return_metadata=True)
|
||||
if model_config is None:
|
||||
_, clip, _, _ = comfy.sd.load_state_dict_guess_config(
|
||||
sd,
|
||||
output_vae=False,
|
||||
output_clip=True,
|
||||
output_model=False,
|
||||
embedding_directory=folder_paths.get_folder_paths("embeddings"),
|
||||
metadata=metadata,
|
||||
)
|
||||
else:
|
||||
scaled_fp8_list = []
|
||||
for key in list(sd.keys()):
|
||||
if key.endswith(".scaled_fp8"):
|
||||
scaled_fp8_list.append(key[:-len(".scaled_fp8")])
|
||||
|
||||
if scaled_fp8_list:
|
||||
clip_source_sd: dict[str, Any] = {}
|
||||
for key, value in sd.items():
|
||||
if any(key.startswith(prefix) for prefix in scaled_fp8_list):
|
||||
continue
|
||||
clip_source_sd[key] = value
|
||||
for prefix in scaled_fp8_list:
|
||||
quant_sd, _ = comfy.utils.convert_old_quants(sd, prefix, metadata=metadata or {})
|
||||
clip_source_sd.update(quant_sd)
|
||||
else:
|
||||
clip_source_sd = sd
|
||||
|
||||
clip_target = model_config.clip_target(state_dict=clip_source_sd)
|
||||
if clip_target is None:
|
||||
clip = None
|
||||
else:
|
||||
clip_sd = model_config.process_clip_state_dict(clip_source_sd)
|
||||
if len(clip_sd) == 0:
|
||||
_LOG.warning("%s: no CLIP/text encoder weights found in %s after selective checkpoint load", loader_name, ckpt_path)
|
||||
clip = None
|
||||
else:
|
||||
parameters = comfy.utils.calculate_parameters(clip_sd)
|
||||
clip = comfy.sd.CLIP(
|
||||
clip_target,
|
||||
embedding_directory=folder_paths.get_folder_paths("embeddings"),
|
||||
tokenizer_data=clip_sd,
|
||||
parameters=parameters,
|
||||
state_dict=clip_sd,
|
||||
model_options={},
|
||||
disable_dynamic=False,
|
||||
)
|
||||
_bind_clip_for_reuse(clip, source_path=ckpt_path, note="checkpoint clip", loader_key=loader_key)
|
||||
return clip
|
||||
|
||||
|
||||
def _load_checkpoint_vae_only(
|
||||
*,
|
||||
ckpt_path: str,
|
||||
policy_override: str | None,
|
||||
loader_name: str,
|
||||
):
|
||||
loader_key = _checkpoint_component_loader_key("vae", policy_override)
|
||||
reused_vae = REGISTRY.lookup_live_object(kind=KIND_VAE, source_path=ckpt_path, loader_key=loader_key)
|
||||
if reused_vae is not None:
|
||||
_LOG.info("%s: reusing live VAE for %s", loader_name, ckpt_path)
|
||||
return reused_vae
|
||||
|
||||
_warn_pickle_checkpoint_gpu_compatibility(loader_name, ckpt_path)
|
||||
is_safetensors = _is_safetensors_path(ckpt_path)
|
||||
header_info = checkpoint_component_info_from_header(ckpt_path) if is_safetensors else None
|
||||
model_config = None if header_info is None else header_info.get("model_config")
|
||||
explicit_device = REGISTRY.explicit_load_device(kind=KIND_VAE, source_path=ckpt_path)
|
||||
with REGISTRY.load_context(kind=KIND_VAE, source_path=ckpt_path, explicit_device=explicit_device, cache_key=loader_key):
|
||||
if is_safetensors and model_config is None:
|
||||
requested_device = explicit_device if explicit_device is not None else torch.device("cpu")
|
||||
sd, metadata, _, _ = load_safetensors_state_dict(
|
||||
ckpt_path,
|
||||
requested_device,
|
||||
return_metadata=True,
|
||||
selected_keys=None,
|
||||
)
|
||||
else:
|
||||
sd, metadata = comfy.utils.load_torch_file(ckpt_path, return_metadata=True)
|
||||
if model_config is None:
|
||||
_, _, vae, _ = comfy.sd.load_state_dict_guess_config(
|
||||
sd,
|
||||
output_vae=True,
|
||||
output_clip=False,
|
||||
output_model=False,
|
||||
embedding_directory=folder_paths.get_folder_paths("embeddings"),
|
||||
metadata=metadata,
|
||||
)
|
||||
else:
|
||||
vae_sd = comfy.utils.state_dict_prefix_replace(
|
||||
sd,
|
||||
{prefix: "" for prefix in model_config.vae_key_prefix},
|
||||
filter_keys=True,
|
||||
)
|
||||
vae_sd = model_config.process_vae_state_dict(vae_sd)
|
||||
if len(vae_sd) == 0:
|
||||
_LOG.warning("%s: no VAE weights found in %s after selective checkpoint load", loader_name, ckpt_path)
|
||||
vae = None
|
||||
else:
|
||||
vae = comfy.sd.VAE(sd=vae_sd, metadata=metadata)
|
||||
_bind_vae_for_reuse(vae, source_path=ckpt_path, note="checkpoint vae", loader_key=loader_key)
|
||||
return vae
|
||||
|
||||
|
||||
def _load_full_checkpoint(
|
||||
*,
|
||||
ckpt_path: str,
|
||||
weight_dtype: str,
|
||||
compute_dtype: str,
|
||||
patch_cublaslinear: bool,
|
||||
sage_attention: str,
|
||||
enable_fp16_accumulation: bool,
|
||||
policy_override: str | None,
|
||||
loader_name: str,
|
||||
):
|
||||
model_key = _checkpoint_model_loader_key(
|
||||
weight_dtype=weight_dtype,
|
||||
compute_dtype=compute_dtype,
|
||||
patch_cublaslinear=patch_cublaslinear,
|
||||
sage_attention=sage_attention,
|
||||
enable_fp16_accumulation=enable_fp16_accumulation,
|
||||
policy_override=policy_override,
|
||||
)
|
||||
clip_key = _checkpoint_component_loader_key("clip", policy_override)
|
||||
vae_key = _checkpoint_component_loader_key("vae", policy_override)
|
||||
|
||||
model = REGISTRY.lookup_live_object(kind=KIND_MODEL, source_path=ckpt_path, loader_key=model_key)
|
||||
clip = REGISTRY.lookup_live_object(kind=KIND_CLIP, source_path=ckpt_path, loader_key=clip_key)
|
||||
vae = REGISTRY.lookup_live_object(kind=KIND_VAE, source_path=ckpt_path, loader_key=vae_key)
|
||||
missing = [name for name, value in (("model", model), ("clip", clip), ("vae", vae)) if value is None]
|
||||
if not missing:
|
||||
_LOG.info("%s: reusing live checkpoint outputs for %s", loader_name, ckpt_path)
|
||||
return model, clip, vae
|
||||
|
||||
if model is None:
|
||||
model = _load_resident_diffusion_model(
|
||||
loader_name=loader_name,
|
||||
cache_scope="checkpoint_model",
|
||||
source_path=ckpt_path,
|
||||
note="checkpoint model",
|
||||
weight_dtype=weight_dtype,
|
||||
compute_dtype=compute_dtype,
|
||||
patch_cublaslinear=patch_cublaslinear,
|
||||
sage_attention=sage_attention,
|
||||
enable_fp16_accumulation=enable_fp16_accumulation,
|
||||
policy_override=policy_override,
|
||||
)
|
||||
if clip is None:
|
||||
clip = _load_checkpoint_clip_only(
|
||||
ckpt_path=ckpt_path,
|
||||
policy_override=policy_override,
|
||||
loader_name=loader_name,
|
||||
)
|
||||
if vae is None:
|
||||
vae = _load_checkpoint_vae_only(
|
||||
ckpt_path=ckpt_path,
|
||||
policy_override=policy_override,
|
||||
loader_name=loader_name,
|
||||
)
|
||||
return model, clip, vae
|
||||
|
||||
|
||||
class DiffusionModelSelectorResident:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
@@ -306,6 +814,13 @@ class DiffusionModelLoaderResident:
|
||||
"forceInput": True,
|
||||
"tooltip": "Optional absolute path to a second state dict merged into the main diffusion state dict before model detection.",
|
||||
},
|
||||
),
|
||||
"policy_override": (
|
||||
"STRING",
|
||||
{
|
||||
"forceInput": True,
|
||||
"tooltip": "Optional residency policy override. Connect Set Global Residency Policy here, not to extra_state_dict.",
|
||||
},
|
||||
)
|
||||
},
|
||||
}
|
||||
@@ -328,28 +843,30 @@ class DiffusionModelLoaderResident:
|
||||
sage_attention: str,
|
||||
enable_fp16_accumulation: bool,
|
||||
extra_state_dict: str | None = None,
|
||||
policy_override: str | None = None,
|
||||
):
|
||||
with _temporary_backend_flags(
|
||||
cublas=patch_cublaslinear,
|
||||
fp16_accumulation=enable_fp16_accumulation,
|
||||
):
|
||||
model_options = _build_model_options(weight_dtype)
|
||||
policy_override, extra_state_dict = _resolve_loader_policy_and_extra_state_dict(
|
||||
loader_name="Diffusion Model Loader Resident",
|
||||
policy_override=policy_override,
|
||||
extra_state_dict=extra_state_dict,
|
||||
)
|
||||
|
||||
with _apply_policy_override(policy_override):
|
||||
unet_path = folder_paths.get_full_path_or_raise("diffusion_models", model_name)
|
||||
explicit_device = REGISTRY.explicit_load_device(kind=KIND_MODEL, source_path=unet_path)
|
||||
|
||||
with REGISTRY.load_context(kind=KIND_MODEL, source_path=unet_path, explicit_device=explicit_device):
|
||||
sd, metadata = comfy.utils.load_torch_file(unet_path, return_metadata=True)
|
||||
if extra_state_dict:
|
||||
extra_sd = comfy.utils.load_torch_file(extra_state_dict)
|
||||
sd.update(extra_sd)
|
||||
del extra_sd
|
||||
|
||||
diffusion_model_prefix = comfy.sd.model_detection.unet_prefix_from_state_dict(sd)
|
||||
sd = comfy.utils.state_dict_prefix_replace(sd, {diffusion_model_prefix: ""}, filter_keys=False)
|
||||
model = comfy.sd.load_diffusion_model_state_dict(sd, model_options=model_options, metadata=metadata)
|
||||
_apply_model_postload_options(model, compute_dtype=compute_dtype, sage_attention=sage_attention)
|
||||
REGISTRY.bind_object(model, source_path=unet_path, kind=KIND_MODEL)
|
||||
return (model,)
|
||||
model = _load_resident_diffusion_model(
|
||||
loader_name="Diffusion Model Loader Resident",
|
||||
cache_scope="diffusion_model",
|
||||
source_path=unet_path,
|
||||
note="diffusion model",
|
||||
weight_dtype=weight_dtype,
|
||||
compute_dtype=compute_dtype,
|
||||
patch_cublaslinear=patch_cublaslinear,
|
||||
sage_attention=sage_attention,
|
||||
enable_fp16_accumulation=enable_fp16_accumulation,
|
||||
extra_state_dict=extra_state_dict,
|
||||
policy_override=policy_override,
|
||||
)
|
||||
return (model,)
|
||||
|
||||
|
||||
class CheckpointLoaderResident:
|
||||
@@ -378,7 +895,16 @@ class CheckpointLoaderResident:
|
||||
"BOOLEAN",
|
||||
{"default": False, "tooltip": "Set torch.backends.cuda.matmul.allow_fp16_accumulation."},
|
||||
),
|
||||
}
|
||||
},
|
||||
"optional": {
|
||||
"policy_override": (
|
||||
"STRING",
|
||||
{
|
||||
"forceInput": True,
|
||||
"tooltip": "Optional residency policy override. Connect Set Global Residency Policy here.",
|
||||
},
|
||||
)
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL", "CLIP", "VAE")
|
||||
@@ -386,7 +912,7 @@ class CheckpointLoaderResident:
|
||||
CATEGORY = "GPU Resident Loader/loaders"
|
||||
DESCRIPTION = (
|
||||
"Checkpoint loader with the KJ DiffusionModelLoader-style tuning knobs plus GPU-resident ingest. "
|
||||
"It loads the whole checkpoint, then binds the model, CLIP, and VAE into the residency registry."
|
||||
"It reuses live checkpoint components when possible and composes missing outputs from the model, CLIP, and VAE loaders instead of materializing a broad checkpoint state dict."
|
||||
)
|
||||
|
||||
def load(
|
||||
@@ -397,30 +923,139 @@ class CheckpointLoaderResident:
|
||||
patch_cublaslinear: bool,
|
||||
sage_attention: str,
|
||||
enable_fp16_accumulation: bool,
|
||||
policy_override: str | None = None,
|
||||
):
|
||||
with _temporary_backend_flags(
|
||||
cublas=patch_cublaslinear,
|
||||
fp16_accumulation=enable_fp16_accumulation,
|
||||
):
|
||||
model_options = _build_model_options(weight_dtype)
|
||||
policy_override, _ = _resolve_loader_policy_and_extra_state_dict(
|
||||
loader_name="Checkpoint Loader Resident",
|
||||
policy_override=policy_override,
|
||||
extra_state_dict=None,
|
||||
)
|
||||
|
||||
with _apply_policy_override(policy_override):
|
||||
ckpt_path = folder_paths.get_full_path_or_raise("checkpoints", ckpt_name)
|
||||
explicit_device = REGISTRY.explicit_load_device(kind=KIND_CHECKPOINT, source_path=ckpt_path)
|
||||
|
||||
with REGISTRY.load_context(kind=KIND_CHECKPOINT, source_path=ckpt_path, explicit_device=explicit_device):
|
||||
sd, metadata = comfy.utils.load_torch_file(ckpt_path, return_metadata=True)
|
||||
|
||||
model, clip, vae, _ = comfy.sd.load_state_dict_guess_config(
|
||||
sd,
|
||||
output_vae=True,
|
||||
output_clip=True,
|
||||
embedding_directory=folder_paths.get_folder_paths("embeddings"),
|
||||
metadata=metadata,
|
||||
model_options=model_options,
|
||||
model, clip, vae = _load_full_checkpoint(
|
||||
ckpt_path=ckpt_path,
|
||||
weight_dtype=weight_dtype,
|
||||
compute_dtype=compute_dtype,
|
||||
patch_cublaslinear=patch_cublaslinear,
|
||||
sage_attention=sage_attention,
|
||||
enable_fp16_accumulation=enable_fp16_accumulation,
|
||||
policy_override=policy_override,
|
||||
loader_name="Checkpoint Loader Resident",
|
||||
)
|
||||
_apply_model_postload_options(model, compute_dtype=compute_dtype, sage_attention=sage_attention)
|
||||
REGISTRY.bind_object(model, source_path=ckpt_path, kind=KIND_MODEL, note="checkpoint model")
|
||||
if clip is not None and getattr(clip, "patcher", None) is not None:
|
||||
REGISTRY.bind_object(clip.patcher, source_path=ckpt_path, kind=KIND_CLIP, note="checkpoint clip")
|
||||
if vae is not None and getattr(vae, "patcher", None) is not None:
|
||||
REGISTRY.bind_object(vae.patcher, source_path=ckpt_path, kind=KIND_VAE, note="checkpoint vae")
|
||||
return model, clip, vae
|
||||
return model, clip, vae
|
||||
|
||||
|
||||
class CheckpointModelLoaderResident:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return CheckpointLoaderResident.INPUT_TYPES()
|
||||
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
FUNCTION = "load"
|
||||
CATEGORY = "GPU Resident Loader/loaders"
|
||||
DESCRIPTION = (
|
||||
"Checkpoint model-only loader. Safetensors checkpoints take the selective UNet fast path, and live equivalent models are reused when available."
|
||||
)
|
||||
|
||||
def load(
|
||||
self,
|
||||
ckpt_name: str,
|
||||
weight_dtype: str,
|
||||
compute_dtype: str,
|
||||
patch_cublaslinear: bool,
|
||||
sage_attention: str,
|
||||
enable_fp16_accumulation: bool,
|
||||
policy_override: str | None = None,
|
||||
):
|
||||
policy_override, _ = _resolve_loader_policy_and_extra_state_dict(
|
||||
loader_name="Checkpoint Model Loader Resident",
|
||||
policy_override=policy_override,
|
||||
extra_state_dict=None,
|
||||
)
|
||||
|
||||
with _apply_policy_override(policy_override):
|
||||
ckpt_path = folder_paths.get_full_path_or_raise("checkpoints", ckpt_name)
|
||||
model = _load_resident_diffusion_model(
|
||||
loader_name="Checkpoint Model Loader Resident",
|
||||
cache_scope="checkpoint_model",
|
||||
source_path=ckpt_path,
|
||||
note="checkpoint model",
|
||||
weight_dtype=weight_dtype,
|
||||
compute_dtype=compute_dtype,
|
||||
patch_cublaslinear=patch_cublaslinear,
|
||||
sage_attention=sage_attention,
|
||||
enable_fp16_accumulation=enable_fp16_accumulation,
|
||||
policy_override=policy_override,
|
||||
)
|
||||
return (model,)
|
||||
|
||||
|
||||
class CheckpointClipLoaderResident:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"ckpt_name": (
|
||||
folder_paths.get_filename_list("checkpoints"),
|
||||
{"tooltip": "Checkpoint file to load."},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"policy_override": (
|
||||
"STRING",
|
||||
{
|
||||
"forceInput": True,
|
||||
"tooltip": "Optional residency policy override. Connect Set Global Residency Policy here.",
|
||||
},
|
||||
)
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CLIP",)
|
||||
FUNCTION = "load"
|
||||
CATEGORY = "GPU Resident Loader/loaders"
|
||||
DESCRIPTION = "Checkpoint CLIP-only loader with live-object reuse. It avoids rebuilding the text encoder when an equivalent CLIP object is already alive."
|
||||
|
||||
def load(self, ckpt_name: str, policy_override: str | None = None):
|
||||
policy_override, _ = _resolve_loader_policy_and_extra_state_dict(
|
||||
loader_name="Checkpoint Clip Loader Resident",
|
||||
policy_override=policy_override,
|
||||
extra_state_dict=None,
|
||||
)
|
||||
|
||||
with _apply_policy_override(policy_override):
|
||||
ckpt_path = folder_paths.get_full_path_or_raise("checkpoints", ckpt_name)
|
||||
clip = _load_checkpoint_clip_only(
|
||||
ckpt_path=ckpt_path,
|
||||
policy_override=policy_override,
|
||||
loader_name="Checkpoint Clip Loader Resident",
|
||||
)
|
||||
return (clip,)
|
||||
|
||||
|
||||
class CheckpointVAELoaderResident:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return CheckpointClipLoaderResident.INPUT_TYPES()
|
||||
|
||||
RETURN_TYPES = ("VAE",)
|
||||
FUNCTION = "load"
|
||||
CATEGORY = "GPU Resident Loader/loaders"
|
||||
DESCRIPTION = "Checkpoint VAE-only loader with live-object reuse. It avoids rebuilding the VAE when an equivalent object is still alive."
|
||||
|
||||
def load(self, ckpt_name: str, policy_override: str | None = None):
|
||||
policy_override, _ = _resolve_loader_policy_and_extra_state_dict(
|
||||
loader_name="Checkpoint VAE Loader Resident",
|
||||
policy_override=policy_override,
|
||||
extra_state_dict=None,
|
||||
)
|
||||
|
||||
with _apply_policy_override(policy_override):
|
||||
ckpt_path = folder_paths.get_full_path_or_raise("checkpoints", ckpt_name)
|
||||
vae = _load_checkpoint_vae_only(
|
||||
ckpt_path=ckpt_path,
|
||||
policy_override=policy_override,
|
||||
loader_name="Checkpoint VAE Loader Resident",
|
||||
)
|
||||
return (vae,)
|
||||
|
||||
@@ -5,7 +5,15 @@ from typing import Any
|
||||
|
||||
import comfy.model_management as model_management
|
||||
|
||||
from .kj_loader import CheckpointLoaderResident, DiffusionModelLoaderResident, DiffusionModelSelectorResident
|
||||
from .cleanup import trim_resident_vram
|
||||
from .kj_loader import (
|
||||
CheckpointClipLoaderResident,
|
||||
CheckpointLoaderResident,
|
||||
CheckpointModelLoaderResident,
|
||||
CheckpointVAELoaderResident,
|
||||
DiffusionModelLoaderResident,
|
||||
DiffusionModelSelectorResident,
|
||||
)
|
||||
from .residency import REGISTRY
|
||||
|
||||
|
||||
@@ -324,10 +332,61 @@ class ReportVAEResidency:
|
||||
return (_entry_report_json(_patcher_for_vae(vae)),)
|
||||
|
||||
|
||||
class TrimResidentVRAM:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"target_free_vram_mb": ("INT", {"default": 4096, "min": 0, "max": 1_048_576, "step": 1}),
|
||||
"respect_sticky": ("BOOLEAN", {"default": True}),
|
||||
"sticky_floor_priority": ("INT", {"default": 0, "min": -100, "max": 1000, "step": 1}),
|
||||
"allow_partial_unload": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
"optional": {
|
||||
"payload": (
|
||||
"*",
|
||||
{
|
||||
"forceInput": True,
|
||||
"tooltip": "Optional passthrough payload so the trim can sit on a stage boundary without changing graph data.",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("*", "STRING")
|
||||
RETURN_NAMES = ("payload", "trim_report")
|
||||
FUNCTION = "trim"
|
||||
CATEGORY = "GPU Resident Loader/residency"
|
||||
DESCRIPTION = (
|
||||
"Resident-aware VRAM trimmer. It frees only enough VRAM to reach the target budget, "
|
||||
"tries partial unload before full detach, orders sticky entries behind non-sticky work unless their priority falls below the floor, "
|
||||
"and returns a JSON report describing whether the target was met, only partially met, or failed."
|
||||
)
|
||||
|
||||
def trim(
|
||||
self,
|
||||
target_free_vram_mb: int,
|
||||
respect_sticky: bool,
|
||||
sticky_floor_priority: int,
|
||||
allow_partial_unload: bool,
|
||||
payload=None,
|
||||
):
|
||||
report = trim_resident_vram(
|
||||
target_free_vram_bytes=max(0, int(target_free_vram_mb)) * 1024 * 1024,
|
||||
respect_sticky=respect_sticky,
|
||||
sticky_floor_priority=sticky_floor_priority,
|
||||
allow_partial_unload=allow_partial_unload,
|
||||
)
|
||||
return payload, json.dumps(report, indent=2, sort_keys=True)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"DiffusionModelSelectorResident": DiffusionModelSelectorResident,
|
||||
"DiffusionModelLoaderResident": DiffusionModelLoaderResident,
|
||||
"CheckpointLoaderResident": CheckpointLoaderResident,
|
||||
"CheckpointModelLoaderResident": CheckpointModelLoaderResident,
|
||||
"CheckpointClipLoaderResident": CheckpointClipLoaderResident,
|
||||
"CheckpointVAELoaderResident": CheckpointVAELoaderResident,
|
||||
"SetGlobalResidencyPolicy": SetGlobalResidencyPolicy,
|
||||
"RegistrySnapshot": RegistrySnapshot,
|
||||
"PinModelResidency": PinModelResidency,
|
||||
@@ -342,12 +401,16 @@ NODE_CLASS_MAPPINGS = {
|
||||
"ReportModelResidency": ReportModelResidency,
|
||||
"ReportClipResidency": ReportClipResidency,
|
||||
"ReportVAEResidency": ReportVAEResidency,
|
||||
"TrimResidentVRAM": TrimResidentVRAM,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"DiffusionModelSelectorResident": "Diffusion Model Selector Resident",
|
||||
"DiffusionModelLoaderResident": "Diffusion Model Loader Resident",
|
||||
"CheckpointLoaderResident": "Checkpoint Loader Resident",
|
||||
"CheckpointModelLoaderResident": "Checkpoint Model Loader Resident",
|
||||
"CheckpointClipLoaderResident": "Checkpoint Clip Loader Resident",
|
||||
"CheckpointVAELoaderResident": "Checkpoint VAE Loader Resident",
|
||||
"SetGlobalResidencyPolicy": "Set Global Residency Policy",
|
||||
"RegistrySnapshot": "Registry Snapshot",
|
||||
"PinModelResidency": "Pin Model Residency",
|
||||
@@ -362,4 +425,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"ReportModelResidency": "Report Model Residency",
|
||||
"ReportClipResidency": "Report CLIP Residency",
|
||||
"ReportVAEResidency": "Report VAE Residency",
|
||||
"TrimResidentVRAM": "Trim Resident VRAM",
|
||||
}
|
||||
|
||||
+315
-73
@@ -1,8 +1,10 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import functools
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import struct
|
||||
from typing import Any, Callable
|
||||
|
||||
import torch
|
||||
@@ -22,6 +24,121 @@ from .residency import (
|
||||
_LOG = logging.getLogger(__name__)
|
||||
_PATCHED = False
|
||||
_ORIGINALS: dict[str, Callable[..., Any]] = {}
|
||||
_METADATA_CPU_KEY_SUFFIXES = ("spiece_model", "tekken_model", "comfy_quant")
|
||||
_UNET_PREFIX_CANDIDATES = ("model.diffusion_model.", "model.model.", "net.")
|
||||
_WARNED_PICKLE_GPU_PATHS: set[str] = set()
|
||||
_SAFE_TENSORS_COMPONENT_CACHE_MAX = 32
|
||||
_SAFETENSORS_DTYPE_MAP = {
|
||||
"BOOL": torch.bool,
|
||||
"U8": torch.uint8,
|
||||
"I8": torch.int8,
|
||||
"I16": torch.int16,
|
||||
"U16": getattr(torch, "uint16", torch.int32),
|
||||
"I32": torch.int32,
|
||||
"U32": getattr(torch, "uint32", torch.int64),
|
||||
"I64": torch.int64,
|
||||
"U64": getattr(torch, "uint64", torch.int64),
|
||||
"F16": torch.float16,
|
||||
"BF16": torch.bfloat16,
|
||||
"F32": torch.float32,
|
||||
"F64": torch.float64,
|
||||
"F8_E4M3FN": getattr(torch, "float8_e4m3fn", torch.float16),
|
||||
"F8_E5M2": getattr(torch, "float8_e5m2", torch.float16),
|
||||
}
|
||||
|
||||
|
||||
def _safetensors_header_cache_key(path: str) -> tuple[str, int, int]:
|
||||
stat = os.stat(path)
|
||||
return (os.path.abspath(path), stat.st_mtime_ns, stat.st_size)
|
||||
|
||||
|
||||
def _read_safetensors_header(path: str) -> tuple[dict[str, dict[str, Any]], dict[str, str] | None]:
|
||||
with open(path, "rb") as handle:
|
||||
header_size = struct.unpack("<Q", handle.read(8))[0]
|
||||
header = json.loads(handle.read(header_size))
|
||||
|
||||
metadata = header.get("__metadata__")
|
||||
tensor_headers = {
|
||||
key: value
|
||||
for key, value in header.items()
|
||||
if key != "__metadata__" and isinstance(value, dict)
|
||||
}
|
||||
return tensor_headers, metadata if isinstance(metadata, dict) else None
|
||||
|
||||
|
||||
def _torch_dtype_from_safetensors_code(code: str | None) -> torch.dtype:
|
||||
if code is None:
|
||||
return torch.float32
|
||||
return _SAFETENSORS_DTYPE_MAP.get(code, torch.float32)
|
||||
|
||||
|
||||
def _build_meta_state_dict_from_header(tensor_headers: dict[str, dict[str, Any]]) -> dict[str, torch.Tensor]:
|
||||
meta_state_dict: dict[str, torch.Tensor] = {}
|
||||
for key, tensor_info in tensor_headers.items():
|
||||
shape = tuple(int(dim) for dim in tensor_info.get("shape", ()))
|
||||
meta_state_dict[key] = torch.empty(
|
||||
shape,
|
||||
dtype=_torch_dtype_from_safetensors_code(tensor_info.get("dtype")),
|
||||
device="meta",
|
||||
)
|
||||
return meta_state_dict
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=_SAFE_TENSORS_COMPONENT_CACHE_MAX)
|
||||
def _cached_component_key_maps(cache_key: tuple[str, int, int]) -> dict[str, Any]:
|
||||
import comfy.sd as comfy_sd
|
||||
|
||||
path = cache_key[0]
|
||||
tensor_headers, metadata = _read_safetensors_header(path)
|
||||
all_keys = tuple(tensor_headers)
|
||||
unet_prefix = infer_unet_prefix_from_keys(all_keys)
|
||||
meta_state_dict = _build_meta_state_dict_from_header(tensor_headers)
|
||||
model_config = comfy_sd.model_detection.model_config_from_unet(meta_state_dict, unet_prefix, metadata=metadata)
|
||||
|
||||
def select_prefixed(prefixes: tuple[str, ...] | list[str] | None) -> tuple[tuple[str, str], ...]:
|
||||
if not prefixes:
|
||||
return ()
|
||||
return tuple(
|
||||
(key, key)
|
||||
for key in all_keys
|
||||
if any(key.startswith(prefix) for prefix in prefixes)
|
||||
)
|
||||
|
||||
return {
|
||||
"metadata": metadata,
|
||||
"unet_prefix": unet_prefix,
|
||||
"model_config": model_config,
|
||||
"model": tuple((key, key[len(unet_prefix):]) for key in all_keys if key.startswith(unet_prefix)),
|
||||
"clip": select_prefixed(getattr(model_config, "text_encoder_key_prefix", None) or ()),
|
||||
"vae": select_prefixed(getattr(model_config, "vae_key_prefix", None) or ()),
|
||||
}
|
||||
|
||||
|
||||
def checkpoint_component_info_from_header(path: str) -> dict[str, Any] | None:
|
||||
try:
|
||||
return _cached_component_key_maps(_safetensors_header_cache_key(path))
|
||||
except Exception as exc:
|
||||
_LOG.warning("GPU Resident Loader: failed to build selective safetensors header map for %s: %s", path, exc)
|
||||
return None
|
||||
|
||||
|
||||
def _selected_component_keys_from_header(path: str, kind: str) -> dict[str, str] | None:
|
||||
if kind not in {KIND_MODEL, KIND_CLIP, KIND_VAE}:
|
||||
return None
|
||||
component_maps = checkpoint_component_info_from_header(path)
|
||||
if component_maps is None:
|
||||
return None
|
||||
|
||||
pairs = component_maps.get(kind, ())
|
||||
return dict(pairs) if pairs else None
|
||||
|
||||
|
||||
def _selected_component_suffix(kind: str | None) -> str | None:
|
||||
return {
|
||||
KIND_MODEL: "model_only",
|
||||
KIND_CLIP: "clip_only",
|
||||
KIND_VAE: "vae_only",
|
||||
}.get(kind)
|
||||
|
||||
|
||||
def _normalize_device(device: Any | None) -> torch.device | None:
|
||||
@@ -60,6 +177,130 @@ def _copy_tensor_if_needed(
|
||||
return tensor.to(device=target_device, copy=True)
|
||||
|
||||
|
||||
def _tensor_key_requires_cpu(key: str) -> bool:
|
||||
return key.endswith(_METADATA_CPU_KEY_SUFFIXES)
|
||||
|
||||
|
||||
def _prepare_loaded_tensor(
|
||||
key: str,
|
||||
tensor: torch.Tensor,
|
||||
requested_device: torch.device,
|
||||
*,
|
||||
disable_mmap: bool,
|
||||
move_to_requested_device: bool = False,
|
||||
) -> torch.Tensor:
|
||||
if _tensor_key_requires_cpu(key):
|
||||
return _copy_tensor_if_needed(tensor, torch.device("cpu"))
|
||||
|
||||
if move_to_requested_device and tensor.device != requested_device:
|
||||
return _copy_tensor_if_needed(tensor, requested_device)
|
||||
|
||||
if disable_mmap and tensor.device.type == "cpu":
|
||||
return _copy_tensor_if_needed(tensor, requested_device, force_copy=True)
|
||||
|
||||
return tensor
|
||||
|
||||
|
||||
def _state_dict_device_summary(sd: dict[str, Any], requested_device: torch.device) -> str:
|
||||
devices: set[str] = set()
|
||||
for value in sd.values():
|
||||
if torch.is_tensor(value):
|
||||
devices.add(str(value.device))
|
||||
if not devices:
|
||||
return str(requested_device)
|
||||
if len(devices) == 1:
|
||||
return next(iter(devices))
|
||||
return ", ".join(sorted(devices))
|
||||
|
||||
|
||||
def _device_summary_from_observed(observed_devices: set[str], requested_device: torch.device) -> str:
|
||||
if not observed_devices:
|
||||
return str(requested_device)
|
||||
if len(observed_devices) == 1:
|
||||
return next(iter(observed_devices))
|
||||
return ", ".join(sorted(observed_devices))
|
||||
|
||||
|
||||
def infer_unet_prefix_from_keys(keys: list[str] | tuple[str, ...]) -> str:
|
||||
counts = {candidate: 0 for candidate in _UNET_PREFIX_CANDIDATES}
|
||||
for key in keys:
|
||||
for candidate in _UNET_PREFIX_CANDIDATES:
|
||||
if key.startswith(candidate):
|
||||
counts[candidate] += 1
|
||||
break
|
||||
top = max(counts, key=counts.get)
|
||||
return top if counts[top] > 5 else "model."
|
||||
|
||||
|
||||
def load_safetensors_state_dict(
|
||||
ckpt: str,
|
||||
requested_device: torch.device,
|
||||
*,
|
||||
return_metadata: bool = False,
|
||||
selected_keys: dict[str, str] | None = None,
|
||||
) -> tuple[dict[str, Any], Any, str, str]:
|
||||
import comfy.memory_management
|
||||
import comfy.utils as comfy_utils
|
||||
|
||||
metadata = None
|
||||
if comfy.memory_management.aimdo_enabled and requested_device.type == "cpu" and selected_keys is None:
|
||||
sd, metadata = comfy_utils.load_safetensors(ckpt)
|
||||
if not return_metadata:
|
||||
metadata = None
|
||||
return sd, metadata, "cpu", "aimdo_cpu"
|
||||
|
||||
disable_mmap = getattr(comfy_utils, "DISABLE_MMAP", False)
|
||||
|
||||
def read_handle(device_arg: Any, *, move_to_requested_device: bool) -> tuple[dict[str, Any], Any, str]:
|
||||
observed_devices: set[str] = set()
|
||||
with safe_open(ckpt, framework="pt", device=device_arg) as handle:
|
||||
key_map = selected_keys if selected_keys is not None else {key: key for key in handle.keys()}
|
||||
sd: dict[str, Any] = {}
|
||||
for source_key, target_key in key_map.items():
|
||||
tensor = handle.get_tensor(source_key)
|
||||
loaded = _prepare_loaded_tensor(
|
||||
source_key,
|
||||
tensor,
|
||||
requested_device,
|
||||
disable_mmap=disable_mmap,
|
||||
move_to_requested_device=move_to_requested_device,
|
||||
)
|
||||
sd[target_key] = loaded
|
||||
if torch.is_tensor(loaded):
|
||||
observed_devices.add(str(loaded.device))
|
||||
handle_metadata = handle.metadata() if return_metadata else None
|
||||
return sd, handle_metadata, _device_summary_from_observed(observed_devices, requested_device)
|
||||
|
||||
try:
|
||||
safe_device = _safe_open_device_arg(requested_device)
|
||||
sd, metadata, actual_device = read_handle(safe_device, move_to_requested_device=False)
|
||||
return sd, metadata, actual_device, "direct"
|
||||
except Exception as exc:
|
||||
if requested_device.type == "cuda":
|
||||
_LOG.warning(
|
||||
"GPU Resident Loader: direct GPU safetensors load failed for %s; falling back to CPU path: %s",
|
||||
ckpt,
|
||||
exc,
|
||||
)
|
||||
try:
|
||||
sd, metadata, actual_device = read_handle("cpu", move_to_requested_device=True)
|
||||
return sd, metadata, actual_device, "cpu_then_copy"
|
||||
except Exception as fallback_exc:
|
||||
raise fallback_exc from exc
|
||||
raise
|
||||
|
||||
|
||||
def _warn_pickle_gpu_compatibility(path: str, requested_device: torch.device) -> None:
|
||||
if requested_device.type != "cuda" or path in _WARNED_PICKLE_GPU_PATHS:
|
||||
return
|
||||
_WARNED_PICKLE_GPU_PATHS.add(path)
|
||||
_LOG.warning(
|
||||
"GPU Resident Loader: %s is not a safetensors file, so GPU-resident loading still goes through CPU-first torch.load(). "
|
||||
"Convert hot models with scripts/convert_checkpoint_to_safetensors.py for the narrow fast path.",
|
||||
path,
|
||||
)
|
||||
|
||||
|
||||
def _resolved_context(kind: str, source_path: str | None) -> tuple[torch.device | None, str, str | None]:
|
||||
ctx = REGISTRY.current_context()
|
||||
if ctx is not None and ctx.explicit_device is not None:
|
||||
@@ -105,32 +346,28 @@ def _patched_load_torch_file(ckpt, safe_load=False, device=None, return_metadata
|
||||
|
||||
if lowered.endswith((".safetensors", ".sft")):
|
||||
try:
|
||||
if comfy.memory_management.aimdo_enabled and requested_device.type == "cpu":
|
||||
sd, metadata = comfy_utils.load_safetensors(ckpt)
|
||||
selected_keys = None
|
||||
selected_suffix = None
|
||||
if ctx is not None and ctx.kind in {KIND_MODEL, KIND_CLIP, KIND_VAE}:
|
||||
selected_keys = _selected_component_keys_from_header(ckpt, ctx.kind)
|
||||
selected_suffix = _selected_component_suffix(ctx.kind)
|
||||
|
||||
sd, metadata, actual_device, load_mode = load_safetensors_state_dict(
|
||||
ckpt,
|
||||
requested_device,
|
||||
return_metadata=return_metadata,
|
||||
selected_keys=selected_keys,
|
||||
)
|
||||
if load_mode == "aimdo_cpu":
|
||||
method = "safetensors_aimdo_cpu"
|
||||
if not return_metadata:
|
||||
metadata = None
|
||||
_record_generic_load(
|
||||
path=ckpt,
|
||||
method=method,
|
||||
requested_device=requested_device,
|
||||
actual_device="cpu",
|
||||
elif selected_keys is not None and selected_suffix is not None:
|
||||
method = f"safetensors_cpu_then_copy_to_cuda_{selected_suffix}" if load_mode == "cpu_then_copy" else (
|
||||
f"safetensors_gpu_direct_{selected_suffix}" if requested_device.type == "cuda" else f"safetensors_cpu_{selected_suffix}"
|
||||
)
|
||||
else:
|
||||
method = "safetensors_cpu_then_copy_to_cuda" if load_mode == "cpu_then_copy" else (
|
||||
"safetensors_gpu_direct" if requested_device.type == "cuda" else "safetensors_cpu"
|
||||
)
|
||||
return (sd, metadata) if return_metadata else sd
|
||||
|
||||
safe_device = _safe_open_device_arg(requested_device)
|
||||
with safe_open(ckpt, framework="pt", device=safe_device) as handle:
|
||||
sd = {}
|
||||
for key in handle.keys():
|
||||
tensor = handle.get_tensor(key)
|
||||
if getattr(comfy_utils, "DISABLE_MMAP", False) and tensor.device.type == "cpu":
|
||||
tensor = _copy_tensor_if_needed(tensor, requested_device, force_copy=True)
|
||||
sd[key] = tensor
|
||||
if return_metadata:
|
||||
metadata = handle.metadata()
|
||||
|
||||
actual_device = str(next(iter(sd.values())).device) if sd else str(requested_device)
|
||||
method = "safetensors_gpu_direct" if requested_device.type == "cuda" else "safetensors_cpu"
|
||||
_record_generic_load(
|
||||
path=ckpt,
|
||||
method=method,
|
||||
@@ -139,37 +376,6 @@ def _patched_load_torch_file(ckpt, safe_load=False, device=None, return_metadata
|
||||
)
|
||||
return (sd, metadata) if return_metadata else sd
|
||||
except Exception as exc:
|
||||
if requested_device.type == "cuda":
|
||||
_LOG.warning(
|
||||
"GPU Resident Loader: direct GPU safetensors load failed for %s; falling back to CPU path: %s",
|
||||
ckpt,
|
||||
exc,
|
||||
)
|
||||
try:
|
||||
with safe_open(ckpt, framework="pt", device="cpu") as handle:
|
||||
sd = {}
|
||||
for key in handle.keys():
|
||||
sd[key] = handle.get_tensor(key).to(requested_device)
|
||||
if return_metadata:
|
||||
metadata = handle.metadata()
|
||||
_record_generic_load(
|
||||
path=ckpt,
|
||||
method="safetensors_cpu_then_copy_to_cuda",
|
||||
requested_device=requested_device,
|
||||
actual_device=str(requested_device),
|
||||
error=str(exc),
|
||||
)
|
||||
return (sd, metadata) if return_metadata else sd
|
||||
except Exception as fallback_exc:
|
||||
_record_generic_load(
|
||||
path=ckpt,
|
||||
method="safetensors_cpu_fallback_failed",
|
||||
requested_device=requested_device,
|
||||
actual_device="error",
|
||||
error=str(fallback_exc),
|
||||
)
|
||||
raise fallback_exc from exc
|
||||
|
||||
if len(getattr(exc, "args", ())) > 0:
|
||||
message = exc.args[0]
|
||||
if isinstance(message, str):
|
||||
@@ -198,15 +404,9 @@ def _patched_load_torch_file(ckpt, safe_load=False, device=None, return_metadata
|
||||
if getattr(comfy_utils, "MMAP_TORCH_FILES", False):
|
||||
torch_args["mmap"] = True
|
||||
|
||||
pl_sd = torch.load(ckpt, map_location=requested_device, weights_only=True, **torch_args)
|
||||
method = "torch_load_cpu_first_to_cuda" if requested_device.type == "cuda" else "torch_load_cpu"
|
||||
_record_generic_load(
|
||||
path=ckpt,
|
||||
method=method,
|
||||
requested_device=requested_device,
|
||||
actual_device=str(requested_device),
|
||||
)
|
||||
|
||||
_warn_pickle_gpu_compatibility(ckpt, requested_device)
|
||||
torch_load_device = torch.device("cpu") if requested_device.type == "cuda" else requested_device
|
||||
pl_sd = torch.load(ckpt, map_location=torch_load_device, weights_only=True, **torch_args)
|
||||
if "state_dict" in pl_sd:
|
||||
sd = pl_sd["state_dict"]
|
||||
else:
|
||||
@@ -217,6 +417,25 @@ def _patched_load_torch_file(ckpt, safe_load=False, device=None, return_metadata
|
||||
sd = pl_sd
|
||||
else:
|
||||
sd = pl_sd
|
||||
|
||||
if isinstance(sd, dict):
|
||||
for key, value in list(sd.items()):
|
||||
if torch.is_tensor(value):
|
||||
sd[key] = _prepare_loaded_tensor(
|
||||
key,
|
||||
value,
|
||||
requested_device,
|
||||
disable_mmap=False,
|
||||
move_to_requested_device=True,
|
||||
)
|
||||
|
||||
method = "torch_load_cpu_first_to_cuda" if requested_device.type == "cuda" else "torch_load_cpu"
|
||||
_record_generic_load(
|
||||
path=ckpt,
|
||||
method=method,
|
||||
requested_device=requested_device,
|
||||
actual_device=_state_dict_device_summary(sd, requested_device) if isinstance(sd, dict) else str(requested_device),
|
||||
)
|
||||
return (sd, metadata) if return_metadata else sd
|
||||
|
||||
|
||||
@@ -308,18 +527,31 @@ def _wrap_free_memory(func: Callable[..., Any]) -> Callable[..., Any]:
|
||||
if REGISTRY.get_policy() == "sticky_gpu":
|
||||
sticky_wrappers = [w for w in REGISTRY.sticky_loaded_wrappers(device) if w not in keep_loaded]
|
||||
|
||||
unloaded = func(memory_required, device, keep_loaded + sticky_wrappers, *args, **kwargs)
|
||||
|
||||
protected_wrappers: list[Any] = []
|
||||
if device is not None and sticky_wrappers:
|
||||
try:
|
||||
free_after = model_management.get_free_memory(device)
|
||||
free_now = model_management.get_free_memory(device)
|
||||
except Exception:
|
||||
free_after = None
|
||||
if free_after is not None and free_after < memory_required:
|
||||
_LOG.warning(
|
||||
"GPU Resident Loader: sticky set exceeded VRAM budget; allowing fallback eviction to satisfy request"
|
||||
free_now = None
|
||||
if free_now is None:
|
||||
protected_wrappers = sticky_wrappers
|
||||
else:
|
||||
unloadable_wrappers = []
|
||||
for loaded in list(model_management.current_loaded_models):
|
||||
if loaded.device == device and loaded not in keep_loaded and not loaded.is_dead():
|
||||
unloadable_wrappers.append(loaded)
|
||||
available_for_protection = max(
|
||||
0,
|
||||
free_now + sum(max(0, loaded.model_loaded_memory()) for loaded in unloadable_wrappers) - memory_required,
|
||||
)
|
||||
unloaded = func(memory_required, device, keep_loaded, *args, **kwargs)
|
||||
protected_memory = 0
|
||||
for loaded in sticky_wrappers:
|
||||
estimated_memory = max(0, loaded.model_loaded_memory())
|
||||
if protected_memory + estimated_memory <= available_for_protection:
|
||||
protected_wrappers.append(loaded)
|
||||
protected_memory += estimated_memory
|
||||
|
||||
unloaded = func(memory_required, device, keep_loaded + protected_wrappers, *args, **kwargs)
|
||||
|
||||
REGISTRY.refresh_runtime_state()
|
||||
return unloaded
|
||||
@@ -334,6 +566,15 @@ def _remember_original(key: str, value: Callable[..., Any]) -> Callable[..., Any
|
||||
def _patch_model_management_devices() -> None:
|
||||
import comfy.model_management as model_management
|
||||
|
||||
kind_by_function = {
|
||||
"unet_offload_device": KIND_MODEL,
|
||||
"unet_inital_load_device": KIND_MODEL,
|
||||
"text_encoder_offload_device": KIND_CLIP,
|
||||
"text_encoder_device": KIND_CLIP,
|
||||
"vae_offload_device": KIND_VAE,
|
||||
"vae_device": KIND_VAE,
|
||||
}
|
||||
|
||||
def wrap_device_func(name: str) -> None:
|
||||
key = f"model_management.{name}"
|
||||
original = _remember_original(key, getattr(model_management, name))
|
||||
@@ -343,7 +584,8 @@ def _patch_model_management_devices() -> None:
|
||||
@functools.wraps(original)
|
||||
def wrapper(*args, **kwargs):
|
||||
result = original(*args, **kwargs)
|
||||
if not REGISTRY.wants_gpu_offload(name):
|
||||
kind = kind_by_function.get(name)
|
||||
if not REGISTRY.wants_gpu_offload(kind):
|
||||
return result
|
||||
gpu_device = model_management.get_torch_device()
|
||||
if getattr(gpu_device, "type", None) == "cpu":
|
||||
|
||||
+160
-36
@@ -40,6 +40,7 @@ class LoadContext:
|
||||
source_path: str | None = None
|
||||
explicit_device: torch.device | None = None
|
||||
note: str | None = None
|
||||
cache_key: str | None = None
|
||||
|
||||
|
||||
@dataclasses.dataclass(slots=True)
|
||||
@@ -82,8 +83,10 @@ class ResidencyEntry:
|
||||
current_device: str | None = None
|
||||
last_method: str | None = None
|
||||
last_report: dict[str, Any] | None = None
|
||||
loader_key: str | None = None
|
||||
notes: list[str] = dataclasses.field(default_factory=list)
|
||||
object_ref: weakref.ReferenceType[Any] | None = None
|
||||
cached_object_ref: weakref.ReferenceType[Any] | None = None
|
||||
|
||||
def is_alive(self) -> bool:
|
||||
return self.object_ref is not None and self.object_ref() is not None
|
||||
@@ -91,6 +94,9 @@ class ResidencyEntry:
|
||||
def object(self) -> Any | None:
|
||||
return None if self.object_ref is None else self.object_ref()
|
||||
|
||||
def cached_object(self) -> Any | None:
|
||||
return None if self.cached_object_ref is None else self.cached_object_ref()
|
||||
|
||||
def as_dict(self) -> dict[str, Any]:
|
||||
basename = os.path.basename(self.source_path) if self.source_path else None
|
||||
return {
|
||||
@@ -109,6 +115,7 @@ class ResidencyEntry:
|
||||
"current_device": self.current_device,
|
||||
"last_method": self.last_method,
|
||||
"last_report": self.last_report,
|
||||
"loader_key": self.loader_key,
|
||||
"notes": list(self.notes),
|
||||
"alive": self.is_alive(),
|
||||
}
|
||||
@@ -120,9 +127,37 @@ class ResidencyRegistry:
|
||||
self._entries: dict[str, ResidencyEntry] = {}
|
||||
self._reports_by_path: dict[str, LoadReport] = {}
|
||||
self._path_to_entry: dict[tuple[str, str], str] = {}
|
||||
self._loader_key_to_entry: dict[tuple[str, str, str], str] = {}
|
||||
self._object_to_entry: weakref.WeakKeyDictionary[Any, str] = weakref.WeakKeyDictionary()
|
||||
self._policy = self._default_policy()
|
||||
|
||||
def _gpu_ingest_kinds(self, policy: str) -> set[str]:
|
||||
if policy in {"prefer_gpu", "sticky_gpu"}:
|
||||
return {KIND_MODEL, KIND_CHECKPOINT, KIND_CLIP, KIND_VAE, KIND_CONTROLNET, KIND_CLIP_VISION}
|
||||
return set()
|
||||
|
||||
def _gpu_offload_kinds(self, policy: str) -> set[str]:
|
||||
if policy == "sticky_gpu":
|
||||
return {KIND_MODEL, KIND_CLIP, KIND_VAE, KIND_CONTROLNET}
|
||||
if policy == "prefer_gpu":
|
||||
return {KIND_MODEL, KIND_CLIP, KIND_CONTROLNET}
|
||||
return set()
|
||||
|
||||
def _autopin_kinds(self, policy: str) -> set[str]:
|
||||
if policy == "sticky_gpu":
|
||||
return {KIND_MODEL, KIND_CLIP}
|
||||
return set()
|
||||
|
||||
def default_priority(self, kind: str) -> int:
|
||||
return {
|
||||
KIND_MODEL: 300,
|
||||
KIND_CHECKPOINT: 300,
|
||||
KIND_CONTROLNET: 200,
|
||||
KIND_CLIP: 150,
|
||||
KIND_VAE: 100,
|
||||
KIND_CLIP_VISION: 50,
|
||||
}.get(kind, 0)
|
||||
|
||||
def _default_policy(self) -> str:
|
||||
env_value = os.environ.get("COMFYUI_GPU_RESIDENT_POLICY", "").strip().lower()
|
||||
if env_value in POLICIES:
|
||||
@@ -154,14 +189,21 @@ class ResidencyRegistry:
|
||||
|
||||
def wants_gpu_ingest(self, kind: str | None = None) -> bool:
|
||||
policy = self.get_policy()
|
||||
return policy in {"prefer_gpu", "sticky_gpu"}
|
||||
if kind is None:
|
||||
return bool(self._gpu_ingest_kinds(policy))
|
||||
return kind in self._gpu_ingest_kinds(policy)
|
||||
|
||||
def wants_gpu_offload(self, kind: str | None = None) -> bool:
|
||||
policy = self.get_policy()
|
||||
return policy in {"prefer_gpu", "sticky_gpu"}
|
||||
if kind is None:
|
||||
return bool(self._gpu_offload_kinds(policy))
|
||||
return kind in self._gpu_offload_kinds(policy)
|
||||
|
||||
def autopin_on_bind(self, kind: str | None = None) -> bool:
|
||||
return self.get_policy() == "sticky_gpu"
|
||||
policy = self.get_policy()
|
||||
if kind is None:
|
||||
return bool(self._autopin_kinds(policy))
|
||||
return kind in self._autopin_kinds(policy)
|
||||
|
||||
def explicit_load_device(self, kind: str, source_path: str | None = None) -> torch.device | None:
|
||||
if not self.wants_gpu_ingest(kind):
|
||||
@@ -183,6 +225,7 @@ class ResidencyRegistry:
|
||||
source_path: str | None = None,
|
||||
explicit_device: torch.device | None = None,
|
||||
note: str | None = None,
|
||||
cache_key: str | None = None,
|
||||
) -> Iterator[None]:
|
||||
token = _LOAD_CONTEXT.set(
|
||||
LoadContext(
|
||||
@@ -190,6 +233,7 @@ class ResidencyRegistry:
|
||||
source_path=source_path,
|
||||
explicit_device=explicit_device,
|
||||
note=note,
|
||||
cache_key=cache_key,
|
||||
)
|
||||
)
|
||||
try:
|
||||
@@ -222,7 +266,12 @@ class ResidencyRegistry:
|
||||
)
|
||||
with self._lock:
|
||||
self._reports_by_path[path] = report
|
||||
entry_id = self._path_to_entry.get((kind, path))
|
||||
entry_id = None
|
||||
ctx = self.current_context()
|
||||
if ctx is not None and ctx.cache_key is not None:
|
||||
entry_id = self._loader_key_to_entry.get((kind, path, ctx.cache_key))
|
||||
else:
|
||||
entry_id = self._path_to_entry.get((kind, path))
|
||||
if entry_id is not None:
|
||||
entry = self._entries.get(entry_id)
|
||||
if entry is not None:
|
||||
@@ -242,6 +291,40 @@ class ResidencyRegistry:
|
||||
basename = os.path.basename(source_path) or "anonymous"
|
||||
return f"{kind}:{basename}:{len(self._entries) + 1}"
|
||||
|
||||
def _clear_object_binding(self, obj: Any | None, entry_id: str) -> None:
|
||||
if obj is None:
|
||||
return
|
||||
try:
|
||||
if self._object_to_entry.get(obj) == entry_id:
|
||||
self._object_to_entry.pop(obj, None)
|
||||
except TypeError:
|
||||
pass
|
||||
try:
|
||||
if getattr(obj, "__gpu_resident_loader_entry_id__", None) == entry_id:
|
||||
delattr(obj, "__gpu_resident_loader_entry_id__")
|
||||
except (AttributeError, TypeError):
|
||||
pass
|
||||
|
||||
def _tag_object_with_entry(self, obj: Any, entry_id: str) -> None:
|
||||
try:
|
||||
setattr(obj, "__gpu_resident_loader_entry_id__", entry_id)
|
||||
try:
|
||||
self._object_to_entry.pop(obj, None)
|
||||
except TypeError:
|
||||
pass
|
||||
return
|
||||
except (AttributeError, TypeError):
|
||||
_LOG.debug(
|
||||
"GPU Resident Loader: could not tag object %r with residency entry id %s",
|
||||
type(obj),
|
||||
entry_id,
|
||||
)
|
||||
|
||||
try:
|
||||
self._object_to_entry[obj] = entry_id
|
||||
except TypeError:
|
||||
pass
|
||||
|
||||
def bind_object(
|
||||
self,
|
||||
obj: Any,
|
||||
@@ -249,15 +332,22 @@ class ResidencyRegistry:
|
||||
source_path: str,
|
||||
kind: str,
|
||||
sticky: bool | None = None,
|
||||
priority: int = 0,
|
||||
priority: int | None = None,
|
||||
note: str | None = None,
|
||||
loader_key: str | None = None,
|
||||
reusable_obj: Any | None = None,
|
||||
) -> ResidencyEntry:
|
||||
if obj is None:
|
||||
raise ValueError("Cannot bind None into residency registry")
|
||||
|
||||
with self._lock:
|
||||
old_key: tuple[str, str] | None = None
|
||||
entry_id = getattr(obj, "__gpu_resident_loader_entry_id__", None)
|
||||
old_loader_key: tuple[str, str, str] | None = None
|
||||
previous_obj: Any | None = None
|
||||
previous_cached_obj: Any | None = None
|
||||
entry_id = self._loader_key_to_entry.get((kind, source_path, loader_key)) if loader_key is not None else None
|
||||
if entry_id is None:
|
||||
entry_id = getattr(obj, "__gpu_resident_loader_entry_id__", None)
|
||||
if entry_id is None:
|
||||
try:
|
||||
entry_id = self._object_to_entry.get(obj)
|
||||
@@ -266,6 +356,10 @@ class ResidencyRegistry:
|
||||
if entry_id is not None and entry_id in self._entries:
|
||||
entry = self._entries[entry_id]
|
||||
old_key = (entry.kind, entry.source_path)
|
||||
previous_obj = entry.object()
|
||||
previous_cached_obj = entry.cached_object()
|
||||
if entry.loader_key is not None:
|
||||
old_loader_key = (entry.kind, entry.source_path, entry.loader_key)
|
||||
else:
|
||||
entry_id = self._make_entry_id(kind, source_path)
|
||||
entry = ResidencyEntry(
|
||||
@@ -273,45 +367,46 @@ class ResidencyRegistry:
|
||||
kind=kind,
|
||||
source_path=source_path,
|
||||
sticky=self.autopin_on_bind(kind) if sticky is None else bool(sticky),
|
||||
priority=int(priority),
|
||||
priority=self.default_priority(kind) if priority is None else int(priority),
|
||||
)
|
||||
self._entries[entry_id] = entry
|
||||
self._path_to_entry[(kind, source_path)] = entry_id
|
||||
try:
|
||||
setattr(obj, "__gpu_resident_loader_entry_id__", entry_id)
|
||||
try:
|
||||
self._object_to_entry.pop(obj, None)
|
||||
except TypeError:
|
||||
pass
|
||||
except (AttributeError, TypeError):
|
||||
_LOG.debug(
|
||||
"GPU Resident Loader: could not tag object %r with residency entry id %s",
|
||||
type(obj),
|
||||
entry_id,
|
||||
)
|
||||
|
||||
cache_obj = obj if reusable_obj is None else reusable_obj
|
||||
for stale_obj in (previous_obj, previous_cached_obj):
|
||||
if stale_obj is None or stale_obj is obj or stale_obj is cache_obj:
|
||||
continue
|
||||
self._clear_object_binding(stale_obj, entry_id)
|
||||
|
||||
try:
|
||||
entry.object_ref = weakref.ref(obj)
|
||||
if getattr(obj, "__gpu_resident_loader_entry_id__", None) is None:
|
||||
try:
|
||||
self._object_to_entry[obj] = entry_id
|
||||
except TypeError:
|
||||
pass
|
||||
except TypeError:
|
||||
entry.object_ref = None
|
||||
_LOG.debug(
|
||||
"GPU Resident Loader: object %r is not weak-referenceable; tracking metadata only",
|
||||
type(obj),
|
||||
)
|
||||
self._tag_object_with_entry(obj, entry_id)
|
||||
try:
|
||||
entry.cached_object_ref = weakref.ref(cache_obj)
|
||||
except TypeError:
|
||||
entry.cached_object_ref = entry.object_ref
|
||||
entry.sticky = entry.sticky if sticky is None else bool(sticky)
|
||||
entry.priority = int(priority)
|
||||
entry.priority = entry.priority if priority is None else int(priority)
|
||||
entry.source_path = source_path
|
||||
entry.kind = kind
|
||||
entry.loader_key = loader_key
|
||||
new_key = (entry.kind, entry.source_path)
|
||||
new_loader_key = (entry.kind, entry.source_path, entry.loader_key) if entry.loader_key is not None else None
|
||||
if old_key is not None and old_key != new_key:
|
||||
if self._path_to_entry.get(old_key) == entry_id:
|
||||
self._path_to_entry.pop(old_key, None)
|
||||
self._path_to_entry[new_key] = entry_id
|
||||
if old_loader_key is not None and old_loader_key != new_loader_key:
|
||||
if self._loader_key_to_entry.get(old_loader_key) == entry_id:
|
||||
self._loader_key_to_entry.pop(old_loader_key, None)
|
||||
if new_loader_key is not None:
|
||||
self._loader_key_to_entry[new_loader_key] = entry_id
|
||||
entry.last_touched = _now()
|
||||
if note:
|
||||
entry.notes.append(note)
|
||||
@@ -321,21 +416,49 @@ class ResidencyRegistry:
|
||||
entry.last_report = report.as_dict()
|
||||
entry.current_device = report.actual_device
|
||||
|
||||
self.refresh_runtime_state()
|
||||
return entry
|
||||
|
||||
def lookup_live_object(self, *, kind: str, source_path: str, loader_key: str) -> Any | None:
|
||||
with self._lock:
|
||||
entry_id = self._loader_key_to_entry.get((kind, source_path, loader_key))
|
||||
if entry_id is None:
|
||||
return None
|
||||
entry = self._entries.get(entry_id)
|
||||
if entry is None:
|
||||
return None
|
||||
obj = entry.cached_object()
|
||||
if obj is None:
|
||||
if entry.cached_object_ref is not None:
|
||||
return None
|
||||
obj = entry.object()
|
||||
if obj is None:
|
||||
return None
|
||||
entry.last_touched = _now()
|
||||
return obj
|
||||
|
||||
def _entry_id_for_object(self, obj: Any) -> str | None:
|
||||
current = obj
|
||||
seen: set[int] = set()
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
entry_id = getattr(current, "__gpu_resident_loader_entry_id__", None)
|
||||
if entry_id is None:
|
||||
try:
|
||||
entry_id = self._object_to_entry.get(current)
|
||||
except TypeError:
|
||||
entry_id = None
|
||||
if entry_id is not None:
|
||||
return entry_id
|
||||
current = getattr(current, "parent", None)
|
||||
return None
|
||||
|
||||
def entry_for_object(self, obj: Any) -> ResidencyEntry | None:
|
||||
if obj is None:
|
||||
return None
|
||||
entry_id = getattr(obj, "__gpu_resident_loader_entry_id__", None)
|
||||
if entry_id is None:
|
||||
try:
|
||||
entry_id = self._object_to_entry.get(obj)
|
||||
except TypeError:
|
||||
entry_id = None
|
||||
if entry_id is None:
|
||||
return None
|
||||
with self._lock:
|
||||
entry_id = self._entry_id_for_object(obj)
|
||||
if entry_id is None:
|
||||
return None
|
||||
return self._entries.get(entry_id)
|
||||
|
||||
def set_sticky(self, obj: Any, sticky: bool, priority: int | None = None) -> ResidencyEntry | None:
|
||||
@@ -369,8 +492,9 @@ class ResidencyRegistry:
|
||||
continue
|
||||
entry = self.entry_for_object(loaded.model)
|
||||
if entry is not None and entry.sticky:
|
||||
output.append(loaded)
|
||||
return output
|
||||
output.append((entry.priority, entry.last_touched, loaded))
|
||||
output.sort(key=lambda item: (item[0], item[1]), reverse=True)
|
||||
return [item[2] for item in output]
|
||||
|
||||
def refresh_runtime_state(self) -> None:
|
||||
try:
|
||||
|
||||
Reference in New Issue
Block a user