Compare commits

...
Author SHA1 Message Date
xmarre f1519f8da3 Merge pull request #21 from xmarre/codex/disable-load-trim-highvram
Disable eager loader trim in sticky/highvram workflows
2026-04-17 10:08:00 +02:00
coderabbitai[bot] 50e8188e5a 📝 Add docstrings to codex/disable-load-trim-highvram
Docstrings generation was requested by @xmarre.

The following files were modified:

* `kj_loader.py`
2026-04-17 08:05:39 +00:00
xmarre 32dee8e647 Disable eager loader trim for sticky/highvram modes 2026-04-17 10:00:30 +02:00
xmarre 9d5d5e3b42 Merge pull request #20 from xmarre/codex/pr20-external-trim-fix
Fix external trim gate for sticky VAE
2026-04-16 08:34:26 +02:00
xmarre 026d8527f8 Make VAE preflight trim respect external opt-in 2026-04-16 08:26:30 +02:00
xmarre 84d10add70 Add second-chance external VAE trim 2026-04-16 08:17:59 +02:00
xmarre 29b0a34865 Fix external trim gate for sticky VAE 2026-04-16 08:13:41 +02:00
xmarre fd9f33f17f Merge pull request #19 from xmarre/codex/vae-preflight-model-load-fix
Budget sticky VAE preflight for model load
2026-04-16 07:37:53 +02:00
xmarre fc7c0f6946 Tighten sticky VAE load preflight accounting 2026-04-16 07:32:50 +02:00
xmarre d56a716418 Budget sticky VAE preflight for model load 2026-04-16 07:25:31 +02:00
xmarre 27aa644775 Merge pull request #18 from xmarre/codex/inpaint-vae-node-fallback
Handle sticky tiled fallback for inpaint VAE encodes
2026-04-16 06:50:48 +02:00
xmarre e345e5f9d0 Avoid caching bound tiled VAE methods 2026-04-16 06:46:49 +02:00
coderabbitai[bot]andCodeRabbit 43fce69abd fix: apply CodeRabbit auto-fixes
Fixed 1 file(s) based on 1 unresolved review comment.

Co-authored-by: CodeRabbit <noreply@coderabbit.ai>
2026-04-16 04:37:55 +00:00
xmarre bafd34d48b Route sticky inpaint VAE encodes through tiled entrypoint 2026-04-16 06:30:29 +02:00
xmarre cf6843fd70 Merge pull request #17 from xmarre/codex/face-detailer-tiled-vae-admission
Wrap tiled VAE memory admission under sticky GPU
2026-04-16 05:41:26 +02:00
xmarre 5c623b9cdb Fix tiled VAE wrapper review regressions 2026-04-16 05:35:27 +02:00
xmarre bd249781b3 Wrap tiled VAE memory admission under sticky GPU 2026-04-16 05:21:56 +02:00
xmarre 1d1fa53828 Merge pull request #16 from xmarre/codex/external-fallback-free-memory
[codex] Fix external free-memory fallback trim
2026-04-16 04:04:34 +02:00
xmarre 6342d4c3ad Protect related external fallback entries 2026-04-16 03:57:10 +02:00
xmarre f60a8f9972 Skip dynamic external fallback trim 2026-04-16 03:54:13 +02:00
coderabbitai[bot] 61ed3e16d4 📝 Add docstrings to codex/external-fallback-free-memory
Docstrings generation was requested by @xmarre.

The following files were modified:

* `cleanup.py`
* `patches.py`
2026-04-16 01:47:19 +00:00
xmarre d827213bb2 Fix external free-memory fallback trim 2026-04-16 03:41:44 +02:00
xmarre 8998be1d78 Merge pull request #15 from xmarre/codex/fix-vae-preload-boundary
Fix sticky VAE preflight at the preload boundary
2026-04-16 03:12:08 +02:00
xmarre 9d34f008ab Fix VAE preload boundary preflight 2026-04-16 03:06:11 +02:00
xmarre 1f282ea6f0 Merge pull request #14 from xmarre/codex/auto-vae-preflight
[codex] Add sticky VAE preflight tiling
2026-04-15 21:47:18 +02:00
xmarre 79b9565867 Add sticky VAE preflight tiling 2026-04-15 21:41:13 +02:00
xmarre 2eb0817c12 Merge pull request #13 from xmarre/codex/seedvr2-pr12-regression-fix
[codex] restore SeedVR2 cache reuse by default
2026-04-15 20:31:42 +02:00
coderabbitai[bot]andCodeRabbit 334390a739 fix: apply CodeRabbit auto-fixes
Fixed 1 file(s) based on 2 unresolved review comments.

Co-authored-by: CodeRabbit <noreply@coderabbit.ai>
2026-04-15 18:27:16 +00:00
xmarre 815590498b fix: close SeedVR2 integration install race 2026-04-15 20:17:27 +02:00
xmarre ac42debea7 fix: restore SeedVR2 cache reuse by default 2026-04-15 20:07:10 +02:00
xmarre baba29152f Update README with external GPU cache details
Enhance README to include details about external GPU model caches and clarify residency system functionalities.
2026-04-15 10:06:35 +02:00
xmarre 37ed1899ae Merge pull request #12 from xmarre/codex/seedvr2-external-cache
[codex] Integrate SeedVR2 external cache eviction
2026-04-15 09:18:39 +02:00
coderabbitai[bot]andCodeRabbit a56a64b62e fix: apply CodeRabbit auto-fixes
Fixed 1 file(s) based on 1 unresolved review comment.

