From d0fe3f5058e6ca94c45eea00a4bb470007defb4f Mon Sep 17 00:00:00 2001 From: xmarre Date: Sun, 12 Apr 2026 03:46:23 +0200 Subject: [PATCH 1/4] Import GPU resident loader implementation --- .gitignore | 23 ++ README.md | 175 ++++++++ __init__.py | 16 + kj_loader.py | 386 +++++++++++++++++ nodes.py | 361 ++++++++++++++++ patches.py | 413 +++++++++++++++++++ pyproject.toml | 19 + requirements.txt | 1 + residency.py | 387 +++++++++++++++++ scripts/convert_checkpoint_to_safetensors.py | 69 ++++ startup.py | 15 + 11 files changed, 1865 insertions(+) create mode 100644 .gitignore create mode 100644 README.md create mode 100644 __init__.py create mode 100644 kj_loader.py create mode 100644 nodes.py create mode 100644 patches.py create mode 100644 pyproject.toml create mode 100644 requirements.txt create mode 100644 residency.py create mode 100644 scripts/convert_checkpoint_to_safetensors.py create mode 100644 startup.py diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..8fba17d --- /dev/null +++ b/.gitignore @@ -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/ diff --git a/README.md b/README.md new file mode 100644 index 0000000..44e421f --- /dev/null +++ b/README.md @@ -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. diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..de038e4 --- /dev/null +++ b/__init__.py @@ -0,0 +1,16 @@ +import os +import sys + +CURRENT_DIR = os.path.dirname(os.path.abspath(__file__)) +if CURRENT_DIR not in sys.path: + sys.path.insert(0, CURRENT_DIR) + +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"] diff --git a/kj_loader.py b/kj_loader.py new file mode 100644 index 0000000..faae2d2 --- /dev/null +++ b/kj_loader.py @@ -0,0 +1,386 @@ +from __future__ import annotations + +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_MODEL, 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: + flag = getattr(torch.backends.cuda.matmul, "allow_fp16_accumulation", None) + if flag is None: + raise RuntimeError( + "Failed to set fp16 accumulation. This requires a PyTorch build exposing " + "torch.backends.cuda.matmul.allow_fp16_accumulation." + ) + torch.backends.cuda.matmul.allow_fp16_accumulation = bool(enabled) + + +def get_sage_func(sage_attention: str, allow_compile: bool = False): + _LOG.info("GPU Resident Loader: using sage attention mode %s", sage_attention) + from sageattention import sageattn + + if sage_attention == "auto": + 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, + ): + _set_cublas_linear(patch_cublaslinear) + _set_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, + ): + _set_cublas_linear(patch_cublaslinear) + _set_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 diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..1d9f778 --- /dev/null +++ b/nodes.py @@ -0,0 +1,361 @@ +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) + 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) + 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", +} diff --git a/patches.py b/patches.py new file mode 100644 index 0000000..f987995 --- /dev/null +++ b/patches.py @@ -0,0 +1,413 @@ +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) -> torch.Tensor: + if tensor.device == target_device: + 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) + sd[key] = tensor + if return_metadata: + metadata = handle.metadata() + + actual_device = next(iter(sd.values())).device.type if sd else requested_device.type + 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=str(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: + pass + + 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 _patch_model_management_devices() -> None: + import comfy.model_management as model_management + + def wrap_device_func(name: str) -> None: + original = getattr(model_management, name) + _ORIGINALS[f"model_management.{name}"] = original + + @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) + + _ORIGINALS["model_management.free_memory"] = model_management.free_memory + model_management.free_memory = _wrap_free_memory(model_management.free_memory) + + _ORIGINALS["model_management.load_models_gpu"] = model_management.load_models_gpu + model_management.load_models_gpu = _wrap_load_models_gpu(model_management.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 + + _ORIGINALS["utils.load_torch_file"] = comfy_utils.load_torch_file + comfy_utils.load_torch_file = _patched_load_torch_file + if hasattr(clip_vision, "load_torch_file"): + clip_vision.load_torch_file = comfy_utils.load_torch_file + + _patch_model_management_devices() + + _ORIGINALS["sd.load_checkpoint_guess_config"] = comfy_sd.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, + )(comfy_sd.load_checkpoint_guess_config) + + _ORIGINALS["sd.load_diffusion_model"] = comfy_sd.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), + )(comfy_sd.load_diffusion_model) + + _ORIGINALS["sd.load_clip"] = comfy_sd.load_clip + comfy_sd.load_clip = _wrap_load_clip(comfy_sd.load_clip) + + _ORIGINALS["clip_vision.load"] = 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), + )(clip_vision.load) + + _ORIGINALS["controlnet.load_controlnet"] = controlnet.load_controlnet + controlnet.load_controlnet = _wrap_with_load_context(KIND_CONTROLNET, path_arg_index=0)(controlnet.load_controlnet) + + _ORIGINALS["diffusers_load.load_diffusers"] = diffusers_load.load_diffusers + diffusers_load.load_diffusers = _wrap_with_load_context( + KIND_CHECKPOINT, + path_arg_index=0, + bind_output=_bind_diffusers_outputs, + )(diffusers_load.load_diffusers) + + REGISTRY.refresh_runtime_state() + _PATCHED = True + _LOG.info("GPU Resident Loader: monkey patches active on ComfyUI loader and residency paths") diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..7724416 --- /dev/null +++ b/pyproject.toml @@ -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", +] diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..fff1530 --- /dev/null +++ b/requirements.txt @@ -0,0 +1 @@ +safetensors>=0.4.3 diff --git a/residency.py b/residency.py new file mode 100644 index 0000000..bd58bd9 --- /dev/null +++ b/residency.py @@ -0,0 +1,387 @@ +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._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: + entry_id = getattr(obj, "__gpu_resident_loader_entry_id__", None) + if entry_id is not None and entry_id in self._entries: + entry = self._entries[entry_id] + 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) + except Exception: + pass + + entry.object_ref = weakref.ref(obj) + entry.sticky = entry.sticky if sticky is None else bool(sticky) + entry.priority = int(priority) + entry.source_path = source_path + entry.kind = kind + 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: + 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 + 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) + + 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() diff --git a/scripts/convert_checkpoint_to_safetensors.py b/scripts/convert_checkpoint_to_safetensors.py new file mode 100644 index 0000000..7d55a0b --- /dev/null +++ b/scripts/convert_checkpoint_to_safetensors.py @@ -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()) diff --git a/startup.py b/startup.py new file mode 100644 index 0000000..4bad7a7 --- /dev/null +++ b/startup.py @@ -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") From 669e6014fdac42d3ab61fc7f0ab40f8d5d64d501 Mon Sep 17 00:00:00 2001 From: xmarre Date: Sun, 12 Apr 2026 04:14:07 +0200 Subject: [PATCH 2/4] Fix package imports and patch guards --- __init__.py | 11 +---- kj_loader.py | 24 +++++++++-- nodes.py | 8 +++- patches.py | 113 +++++++++++++++++++++++++++++++-------------------- residency.py | 17 ++++++-- startup.py | 2 +- 6 files changed, 113 insertions(+), 62 deletions(-) diff --git a/__init__.py b/__init__.py index de038e4..d01feb8 100644 --- a/__init__.py +++ b/__init__.py @@ -1,12 +1,5 @@ -import os -import sys - -CURRENT_DIR = os.path.dirname(os.path.abspath(__file__)) -if CURRENT_DIR not in sys.path: - sys.path.insert(0, CURRENT_DIR) - -from startup import install_patches -from nodes import ( +from .startup import install_patches +from .nodes import ( NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS, ) diff --git a/kj_loader.py b/kj_loader.py index faae2d2..b2c19fc 100644 --- a/kj_loader.py +++ b/kj_loader.py @@ -11,7 +11,7 @@ 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_MODEL, REGISTRY +from .residency import KIND_CHECKPOINT, KIND_CLIP, KIND_MODEL, KIND_VAE, REGISTRY _LOG = logging.getLogger(__name__) @@ -44,12 +44,28 @@ def _set_cublas_linear(enabled: bool) -> None: 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: - raise RuntimeError( - "Failed to set fp16 accumulation. This requires a PyTorch build exposing " - "torch.backends.cuda.matmul.allow_fp16_accumulation." + 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) diff --git a/nodes.py b/nodes.py index 1d9f778..5289978 100644 --- a/nodes.py +++ b/nodes.py @@ -5,8 +5,8 @@ from typing import Any import comfy.model_management as model_management -from kj_loader import CheckpointLoaderResident, DiffusionModelLoaderResident, DiffusionModelSelectorResident -from residency import REGISTRY +from .kj_loader import CheckpointLoaderResident, DiffusionModelLoaderResident, DiffusionModelSelectorResident +from .residency import REGISTRY def _entry_report_json(obj: Any) -> str: @@ -133,6 +133,8 @@ class PinClipResidency: 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,) @@ -154,6 +156,8 @@ class PinVAEResidency: 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,) diff --git a/patches.py b/patches.py index f987995..138a4ef 100644 --- a/patches.py +++ b/patches.py @@ -8,7 +8,7 @@ from typing import Any, Callable import torch from safetensors import safe_open -from residency import ( +from .residency import ( KIND_CHECKPOINT, KIND_CLIP, KIND_CLIP_VISION, @@ -155,8 +155,15 @@ def _patched_load_torch_file(ckpt, safe_load=False, device=None, return_metadata error=str(exc), ) return (sd, metadata) if return_metadata else sd - except Exception: - pass + 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] @@ -315,12 +322,18 @@ def _wrap_free_memory(func: Callable[..., Any]) -> Callable[..., Any]: 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: - original = getattr(model_management, name) - _ORIGINALS[f"model_management.{name}"] = original + 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): @@ -345,11 +358,13 @@ def _patch_model_management_devices() -> None: if hasattr(model_management, name): wrap_device_func(name) - _ORIGINALS["model_management.free_memory"] = model_management.free_memory - model_management.free_memory = _wrap_free_memory(model_management.free_memory) + 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) - _ORIGINALS["model_management.load_models_gpu"] = model_management.load_models_gpu - model_management.load_models_gpu = _wrap_load_models_gpu(model_management.load_models_gpu) + 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: @@ -364,49 +379,61 @@ def install_patches() -> None: import comfy.sd as comfy_sd import comfy.utils as comfy_utils - _ORIGINALS["utils.load_torch_file"] = comfy_utils.load_torch_file - comfy_utils.load_torch_file = _patched_load_torch_file + 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"): - clip_vision.load_torch_file = comfy_utils.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() - _ORIGINALS["sd.load_checkpoint_guess_config"] = comfy_sd.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, - )(comfy_sd.load_checkpoint_guess_config) + 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) - _ORIGINALS["sd.load_diffusion_model"] = comfy_sd.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), - )(comfy_sd.load_diffusion_model) + 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) - _ORIGINALS["sd.load_clip"] = comfy_sd.load_clip - comfy_sd.load_clip = _wrap_load_clip(comfy_sd.load_clip) + 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) - _ORIGINALS["clip_vision.load"] = 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), - )(clip_vision.load) + 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) - _ORIGINALS["controlnet.load_controlnet"] = controlnet.load_controlnet - controlnet.load_controlnet = _wrap_with_load_context(KIND_CONTROLNET, path_arg_index=0)(controlnet.load_controlnet) + 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) - _ORIGINALS["diffusers_load.load_diffusers"] = diffusers_load.load_diffusers - diffusers_load.load_diffusers = _wrap_with_load_context( - KIND_CHECKPOINT, - path_arg_index=0, - bind_output=_bind_diffusers_outputs, - )(diffusers_load.load_diffusers) + 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 diff --git a/residency.py b/residency.py index bd58bd9..e814b89 100644 --- a/residency.py +++ b/residency.py @@ -271,10 +271,21 @@ class ResidencyRegistry: self._path_to_entry[(kind, source_path)] = entry_id try: setattr(obj, "__gpu_resident_loader_entry_id__", entry_id) - except Exception: - pass + except (AttributeError, TypeError): + _LOG.debug( + "GPU Resident Loader: could not tag object %r with residency entry id %s", + type(obj), + entry_id, + ) - entry.object_ref = weakref.ref(obj) + try: + entry.object_ref = weakref.ref(obj) + 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 diff --git a/startup.py b/startup.py index 4bad7a7..145a6bf 100644 --- a/startup.py +++ b/startup.py @@ -1,6 +1,6 @@ import logging -from patches import install_patches as _install_patches +from .patches import install_patches as _install_patches _LOG = logging.getLogger(__name__) _PATCHES_INSTALLED = False From a2bff642072c49bb9cb90752998f50fed604f491 Mon Sep 17 00:00:00 2001 From: xmarre Date: Sun, 12 Apr 2026 04:35:13 +0200 Subject: [PATCH 3/4] Fix loader state restoration and residency bookkeeping --- kj_loader.py | 92 +++++++++++++++++++++++++++++++++------------------- patches.py | 11 +++++-- residency.py | 9 +++++ 3 files changed, 75 insertions(+), 37 deletions(-) diff --git a/kj_loader.py b/kj_loader.py index b2c19fc..348f513 100644 --- a/kj_loader.py +++ b/kj_loader.py @@ -1,5 +1,6 @@ from __future__ import annotations +import contextlib import logging from typing import Any @@ -69,11 +70,32 @@ def _set_fp16_accumulation(enabled: bool) -> None: 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) - from sageattention import sageattn 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": @@ -307,24 +329,25 @@ class DiffusionModelLoaderResident: enable_fp16_accumulation: bool, extra_state_dict: str | None = None, ): - _set_cublas_linear(patch_cublaslinear) - _set_fp16_accumulation(enable_fp16_accumulation) + 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) - 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 - 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) + 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,) @@ -375,25 +398,26 @@ class CheckpointLoaderResident: sage_attention: str, enable_fp16_accumulation: bool, ): - _set_cublas_linear(patch_cublaslinear) - _set_fp16_accumulation(enable_fp16_accumulation) + 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) - 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) - 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) + 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") diff --git a/patches.py b/patches.py index 138a4ef..dfd6868 100644 --- a/patches.py +++ b/patches.py @@ -49,8 +49,13 @@ def _safe_open_device_arg(device: torch.device) -> Any: return device.type -def _copy_tensor_if_needed(tensor: torch.Tensor, target_device: torch.device) -> torch.Tensor: - if tensor.device == target_device: +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) @@ -119,7 +124,7 @@ def _patched_load_torch_file(ckpt, safe_load=False, device=None, return_metadata 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) + tensor = _copy_tensor_if_needed(tensor, requested_device, force_copy=True) sd[key] = tensor if return_metadata: metadata = handle.metadata() diff --git a/residency.py b/residency.py index e814b89..284e847 100644 --- a/residency.py +++ b/residency.py @@ -255,9 +255,11 @@ class ResidencyRegistry: 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 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( @@ -290,6 +292,11 @@ class ResidencyRegistry: 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) @@ -358,12 +365,14 @@ class ResidencyRegistry: 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) From ffe93ec16dc7dbb0ff76cb9e01fd33daefebc137 Mon Sep 17 00:00:00 2001 From: xmarre Date: Sun, 12 Apr 2026 04:45:43 +0200 Subject: [PATCH 4/4] Fix telemetry and weakref entry lookup --- patches.py | 4 ++-- residency.py | 20 ++++++++++++++++++++ 2 files changed, 22 insertions(+), 2 deletions(-) diff --git a/patches.py b/patches.py index dfd6868..e10d3f1 100644 --- a/patches.py +++ b/patches.py @@ -129,13 +129,13 @@ def _patched_load_torch_file(ckpt, safe_load=False, device=None, return_metadata if return_metadata: metadata = handle.metadata() - actual_device = next(iter(sd.values())).device.type if sd else requested_device.type + 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=str(actual_device), + actual_device=actual_device, ) return (sd, metadata) if return_metadata else sd except Exception as exc: diff --git a/residency.py b/residency.py index 284e847..612dd1c 100644 --- a/residency.py +++ b/residency.py @@ -120,6 +120,7 @@ class ResidencyRegistry: 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: @@ -257,6 +258,11 @@ class ResidencyRegistry: 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) @@ -273,6 +279,10 @@ class ResidencyRegistry: 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", @@ -282,6 +292,11 @@ class ResidencyRegistry: 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( @@ -313,6 +328,11 @@ class ResidencyRegistry: 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: