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