Compare commits
43
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e345e5f9d0 | ||
|
|
43fce69abd | ||
|
|
bafd34d48b | ||
|
|
cf6843fd70 | ||
|
|
5c623b9cdb | ||
|
|
bd249781b3 | ||
|
|
1d1fa53828 | ||
|
|
6342d4c3ad | ||
|
|
f60a8f9972 | ||
|
|
61ed3e16d4 | ||
|
|
d827213bb2 | ||
|
|
8998be1d78 | ||
|
|
9d34f008ab | ||
|
|
1f282ea6f0 | ||
|
|
79b9565867 | ||
|
|
2eb0817c12 | ||
|
|
334390a739 | ||
|
|
815590498b | ||
|
|
ac42debea7 | ||
|
|
baba29152f | ||
|
|
37ed1899ae | ||
|
|
a56a64b62e | ||
|
|
9719017900 | ||
|
|
cbe4825973 | ||
|
|
103613ad4f | ||
|
|
9d714f896c | ||
|
|
b0a4c6d7a2 | ||
|
|
5e427f4196 | ||
|
|
e7e4df92bf | ||
|
|
a7dc557376 | ||
|
|
3f1d995138 | ||
|
|
3bbc263a46 | ||
|
|
2dc41ab929 | ||
|
|
7dd33f3675 | ||
|
|
4f36431bf2 | ||
|
|
8716f77749 | ||
|
|
420b27ba4f | ||
|
|
342b80d0ec | ||
|
|
fcd4d80f7b | ||
|
|
59ac769635 | ||
|
|
c121278c65 | ||
|
|
da3f105661 | ||
|
|
1365161b49 |
@@ -1,32 +1,44 @@
|
||||
# ComfyUI GPU Resident Loader
|
||||
|
||||
A ComfyUI custom-node pack that targets **time-to-VRAM** and **sticky GPU residency**, not just lower peak host RAM.
|
||||
A ComfyUI custom-node pack for **faster time-to-VRAM**, **selective safetensors loading**, **sticky GPU residency control**, and **visibility into compatible external GPU model caches**.
|
||||
|
||||
It does two things:
|
||||
This repo does three related jobs:
|
||||
|
||||
1. **Installs startup-time loader and residency patches** before any workflow nodes run.
|
||||
2. Ships **KJ-compatible loader nodes** for diffusion models and checkpoints, plus preload/pin/evict/report nodes for manual residency control.
|
||||
1. **Installs startup-time monkey patches** before any workflow nodes run.
|
||||
2. Ships **KJ-style resident loader nodes** for diffusion models and checkpoints.
|
||||
3. Maintains a **live residency system** for native ComfyUI objects and compatible external caches, with preload / pin / evict / report controls for native tracked objects and automatic snapshot support plus provider-specific eviction for compatible external entries.
|
||||
|
||||
It is not just a “clean RAM” addon. The main target is the path from model file -> tensors -> live ComfyUI object -> VRAM retention, including GPU-resident caches that live outside ComfyUI’s normal loaded-model list.
|
||||
|
||||
## Why this exists
|
||||
|
||||
Stock ComfyUI makes separate decisions for:
|
||||
ComfyUI’s default behavior mixes together two separate concerns:
|
||||
|
||||
- where a model **lives after load**, and
|
||||
- where checkpoint tensors are **materialized first**.
|
||||
- **ingest path** — where tensors are first materialized while a model is being read, and
|
||||
- **residency policy** — where the finished model tends to live afterwards.
|
||||
|
||||
Those are not the same thing.
|
||||
Those are not the same problem.
|
||||
|
||||
This repo targets the second problem directly for `.safetensors` by steering eligible loads toward direct GPU ingest, narrowing resident diffusion-model loads down to UNet tensors only, then targets the first problem by overriding offload policy and by teaching `free_memory()` to protect high-priority sticky entries until the VRAM budget says otherwise.
|
||||
This repo focuses on both:
|
||||
|
||||
## What is included
|
||||
- For **`.safetensors`**, it tries to keep eligible loads on the narrowest, most GPU-friendly path it can.
|
||||
- For **resident diffusion and checkpoint-model loads**, it avoids broad checkpoint materialization by selecting only the detected UNet keys where possible.
|
||||
- For **runtime VRAM pressure**, it adds a sticky-priority registry and teaches ComfyUI’s unload path to protect higher-value resident entries until enough VRAM must be reclaimed.
|
||||
- For **compatible external GPU caches** that bypass `comfy.model_management.current_loaded_models`, it can discover supported providers at runtime and include their entries in snapshot output and provider-specific eviction decisions, with automatic trim kept opt-in.
|
||||
- For **manual control**, it exposes nodes that let you preload, pin, evict, and inspect tracked native models, CLIPs, and VAEs.
|
||||
|
||||
### Startup patcher
|
||||
## What changes at startup
|
||||
|
||||
Installed automatically from `__init__.py` when the custom node loads.
|
||||
`__init__.py` calls `startup.install_patches()`, which applies the core monkey patches exactly once when the custom node is imported.
|
||||
|
||||
Those startup patches cover the built-in ComfyUI load and residency paths below. Compatible external-cache integrations are installed lazily later, on demand, when a supported module is actually present in the running process.
|
||||
|
||||
### Patched functions / methods
|
||||
|
||||
Current patch surface:
|
||||
|
||||
- `comfy.utils.load_torch_file`
|
||||
- `comfy.clip_vision.load_torch_file` (redirected to the patched `comfy.utils.load_torch_file` when present)
|
||||
- `comfy.model_management.free_memory`
|
||||
- `comfy.model_management.load_models_gpu`
|
||||
- `comfy.model_management.unet_offload_device`
|
||||
@@ -35,75 +47,437 @@ Current patch surface:
|
||||
- `comfy.model_management.text_encoder_device`
|
||||
- `comfy.model_management.vae_device`
|
||||
- `comfy.model_management.unet_inital_load_device`
|
||||
- `comfy.model_management.LoadedModel.model_unload`
|
||||
- `comfy.model_patcher.ModelPatcher.detach`
|
||||
- `comfy.sd.load_checkpoint_guess_config`
|
||||
- `comfy.sd.load_diffusion_model`
|
||||
- `comfy.sd.load_clip`
|
||||
- `comfy.sd.VAE.encode`
|
||||
- `comfy.sd.VAE.decode`
|
||||
- `comfy.clip_vision.load`
|
||||
- `comfy.controlnet.load_controlnet`
|
||||
- `comfy.diffusers_load.load_diffusers`
|
||||
|
||||
### Loader nodes
|
||||
### Lazily installed external integrations
|
||||
|
||||
- **Diffusion Model Loader Resident**
|
||||
- **Checkpoint Loader Resident**
|
||||
- **Checkpoint Model Loader Resident**
|
||||
- **Checkpoint Clip Loader Resident**
|
||||
- **Checkpoint VAE Loader Resident**
|
||||
- **Diffusion Model Selector Resident**
|
||||
Current external integration surface:
|
||||
|
||||
`Diffusion Model Loader Resident` mirrors the relevant KJ diffusion-loader feature surface:
|
||||
- compatible **SeedVR2** `src/core/model_cache.py` modules discovered at runtime
|
||||
|
||||
- weight dtype override
|
||||
- compute dtype override
|
||||
- cublas-ops toggle
|
||||
- SageAttention override
|
||||
- fp16 accumulation toggle
|
||||
- optional extra-state-dict merge
|
||||
When a compatible SeedVR2 cache module is present, the repo wraps:
|
||||
|
||||
On `.safetensors`, the resident diffusion-model path now reads only the detected UNet keys and merges only matching keys from any extra state dict. Repeated resident loads also reuse a live equivalent object when the source path and loader-relevant options still match.
|
||||
- `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`
|
||||
|
||||
### Residency nodes
|
||||
That lazy integration lets the loader:
|
||||
|
||||
- **Set Global Residency Policy**
|
||||
- **Registry Snapshot**
|
||||
- **Pin Model/CLIP/VAE Residency**
|
||||
- **Preload Model/CLIP/VAE To GPU**
|
||||
- **Evict Model/CLIP/VAE From GPU**
|
||||
- **Report Model/CLIP/VAE Residency**
|
||||
- mirror SeedVR2-owned cached **DiT** and **VAE** objects into a separate external residency registry
|
||||
- refresh byte / device / claimed-state metadata from the live cached object
|
||||
- evict those entries through SeedVR2’s own removal path instead of assuming they live in `comfy.model_management.current_loaded_models`
|
||||
|
||||
## What those patches do
|
||||
|
||||
### 1) `load_torch_file` becomes residency-aware
|
||||
|
||||
The patched loader:
|
||||
|
||||
- detects the active load context (`model`, `clip`, `vae`, `checkpoint`, etc.)
|
||||
- picks an explicit GPU target device when the active policy wants GPU ingest
|
||||
- attempts **direct safetensors reads** on the requested device
|
||||
- falls back to **CPU read + tensor-by-tensor copy** if direct GPU safetensors loading fails
|
||||
- still uses **CPU-first `torch.load()`** for pickle formats (`.ckpt`, `.pt`, `.pth`, `.bin`)
|
||||
- records the actual load method in the residency registry
|
||||
|
||||
For safetensors loads happening inside a `model` / `clip` / `vae` context, it can select only the detected component keys from the file header instead of pulling the full file into memory first.
|
||||
|
||||
### 2) Loader contexts are attached to stock ComfyUI load paths
|
||||
|
||||
These stock paths are wrapped with registry context and output binding:
|
||||
|
||||
- checkpoint loads
|
||||
- diffusion-model loads
|
||||
- CLIP loads
|
||||
- CLIP Vision loads
|
||||
- ControlNet loads
|
||||
- diffusers loads
|
||||
|
||||
That means the native registry is not limited to the custom resident nodes. Stock ComfyUI loaders that pass through these paths are also tracked.
|
||||
|
||||
### 3) Compatible external caches can join the residency system lazily
|
||||
|
||||
When a compatible SeedVR2 cache module is present, the loader installs cache-level hooks that register SeedVR2-owned cached DiT / VAE objects into a separate external registry.
|
||||
|
||||
Those entries:
|
||||
|
||||
- are refreshed from the live cached object at runtime
|
||||
- appear in **Registry Snapshot** output under `external_entries`
|
||||
- are evicted through SeedVR2’s own cache-removal methods rather than the normal Comfy unload path
|
||||
- remain outside automatic trim unless `COMFYUI_GPU_RESIDENT_TRIM_EXTERNAL=1` is set
|
||||
|
||||
### 4) Device/offload policy is overridden
|
||||
|
||||
Depending on the active policy, the patcher can steer:
|
||||
|
||||
- initial UNet load device
|
||||
- CLIP/Text Encoder device
|
||||
- VAE device
|
||||
- offload devices for UNet / CLIP / VAE
|
||||
|
||||
This is how `prefer_gpu` and `sticky_gpu` keep more of the hot path on the GPU side than stock ComfyUI would.
|
||||
|
||||
### 5) `free_memory()` becomes sticky-aware
|
||||
|
||||
Under `sticky_gpu`, `comfy.model_management.free_memory()` is patched so that:
|
||||
|
||||
- sticky tracked wrappers are considered first
|
||||
- higher-priority sticky entries are protected first
|
||||
- lower-priority or older sticky entries yield first when VRAM must be reclaimed
|
||||
- a transient protection floor is applied so ComfyUI does not immediately tear down high-value resident entries for small requests
|
||||
|
||||
### 6) Clone replacement is hardened
|
||||
|
||||
`load_models_gpu()` is patched to fully unload clone-conflict wrappers before replacement instead of relying on a shallow detach path that can leave base weights patched.
|
||||
|
||||
### 7) Unload / detach is redirected to CPU when needed
|
||||
|
||||
`LoadedModel.model_unload()` and `ModelPatcher.detach()` are patched so that unloads which would otherwise not reclaim VRAM are redirected through a CPU offload target first.
|
||||
|
||||
### 8) VAE encode/decode gets a sticky-safe path
|
||||
|
||||
Under `sticky_gpu`, patched `VAE.encode()` and `VAE.decode()`:
|
||||
|
||||
- cap the working batch count when necessary to preserve transient VRAM headroom
|
||||
- retry with tiled VAE encode/decode on OOM
|
||||
|
||||
That behavior is not a general performance feature toggle. It exists to reduce avoidable VRAM spikes while sticky residency is active.
|
||||
|
||||
## Policies
|
||||
|
||||
The startup patcher exposes four policies:
|
||||
The registry exposes four global policies:
|
||||
|
||||
- `legacy` — leave ingest/offload behavior close to stock ComfyUI.
|
||||
- `balanced` — keep the registry and diagnostics, but do not aggressively steer ingest to GPU.
|
||||
- `prefer_gpu` — prefer GPU ingest, keep UNet/ControlNet/CLIP on the faster side of the device policy, but do not auto-pin tracked objects.
|
||||
- `sticky_gpu` — prefer GPU ingest, auto-pin the highest-value tracked outputs, and let lower-priority sticky entries yield first when VRAM pressure rises.
|
||||
### `legacy`
|
||||
|
||||
Default selection order:
|
||||
Stay closest to stock ComfyUI behavior. Registry tracking still exists, but the patcher does not aggressively steer ingest/offload toward the GPU path.
|
||||
|
||||
1. `COMFYUI_GPU_RESIDENT_POLICY` environment variable, if set.
|
||||
2. `sticky_gpu` when `--gpu_only` is active.
|
||||
3. `sticky_gpu` when `--highvram` is active.
|
||||
4. otherwise `prefer_gpu`.
|
||||
### `balanced`
|
||||
|
||||
## Important scope limits
|
||||
Keep registry tracking and diagnostics without aggressive GPU residency behavior.
|
||||
|
||||
### Best path: `.safetensors`
|
||||
### `prefer_gpu`
|
||||
|
||||
This repo is optimized around `.safetensors`.
|
||||
Prefer GPU ingest for tracked model-like loads and keep the faster side of the device/offload policy for:
|
||||
|
||||
Direct GPU ingest is attempted for `.safetensors` loads. Resident diffusion-model loads take a header-first selective path and fetch only the detected UNet tensors instead of loading the whole file and filtering later. If the direct path fails, the patcher falls back to CPU read + GPU copy and records that fallback in the registry.
|
||||
- diffusion models
|
||||
- CLIP / text encoders
|
||||
- ControlNets
|
||||
|
||||
### `.ckpt` / `.pt` remain CPU-first under PyTorch
|
||||
This policy does **not** auto-pin tracked objects.
|
||||
|
||||
Those formats still go through `torch.load()`. The repo tracks that path and can still keep the resulting model hot in VRAM, but it does **not** claim true direct-to-GPU checkpoint ingest for pickle-based formats. Resident checkpoint nodes now warn about this when a GPU-resident policy is active so the compatibility path is not mistaken for the fast path.
|
||||
### `sticky_gpu`
|
||||
|
||||
Use the included conversion helper to migrate hot models to `.safetensors`.
|
||||
Builds on `prefer_gpu` and additionally:
|
||||
|
||||
- auto-pins newly bound **models** and **CLIPs**
|
||||
- keeps **VAE offload** on the GPU side as well
|
||||
- patches `free_memory()` to protect sticky tracked wrappers by priority
|
||||
- uses the sticky-safe VAE encode/decode behavior
|
||||
|
||||
### Default policy selection
|
||||
|
||||
Selection order is:
|
||||
|
||||
1. `COMFYUI_GPU_RESIDENT_POLICY`, if set to a supported value
|
||||
2. `sticky_gpu` when ComfyUI is started with `--gpu_only`
|
||||
3. `sticky_gpu` when ComfyUI is started with `--highvram`
|
||||
4. otherwise `prefer_gpu`
|
||||
|
||||
Supported values are:
|
||||
|
||||
- `legacy`
|
||||
- `balanced`
|
||||
- `prefer_gpu`
|
||||
- `sticky_gpu`
|
||||
|
||||
## Included nodes
|
||||
|
||||
All nodes live under the `GPU Resident Loader` category.
|
||||
|
||||
### Loader nodes
|
||||
|
||||
#### Diffusion Model Selector Resident
|
||||
|
||||
Returns an absolute path string for a selected diffusion model.
|
||||
|
||||
Notes:
|
||||
|
||||
- resolves from `diffusion_models`
|
||||
- also exposes `text_encoders` entries whose filename contains `connector`
|
||||
|
||||
#### Diffusion Model Loader Resident
|
||||
|
||||
KJ-style diffusion-model loader with these controls:
|
||||
|
||||
- `weight_dtype`
|
||||
- `compute_dtype`
|
||||
- `patch_cublaslinear`
|
||||
- `sage_attention`
|
||||
- `enable_fp16_accumulation`
|
||||
- optional `extra_state_dict`
|
||||
- optional `policy_override`
|
||||
|
||||
Behavior:
|
||||
|
||||
- for `.safetensors`, it loads only the detected UNet portion of the file
|
||||
- if `extra_state_dict` is provided, only matching UNet keys are merged
|
||||
- repeated loads reuse a live equivalent model when the source path and loader-relevant options still match
|
||||
- before GPU-bound loads, it estimates the upcoming footprint and trims only enough lower-priority residency to cover the request plus adaptive headroom
|
||||
|
||||
#### Checkpoint Loader Resident
|
||||
|
||||
Full checkpoint loader that returns:
|
||||
|
||||
- `MODEL`
|
||||
- `CLIP`
|
||||
- `VAE`
|
||||
|
||||
Behavior:
|
||||
|
||||
- shares the same tuning knobs as the resident diffusion-model loader for the model component
|
||||
- reuses already-live equivalent components when possible
|
||||
- composes the final output from model / clip / vae component loaders instead of always rebuilding the whole checkpoint path from scratch
|
||||
|
||||
#### Checkpoint Model Loader Resident
|
||||
|
||||
Model-only checkpoint loader.
|
||||
|
||||
Behavior:
|
||||
|
||||
- takes the same selective safetensors UNet fast path as the diffusion-model loader
|
||||
- reuses a live equivalent model when available
|
||||
- uses the same dtype / attention / cublas / fp16-accumulation knobs as the full checkpoint loader
|
||||
|
||||
#### Checkpoint Clip Loader Resident
|
||||
|
||||
CLIP-only checkpoint loader.
|
||||
|
||||
Behavior:
|
||||
|
||||
- can reuse a live equivalent CLIP object
|
||||
- avoids rebuilding the diffusion model and VAE outputs when only CLIP is needed
|
||||
|
||||
#### Checkpoint VAE Loader Resident
|
||||
|
||||
VAE-only checkpoint loader.
|
||||
|
||||
Behavior:
|
||||
|
||||
- can reuse a live equivalent VAE object
|
||||
- avoids rebuilding the diffusion model and CLIP outputs when only VAE is needed
|
||||
|
||||
### Residency nodes
|
||||
|
||||
#### Set Global Residency Policy
|
||||
|
||||
Sets the active global policy and returns it as a `STRING`.
|
||||
|
||||
The loader nodes also expose an optional `policy_override` string input for one-off loads.
|
||||
|
||||
#### Registry Snapshot
|
||||
|
||||
Returns a composite formatted JSON snapshot with:
|
||||
|
||||
- `policy` for the active global policy
|
||||
- `entries` for native Comfy-managed registry entries
|
||||
- `external_entries` for compatible external cache entries discovered at runtime
|
||||
|
||||
#### Pin Model Residency / Pin CLIP Residency / Pin VAE Residency
|
||||
|
||||
Marks a tracked native object as sticky or non-sticky and optionally changes its priority.
|
||||
|
||||
#### Preload Model To GPU / Preload CLIP To GPU / Preload VAE To GPU
|
||||
|
||||
Calls `load_models_gpu(..., force_full_load=True)` for the selected native object, then updates sticky state / priority in the registry.
|
||||
|
||||
#### Evict Model From GPU / Evict CLIP From GPU / Evict VAE From GPU
|
||||
|
||||
Attempts to unload the selected native object from the current loaded-model set.
|
||||
|
||||
`unpatch_weights=True` performs a full unload path. When eviction succeeds, the node returns `evicted`; otherwise `not_loaded`.
|
||||
|
||||
#### Report Model Residency / Report CLIP Residency / Report VAE Residency
|
||||
|
||||
Returns a JSON report for a single tracked native object.
|
||||
|
||||
If the object is not currently bound in the registry, the node returns a JSON payload with `tracked: false`.
|
||||
|
||||
## Adaptive trimming before resident loads
|
||||
|
||||
The resident loaders now do load-scoped VRAM trimming themselves.
|
||||
|
||||
Before a GPU-bound resident load, the loader estimates required bytes from:
|
||||
|
||||
- the safetensors header when possible
|
||||
- the detected checkpoint component subset when possible
|
||||
- otherwise the source file size as a fallback
|
||||
|
||||
It then requests enough free VRAM for:
|
||||
|
||||
- the estimated load size
|
||||
- adaptive headroom
|
||||
|
||||
Current adaptive headroom policy:
|
||||
|
||||
- ratio: `12.5%` of the estimated load
|
||||
- floor: `256 MiB`
|
||||
- ceiling: `1 GiB`
|
||||
|
||||
The trim path prefers to:
|
||||
|
||||
- unload non-sticky entries first
|
||||
- then lower-priority sticky entries
|
||||
- preserve explicitly kept models
|
||||
- use partial unload where available
|
||||
- include compatible external cache entries in the same candidate search only when `COMFYUI_GPU_RESIDENT_TRIM_EXTERNAL=1`, they are visible, on the same device, and not currently claimed/in use
|
||||
|
||||
This logic lives in the resident loader path. You do not need a separate “target free VRAM” node for it.
|
||||
|
||||
## Registry and observability
|
||||
|
||||
The snapshot now exposes two collections:
|
||||
|
||||
- `entries` for native Comfy-managed tracked objects
|
||||
- `external_entries` for compatible external cache objects
|
||||
|
||||
### Native registry entries
|
||||
|
||||
The native registry tracks residency metadata for bound objects.
|
||||
|
||||
Typical per-entry fields include:
|
||||
|
||||
- `entry_id`
|
||||
- `kind`
|
||||
- `source_path`
|
||||
- `basename`
|
||||
- `sticky`
|
||||
- `priority`
|
||||
- `created_at`
|
||||
- `last_touched`
|
||||
- `loaded_bytes`
|
||||
- `total_bytes`
|
||||
- `load_device`
|
||||
- `offload_device`
|
||||
- `current_device`
|
||||
- `last_method`
|
||||
- `last_report`
|
||||
- `loader_key`
|
||||
- `notes`
|
||||
- `alive`
|
||||
|
||||
The `last_method` / `last_report` fields let you see whether a load actually used:
|
||||
|
||||
- direct safetensors GPU ingest
|
||||
- safetensors CPU -> CUDA fallback
|
||||
- safetensors component-only load
|
||||
- CPU-first `torch.load()` compatibility path
|
||||
- a recorded load failure
|
||||
|
||||
### External registry entries
|
||||
|
||||
Typical external-entry fields include:
|
||||
|
||||
- `entry_id`
|
||||
- `cache_key`
|
||||
- `kind`
|
||||
- `source_path`
|
||||
- `basename`
|
||||
- `sticky`
|
||||
- `priority`
|
||||
- `created_at`
|
||||
- `last_touched`
|
||||
- `loaded_bytes`
|
||||
- `total_bytes`
|
||||
- `load_device`
|
||||
- `offload_device`
|
||||
- `current_device`
|
||||
- `claimed`
|
||||
- `notes`
|
||||
- `alive`
|
||||
- `external`
|
||||
|
||||
For SeedVR2-backed entries, `claimed: true` means the cache object is currently marked in use and is skipped by the external trim candidate search when that opt-in path is enabled.
|
||||
|
||||
## What gets tracked
|
||||
|
||||
### Native tracked/bound paths
|
||||
|
||||
Native tracked paths include:
|
||||
|
||||
- resident node loads from this repo
|
||||
- stock checkpoint loads
|
||||
- stock diffusion-model loads
|
||||
- stock CLIP loads
|
||||
- stock CLIP Vision loads
|
||||
- stock diffusers loads
|
||||
|
||||
ControlNet loads also participate in the patched load context and device-policy path, but this repo does not currently expose dedicated ControlNet residency nodes.
|
||||
|
||||
### External tracked paths
|
||||
|
||||
Current external integration coverage is:
|
||||
|
||||
- SeedVR2 global cached **DiT** entries
|
||||
- SeedVR2 global cached **VAE** entries
|
||||
|
||||
Those entries are discovered lazily from compatible SeedVR2 cache modules at runtime. They are tracked separately from the native registry and participate in snapshot plus provider-specific eviction decisions, with load-scoped trim available only through the explicit external-trim opt-in.
|
||||
|
||||
## Important limits and non-goals
|
||||
|
||||
### Best path is still `.safetensors`
|
||||
|
||||
The narrow fast path is built around `.safetensors`.
|
||||
|
||||
That is where this repo can:
|
||||
|
||||
- inspect headers cheaply
|
||||
- select only model / clip / vae subsets
|
||||
- estimate component bytes more accurately
|
||||
- attempt direct device-targeted reads
|
||||
|
||||
### `.ckpt` / `.pt` / pickle formats are still CPU-first
|
||||
|
||||
For pickle-based formats, PyTorch still goes through `torch.load()` on CPU first.
|
||||
|
||||
The repo can still:
|
||||
|
||||
- track those loads
|
||||
- keep the resulting live objects resident
|
||||
- reuse equivalent live objects later
|
||||
|
||||
It does **not** claim direct-to-GPU ingest for those formats.
|
||||
|
||||
### Cross-process persistence is out of scope
|
||||
|
||||
This repo does **not** keep VRAM contents alive after ComfyUI or WSL exits. CUDA memory lifetime is process/context scoped. Achieving persistence across process shutdown requires a long-lived keeper process or server that owns the CUDA context.
|
||||
This repo does **not** keep VRAM allocations alive after ComfyUI, Python, or WSL exits.
|
||||
|
||||
CUDA memory lifetime is process/context scoped. True persistence across process shutdown would need a separate long-lived keeper process or service that owns the CUDA context.
|
||||
|
||||
### External integrations are compatibility-based, not universal
|
||||
|
||||
The external registry does **not** automatically manage every third-party cache.
|
||||
|
||||
At the moment, the documented external integration target is **SeedVR2**. Other custom nodes with private caches remain invisible until this repo grows a provider-specific integration for them.
|
||||
|
||||
### It does not automatically capture arbitrary custom loader implementations
|
||||
|
||||
The native registry only sees objects that pass through the patched ComfyUI load paths or through this repo’s resident nodes.
|
||||
|
||||
If another custom node loads models through its own private code path and bypasses those patched entry points, that object may never become a tracked native registry entry. In that case, the preload / pin / evict / report nodes from this repo cannot manage it until that external loader is integrated or patched.
|
||||
|
||||
Likewise, even for supported external providers such as SeedVR2, the current external integration is about **observation + provider-specific eviction**, with automatic trim kept opt-in. This repo does **not** yet expose dedicated external preload / pin / report / evict nodes for provider-owned cache entries.
|
||||
|
||||
## Installation
|
||||
|
||||
@@ -113,67 +487,78 @@ Clone into `custom_nodes`:
|
||||
git clone https://github.com/xmarre/ComfyUI-GPU-Resident-Loader ComfyUI/custom_nodes/ComfyUI-GPU-Resident-Loader
|
||||
```
|
||||
|
||||
Install dependencies inside the same Python environment ComfyUI uses:
|
||||
Install dependencies into the same Python environment ComfyUI uses:
|
||||
|
||||
```bash
|
||||
pip install -r ComfyUI/custom_nodes/ComfyUI-GPU-Resident-Loader/requirements.txt
|
||||
```
|
||||
|
||||
Optional SageAttention dependencies are **not** installed by default. Install those separately if you plan to use the SageAttention loader modes.
|
||||
Requirements declared by the repo:
|
||||
|
||||
## Basic usage
|
||||
- Python `>=3.10`
|
||||
- `safetensors>=0.4.3`
|
||||
|
||||
### For direct diffusion-model loading
|
||||
Optional SageAttention dependencies are **not** installed by default. Install those separately if you plan to use a SageAttention mode in the resident loaders.
|
||||
|
||||
Use **Diffusion Model Loader Resident**.
|
||||
## Basic usage patterns
|
||||
|
||||
Recommended on a large VRAM machine:
|
||||
### 1) Large-VRAM, mostly resident workflow
|
||||
|
||||
- policy: `sticky_gpu`
|
||||
- model format: `.safetensors`
|
||||
- preload with **Preload Model To GPU**
|
||||
- inspect with **Report Model Residency** or **Registry Snapshot**
|
||||
Recommended baseline:
|
||||
|
||||
### For full checkpoints
|
||||
- start ComfyUI with `--highvram` or set policy manually to `sticky_gpu`
|
||||
- prefer `.safetensors` for hot models
|
||||
- load diffusion models through **Diffusion Model Loader Resident**
|
||||
- use **Preload ... To GPU** for models you know you will reuse
|
||||
- inspect with **Report ... Residency** or **Registry Snapshot**
|
||||
|
||||
Use **Checkpoint Loader Resident**.
|
||||
### 2) Full checkpoint workflow
|
||||
|
||||
That tracks and binds the resulting diffusion model, CLIP, and VAE independently so they appear in the registry snapshot. If an equivalent live model, CLIP, or VAE already exists, the loader reuses it instead of rebuilding it.
|
||||
Use **Checkpoint Loader Resident** when you want `MODEL + CLIP + VAE` together.
|
||||
|
||||
### For staged checkpoint loads
|
||||
That path can reuse already-live components instead of always rebuilding all three outputs.
|
||||
|
||||
### 3) Staged checkpoint workflow
|
||||
|
||||
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
|
||||
|
||||
### 4) Manual native residency control
|
||||
|
||||
Use:
|
||||
|
||||
- **Checkpoint Model Loader Resident** when the workflow only needs the diffusion model
|
||||
- **Checkpoint Clip Loader Resident** when the workflow only needs the text encoder
|
||||
- **Checkpoint VAE Loader Resident** when the workflow only needs the VAE
|
||||
- **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
|
||||
|
||||
The model-only checkpoint node takes the same selective safetensors UNet fast path as the diffusion-model loader. CLIP-only and VAE-only nodes still use ComfyUI's checkpoint construction logic, but they avoid materializing the other outputs and can reuse an already-live equivalent object.
|
||||
### 5) Mixed workflows with SeedVR2 external caching
|
||||
|
||||
### For manual residency control
|
||||
If SeedVR2 keeps DiT or VAE models in its own global cache, those objects can now show up in **Registry Snapshot** under `external_entries`.
|
||||
|
||||
- use **Pin ... Residency** to mark a tracked object sticky or evictable
|
||||
- use **Preload ... To GPU** to fully materialize it in VRAM immediately
|
||||
- use **Evict ... From GPU** to unload it from the current loaded-model set
|
||||
That means:
|
||||
|
||||
## Observability
|
||||
- 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
|
||||
|
||||
Every tracked load stores:
|
||||
## Notes on compatibility and migration
|
||||
|
||||
- source path
|
||||
- last load method
|
||||
- requested device
|
||||
- actual device
|
||||
- sticky flag
|
||||
- current loaded bytes
|
||||
- total bytes
|
||||
- current/offload/load device
|
||||
### Legacy wiring: `extra_state_dict` used as a policy string
|
||||
|
||||
That data is surfaced through the report nodes and the registry snapshot node.
|
||||
The resident diffusion-model loader contains a compatibility shim for older graphs:
|
||||
|
||||
## Conversion helper
|
||||
- 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
|
||||
|
||||
`scripts/convert_checkpoint_to_safetensors.py` is included for one-time conversion of hot `.ckpt` / `.pt` files into `.safetensors`.
|
||||
New graphs should connect policy strings to **`policy_override`**, not to `extra_state_dict`.
|
||||
|
||||
### Convert hot pickle checkpoints to safetensors
|
||||
|
||||
`scripts/convert_checkpoint_to_safetensors.py` is included for one-time conversion of hot `.ckpt` / `.pt` / `.pth` style checkpoints.
|
||||
|
||||
Example:
|
||||
|
||||
@@ -183,8 +568,13 @@ python ComfyUI/custom_nodes/ComfyUI-GPU-Resident-Loader/scripts/convert_checkpoi
|
||||
--output /path/to/model.safetensors
|
||||
```
|
||||
|
||||
Optional flags:
|
||||
|
||||
- `--state-dict-key <key>` to extract a different top-level dict key
|
||||
- `--allow-non-tensor-values` to skip non-tensor entries instead of failing
|
||||
|
||||
## License
|
||||
|
||||
GPL-3.0-or-later.
|
||||
|
||||
This repo intentionally stays GPL-compatible because it adapts behavior from GPL-licensed ComfyUI and mirrors feature behavior from the GPL-3.0-licensed KJNodes diffusion loader.
|
||||
This repo stays GPL-compatible because it adapts behavior from GPL-licensed ComfyUI and mirrors relevant loader behavior from GPL-3.0-licensed KJNodes.
|
||||
|
||||
+260
-28
@@ -1,17 +1,33 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
from typing import Any
|
||||
|
||||
import comfy.model_management as model_management
|
||||
import torch
|
||||
|
||||
from .external_residency import EXTERNAL_REGISTRY, ensure_external_integrations_installed, external_trim_enabled
|
||||
from .residency import REGISTRY
|
||||
|
||||
_ADAPTIVE_HEADROOM_RATIO = 0.125
|
||||
_ADAPTIVE_HEADROOM_FLOOR_BYTES = 256 * 1024 * 1024
|
||||
_ADAPTIVE_HEADROOM_CEIL_BYTES = 1024 * 1024 * 1024
|
||||
|
||||
|
||||
def _safe_free_memory(device) -> int:
|
||||
return int(model_management.get_free_memory(device))
|
||||
|
||||
|
||||
def _normalize_trim_device(device: str | torch.device | None):
|
||||
if device is None or isinstance(device, torch.device):
|
||||
return device
|
||||
try:
|
||||
return torch.device(device)
|
||||
except Exception:
|
||||
return device
|
||||
|
||||
|
||||
def _safe_is_dead(loaded) -> bool:
|
||||
try:
|
||||
return loaded.is_dead()
|
||||
@@ -19,6 +35,69 @@ def _safe_is_dead(loaded) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
def _device_matches(device_a, device_b) -> bool:
|
||||
normalized_a = _normalize_trim_device(device_a)
|
||||
normalized_b = _normalize_trim_device(device_b)
|
||||
if normalized_a is None or normalized_b is None:
|
||||
return False
|
||||
if isinstance(normalized_a, torch.device) and isinstance(normalized_b, torch.device):
|
||||
if normalized_a.type != normalized_b.type:
|
||||
return False
|
||||
if normalized_a.type == "cuda":
|
||||
index_a = 0 if normalized_a.index is None else normalized_a.index
|
||||
index_b = 0 if normalized_b.index is None else normalized_b.index
|
||||
return index_a == index_b
|
||||
return True
|
||||
return str(normalized_a) == str(normalized_b)
|
||||
|
||||
|
||||
def _should_force_cpu_offload(model: Any, *, active_device=None, force: bool = False) -> bool:
|
||||
if force:
|
||||
return True
|
||||
active_device = _normalize_trim_device(active_device)
|
||||
if not isinstance(active_device, torch.device) or active_device.type != "cuda":
|
||||
return False
|
||||
return _device_matches(active_device, getattr(model, "offload_device", None))
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _temporary_offload_device(model: Any, device):
|
||||
if model is None or device is None or not hasattr(model, "offload_device"):
|
||||
yield
|
||||
return
|
||||
|
||||
original_device = model.offload_device
|
||||
model.offload_device = device
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
model.offload_device = original_device
|
||||
|
||||
|
||||
def unload_loaded_model(
|
||||
loaded,
|
||||
*,
|
||||
active_device: str | torch.device | None = None,
|
||||
force_offload_to_cpu: bool = False,
|
||||
unpatch_weights: bool = True,
|
||||
) -> bool:
|
||||
# This helper is intentionally scoped to full unloads owned by this plugin.
|
||||
model = getattr(loaded, "model", None)
|
||||
if model is None:
|
||||
return False
|
||||
|
||||
active_device = _normalize_trim_device(active_device if active_device is not None else getattr(loaded, "device", None))
|
||||
force_cpu_offload = _should_force_cpu_offload(
|
||||
model,
|
||||
active_device=active_device,
|
||||
force=force_offload_to_cpu,
|
||||
)
|
||||
unload_target = torch.device("cpu") if force_cpu_offload else None
|
||||
|
||||
with _temporary_offload_device(model, unload_target):
|
||||
return loaded.model_unload(None, unpatch_weights=unpatch_weights)
|
||||
|
||||
|
||||
def _sort_key_for_candidate(entry, *, sticky_respected: bool) -> tuple[int, int, float]:
|
||||
if not sticky_respected:
|
||||
return (0, 0, getattr(entry, "last_touched", 0.0) if entry is not None else 0.0)
|
||||
@@ -27,8 +106,53 @@ def _sort_key_for_candidate(entry, *, sticky_respected: bool) -> tuple[int, int,
|
||||
return (1, priority, last_touched)
|
||||
|
||||
|
||||
def _trim_candidates(*, device, respect_sticky: bool, sticky_floor_priority: int) -> list[tuple[Any, Any, bool]]:
|
||||
candidates: list[tuple[Any, Any, bool]] = []
|
||||
def _should_keep_loaded_model(model: Any, keep_models: tuple[Any, ...]) -> bool:
|
||||
if not keep_models:
|
||||
return False
|
||||
for keep in keep_models:
|
||||
if keep is None:
|
||||
continue
|
||||
if model is keep:
|
||||
return True
|
||||
is_clone = getattr(model, "is_clone", None)
|
||||
if callable(is_clone):
|
||||
try:
|
||||
if is_clone(keep):
|
||||
return True
|
||||
except Exception:
|
||||
pass
|
||||
return False
|
||||
|
||||
|
||||
def _trim_candidates(
|
||||
*,
|
||||
device,
|
||||
respect_sticky: bool,
|
||||
sticky_floor_priority: int,
|
||||
keep_models: tuple[Any, ...],
|
||||
include_external: bool | None = None,
|
||||
) -> list[tuple[Any, Any, bool, bool]]:
|
||||
"""
|
||||
Collects and returns eviction/unload candidates from in-memory models and, optionally, external integrations.
|
||||
|
||||
Filters currently loaded models by the given device, excludes dead or missing models, and skips models listed in `keep_models`. If `respect_sticky` is true, entries whose registry metadata mark them as sticky and whose priority meets or exceeds `sticky_floor_priority` are flagged so they are treated as higher-priority to keep. When `include_external` is true (or when `include_external` is None and external trimming is enabled), candidates from the external registry are included.
|
||||
|
||||
Parameters:
|
||||
device: Device filter for candidates; if not None only candidates matching this device are considered.
|
||||
respect_sticky (bool): Whether to respect sticky registry entries when computing candidate priority.
|
||||
sticky_floor_priority (int): Minimum priority value for a registry entry to be considered sticky.
|
||||
keep_models (tuple[Any, ...]): Objects that must not be selected as candidates.
|
||||
include_external (bool | None): If True include external-registry candidates; if False exclude them; if None defer to runtime external_trim_enabled().
|
||||
|
||||
Returns:
|
||||
list[tuple[Any, Any, bool, bool]]: A sorted list of tuples (candidate_obj, registry_entry_or_None, sticky_respected, is_external_candidate).
|
||||
- candidate_obj: The loaded model object (internal) or the external object.
|
||||
- registry_entry_or_None: Registry metadata for the candidate, or None if unavailable.
|
||||
- sticky_respected: `True` when the candidate is marked sticky and meets `sticky_floor_priority`.
|
||||
- is_external_candidate: `True` for candidates originating from the external registry.
|
||||
"""
|
||||
candidates: list[tuple[Any, Any, bool, bool]] = []
|
||||
ensure_external_integrations_installed()
|
||||
for loaded in list(model_management.current_loaded_models):
|
||||
if device is not None and loaded.device != device:
|
||||
continue
|
||||
@@ -38,6 +162,8 @@ def _trim_candidates(*, device, respect_sticky: bool, sticky_floor_priority: int
|
||||
model = getattr(loaded, "model", None)
|
||||
if model is None:
|
||||
continue
|
||||
if _should_keep_loaded_model(model, keep_models):
|
||||
continue
|
||||
|
||||
entry = REGISTRY.entry_for_object(model)
|
||||
sticky_respected = (
|
||||
@@ -46,7 +172,17 @@ def _trim_candidates(*, device, respect_sticky: bool, sticky_floor_priority: int
|
||||
and bool(getattr(entry, "sticky", False))
|
||||
and int(getattr(entry, "priority", 0)) >= int(sticky_floor_priority)
|
||||
)
|
||||
candidates.append((loaded, entry, sticky_respected))
|
||||
candidates.append((loaded, entry, sticky_respected, False))
|
||||
|
||||
should_include_external = external_trim_enabled() if include_external is None else bool(include_external)
|
||||
if should_include_external:
|
||||
for external_obj, entry, sticky_respected in EXTERNAL_REGISTRY.candidates(
|
||||
device=device,
|
||||
respect_sticky=respect_sticky,
|
||||
sticky_floor_priority=sticky_floor_priority,
|
||||
keep_models=keep_models,
|
||||
):
|
||||
candidates.append((external_obj, entry, sticky_respected, True))
|
||||
|
||||
candidates.sort(key=lambda item: _sort_key_for_candidate(item[1], sticky_respected=item[2]))
|
||||
return candidates
|
||||
@@ -54,17 +190,54 @@ def _trim_candidates(*, device, respect_sticky: bool, sticky_floor_priority: int
|
||||
|
||||
def trim_resident_vram(
|
||||
*,
|
||||
device: str | torch.device | None = None,
|
||||
target_free_vram_bytes: int,
|
||||
respect_sticky: bool,
|
||||
sticky_floor_priority: int,
|
||||
allow_partial_unload: bool,
|
||||
keep_models: tuple[Any, ...] = (),
|
||||
include_external: bool | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Trim resident models (and optional external integration objects) until the requested amount of free VRAM is available or a stopping condition occurs.
|
||||
|
||||
Tries to free GPU memory on `device` by evicting or unloading loaded models (and, optionally, external candidates). Respects sticky/priority hints, can perform partial unloads of pinned RAM when supported, and records each attempted action in the returned report.
|
||||
|
||||
Parameters:
|
||||
device (str | torch.device | None): Target device to free (defaults to model_management.get_torch_device()).
|
||||
target_free_vram_bytes (int): Desired amount of free VRAM, in bytes.
|
||||
respect_sticky (bool): If true, prefer protecting entries marked as sticky with sufficient priority.
|
||||
sticky_floor_priority (int): Minimum priority required for a sticky entry to be respected.
|
||||
allow_partial_unload (bool): If true, allow partial unloads and attempts to free pinned host RAM before full eviction.
|
||||
keep_models (tuple[Any, ...]): Sequence of model objects that must not be unloaded (exact matches or recognized clones).
|
||||
include_external (bool | None): If None, use external_trim_enabled() at runtime; otherwise force inclusion/exclusion of external candidates.
|
||||
|
||||
Returns:
|
||||
dict[str, Any]: Report of the trimming operation containing:
|
||||
- status: "met_target", "partial", or "error".
|
||||
- stopped_reason: reason the loop stopped (e.g., "target_met", "no_candidates", "no_progress", "error").
|
||||
- target_met (bool): whether the target free VRAM was reached.
|
||||
- device (str): string form of the device used.
|
||||
- target_free_vram_bytes (int), free_before_bytes (int), free_after_bytes (int).
|
||||
- freed_vram_bytes (int): total freed VRAM during this call.
|
||||
- respect_sticky (bool), sticky_floor_priority (int), allow_partial_unload (bool).
|
||||
- external_trim_enabled (bool): computed flag indicating whether external candidates were considered.
|
||||
- actions (list[dict]): ordered per-candidate action records; each entry includes metadata such as
|
||||
entry_id, basename, tracked, external_candidate, sticky_respected, priority,
|
||||
need_before_bytes, loaded_before_bytes, freed_pinned_ram_bytes,
|
||||
mode (e.g., "full_unload", "partial_unload", "external_evict", "error"),
|
||||
loaded_after_bytes, freed_vram_bytes, free_after_bytes, and any warnings/errors.
|
||||
"""
|
||||
ensure_external_integrations_installed()
|
||||
cleanup_models_gc = getattr(model_management, "cleanup_models_gc", None)
|
||||
if callable(cleanup_models_gc):
|
||||
cleanup_models_gc()
|
||||
|
||||
REGISTRY.refresh_runtime_state()
|
||||
device = model_management.get_torch_device()
|
||||
EXTERNAL_REGISTRY.refresh_runtime_state()
|
||||
device = _normalize_trim_device(device)
|
||||
if device is None:
|
||||
device = model_management.get_torch_device()
|
||||
free_before = _safe_free_memory(device)
|
||||
actions: list[dict[str, Any]] = []
|
||||
soft_empty_cache = getattr(model_management, "soft_empty_cache", None)
|
||||
@@ -81,18 +254,21 @@ def trim_resident_vram(
|
||||
device=device,
|
||||
respect_sticky=respect_sticky,
|
||||
sticky_floor_priority=sticky_floor_priority,
|
||||
keep_models=keep_models,
|
||||
include_external=include_external,
|
||||
)
|
||||
if not candidates:
|
||||
stopped_reason = "no_candidates"
|
||||
break
|
||||
|
||||
loaded, entry, sticky_respected = candidates[0]
|
||||
model = loaded.model
|
||||
loaded_before = int(loaded.model_loaded_memory())
|
||||
candidate, entry, sticky_respected, is_external_candidate = candidates[0]
|
||||
model = candidate if is_external_candidate else candidate.model
|
||||
loaded_before = int(getattr(entry, "loaded_bytes", 0)) if is_external_candidate else int(candidate.model_loaded_memory())
|
||||
action = {
|
||||
"entry_id": getattr(entry, "entry_id", None),
|
||||
"basename": None if entry is None else os.path.basename(getattr(entry, "source_path", "") or ""),
|
||||
"tracked": entry is not None,
|
||||
"external_candidate": bool(is_external_candidate),
|
||||
"sticky_respected": sticky_respected,
|
||||
"priority": None if entry is None else int(getattr(entry, "priority", 0)),
|
||||
"need_before_bytes": need,
|
||||
@@ -100,7 +276,7 @@ def trim_resident_vram(
|
||||
"freed_pinned_ram_bytes": 0,
|
||||
}
|
||||
|
||||
if allow_partial_unload and hasattr(model, "pinned_memory_size") and hasattr(model, "partially_unload_ram"):
|
||||
if (not is_external_candidate and allow_partial_unload and hasattr(model, "pinned_memory_size") and hasattr(model, "partially_unload_ram")):
|
||||
try:
|
||||
pinned_memory = int(model.pinned_memory_size())
|
||||
if pinned_memory > 0:
|
||||
@@ -110,31 +286,42 @@ def trim_resident_vram(
|
||||
except Exception as exc:
|
||||
action["pinned_ram_warning"] = str(exc)
|
||||
|
||||
try:
|
||||
fully_unloaded = loaded.model_unload(need if allow_partial_unload else None)
|
||||
action["mode"] = "full_unload" if fully_unloaded else "partial_unload"
|
||||
except Exception as exc:
|
||||
if allow_partial_unload:
|
||||
try:
|
||||
fully_unloaded = loaded.model_unload(None)
|
||||
action["mode"] = "full_unload_fallback"
|
||||
action["partial_unload_warning"] = str(exc)
|
||||
except Exception as fallback_exc:
|
||||
action["mode"] = "error"
|
||||
action["error"] = str(fallback_exc)
|
||||
actions.append(action)
|
||||
stopped_reason = "error"
|
||||
break
|
||||
else:
|
||||
if is_external_candidate:
|
||||
try:
|
||||
fully_unloaded = EXTERNAL_REGISTRY.evict(entry)
|
||||
action["mode"] = "external_evict"
|
||||
except Exception as exc:
|
||||
action["mode"] = "error"
|
||||
action["error"] = str(exc)
|
||||
actions.append(action)
|
||||
stopped_reason = "error"
|
||||
break
|
||||
|
||||
if fully_unloaded:
|
||||
else:
|
||||
try:
|
||||
model_management.current_loaded_models.remove(loaded)
|
||||
fully_unloaded = candidate.model_unload(need if allow_partial_unload else None)
|
||||
action["mode"] = "full_unload" if fully_unloaded else "partial_unload"
|
||||
except Exception as exc:
|
||||
if allow_partial_unload:
|
||||
try:
|
||||
fully_unloaded = candidate.model_unload(None)
|
||||
action["mode"] = "full_unload_fallback"
|
||||
action["partial_unload_warning"] = str(exc)
|
||||
except Exception as fallback_exc:
|
||||
action["mode"] = "error"
|
||||
action["error"] = str(fallback_exc)
|
||||
actions.append(action)
|
||||
stopped_reason = "error"
|
||||
break
|
||||
else:
|
||||
action["mode"] = "error"
|
||||
action["error"] = str(exc)
|
||||
actions.append(action)
|
||||
stopped_reason = "error"
|
||||
break
|
||||
|
||||
if fully_unloaded and not is_external_candidate:
|
||||
try:
|
||||
model_management.current_loaded_models.remove(candidate)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
@@ -142,8 +329,12 @@ def trim_resident_vram(
|
||||
soft_empty_cache()
|
||||
|
||||
REGISTRY.refresh_runtime_state()
|
||||
EXTERNAL_REGISTRY.refresh_runtime_state()
|
||||
free_after = _safe_free_memory(device)
|
||||
action["loaded_after_bytes"] = 0 if fully_unloaded else int(loaded.model_loaded_memory())
|
||||
if is_external_candidate:
|
||||
action["loaded_after_bytes"] = 0 if fully_unloaded else int(getattr(entry, "loaded_bytes", 0))
|
||||
else:
|
||||
action["loaded_after_bytes"] = 0 if fully_unloaded else int(candidate.model_loaded_memory())
|
||||
action["freed_vram_bytes"] = max(0, free_after - free_now)
|
||||
action["free_after_bytes"] = free_after
|
||||
actions.append(action)
|
||||
@@ -168,5 +359,46 @@ 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,
|
||||
}
|
||||
|
||||
|
||||
def adaptive_headroom_bytes(required_bytes: int) -> int:
|
||||
required = max(0, int(required_bytes))
|
||||
if required == 0:
|
||||
return 0
|
||||
return min(
|
||||
_ADAPTIVE_HEADROOM_CEIL_BYTES,
|
||||
max(_ADAPTIVE_HEADROOM_FLOOR_BYTES, int(required * _ADAPTIVE_HEADROOM_RATIO)),
|
||||
)
|
||||
|
||||
|
||||
def trim_resident_vram_for_load(
|
||||
*,
|
||||
required_bytes: int,
|
||||
reason: str,
|
||||
device: str | torch.device | None = None,
|
||||
respect_sticky: bool = True,
|
||||
sticky_floor_priority: int = 0,
|
||||
allow_partial_unload: bool = True,
|
||||
keep_models: tuple[Any, ...] = (),
|
||||
) -> dict[str, Any]:
|
||||
estimated_load_bytes = max(0, int(required_bytes))
|
||||
headroom_bytes = adaptive_headroom_bytes(estimated_load_bytes)
|
||||
target_free_vram_bytes = estimated_load_bytes + headroom_bytes
|
||||
|
||||
report = trim_resident_vram(
|
||||
device=device,
|
||||
target_free_vram_bytes=target_free_vram_bytes,
|
||||
respect_sticky=respect_sticky,
|
||||
sticky_floor_priority=sticky_floor_priority,
|
||||
allow_partial_unload=allow_partial_unload,
|
||||
keep_models=keep_models,
|
||||
)
|
||||
report["trim_strategy"] = "adaptive_load_request"
|
||||
report["trim_reason"] = str(reason)
|
||||
report["estimated_load_bytes"] = estimated_load_bytes
|
||||
report["adaptive_headroom_bytes"] = headroom_bytes
|
||||
report["kept_loaded_models"] = len([model for model in keep_models if model is not None])
|
||||
return report
|
||||
|
||||
@@ -0,0 +1,726 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
import contextlib
|
||||
import dataclasses
|
||||
import inspect
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
import weakref
|
||||
from typing import Any, Callable
|
||||
|
||||
import torch
|
||||
|
||||
from .residency import KIND_MODEL, KIND_VAE, REGISTRY
|
||||
|
||||
_LOG = logging.getLogger(__name__)
|
||||
_SEEDVR2_PATCHED_CLASS_IDS: set[int] = set()
|
||||
_SEEDVR2_PATCHING_CLASS_IDS: set[int] = set()
|
||||
_SEEDVR2_PATCHING_THREAD_IDS: dict[int, int] = {}
|
||||
_SEEDVR2_PATCH_LOCK = threading.Lock()
|
||||
_SEEDVR2_PATCH_CONDITION = threading.Condition(_SEEDVR2_PATCH_LOCK)
|
||||
|
||||
|
||||
def _now() -> float:
|
||||
return time.time()
|
||||
|
||||
|
||||
def _normalize_device(device: str | torch.device | None):
|
||||
if device is None or isinstance(device, torch.device):
|
||||
return device
|
||||
try:
|
||||
return torch.device(device)
|
||||
except Exception:
|
||||
return device
|
||||
|
||||
|
||||
def _device_matches(device_a, device_b) -> bool:
|
||||
normalized_a = _normalize_device(device_a)
|
||||
normalized_b = _normalize_device(device_b)
|
||||
if normalized_a is None or normalized_b is None:
|
||||
return False
|
||||
if isinstance(normalized_a, torch.device) and isinstance(normalized_b, torch.device):
|
||||
if normalized_a.type != normalized_b.type:
|
||||
return False
|
||||
if normalized_a.type == "cuda":
|
||||
index_a = 0 if normalized_a.index is None else normalized_a.index
|
||||
index_b = 0 if normalized_b.index is None else normalized_b.index
|
||||
return index_a == index_b
|
||||
return True
|
||||
return str(normalized_a) == str(normalized_b)
|
||||
|
||||
|
||||
def _iter_seedvr2_wrapper_chain(model: Any):
|
||||
if model is None:
|
||||
return
|
||||
|
||||
stack = [model]
|
||||
seen: set[int] = set()
|
||||
while stack:
|
||||
current = stack.pop()
|
||||
if current is None:
|
||||
continue
|
||||
current_id = id(current)
|
||||
if current_id in seen:
|
||||
continue
|
||||
seen.add(current_id)
|
||||
yield current
|
||||
for attr in ("_orig_mod", "dit_model"):
|
||||
child = getattr(current, attr, None)
|
||||
if child is not None:
|
||||
stack.append(child)
|
||||
|
||||
|
||||
def _seedvr2_is_claimed(model: Any) -> bool:
|
||||
return any(bool(getattr(current, "_seedvr2_cache_claimed", False)) for current in _iter_seedvr2_wrapper_chain(model))
|
||||
|
||||
|
||||
def _first_tensor_device(model: Any) -> str | None:
|
||||
if model is None:
|
||||
return None
|
||||
try:
|
||||
for param in model.parameters():
|
||||
return str(param.device)
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
for buffer in model.buffers():
|
||||
return str(buffer.device)
|
||||
except Exception:
|
||||
pass
|
||||
return None
|
||||
|
||||
|
||||
def _unique_tensor_nbytes(model: Any) -> int:
|
||||
if model is None:
|
||||
return 0
|
||||
|
||||
total = 0
|
||||
seen_storages: set[tuple[int, int, int]] = set()
|
||||
|
||||
def visit_tensor(tensor: torch.Tensor) -> None:
|
||||
nonlocal total
|
||||
if tensor is None:
|
||||
return
|
||||
try:
|
||||
storage = tensor.untyped_storage()
|
||||
key = (storage.data_ptr(), storage.nbytes(), int(tensor.device.index or 0) if tensor.device.type == "cuda" else -1)
|
||||
except Exception:
|
||||
key = (id(tensor), tensor.numel() * tensor.element_size(), -2)
|
||||
if key in seen_storages:
|
||||
return
|
||||
seen_storages.add(key)
|
||||
total += int(tensor.numel()) * int(tensor.element_size())
|
||||
|
||||
try:
|
||||
for tensor in model.parameters():
|
||||
visit_tensor(tensor)
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
for tensor in model.buffers():
|
||||
visit_tensor(tensor)
|
||||
except Exception:
|
||||
pass
|
||||
return int(total)
|
||||
|
||||
|
||||
def _seedvr2_entry_key(kind: str, node_id: Any) -> str:
|
||||
return f"seedvr2:{kind}:{node_id}"
|
||||
|
||||
|
||||
def _seedvr2_source_path(kind: str, node_id: Any, config: dict[str, Any], model: Any) -> str:
|
||||
model_name = config.get("model") or getattr(model, "_model_name", None) or f"node_{node_id}"
|
||||
return f"seedvr2/{kind}/{node_id}/{model_name}"
|
||||
|
||||
|
||||
def _seedvr2_state_provider(
|
||||
model_ref: Callable[[], Any | None],
|
||||
config: dict[str, Any],
|
||||
) -> Callable[[], dict[str, Any]]:
|
||||
def provider() -> dict[str, Any]:
|
||||
model = model_ref()
|
||||
if model is None:
|
||||
return {}
|
||||
total_bytes = _unique_tensor_nbytes(model)
|
||||
return {
|
||||
"current_device": _first_tensor_device(model),
|
||||
"load_device": config.get("device"),
|
||||
"offload_device": config.get("offload_device"),
|
||||
"claimed": _seedvr2_is_claimed(model),
|
||||
"loaded_bytes": total_bytes,
|
||||
"total_bytes": total_bytes,
|
||||
}
|
||||
|
||||
return provider
|
||||
|
||||
|
||||
def _coerce_external_bytes(value: Any, fallback: int, *, cache_key: str, field_name: str) -> int:
|
||||
try:
|
||||
return int(value)
|
||||
except (TypeError, ValueError) as exc:
|
||||
_LOG.debug(
|
||||
"GPU Resident Loader: ignoring invalid external %s for %s: %r (%s)",
|
||||
field_name,
|
||||
cache_key,
|
||||
value,
|
||||
exc,
|
||||
)
|
||||
return int(fallback)
|
||||
|
||||
|
||||
@dataclasses.dataclass(slots=True)
|
||||
class ExternalResidencyEntry:
|
||||
entry_id: str
|
||||
cache_key: str
|
||||
kind: str
|
||||
source_path: str
|
||||
sticky: bool
|
||||
priority: int
|
||||
created_at: float = dataclasses.field(default_factory=_now)
|
||||
last_touched: float = dataclasses.field(default_factory=_now)
|
||||
loaded_bytes: int = 0
|
||||
total_bytes: int = 0
|
||||
load_device: str | None = None
|
||||
offload_device: str | None = None
|
||||
current_device: str | None = None
|
||||
claimed: bool = False
|
||||
notes: list[str] = dataclasses.field(default_factory=list)
|
||||
object_ref: weakref.ReferenceType[Any] | None = None
|
||||
state_provider: Callable[[], dict[str, Any]] | None = None
|
||||
evict_callback: Callable[[], bool] | None = None
|
||||
|
||||
def object(self) -> Any | None:
|
||||
return None if self.object_ref is None else self.object_ref()
|
||||
|
||||
def is_alive(self) -> bool:
|
||||
return self.object() is not None
|
||||
|
||||
def as_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"entry_id": self.entry_id,
|
||||
"cache_key": self.cache_key,
|
||||
"kind": self.kind,
|
||||
"source_path": self.source_path,
|
||||
"basename": os.path.basename(self.source_path) if self.source_path else None,
|
||||
"sticky": self.sticky,
|
||||
"priority": self.priority,
|
||||
"created_at": self.created_at,
|
||||
"last_touched": self.last_touched,
|
||||
"loaded_bytes": self.loaded_bytes,
|
||||
"total_bytes": self.total_bytes,
|
||||
"load_device": self.load_device,
|
||||
"offload_device": self.offload_device,
|
||||
"current_device": self.current_device,
|
||||
"claimed": self.claimed,
|
||||
"notes": list(self.notes),
|
||||
"alive": self.is_alive(),
|
||||
"external": True,
|
||||
}
|
||||
|
||||
|
||||
class ExternalResidencyRegistry:
|
||||
def __init__(self) -> None:
|
||||
self._lock = threading.RLock()
|
||||
self._entries: dict[str, ExternalResidencyEntry] = {}
|
||||
self._cache_key_to_entry: dict[str, str] = {}
|
||||
self._next_entry_seq = 1
|
||||
|
||||
def _make_entry_id(self, kind: str, source_path: str) -> str:
|
||||
basename = os.path.basename(source_path) or "anonymous"
|
||||
entry_id = f"external:{kind}:{basename}:{self._next_entry_seq}"
|
||||
self._next_entry_seq += 1
|
||||
return entry_id
|
||||
|
||||
def bind(
|
||||
self,
|
||||
*,
|
||||
cache_key: str,
|
||||
obj: Any,
|
||||
kind: str,
|
||||
source_path: str,
|
||||
state_provider: Callable[[], dict[str, Any]],
|
||||
evict_callback: Callable[[], bool],
|
||||
sticky: bool = False,
|
||||
priority: int | None = None,
|
||||
note: str | None = None,
|
||||
) -> ExternalResidencyEntry:
|
||||
if obj is None:
|
||||
raise ValueError("Cannot bind None into external residency registry")
|
||||
|
||||
with self._lock:
|
||||
entry_id = self._cache_key_to_entry.get(cache_key)
|
||||
if entry_id is not None and entry_id in self._entries:
|
||||
entry = self._entries[entry_id]
|
||||
else:
|
||||
entry_id = self._make_entry_id(kind, source_path)
|
||||
entry = ExternalResidencyEntry(
|
||||
entry_id=entry_id,
|
||||
cache_key=cache_key,
|
||||
kind=kind,
|
||||
source_path=source_path,
|
||||
sticky=bool(sticky),
|
||||
priority=REGISTRY.default_priority(kind) if priority is None else int(priority),
|
||||
)
|
||||
self._entries[entry_id] = entry
|
||||
self._cache_key_to_entry[cache_key] = entry_id
|
||||
|
||||
try:
|
||||
entry.object_ref = weakref.ref(obj)
|
||||
except TypeError:
|
||||
entry.object_ref = None
|
||||
entry.cache_key = cache_key
|
||||
entry.kind = kind
|
||||
entry.source_path = source_path
|
||||
entry.sticky = bool(sticky)
|
||||
entry.priority = REGISTRY.default_priority(kind) if priority is None else int(priority)
|
||||
entry.state_provider = state_provider
|
||||
entry.evict_callback = evict_callback
|
||||
entry.last_touched = _now()
|
||||
if note and note not in entry.notes:
|
||||
entry.notes.append(note)
|
||||
|
||||
self.refresh_runtime_state()
|
||||
return entry
|
||||
|
||||
def remove(self, *, cache_key: str) -> bool:
|
||||
with self._lock:
|
||||
entry_id = self._cache_key_to_entry.pop(cache_key, None)
|
||||
if entry_id is None:
|
||||
return False
|
||||
return self._entries.pop(entry_id, None) is not None
|
||||
|
||||
def refresh_runtime_state(self) -> None:
|
||||
ensure_external_integrations_installed()
|
||||
stale_keys: list[str] = []
|
||||
with self._lock:
|
||||
for cache_key, entry_id in list(self._cache_key_to_entry.items()):
|
||||
entry = self._entries.get(entry_id)
|
||||
if entry is None:
|
||||
stale_keys.append(cache_key)
|
||||
continue
|
||||
obj = entry.object()
|
||||
if obj is None:
|
||||
stale_keys.append(cache_key)
|
||||
continue
|
||||
state_provider = entry.state_provider
|
||||
if state_provider is None:
|
||||
continue
|
||||
try:
|
||||
state = state_provider() or {}
|
||||
except Exception as exc:
|
||||
_LOG.debug("GPU Resident Loader: failed to refresh external cache state for %s: %s", cache_key, exc)
|
||||
continue
|
||||
entry.current_device = state.get("current_device")
|
||||
entry.load_device = state.get("load_device")
|
||||
entry.offload_device = state.get("offload_device")
|
||||
entry.claimed = bool(state.get("claimed", False))
|
||||
entry.loaded_bytes = _coerce_external_bytes(
|
||||
state.get("loaded_bytes", entry.loaded_bytes or 0),
|
||||
entry.loaded_bytes or 0,
|
||||
cache_key=cache_key,
|
||||
field_name="loaded_bytes",
|
||||
)
|
||||
entry.total_bytes = _coerce_external_bytes(
|
||||
state.get("total_bytes", entry.total_bytes or entry.loaded_bytes or 0),
|
||||
entry.total_bytes or entry.loaded_bytes or 0,
|
||||
cache_key=cache_key,
|
||||
field_name="total_bytes",
|
||||
)
|
||||
for cache_key in stale_keys:
|
||||
entry_id = self._cache_key_to_entry.pop(cache_key, None)
|
||||
if entry_id is not None:
|
||||
self._entries.pop(entry_id, None)
|
||||
|
||||
def candidates(
|
||||
self,
|
||||
*,
|
||||
device: str | torch.device | None,
|
||||
respect_sticky: bool,
|
||||
sticky_floor_priority: int,
|
||||
keep_models: tuple[Any, ...],
|
||||
) -> list[tuple[Any, ExternalResidencyEntry, bool]]:
|
||||
self.refresh_runtime_state()
|
||||
output: list[tuple[Any, ExternalResidencyEntry, bool]] = []
|
||||
with self._lock:
|
||||
for entry in self._entries.values():
|
||||
obj = entry.object()
|
||||
if obj is None:
|
||||
continue
|
||||
if entry.claimed:
|
||||
continue
|
||||
if device is not None and not _device_matches(device, entry.current_device):
|
||||
continue
|
||||
if any(obj is keep for keep in keep_models if keep is not None):
|
||||
continue
|
||||
sticky_respected = (
|
||||
respect_sticky
|
||||
and bool(entry.sticky)
|
||||
and int(entry.priority) >= int(sticky_floor_priority)
|
||||
)
|
||||
output.append((obj, entry, sticky_respected))
|
||||
return output
|
||||
|
||||
def evict(self, entry: ExternalResidencyEntry) -> bool:
|
||||
callback = entry.evict_callback
|
||||
if callback is None:
|
||||
return False
|
||||
try:
|
||||
result = bool(callback())
|
||||
except Exception as exc:
|
||||
_LOG.warning(
|
||||
"GPU Resident Loader: external eviction callback failed for %s: %s",
|
||||
entry.cache_key,
|
||||
exc,
|
||||
)
|
||||
result = False
|
||||
finally:
|
||||
self.refresh_runtime_state()
|
||||
return result
|
||||
|
||||
def snapshot(self) -> list[dict[str, Any]]:
|
||||
self.refresh_runtime_state()
|
||||
with self._lock:
|
||||
items = [entry.as_dict() for entry in self._entries.values()]
|
||||
items.sort(
|
||||
key=lambda item: (
|
||||
not item["sticky"],
|
||||
item["kind"],
|
||||
item["basename"] or "",
|
||||
)
|
||||
)
|
||||
return items
|
||||
|
||||
|
||||
EXTERNAL_REGISTRY = ExternalResidencyRegistry()
|
||||
|
||||
|
||||
def external_trim_enabled() -> bool:
|
||||
value = os.environ.get("COMFYUI_GPU_RESIDENT_TRIM_EXTERNAL", "").strip().lower()
|
||||
return value in {"1", "true", "yes", "on"}
|
||||
|
||||
|
||||
def external_objects_for_models(models: tuple[Any, ...] | list[Any]) -> tuple[Any, ...]:
|
||||
"""
|
||||
Returns registered external objects that are part of any supplied model wrapper chain.
|
||||
|
||||
This lets callers preserve external cache entries when they are associated with a kept
|
||||
model through wrapper indirection rather than exact object identity.
|
||||
"""
|
||||
related_ids: set[int] = set()
|
||||
for model in models:
|
||||
for related in _iter_seedvr2_wrapper_chain(model):
|
||||
related_ids.add(id(related))
|
||||
|
||||
if not related_ids:
|
||||
return ()
|
||||
|
||||
EXTERNAL_REGISTRY.refresh_runtime_state()
|
||||
matches: list[Any] = []
|
||||
seen_ids: set[int] = set()
|
||||
with EXTERNAL_REGISTRY._lock:
|
||||
for entry in EXTERNAL_REGISTRY._entries.values():
|
||||
obj = entry.object()
|
||||
if obj is None:
|
||||
continue
|
||||
obj_id = id(obj)
|
||||
if obj_id in related_ids and obj_id not in seen_ids:
|
||||
matches.append(obj)
|
||||
seen_ids.add(obj_id)
|
||||
return tuple(matches)
|
||||
|
||||
|
||||
def _call_seedvr2_method_with_optional_expected_model(
|
||||
method: Callable[..., Any],
|
||||
*args: Any,
|
||||
debug: Any = None,
|
||||
expected_model: Any = None,
|
||||
) -> Any:
|
||||
try:
|
||||
parameters = inspect.signature(method).parameters
|
||||
except (TypeError, ValueError):
|
||||
parameters = {}
|
||||
|
||||
accepts_kwargs = any(param.kind == inspect.Parameter.VAR_KEYWORD for param in parameters.values())
|
||||
kwargs: dict[str, Any] = {}
|
||||
if accepts_kwargs or "debug" in parameters:
|
||||
kwargs["debug"] = debug
|
||||
if expected_model is not None and (accepts_kwargs or "expected_model" in parameters):
|
||||
kwargs["expected_model"] = expected_model
|
||||
return method(*args, **kwargs)
|
||||
|
||||
|
||||
def _register_seedvr2_cached_model(global_cache: Any, *, kind: str, config: Any, model: Any) -> ExternalResidencyEntry | None:
|
||||
if not isinstance(config, Mapping):
|
||||
_LOG.debug(
|
||||
"GPU Resident Loader: skipping SeedVR2 %s cache entry with unexpected config type: %s",
|
||||
kind,
|
||||
type(config).__name__,
|
||||
)
|
||||
return None
|
||||
|
||||
node_id = config.get("node_id")
|
||||
if node_id is None or model is None:
|
||||
return None
|
||||
|
||||
try:
|
||||
model_ref = weakref.ref(model)
|
||||
except TypeError:
|
||||
def model_ref() -> Any | None:
|
||||
return None
|
||||
|
||||
source_path = _seedvr2_source_path(kind, node_id, config, model)
|
||||
cache_key = _seedvr2_entry_key(kind, node_id)
|
||||
registry_kind = KIND_MODEL if kind == "dit" else KIND_VAE
|
||||
|
||||
if kind == "dit":
|
||||
def evict_callback() -> bool:
|
||||
return bool(
|
||||
_call_seedvr2_method_with_optional_expected_model(
|
||||
global_cache.remove_dit,
|
||||
{"node_id": node_id},
|
||||
debug=None,
|
||||
expected_model=model_ref(),
|
||||
)
|
||||
)
|
||||
note = f"SeedVR2 cached DiT node {node_id}"
|
||||
else:
|
||||
def evict_callback() -> bool:
|
||||
return bool(
|
||||
_call_seedvr2_method_with_optional_expected_model(
|
||||
global_cache.remove_vae,
|
||||
{"node_id": node_id},
|
||||
debug=None,
|
||||
expected_model=model_ref(),
|
||||
)
|
||||
)
|
||||
note = f"SeedVR2 cached VAE node {node_id}"
|
||||
|
||||
return EXTERNAL_REGISTRY.bind(
|
||||
cache_key=cache_key,
|
||||
obj=model,
|
||||
kind=registry_kind,
|
||||
source_path=source_path,
|
||||
state_provider=_seedvr2_state_provider(model_ref, config),
|
||||
evict_callback=evict_callback,
|
||||
sticky=False,
|
||||
priority=REGISTRY.default_priority(registry_kind),
|
||||
note=note,
|
||||
)
|
||||
|
||||
|
||||
def _install_seedvr2_integration_for_module(module: Any) -> bool:
|
||||
model_cache_cls = getattr(module, "GlobalModelCache", None)
|
||||
get_global_cache = getattr(module, "get_global_cache", None)
|
||||
if model_cache_cls is None or not callable(get_global_cache):
|
||||
return False
|
||||
|
||||
class_id = id(model_cache_cls)
|
||||
thread_id = threading.get_ident()
|
||||
original_set_dit = None
|
||||
original_set_vae = None
|
||||
original_replace_dit = None
|
||||
original_replace_vae = None
|
||||
original_remove_dit = None
|
||||
original_remove_vae = None
|
||||
provisional_cache_keys: set[str] = set()
|
||||
provisional_cache_ownership: dict[str, str] = {}
|
||||
|
||||
try:
|
||||
with _SEEDVR2_PATCH_CONDITION:
|
||||
while (
|
||||
class_id in _SEEDVR2_PATCHING_CLASS_IDS
|
||||
and _SEEDVR2_PATCHING_THREAD_IDS.get(class_id) != thread_id
|
||||
):
|
||||
_SEEDVR2_PATCH_CONDITION.wait()
|
||||
if class_id in _SEEDVR2_PATCHED_CLASS_IDS or _SEEDVR2_PATCHING_THREAD_IDS.get(class_id) == thread_id:
|
||||
return True
|
||||
_SEEDVR2_PATCHING_CLASS_IDS.add(class_id)
|
||||
_SEEDVR2_PATCHING_THREAD_IDS[class_id] = thread_id
|
||||
|
||||
original_set_dit = model_cache_cls.set_dit
|
||||
original_set_vae = model_cache_cls.set_vae
|
||||
original_replace_dit = getattr(model_cache_cls, "replace_dit", None)
|
||||
original_replace_vae = getattr(model_cache_cls, "replace_vae", None)
|
||||
original_remove_dit = model_cache_cls.remove_dit
|
||||
original_remove_vae = model_cache_cls.remove_vae
|
||||
|
||||
def set_dit_wrapper(self, dit_config, model, model_name, debug=None):
|
||||
result = original_set_dit(self, dit_config, model, model_name, debug)
|
||||
if result is not None:
|
||||
_register_seedvr2_cached_model(self, kind="dit", config=dit_config, model=model)
|
||||
return result
|
||||
|
||||
def set_vae_wrapper(self, vae_config, model, model_name, debug=None):
|
||||
result = original_set_vae(self, vae_config, model, model_name, debug)
|
||||
if result is not None:
|
||||
_register_seedvr2_cached_model(self, kind="vae", config=vae_config, model=model)
|
||||
return result
|
||||
|
||||
def replace_dit_wrapper(self, dit_config, model, debug=None, expected_model=None):
|
||||
if original_replace_dit is None:
|
||||
return False
|
||||
result = _call_seedvr2_method_with_optional_expected_model(
|
||||
original_replace_dit,
|
||||
self,
|
||||
dit_config,
|
||||
model,
|
||||
debug=debug,
|
||||
expected_model=expected_model,
|
||||
)
|
||||
if result:
|
||||
_register_seedvr2_cached_model(self, kind="dit", config=dit_config, model=model)
|
||||
return result
|
||||
|
||||
def replace_vae_wrapper(self, vae_config, model, debug=None, expected_model=None):
|
||||
if original_replace_vae is None:
|
||||
return False
|
||||
result = _call_seedvr2_method_with_optional_expected_model(
|
||||
original_replace_vae,
|
||||
self,
|
||||
vae_config,
|
||||
model,
|
||||
debug=debug,
|
||||
expected_model=expected_model,
|
||||
)
|
||||
if result:
|
||||
_register_seedvr2_cached_model(self, kind="vae", config=vae_config, model=model)
|
||||
return result
|
||||
|
||||
def remove_dit_wrapper(self, dit_config, debug=None, expected_model=None):
|
||||
result = _call_seedvr2_method_with_optional_expected_model(
|
||||
original_remove_dit,
|
||||
self,
|
||||
dit_config,
|
||||
debug=debug,
|
||||
expected_model=expected_model,
|
||||
)
|
||||
if result:
|
||||
EXTERNAL_REGISTRY.remove(cache_key=_seedvr2_entry_key("dit", dit_config.get("node_id")))
|
||||
return result
|
||||
|
||||
def remove_vae_wrapper(self, vae_config, debug=None, expected_model=None):
|
||||
result = _call_seedvr2_method_with_optional_expected_model(
|
||||
original_remove_vae,
|
||||
self,
|
||||
vae_config,
|
||||
debug=debug,
|
||||
expected_model=expected_model,
|
||||
)
|
||||
if result:
|
||||
EXTERNAL_REGISTRY.remove(cache_key=_seedvr2_entry_key("vae", vae_config.get("node_id")))
|
||||
return result
|
||||
|
||||
global_cache = get_global_cache()
|
||||
model_cache_lock = getattr(global_cache, "_model_cache_lock", None)
|
||||
lock_context = model_cache_lock if model_cache_lock is not None else contextlib.nullcontext()
|
||||
with lock_context:
|
||||
model_cache_cls.set_dit = set_dit_wrapper
|
||||
model_cache_cls.set_vae = set_vae_wrapper
|
||||
if original_replace_dit is not None:
|
||||
model_cache_cls.replace_dit = replace_dit_wrapper
|
||||
if original_replace_vae is not None:
|
||||
model_cache_cls.replace_vae = replace_vae_wrapper
|
||||
model_cache_cls.remove_dit = remove_dit_wrapper
|
||||
model_cache_cls.remove_vae = remove_vae_wrapper
|
||||
dit_items = list(getattr(global_cache, "_dit_models", {}).items())
|
||||
vae_items = list(getattr(global_cache, "_vae_models", {}).items())
|
||||
for _node_id, entry in dit_items:
|
||||
if not isinstance(entry, tuple) or len(entry) != 2:
|
||||
continue
|
||||
model, config = entry
|
||||
if model is not None:
|
||||
if isinstance(config, Mapping) and config.get("node_id") is not None:
|
||||
cache_key = _seedvr2_entry_key("dit", config.get("node_id"))
|
||||
provisional_cache_keys.add(cache_key)
|
||||
prior_entry_id = EXTERNAL_REGISTRY._cache_key_to_entry.get(cache_key)
|
||||
else:
|
||||
cache_key = None
|
||||
prior_entry_id = None
|
||||
registered_entry = _register_seedvr2_cached_model(global_cache, kind="dit", config=config, model=model)
|
||||
if cache_key is not None and registered_entry is not None:
|
||||
if prior_entry_id != registered_entry.entry_id:
|
||||
provisional_cache_ownership[cache_key] = registered_entry.entry_id
|
||||
for _node_id, entry in vae_items:
|
||||
if not isinstance(entry, tuple) or len(entry) != 2:
|
||||
continue
|
||||
model, config = entry
|
||||
if model is not None:
|
||||
if isinstance(config, Mapping) and config.get("node_id") is not None:
|
||||
cache_key = _seedvr2_entry_key("vae", config.get("node_id"))
|
||||
provisional_cache_keys.add(cache_key)
|
||||
prior_entry_id = EXTERNAL_REGISTRY._cache_key_to_entry.get(cache_key)
|
||||
else:
|
||||
cache_key = None
|
||||
prior_entry_id = None
|
||||
registered_entry = _register_seedvr2_cached_model(global_cache, kind="vae", config=config, model=model)
|
||||
if cache_key is not None and registered_entry is not None:
|
||||
if prior_entry_id != registered_entry.entry_id:
|
||||
provisional_cache_ownership[cache_key] = registered_entry.entry_id
|
||||
with _SEEDVR2_PATCH_CONDITION:
|
||||
_SEEDVR2_PATCHING_CLASS_IDS.discard(class_id)
|
||||
_SEEDVR2_PATCHED_CLASS_IDS.add(class_id)
|
||||
_SEEDVR2_PATCHING_THREAD_IDS.pop(class_id, None)
|
||||
_SEEDVR2_PATCH_CONDITION.notify_all()
|
||||
except Exception:
|
||||
_LOG.debug(
|
||||
"GPU Resident Loader: rolling back SeedVR2 integration for class_id=%s provisional_cache_keys=%s",
|
||||
class_id,
|
||||
sorted(provisional_cache_keys),
|
||||
exc_info=True,
|
||||
)
|
||||
for cache_key in provisional_cache_keys:
|
||||
created_entry_id = provisional_cache_ownership.get(cache_key)
|
||||
if created_entry_id is not None:
|
||||
with EXTERNAL_REGISTRY._lock:
|
||||
current_entry_id = EXTERNAL_REGISTRY._cache_key_to_entry.get(cache_key)
|
||||
if current_entry_id == created_entry_id:
|
||||
EXTERNAL_REGISTRY.remove(cache_key=cache_key)
|
||||
with lock_context:
|
||||
if original_set_dit is not None:
|
||||
model_cache_cls.set_dit = original_set_dit
|
||||
if original_set_vae is not None:
|
||||
model_cache_cls.set_vae = original_set_vae
|
||||
if original_replace_dit is not None:
|
||||
model_cache_cls.replace_dit = original_replace_dit
|
||||
if original_replace_vae is not None:
|
||||
model_cache_cls.replace_vae = original_replace_vae
|
||||
if original_remove_dit is not None:
|
||||
model_cache_cls.remove_dit = original_remove_dit
|
||||
if original_remove_vae is not None:
|
||||
model_cache_cls.remove_vae = original_remove_vae
|
||||
with _SEEDVR2_PATCH_CONDITION:
|
||||
_SEEDVR2_PATCHING_CLASS_IDS.discard(class_id)
|
||||
_SEEDVR2_PATCHED_CLASS_IDS.discard(class_id)
|
||||
_SEEDVR2_PATCHING_THREAD_IDS.pop(class_id, None)
|
||||
_SEEDVR2_PATCH_CONDITION.notify_all()
|
||||
raise
|
||||
|
||||
_LOG.info("GPU Resident Loader: integrated external SeedVR2 cache visibility hooks")
|
||||
return True
|
||||
|
||||
|
||||
def ensure_external_integrations_installed() -> None:
|
||||
for module in list(sys.modules.values()):
|
||||
if module is None:
|
||||
continue
|
||||
module_file = getattr(module, "__file__", None)
|
||||
if not module_file:
|
||||
continue
|
||||
normalized_file = os.path.abspath(module_file).replace("\\", "/")
|
||||
if not normalized_file.endswith("/src/core/model_cache.py"):
|
||||
continue
|
||||
if "seedvr2" not in normalized_file.lower():
|
||||
continue
|
||||
if hasattr(module, "GlobalModelCache") and hasattr(module, "get_global_cache"):
|
||||
try:
|
||||
_install_seedvr2_integration_for_module(module)
|
||||
except Exception as exc:
|
||||
_LOG.warning(
|
||||
"GPU Resident Loader: failed to install SeedVR2 external integration from %s: %s",
|
||||
normalized_file,
|
||||
exc,
|
||||
)
|
||||
+150
-1
@@ -14,7 +14,14 @@ import comfy.utils
|
||||
from comfy.cli_args import PerformanceFeature, args
|
||||
from comfy.ldm.modules.attention import attention_pytorch, wrap_attn
|
||||
|
||||
from .patches import checkpoint_component_info_from_header, infer_unet_prefix_from_keys, load_safetensors_state_dict
|
||||
from .cleanup import trim_resident_vram_for_load
|
||||
from .patches import (
|
||||
checkpoint_component_info_from_header,
|
||||
estimate_checkpoint_component_bytes,
|
||||
estimate_safetensors_tensor_bytes,
|
||||
infer_unet_prefix_from_keys,
|
||||
load_safetensors_state_dict,
|
||||
)
|
||||
from .residency import KIND_CHECKPOINT, KIND_CLIP, KIND_MODEL, KIND_VAE, POLICIES, REGISTRY
|
||||
|
||||
|
||||
@@ -328,6 +335,113 @@ def _warn_pickle_checkpoint_gpu_compatibility(loader_name: str, path: str) -> No
|
||||
)
|
||||
|
||||
|
||||
def _fallback_file_size_bytes(path: str) -> int:
|
||||
try:
|
||||
return int(os.path.getsize(path))
|
||||
except OSError:
|
||||
return 0
|
||||
|
||||
|
||||
def _weight_dtype_override(weight_dtype: str) -> torch.dtype | None:
|
||||
if weight_dtype == "fp8_e4m3fn_fast":
|
||||
return torch.float8_e4m3fn
|
||||
return DTYPE_MAP.get(weight_dtype)
|
||||
|
||||
|
||||
def _normalize_keep_models(*objects: Any) -> tuple[Any, ...]:
|
||||
keep_models: list[Any] = []
|
||||
for obj in objects:
|
||||
if obj is None:
|
||||
continue
|
||||
patcher = getattr(obj, "patcher", None)
|
||||
keep = patcher if patcher is not None else obj
|
||||
if not any(existing is keep for existing in keep_models):
|
||||
keep_models.append(keep)
|
||||
return tuple(keep_models)
|
||||
|
||||
|
||||
def _estimate_extra_state_dict_bytes(extra_state_dict: str | None, *, weight_dtype: str) -> int:
|
||||
if not extra_state_dict:
|
||||
return 0
|
||||
dtype_override = _weight_dtype_override(weight_dtype)
|
||||
if _is_safetensors_path(extra_state_dict):
|
||||
estimated = estimate_safetensors_tensor_bytes(extra_state_dict, dtype_override=dtype_override)
|
||||
if estimated is not None:
|
||||
return int(estimated)
|
||||
return _fallback_file_size_bytes(extra_state_dict)
|
||||
|
||||
|
||||
def _estimate_model_load_bytes(
|
||||
source_path: str,
|
||||
*,
|
||||
cache_scope: str,
|
||||
weight_dtype: str,
|
||||
extra_state_dict: str | None,
|
||||
) -> int:
|
||||
dtype_override = _weight_dtype_override(weight_dtype)
|
||||
estimated = None
|
||||
if _is_safetensors_path(source_path):
|
||||
if cache_scope == "checkpoint_model":
|
||||
estimated = estimate_checkpoint_component_bytes(
|
||||
source_path,
|
||||
KIND_MODEL,
|
||||
dtype_override=dtype_override,
|
||||
)
|
||||
else:
|
||||
estimated = estimate_safetensors_tensor_bytes(source_path, dtype_override=dtype_override)
|
||||
if estimated is None:
|
||||
estimated = _fallback_file_size_bytes(source_path)
|
||||
return int(estimated) + _estimate_extra_state_dict_bytes(extra_state_dict, weight_dtype=weight_dtype)
|
||||
|
||||
|
||||
def _estimate_checkpoint_aux_component_bytes(ckpt_path: str, *, kind: str) -> int:
|
||||
estimated = None
|
||||
if _is_safetensors_path(ckpt_path):
|
||||
estimated = estimate_checkpoint_component_bytes(ckpt_path, kind)
|
||||
if estimated is None:
|
||||
estimated = _fallback_file_size_bytes(ckpt_path)
|
||||
return int(estimated)
|
||||
|
||||
|
||||
def _maybe_trim_before_load(
|
||||
*,
|
||||
loader_name: str,
|
||||
reason: str,
|
||||
explicit_device: torch.device | None,
|
||||
required_bytes: int,
|
||||
keep_models: tuple[Any, ...] = (),
|
||||
) -> None:
|
||||
if explicit_device is None or explicit_device.type != "cuda":
|
||||
return
|
||||
if required_bytes <= 0:
|
||||
return
|
||||
|
||||
report = trim_resident_vram_for_load(
|
||||
required_bytes=required_bytes,
|
||||
reason=reason,
|
||||
device=explicit_device,
|
||||
keep_models=keep_models,
|
||||
)
|
||||
if report["freed_vram_bytes"] > 0:
|
||||
_LOG.info(
|
||||
"%s: adaptively freed %.2f GiB before load (%s, estimated %.2f GiB + %.2f GiB headroom)",
|
||||
loader_name,
|
||||
report["freed_vram_bytes"] / (1024 ** 3),
|
||||
report["stopped_reason"],
|
||||
report["estimated_load_bytes"] / (1024 ** 3),
|
||||
report["adaptive_headroom_bytes"] / (1024 ** 3),
|
||||
)
|
||||
if not report["target_met"]:
|
||||
_LOG.warning(
|
||||
"%s: adaptive trim could not reach the estimated headroom for %s (%s, free %.2f GiB / target %.2f GiB)",
|
||||
loader_name,
|
||||
reason,
|
||||
report["stopped_reason"],
|
||||
report["free_after_bytes"] / (1024 ** 3),
|
||||
report["target_free_vram_bytes"] / (1024 ** 3),
|
||||
)
|
||||
|
||||
|
||||
def _selected_unet_key_map_from_header(
|
||||
path: str,
|
||||
*,
|
||||
@@ -485,6 +599,7 @@ def _load_resident_diffusion_model(
|
||||
enable_fp16_accumulation: bool,
|
||||
extra_state_dict: str | None = None,
|
||||
policy_override: str | None = None,
|
||||
keep_models: tuple[Any, ...] = (),
|
||||
) -> Any:
|
||||
model_options = _build_model_options(weight_dtype)
|
||||
effective_policy = _effective_policy_name(policy_override)
|
||||
@@ -506,6 +621,18 @@ def _load_resident_diffusion_model(
|
||||
|
||||
_warn_pickle_checkpoint_gpu_compatibility(loader_name, source_path)
|
||||
explicit_device = REGISTRY.explicit_load_device(kind=KIND_MODEL, source_path=source_path)
|
||||
_maybe_trim_before_load(
|
||||
loader_name=loader_name,
|
||||
reason=f"model load for {os.path.basename(source_path)}",
|
||||
explicit_device=explicit_device,
|
||||
required_bytes=_estimate_model_load_bytes(
|
||||
source_path,
|
||||
cache_scope=cache_scope,
|
||||
weight_dtype=weight_dtype,
|
||||
extra_state_dict=extra_state_dict,
|
||||
),
|
||||
keep_models=keep_models,
|
||||
)
|
||||
|
||||
with _temporary_backend_flags(
|
||||
cublas=patch_cublaslinear,
|
||||
@@ -568,6 +695,7 @@ def _load_checkpoint_clip_only(
|
||||
ckpt_path: str,
|
||||
policy_override: str | None,
|
||||
loader_name: str,
|
||||
keep_models: tuple[Any, ...] = (),
|
||||
):
|
||||
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)
|
||||
@@ -580,6 +708,13 @@ def _load_checkpoint_clip_only(
|
||||
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")
|
||||
explicit_device = REGISTRY.explicit_load_device(kind=KIND_CLIP, source_path=ckpt_path)
|
||||
_maybe_trim_before_load(
|
||||
loader_name=loader_name,
|
||||
reason=f"CLIP load for {os.path.basename(ckpt_path)}",
|
||||
explicit_device=explicit_device,
|
||||
required_bytes=_estimate_checkpoint_aux_component_bytes(ckpt_path, kind=KIND_CLIP),
|
||||
keep_models=keep_models,
|
||||
)
|
||||
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:
|
||||
requested_device = explicit_device if explicit_device is not None else torch.device("cpu")
|
||||
@@ -646,6 +781,7 @@ def _load_checkpoint_vae_only(
|
||||
ckpt_path: str,
|
||||
policy_override: str | None,
|
||||
loader_name: str,
|
||||
keep_models: tuple[Any, ...] = (),
|
||||
):
|
||||
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)
|
||||
@@ -658,6 +794,13 @@ def _load_checkpoint_vae_only(
|
||||
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")
|
||||
explicit_device = REGISTRY.explicit_load_device(kind=KIND_VAE, source_path=ckpt_path)
|
||||
_maybe_trim_before_load(
|
||||
loader_name=loader_name,
|
||||
reason=f"VAE load for {os.path.basename(ckpt_path)}",
|
||||
explicit_device=explicit_device,
|
||||
required_bytes=_estimate_checkpoint_aux_component_bytes(ckpt_path, kind=KIND_VAE),
|
||||
keep_models=keep_models,
|
||||
)
|
||||
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:
|
||||
requested_device = explicit_device if explicit_device is not None else torch.device("cpu")
|
||||
@@ -724,6 +867,7 @@ def _load_full_checkpoint(
|
||||
_LOG.info("%s: reusing live checkpoint outputs for %s", loader_name, ckpt_path)
|
||||
return model, clip, vae
|
||||
|
||||
keep_models = _normalize_keep_models(model, clip, vae)
|
||||
if model is None:
|
||||
model = _load_resident_diffusion_model(
|
||||
loader_name=loader_name,
|
||||
@@ -736,18 +880,23 @@ def _load_full_checkpoint(
|
||||
sage_attention=sage_attention,
|
||||
enable_fp16_accumulation=enable_fp16_accumulation,
|
||||
policy_override=policy_override,
|
||||
keep_models=keep_models,
|
||||
)
|
||||
keep_models = keep_models + _normalize_keep_models(model)
|
||||
if clip is None:
|
||||
clip = _load_checkpoint_clip_only(
|
||||
ckpt_path=ckpt_path,
|
||||
policy_override=policy_override,
|
||||
loader_name=loader_name,
|
||||
keep_models=keep_models,
|
||||
)
|
||||
keep_models = keep_models + _normalize_keep_models(clip)
|
||||
if vae is None:
|
||||
vae = _load_checkpoint_vae_only(
|
||||
ckpt_path=ckpt_path,
|
||||
policy_override=policy_override,
|
||||
loader_name=loader_name,
|
||||
keep_models=keep_models,
|
||||
)
|
||||
return model, clip, vae
|
||||
|
||||
|
||||
@@ -1,11 +1,13 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
import comfy.model_management as model_management
|
||||
|
||||
from .cleanup import trim_resident_vram
|
||||
from .cleanup import unload_loaded_model
|
||||
from .external_residency import EXTERNAL_REGISTRY, ensure_external_integrations_installed
|
||||
from .kj_loader import (
|
||||
CheckpointClipLoaderResident,
|
||||
CheckpointLoaderResident,
|
||||
@@ -16,6 +18,8 @@ from .kj_loader import (
|
||||
)
|
||||
from .residency import REGISTRY
|
||||
|
||||
_LOG = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _entry_report_json(obj: Any) -> str:
|
||||
entry = REGISTRY.entry_for_object(obj)
|
||||
@@ -60,9 +64,28 @@ def _evict_patcher(patcher, *, unpatch_weights: bool) -> bool:
|
||||
raise RuntimeError("Expected a patcher-capable object, but no patcher was found.")
|
||||
unloaded = False
|
||||
for loaded in list(model_management.current_loaded_models):
|
||||
if loaded.model is patcher or loaded.model.is_clone(patcher):
|
||||
loaded.model_unload(unpatch_weights=unpatch_weights)
|
||||
unloaded = True
|
||||
loaded_model = getattr(loaded, "model", None)
|
||||
if loaded_model is None:
|
||||
continue
|
||||
safe_is_clone = False
|
||||
if loaded_model is not patcher:
|
||||
try:
|
||||
safe_is_clone = loaded_model.is_clone(patcher)
|
||||
except Exception as exc:
|
||||
_LOG.warning("GPU Resident Loader: failed to evaluate clone state during eviction: %s", exc)
|
||||
if loaded_model is patcher or safe_is_clone:
|
||||
fully_unloaded = unload_loaded_model(
|
||||
loaded,
|
||||
active_device=getattr(loaded, "device", None),
|
||||
force_offload_to_cpu=True,
|
||||
unpatch_weights=unpatch_weights,
|
||||
)
|
||||
if fully_unloaded and unpatch_weights:
|
||||
try:
|
||||
model_management.current_loaded_models.remove(loaded)
|
||||
except ValueError:
|
||||
pass
|
||||
unloaded = unloaded or fully_unloaded
|
||||
if unloaded and hasattr(model_management, "soft_empty_cache"):
|
||||
model_management.soft_empty_cache()
|
||||
REGISTRY.refresh_runtime_state()
|
||||
@@ -101,7 +124,13 @@ class RegistrySnapshot:
|
||||
CATEGORY = "GPU Resident Loader/residency"
|
||||
|
||||
def snapshot(self):
|
||||
return (REGISTRY.snapshot_json(),)
|
||||
ensure_external_integrations_installed()
|
||||
payload = {
|
||||
"policy": REGISTRY.get_policy(),
|
||||
"entries": REGISTRY.snapshot(),
|
||||
"external_entries": EXTERNAL_REGISTRY.snapshot(),
|
||||
}
|
||||
return (json.dumps(payload, indent=2, sort_keys=True),)
|
||||
|
||||
|
||||
class PinModelResidency:
|
||||
@@ -332,54 +361,6 @@ class ReportVAEResidency:
|
||||
return (_entry_report_json(_patcher_for_vae(vae)),)
|
||||
|
||||
|
||||
class TrimResidentVRAM:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"target_free_vram_mb": ("INT", {"default": 4096, "min": 0, "max": 1_048_576, "step": 1}),
|
||||
"respect_sticky": ("BOOLEAN", {"default": True}),
|
||||
"sticky_floor_priority": ("INT", {"default": 0, "min": -100, "max": 1000, "step": 1}),
|
||||
"allow_partial_unload": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
"optional": {
|
||||
"payload": (
|
||||
"*",
|
||||
{
|
||||
"forceInput": True,
|
||||
"tooltip": "Optional passthrough payload so the trim can sit on a stage boundary without changing graph data.",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("*", "STRING")
|
||||
RETURN_NAMES = ("payload", "trim_report")
|
||||
FUNCTION = "trim"
|
||||
CATEGORY = "GPU Resident Loader/residency"
|
||||
DESCRIPTION = (
|
||||
"Resident-aware VRAM trimmer. It frees only enough VRAM to reach the target budget, "
|
||||
"tries partial unload before full detach, orders sticky entries behind non-sticky work unless their priority falls below the floor, "
|
||||
"and returns a JSON report describing whether the target was met, only partially met, or failed."
|
||||
)
|
||||
|
||||
def trim(
|
||||
self,
|
||||
target_free_vram_mb: int,
|
||||
respect_sticky: bool,
|
||||
sticky_floor_priority: int,
|
||||
allow_partial_unload: bool,
|
||||
payload=None,
|
||||
):
|
||||
report = trim_resident_vram(
|
||||
target_free_vram_bytes=max(0, int(target_free_vram_mb)) * 1024 * 1024,
|
||||
respect_sticky=respect_sticky,
|
||||
sticky_floor_priority=sticky_floor_priority,
|
||||
allow_partial_unload=allow_partial_unload,
|
||||
)
|
||||
return payload, json.dumps(report, indent=2, sort_keys=True)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"DiffusionModelSelectorResident": DiffusionModelSelectorResident,
|
||||
"DiffusionModelLoaderResident": DiffusionModelLoaderResident,
|
||||
@@ -401,7 +382,6 @@ NODE_CLASS_MAPPINGS = {
|
||||
"ReportModelResidency": ReportModelResidency,
|
||||
"ReportClipResidency": ReportClipResidency,
|
||||
"ReportVAEResidency": ReportVAEResidency,
|
||||
"TrimResidentVRAM": TrimResidentVRAM,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
@@ -425,5 +405,4 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"ReportModelResidency": "Report Model Residency",
|
||||
"ReportClipResidency": "Report CLIP Residency",
|
||||
"ReportVAEResidency": "Report VAE Residency",
|
||||
"TrimResidentVRAM": "Trim Resident VRAM",
|
||||
}
|
||||
|
||||
+960
-5
File diff suppressed because it is too large
Load Diff
@@ -525,7 +525,6 @@ class ResidencyRegistry:
|
||||
entry.loaded_bytes = int(loaded.model_loaded_memory())
|
||||
entry.total_bytes = int(loaded.model_memory())
|
||||
entry.current_device = str(loaded.device)
|
||||
entry.last_touched = _now()
|
||||
|
||||
def snapshot(self) -> list[dict[str, Any]]:
|
||||
self.refresh_runtime_state()
|
||||
|
||||
Reference in New Issue
Block a user