Merge branch 'dev' into 'main'
release: comfyui-web-viewer 1.1.27 See merge request vrch/comfyui/comfyui-web-viewer!62
This commit is contained in:
+1
-1
@@ -1,5 +1,5 @@
|
||||
[bumpversion]
|
||||
current_version = 1.1.26
|
||||
current_version = 1.1.27
|
||||
commit = True
|
||||
tag = True
|
||||
parse = (?P<major>\d+)\.(?P<minor>\d+)\.(?P<patch>\d+)
|
||||
|
||||
@@ -5,6 +5,22 @@ All notable changes to this project will be documented in this file.
|
||||
The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/),
|
||||
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
|
||||
|
||||
## [1.1.27 - 2026-08-05]
|
||||
|
||||
### Added
|
||||
|
||||
- add a checkpoint CLIP-only loader, a CPU-safe ControlNet loader, and a TAESD memory profile node
|
||||
- qualify residual ControlNet TensorRT Engines through the `vrch-tensorrt-controlnet-residual-v1` capability
|
||||
|
||||
### Changed
|
||||
|
||||
- make the TensorRT Auto Loader's PyTorch MODEL input lazy and optional, with explicit model-family and checkpoint fallback settings
|
||||
- reuse one qualified residual Engine across required and optional ControlNet modes without loading a duplicate TensorRT MODEL
|
||||
|
||||
### Fixed
|
||||
|
||||
- preserve CLIP prompt headroom and ControlNet state across TAESD decode in high-VRAM realtime workflows
|
||||
|
||||
## [1.1.26 - 2026-07-31]
|
||||
|
||||
### Added
|
||||
|
||||
@@ -172,6 +172,10 @@ To use the `AUDIO Music to Emotion Detector @ vrch.ai` node, you'll need to inst
|
||||
|
||||
### `Model Nodes`
|
||||
|
||||
- TensorRT Auto Loader with lazy PyTorch fallback and residual ControlNet qualification
|
||||
- Checkpoint CLIP-only Loader for prompt changes without an unused checkpoint UNet
|
||||
- ControlNet Loader with CPU construction/offload safety for high-VRAM TensorRT workflows
|
||||
- TAESD Memory Profile for bounded encode/decode scheduling estimates
|
||||
- Documentation: [Usage of Model nodes](./docs/model_nodes.md)
|
||||
|
||||
### `Text Nodes`
|
||||
|
||||
+7
-1
@@ -14,7 +14,7 @@ from .nodes.audio_music2emo_node import *
|
||||
from .nodes.workflow_export_nodes import *
|
||||
from .nodes.model_nodes import *
|
||||
|
||||
__version__ = "1.1.26"
|
||||
__version__ = "1.1.27"
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"VrchAnyOSCControlNode": VrchAnyOSCControlNode,
|
||||
@@ -34,6 +34,8 @@ NODE_CLASS_MAPPINGS = {
|
||||
"VrchBooleanKeyControlNode": VrchBooleanKeyControlNode,
|
||||
"VrchChannelOSCControlNode": VrchChannelOSCControlNode,
|
||||
"VrchChannelX4OSCControlNode": VrchChannelX4OSCControlNode,
|
||||
"VrchCheckpointClipLoaderNode": VrchCheckpointClipLoaderNode,
|
||||
"VrchControlNetLoaderNode": VrchControlNetLoaderNode,
|
||||
"VrchDelayNode": VrchDelayNode,
|
||||
"VrchDelayOSCControlNode": VrchDelayOSCControlNode,
|
||||
"VrchFloatKeyControlNode": VrchFloatKeyControlNode,
|
||||
@@ -70,6 +72,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"VrchMidiWebSocketChannelLoaderNode": VrchMidiWebSocketChannelLoaderNode,
|
||||
"VrchModelWebViewerNode": VrchModelWebViewerNode,
|
||||
"VrchTensorRTAutoLoaderNode": VrchTensorRTAutoLoaderNode,
|
||||
"VrchTAESDMemoryProfileNode": VrchTAESDMemoryProfileNode,
|
||||
"VrchOSCControlSettingsNode": VrchOSCControlSettingsNode,
|
||||
"VrchQRCodeNode": VrchQRCodeNode,
|
||||
"VrchSwitchOSCControlNode": VrchSwitchOSCControlNode,
|
||||
@@ -108,6 +111,8 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"VrchBooleanKeyControlNode": "BOOLEAN Key Control @ vrch.ai",
|
||||
"VrchChannelOSCControlNode": "CHANNEL OSC Control @ vrch.ai",
|
||||
"VrchChannelX4OSCControlNode": "CHANNEL x4 OSC Control @ vrch.ai",
|
||||
"VrchCheckpointClipLoaderNode": "Checkpoint CLIP-only Loader @ vrch.ai",
|
||||
"VrchControlNetLoaderNode": "ControlNet Loader (CPU Offload) @ vrch.ai",
|
||||
"VrchDelayNode": "DELAY @ vrch.ai",
|
||||
"VrchDelayOSCControlNode": "DELAY OSC Control @ vrch.ai",
|
||||
"VrchFloatKeyControlNode": "FLOAT Key Control @ vrch.ai",
|
||||
@@ -144,6 +149,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"VrchMidiWebSocketChannelLoaderNode": "MIDI WebSocket Channel Loader @ vrch.ai",
|
||||
"VrchModelWebViewerNode": "3D MODEL Web Viewer @ vrch.ai",
|
||||
"VrchTensorRTAutoLoaderNode": "TensorRT Auto Loader @ vrch.ai",
|
||||
"VrchTAESDMemoryProfileNode": "TAESD Memory Profile @ vrch.ai",
|
||||
"VrchOSCControlSettingsNode": "OSC Control Settings @ vrch.ai",
|
||||
"VrchQRCodeNode": "QR Code Generator @ vrch.ai",
|
||||
"VrchSwitchOSCControlNode": "SWITCH OSC Control @ vrch.ai",
|
||||
|
||||
+187
-15
@@ -1,29 +1,201 @@
|
||||
# Model Nodes
|
||||
|
||||
## TensorRT Auto Loader @ vrch.ai
|
||||
The model nodes support TensorRT workflows without forcing an unused PyTorch
|
||||
UNet to remain resident on the GPU. They are backend nodes and do not require
|
||||
custom frontend JavaScript.
|
||||
|
||||
Loads a selected local TensorRT Engine when available while keeping the original ComfyUI `MODEL` as a fallback.
|
||||
## Checkpoint CLIP-only Loader @ vrch.ai
|
||||
|
||||
Loads only the CLIP text encoder from a checkpoint. Use it when TensorRT owns
|
||||
the diffusion model but prompts must remain editable during a live session.
|
||||
|
||||
### Inputs
|
||||
|
||||
- **`model`** (`MODEL`): Original model used by `pytorch` mode and by `auto` fallback.
|
||||
- **`load_mode`**:
|
||||
- `auto`: Load the selected TensorRT Engine; use the original model if loading is unavailable or fails.
|
||||
- `tensorrt`: Require the selected TensorRT Engine and stop the prompt if it cannot be loaded.
|
||||
- `pytorch`: Bypass TensorRT and use the original model.
|
||||
- **`engine_name`**: TensorRT Engine from the host's ComfyUI `output/tensorrt` directory or another registered TensorRT model path.
|
||||
- **`debug`**: Print concise loader, cache, and fallback diagnostics to the ComfyUI server console.
|
||||
- **`ckpt_name`** (`CHECKPOINT`, required): checkpoint selected from ComfyUI's
|
||||
registered `checkpoints` paths. The node does not hard-code a default; a new
|
||||
node uses the first option exposed by ComfyUI.
|
||||
|
||||
### Outputs
|
||||
|
||||
- **`model`** (`MODEL`): TensorRT model or the original model selected by `load_mode`.
|
||||
- **`backend`** (`STRING`): Actual backend, `tensorrt` or `pytorch`.
|
||||
- **`status`** (`STRING`): Loading result or fallback reason.
|
||||
- **`CLIP`**: the checkpoint's CLIP text encoder.
|
||||
|
||||
### Behavior
|
||||
|
||||
The node is independent of Live Console Maintenance and does not require custom frontend JavaScript. Engine choices use ComfyUI's native dropdown. Refresh the ComfyUI frontend after adding a new Engine so the dropdown is rebuilt.
|
||||
- Requests CLIP with `output_clip=True` while explicitly disabling MODEL, VAE,
|
||||
and CLIP Vision output.
|
||||
- Uses ComfyUI's configured embeddings directories.
|
||||
- Keeps ordinary ComfyUI node caching. Changing prompt text re-runs downstream
|
||||
CLIP encoding without reconstructing the checkpoint loader.
|
||||
- Fails the prompt when the checkpoint contains no supported CLIP encoder.
|
||||
|
||||
An Engine name saved by another host or removed later is accepted by the workflow. In `auto` mode it safely falls back to the input model; in `tensorrt` mode it reports an error. TensorRT inference errors that occur after the model has loaded are not retried with PyTorch.
|
||||
The node does not expose device, dtype, cache, VAE, UNet, or CLIP Vision
|
||||
controls. Device and offload behavior remain under ComfyUI model management.
|
||||
|
||||
The current implementation infers the TensorRT model family from the input `MODEL`. The installed `TensorRTLoader` remains responsible for Engine deserialization and compatibility.
|
||||
## ControlNet Loader (CPU Offload) @ vrch.ai
|
||||
|
||||
Loads a ControlNet checkpoint through ComfyUI while forcing construction onto
|
||||
CPU. This prevents `--highvram` from constructing another large model directly
|
||||
on CUDA before ComfyUI can offload or stream it beside a resident TensorRT
|
||||
Engine.
|
||||
|
||||
### Inputs
|
||||
|
||||
- **`control_net_name`** (`CONTROL_NET`, required): checkpoint selected from
|
||||
ComfyUI's registered `controlnet` paths. There is no hard-coded default.
|
||||
|
||||
### Outputs
|
||||
|
||||
- **`CONTROL_NET`**: the loaded ControlNet object.
|
||||
|
||||
### Behavior
|
||||
|
||||
- Temporarily overrides ComfyUI's UNet offload device only for the synchronous
|
||||
ControlNet load call, then restores the original function even on failure.
|
||||
- Serializes loads through a process lock so concurrent calls cannot observe
|
||||
the temporary device override.
|
||||
- Uses ComfyUI's ordinary ControlNet loader and therefore supports any
|
||||
ControlNet checkpoint that the installed ComfyUI version supports. Union
|
||||
type selection, strength, start/end percentages, and conditioning remain in
|
||||
downstream nodes.
|
||||
- Fails the prompt when the checkpoint contains no supported ControlNet model.
|
||||
|
||||
## TAESD Memory Profile @ vrch.ai
|
||||
|
||||
Overrides ComfyUI's full-VAE memory estimate with a fixed estimate appropriate
|
||||
for a TAESD VAE. It does not change VAE math, reserve GPU memory, or unload
|
||||
CLIP.
|
||||
|
||||
### Inputs
|
||||
|
||||
- **`vae`** (`VAE`, required): a TAESD VAE.
|
||||
- **`memory_mib`** (`INT`): reported encode/decode memory requirement.
|
||||
- default: `256`
|
||||
- minimum: `64`
|
||||
- maximum: `1024`
|
||||
- step: `64`
|
||||
|
||||
### Outputs
|
||||
|
||||
- **`VAE`**: the same VAE object with the memory profile applied.
|
||||
|
||||
### Behavior
|
||||
|
||||
- Applies the configured fixed estimate to both encode and decode scheduling.
|
||||
- Adds `vrch_memory_profile` metadata for diagnostics.
|
||||
- Fails the prompt when the input is not a TAESD VAE.
|
||||
|
||||
`64 MiB` is the validated Simple workflow value. The node default remains the
|
||||
more conservative `256 MiB`; lower values should be qualified with the actual
|
||||
resolution and live workload before deployment.
|
||||
|
||||
## TensorRT Auto Loader @ vrch.ai
|
||||
|
||||
Loads a selected local TensorRT Engine and exposes the actual backend and
|
||||
status. PyTorch inputs are lazy, so a healthy TensorRT path does not construct
|
||||
an unused checkpoint UNet.
|
||||
|
||||
### Required inputs
|
||||
|
||||
- **`load_mode`**: backend policy. Default: `auto`.
|
||||
- `auto`: try TensorRT; return a PyTorch fallback when Engine loading is
|
||||
unavailable or fails.
|
||||
- `tensorrt`: require TensorRT and fail the prompt instead of falling back.
|
||||
- `pytorch`: bypass TensorRT.
|
||||
- **`engine_name`**: basename of an `.engine` file from ComfyUI's registered
|
||||
TensorRT paths, including `output/tensorrt`. There is no hard-coded Engine
|
||||
default; `No TensorRT Engine Found` is shown when none are registered.
|
||||
- **`debug`** (`BOOLEAN`): concise loader, cache, and fallback diagnostics.
|
||||
Default: `false`.
|
||||
|
||||
### Optional inputs
|
||||
|
||||
- **`model`** (`MODEL`, lazy): an existing PyTorch model used for model-family
|
||||
inference or fallback only when required.
|
||||
- **`model_type`**: TensorRT model family. Default: `auto`.
|
||||
- `auto`
|
||||
- `sdxl_base`
|
||||
- `sdxl_refiner`
|
||||
- `sd1.x`
|
||||
- `sd2.x-768v`
|
||||
- `svd`
|
||||
- `sd3`
|
||||
- `auraflow`
|
||||
- `flux_dev`
|
||||
- `flux_schnell`
|
||||
- **`fallback_checkpoint`**: checkpoint used to load only a diffusion MODEL
|
||||
when PyTorch fallback is needed. It does not load CLIP, VAE, or CLIP Vision.
|
||||
There is no hard-coded default.
|
||||
- **`require_controlnet`** (`BOOLEAN`): require the installed TensorRT loader
|
||||
and selected Engine to implement the VRCH residual ControlNet contract.
|
||||
Default: `false`.
|
||||
|
||||
### Outputs
|
||||
|
||||
- **`model`** (`MODEL`): TensorRT model or PyTorch fallback.
|
||||
- **`backend`** (`STRING`): actual backend, `tensorrt` or `pytorch`.
|
||||
- **`status`** (`STRING`): selected Engine, residual/control state, or fallback
|
||||
reason.
|
||||
|
||||
### Lazy loading and fallback
|
||||
|
||||
- An explicit `model_type` lets the TensorRT path skip the lazy `model` input.
|
||||
- `model_type=auto` evaluates `model` only to infer the model family.
|
||||
- `auto` or `pytorch` evaluates `model` when no `fallback_checkpoint` is
|
||||
configured. When a fallback checkpoint is configured, fallback loads only
|
||||
its diffusion MODEL on demand.
|
||||
- `tensorrt` mode never hides Engine selection, schema, deserialization, or
|
||||
compatibility errors behind a PyTorch fallback.
|
||||
- `auto` fallback covers loader-time failures. TensorRT inference errors after
|
||||
the MODEL has been returned are not retried with PyTorch.
|
||||
|
||||
Engine choices are host-local. A workflow may contain an Engine name that was
|
||||
saved on another host or later removed; runtime validation handles that as a
|
||||
fallback in `auto` mode or an error in `tensorrt` mode. Engine paths are limited
|
||||
to safe basenames inside registered TensorRT roots.
|
||||
|
||||
### Residual ControlNet contract
|
||||
|
||||
`require_controlnet=true` requires all of the following:
|
||||
|
||||
1. the installed `TensorRTLoader` advertises
|
||||
`vrch-tensorrt-controlnet-residual-v1`;
|
||||
2. the selected Engine contains the residual input schema; and
|
||||
3. an upstream ControlNet produces residuals at inference time.
|
||||
|
||||
A plain Engine is rejected when ControlNet is required. A residual Engine may
|
||||
also serve a ControlNet-OFF workflow when `require_controlnet=false`; missing
|
||||
residuals are zero-filled by the qualified TensorRT loader. Switching the same
|
||||
residual Engine between required and optional modes reuses the cached Engine
|
||||
instead of deserializing a second copy.
|
||||
|
||||
When ComfyUI's `ControlNetApplyAdvanced` receives exact `strength=0`, it may
|
||||
produce no control dictionary. A workflow with `require_controlnet=true` then
|
||||
fails loudly by design instead of silently generating an uncontrolled image.
|
||||
Product workflows should expose an explicit ControlNet ON/OFF state or keep the
|
||||
ON strength above zero.
|
||||
|
||||
### Cache and diagnostics
|
||||
|
||||
The Engine cache key includes model type, Engine name, device/inode identity,
|
||||
size, and modification time. Replacing an Engine invalidates the cached MODEL.
|
||||
The `status` output records the Engine name, whether the residual schema is
|
||||
present, and whether ControlNet is required.
|
||||
|
||||
Refresh the ComfyUI frontend after adding or removing Engine files so native
|
||||
dropdown choices are rebuilt.
|
||||
|
||||
## Recommended TensorRT workflow wiring
|
||||
|
||||
- Route **Checkpoint CLIP-only Loader** to the prompt encoders.
|
||||
- Route **TensorRT Auto Loader** `model` to the sampler.
|
||||
- For ControlNet workflows, route **ControlNet Loader (CPU Offload)** through
|
||||
the appropriate Union/specialized ControlNet configuration and conditioning
|
||||
nodes.
|
||||
- Route a TAESD VAE through **TAESD Memory Profile** before VAE encode/decode.
|
||||
- Use an explicit `model_type` and `fallback_checkpoint` so the healthy
|
||||
TensorRT path never evaluates a full checkpoint MODEL.
|
||||
- Set `require_controlnet=true` only for workflows whose output must be
|
||||
controlled; leave it `false` for ordinary VJ workflows.
|
||||
|
||||
These nodes do not modify image size, prompt, seed, steps, CFG, denoise,
|
||||
ControlNet strength, or other generation defaults.
|
||||
|
||||
+431
-21
@@ -1,5 +1,6 @@
|
||||
"""Model loading and fallback nodes for ComfyUI workflows."""
|
||||
|
||||
import threading
|
||||
from pathlib import Path
|
||||
|
||||
import folder_paths
|
||||
@@ -7,7 +8,11 @@ import folder_paths
|
||||
|
||||
CATEGORY = "vrch.ai/model"
|
||||
NO_ENGINE_OPTION = "No TensorRT Engine Found"
|
||||
NO_CHECKPOINT_OPTION = "No PyTorch Fallback Checkpoint Found"
|
||||
LOAD_MODES = ["auto", "tensorrt", "pytorch"]
|
||||
CONTROLNET_CAPABILITY = "vrch-tensorrt-controlnet-residual-v1"
|
||||
|
||||
print(f"[comfyui-web-viewer] TensorRT capability consumer {CONTROLNET_CAPABILITY}")
|
||||
|
||||
_TENSORRT_MODEL_TYPES = {
|
||||
"SDXL": "sdxl_base",
|
||||
@@ -20,6 +25,60 @@ _TENSORRT_MODEL_TYPES = {
|
||||
"Flux": "flux_dev",
|
||||
"FluxSchnell": "flux_schnell",
|
||||
}
|
||||
_TENSORRT_MODEL_TYPE_OPTIONS = [
|
||||
"auto",
|
||||
*dict.fromkeys(_TENSORRT_MODEL_TYPES.values()),
|
||||
]
|
||||
|
||||
_CONTROLNET_CPU_LOAD_LOCK = threading.Lock()
|
||||
|
||||
|
||||
class VrchCheckpointClipLoaderNode:
|
||||
"""Load only the CLIP component from a checkpoint.
|
||||
|
||||
TensorRT workflows do not use the checkpoint's diffusion model. Loading
|
||||
the complete checkpoint in ``--highvram`` mode can nevertheless place that
|
||||
unused UNet on CUDA and prevent a second TensorRT Engine from fitting.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"ckpt_name": (
|
||||
folder_paths.get_filename_list("checkpoints"),
|
||||
)
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CLIP",)
|
||||
FUNCTION = "load_clip"
|
||||
CATEGORY = CATEGORY
|
||||
|
||||
def load_clip(self, ckpt_name):
|
||||
import comfy.sd
|
||||
|
||||
checkpoint_path = folder_paths.get_full_path_or_raise(
|
||||
"checkpoints",
|
||||
ckpt_name,
|
||||
)
|
||||
result = comfy.sd.load_checkpoint_guess_config(
|
||||
checkpoint_path,
|
||||
output_vae=False,
|
||||
output_clip=True,
|
||||
output_clipvision=False,
|
||||
embedding_directory=folder_paths.get_folder_paths("embeddings"),
|
||||
output_model=False,
|
||||
)
|
||||
if result is None or len(result) < 2 or result[1] is None:
|
||||
raise RuntimeError(
|
||||
"Checkpoint does not contain a supported CLIP text encoder"
|
||||
)
|
||||
print(
|
||||
"[comfyui-web-viewer] Checkpoint CLIP-only load complete: "
|
||||
f"{ckpt_name}"
|
||||
)
|
||||
return (result[1],)
|
||||
|
||||
|
||||
def _register_output_engine_root():
|
||||
@@ -79,6 +138,11 @@ def _engine_options():
|
||||
return engines if engines else [NO_ENGINE_OPTION]
|
||||
|
||||
|
||||
def _checkpoint_options():
|
||||
checkpoints = folder_paths.get_filename_list("checkpoints")
|
||||
return checkpoints if checkpoints else [NO_CHECKPOINT_OPTION]
|
||||
|
||||
|
||||
def _resolve_engine_path(engine_name):
|
||||
_register_output_engine_root()
|
||||
if (
|
||||
@@ -135,6 +199,39 @@ def _infer_tensorrt_model_type(model):
|
||||
return None
|
||||
|
||||
|
||||
def _load_checkpoint_model(checkpoint_name):
|
||||
if (
|
||||
not isinstance(checkpoint_name, str)
|
||||
or not checkpoint_name
|
||||
or checkpoint_name == NO_CHECKPOINT_OPTION
|
||||
):
|
||||
raise RuntimeError("no PyTorch fallback checkpoint was configured")
|
||||
|
||||
import comfy.sd
|
||||
|
||||
checkpoint_path = folder_paths.get_full_path_or_raise(
|
||||
"checkpoints",
|
||||
checkpoint_name,
|
||||
)
|
||||
result = comfy.sd.load_checkpoint_guess_config(
|
||||
checkpoint_path,
|
||||
output_vae=False,
|
||||
output_clip=False,
|
||||
output_clipvision=False,
|
||||
embedding_directory=folder_paths.get_folder_paths("embeddings"),
|
||||
output_model=True,
|
||||
)
|
||||
if result is None or not result or result[0] is None:
|
||||
raise RuntimeError(
|
||||
"fallback checkpoint does not contain a supported diffusion model"
|
||||
)
|
||||
print(
|
||||
"[comfyui-web-viewer] TensorRT lazy PyTorch fallback loaded: "
|
||||
f"{checkpoint_name}"
|
||||
)
|
||||
return result[0]
|
||||
|
||||
|
||||
def _get_tensorrt_loader_class():
|
||||
# Resolve the optional node only when this node executes. This keeps the
|
||||
# vrch.ai node package loadable on hosts without ComfyUI-TensorRT.
|
||||
@@ -148,16 +245,155 @@ def _one_line_error(error):
|
||||
return f"{type(error).__name__}: {message}"[:320]
|
||||
|
||||
|
||||
def _tensorrt_metadata(model):
|
||||
base_model = getattr(model, "model", None)
|
||||
metadata = getattr(base_model, "tensorrt_metadata", None)
|
||||
return metadata if isinstance(metadata, dict) else {}
|
||||
|
||||
|
||||
def _set_tensorrt_control_requirement(model, require_controlnet):
|
||||
required = bool(require_controlnet)
|
||||
base_model = getattr(model, "model", None)
|
||||
diffusion_model = getattr(base_model, "diffusion_model", None)
|
||||
if diffusion_model is not None and hasattr(
|
||||
diffusion_model,
|
||||
"require_controlnet",
|
||||
):
|
||||
diffusion_model.require_controlnet = required
|
||||
metadata = _tensorrt_metadata(model)
|
||||
if metadata:
|
||||
metadata["control_required"] = required
|
||||
|
||||
|
||||
class VrchControlNetLoaderNode:
|
||||
"""Load ControlNet weights on CPU so a resident TRT Engine is not duplicated."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"control_net_name": (
|
||||
folder_paths.get_filename_list("controlnet"),
|
||||
)
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTROL_NET",)
|
||||
FUNCTION = "load_controlnet"
|
||||
CATEGORY = CATEGORY
|
||||
|
||||
def load_controlnet(self, control_net_name):
|
||||
controlnet_path = folder_paths.get_full_path_or_raise(
|
||||
"controlnet",
|
||||
control_net_name,
|
||||
)
|
||||
|
||||
# ComfyUI's --highvram mode normally constructs ControlNet directly on
|
||||
# CUDA. A residual TensorRT Engine can already own most of a 16 GiB
|
||||
# device, so construction itself can OOM before model management gets a
|
||||
# chance to stream or offload weights. Keep this override scoped to the
|
||||
# synchronous load call and restore it even when checkpoint loading
|
||||
# fails. ComfyUI executes model-loader nodes on its single prompt worker;
|
||||
# the lock also prevents overlapping calls through this node.
|
||||
import torch
|
||||
import comfy.controlnet
|
||||
import comfy.model_management
|
||||
|
||||
cpu_device = torch.device("cpu")
|
||||
with _CONTROLNET_CPU_LOAD_LOCK:
|
||||
original_offload_device = (
|
||||
comfy.model_management.unet_offload_device
|
||||
)
|
||||
comfy.model_management.unet_offload_device = lambda: cpu_device
|
||||
try:
|
||||
controlnet = comfy.controlnet.load_controlnet(controlnet_path)
|
||||
finally:
|
||||
comfy.model_management.unet_offload_device = (
|
||||
original_offload_device
|
||||
)
|
||||
|
||||
if controlnet is None:
|
||||
raise RuntimeError(
|
||||
"ControlNet checkpoint is invalid and contains no supported model"
|
||||
)
|
||||
print(
|
||||
"[comfyui-web-viewer] ControlNet CPU-offload load complete: "
|
||||
f"{control_net_name}"
|
||||
)
|
||||
return (controlnet,)
|
||||
|
||||
|
||||
class VrchTAESDMemoryProfileNode:
|
||||
"""Use a TAESD-specific memory estimate instead of the full-VAE estimate."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"vae": ("VAE",),
|
||||
"memory_mib": (
|
||||
"INT",
|
||||
{
|
||||
"default": 256,
|
||||
"min": 64,
|
||||
"max": 1024,
|
||||
"step": 64,
|
||||
},
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("VAE",)
|
||||
FUNCTION = "apply_profile"
|
||||
CATEGORY = CATEGORY
|
||||
|
||||
def apply_profile(
|
||||
self,
|
||||
vae,
|
||||
memory_mib=256,
|
||||
):
|
||||
first_stage_model = getattr(vae, "first_stage_model", None)
|
||||
if type(first_stage_model).__name__ != "TAESD":
|
||||
raise RuntimeError(
|
||||
"TAESD memory profile requires a TAESD VAE"
|
||||
)
|
||||
|
||||
memory_bytes = int(memory_mib) * 1024 * 1024
|
||||
vae.memory_used_encode = (
|
||||
lambda _shape, _dtype: memory_bytes
|
||||
)
|
||||
vae.memory_used_decode = (
|
||||
lambda _shape, _dtype: memory_bytes
|
||||
)
|
||||
vae.vrch_memory_profile = {
|
||||
"kind": "taesd",
|
||||
"memory_mib": int(memory_mib),
|
||||
}
|
||||
print(
|
||||
"[comfyui-web-viewer] TAESD memory profile active: "
|
||||
f"{int(memory_mib)} MiB"
|
||||
)
|
||||
return (vae,)
|
||||
|
||||
|
||||
class VrchTensorRTAutoLoaderNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"load_mode": (LOAD_MODES, {"default": "auto"}),
|
||||
"engine_name": (_engine_options(),),
|
||||
"debug": ("BOOLEAN", {"default": False}),
|
||||
}
|
||||
},
|
||||
"optional": {
|
||||
"model": ("MODEL", {"lazy": True}),
|
||||
"model_type": (
|
||||
_TENSORRT_MODEL_TYPE_OPTIONS,
|
||||
{"default": "auto"},
|
||||
),
|
||||
"fallback_checkpoint": (_checkpoint_options(),),
|
||||
"require_controlnet": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL", "STRING", "STRING")
|
||||
@@ -168,31 +404,92 @@ class VrchTensorRTAutoLoaderNode:
|
||||
def __init__(self):
|
||||
self._cached_key = None
|
||||
self._cached_model = None
|
||||
self._cached_status = None
|
||||
|
||||
@classmethod
|
||||
def VALIDATE_INPUTS(cls, engine_name):
|
||||
def VALIDATE_INPUTS(
|
||||
cls,
|
||||
engine_name,
|
||||
require_controlnet=False,
|
||||
model_type="auto",
|
||||
fallback_checkpoint=None,
|
||||
):
|
||||
# Engine choices are host-local. A workflow saved on another host, or
|
||||
# before an Engine was removed, must reach load_model() so auto mode can
|
||||
# fall back instead of failing ComfyUI's pre-execution COMBO check.
|
||||
return True
|
||||
|
||||
def check_lazy_status(
|
||||
self,
|
||||
model=None,
|
||||
load_mode="auto",
|
||||
engine_name=NO_ENGINE_OPTION,
|
||||
debug=False,
|
||||
require_controlnet=False,
|
||||
model_type="auto",
|
||||
fallback_checkpoint=None,
|
||||
):
|
||||
del engine_name, debug, require_controlnet
|
||||
if model_type == "auto":
|
||||
return ["model"]
|
||||
if load_mode in ("auto", "pytorch") and not fallback_checkpoint:
|
||||
return ["model"]
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(cls, model, load_mode, engine_name, debug=False):
|
||||
def IS_CHANGED(
|
||||
cls,
|
||||
model=None,
|
||||
load_mode="auto",
|
||||
engine_name=NO_ENGINE_OPTION,
|
||||
debug=False,
|
||||
require_controlnet=False,
|
||||
model_type="auto",
|
||||
fallback_checkpoint=None,
|
||||
):
|
||||
del model, debug, fallback_checkpoint
|
||||
if load_mode == "pytorch":
|
||||
return "pytorch"
|
||||
return ("pytorch", model_type, bool(require_controlnet))
|
||||
engine_path = _resolve_engine_path(engine_name)
|
||||
if engine_path is None:
|
||||
return (load_mode, engine_name, "missing")
|
||||
return (
|
||||
load_mode,
|
||||
engine_name,
|
||||
model_type,
|
||||
bool(require_controlnet),
|
||||
"missing",
|
||||
)
|
||||
try:
|
||||
fingerprint = _engine_fingerprint(engine_path)
|
||||
except OSError:
|
||||
return (load_mode, engine_name, "unreadable")
|
||||
return (load_mode, engine_name, fingerprint)
|
||||
return (
|
||||
load_mode,
|
||||
engine_name,
|
||||
model_type,
|
||||
bool(require_controlnet),
|
||||
"unreadable",
|
||||
)
|
||||
return (
|
||||
load_mode,
|
||||
engine_name,
|
||||
model_type,
|
||||
bool(require_controlnet),
|
||||
fingerprint,
|
||||
)
|
||||
|
||||
def load_model(self, model, load_mode, engine_name, debug=False):
|
||||
def load_model(
|
||||
self,
|
||||
model=None,
|
||||
load_mode="auto",
|
||||
engine_name=NO_ENGINE_OPTION,
|
||||
debug=False,
|
||||
require_controlnet=False,
|
||||
model_type="auto",
|
||||
fallback_checkpoint=None,
|
||||
):
|
||||
if load_mode == "pytorch":
|
||||
return self._pytorch_result(
|
||||
model,
|
||||
self._fallback_model(model, fallback_checkpoint),
|
||||
"PyTorch selected",
|
||||
debug,
|
||||
)
|
||||
@@ -204,15 +501,19 @@ class VrchTensorRTAutoLoaderNode:
|
||||
load_mode,
|
||||
f"TensorRT Engine is unavailable: {engine_name}",
|
||||
debug,
|
||||
fallback_checkpoint=fallback_checkpoint,
|
||||
)
|
||||
|
||||
model_type = _infer_tensorrt_model_type(model)
|
||||
if model_type is None:
|
||||
resolved_model_type = model_type
|
||||
if resolved_model_type == "auto":
|
||||
resolved_model_type = _infer_tensorrt_model_type(model)
|
||||
if resolved_model_type not in _TENSORRT_MODEL_TYPE_OPTIONS[1:]:
|
||||
return self._load_failure(
|
||||
model,
|
||||
load_mode,
|
||||
"the input model type is not supported by TensorRTLoader",
|
||||
"the TensorRT model type is unavailable or unsupported",
|
||||
debug,
|
||||
fallback_checkpoint=fallback_checkpoint,
|
||||
)
|
||||
|
||||
loader_class = _get_tensorrt_loader_class()
|
||||
@@ -222,6 +523,18 @@ class VrchTensorRTAutoLoaderNode:
|
||||
load_mode,
|
||||
"TensorRTLoader is not installed or registered",
|
||||
debug,
|
||||
fallback_checkpoint=fallback_checkpoint,
|
||||
)
|
||||
|
||||
loader_capability = getattr(loader_class, "CONTROLNET_CAPABILITY", None)
|
||||
if require_controlnet and loader_capability != CONTROLNET_CAPABILITY:
|
||||
return self._load_failure(
|
||||
model,
|
||||
load_mode,
|
||||
"TensorRTLoader lacks the required residual ControlNet capability "
|
||||
f"{CONTROLNET_CAPABILITY}",
|
||||
debug,
|
||||
fallback_checkpoint=fallback_checkpoint,
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -233,56 +546,153 @@ class VrchTensorRTAutoLoaderNode:
|
||||
f"TensorRT Engine cannot be read: {_one_line_error(error)}",
|
||||
debug,
|
||||
cause=error,
|
||||
fallback_checkpoint=fallback_checkpoint,
|
||||
)
|
||||
|
||||
cache_key = (id(model), model_type, engine_name, fingerprint)
|
||||
cache_key = (
|
||||
resolved_model_type,
|
||||
engine_name,
|
||||
fingerprint,
|
||||
)
|
||||
if cache_key == self._cached_key and self._cached_model is not None:
|
||||
self._debug(debug, f"cache hit engine={engine_name} model_type={model_type}")
|
||||
metadata = _tensorrt_metadata(self._cached_model)
|
||||
residual_schema = bool(metadata.get("residual_schema", False))
|
||||
if require_controlnet and not residual_schema:
|
||||
return self._load_failure(
|
||||
model,
|
||||
load_mode,
|
||||
"ControlNet was required but the TensorRT Engine has no "
|
||||
"residual bindings",
|
||||
debug,
|
||||
fallback_checkpoint=fallback_checkpoint,
|
||||
)
|
||||
_set_tensorrt_control_requirement(
|
||||
self._cached_model,
|
||||
require_controlnet,
|
||||
)
|
||||
self._cached_status = self._active_status(
|
||||
engine_name,
|
||||
residual_schema,
|
||||
require_controlnet,
|
||||
)
|
||||
self._debug(
|
||||
debug,
|
||||
f"cache hit engine={engine_name} model_type={resolved_model_type}",
|
||||
)
|
||||
return (
|
||||
self._cached_model,
|
||||
"tensorrt",
|
||||
f"TensorRT active: {engine_name}",
|
||||
self._cached_status,
|
||||
)
|
||||
|
||||
self._debug(debug, f"loading engine={engine_name} model_type={model_type}")
|
||||
self._debug(
|
||||
debug,
|
||||
f"loading engine={engine_name} model_type={resolved_model_type}",
|
||||
)
|
||||
try:
|
||||
loaded = loader_class().load_unet(engine_name, model_type)
|
||||
loader = loader_class()
|
||||
if loader_capability == CONTROLNET_CAPABILITY:
|
||||
loaded = loader.load_unet(
|
||||
engine_name,
|
||||
resolved_model_type,
|
||||
require_controlnet=bool(require_controlnet),
|
||||
)
|
||||
else:
|
||||
loaded = loader.load_unet(engine_name, resolved_model_type)
|
||||
if not isinstance(loaded, tuple) or not loaded or loaded[0] is None:
|
||||
raise RuntimeError("TensorRTLoader returned no MODEL")
|
||||
tensorrt_model = loaded[0]
|
||||
metadata = _tensorrt_metadata(tensorrt_model)
|
||||
residual_schema = bool(metadata.get("residual_schema", False))
|
||||
if require_controlnet and not residual_schema:
|
||||
raise RuntimeError(
|
||||
"TensorRTLoader returned an Engine without the required residual schema"
|
||||
)
|
||||
_set_tensorrt_control_requirement(
|
||||
tensorrt_model,
|
||||
require_controlnet,
|
||||
)
|
||||
except Exception as error:
|
||||
self._cached_key = None
|
||||
self._cached_model = None
|
||||
self._cached_status = None
|
||||
return self._load_failure(
|
||||
model,
|
||||
load_mode,
|
||||
f"TensorRT load failed: {_one_line_error(error)}",
|
||||
debug,
|
||||
cause=error,
|
||||
fallback_checkpoint=fallback_checkpoint,
|
||||
)
|
||||
|
||||
self._cached_key = cache_key
|
||||
self._cached_model = tensorrt_model
|
||||
self._cached_status = self._active_status(
|
||||
engine_name,
|
||||
residual_schema,
|
||||
require_controlnet,
|
||||
)
|
||||
self._debug(debug, f"TensorRT active engine={engine_name}")
|
||||
return (
|
||||
tensorrt_model,
|
||||
"tensorrt",
|
||||
f"TensorRT active: {engine_name}",
|
||||
self._cached_status,
|
||||
)
|
||||
|
||||
def _load_failure(self, model, load_mode, reason, debug, cause=None):
|
||||
def _load_failure(
|
||||
self,
|
||||
model,
|
||||
load_mode,
|
||||
reason,
|
||||
debug,
|
||||
cause=None,
|
||||
fallback_checkpoint=None,
|
||||
):
|
||||
self._cached_key = None
|
||||
self._cached_model = None
|
||||
self._cached_status = None
|
||||
self._debug(debug, reason)
|
||||
if load_mode == "tensorrt":
|
||||
error = RuntimeError(f"TensorRT Auto Loader: {reason}")
|
||||
if cause is not None:
|
||||
raise error from cause
|
||||
raise error
|
||||
return self._pytorch_result(model, f"PyTorch fallback: {reason}", debug)
|
||||
try:
|
||||
fallback_model = self._fallback_model(
|
||||
model,
|
||||
fallback_checkpoint,
|
||||
)
|
||||
except Exception as fallback_error:
|
||||
error = RuntimeError(
|
||||
"TensorRT Auto Loader: "
|
||||
f"{reason}; PyTorch fallback failed: "
|
||||
f"{_one_line_error(fallback_error)}"
|
||||
)
|
||||
raise error from fallback_error
|
||||
return self._pytorch_result(
|
||||
fallback_model,
|
||||
f"PyTorch fallback: {reason}",
|
||||
debug,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _fallback_model(model, fallback_checkpoint):
|
||||
if model is not None:
|
||||
return model
|
||||
return _load_checkpoint_model(fallback_checkpoint)
|
||||
|
||||
def _pytorch_result(self, model, status, debug):
|
||||
self._debug(debug, status)
|
||||
return (model, "pytorch", status)
|
||||
|
||||
@staticmethod
|
||||
def _active_status(engine_name, residual_schema, require_controlnet):
|
||||
return (
|
||||
f"TensorRT active: {engine_name}; "
|
||||
f"residual_schema={str(bool(residual_schema)).lower()}; "
|
||||
f"control_required={str(bool(require_controlnet)).lower()}"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _debug(enabled, message):
|
||||
if enabled:
|
||||
|
||||
@@ -55,6 +55,8 @@ class FakeFolderPaths:
|
||||
return [str(self.root)]
|
||||
|
||||
def get_filename_list(self, folder_name):
|
||||
if folder_name == "checkpoints":
|
||||
return ["fallback.safetensors"]
|
||||
return sorted(path.name for path in self.root.glob("*.engine"))
|
||||
|
||||
def get_full_path(self, folder_name, filename):
|
||||
@@ -73,21 +75,43 @@ class TestTensorRTAutoLoaderNode(unittest.TestCase):
|
||||
|
||||
self.original_folder_paths = model_nodes.folder_paths
|
||||
self.original_get_loader = model_nodes._get_tensorrt_loader_class
|
||||
self.original_load_checkpoint_model = model_nodes._load_checkpoint_model
|
||||
model_nodes.folder_paths = FakeFolderPaths(self.engine_root)
|
||||
self.addCleanup(setattr, model_nodes, "folder_paths", self.original_folder_paths)
|
||||
self.addCleanup(setattr, model_nodes, "_get_tensorrt_loader_class", self.original_get_loader)
|
||||
self.addCleanup(
|
||||
setattr,
|
||||
model_nodes,
|
||||
"_load_checkpoint_model",
|
||||
self.original_load_checkpoint_model,
|
||||
)
|
||||
|
||||
self.model = FakeModelPatcher()
|
||||
self.tensorrt_model = object()
|
||||
self.tensorrt_model = types.SimpleNamespace(
|
||||
model=types.SimpleNamespace(
|
||||
diffusion_model=types.SimpleNamespace(
|
||||
require_controlnet=False,
|
||||
),
|
||||
tensorrt_metadata={},
|
||||
)
|
||||
)
|
||||
|
||||
def install_loader(self, error=None):
|
||||
def install_loader(self, error=None, capability=None, metadata=None):
|
||||
output_model = self.tensorrt_model
|
||||
if metadata is not None:
|
||||
output_model.model.tensorrt_metadata = metadata
|
||||
|
||||
class FakeLoader:
|
||||
CONTROLNET_CAPABILITY = capability
|
||||
calls = []
|
||||
|
||||
def load_unet(self, engine_name, model_type):
|
||||
self.calls.append((engine_name, model_type))
|
||||
def load_unet(self, engine_name, model_type, require_controlnet=False):
|
||||
if capability is None:
|
||||
self.calls.append((engine_name, model_type))
|
||||
else:
|
||||
self.calls.append(
|
||||
(engine_name, model_type, bool(require_controlnet))
|
||||
)
|
||||
if error is not None:
|
||||
raise error
|
||||
return (output_model,)
|
||||
@@ -98,10 +122,21 @@ class TestTensorRTAutoLoaderNode(unittest.TestCase):
|
||||
def test_node_contract(self):
|
||||
inputs = model_nodes.VrchTensorRTAutoLoaderNode.INPUT_TYPES()["required"]
|
||||
|
||||
self.assertEqual(inputs["model"], ("MODEL",))
|
||||
self.assertEqual(inputs["load_mode"][0], ["auto", "tensorrt", "pytorch"])
|
||||
self.assertEqual(inputs["engine_name"][0], ["test.engine"])
|
||||
self.assertEqual(inputs["debug"], ("BOOLEAN", {"default": False}))
|
||||
optional = model_nodes.VrchTensorRTAutoLoaderNode.INPUT_TYPES()["optional"]
|
||||
self.assertEqual(optional["model"], ("MODEL", {"lazy": True}))
|
||||
self.assertEqual(optional["model_type"][0][0], "auto")
|
||||
self.assertIn("sdxl_base", optional["model_type"][0])
|
||||
self.assertEqual(
|
||||
optional["fallback_checkpoint"][0],
|
||||
["fallback.safetensors"],
|
||||
)
|
||||
self.assertEqual(
|
||||
optional["require_controlnet"],
|
||||
("BOOLEAN", {"default": False}),
|
||||
)
|
||||
self.assertEqual(
|
||||
model_nodes.VrchTensorRTAutoLoaderNode.RETURN_NAMES,
|
||||
("model", "backend", "status"),
|
||||
@@ -110,9 +145,41 @@ class TestTensorRTAutoLoaderNode(unittest.TestCase):
|
||||
def test_stale_engine_validation_only_accepts_engine_name(self):
|
||||
signature = inspect.signature(model_nodes.VrchTensorRTAutoLoaderNode.VALIDATE_INPUTS)
|
||||
|
||||
self.assertEqual(list(signature.parameters), ["engine_name"])
|
||||
self.assertEqual(
|
||||
list(signature.parameters),
|
||||
[
|
||||
"engine_name",
|
||||
"require_controlnet",
|
||||
"model_type",
|
||||
"fallback_checkpoint",
|
||||
],
|
||||
)
|
||||
self.assertTrue(model_nodes.VrchTensorRTAutoLoaderNode.VALIDATE_INPUTS("removed.engine"))
|
||||
|
||||
def test_explicit_model_type_and_checkpoint_keep_model_input_lazy(self):
|
||||
node = model_nodes.VrchTensorRTAutoLoaderNode()
|
||||
|
||||
self.assertEqual(
|
||||
node.check_lazy_status(
|
||||
model=None,
|
||||
load_mode="auto",
|
||||
engine_name="test.engine",
|
||||
model_type="sdxl_base",
|
||||
fallback_checkpoint="fallback.safetensors",
|
||||
),
|
||||
[],
|
||||
)
|
||||
self.assertEqual(
|
||||
node.check_lazy_status(
|
||||
model=None,
|
||||
load_mode="auto",
|
||||
engine_name="test.engine",
|
||||
model_type="auto",
|
||||
fallback_checkpoint="fallback.safetensors",
|
||||
),
|
||||
["model"],
|
||||
)
|
||||
|
||||
def test_pytorch_mode_bypasses_tensorrt(self):
|
||||
loader = self.install_loader()
|
||||
|
||||
@@ -136,6 +203,26 @@ class TestTensorRTAutoLoaderNode(unittest.TestCase):
|
||||
self.assertIs(second[0], self.tensorrt_model)
|
||||
self.assertEqual(loader.calls, [("test.engine", "sdxl_base")])
|
||||
|
||||
def test_auto_mode_uses_explicit_type_without_loading_fallback(self):
|
||||
loader = self.install_loader()
|
||||
fallback_calls = []
|
||||
model_nodes._load_checkpoint_model = (
|
||||
lambda checkpoint: fallback_calls.append(checkpoint)
|
||||
)
|
||||
|
||||
result = model_nodes.VrchTensorRTAutoLoaderNode().load_model(
|
||||
model=None,
|
||||
load_mode="auto",
|
||||
engine_name="test.engine",
|
||||
model_type="sdxl_base",
|
||||
fallback_checkpoint="fallback.safetensors",
|
||||
)
|
||||
|
||||
self.assertIs(result[0], self.tensorrt_model)
|
||||
self.assertEqual(result[1], "tensorrt")
|
||||
self.assertEqual(loader.calls, [("test.engine", "sdxl_base")])
|
||||
self.assertEqual(fallback_calls, [])
|
||||
|
||||
def test_auto_mode_falls_back_when_engine_is_missing(self):
|
||||
result = model_nodes.VrchTensorRTAutoLoaderNode().load_model(
|
||||
self.model, "auto", "removed.engine", False
|
||||
@@ -145,6 +232,27 @@ class TestTensorRTAutoLoaderNode(unittest.TestCase):
|
||||
self.assertEqual(result[1], "pytorch")
|
||||
self.assertIn("fallback", result[2])
|
||||
|
||||
def test_auto_mode_lazily_loads_checkpoint_after_engine_failure(self):
|
||||
fallback_model = FakeModelPatcher()
|
||||
fallback_calls = []
|
||||
|
||||
def load_fallback(checkpoint):
|
||||
fallback_calls.append(checkpoint)
|
||||
return fallback_model
|
||||
|
||||
model_nodes._load_checkpoint_model = load_fallback
|
||||
result = model_nodes.VrchTensorRTAutoLoaderNode().load_model(
|
||||
model=None,
|
||||
load_mode="auto",
|
||||
engine_name="removed.engine",
|
||||
model_type="sdxl_base",
|
||||
fallback_checkpoint="fallback.safetensors",
|
||||
)
|
||||
|
||||
self.assertIs(result[0], fallback_model)
|
||||
self.assertEqual(result[1], "pytorch")
|
||||
self.assertEqual(fallback_calls, ["fallback.safetensors"])
|
||||
|
||||
def test_tensorrt_mode_fails_when_engine_is_missing(self):
|
||||
with self.assertRaisesRegex(RuntimeError, "Engine is unavailable"):
|
||||
model_nodes.VrchTensorRTAutoLoaderNode().load_model(
|
||||
@@ -170,6 +278,141 @@ class TestTensorRTAutoLoaderNode(unittest.TestCase):
|
||||
self.model, "tensorrt", "test.engine", False
|
||||
)
|
||||
|
||||
def test_controlnet_requirement_fails_closed_with_legacy_loader(self):
|
||||
self.install_loader()
|
||||
|
||||
with self.assertRaisesRegex(RuntimeError, "lacks the required residual"):
|
||||
model_nodes.VrchTensorRTAutoLoaderNode().load_model(
|
||||
self.model,
|
||||
"tensorrt",
|
||||
"test.engine",
|
||||
False,
|
||||
require_controlnet=True,
|
||||
)
|
||||
|
||||
def test_controlnet_auto_mode_falls_back_with_legacy_loader(self):
|
||||
self.install_loader()
|
||||
|
||||
result = model_nodes.VrchTensorRTAutoLoaderNode().load_model(
|
||||
self.model,
|
||||
"auto",
|
||||
"test.engine",
|
||||
False,
|
||||
require_controlnet=True,
|
||||
)
|
||||
|
||||
self.assertIs(result[0], self.model)
|
||||
self.assertEqual(result[1], "pytorch")
|
||||
self.assertIn("required residual", result[2])
|
||||
|
||||
def test_controlnet_requirement_loads_qualified_residual_engine(self):
|
||||
loader = self.install_loader(
|
||||
capability=model_nodes.CONTROLNET_CAPABILITY,
|
||||
metadata={"residual_schema": True},
|
||||
)
|
||||
|
||||
result = model_nodes.VrchTensorRTAutoLoaderNode().load_model(
|
||||
self.model,
|
||||
"tensorrt",
|
||||
"test.engine",
|
||||
False,
|
||||
require_controlnet=True,
|
||||
)
|
||||
|
||||
self.assertIs(result[0], self.tensorrt_model)
|
||||
self.assertEqual(result[1], "tensorrt")
|
||||
self.assertIn("residual_schema=true", result[2])
|
||||
self.assertIn("control_required=true", result[2])
|
||||
self.assertEqual(loader.calls, [("test.engine", "sdxl_base", True)])
|
||||
|
||||
def test_controlnet_mode_switch_reuses_one_residual_engine(self):
|
||||
loader = self.install_loader(
|
||||
capability=model_nodes.CONTROLNET_CAPABILITY,
|
||||
metadata={"residual_schema": True},
|
||||
)
|
||||
node = model_nodes.VrchTensorRTAutoLoaderNode()
|
||||
|
||||
off = node.load_model(
|
||||
self.model,
|
||||
"tensorrt",
|
||||
"test.engine",
|
||||
False,
|
||||
require_controlnet=False,
|
||||
)
|
||||
on = node.load_model(
|
||||
FakeModelPatcher(),
|
||||
"tensorrt",
|
||||
"test.engine",
|
||||
False,
|
||||
require_controlnet=True,
|
||||
)
|
||||
off_again = node.load_model(
|
||||
self.model,
|
||||
"tensorrt",
|
||||
"test.engine",
|
||||
False,
|
||||
require_controlnet=False,
|
||||
)
|
||||
|
||||
self.assertIs(off[0], self.tensorrt_model)
|
||||
self.assertIs(on[0], self.tensorrt_model)
|
||||
self.assertIs(off_again[0], self.tensorrt_model)
|
||||
self.assertEqual(
|
||||
loader.calls,
|
||||
[("test.engine", "sdxl_base", False)],
|
||||
)
|
||||
self.assertIn("control_required=true", on[2])
|
||||
self.assertIn("control_required=false", off_again[2])
|
||||
self.assertFalse(
|
||||
self.tensorrt_model.model.diffusion_model.require_controlnet
|
||||
)
|
||||
|
||||
def test_controlnet_mode_switch_rejects_cached_plain_engine(self):
|
||||
loader = self.install_loader(
|
||||
capability=model_nodes.CONTROLNET_CAPABILITY,
|
||||
metadata={"residual_schema": False},
|
||||
)
|
||||
node = model_nodes.VrchTensorRTAutoLoaderNode()
|
||||
node.load_model(
|
||||
self.model,
|
||||
"tensorrt",
|
||||
"test.engine",
|
||||
False,
|
||||
require_controlnet=False,
|
||||
)
|
||||
|
||||
with self.assertRaisesRegex(
|
||||
RuntimeError,
|
||||
"ControlNet was required but the TensorRT Engine has no residual bindings",
|
||||
):
|
||||
node.load_model(
|
||||
self.model,
|
||||
"tensorrt",
|
||||
"test.engine",
|
||||
False,
|
||||
require_controlnet=True,
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
loader.calls,
|
||||
[("test.engine", "sdxl_base", False)],
|
||||
)
|
||||
|
||||
def test_controlnet_requirement_rejects_missing_residual_metadata(self):
|
||||
self.install_loader(
|
||||
capability=model_nodes.CONTROLNET_CAPABILITY,
|
||||
metadata={"residual_schema": False},
|
||||
)
|
||||
|
||||
with self.assertRaisesRegex(RuntimeError, "without the required residual"):
|
||||
model_nodes.VrchTensorRTAutoLoaderNode().load_model(
|
||||
self.model,
|
||||
"tensorrt",
|
||||
"test.engine",
|
||||
False,
|
||||
require_controlnet=True,
|
||||
)
|
||||
|
||||
def test_missing_inventory_uses_placeholder(self):
|
||||
self.engine_path.unlink()
|
||||
|
||||
@@ -178,5 +421,203 @@ class TestTensorRTAutoLoaderNode(unittest.TestCase):
|
||||
self.assertEqual(options, [model_nodes.NO_ENGINE_OPTION])
|
||||
|
||||
|
||||
class TestCheckpointClipLoaderNode(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.original_folder_paths = model_nodes.folder_paths
|
||||
self.original_modules = {
|
||||
name: sys.modules.get(name)
|
||||
for name in ("comfy", "comfy.sd")
|
||||
}
|
||||
self.addCleanup(self.restore_modules)
|
||||
self.addCleanup(
|
||||
setattr,
|
||||
model_nodes,
|
||||
"folder_paths",
|
||||
self.original_folder_paths,
|
||||
)
|
||||
|
||||
def restore_modules(self):
|
||||
for name, module in self.original_modules.items():
|
||||
if module is None:
|
||||
sys.modules.pop(name, None)
|
||||
else:
|
||||
sys.modules[name] = module
|
||||
|
||||
def install_runtime(self, clip=object()):
|
||||
calls = []
|
||||
sd_module = types.ModuleType("comfy.sd")
|
||||
|
||||
def load_checkpoint_guess_config(path, **kwargs):
|
||||
calls.append((path, kwargs))
|
||||
return (None, clip, None, None)
|
||||
|
||||
sd_module.load_checkpoint_guess_config = load_checkpoint_guess_config
|
||||
comfy_module = types.ModuleType("comfy")
|
||||
comfy_module.sd = sd_module
|
||||
sys.modules["comfy"] = comfy_module
|
||||
sys.modules["comfy.sd"] = sd_module
|
||||
model_nodes.folder_paths = types.SimpleNamespace(
|
||||
get_filename_list=lambda _name: ["sdxl.safetensors"],
|
||||
get_full_path_or_raise=lambda folder, name: f"/{folder}/{name}",
|
||||
get_folder_paths=lambda folder: [f"/{folder}"],
|
||||
)
|
||||
return calls, clip
|
||||
|
||||
def test_loads_only_clip_without_constructing_checkpoint_unet(self):
|
||||
calls, clip = self.install_runtime()
|
||||
|
||||
result = model_nodes.VrchCheckpointClipLoaderNode().load_clip(
|
||||
"sdxl.safetensors"
|
||||
)
|
||||
|
||||
self.assertIs(result[0], clip)
|
||||
self.assertEqual(calls[0][0], "/checkpoints/sdxl.safetensors")
|
||||
self.assertFalse(calls[0][1]["output_model"])
|
||||
self.assertFalse(calls[0][1]["output_vae"])
|
||||
self.assertTrue(calls[0][1]["output_clip"])
|
||||
|
||||
def test_rejects_checkpoint_without_clip(self):
|
||||
self.install_runtime(clip=None)
|
||||
|
||||
with self.assertRaisesRegex(RuntimeError, "does not contain"):
|
||||
model_nodes.VrchCheckpointClipLoaderNode().load_clip(
|
||||
"sdxl.safetensors"
|
||||
)
|
||||
|
||||
|
||||
class TestControlNetLoaderNode(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.original_folder_paths = model_nodes.folder_paths
|
||||
self.original_modules = {
|
||||
name: sys.modules.get(name)
|
||||
for name in (
|
||||
"torch",
|
||||
"comfy",
|
||||
"comfy.controlnet",
|
||||
"comfy.model_management",
|
||||
)
|
||||
}
|
||||
self.addCleanup(self.restore_modules)
|
||||
self.addCleanup(
|
||||
setattr,
|
||||
model_nodes,
|
||||
"folder_paths",
|
||||
self.original_folder_paths,
|
||||
)
|
||||
|
||||
def restore_modules(self):
|
||||
for name, module in self.original_modules.items():
|
||||
if module is None:
|
||||
sys.modules.pop(name, None)
|
||||
else:
|
||||
sys.modules[name] = module
|
||||
|
||||
def install_runtime(self, error=None):
|
||||
calls = []
|
||||
gpu_device = object()
|
||||
cpu_device = object()
|
||||
model_management = types.ModuleType("comfy.model_management")
|
||||
original_offload = lambda: gpu_device
|
||||
model_management.unet_offload_device = original_offload
|
||||
|
||||
controlnet_module = types.ModuleType("comfy.controlnet")
|
||||
|
||||
def load_controlnet(path):
|
||||
calls.append((path, model_management.unet_offload_device()))
|
||||
if error is not None:
|
||||
raise error
|
||||
return types.SimpleNamespace(
|
||||
control_model_wrapped=types.SimpleNamespace(
|
||||
offload_device=model_management.unet_offload_device()
|
||||
)
|
||||
)
|
||||
|
||||
controlnet_module.load_controlnet = load_controlnet
|
||||
comfy_module = types.ModuleType("comfy")
|
||||
comfy_module.controlnet = controlnet_module
|
||||
comfy_module.model_management = model_management
|
||||
torch_module = types.ModuleType("torch")
|
||||
torch_module.device = lambda name: cpu_device if name == "cpu" else name
|
||||
|
||||
sys.modules["torch"] = torch_module
|
||||
sys.modules["comfy"] = comfy_module
|
||||
sys.modules["comfy.controlnet"] = controlnet_module
|
||||
sys.modules["comfy.model_management"] = model_management
|
||||
return calls, model_management, original_offload, cpu_device
|
||||
|
||||
def test_loads_controlnet_with_cpu_offload_and_restores_global(self):
|
||||
calls, model_management, original_offload, cpu_device = (
|
||||
self.install_runtime()
|
||||
)
|
||||
model_nodes.folder_paths = types.SimpleNamespace(
|
||||
get_full_path_or_raise=lambda folder, name: f"/{folder}/{name}",
|
||||
)
|
||||
|
||||
result = model_nodes.VrchControlNetLoaderNode().load_controlnet(
|
||||
"union.safetensors"
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
calls,
|
||||
[("/controlnet/union.safetensors", cpu_device)],
|
||||
)
|
||||
self.assertIs(
|
||||
result[0].control_model_wrapped.offload_device,
|
||||
cpu_device,
|
||||
)
|
||||
self.assertIs(
|
||||
model_management.unet_offload_device,
|
||||
original_offload,
|
||||
)
|
||||
|
||||
def test_restores_global_after_load_failure(self):
|
||||
_, model_management, original_offload, _ = self.install_runtime(
|
||||
RuntimeError("broken checkpoint")
|
||||
)
|
||||
model_nodes.folder_paths = types.SimpleNamespace(
|
||||
get_full_path_or_raise=lambda _folder, _name: "/broken.safetensors",
|
||||
)
|
||||
|
||||
with self.assertRaisesRegex(RuntimeError, "broken checkpoint"):
|
||||
model_nodes.VrchControlNetLoaderNode().load_controlnet(
|
||||
"broken.safetensors"
|
||||
)
|
||||
|
||||
self.assertIs(
|
||||
model_management.unet_offload_device,
|
||||
original_offload,
|
||||
)
|
||||
|
||||
|
||||
class TestTAESDMemoryProfileNode(unittest.TestCase):
|
||||
def test_applies_fixed_encode_and_decode_budget(self):
|
||||
class TAESD:
|
||||
pass
|
||||
|
||||
vae = types.SimpleNamespace(
|
||||
first_stage_model=TAESD(),
|
||||
memory_used_encode=lambda _shape, _dtype: 1,
|
||||
memory_used_decode=lambda _shape, _dtype: 2,
|
||||
)
|
||||
|
||||
result = model_nodes.VrchTAESDMemoryProfileNode().apply_profile(
|
||||
vae,
|
||||
256,
|
||||
)
|
||||
|
||||
self.assertIs(result[0], vae)
|
||||
self.assertEqual(vae.memory_used_encode(None, None), 256 * 1024 * 1024)
|
||||
self.assertEqual(vae.memory_used_decode(None, None), 256 * 1024 * 1024)
|
||||
self.assertEqual(
|
||||
vae.vrch_memory_profile,
|
||||
{"kind": "taesd", "memory_mib": 256},
|
||||
)
|
||||
|
||||
def test_rejects_non_taesd_vae(self):
|
||||
vae = types.SimpleNamespace(first_stage_model=object())
|
||||
|
||||
with self.assertRaisesRegex(RuntimeError, "requires a TAESD VAE"):
|
||||
model_nodes.VrchTAESDMemoryProfileNode().apply_profile(vae, 256)
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "comfyui-web-viewer"
|
||||
description = "The ComfyUI Web Viewer by vrch.ai is a custom node collection offering a real-time AI-generated interactive art framework. This utility integrates realtime streaming into ComfyUI workflows, supporting keyboard control nodes, OSC control nodes, sound input nodes, and more. Accessible from any device with a web browser, it enables real time interaction with AI-generated content, making it ideal for interactive visual projects and enhancing ComfyUI workflows with efficient content management and display."
|
||||
version = "1.1.26"
|
||||
version = "1.1.27"
|
||||
license = {file = "LICENSE"}
|
||||
dependencies = ["aiohttp","ffmpeg-python","matplotlib","pydub","audioop-lts; python_version >= '3.13'","python-osc","qrcode[pil]","scikit-learn","srt","websockets"]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user