Co-authored-by: CodeRabbit <noreply@coderabbit.ai>
2026-04-15 07:11:31 +00:00
xmarre 9719017900 Harden external state refresh and SeedVR2 install 2026-04-15 09:00:00 +02:00
xmarre cbe4825973 Harden SeedVR2 external cache binding 2026-04-15 08:49:25 +02:00
xmarre 103613ad4f Harden external residency failure handling 2026-04-15 08:36:19 +02:00
xmarre 9d714f896c Harden SeedVR2 external residency hooks 2026-04-15 08:28:55 +02:00
xmarre b0a4c6d7a2 Integrate SeedVR2 external cache eviction 2026-04-15 08:15:10 +02:00
xmarre 5e427f4196 Revise README for clarity and detail enhancement 2026-04-15 07:43:04 +02:00
xmarre e7e4df92bf Merge pull request #11 from xmarre/codex/fix-sticky-vae-headroom
[codex] Preserve sticky VAE headroom
2026-04-15 07:08:58 +02:00
xmarre a7dc557376 Preserve sticky VAE headroom 2026-04-15 07:03:35 +02:00
xmarre 3f1d995138 Merge pull request #10 from xmarre/codex/fix-face-detailer-vram-overflow
[codex] Fix sticky reclaim before Face Detailer VAE encode
2026-04-15 06:15:52 +02:00
xmarre 3bbc263a46 Fix sticky VRAM reclaim before VAE encode 2026-04-15 06:09:41 +02:00
xmarre 2dc41ab929 Merge pull request #9 from xmarre/codex/pr8-clone-conflict-unload
[codex] Fix clone-conflict VRAM reclamation
2026-04-15 05:30:11 +02:00
xmarre 7dd33f3675 Address PR review feedback 2026-04-15 05:25:27 +02:00
xmarre 4f36431bf2 Fix clone-conflict VRAM reclamation 2026-04-15 05:10:51 +02:00
xmarre 8716f77749 Merge pull request #8 from xmarre/codex/fix-same-device-vram-eviction
[codex] Fix same-device VRAM eviction
2026-04-15 04:07:02 +02:00
xmarre 420b27ba4f Fix same-device VRAM eviction 2026-04-15 04:01:28 +02:00
xmarre 342b80d0ec Merge pull request #7 from xmarre/codex/adaptive-resident-trim
Move VRAM trimming into the resident loader path
2026-04-15 03:21:32 +02:00
xmarre fcd4d80f7b Fallback when selected safetensors keys miss 2026-04-15 03:14:13 +02:00
xmarre 59ac769635 Treat empty component maps as unknown size 2026-04-15 03:09:04 +02:00
xmarre c121278c65 Thread adaptive trim to the target CUDA device 2026-04-15 03:05:32 +02:00
xmarre da3f105661 Move VRAM trimming into resident loaders 2026-04-15 02:57:40 +02:00
xmarre 1365161b49 Merge pull request #6 from xmarre/codex/resident-loader-two-fixes
Add selective checkpoint ingest and resident VRAM trim
2026-04-15 02:38:56 +02:00
xmarre cfb707f082 Fix residency entry lookup locking 2026-04-15 02:36:35 +02:00
xmarre 19e69a8ecc Address PR review comments 2026-04-15 02:29:52 +02:00
xmarre f7b8fe5cc9 Add selective checkpoint ingest and resident VRAM trim 2026-04-15 02:17:16 +02:00
xmarre 2643a1903b Merge pull request #5 from xmarre/codex/narrow-resident-loader-hot-path
[codex] Narrow resident loader hot paths
2026-04-13 06:28:08 +02:00
xmarre d11d8e0a25 Fix loader-keyed load report attribution 2026-04-13 06:24:27 +02:00
xmarre df9128e9ee Fix resident loader reuse edge cases 2026-04-13 06:16:06 +02:00
xmarre 9f4b27eb23 Narrow resident loader hot paths 2026-04-13 06:00:19 +02:00
xmarre b94fd0b3f3 Merge pull request #4 from xmarre/codex/fix-loader-extra-unet-filter
Filter extra diffusion loader state dicts to UNet keys
2026-04-12 10:39:55 +02:00
xmarre 1b80a6a82b Preserve sparse extra UNet overrides 2026-04-12 10:36:40 +02:00
xmarre 6d4bd6c7c4 Filter extra diffusion weights to UNet keys 2026-04-12 10:24:59 +02:00
xmarre d533beb275 Merge pull request #3 from xmarre/codex/policy-override-loader
Add explicit residency policy loader input
2026-04-12 06:49:22 +02:00
xmarre 24ab29880a Scope loader policy override to load 2026-04-12 06:45:32 +02:00
xmarre 791a48afd3 Add explicit residency policy loader input 2026-04-12 06:35:30 +02:00
xmarre 1253b88361 Merge pull request #2 from xmarre/codex/fix-clip-metadata-cpu
[codex] Keep CLIP metadata tensors on CPU
2026-04-12 05:55:53 +02:00
xmarre ca8118de03 fix cpu fallback and torch remap handling 2026-04-12 05:54:14 +02:00
xmarre 752def584f load torch checkpoints on cpu before cuda remap 2026-04-12 05:42:29 +02:00
xmarre 871c144ac9 keep clip metadata tensors on cpu 2026-04-12 05:26:51 +02:00
xmarre ca0d55ce04 Merge pull request #1 from xmarre/codex/import-zip-implementation
[codex] Import GPU resident loader implementation
2026-04-12 04:52:39 +02:00
7 changed files with 4149 additions and 245 deletions
+484 -79
View File
@@ -1,32 +1,44 @@
# ComfyUI GPU Resident Loader
A ComfyUI custom-node pack that targets **time-to-VRAM** and **sticky GPU residency**, not just lower peak host RAM.
A ComfyUI custom-node pack for **faster time-to-VRAM**, **selective safetensors loading**, **sticky GPU residency control**, and **visibility into compatible external GPU model caches**.
It does two things:
This repo does three related jobs:
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.
1. **Installs startup-time monkey patches** before any workflow nodes run.
2. Ships **KJ-style resident loader nodes** for diffusion models and checkpoints.
3. Maintains a **live residency system** for native ComfyUI objects and compatible external caches, with preload / pin / evict / report controls for native tracked objects and automatic snapshot support plus provider-specific eviction for compatible external entries.
It is not just a “clean RAM” addon. The main target is the path from model file -> tensors -> live ComfyUI object -> VRAM retention, including GPU-resident caches that live outside ComfyUI’s normal loaded-model list.
## Why this exists
Stock ComfyUI makes separate decisions for:
ComfyUI’s default behavior mixes together two separate concerns:
- where a model **lives after load**, and
- where checkpoint tensors are **materialized first**.
- **ingest path** — where tensors are first materialized while a model is being read, and
- **residency policy** — where the finished model tends to live afterwards.
Those are not the same thing.
Those are not the same problem.
This repo targets the second problem directly for `.safetensors` by steering eligible loads toward direct GPU ingest, then targets the first problem by overriding offload policy and by teaching `free_memory()` to respect sticky entries until the VRAM budget is actually exceeded.
This repo focuses on both:
## What is included
- For **`.safetensors`**, it tries to keep eligible loads on the narrowest, most GPU-friendly path it can.
- For **resident diffusion and checkpoint-model loads**, it avoids broad checkpoint materialization by selecting only the detected UNet keys where possible.
- For **runtime VRAM pressure**, it adds a sticky-priority registry and teaches ComfyUI’s unload path to protect higher-value resident entries until enough VRAM must be reclaimed.
- For **compatible external GPU caches** that bypass `comfy.model_management.current_loaded_models`, it can discover supported providers at runtime and include their entries in snapshot output and provider-specific eviction decisions, with automatic trim kept opt-in.
- For **manual control**, it exposes nodes that let you preload, pin, evict, and inspect tracked native models, CLIPs, and VAEs.
### Startup patcher
## What changes at startup
Installed automatically from `__init__.py` when the custom node loads.
`__init__.py` calls `startup.install_patches()`, which applies the core monkey patches exactly once when the custom node is imported.
Those startup patches cover the built-in ComfyUI load and residency paths below. Compatible external-cache integrations are installed lazily later, on demand, when a supported module is actually present in the running process.
### Patched functions / methods
Current patch surface:
- `comfy.utils.load_torch_file`
- `comfy.clip_vision.load_torch_file` (redirected to the patched `comfy.utils.load_torch_file` when present)
- `comfy.model_management.free_memory`
- `comfy.model_management.load_models_gpu`
- `comfy.model_management.unet_offload_device`
@@ -35,70 +47,437 @@ Current patch surface:
- `comfy.model_management.text_encoder_device`
- `comfy.model_management.vae_device`
- `comfy.model_management.unet_inital_load_device`
- `comfy.model_management.LoadedModel.model_unload`
- `comfy.model_patcher.ModelPatcher.detach`
- `comfy.sd.load_checkpoint_guess_config`
- `comfy.sd.load_diffusion_model`
- `comfy.sd.load_clip`
- `comfy.sd.VAE.encode`
- `comfy.sd.VAE.decode`
- `comfy.clip_vision.load`
- `comfy.controlnet.load_controlnet`
- `comfy.diffusers_load.load_diffusers`
### Loader nodes
### Lazily installed external integrations
- **Diffusion Model Loader Resident**
- **Checkpoint Loader Resident**
- **Diffusion Model Selector Resident**
Current external integration surface:
`Diffusion Model Loader Resident` mirrors the relevant KJ diffusion-loader feature surface:
- compatible **SeedVR2** `src/core/model_cache.py` modules discovered at runtime
- weight dtype override
- compute dtype override
- cublas-ops toggle
- SageAttention override
- fp16 accumulation toggle
- optional extra-state-dict merge
When a compatible SeedVR2 cache module is present, the repo wraps:
### Residency nodes
- `GlobalModelCache.set_dit`
- `GlobalModelCache.set_vae`
- `GlobalModelCache.replace_dit` (when that method exists in the installed SeedVR2 build)
- `GlobalModelCache.replace_vae` (when that method exists in the installed SeedVR2 build)
- `GlobalModelCache.remove_dit`
- `GlobalModelCache.remove_vae`
- **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**
That lazy integration lets the loader:
- mirror SeedVR2-owned cached **DiT** and **VAE** objects into a separate external residency registry
- refresh byte / device / claimed-state metadata from the live cached object
- evict those entries through SeedVR2’s own removal path instead of assuming they live in `comfy.model_management.current_loaded_models`
## What those patches do
### 1) `load_torch_file` becomes residency-aware
The patched loader:
- detects the active load context (`model`, `clip`, `vae`, `checkpoint`, etc.)
- picks an explicit GPU target device when the active policy wants GPU ingest
- attempts **direct safetensors reads** on the requested device
- falls back to **CPU read + tensor-by-tensor copy** if direct GPU safetensors loading fails
- still uses **CPU-first `torch.load()`** for pickle formats (`.ckpt`, `.pt`, `.pth`, `.bin`)
- records the actual load method in the residency registry
For safetensors loads happening inside a `model` / `clip` / `vae` context, it can select only the detected component keys from the file header instead of pulling the full file into memory first.
### 2) Loader contexts are attached to stock ComfyUI load paths
These stock paths are wrapped with registry context and output binding:
- checkpoint loads
- diffusion-model loads
- CLIP loads
- CLIP Vision loads
- ControlNet loads
- diffusers loads
That means the native registry is not limited to the custom resident nodes. Stock ComfyUI loaders that pass through these paths are also tracked.
### 3) Compatible external caches can join the residency system lazily
When a compatible SeedVR2 cache module is present, the loader installs cache-level hooks that register SeedVR2-owned cached DiT / VAE objects into a separate external registry.
Those entries:
- are refreshed from the live cached object at runtime
- appear in **Registry Snapshot** output under `external_entries`
- are evicted through SeedVR2’s own cache-removal methods rather than the normal Comfy unload path
- remain outside automatic trim unless `COMFYUI_GPU_RESIDENT_TRIM_EXTERNAL=1` is set
### 4) Device/offload policy is overridden
Depending on the active policy, the patcher can steer:
- initial UNet load device
- CLIP/Text Encoder device
- VAE device
- offload devices for UNet / CLIP / VAE
This is how `prefer_gpu` and `sticky_gpu` keep more of the hot path on the GPU side than stock ComfyUI would.
### 5) `free_memory()` becomes sticky-aware
Under `sticky_gpu`, `comfy.model_management.free_memory()` is patched so that:
- sticky tracked wrappers are considered first
- higher-priority sticky entries are protected first
- lower-priority or older sticky entries yield first when VRAM must be reclaimed
- a transient protection floor is applied so ComfyUI does not immediately tear down high-value resident entries for small requests
### 6) Clone replacement is hardened
`load_models_gpu()` is patched to fully unload clone-conflict wrappers before replacement instead of relying on a shallow detach path that can leave base weights patched.
### 7) Unload / detach is redirected to CPU when needed
`LoadedModel.model_unload()` and `ModelPatcher.detach()` are patched so that unloads which would otherwise not reclaim VRAM are redirected through a CPU offload target first.
### 8) VAE encode/decode gets a sticky-safe path
Under `sticky_gpu`, patched `VAE.encode()` and `VAE.decode()`:
- cap the working batch count when necessary to preserve transient VRAM headroom
- retry with tiled VAE encode/decode on OOM
That behavior is not a general performance feature toggle. It exists to reduce avoidable VRAM spikes while sticky residency is active.
## Policies
The startup patcher exposes four policies:
The registry exposes four global 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.
### `legacy`
Default selection order:
Stay closest to stock ComfyUI behavior. Registry tracking still exists, but the patcher does not aggressively steer ingest/offload toward the GPU path.
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`.
### `balanced`
## Important scope limits
Keep registry tracking and diagnostics without aggressive GPU residency behavior.
### Best path: `.safetensors`
### `prefer_gpu`
This repo is optimized around `.safetensors`.
Prefer GPU ingest for tracked model-like loads and keep the faster side of the device/offload policy for:
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.
- diffusion models
- CLIP / text encoders
- ControlNets
### `.ckpt` / `.pt` remain CPU-first under PyTorch
This policy does **not** auto-pin tracked objects.
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.
### `sticky_gpu`
Use the included conversion helper to migrate hot models to `.safetensors`.
Builds on `prefer_gpu` and additionally:
- auto-pins newly bound **models** and **CLIPs**
- keeps **VAE offload** on the GPU side as well
- patches `free_memory()` to protect sticky tracked wrappers by priority
- uses the sticky-safe VAE encode/decode behavior
### Default policy selection
Selection order is:
1. `COMFYUI_GPU_RESIDENT_POLICY`, if set to a supported value
2. `sticky_gpu` when ComfyUI is started with `--gpu_only`
3. `sticky_gpu` when ComfyUI is started with `--highvram`
4. otherwise `prefer_gpu`
Supported values are:
- `legacy`
- `balanced`
- `prefer_gpu`
- `sticky_gpu`
## Included nodes
All nodes live under the `GPU Resident Loader` category.
### Loader nodes
#### Diffusion Model Selector Resident
Returns an absolute path string for a selected diffusion model.
Notes:
- resolves from `diffusion_models`
- also exposes `text_encoders` entries whose filename contains `connector`
#### Diffusion Model Loader Resident
KJ-style diffusion-model loader with these controls:
- `weight_dtype`
- `compute_dtype`
- `patch_cublaslinear`
- `sage_attention`
- `enable_fp16_accumulation`
- optional `extra_state_dict`
- optional `policy_override`
Behavior:
- for `.safetensors`, it loads only the detected UNet portion of the file
- if `extra_state_dict` is provided, only matching UNet keys are merged
- repeated loads reuse a live equivalent model when the source path and loader-relevant options still match
- before GPU-bound loads, it estimates the upcoming footprint and trims only enough lower-priority residency to cover the request plus adaptive headroom
#### Checkpoint Loader Resident
Full checkpoint loader that returns:
- `MODEL`
- `CLIP`
- `VAE`
Behavior:
- shares the same tuning knobs as the resident diffusion-model loader for the model component
- reuses already-live equivalent components when possible
- composes the final output from model / clip / vae component loaders instead of always rebuilding the whole checkpoint path from scratch
#### Checkpoint Model Loader Resident
Model-only checkpoint loader.
Behavior:
- takes the same selective safetensors UNet fast path as the diffusion-model loader
- reuses a live equivalent model when available
- uses the same dtype / attention / cublas / fp16-accumulation knobs as the full checkpoint loader
#### Checkpoint Clip Loader Resident
CLIP-only checkpoint loader.
Behavior:
- can reuse a live equivalent CLIP object
- avoids rebuilding the diffusion model and VAE outputs when only CLIP is needed
#### Checkpoint VAE Loader Resident
VAE-only checkpoint loader.
Behavior:
- can reuse a live equivalent VAE object
- avoids rebuilding the diffusion model and CLIP outputs when only VAE is needed
### Residency nodes
#### Set Global Residency Policy
Sets the active global policy and returns it as a `STRING`.
The loader nodes also expose an optional `policy_override` string input for one-off loads.
#### Registry Snapshot
Returns a composite formatted JSON snapshot with:
- `policy` for the active global policy
- `entries` for native Comfy-managed registry entries
- `external_entries` for compatible external cache entries discovered at runtime
#### Pin Model Residency / Pin CLIP Residency / Pin VAE Residency
Marks a tracked native object as sticky or non-sticky and optionally changes its priority.
#### Preload Model To GPU / Preload CLIP To GPU / Preload VAE To GPU
Calls `load_models_gpu(..., force_full_load=True)` for the selected native object, then updates sticky state / priority in the registry.
#### Evict Model From GPU / Evict CLIP From GPU / Evict VAE From GPU
Attempts to unload the selected native object from the current loaded-model set.
`unpatch_weights=True` performs a full unload path. When eviction succeeds, the node returns `evicted`; otherwise `not_loaded`.
#### Report Model Residency / Report CLIP Residency / Report VAE Residency
Returns a JSON report for a single tracked native object.
If the object is not currently bound in the registry, the node returns a JSON payload with `tracked: false`.
## Adaptive trimming before resident loads
The resident loaders now do load-scoped VRAM trimming themselves.
Before a GPU-bound resident load, the loader estimates required bytes from:
- the safetensors header when possible
- the detected checkpoint component subset when possible
- otherwise the source file size as a fallback
It then requests enough free VRAM for:
- the estimated load size
- adaptive headroom
Current adaptive headroom policy:
- ratio: `12.5%` of the estimated load
- floor: `256 MiB`
- ceiling: `1 GiB`
The trim path prefers to:
- unload non-sticky entries first
- then lower-priority sticky entries
- preserve explicitly kept models
- use partial unload where available
- include compatible external cache entries in the same candidate search only when `COMFYUI_GPU_RESIDENT_TRIM_EXTERNAL=1`, they are visible, on the same device, and not currently claimed/in use
This logic lives in the resident loader path. You do not need a separate “target free VRAM” node for it.
## Registry and observability
The snapshot now exposes two collections:
- `entries` for native Comfy-managed tracked objects
- `external_entries` for compatible external cache objects
### Native registry entries
The native registry tracks residency metadata for bound objects.
Typical per-entry fields include:
- `entry_id`
- `kind`
- `source_path`
- `basename`
- `sticky`
- `priority`
- `created_at`
- `last_touched`
- `loaded_bytes`
- `total_bytes`
- `load_device`
- `offload_device`
- `current_device`
- `last_method`
- `last_report`
- `loader_key`
- `notes`
- `alive`
The `last_method` / `last_report` fields let you see whether a load actually used:
- direct safetensors GPU ingest
- safetensors CPU -> CUDA fallback
- safetensors component-only load
- CPU-first `torch.load()` compatibility path
- a recorded load failure
### External registry entries
Typical external-entry fields include:
- `entry_id`
- `cache_key`
- `kind`
- `source_path`
- `basename`
- `sticky`
- `priority`
- `created_at`
- `last_touched`
- `loaded_bytes`
- `total_bytes`
- `load_device`
- `offload_device`
- `current_device`
- `claimed`
- `notes`
- `alive`
- `external`
For SeedVR2-backed entries, `claimed: true` means the cache object is currently marked in use and is skipped by the external trim candidate search when that opt-in path is enabled.
## What gets tracked
### Native tracked/bound paths
Native tracked paths include:
- resident node loads from this repo
- stock checkpoint loads
- stock diffusion-model loads
- stock CLIP loads
- stock CLIP Vision loads
- stock diffusers loads
ControlNet loads also participate in the patched load context and device-policy path, but this repo does not currently expose dedicated ControlNet residency nodes.
### External tracked paths
Current external integration coverage is:
- SeedVR2 global cached **DiT** entries
- SeedVR2 global cached **VAE** entries
Those entries are discovered lazily from compatible SeedVR2 cache modules at runtime. They are tracked separately from the native registry and participate in snapshot plus provider-specific eviction decisions, with load-scoped trim available only through the explicit external-trim opt-in.
## Important limits and non-goals
### Best path is still `.safetensors`
The narrow fast path is built around `.safetensors`.
That is where this repo can:
- inspect headers cheaply
- select only model / clip / vae subsets
- estimate component bytes more accurately
- attempt direct device-targeted reads
### `.ckpt` / `.pt` / pickle formats are still CPU-first
For pickle-based formats, PyTorch still goes through `torch.load()` on CPU first.
The repo can still:
- track those loads
- keep the resulting live objects resident
- reuse equivalent live objects later
It does **not** claim direct-to-GPU ingest for those formats.
### 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.
This repo does **not** keep VRAM allocations alive after ComfyUI, Python, or WSL exits.
CUDA memory lifetime is process/context scoped. True persistence across process shutdown would need a separate long-lived keeper process or service that owns the CUDA context.
### External integrations are compatibility-based, not universal
The external registry does **not** automatically manage every third-party cache.
At the moment, the documented external integration target is **SeedVR2**. Other custom nodes with private caches remain invisible until this repo grows a provider-specific integration for them.
### It does not automatically capture arbitrary custom loader implementations
The native registry only sees objects that pass through the patched ComfyUI load paths or through this repo’s resident nodes.
If another custom node loads models through its own private code path and bypasses those patched entry points, that object may never become a tracked native registry entry. In that case, the preload / pin / evict / report nodes from this repo cannot manage it until that external loader is integrated or patched.
Likewise, even for supported external providers such as SeedVR2, the current external integration is about **observation + provider-specific eviction**, with automatic trim kept opt-in. This repo does **not** yet expose dedicated external preload / pin / report / evict nodes for provider-owned cache entries.
## Installation
@@ -108,57 +487,78 @@ Clone into `custom_nodes`:
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:
Install dependencies into 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.
Requirements declared by the repo:
## Basic usage
- Python `>=3.10`
- `safetensors>=0.4.3`
### For direct diffusion-model loading
Optional SageAttention dependencies are **not** installed by default. Install those separately if you plan to use a SageAttention mode in the resident loaders.
Use **Diffusion Model Loader Resident**.
## Basic usage patterns
Recommended on a large VRAM machine:
### 1) Large-VRAM, mostly resident workflow
- policy: `sticky_gpu`
- model format: `.safetensors`
- preload with **Preload Model To GPU**
- inspect with **Report Model Residency** or **Registry Snapshot**
Recommended baseline:
### For full checkpoints
- start ComfyUI with `--highvram` or set policy manually to `sticky_gpu`
- prefer `.safetensors` for hot models
- load diffusion models through **Diffusion Model Loader Resident**
- use **Preload ... To GPU** for models you know you will reuse
- inspect with **Report ... Residency** or **Registry Snapshot**
Use **Checkpoint Loader Resident**.
### 2) Full checkpoint workflow
That tracks and binds the resulting diffusion model, CLIP, and VAE independently so they appear in the registry snapshot.
Use **Checkpoint Loader Resident** when you want `MODEL + CLIP + VAE` together.
### For manual residency control
That path can reuse already-live components instead of always rebuilding all three outputs.
- 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
### 3) Staged checkpoint workflow
## Observability
Use component loaders when the graph does not need the whole checkpoint at once:
Every tracked load stores:
- **Checkpoint Model Loader Resident** for diffusion model only
- **Checkpoint Clip Loader Resident** for CLIP only
- **Checkpoint VAE Loader Resident** for VAE only
- source path
- last load method
- requested device
- actual device
- sticky flag
- current loaded bytes
- total bytes
- current/offload/load device
### 4) Manual native residency control
That data is surfaced through the report nodes and the registry snapshot node.
Use:
## Conversion helper
- **Pin ... Residency** to mark a tracked native entry sticky / non-sticky
- **Preload ... To GPU** to force a full live native load now
- **Evict ... From GPU** to unload it from the current native loaded-model set
`scripts/convert_checkpoint_to_safetensors.py` is included for one-time conversion of hot `.ckpt` / `.pt` files into `.safetensors`.
### 5) Mixed workflows with SeedVR2 external caching
If SeedVR2 keeps DiT or VAE models in its own global cache, those objects can now show up in **Registry Snapshot** under `external_entries`.
That means:
- you can see that those bytes exist even though they are outside `current_loaded_models`
- resident loader trim only reclaims them when `COMFYUI_GPU_RESIDENT_TRIM_EXTERNAL=1`, they are visible, on the same device, and not currently claimed/in use
- eviction goes through SeedVR2’s own cache-removal path instead of a normal Comfy wrapper unload
## Notes on compatibility and migration
### Legacy wiring: `extra_state_dict` used as a policy string
The resident diffusion-model loader contains a compatibility shim for older graphs:
- if `extra_state_dict` receives one of the known policy names
- and that value is **not** an existing file path
- it is interpreted as `policy_override` instead
New graphs should connect policy strings to **`policy_override`**, not to `extra_state_dict`.
### Convert hot pickle checkpoints to safetensors
`scripts/convert_checkpoint_to_safetensors.py` is included for one-time conversion of hot `.ckpt` / `.pt` / `.pth` style checkpoints.
Example:
@@ -168,8 +568,13 @@ python ComfyUI/custom_nodes/ComfyUI-GPU-Resident-Loader/scripts/convert_checkpoi
--output /path/to/model.safetensors
```
Optional flags:
- `--state-dict-key <key>` to extract a different top-level dict key
- `--allow-non-tensor-values` to skip non-tensor entries instead of failing
## 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.
This repo stays GPL-compatible because it adapts behavior from GPL-licensed ComfyUI and mirrors relevant loader behavior from GPL-3.0-licensed KJNodes.
+404
View File
@@ -0,0 +1,404 @@
from __future__ import annotations
import os
from contextlib import contextmanager
from typing import Any
import comfy.model_management as model_management
import torch
from .external_residency import EXTERNAL_REGISTRY, ensure_external_integrations_installed, external_trim_enabled
from .residency import REGISTRY
_ADAPTIVE_HEADROOM_RATIO = 0.125
_ADAPTIVE_HEADROOM_FLOOR_BYTES = 256 * 1024 * 1024
_ADAPTIVE_HEADROOM_CEIL_BYTES = 1024 * 1024 * 1024
def _safe_free_memory(device) -> int:
return int(model_management.get_free_memory(device))
def _normalize_trim_device(device: str | torch.device | None):
if device is None or isinstance(device, torch.device):
return device
try:
return torch.device(device)
except Exception:
return device
def _safe_is_dead(loaded) -> bool:
try:
return loaded.is_dead()
except Exception:
return False
def _device_matches(device_a, device_b) -> bool:
normalized_a = _normalize_trim_device(device_a)
normalized_b = _normalize_trim_device(device_b)
if normalized_a is None or normalized_b is None:
return False
if isinstance(normalized_a, torch.device) and isinstance(normalized_b, torch.device):
if normalized_a.type != normalized_b.type:
return False
if normalized_a.type == "cuda":
index_a = 0 if normalized_a.index is None else normalized_a.index
index_b = 0 if normalized_b.index is None else normalized_b.index
return index_a == index_b
return True
return str(normalized_a) == str(normalized_b)
def _should_force_cpu_offload(model: Any, *, active_device=None, force: bool = False) -> bool:
if force:
return True
active_device = _normalize_trim_device(active_device)
if not isinstance(active_device, torch.device) or active_device.type != "cuda":
return False
return _device_matches(active_device, getattr(model, "offload_device", None))
@contextmanager
def _temporary_offload_device(model: Any, device):
if model is None or device is None or not hasattr(model, "offload_device"):
yield
return
original_device = model.offload_device
model.offload_device = device
try:
yield
finally:
model.offload_device = original_device
def unload_loaded_model(
loaded,
*,
active_device: str | torch.device | None = None,
force_offload_to_cpu: bool = False,
unpatch_weights: bool = True,
) -> bool:
# This helper is intentionally scoped to full unloads owned by this plugin.
model = getattr(loaded, "model", None)
if model is None:
return False
active_device = _normalize_trim_device(active_device if active_device is not None else getattr(loaded, "device", None))
force_cpu_offload = _should_force_cpu_offload(
model,
active_device=active_device,
force=force_offload_to_cpu,
)
unload_target = torch.device("cpu") if force_cpu_offload else None
with _temporary_offload_device(model, unload_target):
return loaded.model_unload(None, unpatch_weights=unpatch_weights)
def _sort_key_for_candidate(entry, *, sticky_respected: bool) -> tuple[int, int, float]:
if not sticky_respected:
return (0, 0, getattr(entry, "last_touched", 0.0) if entry is not None else 0.0)
priority = int(getattr(entry, "priority", 0)) if entry is not None else 0
last_touched = float(getattr(entry, "last_touched", 0.0)) if entry is not None else 0.0
return (1, priority, last_touched)
def _should_keep_loaded_model(model: Any, keep_models: tuple[Any, ...]) -> bool:
if not keep_models:
return False
for keep in keep_models:
if keep is None:
continue
if model is keep:
return True
is_clone = getattr(model, "is_clone", None)
if callable(is_clone):
try:
if is_clone(keep):
return True
except Exception:
pass
return False
def _trim_candidates(
*,
device,
respect_sticky: bool,
sticky_floor_priority: int,
keep_models: tuple[Any, ...],
include_external: bool | None = None,
) -> list[tuple[Any, Any, bool, bool]]:
"""
Collects and returns eviction/unload candidates from in-memory models and, optionally, external integrations.
Filters currently loaded models by the given device, excludes dead or missing models, and skips models listed in `keep_models`. If `respect_sticky` is true, entries whose registry metadata mark them as sticky and whose priority meets or exceeds `sticky_floor_priority` are flagged so they are treated as higher-priority to keep. When `include_external` is true (or when `include_external` is None and external trimming is enabled), candidates from the external registry are included.
Parameters:
device: Device filter for candidates; if not None only candidates matching this device are considered.
respect_sticky (bool): Whether to respect sticky registry entries when computing candidate priority.
sticky_floor_priority (int): Minimum priority value for a registry entry to be considered sticky.
keep_models (tuple[Any, ...]): Objects that must not be selected as candidates.
include_external (bool | None): If True include external-registry candidates; if False exclude them; if None defer to runtime external_trim_enabled().
Returns:
list[tuple[Any, Any, bool, bool]]: A sorted list of tuples (candidate_obj, registry_entry_or_None, sticky_respected, is_external_candidate).
- candidate_obj: The loaded model object (internal) or the external object.
- registry_entry_or_None: Registry metadata for the candidate, or None if unavailable.
- sticky_respected: `True` when the candidate is marked sticky and meets `sticky_floor_priority`.
- is_external_candidate: `True` for candidates originating from the external registry.
"""
candidates: list[tuple[Any, Any, bool, bool]] = []
ensure_external_integrations_installed()
for loaded in list(model_management.current_loaded_models):
if device is not None and loaded.device != device:
continue
if _safe_is_dead(loaded):
continue
model = getattr(loaded, "model", None)
if model is None:
continue
if _should_keep_loaded_model(model, keep_models):
continue
entry = REGISTRY.entry_for_object(model)
sticky_respected = (
respect_sticky
and entry is not None
and bool(getattr(entry, "sticky", False))
and int(getattr(entry, "priority", 0)) >= int(sticky_floor_priority)
)
candidates.append((loaded, entry, sticky_respected, False))
should_include_external = external_trim_enabled() if include_external is None else bool(include_external)
if should_include_external:
for external_obj, entry, sticky_respected in EXTERNAL_REGISTRY.candidates(
device=device,
respect_sticky=respect_sticky,
sticky_floor_priority=sticky_floor_priority,
keep_models=keep_models,
):
candidates.append((external_obj, entry, sticky_respected, True))
candidates.sort(key=lambda item: _sort_key_for_candidate(item[1], sticky_respected=item[2]))
return candidates
def trim_resident_vram(
*,
device: str | torch.device | None = None,
target_free_vram_bytes: int,
respect_sticky: bool,
sticky_floor_priority: int,
allow_partial_unload: bool,
keep_models: tuple[Any, ...] = (),
include_external: bool | None = None,
) -> dict[str, Any]:
"""
Trim resident models (and optional external integration objects) until the requested amount of free VRAM is available or a stopping condition occurs.
Tries to free GPU memory on `device` by evicting or unloading loaded models (and, optionally, external candidates). Respects sticky/priority hints, can perform partial unloads of pinned RAM when supported, and records each attempted action in the returned report.
Parameters:
device (str | torch.device | None): Target device to free (defaults to model_management.get_torch_device()).
target_free_vram_bytes (int): Desired amount of free VRAM, in bytes.
respect_sticky (bool): If true, prefer protecting entries marked as sticky with sufficient priority.
sticky_floor_priority (int): Minimum priority required for a sticky entry to be respected.
allow_partial_unload (bool): If true, allow partial unloads and attempts to free pinned host RAM before full eviction.
keep_models (tuple[Any, ...]): Sequence of model objects that must not be unloaded (exact matches or recognized clones).
include_external (bool | None): If None, use external_trim_enabled() at runtime; otherwise force inclusion/exclusion of external candidates.
Returns:
dict[str, Any]: Report of the trimming operation containing:
- status: "met_target", "partial", or "error".
- stopped_reason: reason the loop stopped (e.g., "target_met", "no_candidates", "no_progress", "error").
- target_met (bool): whether the target free VRAM was reached.
- device (str): string form of the device used.
- target_free_vram_bytes (int), free_before_bytes (int), free_after_bytes (int).
- freed_vram_bytes (int): total freed VRAM during this call.
- respect_sticky (bool), sticky_floor_priority (int), allow_partial_unload (bool).
- external_trim_enabled (bool): computed flag indicating whether external candidates were considered.
- actions (list[dict]): ordered per-candidate action records; each entry includes metadata such as
entry_id, basename, tracked, external_candidate, sticky_respected, priority,
need_before_bytes, loaded_before_bytes, freed_pinned_ram_bytes,
mode (e.g., "full_unload", "partial_unload", "external_evict", "error"),
loaded_after_bytes, freed_vram_bytes, free_after_bytes, and any warnings/errors.
"""
ensure_external_integrations_installed()
cleanup_models_gc = getattr(model_management, "cleanup_models_gc", None)
if callable(cleanup_models_gc):
cleanup_models_gc()
REGISTRY.refresh_runtime_state()
EXTERNAL_REGISTRY.refresh_runtime_state()
device = _normalize_trim_device(device)
if device is None:
device = model_management.get_torch_device()
free_before = _safe_free_memory(device)
actions: list[dict[str, Any]] = []
soft_empty_cache = getattr(model_management, "soft_empty_cache", None)
stopped_reason = "target_met"
while True:
free_now = _safe_free_memory(device)
need = int(target_free_vram_bytes) - free_now
if need <= 0:
stopped_reason = "target_met"
break
candidates = _trim_candidates(
device=device,
respect_sticky=respect_sticky,
sticky_floor_priority=sticky_floor_priority,
keep_models=keep_models,
include_external=include_external,
)
if not candidates:
stopped_reason = "no_candidates"
break
candidate, entry, sticky_respected, is_external_candidate = candidates[0]
model = candidate if is_external_candidate else candidate.model
loaded_before = int(getattr(entry, "loaded_bytes", 0)) if is_external_candidate else int(candidate.model_loaded_memory())
action = {
"entry_id": getattr(entry, "entry_id", None),
"basename": None if entry is None else os.path.basename(getattr(entry, "source_path", "") or ""),
"tracked": entry is not None,
"external_candidate": bool(is_external_candidate),
"sticky_respected": sticky_respected,
"priority": None if entry is None else int(getattr(entry, "priority", 0)),
"need_before_bytes": need,
"loaded_before_bytes": loaded_before,
"freed_pinned_ram_bytes": 0,
}
if (not is_external_candidate and allow_partial_unload and hasattr(model, "pinned_memory_size") and hasattr(model, "partially_unload_ram")):
try:
pinned_memory = int(model.pinned_memory_size())
if pinned_memory > 0:
pinned_budget = min(pinned_memory, max(need, 0))
model.partially_unload_ram(pinned_budget)
action["freed_pinned_ram_bytes"] = pinned_budget
except Exception as exc:
action["pinned_ram_warning"] = str(exc)
if is_external_candidate:
try:
fully_unloaded = EXTERNAL_REGISTRY.evict(entry)
action["mode"] = "external_evict"
except Exception as exc:
action["mode"] = "error"
action["error"] = str(exc)
actions.append(action)
stopped_reason = "error"
break
else:
try:
fully_unloaded = candidate.model_unload(need if allow_partial_unload else None)
action["mode"] = "full_unload" if fully_unloaded else "partial_unload"
except Exception as exc:
if allow_partial_unload:
try:
fully_unloaded = candidate.model_unload(None)
action["mode"] = "full_unload_fallback"
action["partial_unload_warning"] = str(exc)
except Exception as fallback_exc:
action["mode"] = "error"
action["error"] = str(fallback_exc)
actions.append(action)
stopped_reason = "error"
break
else:
action["mode"] = "error"
action["error"] = str(exc)
actions.append(action)
stopped_reason = "error"
break
if fully_unloaded and not is_external_candidate:
try:
model_management.current_loaded_models.remove(candidate)
except ValueError:
pass
if callable(soft_empty_cache):
soft_empty_cache()
REGISTRY.refresh_runtime_state()
EXTERNAL_REGISTRY.refresh_runtime_state()
free_after = _safe_free_memory(device)
if is_external_candidate:
action["loaded_after_bytes"] = 0 if fully_unloaded else int(getattr(entry, "loaded_bytes", 0))
else:
action["loaded_after_bytes"] = 0 if fully_unloaded else int(candidate.model_loaded_memory())
action["freed_vram_bytes"] = max(0, free_after - free_now)
action["free_after_bytes"] = free_after
actions.append(action)
if action["freed_vram_bytes"] <= 0 and not fully_unloaded:
stopped_reason = "no_progress"
break
free_after = _safe_free_memory(device)
target_met = free_after >= int(target_free_vram_bytes)
if target_met:
stopped_reason = "target_met"
return {
"status": "met_target" if target_met else ("error" if stopped_reason == "error" else "partial"),
"stopped_reason": stopped_reason,
"target_met": target_met,
"device": str(device),
"target_free_vram_bytes": int(target_free_vram_bytes),
"free_before_bytes": free_before,
"free_after_bytes": free_after,
"freed_vram_bytes": max(0, free_after - free_before),
"respect_sticky": bool(respect_sticky),
"sticky_floor_priority": int(sticky_floor_priority),
"allow_partial_unload": bool(allow_partial_unload),
"external_trim_enabled": bool(external_trim_enabled() if include_external is None else include_external),
"actions": actions,
}
def adaptive_headroom_bytes(required_bytes: int) -> int:
required = max(0, int(required_bytes))
if required == 0:
return 0
return min(
_ADAPTIVE_HEADROOM_CEIL_BYTES,
max(_ADAPTIVE_HEADROOM_FLOOR_BYTES, int(required * _ADAPTIVE_HEADROOM_RATIO)),
)
def trim_resident_vram_for_load(
*,
required_bytes: int,
reason: str,
device: str | torch.device | None = None,
respect_sticky: bool = True,
sticky_floor_priority: int = 0,
allow_partial_unload: bool = True,
keep_models: tuple[Any, ...] = (),
) -> dict[str, Any]:
estimated_load_bytes = max(0, int(required_bytes))
headroom_bytes = adaptive_headroom_bytes(estimated_load_bytes)
target_free_vram_bytes = estimated_load_bytes + headroom_bytes
report = trim_resident_vram(
device=device,
target_free_vram_bytes=target_free_vram_bytes,
respect_sticky=respect_sticky,
sticky_floor_priority=sticky_floor_priority,
allow_partial_unload=allow_partial_unload,
keep_models=keep_models,
)
report["trim_strategy"] = "adaptive_load_request"
report["trim_reason"] = str(reason)
report["estimated_load_bytes"] = estimated_load_bytes
report["adaptive_headroom_bytes"] = headroom_bytes
report["kept_loaded_models"] = len([model for model in keep_models if model is not None])
return report
+726
View File
@@ -0,0 +1,726 @@
from __future__ import annotations
from collections.abc import Mapping
import contextlib
import dataclasses
import inspect
import logging
import os
import sys
import threading
import time
import weakref
from typing import Any, Callable
import torch
from .residency import KIND_MODEL, KIND_VAE, REGISTRY
_LOG = logging.getLogger(__name__)
_SEEDVR2_PATCHED_CLASS_IDS: set[int] = set()
_SEEDVR2_PATCHING_CLASS_IDS: set[int] = set()
_SEEDVR2_PATCHING_THREAD_IDS: dict[int, int] = {}
_SEEDVR2_PATCH_LOCK = threading.Lock()
_SEEDVR2_PATCH_CONDITION = threading.Condition(_SEEDVR2_PATCH_LOCK)
def _now() -> float:
return time.time()
def _normalize_device(device: str | torch.device | None):
if device is None or isinstance(device, torch.device):
return device
try:
return torch.device(device)
except Exception:
return device
def _device_matches(device_a, device_b) -> bool:
normalized_a = _normalize_device(device_a)
normalized_b = _normalize_device(device_b)
if normalized_a is None or normalized_b is None:
return False
if isinstance(normalized_a, torch.device) and isinstance(normalized_b, torch.device):
if normalized_a.type != normalized_b.type:
return False
if normalized_a.type == "cuda":
index_a = 0 if normalized_a.index is None else normalized_a.index
index_b = 0 if normalized_b.index is None else normalized_b.index
return index_a == index_b
return True
return str(normalized_a) == str(normalized_b)
def _iter_seedvr2_wrapper_chain(model: Any):
if model is None:
return
stack = [model]
seen: set[int] = set()
while stack:
current = stack.pop()
if current is None:
continue
current_id = id(current)
if current_id in seen:
continue
seen.add(current_id)
yield current
for attr in ("_orig_mod", "dit_model"):
child = getattr(current, attr, None)
if child is not None:
stack.append(child)
def _seedvr2_is_claimed(model: Any) -> bool:
return any(bool(getattr(current, "_seedvr2_cache_claimed", False)) for current in _iter_seedvr2_wrapper_chain(model))
def _first_tensor_device(model: Any) -> str | None:
if model is None:
return None
try:
for param in model.parameters():
return str(param.device)
except Exception:
pass
try:
for buffer in model.buffers():
return str(buffer.device)
except Exception:
pass
return None
def _unique_tensor_nbytes(model: Any) -> int:
if model is None:
return 0
total = 0
seen_storages: set[tuple[int, int, int]] = set()
def visit_tensor(tensor: torch.Tensor) -> None:
nonlocal total
if tensor is None:
return
try:
storage = tensor.untyped_storage()
key = (storage.data_ptr(), storage.nbytes(), int(tensor.device.index or 0) if tensor.device.type == "cuda" else -1)
except Exception:
key = (id(tensor), tensor.numel() * tensor.element_size(), -2)
if key in seen_storages:
return
seen_storages.add(key)
total += int(tensor.numel()) * int(tensor.element_size())
try:
for tensor in model.parameters():
visit_tensor(tensor)
except Exception:
pass
try:
for tensor in model.buffers():
visit_tensor(tensor)
except Exception:
pass
return int(total)
def _seedvr2_entry_key(kind: str, node_id: Any) -> str:
return f"seedvr2:{kind}:{node_id}"
def _seedvr2_source_path(kind: str, node_id: Any, config: dict[str, Any], model: Any) -> str:
model_name = config.get("model") or getattr(model, "_model_name", None) or f"node_{node_id}"
return f"seedvr2/{kind}/{node_id}/{model_name}"
def _seedvr2_state_provider(
model_ref: Callable[[], Any | None],
config: dict[str, Any],
) -> Callable[[], dict[str, Any]]:
def provider() -> dict[str, Any]:
model = model_ref()
if model is None:
return {}
total_bytes = _unique_tensor_nbytes(model)
return {
"current_device": _first_tensor_device(model),
"load_device": config.get("device"),
"offload_device": config.get("offload_device"),
"claimed": _seedvr2_is_claimed(model),
"loaded_bytes": total_bytes,
"total_bytes": total_bytes,
}
return provider
def _coerce_external_bytes(value: Any, fallback: int, *, cache_key: str, field_name: str) -> int:
try:
return int(value)
except (TypeError, ValueError) as exc:
_LOG.debug(
"GPU Resident Loader: ignoring invalid external %s for %s: %r (%s)",
field_name,
cache_key,
value,
exc,
)
return int(fallback)
@dataclasses.dataclass(slots=True)
class ExternalResidencyEntry:
entry_id: str
cache_key: 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
claimed: bool = False
notes: list[str] = dataclasses.field(default_factory=list)
object_ref: weakref.ReferenceType[Any] | None = None
state_provider: Callable[[], dict[str, Any]] | None = None
evict_callback: Callable[[], bool] | None = None
def object(self) -> Any | None:
return None if self.object_ref is None else self.object_ref()
def is_alive(self) -> bool:
return self.object() is not None
def as_dict(self) -> dict[str, Any]:
return {
"entry_id": self.entry_id,
"cache_key": self.cache_key,
"kind": self.kind,
"source_path": self.source_path,
"basename": os.path.basename(self.source_path) if self.source_path else None,
"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,
"claimed": self.claimed,
"notes": list(self.notes),
"alive": self.is_alive(),
"external": True,
}
class ExternalResidencyRegistry:
def __init__(self) -> None:
self._lock = threading.RLock()
self._entries: dict[str, ExternalResidencyEntry] = {}
self._cache_key_to_entry: dict[str, str] = {}
self._next_entry_seq = 1
def _make_entry_id(self, kind: str, source_path: str) -> str:
basename = os.path.basename(source_path) or "anonymous"
entry_id = f"external:{kind}:{basename}:{self._next_entry_seq}"
self._next_entry_seq += 1
return entry_id
def bind(
self,
*,
cache_key: str,
obj: Any,
kind: str,
source_path: str,
state_provider: Callable[[], dict[str, Any]],
evict_callback: Callable[[], bool],
sticky: bool = False,
priority: int | None = None,
note: str | None = None,
) -> ExternalResidencyEntry:
if obj is None:
raise ValueError("Cannot bind None into external residency registry")
with self._lock:
entry_id = self._cache_key_to_entry.get(cache_key)
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 = ExternalResidencyEntry(
entry_id=entry_id,
cache_key=cache_key,
kind=kind,
source_path=source_path,
sticky=bool(sticky),
priority=REGISTRY.default_priority(kind) if priority is None else int(priority),
)
self._entries[entry_id] = entry
self._cache_key_to_entry[cache_key] = entry_id
try:
entry.object_ref = weakref.ref(obj)
except TypeError:
entry.object_ref = None
entry.cache_key = cache_key
entry.kind = kind
entry.source_path = source_path
entry.sticky = bool(sticky)
entry.priority = REGISTRY.default_priority(kind) if priority is None else int(priority)
entry.state_provider = state_provider
entry.evict_callback = evict_callback
entry.last_touched = _now()
if note and note not in entry.notes:
entry.notes.append(note)
self.refresh_runtime_state()
return entry
def remove(self, *, cache_key: str) -> bool:
with self._lock:
entry_id = self._cache_key_to_entry.pop(cache_key, None)
if entry_id is None:
return False
return self._entries.pop(entry_id, None) is not None
def refresh_runtime_state(self) -> None:
ensure_external_integrations_installed()
stale_keys: list[str] = []
with self._lock:
for cache_key, entry_id in list(self._cache_key_to_entry.items()):
entry = self._entries.get(entry_id)
if entry is None:
stale_keys.append(cache_key)
continue
obj = entry.object()
if obj is None:
stale_keys.append(cache_key)
continue
state_provider = entry.state_provider
if state_provider is None:
continue
try:
state = state_provider() or {}
except Exception as exc:
_LOG.debug("GPU Resident Loader: failed to refresh external cache state for %s: %s", cache_key, exc)
continue
entry.current_device = state.get("current_device")
entry.load_device = state.get("load_device")
entry.offload_device = state.get("offload_device")
entry.claimed = bool(state.get("claimed", False))
entry.loaded_bytes = _coerce_external_bytes(
state.get("loaded_bytes", entry.loaded_bytes or 0),
entry.loaded_bytes or 0,
cache_key=cache_key,
field_name="loaded_bytes",
)
entry.total_bytes = _coerce_external_bytes(
state.get("total_bytes", entry.total_bytes or entry.loaded_bytes or 0),
entry.total_bytes or entry.loaded_bytes or 0,
cache_key=cache_key,
field_name="total_bytes",
)
for cache_key in stale_keys:
entry_id = self._cache_key_to_entry.pop(cache_key, None)
if entry_id is not None:
self._entries.pop(entry_id, None)
def candidates(
self,
*,
device: str | torch.device | None,
respect_sticky: bool,
sticky_floor_priority: int,
keep_models: tuple[Any, ...],
) -> list[tuple[Any, ExternalResidencyEntry, bool]]:
self.refresh_runtime_state()
output: list[tuple[Any, ExternalResidencyEntry, bool]] = []
with self._lock:
for entry in self._entries.values():
obj = entry.object()
if obj is None:
continue
if entry.claimed:
continue
if device is not None and not _device_matches(device, entry.current_device):
continue
if any(obj is keep for keep in keep_models if keep is not None):
continue
sticky_respected = (
respect_sticky
and bool(entry.sticky)
and int(entry.priority) >= int(sticky_floor_priority)
)
output.append((obj, entry, sticky_respected))
return output
def evict(self, entry: ExternalResidencyEntry) -> bool:
callback = entry.evict_callback
if callback is None:
return False
try:
result = bool(callback())
except Exception as exc:
_LOG.warning(
"GPU Resident Loader: external eviction callback failed for %s: %s",
entry.cache_key,
exc,
)
result = False
finally:
self.refresh_runtime_state()
return result
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
EXTERNAL_REGISTRY = ExternalResidencyRegistry()
def external_trim_enabled() -> bool:
value = os.environ.get("COMFYUI_GPU_RESIDENT_TRIM_EXTERNAL", "").strip().lower()
return value in {"1", "true", "yes", "on"}
def external_objects_for_models(models: tuple[Any, ...] | list[Any]) -> tuple[Any, ...]:
"""
Returns registered external objects that are part of any supplied model wrapper chain.
This lets callers preserve external cache entries when they are associated with a kept
model through wrapper indirection rather than exact object identity.
"""
related_ids: set[int] = set()
for model in models:
for related in _iter_seedvr2_wrapper_chain(model):
related_ids.add(id(related))
if not related_ids:
return ()
EXTERNAL_REGISTRY.refresh_runtime_state()
matches: list[Any] = []
seen_ids: set[int] = set()
with EXTERNAL_REGISTRY._lock:
for entry in EXTERNAL_REGISTRY._entries.values():
obj = entry.object()
if obj is None:
continue
obj_id = id(obj)
if obj_id in related_ids and obj_id not in seen_ids:
matches.append(obj)
seen_ids.add(obj_id)
return tuple(matches)
def _call_seedvr2_method_with_optional_expected_model(
method: Callable[..., Any],
*args: Any,
debug: Any = None,
expected_model: Any = None,
) -> Any:
try:
parameters = inspect.signature(method).parameters
except (TypeError, ValueError):
parameters = {}
accepts_kwargs = any(param.kind == inspect.Parameter.VAR_KEYWORD for param in parameters.values())
kwargs: dict[str, Any] = {}
if accepts_kwargs or "debug" in parameters:
kwargs["debug"] = debug
if expected_model is not None and (accepts_kwargs or "expected_model" in parameters):
kwargs["expected_model"] = expected_model
return method(*args, **kwargs)
def _register_seedvr2_cached_model(global_cache: Any, *, kind: str, config: Any, model: Any) -> ExternalResidencyEntry | None:
if not isinstance(config, Mapping):
_LOG.debug(
"GPU Resident Loader: skipping SeedVR2 %s cache entry with unexpected config type: %s",
kind,
type(config).__name__,
)
return None
node_id = config.get("node_id")
if node_id is None or model is None:
return None
try:
model_ref = weakref.ref(model)
except TypeError:
def model_ref() -> Any | None:
return None
source_path = _seedvr2_source_path(kind, node_id, config, model)
cache_key = _seedvr2_entry_key(kind, node_id)
registry_kind = KIND_MODEL if kind == "dit" else KIND_VAE
if kind == "dit":
def evict_callback() -> bool:
return bool(
_call_seedvr2_method_with_optional_expected_model(
global_cache.remove_dit,
{"node_id": node_id},
debug=None,
expected_model=model_ref(),
)
)
note = f"SeedVR2 cached DiT node {node_id}"
else:
def evict_callback() -> bool:
return bool(
_call_seedvr2_method_with_optional_expected_model(
global_cache.remove_vae,
{"node_id": node_id},
debug=None,
expected_model=model_ref(),
)
)
note = f"SeedVR2 cached VAE node {node_id}"
return EXTERNAL_REGISTRY.bind(
cache_key=cache_key,
obj=model,
kind=registry_kind,
source_path=source_path,
state_provider=_seedvr2_state_provider(model_ref, config),
evict_callback=evict_callback,
sticky=False,
priority=REGISTRY.default_priority(registry_kind),
note=note,
)
def _install_seedvr2_integration_for_module(module: Any) -> bool:
model_cache_cls = getattr(module, "GlobalModelCache", None)
get_global_cache = getattr(module, "get_global_cache", None)
if model_cache_cls is None or not callable(get_global_cache):
return False
class_id = id(model_cache_cls)
thread_id = threading.get_ident()
original_set_dit = None
original_set_vae = None
original_replace_dit = None
original_replace_vae = None
original_remove_dit = None
original_remove_vae = None
provisional_cache_keys: set[str] = set()
provisional_cache_ownership: dict[str, str] = {}
try:
with _SEEDVR2_PATCH_CONDITION:
while (
class_id in _SEEDVR2_PATCHING_CLASS_IDS
and _SEEDVR2_PATCHING_THREAD_IDS.get(class_id) != thread_id
):
_SEEDVR2_PATCH_CONDITION.wait()
if class_id in _SEEDVR2_PATCHED_CLASS_IDS or _SEEDVR2_PATCHING_THREAD_IDS.get(class_id) == thread_id:
return True
_SEEDVR2_PATCHING_CLASS_IDS.add(class_id)
_SEEDVR2_PATCHING_THREAD_IDS[class_id] = thread_id
original_set_dit = model_cache_cls.set_dit
original_set_vae = model_cache_cls.set_vae
original_replace_dit = getattr(model_cache_cls, "replace_dit", None)
original_replace_vae = getattr(model_cache_cls, "replace_vae", None)
original_remove_dit = model_cache_cls.remove_dit
original_remove_vae = model_cache_cls.remove_vae
def set_dit_wrapper(self, dit_config, model, model_name, debug=None):
result = original_set_dit(self, dit_config, model, model_name, debug)
if result is not None:
_register_seedvr2_cached_model(self, kind="dit", config=dit_config, model=model)
return result
def set_vae_wrapper(self, vae_config, model, model_name, debug=None):
result = original_set_vae(self, vae_config, model, model_name, debug)
if result is not None:
_register_seedvr2_cached_model(self, kind="vae", config=vae_config, model=model)
return result
def replace_dit_wrapper(self, dit_config, model, debug=None, expected_model=None):
if original_replace_dit is None:
return False
result = _call_seedvr2_method_with_optional_expected_model(
original_replace_dit,
self,
dit_config,
model,
debug=debug,
expected_model=expected_model,
)
if result:
_register_seedvr2_cached_model(self, kind="dit", config=dit_config, model=model)
return result
def replace_vae_wrapper(self, vae_config, model, debug=None, expected_model=None):
if original_replace_vae is None:
return False
result = _call_seedvr2_method_with_optional_expected_model(
original_replace_vae,
self,
vae_config,
model,
debug=debug,
expected_model=expected_model,
)
if result:
_register_seedvr2_cached_model(self, kind="vae", config=vae_config, model=model)
return result
def remove_dit_wrapper(self, dit_config, debug=None, expected_model=None):
result = _call_seedvr2_method_with_optional_expected_model(
original_remove_dit,
self,
dit_config,
debug=debug,
expected_model=expected_model,
)
if result:
EXTERNAL_REGISTRY.remove(cache_key=_seedvr2_entry_key("dit", dit_config.get("node_id")))
return result
def remove_vae_wrapper(self, vae_config, debug=None, expected_model=None):
result = _call_seedvr2_method_with_optional_expected_model(
original_remove_vae,
self,
vae_config,
debug=debug,
expected_model=expected_model,
)
if result:
EXTERNAL_REGISTRY.remove(cache_key=_seedvr2_entry_key("vae", vae_config.get("node_id")))
return result
global_cache = get_global_cache()
model_cache_lock = getattr(global_cache, "_model_cache_lock", None)
lock_context = model_cache_lock if model_cache_lock is not None else contextlib.nullcontext()
with lock_context:
model_cache_cls.set_dit = set_dit_wrapper
model_cache_cls.set_vae = set_vae_wrapper
if original_replace_dit is not None:
model_cache_cls.replace_dit = replace_dit_wrapper
if original_replace_vae is not None:
model_cache_cls.replace_vae = replace_vae_wrapper
model_cache_cls.remove_dit = remove_dit_wrapper
model_cache_cls.remove_vae = remove_vae_wrapper
dit_items = list(getattr(global_cache, "_dit_models", {}).items())
vae_items = list(getattr(global_cache, "_vae_models", {}).items())
for _node_id, entry in dit_items:
if not isinstance(entry, tuple) or len(entry) != 2:
continue
model, config = entry
if model is not None:
if isinstance(config, Mapping) and config.get("node_id") is not None:
cache_key = _seedvr2_entry_key("dit", config.get("node_id"))
provisional_cache_keys.add(cache_key)
prior_entry_id = EXTERNAL_REGISTRY._cache_key_to_entry.get(cache_key)
else:
cache_key = None
prior_entry_id = None
registered_entry = _register_seedvr2_cached_model(global_cache, kind="dit", config=config, model=model)
if cache_key is not None and registered_entry is not None:
if prior_entry_id != registered_entry.entry_id:
provisional_cache_ownership[cache_key] = registered_entry.entry_id
for _node_id, entry in vae_items:
if not isinstance(entry, tuple) or len(entry) != 2:
continue
model, config = entry
if model is not None:
if isinstance(config, Mapping) and config.get("node_id") is not None:
cache_key = _seedvr2_entry_key("vae", config.get("node_id"))
provisional_cache_keys.add(cache_key)
prior_entry_id = EXTERNAL_REGISTRY._cache_key_to_entry.get(cache_key)
else:
cache_key = None
prior_entry_id = None
registered_entry = _register_seedvr2_cached_model(global_cache, kind="vae", config=config, model=model)
if cache_key is not None and registered_entry is not None:
if prior_entry_id != registered_entry.entry_id:
provisional_cache_ownership[cache_key] = registered_entry.entry_id
with _SEEDVR2_PATCH_CONDITION:
_SEEDVR2_PATCHING_CLASS_IDS.discard(class_id)
_SEEDVR2_PATCHED_CLASS_IDS.add(class_id)
_SEEDVR2_PATCHING_THREAD_IDS.pop(class_id, None)
_SEEDVR2_PATCH_CONDITION.notify_all()
except Exception:
_LOG.debug(
"GPU Resident Loader: rolling back SeedVR2 integration for class_id=%s provisional_cache_keys=%s",
class_id,
sorted(provisional_cache_keys),
exc_info=True,
)
for cache_key in provisional_cache_keys:
created_entry_id = provisional_cache_ownership.get(cache_key)
if created_entry_id is not None:
with EXTERNAL_REGISTRY._lock:
current_entry_id = EXTERNAL_REGISTRY._cache_key_to_entry.get(cache_key)
if current_entry_id == created_entry_id:
EXTERNAL_REGISTRY.remove(cache_key=cache_key)
with lock_context:
if original_set_dit is not None:
model_cache_cls.set_dit = original_set_dit
if original_set_vae is not None:
model_cache_cls.set_vae = original_set_vae
if original_replace_dit is not None:
model_cache_cls.replace_dit = original_replace_dit
if original_replace_vae is not None:
model_cache_cls.replace_vae = original_replace_vae
if original_remove_dit is not None:
model_cache_cls.remove_dit = original_remove_dit
if original_remove_vae is not None:
model_cache_cls.remove_vae = original_remove_vae
with _SEEDVR2_PATCH_CONDITION:
_SEEDVR2_PATCHING_CLASS_IDS.discard(class_id)
_SEEDVR2_PATCHED_CLASS_IDS.discard(class_id)
_SEEDVR2_PATCHING_THREAD_IDS.pop(class_id, None)
_SEEDVR2_PATCH_CONDITION.notify_all()
raise
_LOG.info("GPU Resident Loader: integrated external SeedVR2 cache visibility hooks")
return True
def ensure_external_integrations_installed() -> None:
for module in list(sys.modules.values()):
if module is None:
continue
module_file = getattr(module, "__file__", None)
if not module_file:
continue
normalized_file = os.path.abspath(module_file).replace("\\", "/")
if not normalized_file.endswith("/src/core/model_cache.py"):
continue
if "seedvr2" not in normalized_file.lower():
continue
if hasattr(module, "GlobalModelCache") and hasattr(module, "get_global_cache"):
try:
_install_seedvr2_integration_for_module(module)
except Exception as exc:
_LOG.warning(
"GPU Resident Loader: failed to install SeedVR2 external integration from %s: %s",
normalized_file,
exc,
)
+960 -47
View File
File diff suppressed because it is too large Load Diff
+48 -5
View File
@@ -1,13 +1,25 @@
from __future__ import annotations
import json
import logging
from typing import Any
import comfy.model_management as model_management
from .kj_loader import CheckpointLoaderResident, DiffusionModelLoaderResident, DiffusionModelSelectorResident
from .cleanup import unload_loaded_model
from .external_residency import EXTERNAL_REGISTRY, ensure_external_integrations_installed
from .kj_loader import (
CheckpointClipLoaderResident,
CheckpointLoaderResident,
CheckpointModelLoaderResident,
CheckpointVAELoaderResident,
DiffusionModelLoaderResident,
DiffusionModelSelectorResident,
)
from .residency import REGISTRY
_LOG = logging.getLogger(__name__)
def _entry_report_json(obj: Any) -> str:
entry = REGISTRY.entry_for_object(obj)
@@ -52,9 +64,28 @@ def _evict_patcher(patcher, *, unpatch_weights: bool) -> bool:
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
loaded_model = getattr(loaded, "model", None)
if loaded_model is None:
continue
safe_is_clone = False
if loaded_model is not patcher:
try:
safe_is_clone = loaded_model.is_clone(patcher)
except Exception as exc:
_LOG.warning("GPU Resident Loader: failed to evaluate clone state during eviction: %s", exc)
if loaded_model is patcher or safe_is_clone:
fully_unloaded = unload_loaded_model(
loaded,
active_device=getattr(loaded, "device", None),
force_offload_to_cpu=True,
unpatch_weights=unpatch_weights,
)
if fully_unloaded and unpatch_weights:
try:
model_management.current_loaded_models.remove(loaded)
except ValueError:
pass
unloaded = unloaded or fully_unloaded
if unloaded and hasattr(model_management, "soft_empty_cache"):
model_management.soft_empty_cache()
REGISTRY.refresh_runtime_state()
@@ -93,7 +124,13 @@ class RegistrySnapshot:
CATEGORY = "GPU Resident Loader/residency"
def snapshot(self):
return (REGISTRY.snapshot_json(),)
ensure_external_integrations_installed()
payload = {
"policy": REGISTRY.get_policy(),
"entries": REGISTRY.snapshot(),
"external_entries": EXTERNAL_REGISTRY.snapshot(),
}
return (json.dumps(payload, indent=2, sort_keys=True),)
class PinModelResidency:
@@ -328,6 +365,9 @@ NODE_CLASS_MAPPINGS = {
"DiffusionModelSelectorResident": DiffusionModelSelectorResident,
"DiffusionModelLoaderResident": DiffusionModelLoaderResident,
"CheckpointLoaderResident": CheckpointLoaderResident,
"CheckpointModelLoaderResident": CheckpointModelLoaderResident,
"CheckpointClipLoaderResident": CheckpointClipLoaderResident,
"CheckpointVAELoaderResident": CheckpointVAELoaderResident,
"SetGlobalResidencyPolicy": SetGlobalResidencyPolicy,
"RegistrySnapshot": RegistrySnapshot,
"PinModelResidency": PinModelResidency,
@@ -348,6 +388,9 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"DiffusionModelSelectorResident": "Diffusion Model Selector Resident",
"DiffusionModelLoaderResident": "Diffusion Model Loader Resident",
"CheckpointLoaderResident": "Checkpoint Loader Resident",
"CheckpointModelLoaderResident": "Checkpoint Model Loader Resident",
"CheckpointClipLoaderResident": "Checkpoint Clip Loader Resident",
"CheckpointVAELoaderResident": "Checkpoint VAE Loader Resident",
"SetGlobalResidencyPolicy": "Set Global Residency Policy",
"RegistrySnapshot": "Registry Snapshot",
"PinModelResidency": "Pin Model Residency",
+1367 -77
View File
File diff suppressed because it is too large Load Diff
+160 -37
View File
@@ -40,6 +40,7 @@ class LoadContext:
source_path: str | None = None
explicit_device: torch.device | None = None
note: str | None = None
cache_key: str | None = None
@dataclasses.dataclass(slots=True)
@@ -82,8 +83,10 @@ class ResidencyEntry:
current_device: str | None = None
last_method: str | None = None
last_report: dict[str, Any] | None = None
loader_key: str | None = None
notes: list[str] = dataclasses.field(default_factory=list)
object_ref: weakref.ReferenceType[Any] | None = None
cached_object_ref: weakref.ReferenceType[Any] | None = None
def is_alive(self) -> bool:
return self.object_ref is not None and self.object_ref() is not None
@@ -91,6 +94,9 @@ class ResidencyEntry:
def object(self) -> Any | None:
return None if self.object_ref is None else self.object_ref()
def cached_object(self) -> Any | None:
return None if self.cached_object_ref is None else self.cached_object_ref()
def as_dict(self) -> dict[str, Any]:
basename = os.path.basename(self.source_path) if self.source_path else None
return {
@@ -109,6 +115,7 @@ class ResidencyEntry:
"current_device": self.current_device,
"last_method": self.last_method,
"last_report": self.last_report,
"loader_key": self.loader_key,
"notes": list(self.notes),
"alive": self.is_alive(),
}
@@ -120,9 +127,37 @@ class ResidencyRegistry:
self._entries: dict[str, ResidencyEntry] = {}
self._reports_by_path: dict[str, LoadReport] = {}
self._path_to_entry: dict[tuple[str, str], str] = {}
self._loader_key_to_entry: dict[tuple[str, str, str], str] = {}
self._object_to_entry: weakref.WeakKeyDictionary[Any, str] = weakref.WeakKeyDictionary()
self._policy = self._default_policy()
def _gpu_ingest_kinds(self, policy: str) -> set[str]:
if policy in {"prefer_gpu", "sticky_gpu"}:
return {KIND_MODEL, KIND_CHECKPOINT, KIND_CLIP, KIND_VAE, KIND_CONTROLNET, KIND_CLIP_VISION}
return set()
def _gpu_offload_kinds(self, policy: str) -> set[str]:
if policy == "sticky_gpu":
return {KIND_MODEL, KIND_CLIP, KIND_VAE, KIND_CONTROLNET}
if policy == "prefer_gpu":
return {KIND_MODEL, KIND_CLIP, KIND_CONTROLNET}
return set()
def _autopin_kinds(self, policy: str) -> set[str]:
if policy == "sticky_gpu":
return {KIND_MODEL, KIND_CLIP}
return set()
def default_priority(self, kind: str) -> int:
return {
KIND_MODEL: 300,
KIND_CHECKPOINT: 300,
KIND_CONTROLNET: 200,
KIND_CLIP: 150,
KIND_VAE: 100,
KIND_CLIP_VISION: 50,
}.get(kind, 0)
def _default_policy(self) -> str:
env_value = os.environ.get("COMFYUI_GPU_RESIDENT_POLICY", "").strip().lower()
if env_value in POLICIES:
@@ -154,14 +189,21 @@ class ResidencyRegistry:
def wants_gpu_ingest(self, kind: str | None = None) -> bool:
policy = self.get_policy()
return policy in {"prefer_gpu", "sticky_gpu"}
if kind is None:
return bool(self._gpu_ingest_kinds(policy))
return kind in self._gpu_ingest_kinds(policy)
def wants_gpu_offload(self, kind: str | None = None) -> bool:
policy = self.get_policy()
return policy in {"prefer_gpu", "sticky_gpu"}
if kind is None:
return bool(self._gpu_offload_kinds(policy))
return kind in self._gpu_offload_kinds(policy)
def autopin_on_bind(self, kind: str | None = None) -> bool:
return self.get_policy() == "sticky_gpu"
policy = self.get_policy()
if kind is None:
return bool(self._autopin_kinds(policy))
return kind in self._autopin_kinds(policy)
def explicit_load_device(self, kind: str, source_path: str | None = None) -> torch.device | None:
if not self.wants_gpu_ingest(kind):
@@ -183,6 +225,7 @@ class ResidencyRegistry:
source_path: str | None = None,
explicit_device: torch.device | None = None,
note: str | None = None,
cache_key: str | None = None,
) -> Iterator[None]:
token = _LOAD_CONTEXT.set(
LoadContext(
@@ -190,6 +233,7 @@ class ResidencyRegistry:
source_path=source_path,
explicit_device=explicit_device,
note=note,
cache_key=cache_key,
)
)
try:
@@ -222,7 +266,12 @@ class ResidencyRegistry:
)
with self._lock:
self._reports_by_path[path] = report
entry_id = self._path_to_entry.get((kind, path))
entry_id = None
ctx = self.current_context()
if ctx is not None and ctx.cache_key is not None:
entry_id = self._loader_key_to_entry.get((kind, path, ctx.cache_key))
else:
entry_id = self._path_to_entry.get((kind, path))
if entry_id is not None:
entry = self._entries.get(entry_id)
if entry is not None:
@@ -242,6 +291,40 @@ class ResidencyRegistry:
basename = os.path.basename(source_path) or "anonymous"
return f"{kind}:{basename}:{len(self._entries) + 1}"
def _clear_object_binding(self, obj: Any | None, entry_id: str) -> None:
if obj is None:
return
try:
if self._object_to_entry.get(obj) == entry_id:
self._object_to_entry.pop(obj, None)
except TypeError:
pass
try:
if getattr(obj, "__gpu_resident_loader_entry_id__", None) == entry_id:
delattr(obj, "__gpu_resident_loader_entry_id__")
except (AttributeError, TypeError):
pass
def _tag_object_with_entry(self, obj: Any, entry_id: str) -> None:
try:
setattr(obj, "__gpu_resident_loader_entry_id__", entry_id)
try:
self._object_to_entry.pop(obj, None)
except TypeError:
pass
return
except (AttributeError, TypeError):
_LOG.debug(
"GPU Resident Loader: could not tag object %r with residency entry id %s",
type(obj),
entry_id,
)
try:
self._object_to_entry[obj] = entry_id
except TypeError:
pass
def bind_object(
self,
obj: Any,
@@ -249,15 +332,22 @@ class ResidencyRegistry:
source_path: str,
kind: str,
sticky: bool | None = None,
priority: int = 0,
priority: int | None = None,
note: str | None = None,
loader_key: str | None = None,
reusable_obj: Any | None = None,
) -> ResidencyEntry:
if obj is None:
raise ValueError("Cannot bind None into residency registry")
with self._lock:
old_key: tuple[str, str] | None = None
entry_id = getattr(obj, "__gpu_resident_loader_entry_id__", None)
old_loader_key: tuple[str, str, str] | None = None
previous_obj: Any | None = None
previous_cached_obj: Any | None = None
entry_id = self._loader_key_to_entry.get((kind, source_path, loader_key)) if loader_key is not None else None
if entry_id is None:
entry_id = getattr(obj, "__gpu_resident_loader_entry_id__", None)
if entry_id is None:
try:
entry_id = self._object_to_entry.get(obj)
@@ -266,6 +356,10 @@ class ResidencyRegistry:
if entry_id is not None and entry_id in self._entries:
entry = self._entries[entry_id]
old_key = (entry.kind, entry.source_path)
previous_obj = entry.object()
previous_cached_obj = entry.cached_object()
if entry.loader_key is not None:
old_loader_key = (entry.kind, entry.source_path, entry.loader_key)
else:
entry_id = self._make_entry_id(kind, source_path)
entry = ResidencyEntry(
@@ -273,45 +367,46 @@ class ResidencyRegistry:
kind=kind,
source_path=source_path,
sticky=self.autopin_on_bind(kind) if sticky is None else bool(sticky),
priority=int(priority),
priority=self.default_priority(kind) if priority is None else int(priority),
)
self._entries[entry_id] = entry
self._path_to_entry[(kind, source_path)] = entry_id
try:
setattr(obj, "__gpu_resident_loader_entry_id__", entry_id)
try:
self._object_to_entry.pop(obj, None)
except TypeError:
pass
except (AttributeError, TypeError):
_LOG.debug(
"GPU Resident Loader: could not tag object %r with residency entry id %s",
type(obj),
entry_id,
)
cache_obj = obj if reusable_obj is None else reusable_obj
for stale_obj in (previous_obj, previous_cached_obj):
if stale_obj is None or stale_obj is obj or stale_obj is cache_obj:
continue
self._clear_object_binding(stale_obj, entry_id)
try:
entry.object_ref = weakref.ref(obj)
if getattr(obj, "__gpu_resident_loader_entry_id__", None) is None:
try:
self._object_to_entry[obj] = entry_id
except TypeError:
pass
except TypeError:
entry.object_ref = None
_LOG.debug(
"GPU Resident Loader: object %r is not weak-referenceable; tracking metadata only",
type(obj),
)
self._tag_object_with_entry(obj, entry_id)
try:
entry.cached_object_ref = weakref.ref(cache_obj)
except TypeError:
entry.cached_object_ref = entry.object_ref
entry.sticky = entry.sticky if sticky is None else bool(sticky)
entry.priority = int(priority)
entry.priority = entry.priority if priority is None else int(priority)
entry.source_path = source_path
entry.kind = kind
entry.loader_key = loader_key
new_key = (entry.kind, entry.source_path)
new_loader_key = (entry.kind, entry.source_path, entry.loader_key) if entry.loader_key is not None else None
if old_key is not None and old_key != new_key:
if self._path_to_entry.get(old_key) == entry_id:
self._path_to_entry.pop(old_key, None)
self._path_to_entry[new_key] = entry_id
if old_loader_key is not None and old_loader_key != new_loader_key:
if self._loader_key_to_entry.get(old_loader_key) == entry_id:
self._loader_key_to_entry.pop(old_loader_key, None)
if new_loader_key is not None:
self._loader_key_to_entry[new_loader_key] = entry_id
entry.last_touched = _now()
if note:
entry.notes.append(note)
@@ -321,21 +416,49 @@ class ResidencyRegistry:
entry.last_report = report.as_dict()
entry.current_device = report.actual_device
self.refresh_runtime_state()
return entry
def lookup_live_object(self, *, kind: str, source_path: str, loader_key: str) -> Any | None:
with self._lock:
entry_id = self._loader_key_to_entry.get((kind, source_path, loader_key))
if entry_id is None:
return None
entry = self._entries.get(entry_id)
if entry is None:
return None
obj = entry.cached_object()
if obj is None:
if entry.cached_object_ref is not None:
return None
obj = entry.object()
if obj is None:
return None
entry.last_touched = _now()
return obj
def _entry_id_for_object(self, obj: Any) -> str | None:
current = obj
seen: set[int] = set()
while current is not None and id(current) not in seen:
seen.add(id(current))
entry_id = getattr(current, "__gpu_resident_loader_entry_id__", None)
if entry_id is None:
try:
entry_id = self._object_to_entry.get(current)
except TypeError:
entry_id = None
if entry_id is not None:
return entry_id
current = getattr(current, "parent", None)
return None
def entry_for_object(self, obj: Any) -> ResidencyEntry | None:
if obj is None:
return None
entry_id = getattr(obj, "__gpu_resident_loader_entry_id__", None)
if entry_id is None:
try:
entry_id = self._object_to_entry.get(obj)
except TypeError:
entry_id = None
if entry_id is None:
return None
with self._lock:
entry_id = self._entry_id_for_object(obj)
if entry_id is None:
return None
return self._entries.get(entry_id)
def set_sticky(self, obj: Any, sticky: bool, priority: int | None = None) -> ResidencyEntry | None:
@@ -369,8 +492,9 @@ class ResidencyRegistry:
continue
entry = self.entry_for_object(loaded.model)
if entry is not None and entry.sticky:
output.append(loaded)
return output
output.append((entry.priority, entry.last_touched, loaded))
output.sort(key=lambda item: (item[0], item[1]), reverse=True)
return [item[2] for item in output]
def refresh_runtime_state(self) -> None:
try:
@@ -401,7 +525,6 @@ class ResidencyRegistry:
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()