Merge pull request #1 from xmarre/codex/import-zip-implementation

[codex] Import GPU resident loader implementation
This commit is contained in:
xmarre
2026-04-12 04:52:39 +02:00
committed by GitHub
11 changed files with 1974 additions and 0 deletions
+23
View File
@@ -0,0 +1,23 @@
__pycache__/
*.py[cod]
*.pyo
*.pyd
.Python
.venv/
venv/
env/
ENV/
build/
dist/
*.egg-info/
.mypy_cache/
.ruff_cache/
.pytest_cache/
.DS_Store
.vs/
.v17/
.codex/
+175
View File
@@ -0,0 +1,175 @@
# ComfyUI GPU Resident Loader
A ComfyUI custom-node pack that targets **time-to-VRAM** and **sticky GPU residency**, not just lower peak host RAM.
It does two things:
1. **Installs startup-time loader and residency patches** before any workflow nodes run.
2. Ships **KJ-compatible loader nodes** for diffusion models and checkpoints, plus preload/pin/evict/report nodes for manual residency control.
## Why this exists
Stock ComfyUI makes separate decisions for:
- where a model **lives after load**, and
- where checkpoint tensors are **materialized first**.
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.
## What is included
### Startup patcher
Installed automatically from `__init__.py` when the custom node loads.
Current patch surface:
- `comfy.utils.load_torch_file`
- `comfy.model_management.free_memory`
- `comfy.model_management.load_models_gpu`
- `comfy.model_management.unet_offload_device`
- `comfy.model_management.text_encoder_offload_device`
- `comfy.model_management.vae_offload_device`
- `comfy.model_management.text_encoder_device`
- `comfy.model_management.vae_device`
- `comfy.model_management.unet_inital_load_device`
- `comfy.sd.load_checkpoint_guess_config`
- `comfy.sd.load_diffusion_model`
- `comfy.sd.load_clip`
- `comfy.clip_vision.load`
- `comfy.controlnet.load_controlnet`
- `comfy.diffusers_load.load_diffusers`
### Loader nodes
- **Diffusion Model Loader Resident**
- **Checkpoint Loader Resident**
- **Diffusion Model Selector Resident**
`Diffusion Model Loader Resident` mirrors the relevant KJ diffusion-loader feature surface:
- weight dtype override
- compute dtype override
- cublas-ops toggle
- SageAttention override
- fp16 accumulation toggle
- optional extra-state-dict merge
### Residency nodes
- **Set Global Residency Policy**
- **Registry Snapshot**
- **Pin Model/CLIP/VAE Residency**
- **Preload Model/CLIP/VAE To GPU**
- **Evict Model/CLIP/VAE From GPU**
- **Report Model/CLIP/VAE Residency**
## Policies
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.
Default selection order:
1. `COMFYUI_GPU_RESIDENT_POLICY` environment variable, if set.
2. `sticky_gpu` when `--gpu_only` is active.
3. `sticky_gpu` when `--highvram` is active.
4. otherwise `prefer_gpu`.
## Important scope limits
### Best path: `.safetensors`
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.
### `.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.
Use the included conversion helper to migrate hot models to `.safetensors`.
### Cross-process persistence is out of scope
This repo does **not** keep VRAM contents alive after ComfyUI or WSL exits. CUDA memory lifetime is process/context scoped. Achieving persistence across process shutdown requires a long-lived keeper process or server that owns the CUDA context.
## Installation
Clone into `custom_nodes`:
```bash
git clone https://github.com/xmarre/ComfyUI-GPU-Resident-Loader ComfyUI/custom_nodes/ComfyUI-GPU-Resident-Loader
```
Install dependencies inside the same Python environment ComfyUI uses:
```bash
pip install -r ComfyUI/custom_nodes/ComfyUI-GPU-Resident-Loader/requirements.txt
```
Optional SageAttention dependencies are **not** installed by default. Install those separately if you plan to use the SageAttention loader modes.
## Basic usage
### For direct diffusion-model loading
Use **Diffusion Model Loader Resident**.
Recommended on a large VRAM machine:
- policy: `sticky_gpu`
- model format: `.safetensors`
- preload with **Preload Model To GPU**
- inspect with **Report Model Residency** or **Registry Snapshot**
### For full checkpoints
Use **Checkpoint Loader Resident**.
That tracks and binds the resulting diffusion model, CLIP, and VAE independently so they appear in the registry snapshot.
### For manual residency control
- use **Pin ... Residency** to mark a tracked object sticky or evictable
- use **Preload ... To GPU** to fully materialize it in VRAM immediately
- use **Evict ... From GPU** to unload it from the current loaded-model set
## Observability
Every tracked load stores:
- source path
- last load method
- requested device
- actual device
- sticky flag
- current loaded bytes
- total bytes
- current/offload/load device
That data is surfaced through the report nodes and the registry snapshot node.
## Conversion helper
`scripts/convert_checkpoint_to_safetensors.py` is included for one-time conversion of hot `.ckpt` / `.pt` files into `.safetensors`.
Example:
```bash
python ComfyUI/custom_nodes/ComfyUI-GPU-Resident-Loader/scripts/convert_checkpoint_to_safetensors.py \
--input /path/to/model.ckpt \
--output /path/to/model.safetensors
```
## License
GPL-3.0-or-later.
This repo intentionally stays GPL-compatible because it adapts behavior from GPL-licensed ComfyUI and mirrors feature behavior from the GPL-3.0-licensed KJNodes diffusion loader.
+9
View File
@@ -0,0 +1,9 @@
from .startup import install_patches
from .nodes import (
NODE_CLASS_MAPPINGS,
NODE_DISPLAY_NAME_MAPPINGS,
)
install_patches()
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+426
View File
@@ -0,0 +1,426 @@
from __future__ import annotations
import contextlib
import logging
from typing import Any
import folder_paths
import torch
import comfy.sd
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
_LOG = logging.getLogger(__name__)
SAGE_ATTN_MODES = [
"disabled",
"auto",
"sageattn_qk_int8_pv_fp16_cuda",
"sageattn_qk_int8_pv_fp16_triton",
"sageattn_qk_int8_pv_fp8_cuda",
"sageattn_qk_int8_pv_fp8_cuda++",
"sageattn3",
"sageattn3_per_block_mean",
]
DTYPE_MAP = {
"fp8_e4m3fn": torch.float8_e4m3fn,
"fp8_e5m2": torch.float8_e5m2,
"fp16": torch.float16,
"bf16": torch.bfloat16,
"fp32": torch.float32,
}
def _set_cublas_linear(enabled: bool) -> None:
if enabled:
args.fast.add(PerformanceFeature.CublasOps)
else:
args.fast.discard(PerformanceFeature.CublasOps)
def _set_fp16_accumulation(enabled: bool) -> None:
if not hasattr(torch.backends.cuda, "matmul"):
if enabled:
raise RuntimeError(
"Failed to enable fp16 accumulation. This requires a PyTorch build exposing "
"torch.backends.cuda.matmul.allow_fp16_accumulation."
)
_LOG.warning(
"GPU Resident Loader: fp16 accumulation toggle is unavailable on this PyTorch build; leaving default behavior."
)
return
flag = getattr(torch.backends.cuda.matmul, "allow_fp16_accumulation", None)
if flag is None:
if enabled:
raise RuntimeError(
"Failed to enable fp16 accumulation. This requires a PyTorch build exposing "
"torch.backends.cuda.matmul.allow_fp16_accumulation."
)
_LOG.warning(
"GPU Resident Loader: fp16 accumulation toggle is unavailable on this PyTorch build; leaving default behavior."
)
return
torch.backends.cuda.matmul.allow_fp16_accumulation = bool(enabled)
def _get_fp16_accumulation_state() -> bool | None:
matmul = getattr(torch.backends.cuda, "matmul", None)
return getattr(matmul, "allow_fp16_accumulation", None)
@contextlib.contextmanager
def _temporary_backend_flags(*, cublas: bool, fp16_accumulation: bool):
prev_cublas = PerformanceFeature.CublasOps in args.fast
prev_fp16 = _get_fp16_accumulation_state()
try:
_set_cublas_linear(cublas)
_set_fp16_accumulation(fp16_accumulation)
yield
finally:
_set_cublas_linear(prev_cublas)
if prev_fp16 is not None:
torch.backends.cuda.matmul.allow_fp16_accumulation = prev_fp16
def get_sage_func(sage_attention: str, allow_compile: bool = False):
_LOG.info("GPU Resident Loader: using sage attention mode %s", sage_attention)
if sage_attention == "auto":
from sageattention import sageattn
def sage_func(q, k, v, is_causal=False, attn_mask=None, tensor_layout="NHD"):
return sageattn(q, k, v, is_causal=is_causal, attn_mask=attn_mask, tensor_layout=tensor_layout)
elif sage_attention == "sageattn_qk_int8_pv_fp16_cuda":
from sageattention import sageattn_qk_int8_pv_fp16_cuda
def sage_func(q, k, v, is_causal=False, attn_mask=None, tensor_layout="NHD"):
return sageattn_qk_int8_pv_fp16_cuda(
q,
k,
v,
is_causal=is_causal,
attn_mask=attn_mask,
pv_accum_dtype="fp32",
tensor_layout=tensor_layout,
)
elif sage_attention == "sageattn_qk_int8_pv_fp16_triton":
from sageattention import sageattn_qk_int8_pv_fp16_triton
def sage_func(q, k, v, is_causal=False, attn_mask=None, tensor_layout="NHD"):
return sageattn_qk_int8_pv_fp16_triton(
q,
k,
v,
is_causal=is_causal,
attn_mask=attn_mask,
tensor_layout=tensor_layout,
)
elif sage_attention == "sageattn_qk_int8_pv_fp8_cuda":
from sageattention import sageattn_qk_int8_pv_fp8_cuda
def sage_func(q, k, v, is_causal=False, attn_mask=None, tensor_layout="NHD"):
return sageattn_qk_int8_pv_fp8_cuda(
q,
k,
v,
is_causal=is_causal,
attn_mask=attn_mask,
pv_accum_dtype="fp32+fp32",
tensor_layout=tensor_layout,
)
elif sage_attention == "sageattn_qk_int8_pv_fp8_cuda++":
from sageattention import sageattn_qk_int8_pv_fp8_cuda
def sage_func(q, k, v, is_causal=False, attn_mask=None, tensor_layout="NHD"):
return sageattn_qk_int8_pv_fp8_cuda(
q,
k,
v,
is_causal=is_causal,
attn_mask=attn_mask,
pv_accum_dtype="fp32+fp16",
tensor_layout=tensor_layout,
)
elif "sageattn3" in sage_attention:
from sageattn3 import sageattn3_blackwell
def sage_func(q, k, v, is_causal=False, attn_mask=None, tensor_layout="NHD", **kwargs):
q, k, v = [x.transpose(1, 2) if tensor_layout == "NHD" else x for x in (q, k, v)]
out = sageattn3_blackwell(
q,
k,
v,
is_causal=is_causal,
attn_mask=attn_mask,
per_block_mean=(sage_attention == "sageattn3_per_block_mean"),
)
return out.transpose(1, 2) if tensor_layout == "NHD" else out
else:
raise ValueError(f"Unsupported sage attention mode: {sage_attention}")
compiler = getattr(torch, "compiler", None)
disable = getattr(compiler, "disable", None)
if not allow_compile and callable(disable):
sage_func = disable()(sage_func)
@wrap_attn
def attention_sage(q, k, v, heads, mask=None, attn_precision=None, skip_reshape=False, skip_output_reshape=False, **kwargs):
if kwargs.get("low_precision_attention", True) is False:
return attention_pytorch(
q,
k,
v,
heads,
mask=mask,
skip_reshape=skip_reshape,
skip_output_reshape=skip_output_reshape,
**kwargs,
)
in_dtype = v.dtype
if q.dtype == torch.float32 or k.dtype == torch.float32 or v.dtype == torch.float32:
q, k, v = q.to(torch.float16), k.to(torch.float16), v.to(torch.float16)
if skip_reshape:
batch, _, _, dim_head = q.shape
tensor_layout = "HND"
else:
batch, _, dim_head = q.shape
dim_head //= heads
q, k, v = [tensor.view(batch, -1, heads, dim_head) for tensor in (q, k, v)]
tensor_layout = "NHD"
if mask is not None:
if mask.ndim == 2:
mask = mask.unsqueeze(0)
if mask.ndim == 3:
mask = mask.unsqueeze(1)
out = sage_func(q, k, v, attn_mask=mask, is_causal=False, tensor_layout=tensor_layout).to(in_dtype)
if tensor_layout == "HND":
if not skip_output_reshape:
out = out.transpose(1, 2).reshape(batch, -1, heads * dim_head)
else:
if skip_output_reshape:
out = out.transpose(1, 2)
else:
out = out.reshape(batch, -1, heads * dim_head)
return out
return attention_sage
def _apply_model_postload_options(model, *, compute_dtype: str, sage_attention: str) -> None:
if dtype := DTYPE_MAP.get(compute_dtype):
model.set_model_compute_dtype(dtype)
model.force_cast_weights = False
_LOG.info("GPU Resident Loader: set compute dtype to %s", dtype)
if sage_attention != "disabled":
new_attention = get_sage_func(sage_attention)
def attention_override_sage(func, *args, **kwargs):
return new_attention.__wrapped__(*args, **kwargs)
model.model_options.setdefault("transformer_options", {})["optimized_attention_override"] = attention_override_sage
def _build_model_options(weight_dtype: str) -> dict[str, Any]:
model_options: dict[str, Any] = {}
if dtype := DTYPE_MAP.get(weight_dtype):
model_options["dtype"] = dtype
_LOG.info("GPU Resident Loader: set weight dtype to %s", dtype)
if weight_dtype == "fp8_e4m3fn_fast":
model_options["dtype"] = torch.float8_e4m3fn
model_options["fp8_optimizations"] = True
return model_options
class DiffusionModelSelectorResident:
@classmethod
def INPUT_TYPES(cls):
ltx2_connector_models = folder_paths.get_filename_list("text_encoders")
ltx2_connector_models = [m for m in ltx2_connector_models if "connector" in m.lower()]
return {
"required": {
"model_name": (
folder_paths.get_filename_list("diffusion_models") + ltx2_connector_models,
{"tooltip": "The name of the diffusion model or connector to resolve."},
),
}
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("model_path",)
FUNCTION = "get_path"
CATEGORY = "GPU Resident Loader/loaders"
DESCRIPTION = "Returns the absolute model path as a string. Mirrors the selector behavior of KJ's diffusion model selector."
def get_path(self, model_name: str):
if "connector" in model_name.lower():
model_path = folder_paths.get_full_path_or_raise("text_encoders", model_name)
else:
model_path = folder_paths.get_full_path_or_raise("diffusion_models", model_name)
return (model_path,)
class DiffusionModelLoaderResident:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model_name": (
folder_paths.get_filename_list("diffusion_models"),
{"tooltip": "The diffusion model file to load."},
),
"weight_dtype": (["default", "fp8_e4m3fn", "fp8_e4m3fn_fast", "fp8_e5m2", "fp16", "bf16", "fp32"],),
"compute_dtype": (
["default", "fp16", "bf16", "fp32"],
{"default": "default", "tooltip": "Compute dtype to apply after model creation."},
),
"patch_cublaslinear": (
"BOOLEAN",
{"default": False, "tooltip": "Toggle ComfyUI's cublas_ops performance feature."},
),
"sage_attention": (
SAGE_ATTN_MODES,
{"default": "disabled", "tooltip": "Patch optimized attention override to a SageAttention variant."},
),
"enable_fp16_accumulation": (
"BOOLEAN",
{"default": False, "tooltip": "Set torch.backends.cuda.matmul.allow_fp16_accumulation."},
),
},
"optional": {
"extra_state_dict": (
"STRING",
{
"forceInput": True,
"tooltip": "Optional absolute path to a second state dict merged into the main diffusion state dict before model detection.",
},
)
},
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "patch_and_load"
CATEGORY = "GPU Resident Loader/loaders"
DESCRIPTION = (
"KJ-compatible diffusion-model loader with GPU-resident ingest and residency tracking. "
"It mirrors the KJ node's weight dtype, compute dtype, cublas_ops, SageAttention, fp16 accumulation, "
"and optional extra-state-dict features."
)
def patch_and_load(
self,
model_name: str,
weight_dtype: str,
compute_dtype: str,
patch_cublaslinear: bool,
sage_attention: str,
enable_fp16_accumulation: bool,
extra_state_dict: str | None = None,
):
with _temporary_backend_flags(
cublas=patch_cublaslinear,
fp16_accumulation=enable_fp16_accumulation,
):
model_options = _build_model_options(weight_dtype)
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,)
class CheckpointLoaderResident:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"ckpt_name": (
folder_paths.get_filename_list("checkpoints"),
{"tooltip": "Checkpoint file to load."},
),
"weight_dtype": (["default", "fp8_e4m3fn", "fp8_e4m3fn_fast", "fp8_e5m2", "fp16", "bf16", "fp32"],),
"compute_dtype": (
["default", "fp16", "bf16", "fp32"],
{"default": "default", "tooltip": "Compute dtype to apply to the diffusion model patcher after load."},
),
"patch_cublaslinear": (
"BOOLEAN",
{"default": False, "tooltip": "Toggle ComfyUI's cublas_ops performance feature."},
),
"sage_attention": (
SAGE_ATTN_MODES,
{"default": "disabled", "tooltip": "Patch optimized attention override on the loaded model."},
),
"enable_fp16_accumulation": (
"BOOLEAN",
{"default": False, "tooltip": "Set torch.backends.cuda.matmul.allow_fp16_accumulation."},
),
}
}
RETURN_TYPES = ("MODEL", "CLIP", "VAE")
FUNCTION = "load"
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."
)
def load(
self,
ckpt_name: str,
weight_dtype: str,
compute_dtype: str,
patch_cublaslinear: bool,
sage_attention: str,
enable_fp16_accumulation: bool,
):
with _temporary_backend_flags(
cublas=patch_cublaslinear,
fp16_accumulation=enable_fp16_accumulation,
):
model_options = _build_model_options(weight_dtype)
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,
)
_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
+365
View File
@@ -0,0 +1,365 @@
from __future__ import annotations
import json
from typing import Any
import comfy.model_management as model_management
from .kj_loader import CheckpointLoaderResident, DiffusionModelLoaderResident, DiffusionModelSelectorResident
from .residency import REGISTRY
def _entry_report_json(obj: Any) -> str:
entry = REGISTRY.entry_for_object(obj)
if entry is None:
return json.dumps(
{
"tracked": False,
"policy": REGISTRY.get_policy(),
"message": "Object is not currently bound in the GPU Resident Loader registry.",
},
indent=2,
sort_keys=True,
)
REGISTRY.refresh_runtime_state()
payload = {
"tracked": True,
"policy": REGISTRY.get_policy(),
"entry": entry.as_dict(),
}
return json.dumps(payload, indent=2, sort_keys=True)
def _patcher_for_clip(clip):
return getattr(clip, "patcher", None)
def _patcher_for_vae(vae):
return getattr(vae, "patcher", None)
def _preload_patcher(patcher, *, sticky: bool, priority: int) -> None:
if patcher is None:
raise RuntimeError("Expected a patcher-capable object, but no patcher was found.")
model_management.load_models_gpu([patcher], force_full_load=True)
REGISTRY.set_sticky(patcher, sticky=sticky, priority=priority)
REGISTRY.touch(patcher)
REGISTRY.refresh_runtime_state()
def _evict_patcher(patcher, *, unpatch_weights: bool) -> bool:
if patcher is None:
raise RuntimeError("Expected a patcher-capable object, but no patcher was found.")
unloaded = False
for loaded in list(model_management.current_loaded_models):
if loaded.model is patcher or loaded.model.is_clone(patcher):
loaded.model_unload(unpatch_weights=unpatch_weights)
unloaded = True
if unloaded and hasattr(model_management, "soft_empty_cache"):
model_management.soft_empty_cache()
REGISTRY.refresh_runtime_state()
return unloaded
class SetGlobalResidencyPolicy:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"policy": (
["legacy", "balanced", "prefer_gpu", "sticky_gpu"],
{"default": "sticky_gpu", "tooltip": "Select the global ingest/offload policy used by the startup patcher."},
)
}
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("active_policy",)
FUNCTION = "set_policy"
CATEGORY = "GPU Resident Loader/residency"
def set_policy(self, policy: str):
return (REGISTRY.set_policy(policy),)
class RegistrySnapshot:
@classmethod
def INPUT_TYPES(cls):
return {"required": {}}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("registry_json",)
FUNCTION = "snapshot"
CATEGORY = "GPU Resident Loader/residency"
def snapshot(self):
return (REGISTRY.snapshot_json(),)
class PinModelResidency:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": ("MODEL",),
"sticky": ("BOOLEAN", {"default": True}),
"priority": ("INT", {"default": 0, "min": -100, "max": 100, "step": 1}),
}
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "pin"
CATEGORY = "GPU Resident Loader/residency"
def pin(self, model, sticky: bool, priority: int):
REGISTRY.set_sticky(model, sticky=sticky, priority=priority)
return (model,)
class PinClipResidency:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"clip": ("CLIP",),
"sticky": ("BOOLEAN", {"default": True}),
"priority": ("INT", {"default": 0, "min": -100, "max": 100, "step": 1}),
}
}
RETURN_TYPES = ("CLIP",)
FUNCTION = "pin"
CATEGORY = "GPU Resident Loader/residency"
def pin(self, clip, sticky: bool, priority: int):
patcher = _patcher_for_clip(clip)
if patcher is None:
raise RuntimeError("Expected a CLIP object with a patcher, but no patcher was found.")
REGISTRY.set_sticky(patcher, sticky=sticky, priority=priority)
return (clip,)
class PinVAEResidency:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"vae": ("VAE",),
"sticky": ("BOOLEAN", {"default": True}),
"priority": ("INT", {"default": 0, "min": -100, "max": 100, "step": 1}),
}
}
RETURN_TYPES = ("VAE",)
FUNCTION = "pin"
CATEGORY = "GPU Resident Loader/residency"
def pin(self, vae, sticky: bool, priority: int):
patcher = _patcher_for_vae(vae)
if patcher is None:
raise RuntimeError("Expected a VAE object with a patcher, but no patcher was found.")
REGISTRY.set_sticky(patcher, sticky=sticky, priority=priority)
return (vae,)
class PreloadModelToGPU:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": ("MODEL",),
"sticky": ("BOOLEAN", {"default": True}),
"priority": ("INT", {"default": 0, "min": -100, "max": 100, "step": 1}),
}
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "preload"
CATEGORY = "GPU Resident Loader/residency"
def preload(self, model, sticky: bool, priority: int):
_preload_patcher(model, sticky=sticky, priority=priority)
return (model,)
class PreloadClipToGPU:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"clip": ("CLIP",),
"sticky": ("BOOLEAN", {"default": True}),
"priority": ("INT", {"default": 0, "min": -100, "max": 100, "step": 1}),
}
}
RETURN_TYPES = ("CLIP",)
FUNCTION = "preload"
CATEGORY = "GPU Resident Loader/residency"
def preload(self, clip, sticky: bool, priority: int):
_preload_patcher(_patcher_for_clip(clip), sticky=sticky, priority=priority)
return (clip,)
class PreloadVAEToGPU:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"vae": ("VAE",),
"sticky": ("BOOLEAN", {"default": True}),
"priority": ("INT", {"default": 0, "min": -100, "max": 100, "step": 1}),
}
}
RETURN_TYPES = ("VAE",)
FUNCTION = "preload"
CATEGORY = "GPU Resident Loader/residency"
def preload(self, vae, sticky: bool, priority: int):
_preload_patcher(_patcher_for_vae(vae), sticky=sticky, priority=priority)
return (vae,)
class EvictModelFromGPU:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": ("MODEL",),
"unpatch_weights": ("BOOLEAN", {"default": True}),
}
}
RETURN_TYPES = ("MODEL", "STRING")
RETURN_NAMES = ("model", "eviction_status")
FUNCTION = "evict"
CATEGORY = "GPU Resident Loader/residency"
def evict(self, model, unpatch_weights: bool):
unloaded = _evict_patcher(model, unpatch_weights=unpatch_weights)
return model, ("evicted" if unloaded else "not_loaded")
class EvictClipFromGPU:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"clip": ("CLIP",),
"unpatch_weights": ("BOOLEAN", {"default": True}),
}
}
RETURN_TYPES = ("CLIP", "STRING")
RETURN_NAMES = ("clip", "eviction_status")
FUNCTION = "evict"
CATEGORY = "GPU Resident Loader/residency"
def evict(self, clip, unpatch_weights: bool):
unloaded = _evict_patcher(_patcher_for_clip(clip), unpatch_weights=unpatch_weights)
return clip, ("evicted" if unloaded else "not_loaded")
class EvictVAEFromGPU:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"vae": ("VAE",),
"unpatch_weights": ("BOOLEAN", {"default": True}),
}
}
RETURN_TYPES = ("VAE", "STRING")
RETURN_NAMES = ("vae", "eviction_status")
FUNCTION = "evict"
CATEGORY = "GPU Resident Loader/residency"
def evict(self, vae, unpatch_weights: bool):
unloaded = _evict_patcher(_patcher_for_vae(vae), unpatch_weights=unpatch_weights)
return vae, ("evicted" if unloaded else "not_loaded")
class ReportModelResidency:
@classmethod
def INPUT_TYPES(cls):
return {"required": {"model": ("MODEL",)}}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("report_json",)
FUNCTION = "report"
CATEGORY = "GPU Resident Loader/residency"
def report(self, model):
return (_entry_report_json(model),)
class ReportClipResidency:
@classmethod
def INPUT_TYPES(cls):
return {"required": {"clip": ("CLIP",)}}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("report_json",)
FUNCTION = "report"
CATEGORY = "GPU Resident Loader/residency"
def report(self, clip):
return (_entry_report_json(_patcher_for_clip(clip)),)
class ReportVAEResidency:
@classmethod
def INPUT_TYPES(cls):
return {"required": {"vae": ("VAE",)}}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("report_json",)
FUNCTION = "report"
CATEGORY = "GPU Resident Loader/residency"
def report(self, vae):
return (_entry_report_json(_patcher_for_vae(vae)),)
NODE_CLASS_MAPPINGS = {
"DiffusionModelSelectorResident": DiffusionModelSelectorResident,
"DiffusionModelLoaderResident": DiffusionModelLoaderResident,
"CheckpointLoaderResident": CheckpointLoaderResident,
"SetGlobalResidencyPolicy": SetGlobalResidencyPolicy,
"RegistrySnapshot": RegistrySnapshot,
"PinModelResidency": PinModelResidency,
"PinClipResidency": PinClipResidency,
"PinVAEResidency": PinVAEResidency,
"PreloadModelToGPU": PreloadModelToGPU,
"PreloadClipToGPU": PreloadClipToGPU,
"PreloadVAEToGPU": PreloadVAEToGPU,
"EvictModelFromGPU": EvictModelFromGPU,
"EvictClipFromGPU": EvictClipFromGPU,
"EvictVAEFromGPU": EvictVAEFromGPU,
"ReportModelResidency": ReportModelResidency,
"ReportClipResidency": ReportClipResidency,
"ReportVAEResidency": ReportVAEResidency,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"DiffusionModelSelectorResident": "Diffusion Model Selector Resident",
"DiffusionModelLoaderResident": "Diffusion Model Loader Resident",
"CheckpointLoaderResident": "Checkpoint Loader Resident",
"SetGlobalResidencyPolicy": "Set Global Residency Policy",
"RegistrySnapshot": "Registry Snapshot",
"PinModelResidency": "Pin Model Residency",
"PinClipResidency": "Pin CLIP Residency",
"PinVAEResidency": "Pin VAE Residency",
"PreloadModelToGPU": "Preload Model To GPU",
"PreloadClipToGPU": "Preload CLIP To GPU",
"PreloadVAEToGPU": "Preload VAE To GPU",
"EvictModelFromGPU": "Evict Model From GPU",
"EvictClipFromGPU": "Evict CLIP From GPU",
"EvictVAEFromGPU": "Evict VAE From GPU",
"ReportModelResidency": "Report Model Residency",
"ReportClipResidency": "Report CLIP Residency",
"ReportVAEResidency": "Report VAE Residency",
}
+445
View File
@@ -0,0 +1,445 @@
from __future__ import annotations
import functools
import logging
import os
from typing import Any, Callable
import torch
from safetensors import safe_open
from .residency import (
KIND_CHECKPOINT,
KIND_CLIP,
KIND_CLIP_VISION,
KIND_CONTROLNET,
KIND_MODEL,
KIND_VAE,
REGISTRY,
)
_LOG = logging.getLogger(__name__)
_PATCHED = False
_ORIGINALS: dict[str, Callable[..., Any]] = {}
def _normalize_device(device: Any | None) -> torch.device | None:
if device is None:
return None
if isinstance(device, torch.device):
return device
try:
return torch.device(device)
except Exception:
return None
def _device_string(device: torch.device | None) -> str:
if device is None:
return "auto"
return str(device)
def _safe_open_device_arg(device: torch.device) -> Any:
if device.type == "cuda":
return device.index if device.index is not None else torch.cuda.current_device()
if device.type == "cpu":
return "cpu"
return device.type
def _copy_tensor_if_needed(
tensor: torch.Tensor,
target_device: torch.device,
*,
force_copy: bool = False,
) -> torch.Tensor:
if tensor.device == target_device and not force_copy:
return tensor
return tensor.to(device=target_device, copy=True)
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:
return ctx.explicit_device, ctx.kind, ctx.note
return REGISTRY.explicit_load_device(kind=kind, source_path=source_path), kind, None
def _record_generic_load(
*,
path: str,
method: str,
requested_device: torch.device | None,
actual_device: str,
note: str | None = None,
error: str | None = None,
) -> None:
ctx = REGISTRY.current_context()
kind = ctx.kind if ctx is not None else "unknown"
REGISTRY.record_load(
path=path,
kind=kind,
method=method,
requested_device=_device_string(requested_device),
actual_device=actual_device,
note=note or (ctx.note if ctx is not None else None),
error=error,
)
def _patched_load_torch_file(ckpt, safe_load=False, device=None, return_metadata=False):
import comfy.memory_management
import comfy.utils as comfy_utils
requested_device = _normalize_device(device)
ctx = REGISTRY.current_context()
if requested_device is None and ctx is not None and ctx.explicit_device is not None:
requested_device = ctx.explicit_device
if requested_device is None:
requested_device = torch.device("cpu")
metadata = None
lowered = str(ckpt).lower()
if lowered.endswith((".safetensors", ".sft")):
try:
if comfy.memory_management.aimdo_enabled and requested_device.type == "cpu":
sd, metadata = comfy_utils.load_safetensors(ckpt)
method = "safetensors_aimdo_cpu"
if not return_metadata:
metadata = None
_record_generic_load(
path=ckpt,
method=method,
requested_device=requested_device,
actual_device="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,
requested_device=requested_device,
actual_device=actual_device,
)
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):
if "HeaderTooLarge" in message:
raise ValueError(
f"{message}\n\nFile path: {ckpt}\n\n"
"The safetensors file is corrupt or invalid. Make sure this is actually a "
"safetensors file and not a ckpt or pt or other filetype."
) from exc
if "MetadataIncompleteBuffer" in message:
raise ValueError(
f"{message}\n\nFile path: {ckpt}\n\n"
"The safetensors file is corrupt/incomplete. Check the file size and make sure "
"you have copied/downloaded it correctly."
) from exc
_record_generic_load(
path=ckpt,
method="safetensors_load_failed",
requested_device=requested_device,
actual_device="error",
error=str(exc),
)
raise
torch_args = {}
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),
)
if "state_dict" in pl_sd:
sd = pl_sd["state_dict"]
else:
if len(pl_sd) == 1:
key = list(pl_sd.keys())[0]
sd = pl_sd[key]
if not isinstance(sd, dict):
sd = pl_sd
else:
sd = pl_sd
return (sd, metadata) if return_metadata else sd
def _bind_checkpoint_outputs(result, source_path: str) -> None:
if not result:
return
model = result[0] if len(result) > 0 else None
clip = result[1] if len(result) > 1 else None
vae = result[2] if len(result) > 2 else None
if model is not None:
REGISTRY.bind_object(model, source_path=source_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=source_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=source_path, kind=KIND_VAE, note="checkpoint vae")
def _bind_diffusers_outputs(result, source_path: str) -> None:
if not result:
return
model = result[0] if len(result) > 0 else None
clip = result[1] if len(result) > 1 else None
vae = result[2] if len(result) > 2 else None
if model is not None:
REGISTRY.bind_object(model, source_path=source_path, kind=KIND_MODEL, note="diffusers model")
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="diffusers clip")
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="diffusers vae")
def _wrap_with_load_context(kind: str, path_arg_index: int = 0, bind_output: Callable[[Any, str], None] | None = None):
def decorator(func: Callable[..., Any]) -> Callable[..., Any]:
@functools.wraps(func)
def wrapper(*args, **kwargs):
source_path = None
if len(args) > path_arg_index:
source_path = args[path_arg_index]
explicit_device = REGISTRY.explicit_load_device(kind=kind, source_path=source_path)
with REGISTRY.load_context(kind=kind, source_path=source_path, explicit_device=explicit_device):
result = func(*args, **kwargs)
if bind_output is not None and source_path is not None:
bind_output(result, source_path)
return result
return wrapper
return decorator
def _wrap_load_clip(func: Callable[..., Any]) -> Callable[..., Any]:
@functools.wraps(func)
def wrapper(*args, **kwargs):
ckpt_paths = args[0] if args else kwargs.get("ckpt_paths")
source_path = None
if isinstance(ckpt_paths, (list, tuple)) and ckpt_paths:
source_path = ckpt_paths[0]
explicit_device = REGISTRY.explicit_load_device(kind=KIND_CLIP, source_path=source_path)
with REGISTRY.load_context(kind=KIND_CLIP, source_path=source_path, explicit_device=explicit_device):
clip = func(*args, **kwargs)
if clip is not None and getattr(clip, "patcher", None) is not None and source_path is not None:
REGISTRY.bind_object(clip.patcher, source_path=source_path, kind=KIND_CLIP)
return clip
return wrapper
def _wrap_load_models_gpu(func: Callable[..., Any]) -> Callable[..., Any]:
@functools.wraps(func)
def wrapper(models, *args, **kwargs):
result = func(models, *args, **kwargs)
for model in list(models):
REGISTRY.touch(model)
REGISTRY.refresh_runtime_state()
return result
return wrapper
def _wrap_free_memory(func: Callable[..., Any]) -> Callable[..., Any]:
@functools.wraps(func)
def wrapper(memory_required, device, keep_loaded=None, *args, **kwargs):
import comfy.model_management as model_management
keep_loaded = list(keep_loaded or [])
sticky_wrappers = []
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)
if device is not None and sticky_wrappers:
try:
free_after = 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"
)
unloaded = func(memory_required, device, keep_loaded, *args, **kwargs)
REGISTRY.refresh_runtime_state()
return unloaded
return wrapper
def _remember_original(key: str, value: Callable[..., Any]) -> Callable[..., Any]:
return _ORIGINALS.setdefault(key, value)
def _patch_model_management_devices() -> None:
import comfy.model_management as model_management
def wrap_device_func(name: str) -> None:
key = f"model_management.{name}"
original = _remember_original(key, getattr(model_management, name))
if getattr(model_management, name) is not original:
return
@functools.wraps(original)
def wrapper(*args, **kwargs):
result = original(*args, **kwargs)
if not REGISTRY.wants_gpu_offload(name):
return result
gpu_device = model_management.get_torch_device()
if getattr(gpu_device, "type", None) == "cpu":
return result
return gpu_device
setattr(model_management, name, wrapper)
for name in (
"unet_offload_device",
"text_encoder_offload_device",
"vae_offload_device",
"text_encoder_device",
"vae_device",
"unet_inital_load_device",
):
if hasattr(model_management, name):
wrap_device_func(name)
original_free_memory = _remember_original("model_management.free_memory", model_management.free_memory)
if model_management.free_memory is original_free_memory:
model_management.free_memory = _wrap_free_memory(original_free_memory)
original_load_models_gpu = _remember_original("model_management.load_models_gpu", model_management.load_models_gpu)
if model_management.load_models_gpu is original_load_models_gpu:
model_management.load_models_gpu = _wrap_load_models_gpu(original_load_models_gpu)
def install_patches() -> None:
global _PATCHED
if _PATCHED:
return
import comfy.clip_vision as clip_vision
import comfy.controlnet as controlnet
import comfy.diffusers_load as diffusers_load
import comfy.model_management as model_management
import comfy.sd as comfy_sd
import comfy.utils as comfy_utils
original_load_torch_file = _remember_original("utils.load_torch_file", comfy_utils.load_torch_file)
if comfy_utils.load_torch_file is original_load_torch_file:
comfy_utils.load_torch_file = _patched_load_torch_file
if hasattr(clip_vision, "load_torch_file"):
original_clip_vision_load_torch_file = _remember_original("clip_vision.load_torch_file", clip_vision.load_torch_file)
if clip_vision.load_torch_file is original_clip_vision_load_torch_file:
clip_vision.load_torch_file = comfy_utils.load_torch_file
_patch_model_management_devices()
original_load_checkpoint_guess_config = _remember_original(
"sd.load_checkpoint_guess_config",
comfy_sd.load_checkpoint_guess_config,
)
if comfy_sd.load_checkpoint_guess_config is original_load_checkpoint_guess_config:
comfy_sd.load_checkpoint_guess_config = _wrap_with_load_context(
KIND_CHECKPOINT,
path_arg_index=0,
bind_output=_bind_checkpoint_outputs,
)(original_load_checkpoint_guess_config)
original_load_diffusion_model = _remember_original("sd.load_diffusion_model", comfy_sd.load_diffusion_model)
if comfy_sd.load_diffusion_model is original_load_diffusion_model:
comfy_sd.load_diffusion_model = _wrap_with_load_context(
KIND_MODEL,
path_arg_index=0,
bind_output=lambda model, source_path: model is not None
and REGISTRY.bind_object(model, source_path=source_path, kind=KIND_MODEL),
)(original_load_diffusion_model)
original_load_clip = _remember_original("sd.load_clip", comfy_sd.load_clip)
if comfy_sd.load_clip is original_load_clip:
comfy_sd.load_clip = _wrap_load_clip(original_load_clip)
original_clip_vision_load = _remember_original("clip_vision.load", clip_vision.load)
if clip_vision.load is original_clip_vision_load:
clip_vision.load = _wrap_with_load_context(
KIND_CLIP_VISION,
path_arg_index=0,
bind_output=lambda result, source_path: result is not None
and getattr(result, "patcher", None) is not None
and REGISTRY.bind_object(result.patcher, source_path=source_path, kind=KIND_CLIP_VISION),
)(original_clip_vision_load)
original_load_controlnet = _remember_original("controlnet.load_controlnet", controlnet.load_controlnet)
if controlnet.load_controlnet is original_load_controlnet:
controlnet.load_controlnet = _wrap_with_load_context(KIND_CONTROLNET, path_arg_index=0)(original_load_controlnet)
original_load_diffusers = _remember_original("diffusers_load.load_diffusers", diffusers_load.load_diffusers)
if diffusers_load.load_diffusers is original_load_diffusers:
diffusers_load.load_diffusers = _wrap_with_load_context(
KIND_CHECKPOINT,
path_arg_index=0,
bind_output=_bind_diffusers_outputs,
)(original_load_diffusers)
REGISTRY.refresh_runtime_state()
_PATCHED = True
_LOG.info("GPU Resident Loader: monkey patches active on ComfyUI loader and residency paths")
+19
View File
@@ -0,0 +1,19 @@
[project]
name = "comfyui-gpu-resident-loader"
version = "0.1.0"
description = "ComfyUI custom nodes and startup patches for GPU-resident model loading and sticky VRAM residency"
readme = "README.md"
license = { text = "GPL-3.0-or-later" }
requires-python = ">=3.10"
dependencies = [
"safetensors>=0.4.3",
]
[tool.setuptools]
py-modules = [
"startup",
"patches",
"residency",
"kj_loader",
"nodes",
]
+1
View File
@@ -0,0 +1 @@
safetensors>=0.4.3
+427
View File
@@ -0,0 +1,427 @@
from __future__ import annotations
import contextlib
import contextvars
import dataclasses
import json
import logging
import os
import threading
import time
import weakref
from collections.abc import Iterator
from typing import Any
import torch
_LOG = logging.getLogger(__name__)
_LOAD_CONTEXT: contextvars.ContextVar[LoadContext | None] = contextvars.ContextVar(
"gpu_resident_loader_load_context",
default=None,
)
POLICIES = ("legacy", "balanced", "prefer_gpu", "sticky_gpu")
KIND_MODEL = "model"
KIND_CLIP = "clip"
KIND_VAE = "vae"
KIND_CHECKPOINT = "checkpoint"
KIND_CLIP_VISION = "clip_vision"
KIND_CONTROLNET = "controlnet"
def _now() -> float:
return time.time()
@dataclasses.dataclass(slots=True)
class LoadContext:
kind: str
source_path: str | None = None
explicit_device: torch.device | None = None
note: str | None = None
@dataclasses.dataclass(slots=True)
class LoadReport:
path: str
kind: str
method: str
requested_device: str
actual_device: str
timestamp: float = dataclasses.field(default_factory=_now)
note: str | None = None
error: str | None = None
def as_dict(self) -> dict[str, Any]:
return {
"path": self.path,
"kind": self.kind,
"method": self.method,
"requested_device": self.requested_device,
"actual_device": self.actual_device,
"timestamp": self.timestamp,
"note": self.note,
"error": self.error,
}
@dataclasses.dataclass(slots=True)
class ResidencyEntry:
entry_id: str
kind: str
source_path: str
sticky: bool
priority: int
created_at: float = dataclasses.field(default_factory=_now)
last_touched: float = dataclasses.field(default_factory=_now)
loaded_bytes: int = 0
total_bytes: int = 0
load_device: str | None = None
offload_device: str | None = None
current_device: str | None = None
last_method: str | None = None
last_report: dict[str, Any] | None = None
notes: list[str] = dataclasses.field(default_factory=list)
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
def object(self) -> Any | None:
return None if self.object_ref is None else self.object_ref()
def as_dict(self) -> dict[str, Any]:
basename = os.path.basename(self.source_path) if self.source_path else None
return {
"entry_id": self.entry_id,
"kind": self.kind,
"source_path": self.source_path,
"basename": basename,
"sticky": self.sticky,
"priority": self.priority,
"created_at": self.created_at,
"last_touched": self.last_touched,
"loaded_bytes": self.loaded_bytes,
"total_bytes": self.total_bytes,
"load_device": self.load_device,
"offload_device": self.offload_device,
"current_device": self.current_device,
"last_method": self.last_method,
"last_report": self.last_report,
"notes": list(self.notes),
"alive": self.is_alive(),
}
class ResidencyRegistry:
def __init__(self) -> None:
self._lock = threading.RLock()
self._entries: dict[str, ResidencyEntry] = {}
self._reports_by_path: dict[str, LoadReport] = {}
self._path_to_entry: dict[tuple[str, str], str] = {}
self._object_to_entry: weakref.WeakKeyDictionary[Any, str] = weakref.WeakKeyDictionary()
self._policy = self._default_policy()
def _default_policy(self) -> str:
env_value = os.environ.get("COMFYUI_GPU_RESIDENT_POLICY", "").strip().lower()
if env_value in POLICIES:
return env_value
try:
from comfy.cli_args import args
except Exception:
return "prefer_gpu"
if getattr(args, "gpu_only", False):
return "sticky_gpu"
if getattr(args, "highvram", False):
return "sticky_gpu"
return "prefer_gpu"
def get_policy(self) -> str:
with self._lock:
return self._policy
def set_policy(self, policy: str) -> str:
normalized = str(policy).strip().lower()
if normalized not in POLICIES:
raise ValueError(f"Unsupported residency policy: {policy}")
with self._lock:
self._policy = normalized
_LOG.info("GPU Resident Loader: policy set to %s", normalized)
return normalized
def wants_gpu_ingest(self, kind: str | None = None) -> bool:
policy = self.get_policy()
return policy in {"prefer_gpu", "sticky_gpu"}
def wants_gpu_offload(self, kind: str | None = None) -> bool:
policy = self.get_policy()
return policy in {"prefer_gpu", "sticky_gpu"}
def autopin_on_bind(self, kind: str | None = None) -> bool:
return self.get_policy() == "sticky_gpu"
def explicit_load_device(self, kind: str, source_path: str | None = None) -> torch.device | None:
if not self.wants_gpu_ingest(kind):
return None
try:
import comfy.model_management as model_management
except Exception:
return None
dev = model_management.get_torch_device()
if getattr(dev, "type", None) == "cpu":
return None
return dev
@contextlib.contextmanager
def load_context(
self,
*,
kind: str,
source_path: str | None = None,
explicit_device: torch.device | None = None,
note: str | None = None,
) -> Iterator[None]:
token = _LOAD_CONTEXT.set(
LoadContext(
kind=kind,
source_path=source_path,
explicit_device=explicit_device,
note=note,
)
)
try:
yield
finally:
_LOAD_CONTEXT.reset(token)
def current_context(self) -> LoadContext | None:
return _LOAD_CONTEXT.get()
def record_load(
self,
*,
path: str,
kind: str,
method: str,
requested_device: str,
actual_device: str,
note: str | None = None,
error: str | None = None,
) -> LoadReport:
report = LoadReport(
path=path,
kind=kind,
method=method,
requested_device=requested_device,
actual_device=actual_device,
note=note,
error=error,
)
with self._lock:
self._reports_by_path[path] = report
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:
entry.last_method = method
entry.last_report = report.as_dict()
entry.last_touched = _now()
entry.current_device = actual_device
return report
def latest_report_for_path(self, path: str | None) -> LoadReport | None:
if not path:
return None
with self._lock:
return self._reports_by_path.get(path)
def _make_entry_id(self, kind: str, source_path: str) -> str:
basename = os.path.basename(source_path) or "anonymous"
return f"{kind}:{basename}:{len(self._entries) + 1}"
def bind_object(
self,
obj: Any,
*,
source_path: str,
kind: str,
sticky: bool | None = None,
priority: int = 0,
note: str | 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)
if entry_id is None:
try:
entry_id = self._object_to_entry.get(obj)
except TypeError:
entry_id = None
if entry_id is not None and entry_id in self._entries:
entry = self._entries[entry_id]
old_key = (entry.kind, entry.source_path)
else:
entry_id = self._make_entry_id(kind, source_path)
entry = ResidencyEntry(
entry_id=entry_id,
kind=kind,
source_path=source_path,
sticky=self.autopin_on_bind(kind) if sticky is None else bool(sticky),
priority=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,
)
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),
)
entry.sticky = entry.sticky if sticky is None else bool(sticky)
entry.priority = int(priority)
entry.source_path = source_path
entry.kind = kind
new_key = (entry.kind, entry.source_path)
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
entry.last_touched = _now()
if note:
entry.notes.append(note)
report = self._reports_by_path.get(source_path)
if report is not None:
entry.last_method = report.method
entry.last_report = report.as_dict()
entry.current_device = report.actual_device
self.refresh_runtime_state()
return entry
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:
return self._entries.get(entry_id)
def set_sticky(self, obj: Any, sticky: bool, priority: int | None = None) -> ResidencyEntry | None:
entry = self.entry_for_object(obj)
if entry is None:
return None
with self._lock:
entry.sticky = bool(sticky)
if priority is not None:
entry.priority = int(priority)
entry.last_touched = _now()
return entry
def touch(self, obj: Any) -> None:
entry = self.entry_for_object(obj)
if entry is None:
return
with self._lock:
entry.last_touched = _now()
def sticky_loaded_wrappers(self, device: torch.device | None) -> list[Any]:
try:
import comfy.model_management as model_management
except Exception:
return []
output: list[Any] = []
with self._lock:
for loaded in list(model_management.current_loaded_models):
if device is not None and loaded.device != device:
continue
entry = self.entry_for_object(loaded.model)
if entry is not None and entry.sticky:
output.append(loaded)
return output
def refresh_runtime_state(self) -> None:
try:
import comfy.model_management as model_management
except Exception:
return
with self._lock:
for entry in self._entries.values():
if not entry.is_alive():
continue
obj = entry.object()
if obj is None:
continue
entry.loaded_bytes = 0
load_device = getattr(obj, "load_device", None)
offload_device = getattr(obj, "offload_device", None)
if load_device is not None:
entry.load_device = str(load_device)
if offload_device is not None:
entry.offload_device = str(offload_device)
entry.current_device = entry.offload_device
for loaded in list(model_management.current_loaded_models):
entry = self.entry_for_object(loaded.model)
if entry is None:
continue
entry.loaded_bytes = int(loaded.model_loaded_memory())
entry.total_bytes = int(loaded.model_memory())
entry.current_device = str(loaded.device)
entry.last_touched = _now()
def snapshot(self) -> list[dict[str, Any]]:
self.refresh_runtime_state()
with self._lock:
items = [entry.as_dict() for entry in self._entries.values()]
items.sort(
key=lambda item: (
not item["sticky"],
item["kind"],
item["basename"] or "",
)
)
return items
def snapshot_json(self) -> str:
payload = {
"policy": self.get_policy(),
"entries": self.snapshot(),
}
return json.dumps(payload, indent=2, sort_keys=True)
REGISTRY = ResidencyRegistry()
@@ -0,0 +1,69 @@
#!/usr/bin/env python3
from __future__ import annotations
import argparse
import os
from pathlib import Path
import torch
from safetensors.torch import save_file
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Convert a PyTorch checkpoint into a safetensors state dict file.")
parser.add_argument("--input", required=True, help="Input .ckpt/.pt/.bin/.pth path")
parser.add_argument("--output", required=True, help="Output .safetensors path")
parser.add_argument(
"--state-dict-key",
default="state_dict",
help="Top-level key to extract when the checkpoint is a wrapper object. Default: state_dict",
)
parser.add_argument(
"--allow-non-tensor-values",
action="store_true",
help="Ignore non-tensor items instead of failing.",
)
return parser.parse_args()
def main() -> int:
args = parse_args()
input_path = Path(args.input)
output_path = Path(args.output)
if not input_path.exists():
raise FileNotFoundError(f"Input file does not exist: {input_path}")
checkpoint = torch.load(str(input_path), map_location="cpu", weights_only=True)
if isinstance(checkpoint, dict) and args.state_dict_key in checkpoint and isinstance(checkpoint[args.state_dict_key], dict):
state_dict = checkpoint[args.state_dict_key]
elif isinstance(checkpoint, dict):
state_dict = checkpoint
else:
raise TypeError("Checkpoint is not a dictionary and no state_dict could be extracted.")
output_tensors = {}
skipped = []
for key, value in state_dict.items():
if isinstance(value, torch.Tensor):
output_tensors[key] = value.detach().cpu().contiguous()
elif args.allow_non_tensor_values:
skipped.append(key)
else:
raise TypeError(f"Key {key!r} is not a tensor. Re-run with --allow-non-tensor-values to skip it.")
output_path.parent.mkdir(parents=True, exist_ok=True)
metadata = {
"converted_from": os.path.basename(str(input_path)),
"converter": "ComfyUI-GPU-Resident-Loader",
}
save_file(output_tensors, str(output_path), metadata=metadata)
print(f"Wrote {len(output_tensors)} tensors to {output_path}")
if skipped:
print(f"Skipped {len(skipped)} non-tensor keys")
return 0
if __name__ == "__main__":
raise SystemExit(main())
+15
View File
@@ -0,0 +1,15 @@
import logging
from .patches import install_patches as _install_patches
_LOG = logging.getLogger(__name__)
_PATCHES_INSTALLED = False
def install_patches() -> None:
global _PATCHES_INSTALLED
if _PATCHES_INSTALLED:
return
_install_patches()
_PATCHES_INSTALLED = True
_LOG.info("GPU Resident Loader: startup patches installed")