Merge pull request #1 from xmarre/codex/import-zip-implementation
[codex] Import GPU resident loader implementation
This commit is contained in:
+23
@@ -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/
|
||||
@@ -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.
|
||||
@@ -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
@@ -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
|
||||
@@ -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
@@ -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")
|
||||
@@ -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",
|
||||
]
|
||||
@@ -0,0 +1 @@
|
||||
safetensors>=0.4.3
|
||||
+427
@@ -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
@@ -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")
|
||||
Reference in New Issue
Block a user