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
5 changed files with 1217 additions and 226 deletions
+218 -117
View File
@@ -1,14 +1,14 @@
# ComfyUI GPU Resident Loader
A ComfyUI custom-node pack for **faster time-to-VRAM**, **selective safetensors loading**, and **sticky GPU residency control**.
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**.
This repo does three related jobs:
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 registry** with preload / pin / evict / report controls for tracked objects.
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.
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
@@ -24,11 +24,14 @@ This repo focuses on both:
- 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 **manual control**, it exposes nodes that let you preload, pin, evict, and inspect tracked models, CLIPs, and VAEs.
- 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.
## What changes at startup
`__init__.py` calls `startup.install_patches()`, which applies the monkey patches exactly once when the custom node is imported.
`__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
@@ -55,68 +58,100 @@ Current patch surface:
- `comfy.controlnet.load_controlnet`
- `comfy.diffusers_load.load_diffusers`
### What those patches do
### Lazily installed external integrations
#### 1) `load_torch_file` becomes residency-aware
Current external integration surface:
- compatible **SeedVR2** `src/core/model_cache.py` modules discovered at runtime
When a compatible SeedVR2 cache module is present, the repo wraps:
- `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`
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.
- 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
### 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.
- checkpoint loads
- diffusion-model loads
- CLIP loads
- CLIP Vision loads
- ControlNet loads
- diffusers loads
That means the registry is not limited to the custom resident nodes. Stock ComfyUI loaders that pass through these paths are also tracked.
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) Device/offload policy is overridden
### 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.
- 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.
#### 4) `free_memory()` becomes sticky-aware
### 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.
- 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
#### 5) Clone replacement is hardened
### 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.
#### 6) Unload / detach is redirected to CPU when needed
### 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.
#### 7) VAE encode/decode gets a sticky-safe path
### 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, and
- retry with tiled VAE encode/decode on OOM.
- 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.
@@ -136,9 +171,9 @@ Keep registry tracking and diagnostics without aggressive GPU residency behavior
Prefer GPU ingest for tracked model-like loads and keep the faster side of the device/offload policy for:
- diffusion models,
- CLIP / text encoders,
- ControlNets.
- diffusion models
- CLIP / text encoders
- ControlNets
This policy does **not** auto-pin tracked objects.
@@ -146,19 +181,19 @@ This policy does **not** auto-pin tracked objects.
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.
- 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`.
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:
@@ -179,8 +214,8 @@ Returns an absolute path string for a selected diffusion model.
Notes:
- resolves from `diffusion_models`, and
- also exposes `text_encoders` entries whose filename contains `connector`.
- resolves from `diffusion_models`
- also exposes `text_encoders` entries whose filename contains `connector`
#### Diffusion Model Loader Resident
@@ -196,10 +231,10 @@ KJ-style diffusion-model loader with these controls:
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.
- 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
@@ -211,9 +246,9 @@ Full checkpoint loader that returns:
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.
- 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
@@ -221,9 +256,9 @@ 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.
- 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
@@ -231,8 +266,8 @@ 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.
- can reuse a live equivalent CLIP object
- avoids rebuilding the diffusion model and VAE outputs when only CLIP is needed
#### Checkpoint VAE Loader Resident
@@ -240,8 +275,8 @@ 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.
- can reuse a live equivalent VAE object
- avoids rebuilding the diffusion model and CLIP outputs when only VAE is needed
### Residency nodes
@@ -253,25 +288,29 @@ The loader nodes also expose an optional `policy_override` string input for one-
#### Registry Snapshot
Returns the whole registry as formatted JSON.
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 object as sticky or non-sticky and optionally changes its priority.
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 object, then updates sticky state / priority in the registry.
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 object from the current loaded-model set.
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 object.
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`.
@@ -281,33 +320,41 @@ 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.
- 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, plus
- adaptive headroom.
- the estimated load size
- adaptive headroom
Current adaptive headroom policy:
- ratio: `12.5%` of the estimated load,
- floor: `256 MiB`,
- ceiling: `1 GiB`.
- 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.
- 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 registry tracks residency metadata for bound objects.
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:
@@ -332,25 +379,61 @@ Typical per-entry fields include:
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,
- or a recorded load failure.
- 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
Tracked/bound paths include:
### Native tracked/bound paths
- resident node loads from this repo,
- stock checkpoint loads,
- stock diffusion-model loads,
- stock CLIP loads,
- stock CLIP Vision loads,
- stock diffusers loads.
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`
@@ -359,10 +442,10 @@ 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.
- 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
@@ -370,9 +453,9 @@ 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.
- 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.
@@ -382,11 +465,19 @@ This repo does **not** keep VRAM allocations alive after ComfyUI, Python, or WSL
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 registry only sees objects that pass through the patched ComfyUI load paths or through this repo’s resident nodes.
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 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.
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
@@ -415,11 +506,11 @@ Optional SageAttention dependencies are **not** installed by default. Install th
Recommended baseline:
- 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**.
- 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**
### 2) Full checkpoint workflow
@@ -431,17 +522,27 @@ That path can reuse already-live components instead of always rebuilding all thr
Use component loaders when the graph does not need the whole checkpoint at once:
- **Checkpoint Model Loader Resident** for diffusion model only,
- **Checkpoint Clip Loader Resident** for CLIP only,
- **Checkpoint VAE Loader Resident** for VAE only.
- **Checkpoint Model Loader Resident** for diffusion model only
- **Checkpoint Clip Loader Resident** for CLIP only
- **Checkpoint VAE Loader Resident** for VAE only
### 4) Manual residency control
### 4) Manual native residency control
Use:
- **Pin ... Residency** to mark a tracked entry sticky / non-sticky,
- **Preload ... To GPU** to force a full live load now,
- **Evict ... From GPU** to unload it from the current loaded-model set.
- **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
### 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
@@ -449,9 +550,9 @@ Use:
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.
- 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`.
+63 -8
View File
@@ -7,7 +7,7 @@ from typing import Any
import comfy.model_management as model_management
import torch
from .external_residency import EXTERNAL_REGISTRY, ensure_external_integrations_installed
from .external_residency import EXTERNAL_REGISTRY, ensure_external_integrations_installed, external_trim_enabled
from .residency import REGISTRY
_ADAPTIVE_HEADROOM_RATIO = 0.125
@@ -130,7 +130,27 @@ def _trim_candidates(
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):
@@ -154,13 +174,15 @@ def _trim_candidates(
)
candidates.append((loaded, entry, sticky_respected, False))
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))
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
@@ -174,7 +196,38 @@ def trim_resident_vram(
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):
@@ -202,6 +255,7 @@ def trim_resident_vram(
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"
@@ -305,6 +359,7 @@ def trim_resident_vram(
"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,
}
+135 -53
View File
@@ -280,7 +280,7 @@ class ExternalResidencyRegistry:
entry.state_provider = state_provider
entry.evict_callback = evict_callback
entry.last_touched = _now()
if note:
if note and note not in entry.notes:
entry.notes.append(note)
self.refresh_runtime_state()
@@ -398,6 +398,41 @@ class ExternalResidencyRegistry:
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,
@@ -418,18 +453,18 @@ def _call_seedvr2_method_with_optional_expected_model(
return method(*args, **kwargs)
def _register_seedvr2_cached_model(global_cache: Any, *, kind: str, config: Any, model: Any) -> None:
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
return None
node_id = config.get("node_id")
if node_id is None or model is None:
return
return None
try:
model_ref = weakref.ref(model)
@@ -464,7 +499,7 @@ def _register_seedvr2_cached_model(global_cache: Any, *, kind: str, config: Any,
)
note = f"SeedVR2 cached VAE node {node_id}"
EXTERNAL_REGISTRY.bind(
return EXTERNAL_REGISTRY.bind(
cache_key=cache_key,
obj=model,
kind=registry_kind,
@@ -485,24 +520,33 @@ def _install_seedvr2_integration_for_module(module: Any) -> bool:
class_id = id(model_cache_cls)
thread_id = threading.get_ident()
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 = 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:
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
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)
@@ -570,54 +614,92 @@ def _install_seedvr2_integration_for_module(module: Any) -> bool:
EXTERNAL_REGISTRY.remove(cache_key=_seedvr2_entry_key("vae", vae_config.get("node_id")))
return result
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
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:
_register_seedvr2_cached_model(global_cache, kind="dit", config=config, model=model)
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:
_register_seedvr2_cached_model(global_cache, kind="vae", config=config, model=model)
except Exception:
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:
model_cache_cls.set_dit = original_set_dit
model_cache_cls.set_vae = original_set_vae
_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
model_cache_cls.remove_dit = original_remove_dit
model_cache_cls.remove_vae = original_remove_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
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()
_LOG.info("GPU Resident Loader: integrated external SeedVR2 cache eviction hooks")
_LOG.info("GPU Resident Loader: integrated external SeedVR2 cache visibility hooks")
return True
@@ -641,4 +723,4 @@ def ensure_external_integrations_installed() -> None:
"GPU Resident Loader: failed to install SeedVR2 external integration from %s: %s",
normalized_file,
exc,
)
)
+129
View File
@@ -395,6 +395,21 @@ def _estimate_model_load_bytes(
def _estimate_checkpoint_aux_component_bytes(ckpt_path: str, *, kind: str) -> int:
"""
Estimate the byte size of a specific checkpoint auxiliary component (e.g., CLIP or VAE).
If `ckpt_path` is a safetensors file this attempts a component-specific estimate via
`estimate_checkpoint_component_bytes`. If that returns `None` or the file is not
safetensors, the function falls back to the file size on disk.
Parameters:
ckpt_path (str): Path to the checkpoint file.
kind (str): Component kind to estimate (for example `"clip"`, `"vae"`, or other
checkpoint component identifiers accepted by `estimate_checkpoint_component_bytes`).
Returns:
int: Estimated number of bytes required by the requested component.
"""
estimated = None
if _is_safetensors_path(ckpt_path):
estimated = estimate_checkpoint_component_bytes(ckpt_path, kind)
@@ -403,6 +418,46 @@ def _estimate_checkpoint_aux_component_bytes(ckpt_path: str, *, kind: str) -> in
return int(estimated)
def _env_bool(name: str) -> bool | None:
"""
Parse a boolean-like environment variable value.
Parameters:
name (str): Environment variable name to read.
Returns:
bool | None: `True` if the variable value is one of "1", "true", "yes", or "on";
`False` if it is one of "0", "false", "no", or "off"; `None` if the variable is unset
or contains an unrecognized value.
"""
value = os.environ.get(name, "").strip().lower()
if value in {"1", "true", "yes", "on"}:
return True
if value in {"0", "false", "no", "off"}:
return False
return None
def _should_trim_before_load(*, effective_policy: str) -> bool:
"""
Decides whether to perform adaptive VRAM trimming before loading based on environment overrides, the effective residency policy, and runtime flags.
Parameters:
effective_policy (str): The resolved residency policy name to evaluate (e.g., "sticky_gpu").
Returns:
True if trimming should run before load, False otherwise.
"""
env_override = _env_bool("COMFYUI_GPU_RESIDENT_LOAD_TRIM")
if env_override is not None:
return env_override
if effective_policy == "sticky_gpu":
return False
if getattr(args, "highvram", False) or getattr(args, "gpu_only", False):
return False
return True
def _maybe_trim_before_load(
*,
loader_name: str,
@@ -410,7 +465,26 @@ def _maybe_trim_before_load(
explicit_device: torch.device | None,
required_bytes: int,
keep_models: tuple[Any, ...] = (),
enabled: bool = True,
) -> None:
"""
Attempt to free GPU VRAM proactively to create headroom for a forthcoming load.
If trimming is enabled and an explicit CUDA device is provided and `required_bytes` > 0,
calls the adaptive VRAM trimmer to free memory while preserving any `keep_models`.
Logs an informational message when memory was freed and a warning if the trimmer
could not reach the target headroom.
Parameters:
loader_name (str): Short name used in log messages for the entity requesting the trim.
reason (str): Human-readable reason for the trim (included in logs).
explicit_device (torch.device | None): The explicit device to trim on; trimming is skipped if `None` or not CUDA.
required_bytes (int): Estimated number of bytes needed for the upcoming load; trimming is skipped if <= 0.
keep_models (tuple[Any, ...], optional): Objects to preserve from eviction while trimming (defaults to ()).
enabled (bool, optional): If `False`, the function is a no-op (defaults to True).
"""
if not enabled:
return
if explicit_device is None or explicit_device.type != "cuda":
return
if required_bytes <= 0:
@@ -601,6 +675,26 @@ def _load_resident_diffusion_model(
policy_override: str | None = None,
keep_models: tuple[Any, ...] = (),
) -> Any:
"""
Load a diffusion model state dict from `source_path`, applying residency-aware VRAM trimming, optional extra UNet state merging, backend flags, and post-load model patches, and bind or reuse the resulting model for future loads.
Parameters:
loader_name (str): Human-readable loader identifier used for logging.
cache_scope (str): Logical cache scope used when constructing the loader key.
source_path (str): Filesystem path to the diffusion model checkpoint.
note (str): Short note describing the load context (used when binding).
weight_dtype (str): Weight dtype selection used to build model options.
compute_dtype (str): Compute dtype selection applied to the loaded model.
patch_cublaslinear (bool): Whether to enable the cublas linear optimization during load.
sage_attention (str): SageAttention mode to apply to the model after loading.
enable_fp16_accumulation (bool): If true, enable fp16 accumulation backend flag during load.
extra_state_dict (str | None): Optional path to an additional state dict whose matching UNet keys will be merged into the main state dict.
policy_override (str | None): Optional residency policy override name (affects trimming decision).
keep_models (tuple[Any, ...]): Iterable of already-loaded objects whose residency should be preserved during trimming.
Returns:
The loaded diffusion model object.
"""
model_options = _build_model_options(weight_dtype)
effective_policy = _effective_policy_name(policy_override)
loader_key = _make_loader_key(
@@ -632,6 +726,7 @@ def _load_resident_diffusion_model(
extra_state_dict=extra_state_dict,
),
keep_models=keep_models,
enabled=_should_trim_before_load(effective_policy=effective_policy),
)
with _temporary_backend_flags(
@@ -697,6 +792,20 @@ def _load_checkpoint_clip_only(
loader_name: str,
keep_models: tuple[Any, ...] = (),
):
"""
Load or reuse the CLIP (text encoder) component from a checkpoint file.
This resolves a loader cache key, attempts to reuse a previously loaded CLIP, and if absent loads the checkpoint (safetensors or torch format), optionally applies model-config-based extraction or older-quant conversion, constructs a CLIP object when weights are present, and binds the result for reuse.
Parameters:
ckpt_path (str): Filesystem path to the checkpoint file.
policy_override (str | None): Optional residency policy override used to form the loader key.
loader_name (str): Human-readable name used in log messages and trimming decisions.
keep_models (tuple[Any, ...]): Sequence of model-like objects to preserve during any pre-load VRAM trimming.
Returns:
clip (comfy.sd.CLIP | None): A constructed CLIP/text-encoder instance when weights are available, or `None` if no CLIP weights were found.
"""
loader_key = _checkpoint_component_loader_key("clip", policy_override)
reused_clip = REGISTRY.lookup_live_object(kind=KIND_CLIP, source_path=ckpt_path, loader_key=loader_key)
if reused_clip is not None:
@@ -707,6 +816,7 @@ def _load_checkpoint_clip_only(
is_safetensors = _is_safetensors_path(ckpt_path)
header_info = checkpoint_component_info_from_header(ckpt_path) if is_safetensors else None
model_config = None if header_info is None else header_info.get("model_config")
effective_policy = _effective_policy_name(policy_override)
explicit_device = REGISTRY.explicit_load_device(kind=KIND_CLIP, source_path=ckpt_path)
_maybe_trim_before_load(
loader_name=loader_name,
@@ -714,6 +824,7 @@ def _load_checkpoint_clip_only(
explicit_device=explicit_device,
required_bytes=_estimate_checkpoint_aux_component_bytes(ckpt_path, kind=KIND_CLIP),
keep_models=keep_models,
enabled=_should_trim_before_load(effective_policy=effective_policy),
)
with REGISTRY.load_context(kind=KIND_CLIP, source_path=ckpt_path, explicit_device=explicit_device, cache_key=loader_key):
if is_safetensors and model_config is None:
@@ -783,6 +894,22 @@ def _load_checkpoint_vae_only(
loader_name: str,
keep_models: tuple[Any, ...] = (),
):
"""
Load only the VAE component from a checkpoint, reusing a live VAE when available and binding the loaded VAE for reuse.
Parameters:
ckpt_path (str): Filesystem path to the checkpoint (safetensors or torch file).
policy_override (str | None): Optional residency policy override used to resolve trimming/loading behavior.
loader_name (str): Human-readable name used in logging messages.
keep_models (tuple[Any, ...]): Objects to preserve from eviction when freeing VRAM before loading.
Returns:
vae (comfy.sd.VAE | None): The loaded VAE object, or `None` if no VAE weights were found in the checkpoint.
Side effects:
- May trigger adaptive VRAM trimming before loading depending on the effective policy and runtime flags.
- Binds the resulting VAE (if any) into the live-object registry for reuse.
"""
loader_key = _checkpoint_component_loader_key("vae", policy_override)
reused_vae = REGISTRY.lookup_live_object(kind=KIND_VAE, source_path=ckpt_path, loader_key=loader_key)
if reused_vae is not None:
@@ -793,6 +920,7 @@ def _load_checkpoint_vae_only(
is_safetensors = _is_safetensors_path(ckpt_path)
header_info = checkpoint_component_info_from_header(ckpt_path) if is_safetensors else None
model_config = None if header_info is None else header_info.get("model_config")
effective_policy = _effective_policy_name(policy_override)
explicit_device = REGISTRY.explicit_load_device(kind=KIND_VAE, source_path=ckpt_path)
_maybe_trim_before_load(
loader_name=loader_name,
@@ -800,6 +928,7 @@ def _load_checkpoint_vae_only(
explicit_device=explicit_device,
required_bytes=_estimate_checkpoint_aux_component_bytes(ckpt_path, kind=KIND_VAE),
keep_models=keep_models,
enabled=_should_trim_before_load(effective_policy=effective_policy),
)
with REGISTRY.load_context(kind=KIND_VAE, source_path=ckpt_path, explicit_device=explicit_device, cache_key=loader_key):
if is_safetensors and model_config is None:
+672 -48
View File
@@ -2,10 +2,12 @@ from __future__ import annotations
import contextlib
import functools
import inspect
import json
import logging
import os
import struct
import threading
from typing import Any, Callable
import torch
@@ -18,6 +20,7 @@ from .cleanup import (
trim_resident_vram,
unload_loaded_model,
)
from .external_residency import EXTERNAL_REGISTRY, external_objects_for_models, external_trim_enabled
from .residency import (
KIND_CHECKPOINT,
KIND_CLIP,
@@ -38,6 +41,8 @@ _WARNED_PICKLE_GPU_PATHS: set[str] = set()
_SAFE_TENSORS_COMPONENT_CACHE_MAX = 32
_STICKY_PROTECTION_VRAM_FLOOR_RATIO = 0.125
_STICKY_PROTECTION_VRAM_FLOOR_CEIL_BYTES = 16 * 1024 ** 3
_TILED_VAE_MEMORY_LOCK_ATTR = "_gpu_resident_loader_tiled_memory_lock"
_TILED_VAE_LOCK_INIT = threading.Lock()
_SAFETENSORS_DTYPE_MAP = {
"BOOL": torch.bool,
"U8": torch.uint8,
@@ -661,6 +666,173 @@ def _sticky_safe_batch_number(*, batch_count: int, free_memory: int, memory_used
return max(1, capped)
def _scaled_batch_memory(total_memory_used: int, total_batch_count: int, batch_number: int) -> int:
total_memory = max(1, int(total_memory_used))
total_batches = max(1, int(total_batch_count))
current_batch = max(1, min(int(batch_number), total_batches))
return max(1, (total_memory * current_batch + total_batches - 1) // total_batches)
def _sticky_vae_free_memory(*, device: Any, patcher: Any) -> int:
import comfy.model_management as model_management
get_free_memory = getattr(model_management, "get_free_memory", None)
if callable(get_free_memory):
try:
return max(0, int(get_free_memory(device)))
except Exception:
pass
return max(0, int(patcher.get_free_memory(device)))
def _sticky_model_load_requirement(*, device: Any, models: tuple[Any, ...]) -> int | None:
import comfy.model_management as model_management
loaded_model_cls = getattr(model_management, "LoadedModel", None)
is_device_cpu = getattr(model_management, "is_device_cpu", None)
if device is None or loaded_model_cls is None:
return None
if callable(is_device_cpu) and is_device_cpu(device):
return 0
requested_models: set[Any] = set()
for model in models:
if model is None:
continue
requested_models.add(model)
additional_models = getattr(model, "model_patches_models", None)
if callable(additional_models):
try:
for additional in additional_models():
if additional is not None:
requested_models.add(additional)
except Exception as exc:
_LOG.debug(
"GPU Resident Loader: failed to enumerate patched models for sticky VAE preflight: %s",
exc,
)
return None
total_required = 0
for model in requested_models:
try:
loaded_model = loaded_model_cls(model)
loaded_device = getattr(loaded_model, "device", None)
if loaded_device != device:
continue
if callable(is_device_cpu) and is_device_cpu(loaded_device):
continue
total_required += max(0, int(loaded_model.model_memory_required(loaded_model.device)))
except Exception:
return None
return total_required
def _sticky_model_load_target(model_load_required: int | None) -> int | None:
if model_load_required is None:
return None
required = max(0, int(model_load_required))
return (required * 11 + 9) // 10
def _prepare_sticky_vae_batch(
*,
device: Any,
patcher: Any,
total_memory_used: int,
total_batch_count: int,
) -> tuple[int, int, bool]:
free_memory = _sticky_vae_free_memory(device=device, patcher=patcher)
model_load_required = _sticky_model_load_requirement(device=device, models=(patcher,))
model_load_target = _sticky_model_load_target(model_load_required)
batch_budget = 0 if model_load_target is None else max(0, free_memory - model_load_target)
batch_number = _sticky_safe_batch_number(
batch_count=total_batch_count,
free_memory=batch_budget,
memory_used=total_memory_used,
device=device,
)
batch_memory_used = _scaled_batch_memory(total_memory_used, total_batch_count, batch_number)
if REGISTRY.get_policy() != "sticky_gpu" or device is None:
return batch_number, batch_memory_used, False
target_free = (
free_memory + 1
if model_load_target is None
else model_load_target + _sticky_protection_target(batch_memory_used, device)
)
if free_memory < target_free:
try:
trim_resident_vram(
device=device,
target_free_vram_bytes=target_free,
respect_sticky=True,
sticky_floor_priority=0,
allow_partial_unload=True,
keep_models=(patcher,),
include_external=False,
)
except Exception as exc:
_LOG.debug(
"GPU Resident Loader: proactive VAE trim failed for batch=%s load=%s bytes: %s",
batch_memory_used,
model_load_required,
exc,
)
free_memory = _sticky_vae_free_memory(device=device, patcher=patcher)
if free_memory < target_free and external_trim_enabled():
try:
trim_resident_vram(
device=device,
target_free_vram_bytes=target_free,
respect_sticky=True,
sticky_floor_priority=0,
allow_partial_unload=True,
keep_models=(patcher,),
include_external=True,
)
except Exception as exc:
_LOG.debug(
"GPU Resident Loader: second-chance external VAE trim failed for batch=%s load=%s bytes: %s",
batch_memory_used,
model_load_required,
exc,
)
else:
_LOG.info(
"GPU Resident Loader: native sticky VAE trim was insufficient; retried with external candidates enabled"
)
free_memory = _sticky_vae_free_memory(device=device, patcher=patcher)
batch_budget = 0 if model_load_target is None else max(0, free_memory - model_load_target)
batch_number = _sticky_safe_batch_number(
batch_count=total_batch_count,
free_memory=batch_budget,
memory_used=total_memory_used,
device=device,
)
batch_memory_used = _scaled_batch_memory(total_memory_used, total_batch_count, batch_number)
target_free = (
free_memory + 1
if model_load_target is None
else model_load_target + _sticky_protection_target(batch_memory_used, device)
)
should_tile = free_memory < target_free and batch_number <= 1
if should_tile:
_LOG.info(
"GPU Resident Loader: skipping regular VAE pass and switching directly to tiled mode; free=%s target=%s batch_memory=%s model_load=%s",
free_memory,
target_free,
batch_memory_used,
model_load_required,
)
return batch_number, batch_memory_used, should_tile
def _wrap_vae_encode(func: Callable[..., Any]) -> Callable[..., Any]:
@functools.wraps(func)
def wrapper(self, pixel_samples):
@@ -680,30 +852,36 @@ def _wrap_vae_encode(func: Callable[..., Any]) -> Callable[..., Any]:
pixel_samples = pixel_samples.unsqueeze(2)
try:
memory_used = self.memory_used_encode(pixel_samples.shape, self.vae_dtype)
model_management.load_models_gpu([self.patcher], memory_required=memory_used, force_full_load=self.disable_offload)
free_memory = self.patcher.get_free_memory(self.device)
batch_number = _sticky_safe_batch_number(
batch_count=pixel_samples.shape[0],
free_memory=free_memory,
memory_used=memory_used,
batch_number, batch_memory_used, should_tile = _prepare_sticky_vae_batch(
device=self.device,
patcher=self.patcher,
total_memory_used=memory_used,
total_batch_count=pixel_samples.shape[0],
)
samples = None
for x in range(0, pixel_samples.shape[0], batch_number):
pixels_in = self.process_input(pixel_samples[x:x + batch_number]).to(self.vae_dtype)
if getattr(self.first_stage_model, "comfy_has_chunked_io", False):
out = self.first_stage_model.encode(pixels_in, device=self.device)
else:
pixels_in = pixels_in.to(self.device)
out = self.first_stage_model.encode(pixels_in)
out = out.to(self.output_device).to(dtype=self.vae_output_dtype())
if samples is None:
samples = torch.empty(
(pixel_samples.shape[0],) + tuple(out.shape[1:]),
device=self.output_device,
dtype=self.vae_output_dtype(),
)
samples[x:x + batch_number] = out
if should_tile:
do_tile = True
else:
model_management.load_models_gpu(
[self.patcher],
memory_required=batch_memory_used,
force_full_load=self.disable_offload,
)
samples = None
for x in range(0, pixel_samples.shape[0], batch_number):
pixels_in = self.process_input(pixel_samples[x:x + batch_number]).to(self.vae_dtype)
if getattr(self.first_stage_model, "comfy_has_chunked_io", False):
out = self.first_stage_model.encode(pixels_in, device=self.device)
else:
pixels_in = pixels_in.to(self.device)
out = self.first_stage_model.encode(pixels_in)
out = out.to(self.output_device).to(dtype=self.vae_output_dtype())
if samples is None:
samples = torch.empty(
(pixel_samples.shape[0],) + tuple(out.shape[1:]),
device=self.output_device,
dtype=self.vae_output_dtype(),
)
samples[x:x + batch_number] = out
except Exception as e:
model_management.raise_non_oom(e)
_LOG.warning("Warning: Ran out of memory when regular VAE encoding, retrying with tiled VAE encoding.")
@@ -725,6 +903,364 @@ def _wrap_vae_encode(func: Callable[..., Any]) -> Callable[..., Any]:
return wrapper
def _default_tiled_vae_axes(
*,
latent_dim: int,
extra_1d_channel: Any,
tile_x: int | None,
tile_y: int | None,
tile_t: int | None,
decode: bool,
) -> tuple[int | None, int | None, int | None]:
if latent_dim == 3:
default_tile_x = 32 if decode else 512
default_tile_y = 32 if decode else 512
default_tile_t = 999 if decode else 9999
elif latent_dim == 1 or extra_1d_channel is not None:
default_tile_x = 256 * 2048
default_tile_y = None
default_tile_t = None
else:
default_tile_x = 64 if decode else 512
default_tile_y = 64 if decode else 512
default_tile_t = None
resolved_tile_x = default_tile_x if tile_x is None else max(1, int(tile_x))
resolved_tile_y = default_tile_y if tile_y is None else max(1, int(tile_y))
resolved_tile_t = default_tile_t if tile_t is None else max(1, int(tile_t))
return resolved_tile_x, resolved_tile_y, resolved_tile_t
def _shape_with_capped_tail(shape: tuple[int, ...], tail_caps: dict[int, int | None]) -> tuple[int, ...]:
capped = list(shape)
for index, cap in tail_caps.items():
if cap is None:
continue
capped[index] = min(int(capped[index]), max(1, int(cap)))
return tuple(capped)
def _tiled_vae_memory_shapes(
*,
shape: tuple[int, ...],
latent_dim: int,
extra_1d_channel: Any,
tile_x: int | None,
tile_y: int | None,
tile_t: int | None,
decode: bool,
) -> list[tuple[int, ...]]:
resolved_tile_x, resolved_tile_y, resolved_tile_t = _default_tiled_vae_axes(
latent_dim=latent_dim,
extra_1d_channel=extra_1d_channel,
tile_x=tile_x,
tile_y=tile_y,
tile_t=tile_t,
decode=decode,
)
if latent_dim == 3:
return [
_shape_with_capped_tail(
shape,
{
len(shape) - 3: resolved_tile_t,
len(shape) - 2: resolved_tile_y,
len(shape) - 1: resolved_tile_x,
},
)
]
if latent_dim == 1 or extra_1d_channel is not None:
return [_shape_with_capped_tail(shape, {len(shape) - 1: resolved_tile_x})]
if decode:
return [
_shape_with_capped_tail(
shape,
{
len(shape) - 2: resolved_tile_y,
len(shape) - 1: resolved_tile_x,
},
)
]
return [
_shape_with_capped_tail(
shape,
{
len(shape) - 2: resolved_tile_y,
len(shape) - 1: resolved_tile_x,
},
),
_shape_with_capped_tail(
shape,
{
len(shape) - 2: max(1, resolved_tile_y // 2),
len(shape) - 1: max(1, resolved_tile_x * 2),
},
),
_shape_with_capped_tail(
shape,
{
len(shape) - 2: max(1, resolved_tile_y * 2),
len(shape) - 1: max(1, resolved_tile_x // 2),
},
),
]
@contextlib.contextmanager
def _temporary_tiled_vae_memory_estimate(
self,
*,
decode: bool,
tile_x: int | None,
tile_y: int | None,
tile_t: int | None,
) -> Any:
memory_attr = "memory_used_decode" if decode else "memory_used_encode"
original = getattr(self, memory_attr, None)
if not callable(original):
yield
return
had_instance_attr = memory_attr in getattr(self, "__dict__", {})
lock = getattr(self, _TILED_VAE_MEMORY_LOCK_ATTR, None)
if lock is None:
with _TILED_VAE_LOCK_INIT:
lock = getattr(self, _TILED_VAE_MEMORY_LOCK_ATTR, None)
if lock is None:
lock = threading.RLock()
setattr(self, _TILED_VAE_MEMORY_LOCK_ATTR, lock)
def estimated(shape, dtype, *args, **kwargs):
shapes = _tiled_vae_memory_shapes(
shape=tuple(int(dim) for dim in shape),
latent_dim=int(getattr(self, "latent_dim", 2)),
extra_1d_channel=getattr(self, "extra_1d_channel", None),
tile_x=tile_x,
tile_y=tile_y,
tile_t=tile_t,
decode=decode,
)
return max(int(original(candidate, dtype, *args, **kwargs)) for candidate in shapes)
with lock:
setattr(self, memory_attr, estimated)
try:
yield
finally:
if had_instance_attr:
setattr(self, memory_attr, original)
else:
delattr(self, memory_attr)
@functools.lru_cache(maxsize=None)
def _tiled_vae_supported_kwargs(func: Callable[..., Any]) -> frozenset[str]:
return frozenset(inspect.signature(func).parameters)
def _call_tiled_vae(
func: Callable[..., Any],
self,
data,
*,
tile_x=None,
tile_y=None,
overlap=None,
tile_t=None,
overlap_t=None,
):
kwargs = {}
supported_kwargs = _tiled_vae_supported_kwargs(func)
if "tile_x" in supported_kwargs:
kwargs["tile_x"] = tile_x
if "tile_y" in supported_kwargs:
kwargs["tile_y"] = tile_y
if "overlap" in supported_kwargs:
kwargs["overlap"] = overlap
if "tile_t" in supported_kwargs:
kwargs["tile_t"] = tile_t
if "overlap_t" in supported_kwargs:
kwargs["overlap_t"] = overlap_t
return func(self, data, **kwargs)
def _should_prefer_tiled_vae_encode(vae: Any, pixel_samples: Any) -> bool:
if REGISTRY.get_policy() != "sticky_gpu":
return False
if vae is None or pixel_samples is None:
return False
try:
vae.throw_exception_if_invalid()
prepared = vae.vae_encode_crop_pixels(pixel_samples)
prepared = prepared.movedim(-1, 1)
if int(getattr(vae, "latent_dim", 2)) == 3 and prepared.ndim < 5:
if not getattr(vae, "not_video", False):
prepared = prepared.movedim(1, 0).unsqueeze(0)
else:
prepared = prepared.unsqueeze(2)
memory_used = vae.memory_used_encode(prepared.shape, vae.vae_dtype)
_, _, should_tile = _prepare_sticky_vae_batch(
device=getattr(vae, "device", None),
patcher=getattr(vae, "patcher", None),
total_memory_used=memory_used,
total_batch_count=prepared.shape[0],
)
return bool(should_tile)
except Exception as exc:
_LOG.debug("GPU Resident Loader: failed to preflight sticky VAE encode preference: %s", exc)
return False
def _call_bound_tiled_vae(func: Callable[..., Any], pixel_samples: Any, *args: Any, **kwargs: Any) -> Any:
supported_kwargs = _tiled_vae_supported_kwargs(getattr(func, "__func__", func))
filtered_kwargs = {key: value for key, value in kwargs.items() if key in supported_kwargs}
return func(pixel_samples, *args, **filtered_kwargs)
@contextlib.contextmanager
def _temporary_prefer_tiled_vae_encode(vae: Any):
original_encode = getattr(vae, "encode", None)
encode_tiled = getattr(vae, "encode_tiled", None)
if not callable(original_encode) or not callable(encode_tiled):
yield
return
had_instance_attr = "encode" in getattr(vae, "__dict__", {})
lock = getattr(vae, _TILED_VAE_MEMORY_LOCK_ATTR, None)
if lock is None:
with _TILED_VAE_LOCK_INIT:
lock = getattr(vae, _TILED_VAE_MEMORY_LOCK_ATTR, None)
if lock is None:
lock = threading.RLock()
setattr(vae, _TILED_VAE_MEMORY_LOCK_ATTR, lock)
def prefer_encode(pixel_samples, *args, **kwargs):
if _should_prefer_tiled_vae_encode(vae, pixel_samples):
return _call_bound_tiled_vae(encode_tiled, pixel_samples, *args, **kwargs)
return original_encode(pixel_samples, *args, **kwargs)
with lock:
setattr(vae, "encode", prefer_encode)
try:
yield
finally:
if had_instance_attr:
setattr(vae, "encode", original_encode)
else:
delattr(vae, "encode")
def _wrap_vae_encode_for_inpaint_node(func: Callable[..., Any]) -> Callable[..., Any]:
supported_kwargs = frozenset(inspect.signature(func).parameters)
@functools.wraps(func)
def wrapper(self, *args, **kwargs):
# Filter kwargs to only include those supported by the wrapped function
filtered_kwargs = {key: value for key, value in kwargs.items() if key in supported_kwargs}
# Supply default grow_mask_by if not present and supported
if "grow_mask_by" in supported_kwargs and "grow_mask_by" not in filtered_kwargs:
filtered_kwargs["grow_mask_by"] = 6
if REGISTRY.get_policy() != "sticky_gpu":
return func(self, *args, **filtered_kwargs)
# Extract vae from args for the context manager
vae = args[0] if args else kwargs.get("vae")
with _temporary_prefer_tiled_vae_encode(vae):
return func(self, *args, **filtered_kwargs)
return wrapper
def _wrap_inpaint_model_conditioning_node(func: Callable[..., Any]) -> Callable[..., Any]:
@functools.wraps(func)
def wrapper(self, positive, negative, pixels, vae, mask, noise_mask=True):
if REGISTRY.get_policy() != "sticky_gpu":
return func(self, positive, negative, pixels, vae, mask, noise_mask=noise_mask)
with _temporary_prefer_tiled_vae_encode(vae):
return func(self, positive, negative, pixels, vae, mask, noise_mask=noise_mask)
return wrapper
def _wrap_vae_encode_tiled(func: Callable[..., Any]) -> Callable[..., Any]:
@functools.wraps(func)
def wrapper(self, pixel_samples, tile_x=None, tile_y=None, overlap=None, tile_t=None, overlap_t=None):
if REGISTRY.get_policy() != "sticky_gpu":
return _call_tiled_vae(
func,
self,
pixel_samples,
tile_x=tile_x,
tile_y=tile_y,
overlap=overlap,
tile_t=tile_t,
overlap_t=overlap_t,
)
with _temporary_tiled_vae_memory_estimate(
self,
decode=False,
tile_x=tile_x,
tile_y=tile_y,
tile_t=tile_t,
):
return _call_tiled_vae(
func,
self,
pixel_samples,
tile_x=tile_x,
tile_y=tile_y,
overlap=overlap,
tile_t=tile_t,
overlap_t=overlap_t,
)
return wrapper
def _wrap_vae_decode_tiled(func: Callable[..., Any]) -> Callable[..., Any]:
@functools.wraps(func)
def wrapper(self, samples, tile_x=None, tile_y=None, overlap=None, tile_t=None, overlap_t=None):
if REGISTRY.get_policy() != "sticky_gpu":
return _call_tiled_vae(
func,
self,
samples,
tile_x=tile_x,
tile_y=tile_y,
overlap=overlap,
tile_t=tile_t,
overlap_t=overlap_t,
)
with _temporary_tiled_vae_memory_estimate(
self,
decode=True,
tile_x=tile_x,
tile_y=tile_y,
tile_t=tile_t,
):
return _call_tiled_vae(
func,
self,
samples,
tile_x=tile_x,
tile_y=tile_y,
overlap=overlap,
tile_t=tile_t,
overlap_t=overlap_t,
)
return wrapper
def _wrap_vae_decode(func: Callable[..., Any]) -> Callable[..., Any]:
@functools.wraps(func)
def wrapper(self, samples_in, vae_options={}):
@@ -740,43 +1276,48 @@ def _wrap_vae_decode(func: Callable[..., Any]) -> Callable[..., Any]:
samples_in = samples_in[:, :, 0]
try:
memory_used = self.memory_used_decode(samples_in.shape, self.vae_dtype)
model_management.load_models_gpu([self.patcher], memory_required=memory_used, force_full_load=self.disable_offload)
free_memory = self.patcher.get_free_memory(self.device)
batch_number = _sticky_safe_batch_number(
batch_count=samples_in.shape[0],
free_memory=free_memory,
memory_used=memory_used,
batch_number, batch_memory_used, should_tile = _prepare_sticky_vae_batch(
device=self.device,
patcher=self.patcher,
total_memory_used=memory_used,
total_batch_count=samples_in.shape[0],
)
preallocated = False
if getattr(self.first_stage_model, "comfy_has_chunked_io", False):
pixel_samples = torch.empty(
self.first_stage_model.decode_output_shape(samples_in.shape),
device=self.output_device,
dtype=self.vae_output_dtype(),
if should_tile:
do_tile = True
else:
model_management.load_models_gpu(
[self.patcher],
memory_required=batch_memory_used,
force_full_load=self.disable_offload,
)
preallocated = True
for x in range(0, samples_in.shape[0], batch_number):
samples = samples_in[x:x + batch_number].to(device=self.device, dtype=self.vae_dtype)
if preallocated:
self.first_stage_model.decode(samples, output_buffer=pixel_samples[x:x + batch_number], **vae_options)
else:
out = self.first_stage_model.decode(samples, **vae_options).to(
if getattr(self.first_stage_model, "comfy_has_chunked_io", False):
pixel_samples = torch.empty(
self.first_stage_model.decode_output_shape(samples_in.shape),
device=self.output_device,
dtype=self.vae_output_dtype(),
copy=True,
)
if pixel_samples is None:
pixel_samples = torch.empty(
(samples_in.shape[0],) + tuple(out.shape[1:]),
preallocated = True
for x in range(0, samples_in.shape[0], batch_number):
samples = samples_in[x:x + batch_number].to(device=self.device, dtype=self.vae_dtype)
if preallocated:
self.first_stage_model.decode(samples, output_buffer=pixel_samples[x:x + batch_number], **vae_options)
else:
out = self.first_stage_model.decode(samples, **vae_options).to(
device=self.output_device,
dtype=self.vae_output_dtype(),
copy=True,
)
pixel_samples[x:x + batch_number].copy_(out)
del out
self.process_output(pixel_samples[x:x + batch_number])
if pixel_samples is None:
pixel_samples = torch.empty(
(samples_in.shape[0],) + tuple(out.shape[1:]),
device=self.output_device,
dtype=self.vae_output_dtype(),
)
pixel_samples[x:x + batch_number].copy_(out)
del out
self.process_output(pixel_samples[x:x + batch_number])
except Exception as e:
model_management.raise_non_oom(e)
_LOG.warning("Warning: Ran out of memory when regular VAE decoding, retrying with tiled VAE decoding.")
@@ -907,6 +1448,28 @@ def _wrap_model_patcher_detach(func: Callable[..., Any]) -> Callable[..., Any]:
def _wrap_free_memory(func: Callable[..., Any]) -> Callable[..., Any]:
"""
Wraps a free-memory function to enforce sticky-GPU protection and external fallback trimming.
When the registry policy is "sticky_gpu" and a device is provided, the wrapper:
- Reserves VRAM for sticky-loaded models by attempting a pre-trim to a computed protection target.
- Protects a subset of sticky-loaded wrappers from unloading when calling the original function by adding them to `keep_loaded`.
- After the original free-memory call, refreshes both REGISTRY and EXTERNAL_REGISTRY runtime state.
- If external trimming is available and still needed, attempts a fallback trim that includes external residency.
Dynamic free-memory calls are excluded because Comfy reduces their effective target internally.
Parameters:
memory_required: Number of bytes the caller needs to free.
device: Target device for which memory is being freed (may be None).
keep_loaded: Iterable of loaded-wrapper objects that must be kept; the wrapper may extend this list with additional protected wrappers.
Returns:
The value returned by the wrapped `func`.
Notes:
- The wrapper may call `trim_resident_vram` and `model_management.get_free_memory`; exceptions from trimming or free-memory queries are caught and logged, not propagated.
- Side effects include invoking trims and refreshing runtime state on REGISTRY and EXTERNAL_REGISTRY.
"""
@functools.wraps(func)
def wrapper(memory_required, device, keep_loaded=None, *args, **kwargs):
import comfy.model_management as model_management
@@ -961,6 +1524,42 @@ def _wrap_free_memory(func: Callable[..., Any]) -> Callable[..., Any]:
unloaded = func(memory_required, device, keep_loaded + protected_wrappers, *args, **kwargs)
REGISTRY.refresh_runtime_state()
EXTERNAL_REGISTRY.refresh_runtime_state()
for_dynamic = bool(kwargs.get("for_dynamic", args[0] if args else False))
if device is not None and external_trim_enabled() and not for_dynamic:
fallback_target = memory_required
if REGISTRY.get_policy() == "sticky_gpu":
fallback_target = max(fallback_target, _sticky_protection_target(memory_required, device))
try:
free_now = model_management.get_free_memory(device)
except Exception:
free_now = None
if free_now is not None and int(free_now) < int(fallback_target):
protected_models = tuple(
model
for model in (getattr(loaded_wrapper, "model", None) for loaded_wrapper in keep_loaded + protected_wrappers)
if model is not None
)
keep_models = protected_models + external_objects_for_models(protected_models)
try:
trim_resident_vram(
device=device,
target_free_vram_bytes=int(fallback_target),
respect_sticky=True,
sticky_floor_priority=0,
allow_partial_unload=True,
keep_models=keep_models,
include_external=True,
)
except Exception as exc:
_LOG.debug(
"GPU Resident Loader: external fallback trim failed for free_memory(%s): %s",
memory_required,
exc,
)
REGISTRY.refresh_runtime_state()
EXTERNAL_REGISTRY.refresh_runtime_state()
return unloaded
return wrapper
@@ -1037,6 +1636,7 @@ def install_patches() -> None:
import comfy.model_patcher as model_patcher
import comfy.sd as comfy_sd
import comfy.utils as comfy_utils
import nodes as comfy_nodes
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:
@@ -1084,6 +1684,30 @@ def install_patches() -> None:
if comfy_sd.VAE.decode is original_vae_decode:
comfy_sd.VAE.decode = _wrap_vae_decode(original_vae_decode)
original_vae_encode_tiled = _remember_original("sd.VAE.encode_tiled", comfy_sd.VAE.encode_tiled)
if comfy_sd.VAE.encode_tiled is original_vae_encode_tiled:
comfy_sd.VAE.encode_tiled = _wrap_vae_encode_tiled(original_vae_encode_tiled)
original_vae_decode_tiled = _remember_original("sd.VAE.decode_tiled", comfy_sd.VAE.decode_tiled)
if comfy_sd.VAE.decode_tiled is original_vae_decode_tiled:
comfy_sd.VAE.decode_tiled = _wrap_vae_decode_tiled(original_vae_decode_tiled)
if hasattr(comfy_nodes, "VAEEncodeForInpaint") and hasattr(comfy_nodes.VAEEncodeForInpaint, "encode"):
original_vae_encode_for_inpaint = _remember_original(
"nodes.VAEEncodeForInpaint.encode",
comfy_nodes.VAEEncodeForInpaint.encode,
)
if comfy_nodes.VAEEncodeForInpaint.encode is original_vae_encode_for_inpaint:
comfy_nodes.VAEEncodeForInpaint.encode = _wrap_vae_encode_for_inpaint_node(original_vae_encode_for_inpaint)
if hasattr(comfy_nodes, "InpaintModelConditioning") and hasattr(comfy_nodes.InpaintModelConditioning, "encode"):
original_inpaint_model_conditioning = _remember_original(
"nodes.InpaintModelConditioning.encode",
comfy_nodes.InpaintModelConditioning.encode,
)
if comfy_nodes.InpaintModelConditioning.encode is original_inpaint_model_conditioning:
comfy_nodes.InpaintModelConditioning.encode = _wrap_inpaint_model_conditioning_node(original_inpaint_model_conditioning)
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(