Compare commits
32
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f1519f8da3 | ||
|
|
50e8188e5a | ||
|
|
32dee8e647 | ||
|
|
9d5d5e3b42 | ||
|
|
026d8527f8 | ||
|
|
84d10add70 | ||
|
|
29b0a34865 | ||
|
|
fd9f33f17f | ||
|
|
fc7c0f6946 | ||
|
|
d56a716418 | ||
|
|
27aa644775 | ||
|
|
e345e5f9d0 | ||
|
|
43fce69abd | ||
|
|
bafd34d48b | ||
|
|
cf6843fd70 | ||
|
|
5c623b9cdb | ||
|
|
bd249781b3 | ||
|
|
1d1fa53828 | ||
|
|
6342d4c3ad | ||
|
|
f60a8f9972 | ||
|
|
61ed3e16d4 | ||
|
|
d827213bb2 | ||
|
|
8998be1d78 | ||
|
|
9d34f008ab | ||
|
|
1f282ea6f0 | ||
|
|
79b9565867 | ||
|
|
2eb0817c12 | ||
|
|
334390a739 | ||
|
|
815590498b | ||
|
|
ac42debea7 | ||
|
|
baba29152f | ||
|
|
37ed1899ae |
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user