Compare commits

...
Author SHA1 Message Date
xmarre cfb707f082 Fix residency entry lookup locking 2026-04-15 02:36:35 +02:00
xmarre 19e69a8ecc Address PR review comments 2026-04-15 02:29:52 +02:00
xmarre f7b8fe5cc9 Add selective checkpoint ingest and resident VRAM trim 2026-04-15 02:17:16 +02:00
xmarre 2643a1903b Merge pull request #5 from xmarre/codex/narrow-resident-loader-hot-path
[codex] Narrow resident loader hot paths
2026-04-13 06:28:08 +02:00
xmarre d11d8e0a25 Fix loader-keyed load report attribution 2026-04-13 06:24:27 +02:00
xmarre df9128e9ee Fix resident loader reuse edge cases 2026-04-13 06:16:06 +02:00
xmarre 9f4b27eb23 Narrow resident loader hot paths 2026-04-13 06:00:19 +02:00
xmarre b94fd0b3f3 Merge pull request #4 from xmarre/codex/fix-loader-extra-unet-filter
Filter extra diffusion loader state dicts to UNet keys
2026-04-12 10:39:55 +02:00
xmarre 1b80a6a82b Preserve sparse extra UNet overrides 2026-04-12 10:36:40 +02:00
xmarre 6d4bd6c7c4 Filter extra diffusion weights to UNet keys 2026-04-12 10:24:59 +02:00
xmarre d533beb275 Merge pull request #3 from xmarre/codex/policy-override-loader
Add explicit residency policy loader input
2026-04-12 06:49:22 +02:00
xmarre 24ab29880a Scope loader policy override to load 2026-04-12 06:45:32 +02:00
xmarre 791a48afd3 Add explicit residency policy loader input 2026-04-12 06:35:30 +02:00
xmarre 1253b88361 Merge pull request #2 from xmarre/codex/fix-clip-metadata-cpu
[codex] Keep CLIP metadata tensors on CPU
2026-04-12 05:55:53 +02:00
xmarre ca8118de03 fix cpu fallback and torch remap handling 2026-04-12 05:54:14 +02:00
xmarre 752def584f load torch checkpoints on cpu before cuda remap 2026-04-12 05:42:29 +02:00
xmarre 871c144ac9 keep clip metadata tensors on cpu 2026-04-12 05:26:51 +02:00
xmarre ca0d55ce04 Merge pull request #1 from xmarre/codex/import-zip-implementation
[codex] Import GPU resident loader implementation
2026-04-12 04:52:39 +02:00
6 changed files with 1415 additions and 163 deletions
+21 -6
View File
@@ -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
View File
@@ -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
View File
@@ -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,)
+65 -1
View File
@@ -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
View File
@@ -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
View File
@@ -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: