diff --git a/.bumpversion.cfg b/.bumpversion.cfg index d05e644..d5c6907 100644 --- a/.bumpversion.cfg +++ b/.bumpversion.cfg @@ -1,5 +1,5 @@ [bumpversion] -current_version = 1.1.26 +current_version = 1.1.27 commit = True tag = True parse = (?P\d+)\.(?P\d+)\.(?P\d+) diff --git a/CHANGELOG.md b/CHANGELOG.md index 99ec2e6..0879836 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -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 diff --git a/README.md b/README.md index 5724d02..36f0c84 100644 --- a/README.md +++ b/README.md @@ -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` diff --git a/__init__.py b/__init__.py index 11c48eb..444e4a3 100644 --- a/__init__.py +++ b/__init__.py @@ -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", diff --git a/docs/model_nodes.md b/docs/model_nodes.md index a087950..5f1ad80 100644 --- a/docs/model_nodes.md +++ b/docs/model_nodes.md @@ -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. diff --git a/nodes/model_nodes.py b/nodes/model_nodes.py index e1b00b9..f67fec5 100644 --- a/nodes/model_nodes.py +++ b/nodes/model_nodes.py @@ -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: diff --git a/nodes/tests/model_nodes_test.py b/nodes/tests/model_nodes_test.py index aff4833..6edd177 100644 --- a/nodes/tests/model_nodes_test.py +++ b/nodes/tests/model_nodes_test.py @@ -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() diff --git a/pyproject.toml b/pyproject.toml index 1b5d8ad..7a2e6f8 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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"]