Author SHA1 Message Date
gokayfem 3f9612774e fix: support Transformers 4.x Grounding DINO loading 2026-08-01 14:00:56 +03:00
gokayfem 7a74f5a079 docs: add citation metadata 2026-08-01 03:18:57 +03:00
Gökay Aydoğan 67344abe6a Merge pull request #167 from octo-patch/octo/20260731-music-generation-tool-recvqgT4IrGB9O
Add MiniMax music generation and cover node
2026-07-31 19:25:50 +03:00
octo-patch e79f316908 Add MiniMax music generation node 2026-07-31 20:25:20 +08:00
Gökay Aydoğan f4bc8b9eef Merge pull request #166 from gokayfem/codex/robotics-vla
feat: add robotics VLA policy toolkit
2026-07-31 01:42:15 +03:00
gokayfem 5779c50b20 test: keep robotics clients optional 2026-07-31 01:37:50 +03:00
gokayfem 48101541a3 feat: add robotics VLA policy toolkit 2026-07-31 01:34:48 +03:00
Gökay Aydoğan fcfdf7b210 Merge pull request #165 from gokayfem/dependabot/github_actions/actions-674967a53d
ci: bump actions/upload-artifact from 4 to 7 in the actions group
2026-07-30 13:33:38 +03:00
dependabot[bot] 56cdd25aa9 ci: bump actions/upload-artifact from 4 to 7 in the actions group
Bumps the actions group with 1 update: [actions/upload-artifact](https://github.com/actions/upload-artifact).


Updates `actions/upload-artifact` from 4 to 7
- [Release notes](https://github.com/actions/upload-artifact/releases)
- [Commits](https://github.com/actions/upload-artifact/compare/v4...v7)

---
updated-dependencies:
- dependency-name: actions/upload-artifact
  dependency-version: '7'
  dependency-type: direct:production
  update-type: version-update:semver-major
  dependency-group: actions
...

Signed-off-by: dependabot[bot] <support@github.com>
2026-07-30 10:21:19 +00:00
gokayfem 9aeca11c35 Merge pull request #164 from gokayfem/chore/repo-hygiene-audit 2026-07-30 13:20:35 +03:00
gokayfem 79a929c1ca fix: make node reference schemas exact 2026-07-30 13:14:54 +03:00
Gokay Aydogan e2b20cde13 test: add real-weight node smoke test and share the path bootstrap
The offline suite stubs the llama.cpp boundary, so it proves what the nodes
send and how they handle responses, but never that inference works. Adds an
opt-in real-weight script that drives the node classes themselves:

- LLMLoader resolving a real GGUF through ComfyUI folder_paths, and
  returning before any weights load
- LLMSampler generating real text, asserted deterministic for a fixed seed
  at temperature 0 rather than asserting on model knowledge
- StructuredOutput constraining a real model to a generated JSON Schema.
  This is the llama.cpp grammar path, which a stub cannot verify at all.
- LLMOptionalMemoryFreeSimple releasing a real llama.cpp allocation

Verified against ggml-org/Qwen3.5-0.8B-GGUF (563 MB, Q4_0) on Metal with
llama-cpp-python 0.3.34: all checks pass.

Also extracts tests/_bootstrap.py. Four of the six manual scripts imported
the package only when the checkout directory was named ComfyUI_VLM_nodes,
which conftest.py already notes is not safe for worktrees named after a
branch. All six now share one helper.
2026-07-30 11:17:11 +03:00
Gokay Aydogan 04275b57cb docs: point pre-3.3.0 changelog links at commits
Compare links assumed tags that do not exist. Retroactively tagging the
2.x/3.x versions would run current CI against code that predates it, so
tagging starts at 3.3.0 and earlier entries link to the commit that
declared each version.
2026-07-30 03:10:24 +03:00
Gokay Aydogan 2e41de6ac2 docs: add complete node reference and documentation guards
47 of the 78 registered node classes were never named in the README, so
there was no way to look up a node seen on a canvas. Adds a reference of
all 78, grouped by menu category, with the class_type that appears in
workflow JSON.

Guards it with tests so it cannot drift again: every registered node must
appear in the README, declared license must match LICENSE, and the current
version must have a changelog entry. All three were verified to fail when
violated.
2026-07-30 03:08:52 +03:00
Gokay Aydogan 8b4226474a docs: add changelog, contributing guide, and templates
The repo had no releases, no tags, and no changelog despite being on
version 3.3.0, so users could not pin a version, roll back, or tell that a
bug they filed had been fixed.

CHANGELOG.md reconstructs the 2.x/3.x history from the version-to-commit
mapping in git, and cites the issues each change resolved.

CONTRIBUTING.md documents the constraints that are easy to break: no work
at import time, never reorder existing widgets, no forceInput, optional
dependencies must fail only their own node, and never install torch.

Issue templates require the environment detail that the historically
unresolvable reports lacked, and route llama-cpp-python build failures to
the upstream install guide.
2026-07-30 03:08:52 +03:00
Gokay Aydogan 8a432184d3 ci: add ruff lint job and a coverage floor
CI previously ran pytest, compileall, and build, with no linting and no
coverage measurement. Adds a fast lint job and a --cov-fail-under=70 gate
on the Linux/Python 3.13 leg (currently 73%), plus a coverage artifact.
2026-07-30 03:08:52 +03:00
Gokay Aydogan da5f4d5787 style: apply ruff autofixes
Mechanical only: import ordering, typing -> collections.abc imports,
PEP 604 unions, dict.fromkeys, one unused import, and two over-long
tooltip strings rewrapped by hand. No behaviour change.
2026-07-30 03:08:52 +03:00
Gokay Aydogan 45d21d0642 fix: declare Apache-2.0 in package metadata, add tooling config
pyproject.toml declared license = "MIT" while LICENSE has been Apache-2.0
since the initial commit in 2024-01. The MIT string was introduced in
39fc116 one day earlier, so this is a fresh regression, and Apache-2.0 is
the real license: 2.5 years of outside contributions landed under it.

This metadata is published to the Comfy Registry and into any built wheel,
so the wrong license was being advertised downstream.

Also adds ruff and pytest configuration, and requirements-dev.txt for the
lint/coverage tooling. Vendored nodes/joytagger is excluded from lint; the
three ignored rules are documented inline with their reasons.

Bumps version to 3.3.1. Two fixes (b8ae298, c13ee23) landed after 3.3.0
without a version bump, so the publish workflow saw 3.3.0 already on the
registry and skipped them. They reach Registry users with this release.
2026-07-30 03:08:51 +03:00
Gokay Aydogan 8e55c81b34 test: add contract tests for GGUF text and multimodal nodes
nodes/suggest.py (1011 LOC, 11 node classes) and nodes/llavaloader.py
(599 LOC, 6 node classes) had no test references at all, despite carrying
the longest bug history in the pack.

Coverage: suggest.py 0% -> 98%, llavaloader.py 0% -> 99%.

The cases pin the behaviours the historical reports depended on:
- widget order, which Comfy serializes by position (#156)
- sampling kwargs reaching create_chat_completion (#144)
- handle reuse, and teardown on both success and failure (#137)
- structured-output JSON Schema construction and its error paths
- ChatMusician owning the 'respond in ABC notation' instruction (#149)

No llama.cpp wheel, GGUF weights, or GPU are required.
2026-07-30 03:08:51 +03:00
Gökay Aydoğan f06a2a3e6c Merge pull request #163 from gokayfem/codex/fix-smolvlm-llama-cpp-setup
Fix SmolVLM setup dependencies
2026-07-30 02:27:56 +03:00
Gokay Aydogan c13ee2364e Fix SmolVLM setup dependencies 2026-07-30 02:23:06 +03:00
gokayfem b8ae298abf Fix Moondream 2 and 3.1 local inference 2026-07-29 23:17:59 +03:00
Gökay Aydoğan bb7d51777c Merge pull request #162 from gokayfem/codex/moondream-3-1
Add Moondream Photon and universal VLM acceleration
2026-07-29 17:52:50 +03:00
gokayfem 102f1662ac Add Moondream Photon and universal VLM acceleration 2026-07-29 17:47:35 +03:00
Gökay Aydoğan 44fefcb57a Add adaptive video intelligence and text toolkit (#161) 2026-07-29 16:05:41 +03:00
Gökay Aydoğan 505b324f66 Modernize and secure hosted LLM and VLM APIs (#160)
* Modernize and secure hosted LLM and VLM APIs

* Add web search and portable structured VLM output
2026-07-29 14:59:34 +03:00
Gökay Aydoğan 39fc116341 Add unified VLM vision, segmentation, tracking, and creator mask tools (#159)
* Add unified vision detection segmentation and tracking

* Add creator-ready mask and compositing tools
2026-07-29 13:51:35 +03:00
81 changed files with 18229 additions and 554 deletions
+104
View File
@@ -0,0 +1,104 @@
name: Bug report
description: A node fails, errors, or produces wrong output.
labels: ["bug"]
body:
- type: markdown
attributes:
value: |
Most unresolvable reports are missing the environment details below.
Please run the **VLM Runtime Diagnostics** node and paste its output —
it captures your OS, Python, PyTorch, accelerator backend, and which
optional backends are installed.
- type: input
id: version
attributes:
label: Node pack version
description: From ComfyUI Manager, or the `version` in `pyproject.toml`.
placeholder: "3.3.1"
validations:
required: true
- type: dropdown
id: install
attributes:
label: How did you install it?
options:
- ComfyUI Manager
- Comfy Registry
- git clone into custom_nodes
- Other (describe below)
validations:
required: true
- type: dropdown
id: comfy
attributes:
label: ComfyUI flavour
options:
- ComfyUI Desktop
- ComfyUI Portable (python_embeded)
- Manual install (venv)
- Manual install (conda)
- Cloud / RunPod / other host
validations:
required: true
- type: textarea
id: diagnostics
attributes:
label: VLM Runtime Diagnostics output
description: Add the node to any workflow, run it, and paste the result.
render: text
validations:
required: true
- type: input
id: node
attributes:
label: Which node fails?
placeholder: "LLMSampler, LLavaSamplerSimple, ModernVLM, ..."
validations:
required: true
- type: input
id: model
attributes:
label: Which model / GGUF file?
description: Include the exact filename or Hugging Face repo id.
placeholder: "Qwen 3 VL 4B Instruct, or llava-1.6-mistral-7b.Q4_K_M.gguf"
validations:
required: true
- type: textarea
id: expected
attributes:
label: What did you expect, and what happened instead?
validations:
required: true
- type: textarea
id: traceback
attributes:
label: Full console output
description: |
The complete traceback from the ComfyUI terminal, not just the last
line. Include the startup log if the pack failed to import.
render: shell
validations:
required: true
- type: checkboxes
id: checks
attributes:
label: Before submitting
options:
- label: I updated to the latest version of this node pack and ComfyUI.
required: true
- label: I searched existing open and closed issues.
required: true
- label: >-
If this involves GGUF or `llama-cpp-python`, I installed it with
the arguments for my accelerator from the
[llama-cpp-python install guide](https://github.com/abetlen/llama-cpp-python#installation).
required: false
+11
View File
@@ -0,0 +1,11 @@
blank_issues_enabled: false
contact_links:
- name: llama-cpp-python installation help
url: https://github.com/abetlen/llama-cpp-python#installation
about: >-
Build or GPU-offload failures for GGUF nodes are almost always
llama-cpp-python installation issues. Install the wheel matching your
accelerator first.
- name: ComfyUI Manager and installation problems
url: https://github.com/Comfy-Org/ComfyUI-Manager/issues
about: For problems installing or updating custom nodes in general.
+46
View File
@@ -0,0 +1,46 @@
name: Model or feature request
description: Ask for support for a new VLM/LLM, or a new node.
labels: ["enhancement"]
body:
- type: textarea
id: what
attributes:
label: What would you like added?
validations:
required: true
- type: input
id: model
attributes:
label: Model repository (if requesting a model)
description: A Hugging Face repo id, so the architecture can be checked.
placeholder: "Qwen/Qwen3-VL-8B-Instruct"
- type: dropdown
id: backend
attributes:
label: Which backend would it use?
options:
- transformers (safetensors)
- llama.cpp (GGUF)
- Hosted API
- Not sure
validations:
required: true
- type: textarea
id: why
attributes:
label: What does it let you do that current nodes cannot?
validations:
required: true
- type: checkboxes
id: checks
attributes:
label: Before submitting
options:
- label: >-
I checked the README node reference to confirm this is not already
supported.
required: true
+33
View File
@@ -0,0 +1,33 @@
## What does this change?
<!-- One or two sentences. Link any issue it closes: "Closes #123". -->
## Type of change
- [ ] Bug fix
- [ ] New model support
- [ ] New node
- [ ] Refactor / maintenance
- [ ] Documentation
## Checklist
- [ ] `python -m pytest -q` passes.
- [ ] `python -m ruff check .` passes.
- [ ] Importing the pack still performs no network access, compilation, or
package install.
- [ ] If a node schema changed, existing widget order is preserved (Comfy
serializes widget values by position, so reordering breaks saved
workflows).
- [ ] New optional dependencies fail only the node that needs them, with an
actionable error.
- [ ] `pyproject.toml` `version` is bumped if this is user-visible, and
`CHANGELOG.md` has an entry. Releases only publish on a version change.
## Testing
<!--
Which nodes did you run, on which backend (CUDA / ROCm / Metal / XPU / CPU),
and with which model? Real-weight checks are opt-in:
python tests/manual_model_smoke.py --model "Qwen 3 VL 4B Instruct"
-->
+14
View File
@@ -0,0 +1,14 @@
version: 2
updates:
# Action versions only. Python dependency ranges are deliberately loose
# because ComfyUI owns torch, numpy, and Pillow in the shared environment.
- package-ecosystem: github-actions
directory: "/"
schedule:
interval: monthly
open-pull-requests-limit: 5
commit-message:
prefix: "ci"
groups:
actions:
patterns: ["*"]
+35 -1
View File
@@ -8,6 +8,22 @@ permissions:
contents: read
jobs:
lint:
name: Lint
runs-on: ubuntu-latest
timeout-minutes: 10
steps:
- uses: actions/checkout@v7
- uses: actions/setup-python@v7
with:
python-version: "3.12"
cache: pip
cache-dependency-path: requirements-dev.txt
- name: Install lint tooling
run: python -m pip install -r requirements-dev.txt
- name: Ruff
run: python -m ruff check --output-format github .
test:
name: ${{ matrix.label }}
runs-on: ${{ matrix.os }}
@@ -20,18 +36,22 @@ jobs:
os: ubuntu-latest
python: "3.10"
cpu_index: true
coverage: false
- label: Linux / Python 3.13
os: ubuntu-latest
python: "3.13"
cpu_index: true
coverage: true
- label: Windows / Python 3.12
os: windows-latest
python: "3.12"
cpu_index: true
coverage: false
- label: macOS / Python 3.12
os: macos-14
python: "3.12"
cpu_index: false
coverage: false
steps:
- uses: actions/checkout@v7
- uses: actions/setup-python@v7
@@ -52,10 +72,24 @@ jobs:
- name: Install ComfyUI and node dependencies
run: |
git clone --depth 1 https://github.com/Comfy-Org/ComfyUI.git ../ComfyUI
python -m pip install pytest packaging build
python -m pip install -r requirements-dev.txt
python -m pip install -r ../ComfyUI/requirements.txt -r requirements.txt
- name: Test
if: matrix.coverage == false
run: python -m pytest -q
- name: Test with coverage
if: matrix.coverage == true
run: >-
python -m pytest -q
--cov=nodes --cov-report=term-missing:skip-covered
--cov-report=xml --cov-fail-under=70
- name: Upload coverage report
if: matrix.coverage == true && always()
uses: actions/upload-artifact@v7
with:
name: coverage-xml
path: coverage.xml
if-no-files-found: warn
- name: Compile
run: python -m compileall -q .
- name: Build distribution
+199
View File
@@ -0,0 +1,199 @@
# Changelog
All notable changes to this project are documented here.
The format follows [Keep a Changelog](https://keepachangelog.com/en/1.1.0/),
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
Versions are published to the [Comfy Registry](https://registry.comfy.org/)
from `pyproject.toml`. A release is only published when `version` changes, so
every user-visible fix needs a version bump.
## [3.5.0] - 2026-07-31
### Added
- A MiniMax music node with fixed global and China endpoints, generation and
cover model selection, regional request fields, URL and hexadecimal response
decoding, and MP3, WAV, and PCM output through the existing audio contract.
### Security
- MiniMax credentials are read only from `MINIMAX_API_KEY`; workflows cannot
supply a key or redirect it to a custom endpoint, and request errors redact
the resolved value before reaching ComfyUI.
## [3.4.0] - 2026-07-31
### Added
- A robotics-safe VLA layer with typed embodiment, observation, and action
contracts; bounded multi-camera history; trajectory inspection and preview;
action-chunk replanning; and explicit bounds, rate, dimension, horizon, and
non-finite-value checks before handoff.
- Native policy clients for OpenPI's WebSocket protocol and NVIDIA Isaac
GR00T's ZeroMQ protocol, plus a portable authenticated HTTP/JPEG protocol for
isolated policy runtimes.
- An isolated current-LeRobot policy server with pre/postprocessor support,
serialized inference, optional idle CPU offload, checkpoint feature metadata,
and environment-only bearer authentication.
- A curated 15-model VLA catalog covering SmolVLA, X-VLA, the OpenPI family,
GR00T N1.7, WALL-OSS, MolmoAct2, VLA-JEPA, LingBot-VA, FastWAM, EO-1, EVO-1,
OpenVLA-OFT, and Octo with explicit readiness and fine-tuning requirements.
- A complete API workflow, setup guide, compatibility matrix, security
guidance, and real-weight SmolVLA validation on an RTX 3090.
### Security
- Workflow JSON never stores robotics API keys. The clients read only
`VLA_POLICY_TOKEN`, `OPENPI_API_KEY`, or `GROOT_API_TOKEN` from the
environment, redact them from errors/reports, reject embedded URL
credentials, and require encrypted transports plus explicit opt-in for
remote endpoints where the upstream protocol supports encryption.
- The included policy server bounds request, camera, history, and response
sizes and never uses pickle across the network.
## [3.3.1] - 2026-07-30
### Fixed
- Package metadata declared `license = "MIT"` while the bundled `LICENSE` has
been Apache-2.0 since the initial commit. Built wheels therefore contained
contradictory MIT metadata and Apache-2.0 license text. The Registry already
referenced the license file and was unaffected. Metadata now says
`Apache-2.0`.
- Moondream 2 and Moondream 3.1 local inference (`b8ae298`).
- SmolVLM setup dependencies (`c13ee23`).
The two fixes above landed on `main` after 3.3.0 without a version bump, so
the Registry publish workflow saw 3.3.0 already published and skipped them.
They reach Registry users for the first time in 3.3.1.
### Added
- Test coverage for the GGUF text and multimodal node families, which
previously had none: `nodes/suggest.py` (0% to 98%) and
`nodes/llavaloader.py` (0% to 99%). The new cases pin the behaviours behind
the pack's longest-running bug reports: widget ordering (#156), sampling
kwarg plumbing (#144), and handle teardown on both success and failure
(#137).
- `ruff` lint gate and a coverage floor in CI, plus `requirements-dev.txt`
for the tooling.
- `CHANGELOG.md`, `CONTRIBUTING.md`, issue and pull request templates, and a
Dependabot configuration.
- A complete node reference in the README covering all 78 registered nodes.
## [3.3.0] - 2026-07-29
### Added
- Moondream Photon support and universal VLM acceleration utilities, including
the image pixel-budget and performance-profile nodes (`102f166`).
## [3.2.0] - 2026-07-29
### Added
- Adaptive video intelligence with temporal reasoning, plus the text workflow
toolkit (join, template, clean, replace, split, JSON extract, inspect)
(`44fefcb`).
## [3.1.0] - 2026-07-29
### Changed
- Hosted LLM and VLM API nodes modernized and hardened, with provider profiles
for OpenAI, Google Gemini, Anthropic, xAI, DeepSeek, and others (`505b324`).
## [3.0.0] - 2026-07-29
### Added
- Unified vision stack: open-vocabulary detection (Grounding DINO, OWLv2,
OmDet), SAM2.1 and SAM3.1 segmentation, tracking, and creator mask tools,
with structured detection/segmentation schemas (`39fc116`).
### Changed
- **Breaking:** detection and segmentation nodes now emit structured data
types rather than loose strings. Workflows wiring these outputs into text
nodes need the new converter utilities.
## [2.3.0] - 2026-07-29
### Added
- Reliable streaming VLM text output (`239c904`).
## [2.2.0] - 2026-07-29
### Changed
- llama.cpp GGUF runtime modernized. `llama-cpp-agent` was removed in favour
of llama-cpp-python's native JSON Schema support, which resolves the
unstable wrapper API behind the `unexpected keyword argument 'temperature'`
crashes (#144).
## [2.1.0] - 2026-07-28
### Added
- Cross-platform runtime support across NVIDIA CUDA, AMD ROCm, Apple Metal,
Intel XPU, and CPU, without replacing ComfyUI's PyTorch (`4c200c4`).
## [2.0.1] - 2026-07-28
### Added
- Small VLM catalog and real-weight model validation evidence
(see `MODEL_VALIDATION.md`) (`460b27a`).
## [2.0.0] - 2026-07-28
### Changed
- **Breaking:** node pack modernized with an explicit GPU lifecycle. Models
now load lazily on first execution and register with ComfyUI's model manager
so they participate in smart VRAM offloading, which addresses models
remaining resident after generation (#137) (`b89f628`).
- **Breaking:** `forceInput` string hacks removed from node schemas. They
corrupted the widget index during serialization and shifted inputs on saved
workflows (#156). Use the native right-click "Convert to Input" instead.
- Import is now failure-isolated: a broken optional model cannot prevent
unrelated nodes from loading (#94, #145).
- `numpy` is no longer pinned. The old `numpy<2.0.0` pin crashed startup on
NumPy 2.x environments (#157).
- Model coverage moved to current releases, including Qwen 3 / 3.5 VL,
SmolVLM2, InternVL, Granite Vision, and Gemma 3 (#148, #151). The
unmaintained InternLM-XComposer2 nodes were dropped (#139).
### Removed
- **Breaking:** `llama-cpp-agent` dependency (see 2.2.0).
- **Breaking:** InternLM-XComposer2 nodes, which depended on an AutoGPTQ stack
that pinned incompatible PyTorch versions (#139).
## 1.0.0 - 1.0.6 (2024-05-20 to 2024-11-03)
Initial packaged releases, predating changelog tracking. This line covered
LLaVA GGUF loaders and samplers, Moondream, Kosmos-2, JoyTag, UForm,
MiniCPM-V, PaLI-Gemma, Florence-2, Molmo, Qwen2-VL, the LLM prompt and
suggestion generators, AudioLDM2, and ChatMusician. See the
[commit history](https://github.com/gokayfem/ComfyUI_VLM_nodes/commits/main)
for detail.
Tagging began at 3.3.0. Earlier versions link to the commit that declared
them, because retroactively tagging them would run current CI against code
that predates it.
[3.4.0]: https://github.com/gokayfem/ComfyUI_VLM_nodes/compare/v3.3.1...v3.4.0
[3.3.1]: https://github.com/gokayfem/ComfyUI_VLM_nodes/compare/v3.3.0...v3.3.1
[3.3.0]: https://github.com/gokayfem/ComfyUI_VLM_nodes/releases/tag/v3.3.0
[3.2.0]: https://github.com/gokayfem/ComfyUI_VLM_nodes/commit/44fefcb
[3.1.0]: https://github.com/gokayfem/ComfyUI_VLM_nodes/commit/505b324
[3.0.0]: https://github.com/gokayfem/ComfyUI_VLM_nodes/commit/39fc116
[2.3.0]: https://github.com/gokayfem/ComfyUI_VLM_nodes/commit/239c904
[2.2.0]: https://github.com/gokayfem/ComfyUI_VLM_nodes/commit/0da5070
[2.1.0]: https://github.com/gokayfem/ComfyUI_VLM_nodes/commit/4c200c4
[2.0.1]: https://github.com/gokayfem/ComfyUI_VLM_nodes/commit/460b27a
[2.0.0]: https://github.com/gokayfem/ComfyUI_VLM_nodes/commit/b89f628
+21
View File
@@ -0,0 +1,21 @@
cff-version: 1.2.0
message: "If you use ComfyUI VLM Nodes in your work, please cite it using the metadata below."
type: software
title: "ComfyUI VLM Nodes"
version: "3.5.0"
date-released: 2026-07-31
authors:
- family-names: "Aydoğan"
given-names: "Gökay"
orcid: "https://orcid.org/0000-0002-2343-9433"
abstract: "Production-ready local and API vision-language, structured prompting, audio, and utility nodes for ComfyUI."
keywords:
- ComfyUI
- vision-language models
- multimodal AI
- image understanding
- video understanding
- generative AI
license: Apache-2.0
repository-code: "https://github.com/gokayfem/ComfyUI_VLM_nodes"
url: "https://github.com/gokayfem/ComfyUI_VLM_nodes"
+76
View File
@@ -49,6 +49,33 @@ cannot execute BF16. This is a portability fallback, not proof that every
model family has been run on every vendor device. See
[MODEL_VALIDATION.md](MODEL_VALIDATION.md) for real-hardware evidence.
### Moondream 3 / 3.1 Photon
Moondream Photon is deliberately isolated from ComfyUI's main Python environment
because `moondream==1.3.0` requires Pillow 10 while current ComfyUI uses a
newer Pillow. Its worker cache, virtual environment, and logs live under
`models/LLavacheckpoints/moondream31-runtime`; it never replaces ComfyUI's
PyTorch or Pillow.
| Platform | Official local Photon support | This integration |
| --- | --- | --- |
| Linux/WSL + NVIDIA Ampere or newer | Supported | 3.1 query/caption/detection/pointing; 3 Preview SVG segmentation |
| Windows + NVIDIA Ampere or newer | Supported | Same isolated worker contract |
| Apple Silicon macOS 13+ | Supported with MPS | Same contract; use a conservative KV-cache profile on low-memory systems |
| AMD ROCm, Intel GPU, CPU | Not currently provided upstream | Node stays importable and fails before model work with an actionable support message |
The final Moondream 3.1 model card lists query, caption, detect, and point; it
does not list segment. Native SVG segment uses `moondream3-preview`, and the
loader rejects a 3.1/segment mismatch before inference.
`max_batch_size` controls Photon's scheduler capacity. The detection, point,
and preview-segmentation nodes issue `parallel_requests` frame requests concurrently,
allowing Photon to build GPU batches. `frame_stride` bounds work for high-frame
rate sources. Performance JSON records warm worker time, end-to-end time,
processed/skipped frames, worker/sustained FPS, target sampled FPS, and
real-time factor; it is a measurement from the current run, not a universal
benchmark claim.
### Video memory and chunking
- Core `Video Slice` should bound work before `GetVideoComponents` materializes
@@ -81,6 +108,10 @@ model card before redistributing weights or outputs.
`ComfyUI/models/checkpoints`.
- `HF_TOKEN` is used when Hugging Face requires authenticated access. Tokens
must be supplied by the environment and must not be embedded in workflows.
- Moondream 3.1 uses the Moondream Model License 1.0. The Loader requires an
explicit workflow acknowledgement. The license permits local product use
but restricts offering general-purpose hosted Moondream access; review the
current upstream terms for the intended deployment.
Authoritative references:
@@ -88,6 +119,8 @@ Authoritative references:
- [Meta SAM3 license](https://huggingface.co/facebook/sam3/blob/main/LICENSE)
- [ComfyUI SAM3.1 checkpoint](https://huggingface.co/Comfy-Org/sam3.1)
- [SAM2.1 Hiera Tiny model card](https://huggingface.co/facebook/sam2.1-hiera-tiny)
- [Moondream 3.1 model card](https://huggingface.co/moondream/moondream3.1-9B-A2B)
- [Moondream Model License 1.0](https://moondream.ai/licenses/model/1.0)
## Dependency behavior
@@ -99,9 +132,43 @@ Authoritative references:
from blocking the whole node pack.
- `requirements-quantization.txt` is available for an explicit quantization
install or source-build environment.
- `requirements-moondream31.txt` belongs only in the isolated Photon sidecar;
installing it into ComfyUI's environment would create a Pillow conflict.
- Model downloads, imports, and package compilation never occur during node
discovery.
## Robotics / VLA policy compatibility
ComfyUI's robotics schemas, safety gate, trajectory tools, and universal HTTP
client run wherever this node pack runs. Policy runtime compatibility is
separate:
| Policy route | ComfyUI client | Policy environment | Practical boundary |
| --- | --- | --- | --- |
| Universal VLA HTTP | Windows, Linux, macOS; CUDA, ROCm, Metal, XPU, CPU | Any host that implements `comfyui-vla-http-v1` | Loopback HTTP or trusted HTTPS; no pickle |
| LeRobot sidecar | Same universal client | Current LeRobot supports Linux, Windows, and macOS; individual policy extras/operators vary | Python/PyTorch live outside ComfyUI; fine-tuned checkpoint required for the target embodiment |
| openpi WebSocket | Lightweight optional client on every ComfyUI platform | Upstream currently tests Ubuntu 22.04 + NVIDIA, inference above 8 GB VRAM | Use WSL/Docker/Linux server; remote transport must be WSS |
| Isaac-GR00T N1.7 ZMQ | Lightweight optional client on every ComfyUI platform | NVIDIA CUDA/Jetson Linux according to upstream deployment matrix | ZMQ has no transport encryption; use a private network/tunnel |
| OpenVLA-OFT | Universal client with a project-specific bridge | Upstream PyTorch/CUDA environment | OFT is the preferred high-frequency multi-image OpenVLA route |
| Octo | Universal client with a project-specific bridge | Isolated JAX environment | Kept as a lightweight research baseline, not the default maintained runtime |
Install only the native client protocols into ComfyUI:
```bash
python -m pip install -r requirements-robotics-client.txt
```
Do not install `lerobot[all]`, openpi, Isaac-GR00T, OpenVLA, or JAX into
ComfyUI's Python. The included LeRobot HTTP sidecar belongs in its own
environment and optionally moves its owned policy to CPU after an idle
interval. It does not flush ComfyUI's accelerator cache.
An embodiment profile is a workflow contract, not a hardware certification.
The supplied profiles are visibly labeled templates. Before real deployment,
replace action bounds/deltas with the trained dataset's semantics and the
manufacturer/controller limits. ComfyUI never opens ROS, serial, CAN, or robot
SDK transports.
Install manually:
```bash
@@ -138,6 +205,15 @@ python -m pip install llama-cpp-python \
--extra-index-url https://abetlen.github.io/llama-cpp-python/whl/vulkan
```
If an Apple Metal wheel is unavailable or fails archive validation, build the
same optional requirement from source:
```bash
CMAKE_ARGS="-DGGML_METAL=on" python -m pip install \
--no-cache-dir --no-binary llama-cpp-python \
-r ComfyUI/custom_nodes/ComfyUI_VLM_nodes/requirements-llama-cpp.txt
```
The official Windows HIP Radeon index is:
```powershell
+106
View File
@@ -0,0 +1,106 @@
# Contributing
Thanks for helping out. This pack runs inside other people's ComfyUI installs
on five accelerator backends, so a few rules exist to keep it from breaking
them.
## The rules that matter most
**Importing the pack must never download a model, install a package, compile
anything, or allocate VRAM.** Models load on first execution. This is enforced
by `tests/test_nodes.py`, which asserts the source contains no `pip install`,
no `subprocess.run`, and no direct `torch.cuda.empty_cache`.
**Never reorder or insert widgets in an existing node's `INPUT_TYPES`.** Comfy
serializes widget values by position, so a reordered schema silently rebinds
every saved workflow. Add new inputs to `optional` at the end. The widget order
of the long-lived nodes is pinned by tests; if a test fails because you moved a
widget, the test is right.
**Never use `forceInput`.** It corrupts the widget index during serialization.
Users get the same result from the native right-click "Convert to Input".
**An optional dependency must fail only the node that needs it.** Use
`require_module()` from `nodes/runtime.py`, which raises an actionable error at
execution time rather than at import time.
**Do not install or replace `torch`.** ComfyUI's own installer picks the CUDA,
ROCm, XPU, Metal, or CPU build. The same applies to `numpy` and `Pillow`.
## Setting up
```bash
cd ComfyUI/custom_nodes
git clone https://github.com/gokayfem/ComfyUI_VLM_nodes.git
cd ComfyUI_VLM_nodes
python -m pip install -r requirements.txt -r requirements-dev.txt
```
Use ComfyUI's Python. On ComfyUI Portable there is no `activate` script, so
call the interpreter directly:
```
..\..\python_embeded\python.exe -m pip install -r requirements.txt
```
## Running checks
```bash
PYTHONPATH=/path/to/custom_nodes:/path/to/ComfyUI python -m pytest -q
python -m ruff check .
```
`PYTHONPATH` needs the directory *containing* this checkout plus ComfyUI
itself, because the tests import `ComfyUI_VLM_nodes` as a package and the nodes
import ComfyUI's `folder_paths`.
CI additionally enforces a coverage floor on Linux/Python 3.13:
```bash
python -m pytest -q --cov=nodes --cov-fail-under=70
```
Real-weight tests are opt-in because they download multi-gigabyte checkpoints,
and are never run in CI:
```bash
python tests/manual_model_smoke.py --model "Qwen 3 VL 4B Instruct"
python tests/manual_specialized_smoke.py --backend florence-large
python tests/manual_llama_cpp_smoke.py --download
```
## Writing tests
Tests must pass without model weights, without a GPU, and without
`llama-cpp-python`. Stub the model boundary instead: see
`tests/test_suggest.py` and `tests/test_llavaloader.py` for the pattern of
faking `LlamaHandle` and `create_chat_completion` to assert what the node sends
to the backend.
`nodes/joytagger/` is vendored upstream code kept byte-compatible with its
source. It is excluded from lint; please don't reformat it.
## Adding a model
1. Prefer adding an entry to the catalog in `nodes/modern_vlm.py` over a new
node. Most current VLMs work through the shared `transformers` path.
2. If it needs a bespoke loader, follow `nodes/minicpm.py` as the smallest
complete example.
3. Register the module in the `node_list` in `__init__.py`.
4. Record what you actually ran in `MODEL_VALIDATION.md`. Catalog entries that
were never executed against real weights must be marked as such.
5. Add the node to the reference table in `README.md`.
## Releasing
The Comfy Registry publishes from `pyproject.toml`, and only when `version`
changes. A fix merged without a version bump never reaches Registry users. So:
- bump `version` in `pyproject.toml`,
- add a `CHANGELOG.md` entry,
- tag the merge commit `vX.Y.Z`.
## Commit messages
Short imperative subject, one logical change per commit. Reference the issue it
closes in the body.
+65
View File
@@ -21,6 +21,71 @@ One checkpoint covers sibling sizes that use the same architecture and loader.
The node does not download every size simply to repeat the same integration
test.
## Robotics VLA pass
Validated on 2026-07-31 through the included isolated LeRobot HTTP policy
server, entirely from WSL and D-drive storage:
- Runtime: Python 3.12.12, LeRobot 0.6.1 from current upstream source,
PyTorch 2.11.0+cu128, and an NVIDIA RTX 3090.
- Checkpoint: `lerobot/smolvla_base` (about 2.5 GiB of D-drive cache), backed
by `HuggingFaceTB/SmolVLM2-500M-Video-Instruct`.
- Real input: local `image (23).png`, a 256x256 outdoor photograph, repeated
across the checkpoint's three declared camera keys with a six-value state
vector and the task “Move the end effector toward the backpack and prepare
to grasp it.”
- Contract: three camera tensors, `observation.state`, the LeRobot
preprocessor, `predict_action_chunk`, the checkpoint postprocessor, bounded
JSON/JPEG transport, action parsing, and the ComfyUI safety layer all ran.
The native checkpoint advertises a 50-step chunk; the server returned four
steps of six actions for this test.
- Five warm requests after one discarded warm-up measured 241.374 ms mean
server inference (242.957 ms median, 234.726–249.980 ms range) and
270.404 ms mean HTTP client time (271.483 ms median,
261.477–282.218 ms range).
- The final raw action chunk was:
```json
[
[0.06258623, -0.11250310, -0.13713294, -0.06168950, -0.00926633, -0.08506130],
[0.15420279, -0.05678255, -0.20159233, 0.06734322, -0.00563951, -0.09575561],
[0.16482556, -0.07453565, -0.17410603, 0.02461835, -0.00256573, 0.15842065],
[0.27048433, -0.09272483, -0.19934477, 0.05491992, 0.05286619, 0.07594281]
]
```
Applying the SO-100/SO-101 template limits from an all-zero previous action
found five per-step delta violations, no bounds violations, and no
non-finite values. `Clamp safely` produced:
```json
[
[0.06258623, -0.1, -0.1, -0.06168950, -0.00926633, -0.08506130],
[0.15420279, -0.05678255, -0.2, 0.03831051, -0.00563951, -0.09575561],
[0.16482556, -0.07453565, -0.17410603, 0.02461835, -0.00256573, 0.05424440],
[0.26482555, -0.09272483, -0.19934477, 0.05491992, 0.05286619, 0.07594281]
]
```
This is an end-to-end loading, preprocessing, inference, transport, parsing,
and safety-contract pass. It is not evidence that a base SmolVLA checkpoint can
control an SO-100 from an arbitrary Internet-style photograph. Actual robot
deployment still requires embodiment-matched fine-tuning, calibrated state and
camera inputs, hardware-certified limits, a deadman/watchdog, collision
handling, and an external emergency stop.
The same real checkpoint was then exercised through ComfyUI's actual local
`POST /prompt` API, not by calling the Python node directly. The graph loaded
and center-cropped the real image to 256x256, constructed the three-camera
checkpoint contract, called the isolated GPU policy, applied the SO-100/SO-101
template safety gate, rendered a 960x480 trajectory preview, and emitted all
three text reports. Final prompt
`86b31f5a-5a8c-4abc-abdb-634e46da5c93` completed successfully: the policy
returned `[4, 6]` actions, the safety gate found three rate violations and
clamped them, `safe_for_handoff` was true under the declared template, and
ComfyUI wrote preview `ComfyUI_temp_icynu_00001_.png`. The reusable acceptance
harness is `tests/manual_robotics_smoke.py`.
## ComfyUI API pass
ComfyUI started from the D-drive WSL installation with all four repaired custom
+571 -20
View File
@@ -1,7 +1,7 @@
# ComfyUI VLM Nodes
Production-oriented vision-language, structured prompting, audio, and utility
nodes for ComfyUI. Version 2.3 supports ComfyUI's selected NVIDIA CUDA, AMD
nodes for ComfyUI. Version 3.4 supports ComfyUI's selected NVIDIA CUDA, AMD
ROCm, Apple Metal, Intel XPU, and CPU device without replacing its PyTorch
build. It removes startup installers and global accelerator cache flushes,
adds real image/video batches and live token streaming, and uses ComfyUI model
@@ -9,19 +9,40 @@ residency and offloading.
## Modern model coverage
The **Modern VLM** node provides one stable interface for:
The **Modern VLM** node provides one stable interface with a deliberately
small, 12-choice production picker:
- Qwen 3.5 0.8B, 2B, 4B, 9B, 27B, and 35B-A3B
- Qwen 3.6 27B
- Qwen 3 VL 2B, 4B, 8B, and 30B-A3B Instruct
- Qwen 2.5 VL 3B and 7B for existing workflows
- Gemma 3 4B, 12B, and 27B IT
- SmolVLM2 256M, 500M, and 2.2B video models
- Liquid LFM2.5-VL 450M and 1.6B edge models
- InternVL 3.5 1B and 2B standard Hugging Face checkpoints
- Granite Vision 3.3 2B and 4.1 4B for documents, charts, and OCR
- Qwen 3.5 0.8B and 4B
- Qwen 3 VL 2B, 4B, and 8B Instruct
- SmolVLM2 500M and 2.2B Video
- Liquid LFM2.5-VL 450M
- InternVL 3.5 1B
- Granite Vision 4.1 4B
- Gemma 3 4B IT
- a compatible custom Hugging Face image-to-text repository
The separate **[Legacy] Modern VLM Compatibility** node contains redundant,
superseded, experimental, and very large tiers:
- Qwen 3.5 2B, 9B, 27B, and 35B-A3B
- Qwen 3.6 27B
- Qwen 3 VL 30B-A3B Instruct
- Qwen 2.5 VL 3B and 7B for existing workflows
- Gemma 3 12B and 27B IT
- SmolVLM2 256M Video
- Liquid LFM2.5-VL 1.6B
- InternVL 3.5 2B
- Granite Vision 3.3 2B
Previously saved `ModernVLM` workflows remain valid even when their selected
model moved to Legacy. The server accepts every known catalog value for
backward compatibility; only the visible new-workflow picker is curated.
Dedicated Molmo, PaLI-Gemma, Qwen2-VL, MiniCPM-V, Kosmos-2, MC-LLaVA, UForm,
and script-style MoonDream nodes are also collected under
`VLM Nodes/Legacy/Model Loaders`. Maintained creator-facing Florence-2,
Moondream2, JoyTag, llama.cpp/GGUF, detection, segmentation, tracking, API,
and video-intelligence nodes stay in their functional categories.
Sixteen curated sub-4B/low-VRAM choices are marked internally as the
small-and-fast tier. The default is Qwen 3 VL 2B: it is much quicker to load
than larger checkpoints while retaining broad image and video understanding.
@@ -41,18 +62,54 @@ when ComfyUI rehydrates workflow output history. Disable `stream_output` for
API-only or headless runs that do not need incremental UI updates. Streaming is
best-effort and never changes the final `STRING` output or makes inference fail.
## Text workflow toolkit
The original `SimpleText`, `JsonToText`, and `ViewText` node IDs and their
first `STRING` outputs remain stable for saved workflows. They now live in
organized `VLM Nodes/Text` subcategories and expose descriptive names, search
aliases, tooltips, appended metrics, and strict error messages:
| Node | Purpose |
| --- | --- |
| `Text` (`SimpleText`) | Multiline/dynamic prompt source with optional edge/newline normalization and character, word, and line outputs |
| `View Text (Streaming)` | Read-only live output with counts, copy, UTF-8 download, line wrapping, stream following, reroute traversal, and history rehydration |
| `JSON to Text` | Plain or fenced JSON parsing with readable, values-only, key/value, pretty, and compact render modes |
| `Text Join` | Join up to eight prompt/context values with empty-value removal and stable deduplication |
| `Text Template` | Safe named placeholders from a JSON object plus four convenient live text sockets, with explicit missing-key policy |
| `Text Clean` | Unicode NFC/NFKC, newline/whitespace cleanup, enclosing Markdown-fence removal, line deduplication, and deterministic length caps |
| `Text Replace` | Literal or regex substitution with case, count, and missing-pattern controls |
| `JSON Extract` | JSONPath-lite (`$.items[0]`) and RFC 6901 JSON Pointer extraction from plain or fenced model responses |
| `Text Split / Batch` | Lines, paragraphs, delimiters, regex, CSV, or JSON arrays converted to a real mapped Comfy `STRING` list |
| `Text Inspector` | Pass-through text plus characters, UTF-8 bytes, words, lines, rough token budget, SHA-256, and JSON metadata |
The JSON utilities never evaluate code, follow references, access files, or
make network requests. Template fields are direct names rather than Python
attribute/index expressions. `approx_tokens` is deliberately labeled as a
rough UTF-8 budget estimate; use the target model tokenizer when exact billing
or context accounting matters.
Specialized nodes remain available where a generic chat node would discard
useful model capabilities:
- **Moondream 3.1 9B-A2B**: official 2B-active Photon runtime with query,
caption, and high-throughput image/video detection and pointing.
- **Moondream 3 Preview segment**: native SVG segmentation through the same
isolated Photon loader. The SVG is preserved and also converted into antialiased
`MASK`, black/white previews, foreground cutouts, overlays, polygons,
canonical `VLM_DETECTIONS`, and core bounding boxes. Detection/pointing
submit frames concurrently so Photon can dynamically batch them; every run
reports measured worker FPS, end-to-end FPS, and real-time factor.
- **Florence-2**: captioning, OCR, detection, region captioning, and referring
expression segmentation, with structured JSON, mask, and overlay outputs.
- **PaLI-Gemma**: caption/VQA plus the official 16-token VQ-VAE segmentation
decoder; segmentation tokens are no longer misinterpreted as polygon points.
- **Moondream2**: pinned query API with explicit decoding controls. Its current
checkpoint is not marked passed on the tested Torch/Transformers stack; use a
small Modern VLM preset for production.
- **Moondream2**: pinned query API with explicit decoding controls. The official
checkpoint is loaded through its native safetensors state dict, avoiding the
silent empty-output regression in Transformers 5 while retaining ComfyUI
managed loading and unloading.
- **Qwen2-VL**: image batches and real video-frame batches.
- **Molmo, Kosmos-2, UForm, MCLLaVA, JoyTag, and MiniCPM-V 2.6 GGUF**.
- **Legacy Molmo, Kosmos-2, UForm, MCLLaVA, and MiniCPM-V 2.6 GGUF**, plus
maintained JoyTag.
- **llama.cpp LLaVA/GGUF**, structured prompt suggestions, OpenAI-compatible
prompting, and AudioLDM2.
@@ -67,6 +124,8 @@ lists between nodes:
| `VLM_TRACKS` | `comfyui-vlm/tracks`, version 1 | Durable object IDs with ordered observations over time |
| `VLM_POINTS` | `comfyui-vlm/points`, version 1 | Pixel-coordinate points, including detection centers |
| `VLM_EVENTS` | `comfyui-vlm/events`, version 1 | Ordered temporal events for downstream video analysis |
| `VLM_VIDEO_SELECTION` | `comfyui-vlm/video-selection`, version 1 | Exact mapping from sampled images to source frame indices and timestamps |
| `VLM_SCENE_STATE` | `comfyui-vlm/scene-state`, version 1 | Compact persistent objects, motion, visibility, and validated events |
All spatial coordinates are source-image pixels. Bounding boxes are
`[x1, y1, x2, y2]` with an exclusive right/bottom edge; polygons contain at
@@ -100,6 +159,78 @@ The utility layer converts without model-specific glue:
- `VLMDetectionsFromJSON` and `VLMDetectionsToJSON` are the explicit API and
persistence boundary for the versioned detection schema.
### Universal VLM performance utilities
The performance nodes sit before any local or hosted VLM, so their savings do
not depend on CUDA, ROCm, MPS, XPU, CPU, Transformers, llama.cpp, or Photon:
- `VLM Performance Profile` emits coherent `max_frames`, pixel budget,
longest-edge, batch-size, and `unload_after` values. `Live / robotics`,
`Fast video`, `Balanced`, `High detail`, and `Low VRAM handoff` are explicit
starting points rather than hidden global flags.
- `VLM Adaptive Frame Sampler` is the existing track-aware temporal gate. It
combines uniform coverage, scene changes, motion, and optional track changes
while preserving source frame indices and timestamps.
- `VLM Image Pixel Budget` downsizes the selected analysis copy once, preserves
aspect ratio, never upscales, and can align dimensions to 14/28-pixel VLM
patches or 32-pixel detector backbones. Fast area and antialiased bicubic
modes are available.
The recommended order is `Video Slice` → `VLM Adaptive Frame Sampler` →
`VLM Image Pixel Budget` → any VLM. A model's own official processor still
performs its required normalization/crop; the pixel-budget node simply prevents
every downstream model from repeatedly receiving unnecessary source pixels.
Local torch models remain registered with ComfyUI's smart model manager, while
external allocators reserve space before loading and close only the handle they
own.
On the real `vlm_api_people_birds.mp4` input in this repository's D-drive test
environment, the utilities selected 10 of 60 1280×720 frames and resized them
to 938×518 in about 0.44 seconds on a cold WSL run. That reduced the
frame×pixel analysis workload by 11.38× before model inference. This is an
input-work reduction measurement, not a claim that every model runs 11.38×
faster; token generation and model-specific vision encoders still determine
end-to-end speed.
### Adaptive video intelligence
The video-intelligence layer keeps generative VLM inference out of the
per-frame loop:
- `VLMAdaptiveFrameSampler` combines scene-change, motion, track-change, and
uniform-coverage signals. It always preserves the real source frame index
and timestamp, enforces a frame budget, and returns selection/diagnostic
JSON. `Uniform coverage`, motion, scene, and track-priority modes remain
available for deterministic experiments.
- `VLMVideoTemporalReasoner` is the one-node path. It adaptively samples the
input, downsizes only the VLM analysis copy (448-pixel longest side by
default), runs a recommended video-capable model, parses the result into
validated `VLM_EVENTS`, and returns summary, events, selection, sampled
previews, raw response, diagnostics, event JSON, and selection JSON.
- `VLMVideoReasoningPrompt` and `VLMEventsFromVideoJSON` expose the same strict
timestamp/evidence contract for custom local or hosted VLM workflows.
- `VLMTrackAwareCrops` chooses representative observations for each durable
track, adds configurable context, and letterboxes crops to one batch size.
This lets a VLM label identities without rereading every full frame.
- `VLMBuildSceneState` converts tracks plus optional events into a compact
persistent world-state summary with first/last observation, current box,
confidence, state, and pixel velocity.
Small VLMs commonly return evidence as positions in the supplied image batch
even when asked for source indices. The parser accepts that form only when
every value is an unambiguous valid supplied-image position, maps it back to
the immutable source selection, and records the normalization mode. Arbitrary
or unsupplied evidence frames, out-of-range timestamps, invalid confidence,
duplicate evidence, malformed JSON, and non-finite values fail validation.
On the repository's real-data smoke test (RTX 3090, Qwen3-VL 2B, 157-frame
896x448 H.264 clip), hybrid sampling selected 12 frames in 0.30 seconds,
reduced temporal inputs by 92.36%, reduced analysis pixels by 75%, used
4.24 GiB peak allocated VRAM in the standalone runner, and produced a valid
timestamped result in 35.17 seconds. The equivalent live ComfyUI `/prompt`
graph completed in 37.45 seconds. These are one-machine measurements, not
portable performance guarantees.
### Open-vocabulary image and video detection
`VLMOpenVocabularyDetection` exposes one interface for:
@@ -200,13 +331,240 @@ image visualization. Region tasks reject ambiguous multi-box input; use
API-format examples are in [`examples/vision`](examples/vision):
- [`grounding_dino_image_api.json`](examples/vision/grounding_dino_image_api.json)
- [`moondream3_preview_svg_segment_api.json`](examples/vision/moondream3_preview_svg_segment_api.json)
- [`moondream31_video_detect_api.json`](examples/vision/moondream31_video_detect_api.json)
- [`sam2_video_tracking_api.json`](examples/vision/sam2_video_tracking_api.json)
- [`sam3_core_adapter_blueprint_api.json`](examples/vision/sam3_core_adapter_blueprint_api.json)
- [`video_temporal_reasoning_api.json`](examples/vision/video_temporal_reasoning_api.json)
- [`vlm_performance_preflight_api.json`](examples/vision/vlm_performance_preflight_api.json)
The dependency-free text-toolkit example is
[`examples/text_toolkit_api.json`](examples/text_toolkit_api.json).
Robotics policy, safety, and sidecar examples are in
[`examples/robotics`](examples/robotics), including a complete universal
HTTP policy graph.
Upload the named media to ComfyUI's input directory, adjust the filenames and
labels, then submit the JSON object as the `prompt` value to `/prompt`. These
are API graphs, not frontend workflow-export JSON.
## Node reference
All 89 registered nodes, grouped by their menu category. The **Node ID** is the
`class_type` written into workflow and API JSON — search for that string when
you need to find a node you saw on a canvas.
### Modern VLM
The main entry point for current vision-language models.
| Node | Node ID | Outputs |
| --- | --- | --- |
| Modern VLM (Qwen / SmolVLM2 / LFM / InternVL / Granite / Gemma) | `ModernVLM` | `STRING` |
| Moondream 2 | `Moondream2model` | `STRING` |
### Moondream 3
Moondream 3 / 3.1 in an isolated Photon runtime. Load once, then reuse the
`MOONDREAM31_MODEL` output across the task nodes.
| Node | Node ID | Outputs |
| --- | --- | --- |
| Moondream 3 / 3.1 Loader (Isolated Photon) | `Moondream31Loader` | `MOONDREAM31_MODEL`, `STRING` |
| Moondream 3 / 3.1 Caption | `Moondream31Caption` | `STRING`, `STRING` |
| Moondream 3 / 3.1 Query | `Moondream31Query` | `STRING`, `STRING`, `STRING` |
| Moondream 3 / 3.1 Detect (Image / Video) | `Moondream31Detect` | `VLM_DETECTIONS`, `STRING`, `IMAGE`, `MASK`, `BOUNDING_BOX`, `BOUNDING_BOXES`, `STRING` |
| Moondream 3 / 3.1 Point (Image / Video) | `Moondream31Point` | `VLM_POINTS`, `STRING`, `IMAGE`, `STRING` |
| Moondream 3 Preview SVG Segment (Image / Video) | `Moondream31Segment` | `VLM_DETECTIONS`, `STRING`, `STRING`, `MASK`, `IMAGE`, `IMAGE`, `IMAGE`, `BOUNDING_BOX`, `BOUNDING_BOXES`, `STRING` |
### Florence-2
| Node | Node ID | Outputs |
| --- | --- | --- |
| Florence-2 Multitask Vision | `Florence2` | `STRING`, `STRING`, `MASK`, `IMAGE` |
### Vision: detection, segmentation, tracking
Open-vocabulary detection and video segmentation. These emit the structured
`VLM_DETECTIONS` / `VLM_POINTS` / `VLM_TRACKS` types rather than loose strings.
| Node | Node ID | Outputs |
| --- | --- | --- |
| VLM Open-Vocabulary Detection | `VLMOpenVocabularyDetection` | `VLM_DETECTIONS`, `STRING`, `IMAGE`, `MASK`, `BOUNDING_BOX`, `BOUNDING_BOXES` |
| VLM SAM2.1 Video Segmentation | `VLMSAM2VideoSegmentation` | `VLM_TRACKS`, `STRING`, `MASK`, `MASK`, `IMAGE` |
| VLM SAM3 Track Adapter | `VLMSAM3TrackAdapter` | `VLM_TRACKS`, `SAM3_TRACK_DATA` |
| VLM Track Detections | `VLMTrackDetections` | `VLM_TRACKS` |
| VLM Track Report | `VLMTrackReport` | `STRING`, `STRING` |
| JoyTag | `Joytag` | `STRING` |
### Vision: spatial reasoning
| Node | Node ID | Outputs |
| --- | --- | --- |
| VLM Spatial Prompt Builder | `VLMSpatialPromptBuilder` | `STRING` |
| VLM Structured Spatial Parser | `VLMStructuredSpatialParser` | `VLM_DETECTIONS`, `VLM_POINTS`, `STRING` |
### Vision: detection utilities
Converters and filters between structured detections and ordinary Comfy types.
| Node | Node ID | Outputs |
| --- | --- | --- |
| Filter VLM Detections | `VLMFilterDetections` | `VLM_DETECTIONS` |
| Select VLM Detection | `VLMSelectDetection` | `VLM_DETECTIONS` |
| Crop VLM Detections | `VLMCropDetections` | `IMAGE`, `STRING` |
| Render VLM Detections | `VLMRenderDetections` | `IMAGE` |
| VLM Detection Centers | `VLMDetectionsToPoints` | `VLM_POINTS`, `STRING` |
| VLM Detections from JSON | `VLMDetectionsFromJSON` | `VLM_DETECTIONS` |
| VLM Detections to JSON | `VLMDetectionsToJSON` | `STRING` |
| VLM Detections to Bounding Boxes | `VLMDetectionsToBoundingBoxes` | `BOUNDING_BOXES`, `STRING` |
| VLM Detections to Masks | `VLMDetectionsToMasks` | `MASK`, `MASK`, `STRING`, `MASK`, `IMAGE`, `IMAGE`, `IMAGE` |
### Vision: mask tools
| Node | Node ID | Outputs |
| --- | --- | --- |
| VLM Mask Processor | `VLMMaskProcessor` | `MASK`, `MASK`, `MASK`, `IMAGE` |
| VLM Mask Composite | `VLMMaskComposite` | `IMAGE`, `IMAGE`, `IMAGE`, `IMAGE` |
### Video intelligence
Adaptive frame selection and temporal reasoning for long videos.
| Node | Node ID | Outputs |
| --- | --- | --- |
| VLM Adaptive Frame Sampler | `VLMAdaptiveFrameSampler` | `IMAGE`, `VLM_VIDEO_SELECTION`, `STRING`, `STRING` |
| VLM Video Reasoning Prompt | `VLMVideoReasoningPrompt` | `STRING`, `STRING` |
| VLM Video Temporal Reasoner | `VLMVideoTemporalReasoner` | `STRING`, `VLM_EVENTS`, `VLM_VIDEO_SELECTION`, `IMAGE`, `STRING`, `STRING`, `STRING`, `STRING` |
| VLM Temporal Events From JSON | `VLMEventsFromVideoJSON` | `VLM_EVENTS`, `STRING`, `STRING` |
| VLM Persistent Scene State | `VLMBuildSceneState` | `VLM_SCENE_STATE`, `STRING`, `STRING` |
| VLM Track-Aware Semantic Crops | `VLMTrackAwareCrops` | `IMAGE`, `STRING` |
### LLM (local GGUF)
llama.cpp text models. `LLM Loader (GGUF)` produces the `CUSTOM` model handle
the samplers consume; the *Managed Cache* variants own their own handle and can
release it after each run.
| Node | Node ID | Outputs |
| --- | --- | --- |
| LLM Loader (GGUF) | `LLMLoader` | `CUSTOM` |
| LLM Sampler | `LLMSampler` | `STRING` |
| LLM Prompt Generator | `LLMPromptGenerator` | `STRING` |
| LLM (Managed Cache) | `LLMOptionalMemoryFreeSimple` | `STRING` |
| LLM (Managed Cache, Advanced) | `LLMOptionalMemoryFreeAdvanced` | `STRING` |
| Structured Output | `StructuredOutput` | `STRING` |
| Structured Keyword Extraction | `KeywordExtraction` | `STRING` |
| Structured Prompt Generator | `LLavaPromptGenerator` | `STRING` |
| Creative Art Prompt Generator | `CreativeArtPromptGenerator` | `STRING` |
| Prompt Suggester | `Suggester` | `STRING` |
### LLaVA (local GGUF multimodal)
Vision models through llama.cpp. These need both a GGUF and its vision
projector (mmproj).
| Node | Node ID | Outputs |
| --- | --- | --- |
| LLaVA Loader | `LLava Loader Simple` | `CUSTOM` |
| LLaVA Vision Projector Loader | `LlavaClipLoader` | `CUSTOM` |
| LLaVA Sampler | `LLavaSamplerSimple` | `STRING` |
| LLaVA Sampler (Advanced) | `LLavaSamplerAdvanced` | `STRING` |
| LLaVA (Managed Cache) | `LLavaOptionalMemoryFreeSimple` | `STRING` |
| LLaVA (Managed Cache, Advanced) | `LLavaOptionalMemoryFreeAdvanced` | `STRING` |
### Hosted APIs
| Node | Node ID | Outputs |
| --- | --- | --- |
| Hosted VLM API (Secure) | `HostedVLMAPI` | `STRING`, `STRING`, `INT` |
| Hosted LLM API (Secure) | `PromptGenerateAPI` | `STRING` |
### Robotics / VLA policies
These nodes build and inspect policy observations/actions. They never send
commands to robot hardware. Heavy policy runtimes stay in isolated LeRobot,
openpi, GR00T, OpenVLA/OFT, or JAX environments.
| Node | Node ID | Outputs |
| --- | --- | --- |
| VLA Embodiment Profile | `VLAEmbodimentProfile` | `VLA_EMBODIMENT`, `STRING`, `INT`, `INT` |
| VLA Observation Builder | `VLAObservationBuilder` | `VLA_OBSERVATION`, `STRING`, `INT` |
| VLA Policy — Universal HTTP | `VLAHTTPPolicy` | `VLA_ACTIONS`, `STRING` |
| VLA Policy — OpenPI WebSocket | `VLAOpenPIWebSocketPolicy` | `VLA_ACTIONS`, `STRING` |
| VLA Policy — GR00T N1.7 ZMQ | `VLAGr00tZMQPolicy` | `VLA_ACTIONS`, `STRING` |
| VLA Action Safety Gate | `VLAActionSafety` | `VLA_ACTIONS`, `STRING`, `BOOLEAN` |
| VLA Actions From JSON | `VLAActionsFromJSON` | `VLA_ACTIONS`, `STRING` |
| VLA Action Chunk Replan | `VLAActionChunkReplan` | `VLA_ACTIONS`, `STRING` |
| VLA Action Inspect | `VLAActionInspect` | `STRING`, `STRING`, `INT`, `INT` |
| VLA Trajectory Preview | `VLATrajectoryPreview` | `IMAGE` |
| VLA Model Catalog | `VLAModelCatalog` | `STRING`, `STRING`, `STRING`, `STRING` |
### Text toolkit
Dependency-free string handling, so a VLM response can be shaped without an
extra node pack.
| Node | Node ID | Outputs |
| --- | --- | --- |
| Text | `SimpleText` | `STRING`, `INT`, `INT`, `INT` |
| Text Join | `VLMTextJoin` | `STRING`, `STRING`, `INT` |
| Text Template | `VLMTextTemplate` | `STRING`, `STRING`, `STRING` |
| Text Clean | `VLMTextClean` | `STRING`, `STRING` |
| Text Replace | `VLMTextReplace` | `STRING`, `INT`, `STRING` |
| Text Split / Batch | `VLMTextSplit` | `STRING`, `STRING`, `INT` |
| Text Inspector | `VLMTextInspect` | `STRING`, `INT`, `INT`, `INT`, `INT`, `INT`, `STRING`, `STRING` |
| View Text (Streaming) | `ViewText` | `STRING`, `INT`, `INT`, `INT`, `STRING` |
| JSON Extract | `VLMJSONExtract` | `STRING`, `BOOLEAN`, `STRING`, `STRING` |
| JSON to Text | `JsonToText` | `STRING`, `STRING`, `INT` |
### Performance and diagnostics
Run **VLM Runtime Diagnostics** before reporting a bug — it reports your
device, backend, and which optional packages are installed.
| Node | Node ID | Outputs |
| --- | --- | --- |
| VLM Runtime Diagnostics | `VLMRuntimeDiagnostics` | `STRING` |
| VLM Performance Profile | `VLMPerformanceProfile` | `INT`, `FLOAT`, `INT`, `INT`, `BOOLEAN`, `STRING` |
| VLM Image Pixel Budget | `VLMImagePixelBudget` | `IMAGE`, `INT`, `INT`, `STRING` |
### Audio
| Node | Node ID | Outputs |
| --- | --- | --- |
| AudioLDM2 | `AudioLDM2Node` | `*`, `INT`, `AUDIO` |
| Chat Musician | `ChatMusician` | `STRING`, `*`, `INT`, `AUDIO` |
| MiniMax Music | `MiniMaxMusicNode` | `*`, `INT`, `AUDIO` |
| PlayMusic Node | `PlayMusic` | `*` |
| Save Audio | `SaveAudioNode` | — |
MiniMax Music reads `MINIMAX_API_KEY` only from the ComfyUI server
environment. It uses fixed `global_en` and `cn_zh` endpoints, supports music
generation and cover models, decodes URL or hexadecimal responses, and emits
MP3, WAV, or PCM results through the existing waveform and `AUDIO` sockets.
The `aigc_watermark` field is sent only for `cn_zh` requests. See the official
[global](https://platform.minimax.io/docs/api-reference/music-generation) or
[China](https://platform.minimaxi.com/docs/api-reference/music-generation)
music API reference for account and content requirements.
### Legacy model loaders
Kept for existing workflows. New graphs should prefer **Modern VLM**, which
covers most of these architectures through one interface.
| Node | Node ID | Outputs |
| --- | --- | --- |
| Qwen2-VL | `Qwen2VLNode` | `STRING` |
| MiniCPM-V 2.6 (GGUF) | `MiniCPMNode` | `STRING` |
| Molmo Vision-Language Model | `MolmoNode` | `STRING` |
| PaLI-Gemma (Official Segmentation) | `Paligemma` | `STRING`, `MASK`, `IMAGE` |
| Kosmos-2 | `Kosmos2model` | `STRING` |
| MC-LLaVA | `MCLLaVAModel` | `STRING` |
| UForm Gen2 Qwen | `UformGen2QwenNode` | `STRING` |
| MoonDream (Moondream 2) | `MoonDream` | `STRING` |
| [Legacy] Modern VLM Compatibility | `LegacyModernVLM` | `STRING` |
## Install
Install through ComfyUI Manager, or clone into `ComfyUI/custom_nodes` and run:
@@ -221,6 +579,80 @@ Current official bitsandbytes wheels are installed automatically only on their
supported OS/architecture combinations. Unsupported machines retain all
non-quantized nodes.
### Robotics / VLA isolated runtimes
The robotics nodes keep policy dependencies outside ComfyUI. The universal
HTTP client works without another package. Native openpi WebSocket and
GR00T ZeroMQ clients use the lightweight optional extra:
```bash
python -m pip install \
-r ComfyUI/custom_nodes/ComfyUI_VLM_nodes/requirements-robotics-client.txt
```
`VLA Model Catalog` covers current SmolVLA, X-VLA, π0/π0-FAST/π0.5,
GR00T N1.7, WALL-OSS, MolmoAct2, VLA-JEPA, LingBot-VA, FastWAM, EO-1,
EVO-1, OpenVLA-OFT, and Octo routes. “Available” means a supported isolated
runtime/checkpoint path; base and architecture-only entries still require
embodiment-specific training and transforms.
Start with SmolVLA for small consumer hardware. The included authenticated
LeRobot sidecar loads one chosen policy, uses its serialized processors,
returns action chunks over bounded JSON/JPEG, keeps it resident for speed,
and can offload it to CPU after an idle timeout. Remote policy URLs require
encrypted transport and explicit opt-in. Tokens are fixed environment
variables (`VLA_POLICY_TOKEN`, `OPENPI_API_KEY`, or `GROOT_API_TOKEN`) and are
never workflow inputs.
See [`examples/robotics/README.md`](examples/robotics/README.md) for D-drive
WSL setup, platform boundaries, current model readiness, observation schemas,
action safety semantics, and the runnable API example.
### Moondream 3 / 3.1 isolated runtime
Moondream's official Photon package pins Pillow below version 11 while
current ComfyUI uses a newer Pillow. It therefore runs in a dedicated sidecar
environment and never changes ComfyUI's Python packages. Read and accept the
[Moondream Model License 1.0](https://moondream.ai/licenses/model/1.0), then
create the environment under the registered `LLavacheckpoints` model folder.
Linux/WSL/macOS:
```bash
runtime="ComfyUI/models/LLavacheckpoints/moondream31-runtime"
uv venv "$runtime/.venv" --python 3.12
uv pip install --python "$runtime/.venv/bin/python" \
-r ComfyUI/custom_nodes/ComfyUI_VLM_nodes/requirements-moondream31.txt
```
Windows PowerShell:
```powershell
$runtime = "ComfyUI\models\LLavacheckpoints\moondream31-runtime"
uv venv "$runtime\.venv" --python 3.12
uv pip install --python "$runtime\.venv\Scripts\python.exe" `
-r "ComfyUI\custom_nodes\ComfyUI_VLM_nodes\requirements-moondream31.txt"
```
The first Loader execution downloads the selected official model below that
runtime's `cache` directory. Use `moondream3.1-9B-A2B` for query, caption,
detection, and pointing. Use `moondream3-preview` only for the SVG segment
skill; the final 3.1 model card does not list segment. Set the server-side
`MOONDREAM_PYTHON` environment variable
when using a different isolated environment. Do not put this path or any
credential in a workflow.
Official Photon local inference currently supports NVIDIA Ampere-or-newer on
Linux/Windows and Apple Silicon on macOS 13 or newer. It does not currently
provide local ROCm, Intel GPU, or CPU execution. Those platforms retain every
portable Transformers, GGUF, API, and vision utility node in this pack.
On CUDA 12 x86-64 systems the isolated requirements deliberately install
`nvidia-cuda-runtime-cu12==12.9.79`. Kestrel 0.4.6's AOT kernels require the
`cudaLibraryLoadData` entry point, which is absent from the CUDA 12.6 runtime
bundled by cu126 PyTorch. This pin updates only Photon's private runtime; it
does not replace ComfyUI's PyTorch build or the host NVIDIA driver.
GGUF nodes use optional `llama-cpp-python`. Install a wheel built for the
desired CUDA, ROCm/HIP, Metal, Vulkan, SYCL, or CPU backend:
@@ -264,7 +696,20 @@ Gemma 3 and PaLI-Gemma require accepting their model licenses on Hugging Face.
The runtime reports llama.cpp's own compiled backend, GPU-offload, mmap, and
mlock capabilities in **VLM Runtime Diagnostics**.
- `unload_after=false` caches one model per node instance for fast repeated
queues. Turn it on for maximum reclamation between prompts.
queues. Cache creation is serialized, so concurrent API work cannot make the
same node allocate duplicate model handles. Turn it on for maximum
reclamation between prompts.
- Moondream Photon asks ComfyUI to make room before it starts, then owns one
exact isolated process. `unload_after=true` gracefully shuts it down and
terminates that process if necessary, which releases Photon model, KV-cache,
and CUDA-graph allocations without flushing unrelated ComfyUI models. The
sidecar intentionally does not inherit ComfyUI's PyTorch allocator override;
Photon's CUDA-graph capture uses the native allocator in its own process. The
worker does not inherit unrelated provider keys or proxy credentials; only
`HF_TOKEN`, and `MOONDREAM_API_KEY` for an explicitly selected adapter, may
cross into its server-side environment. Base-model sidecars honor
`DO_NOT_TRACK` locally and do not start Kestrel's anonymous telemetry task.
Its random IPC secret is not placed on the process command line.
- A connected `video_frames` batch becomes the primary visual input. The
optional still-image socket is ignored for video inference so smaller models
cannot silently answer from the wrong media.
@@ -285,10 +730,96 @@ are not available for the installed PyTorch/backend combination.
## API nodes
`PromptGenerateAPI` supports the current OpenAI Responses API, the legacy Chat
Completions API, and compatible base URLs. API keys can be supplied by node or
environment (`OPENAI_API_KEY`, `ANTHROPIC_API_KEY`, `GEMINI_API_KEY`,
`GROQ_API_KEY`). Keys are never persisted by this repository.
**Hosted LLM API (Secure)** and **Hosted VLM API (Secure)** share a provider
layer built around the current OpenAI Responses and Chat Completions request
shapes, with Anthropic using its native Messages/vision contract and Gemini
switching to its native multimodal contract for grounded or structured calls.
The VLM node
accepts a still image or a video-frame batch, samples
frames uniformly, resizes and JPEG-compresses them, and enforces per-image and
total request limits before upload. Both nodes can stream text into a connected
`ViewText` node.
Both API nodes also expose:
- **Native web search** for OpenAI, Gemini, Anthropic, xAI, and any compatible
model routed through OpenRouter. Unsupported presets fail clearly before a
model request instead of silently pretending to search. Search can add
provider cost and has provider-specific data terms, so it is off by default.
- **JSON object** and **JSON Schema** output. Completed JSON is always parsed
locally, JSON Schema results are validated locally, and invalid results fail
the node instead of flowing into downstream automation.
- **Open-source structured VLM output** through Custom / Local endpoints.
OpenAI-standard mode supports vLLM, Ollama, and compatible servers;
`llama.cpp JSON Schema` emits llama.cpp's direct schema dialect; and
`JSON object + local validation` is a portable fallback for servers that
implement only JSON mode.
User-provided schemas are capped at 64,000 characters, bounded by depth/node
count, checked against their declared JSON Schema draft, and may use only local
fragment `$ref` values. Remote/file references are rejected so validation can
never turn into an unexpected network or filesystem lookup.
Curated production profiles include:
| Provider | Presets | Server environment variable |
| --- | --- | --- |
| OpenAI | GPT-5.6 Terra, Sol, Luna | `OPENAI_API_KEY` |
| Google | Gemini 3.6 Flash, 3.5 Flash, 3.5 Flash-Lite | `GEMINI_API_KEY` |
| Anthropic | Claude Fable 5, Opus 5, Sonnet 5, Haiku 4.5 | `ANTHROPIC_API_KEY` |
| xAI | Grok 4.5 | `XAI_API_KEY` |
| DeepSeek | V4 Flash, V4 Pro | `DEEPSEEK_API_KEY` |
| Groq | Qwen 3.6 27B Vision, GPT-OSS 20B | `GROQ_API_KEY` |
| Mistral | Mistral Large, Mistral Small, Ministral 14B | `MISTRAL_API_KEY` |
| Together AI | Kimi K2.5, Qwen 3.5 9B | `TOGETHER_API_KEY` |
| OpenRouter | Any compatible model ID | `OPENROUTER_API_KEY` |
| Custom/local | OpenAI-compatible endpoint | `CUSTOM_API_KEY` |
Preset IDs were reviewed on 2026-07-29 against the official
[OpenAI](https://developers.openai.com/api/docs/models),
[Gemini](https://ai.google.dev/gemini-api/docs/models),
[Claude](https://platform.claude.com/docs/en/about-claude/models/overview),
[xAI](https://docs.x.ai/developers/models),
[DeepSeek](https://api-docs.deepseek.com/updates/),
[Groq](https://console.groq.com/docs/models),
[Mistral](https://docs.mistral.ai/models/), and
[Together](https://docs.together.ai/docs/inference/recommended-models), plus
[OpenRouter's multimodal compatibility](https://openrouter.ai/docs/guides/overview/multimodal/overview)
catalogs. Use `model_override` when a provider exposes a newer compatible model
before the next node-pack release.
The capability routing follows the current official
[OpenAI web-search](https://developers.openai.com/api/docs/guides/tools-web-search)
and [structured-output](https://developers.openai.com/api/docs/guides/structured-outputs)
contracts,
[Gemini grounding](https://ai.google.dev/gemini-api/docs/google-search) and
[structured output](https://ai.google.dev/gemini-api/docs/structured-output),
[Claude web-search](https://platform.claude.com/docs/en/agents-and-tools/tool-use/web-search-tool)
and [structured-output](https://platform.claude.com/docs/en/build-with-claude/structured-outputs)
contracts, [xAI web search](https://docs.x.ai/developers/tools/web-search) and
[structured outputs](https://docs.x.ai/developers/model-capabilities/text/structured-outputs),
and [OpenRouter server-side search](https://openrouter.ai/docs/guides/features/server-tools/web-search).
The local dialect is based on the
[llama.cpp server API](https://github.com/ggml-org/llama.cpp/blob/master/tools/server/README.md).
API keys are not node inputs. A workflow contains only the provider selection,
and the server resolves that provider's fixed environment variable at execution
time. Built-in credentials are pinned to the provider's official HTTPS host;
only the custom profile accepts a URL, and it can read only `CUSTOM_API_KEY`.
Remote custom URLs require HTTPS, while keyless HTTP is restricted to
`localhost`/loopback. Redirect following and environment proxies are disabled
by default, API calls are stateless, OpenAI Responses explicitly use
`store=false`, and provider exceptions are redacted before ComfyUI receives
them.
Web search sends the prompt (and, where supported, the same multimodal request)
to the selected provider's server-side search system. Do not enable it for
content that must not be processed under that provider's search terms.
Opening an older `PromptGenerateAPI` workflow automatically clears its former
plaintext key widget before the graph is configured. Save the migrated workflow
to overwrite the old file, and rotate any key that was previously saved or
shared. See [SECURITY.md](SECURITY.md) for setup and the exact threat model.
## Reliability guarantees
@@ -322,3 +853,23 @@ catalog-only evidence matrix.
Please report reproducible bugs at the
[issue tracker](https://github.com/gokayfem/ComfyUI_VLM_nodes/issues).
<details>
<summary><strong>Cite this project</strong></summary>
If ComfyUI VLM Nodes supports your work, please cite the software. GitHub also
provides ready-to-copy APA and BibTeX entries via **Cite this repository**.
```bibtex
@software{Aydogan_ComfyUI_VLM_Nodes_2026,
author = {Aydoğan, Gökay},
title = {ComfyUI VLM Nodes},
version = {3.5.0},
year = {2026},
url = {https://github.com/gokayfem/ComfyUI_VLM_nodes}
}
```
[ORCID](https://orcid.org/0000-0002-2343-9433) · [Citation metadata](CITATION.cff)
</details>
+120
View File
@@ -0,0 +1,120 @@
# API credential security
## Guarantees
- API keys are never accepted as node inputs, widget values, workflow fields,
outputs, metadata, or log messages.
- Each built-in provider reads only its standard server-side environment
variable and sends it only to that provider's fixed official HTTPS endpoint.
- A built-in provider key cannot be combined with a workflow-supplied URL.
- The custom endpoint reads only `CUSTOM_API_KEY`. Remote custom endpoints must
use HTTPS; unencrypted and keyless requests are limited to loopback.
- HTTP redirects and environment proxies are disabled by default. Proxy use is
an explicit non-secret node option for installations that require it.
- Hosted calls are stateless. No Python node-instance conversation history is
retained, and OpenAI Responses requests set `store=false`.
- Exceptions are bounded and redact the resolved key, URL-encoded variants,
bearer tokens, common provider-key formats, authorization fields, and URL
user-info before the message reaches ComfyUI.
- Local image/video-frame uploads are uniformly sampled, resized,
JPEG-compressed, limited to 4 MiB per image, and limited to 24 MiB total.
- User JSON Schemas are size/depth/node bounded and may contain only local
fragment `$ref` values. Remote URLs and file references are rejected before
validation, preventing schema resolution from becoming an SSRF or local-file
access path.
## Configure credentials
Set the matching variable in the environment that launches ComfyUI, then
restart ComfyUI:
| Provider | Variable |
| --- | --- |
| OpenAI | `OPENAI_API_KEY` |
| Google Gemini | `GEMINI_API_KEY` |
| Anthropic | `ANTHROPIC_API_KEY` |
| xAI | `XAI_API_KEY` |
| DeepSeek | `DEEPSEEK_API_KEY` |
| Groq | `GROQ_API_KEY` |
| Mistral | `MISTRAL_API_KEY` |
| Together AI | `TOGETHER_API_KEY` |
| OpenRouter | `OPENROUTER_API_KEY` |
| MiniMax | `MINIMAX_API_KEY` |
| Custom remote endpoint | `CUSTOM_API_KEY` |
| Universal VLA policy server | `VLA_POLICY_TOKEN` |
| openpi WebSocket server | `OPENPI_API_KEY` |
| Isaac-GR00T ZMQ server | `GROOT_API_TOKEN` |
For an interactive POSIX/WSL session, this avoids putting the value in shell
history:
```bash
read -rsp "Provider API key: " OPENAI_API_KEY
export OPENAI_API_KEY
python main.py
```
Use the equivalent secret manager or service environment mechanism for a
persistent installation. Do not commit a `.env` file, workflow containing an
old key, shell script containing a key, or copied ComfyUI log.
Web search is disabled by default. Enabling it sends the request content to the
selected provider's server-side search system and may have separate retention,
regional-availability, and billing terms. Treat it as an explicit data-sharing
choice; do not enable it for content that is outside those terms.
## Robotics policy endpoints
Robotics tokens are also server-side only. Workflow nodes select an endpoint,
but cannot select an arbitrary environment variable or contain the secret
value.
- The universal policy client permits unencrypted HTTP only on loopback.
Remote use requires HTTPS plus `allow_remote=true`; redirects and
environment proxies are disabled.
- The openpi client permits unencrypted WebSocket only on loopback. Remote use
requires WSS plus `allow_remote=true`.
- GR00T's official ZeroMQ protocol has token authentication but no built-in
transport encryption. Keep it on loopback/private infrastructure or place it
inside an authenticated encrypted tunnel. Never expose its port directly to
the public internet.
- Camera payloads are JPEG-compressed and bounded per frame and per request.
Response sizes, camera count, observation history, state/action dimensions,
and action horizons are bounded before use.
- MessagePack ndarray decoders reject object/void dtypes and never deserialize
pickle. The included HTTP sidecar uses bounded JSON instead of LeRobot's
pickle-based asynchronous transport.
- Errors redact the resolved token and authorization-like values. Reports
include only endpoint scheme/host/port, not request headers, full camera
payloads, or state data.
Robot observations may expose people, homes, workplaces, proprietary tasks,
and physical state. Treat them as sensitive even when no API key is present.
The safety node is a data validation gate, not a certified control system.
This package intentionally contains no ROS, serial, CAN, motor, or robot SDK
transport; a separate controller must enforce emergency stop, deadman,
watchdog, collision/workspace, command-age, and manufacturer limits.
## Legacy workflows
Versions before this security update exposed an `api_key` text widget.
The frontend migration clears position 3 of every serialized
`PromptGenerateAPI` node before LiteGraph creates the active node, including
nodes inside saved subgraph definitions. The backend independently rejects any
value that is not one of the two safe credential-source choices.
The source workflow file is not rewritten merely by opening it. Save the
migrated workflow, securely remove old copies, and rotate any credential that
was ever saved, shared, committed, backed up, or placed in an exported PNG.
## Threat boundary
ComfyUI custom nodes execute Python code with the permissions of the ComfyUI
process. Another untrusted custom-node package can read the same process
environment regardless of protections in this repository. Install only trusted
node packs, keep ComfyUI authenticated and bound to a trusted interface, and do
not expose an unauthenticated server to the public internet.
If a key may have been exposed, revoke it with the provider immediately, review
usage, create a replacement with the minimum needed project permissions and
spend limit, and restart ComfyUI with the replacement.
+6
View File
@@ -7,22 +7,27 @@ LOGGER = logging.getLogger("ComfyUI_VLM_nodes")
register_model_folder()
node_list = [
"acceleration",
"audioldm2",
"diagnostics",
"florence2",
"grounding",
"hosted_api",
"joytag",
"kosmos2",
"llavaloader",
"mcllava",
"minicpm",
"minimax_music",
"modern_vlm",
"molmo",
"moondream31",
"moondream2",
"moondream_script",
"paligemma",
"playmusic",
"qwen2vl",
"robotics",
"sam2",
"sam3_adapter",
"simpletext",
@@ -30,6 +35,7 @@ node_list = [
"suggest",
"tracking",
"uform",
"video_intelligence",
"vision_utils",
]
+266
View File
@@ -0,0 +1,266 @@
# Robotics and VLA workflows
The robotics nodes make ComfyUI a policy-development, inspection, and
simulation surface. They do **not** send commands to motors, ROS, CAN, serial,
or a robot SDK.
The boundary is intentional:
```text
camera/state/task
|
v
VLA Observation Builder
|
v
isolated policy server ---> raw action chunk
|
v
VLA Action Safety Gate
|
+-------------+-------------+
| |
v v
inspect / plot / record simulator or your own
supervised controller bridge
```
A real controller bridge must independently enforce a deadman, watchdog,
emergency stop, collision/workspace limits, timestamps, command freshness, and
the manufacturer's limits. A `safe_for_handoff=true` workflow result only
means that the declared ComfyUI profile checks passed.
## Why policy runtimes are isolated
LeRobot, openpi, Isaac-GR00T, OpenVLA/OFT, and Octo use different PyTorch/JAX,
CUDA, Transformers, compiler, and operating-system combinations. Installing
all of those into ComfyUI would replace or constrain the working accelerator
stack and make Windows, macOS, ROCm, and XPU support worse.
The ComfyUI package therefore contains only:
- typed state/action/camera contracts;
- bounded image serialization;
- a dependency-light universal HTTPS/loopback HTTP client;
- exact clients for the official openpi MessagePack WebSocket and GR00T
MessagePack/ZeroMQ protocols;
- action validation, horizon control, inspection, and plotting.
The heavyweight policy stays in its own process, container, WSL distribution,
Linux machine, Mac, or GPU server. This also allows ComfyUI to use AMD ROCm,
Apple Metal, Intel XPU, or CPU while a policy runs on an NVIDIA Linux server.
## Fast path: SmolVLA through the universal sidecar
Use a separate LeRobot environment. Current LeRobot documentation recommends
Python 3.12 and exposes policy-specific extras. On this computer, keep it on
the D drive:
```bash
# WSL
python3.12 -m venv /mnt/d/vla-runtime/lerobot-smolvla
source /mnt/d/vla-runtime/lerobot-smolvla/bin/activate
python -m pip install --upgrade pip
python -m pip install "lerobot[smolvla]"
export VLA_POLICY_TOKEN="$(python -c 'import secrets; print(secrets.token_urlsafe(32))')"
python /mnt/d/ComfyUI_windows_portable/ComfyUI/custom_nodes/ComfyUI_VLM_nodes/examples/robotics/lerobot_policy_server.py \
--policy-type smolvla \
--policy-path YOUR_FINE_TUNED_SMOLVLA_CHECKPOINT \
--device auto \
--actions-per-chunk 16 \
--idle-offload-seconds 300
```
Set the same `VLA_POLICY_TOKEN` in the environment that launches ComfyUI.
Never put it in a workflow. In `VLA Policy — Universal HTTP`, use
`http://127.0.0.1:8787`.
`lerobot/smolvla_base` is a base model. It is a useful fine-tuning starting
point, not a universal zero-shot controller. Use an embodiment-specific
checkpoint whose feature names, action dimensions, state dimensions,
normalization statistics, and camera keys match the workflow.
The sidecar:
- loads only the chosen policy and its serialized pre/post-processors;
- uses `predict_action_chunk` when provided and falls back to `select_action`;
- keeps the model resident by default for low latency;
- can move it to CPU after an idle interval and move it back on demand;
- accepts one request at a time per policy, preventing stateful policy races;
- uses bounded JSON/JPEG rather than pickle;
- never returns tracebacks, environment variables, request data, or
authorization headers.
For a real local API acceptance run, start ComfyUI and the policy sidecar, put
an image in ComfyUI's `input` directory, then run:
```bash
python tests/manual_robotics_smoke.py \
--comfy-url http://127.0.0.1:8188 \
--policy-url http://127.0.0.1:8787 \
--image robot_front.png
```
The script queues the graph through `POST /prompt`, waits on its history entry,
and prints the policy report, safety report, first action, and preview filename.
Install the relevant official LeRobot extra for another policy. Examples are
`lerobot[pi]` for π0/π0.5/π0-FAST and `lerobot[smolvla]` for SmolVLA. Some
newer policy integrations may require installing current LeRobot from source
with their documented extra.
## Native openpi server
Install the small ComfyUI client dependencies:
```bash
python -m pip install -r requirements-robotics-client.txt
```
Run the official openpi policy WebSocket server in its own supported
environment. The upstream runtime is currently tested on Ubuntu 22.04 and an
NVIDIA GPU with more than 8 GB for inference; use WSL/Docker or a remote Linux
server rather than forcing it into a macOS/Windows ComfyUI environment.
Use:
- `Flat keys (DROID / LIBERO)` for observations such as
`observation/image`, `observation/wrist_image`, and `observation/state`;
- `Nested images (ALOHA)` for `state`, an `images` mapping such as
`cam_high`/wrist cameras, and `prompt`.
The workflow supplies key names, but the policy's own transform still defines
the exact shapes and normalization. `OPENPI_API_KEY` is read only from the
ComfyUI server environment. Remote endpoints require WSS and explicit
`allow_remote=true`.
## Native Isaac-GR00T N1.7 server
Install the same lightweight robotics client requirements in ComfyUI. Run the
official GR00T `PolicyServer` beside an embodiment-compatible `Gr00tPolicy`.
The node sends the documented nested contract:
```text
video.<camera> uint8 [batch=1, history, height, width, RGB=3]
state.state float32[batch=1, history, state_dim]
language.task string [batch=1, 1]
```
The official server returns one or more physical-unit action streams with
shape `[batch, horizon, dimension]`. The node flattens those streams while
preserving their named slices. `GROOT_API_TOKEN` remains in the ComfyUI
environment.
GR00T N1.7 currently targets NVIDIA CUDA/Jetson Linux and needs an
embodiment-compatible base or post-trained checkpoint. A ComfyUI client on
Windows, macOS, ROCm, or another machine may call that server over a trusted
network, but remote access must be explicitly enabled. Native GR00T ZMQ does
not encrypt traffic; use a private authenticated network/tunnel. Prefer the
universal HTTPS bridge when transport-layer encryption is required.
## Model catalog: what “available” means
`VLA Model Catalog` distinguishes these states:
| Family | Example checkpoint | Route | Important qualification |
| --- | --- | --- | --- |
| SmolVLA | `lerobot/smolvla_base` | LeRobot HTTP sidecar | 450M and the best small starting point; fine-tune for the robot |
| X-VLA | `lerobot/xvla-base` | LeRobot HTTP sidecar | 0.9B cross-embodiment base; use a matching domain checkpoint |
| π0 | `lerobot/pi0_base` | LeRobot or openpi | Base/fine-tuning model, not a universal drop-in controller |
| π0-FAST | `lerobot/pi0fast-base` | LeRobot or openpi | Faster tokenized action generation |
| π0.5 | `lerobot/pi05_base` | LeRobot or openpi | Open-world generalization; still embodiment-specific |
| GR00T N1.7 | `nvidia/GR00T-N1.7-3B` | GR00T ZMQ or LeRobot | Base has specific zero-shot tags; other robots need post-training |
| X-Square WALL-OSS | `x-square-robot/wall-oss-flow` | LeRobot HTTP sidecar | MoE research model; validate checkpoint terms and embodiment |
| MolmoAct2 | `lerobot/MolmoAct2-SO100_101-LeRobot` | LeRobot HTTP sidecar | Converted SO-100/SO-101 checkpoint |
| VLA-JEPA | `lerobot/VLA-JEPA-Pretrain` | LeRobot HTTP sidecar | DROID pretrain plus LIBERO/SimplerEnv checkpoints |
| LingBot-VA | `lerobot/lingbot_va_base` | LeRobot HTTP sidecar | Prefer its LIBERO-Long/RoboTwin post-train where applicable |
| FastWAM | released LIBERO checkpoint | LeRobot HTTP sidecar | Heavy world-action research runtime |
| EO-1 / EVO-1 | your trained checkpoint | LeRobot HTTP sidecar | Architecture support, not a universal ready-made controller |
| OpenVLA-OFT | compatible OFT fine-tune | dedicated sidecar | OFT is the faster multi-image/high-frequency OpenVLA route |
| Octo small | Octo small 27M | dedicated JAX sidecar | Lightweight legacy research baseline |
The catalog is a verified runtime/checkpoint map, not a promise that a base
checkpoint understands an arbitrary robot. Exact data transforms and
fine-tuning are part of the policy.
## Observation history and real-time use
Connect an `IMAGE` batch to a camera input to represent temporal history. All
camera batches must have the same length, although a one-frame camera may
broadcast. Use:
`Video Slice` → `VLM Adaptive Frame Sampler` or a live capture source →
`VLM Image Pixel Budget` → `VLA Observation Builder`
For closed-loop robotics, do not run an unbounded ComfyUI queue for each motor
tick. Use ComfyUI to prototype and inspect the observation/policy/safety
contract, and use the policy runtime's asynchronous control support for the
actual high-frequency loop. LeRobot supports asynchronous action chunks and
GR00T supports TensorRT deployment; both are better places for timing-critical
execution.
## Action safety semantics
The `VLA Action Safety Gate` checks:
- policy action dimension against the embodiment;
- NaN and infinity;
- minimum and maximum values;
- maximum change per action dimension and control step;
- the requested execution horizon.
Modes:
- `Block unsafe`: raise and stop the workflow on any violation.
- `Clamp safely`: replace non-finite values conservatively, then clamp bounds
and sequential per-step deltas.
- `Hold position on unsafe`: replace the whole chunk with the explicitly
supplied previous/current command.
- `Report only`: preserve the raw trajectory and set
`safe_for_handoff=false`.
For delta-action policies, `previous_action_json` means the previous delta
command, not an absolute joint pose. Define the profile in the same units and
semantics as the policy output.
`VLA Actions From JSON` imports recorded/simulator trajectories without a
network policy. `VLA Action Chunk Replan` blends the unexecuted edge of an old
chunk into a new chunk to reduce discontinuities, then the result should pass
through the safety gate again. This deterministic blend is useful for workflow
experiments but does not replace LeRobot's asynchronous controller or a
policy-specific real-time chunking implementation.
## API example
`vla_http_policy_safety_api.json` is a ComfyUI API prompt graph. Put
`robot_front.png` in `ComfyUI/input`, start a compatible sidecar, then POST:
```json
{"prompt": {"...": "contents of vla_http_policy_safety_api.json"}}
```
It builds an observation, calls the policy, clamps it against the explicit
profile, renders the trajectory, and outputs both inference and safety JSON.
## Security checklist
- Keep all policy tokens in environment variables.
- Leave `allow_remote=false` for local servers.
- Remote universal endpoints must use HTTPS; remote openpi endpoints must use
WSS.
- Never expose GR00T ZMQ directly to an untrusted network.
- Pin checkpoint revisions when reproducibility matters.
- Treat camera images, task language, and robot state as sensitive data.
- Do not connect action JSON directly to hardware without a separate
supervised controller bridge and independent safety system.
Authoritative upstream documentation:
- [LeRobot installation](https://huggingface.co/docs/lerobot/main/en/installation)
- [LeRobot SmolVLA](https://huggingface.co/docs/lerobot/smolvla)
- [LeRobot asynchronous inference](https://huggingface.co/docs/lerobot/async)
- [Physical Intelligence openpi](https://github.com/Physical-Intelligence/openpi)
- [NVIDIA Isaac-GR00T](https://github.com/NVIDIA/Isaac-GR00T)
- [OpenVLA and OFT](https://github.com/openvla/openvla)
- [Octo](https://github.com/octo-models/octo)
+1
View File
@@ -0,0 +1 @@
"""Runnable, dependency-isolated robotics policy bridge examples."""
+434
View File
@@ -0,0 +1,434 @@
#!/usr/bin/env python
"""Isolated LeRobot policy server for the ComfyUI VLA HTTP node.
Run this file in a dedicated environment that contains LeRobot and the
policy-specific dependencies. Do not install LeRobot's full dependency stack
into ComfyUI merely to use this bridge.
"""
from __future__ import annotations
import argparse
import base64
import hmac
import io
import json
import os
import threading
import time
from http import HTTPStatus
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from typing import Any
import numpy as np
import torch
from PIL import Image
MAX_REQUEST_BYTES = 64 * 1024 * 1024
MAX_CAMERAS = 16
MAX_FRAMES_PER_CAMERA = 256
MAX_IMAGE_BYTES = 8 * 1024 * 1024
MAX_IMAGE_PIXELS = 16 * 1024 * 1024
MAX_STATE_DIM = 2_048
MAX_ACTION_DIM = 2_048
MAX_TASK_CHARS = 16_384
def _json_bytes(value: Any) -> bytes:
return json.dumps(
value,
ensure_ascii=False,
allow_nan=False,
separators=(",", ":"),
).encode("utf-8")
def _device(value: str) -> str:
if value != "auto":
return value
if torch.cuda.is_available():
return "cuda"
if hasattr(torch, "xpu") and torch.xpu.is_available():
return "xpu"
if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
return "mps"
return "cpu"
def _decode_image(frame: dict[str, Any]) -> np.ndarray:
if frame.get("encoding") != "base64-jpeg":
raise ValueError("Only base64-jpeg camera frames are supported.")
raw = base64.b64decode(frame["data"], validate=True)
if len(raw) > MAX_IMAGE_BYTES:
raise ValueError("Encoded camera frame exceeds the 8 MiB safety limit.")
with Image.open(io.BytesIO(raw)) as image:
if image.width * image.height > MAX_IMAGE_PIXELS:
raise ValueError("Decoded camera frame exceeds the pixel safety limit.")
return np.asarray(image.convert("RGB"), dtype=np.uint8).copy()
def _decode_observation(payload: dict[str, Any]) -> dict[str, Any]:
if payload.get("schema") != "comfyui-vlm/robot-observation":
raise ValueError("Unsupported observation schema.")
if int(payload.get("version", 0)) != 1:
raise ValueError("Unsupported observation schema version.")
cameras = payload.get("cameras")
if not isinstance(cameras, dict) or not 1 <= len(cameras) <= MAX_CAMERAS:
raise ValueError("cameras must contain between 1 and 16 entries.")
observation: dict[str, Any] = {}
for key, encoded_frames in cameras.items():
key = str(key).strip()
if not key or len(key) > 256 or any(ord(char) < 32 for char in key):
raise ValueError("Camera names must contain 1 to 256 printable characters.")
if not isinstance(encoded_frames, list) or not (
1 <= len(encoded_frames) <= MAX_FRAMES_PER_CAMERA
):
raise ValueError(f"Camera {key!r} has an invalid history.")
# Current LeRobot policy processors accept one current observation.
# ComfyUI may send history for servers/models that use it; this generic
# bridge deliberately selects the latest frame.
array = _decode_image(encoded_frames[-1])
tensor = torch.from_numpy(array).permute(2, 0, 1).to(torch.float32) / 255.0
observation[key] = tensor.unsqueeze(0)
state = np.asarray(payload.get("state"), dtype=np.float32)
if (
state.ndim != 1
or not 1 <= state.size <= MAX_STATE_DIM
or not np.isfinite(state).all()
):
raise ValueError(f"state must contain 1 to {MAX_STATE_DIM} finite values.")
observation["observation.state"] = torch.from_numpy(state).unsqueeze(0)
task = str(payload.get("task", "")).strip()
if not task or len(task) > MAX_TASK_CHARS:
raise ValueError(f"task must contain 1 to {MAX_TASK_CHARS} characters.")
observation["task"] = task
return observation
def _postprocess_chunk(postprocessor, action: torch.Tensor) -> torch.Tensor:
if action.ndim == 1:
action = action.unsqueeze(0)
if action.ndim == 2:
# select_action normally returns [batch, dim].
processed = postprocessor(action)
if processed.ndim == 1:
processed = processed.unsqueeze(0)
return processed.unsqueeze(1) if processed.ndim == 2 else processed
if action.ndim != 3:
raise ValueError(f"Policy returned unsupported action shape {tuple(action.shape)}.")
processed_steps = [postprocessor(action[:, index, :]) for index in range(action.shape[1])]
return torch.stack(processed_steps, dim=1)
def _feature_metadata(features: Any) -> dict[str, dict[str, Any]]:
"""Return the portable part of a LeRobot policy feature contract."""
result: dict[str, dict[str, Any]] = {}
for key, feature in (features or {}).items():
if isinstance(feature, dict):
feature_type = feature.get("type")
shape = feature.get("shape", ())
else:
feature_type = getattr(feature, "type", None)
shape = getattr(feature, "shape", ())
feature_type = getattr(feature_type, "value", feature_type)
dimensions: list[int | str | None] = []
for dimension in shape or ():
if dimension is None:
dimensions.append(None)
continue
try:
dimensions.append(int(dimension))
except (TypeError, ValueError):
dimensions.append(str(dimension))
result[str(key)] = {
"type": str(feature_type) if feature_type is not None else "UNKNOWN",
"shape": dimensions,
}
return result
def _optional_config_int(config: Any, name: str) -> int | None:
value = getattr(config, name, None)
try:
return None if value is None else int(value)
except (TypeError, ValueError):
return None
class PolicyRuntime:
def __init__(
self,
*,
policy_type: str,
policy_path: str,
revision: str | None,
device: str,
actions_per_chunk: int,
idle_offload_seconds: float,
):
self.policy_type = policy_type
self.policy_path = policy_path
self.revision = revision
self.device = _device(device)
self.actions_per_chunk = actions_per_chunk
self.idle_offload_seconds = idle_offload_seconds
self.lock = threading.Lock()
self.policy = None
self.preprocessor = None
self.postprocessor = None
self.resident_device = "unloaded"
self.last_request = 0.0
self.load_seconds = 0.0
self._load()
if idle_offload_seconds > 0 and self.device != "cpu":
threading.Thread(target=self._idle_worker, daemon=True).start()
def _load(self) -> None:
from lerobot.policies import get_policy_class, make_pre_post_processors
started = time.perf_counter()
policy_class = get_policy_class(self.policy_type)
kwargs = {}
if self.revision:
kwargs["revision"] = self.revision
self.policy = policy_class.from_pretrained(self.policy_path, **kwargs)
self.policy.eval()
self.policy.to(self.device)
overrides = {"device": self.device}
self.preprocessor, self.postprocessor = make_pre_post_processors(
self.policy.config,
pretrained_path=self.policy_path,
pretrained_revision=self.revision,
preprocessor_overrides={"device_processor": overrides},
postprocessor_overrides={"device_processor": overrides},
)
self.resident_device = self.device
self.last_request = time.monotonic()
self.load_seconds = time.perf_counter() - started
def _ensure_resident(self) -> None:
if self.resident_device != self.device:
self.policy.to(self.device)
self.resident_device = self.device
def _idle_worker(self) -> None:
interval = min(max(self.idle_offload_seconds / 4, 1.0), 30.0)
while True:
time.sleep(interval)
if time.monotonic() - self.last_request < self.idle_offload_seconds:
continue
if not self.lock.acquire(blocking=False):
continue
try:
if (
self.resident_device != "cpu"
and time.monotonic() - self.last_request >= self.idle_offload_seconds
):
self.policy.to("cpu")
self.resident_device = "cpu"
finally:
self.lock.release()
def metadata(self) -> dict[str, Any]:
config = self.policy.config
return {
"protocol": "comfyui-vla-http-v1",
"policy_type": self.policy_type,
"policy_path": self.policy_path,
"revision": self.revision,
"configured_device": self.device,
"resident_device": self.resident_device,
"actions_per_chunk": self.actions_per_chunk,
"idle_offload_seconds": self.idle_offload_seconds,
"load_seconds": self.load_seconds,
"policy_contract": {
"input_features": _feature_metadata(
getattr(config, "input_features", None)
),
"output_features": _feature_metadata(
getattr(config, "output_features", None)
),
"observation_steps": _optional_config_int(config, "n_obs_steps"),
"native_chunk_size": _optional_config_int(config, "chunk_size"),
"native_action_steps": _optional_config_int(config, "n_action_steps"),
},
}
def infer(self, payload: dict[str, Any]) -> dict[str, Any]:
observation = _decode_observation(payload)
with self.lock:
self._ensure_resident()
started = time.perf_counter()
processed = self.preprocessor(observation)
preprocess_ms = (time.perf_counter() - started) * 1000
started_inference = time.perf_counter()
with torch.inference_mode():
predictor = getattr(self.policy, "predict_action_chunk", None)
if callable(predictor):
action = predictor(processed)
else:
action = self.policy.select_action(processed)
inference_ms = (time.perf_counter() - started_inference) * 1000
started_postprocess = time.perf_counter()
action = _postprocess_chunk(self.postprocessor, action)
if action.ndim == 3:
if action.shape[0] != 1:
raise ValueError("Only policy batch size 1 is supported.")
action = action[0]
elif action.ndim == 1:
action = action.unsqueeze(0)
if action.ndim != 2:
raise ValueError(f"Unexpected final action shape {tuple(action.shape)}.")
action = action[: self.actions_per_chunk].detach().to("cpu", torch.float32)
if not 1 <= int(action.shape[1]) <= MAX_ACTION_DIM:
raise ValueError(
f"Policy action dimension must be in [1, {MAX_ACTION_DIM}]."
)
if not torch.isfinite(action).all():
# Preserve the response for ComfyUI's safety node, but do not
# serialize non-standard JSON numbers.
raise ValueError("Policy returned NaN or infinite action values.")
postprocess_ms = (time.perf_counter() - started_postprocess) * 1000
self.last_request = time.monotonic()
return {
"actions": action.tolist(),
"server_timing": {
"preprocess_ms": preprocess_ms,
"infer_ms": inference_ms,
"postprocess_ms": postprocess_ms,
},
"policy": {
"type": self.policy_type,
"path": self.policy_path,
"device": self.device,
},
}
class PolicyHandler(BaseHTTPRequestHandler):
server_version = "ComfyUI-VLA-Policy/1"
def log_message(self, format_string: str, *args: Any) -> None:
# The request path is safe to log. Headers and bodies may contain
# credentials or camera/state data and are intentionally excluded.
print(f"{self.address_string()} - {format_string % args}")
@property
def runtime(self) -> PolicyRuntime:
return self.server.runtime
@property
def token(self) -> str:
return self.server.token
def _authorized(self) -> bool:
if not self.token:
return True
supplied = self.headers.get("Authorization", "")
expected = f"Bearer {self.token}"
return hmac.compare_digest(supplied, expected)
def _send(self, status: HTTPStatus, value: Any) -> None:
body = _json_bytes(value)
self.send_response(status.value)
self.send_header("Content-Type", "application/json; charset=utf-8")
self.send_header("Content-Length", str(len(body)))
self.send_header("Cache-Control", "no-store")
self.send_header("X-Content-Type-Options", "nosniff")
self.end_headers()
self.wfile.write(body)
def do_GET(self) -> None: # noqa: N802
if self.path not in {"/healthz", "/v1/metadata"}:
self._send(HTTPStatus.NOT_FOUND, {"error": "not_found"})
return
if not self._authorized():
self._send(HTTPStatus.UNAUTHORIZED, {"error": "unauthorized"})
return
if self.path == "/healthz":
self._send(HTTPStatus.OK, {"status": "ok"})
else:
self._send(HTTPStatus.OK, self.runtime.metadata())
def do_POST(self) -> None: # noqa: N802
if self.path != "/v1/infer":
self._send(HTTPStatus.NOT_FOUND, {"error": "not_found"})
return
if not self._authorized():
self._send(HTTPStatus.UNAUTHORIZED, {"error": "unauthorized"})
return
try:
content_length = int(self.headers.get("Content-Length", "0"))
if not 1 <= content_length <= MAX_REQUEST_BYTES:
raise ValueError("Request body size is invalid.")
body = self.rfile.read(content_length)
payload = json.loads(body)
if not isinstance(payload, dict):
raise ValueError("Request body must be a JSON object.")
result = self.runtime.infer(payload)
except (TypeError, ValueError, json.JSONDecodeError) as exc:
self._send(HTTPStatus.BAD_REQUEST, {"error": str(exc)[:1000]})
return
except Exception as exc:
# Do not return tracebacks, request data, environment variables, or
# authorization headers across the network.
self._send(
HTTPStatus.INTERNAL_SERVER_ERROR,
{"error": f"{type(exc).__name__}: {str(exc)[:800]}"},
)
return
self._send(HTTPStatus.OK, result)
def main() -> None:
parser = argparse.ArgumentParser(
description="Serve one LeRobot policy through the ComfyUI VLA HTTP protocol."
)
parser.add_argument("--policy-type", required=True, help="LeRobot policy type, e.g. smolvla")
parser.add_argument("--policy-path", required=True, help="Hub repo id or local checkpoint")
parser.add_argument("--revision", default=None, help="Optional immutable Hub revision")
parser.add_argument("--device", default="auto", help="auto, cuda, mps, xpu, or cpu")
parser.add_argument("--host", default="127.0.0.1")
parser.add_argument("--port", type=int, default=8787)
parser.add_argument("--actions-per-chunk", type=int, default=16)
parser.add_argument(
"--idle-offload-seconds",
type=float,
default=0.0,
help="Move the policy to CPU after this idle period; 0 keeps it resident.",
)
args = parser.parse_args()
if not 1 <= args.port <= 65_535:
parser.error("--port must be in [1, 65535]")
if not 1 <= args.actions_per_chunk <= 4096:
parser.error("--actions-per-chunk must be in [1, 4096]")
if args.idle_offload_seconds < 0:
parser.error("--idle-offload-seconds must be non-negative")
runtime = PolicyRuntime(
policy_type=args.policy_type,
policy_path=args.policy_path,
revision=args.revision,
device=args.device,
actions_per_chunk=args.actions_per_chunk,
idle_offload_seconds=args.idle_offload_seconds,
)
token = os.environ.get("VLA_POLICY_TOKEN", "").strip()
server = ThreadingHTTPServer((args.host, args.port), PolicyHandler)
server.runtime = runtime
server.token = token
print(
f"Policy ready at http://{args.host}:{args.port}/v1/infer "
f"(type={args.policy_type}, device={runtime.device}, auth={'on' if token else 'off'})"
)
try:
server.serve_forever()
except KeyboardInterrupt:
pass
finally:
server.server_close()
if __name__ == "__main__":
main()
@@ -0,0 +1,130 @@
{
"1": {
"class_type": "LoadImage",
"inputs": {
"image": "robot_front.png"
}
},
"2": {
"class_type": "VLAEmbodimentProfile",
"inputs": {
"preset": "Generic 7-DoF joint + gripper",
"control_hz": 20.0,
"state_names_json": "",
"action_names_json": "",
"action_min_json": "",
"action_max_json": "",
"max_delta_json": "",
"camera_names_json": "",
"action_mode_override": ""
}
},
"3": {
"class_type": "VLAObservationBuilder",
"inputs": {
"task": "Pick up the blue cube and place it in the tray.",
"state_json": "[0, 0, 0, 0, 0, 0, 0, 0]",
"primary_image": [
"1",
0
],
"primary_camera": "observation.images.front",
"history_fps": 10.0,
"timestamp": 0.0,
"embodiment": [
"2",
0
]
}
},
"4": {
"class_type": "VLAHTTPPolicy",
"inputs": {
"observation": [
"3",
0
],
"endpoint": "http://127.0.0.1:8787",
"timeout_seconds": 120.0,
"include_history": true,
"allow_remote": false
}
},
"5": {
"class_type": "VLAActionSafety",
"inputs": {
"actions": [
"4",
0
],
"embodiment": [
"2",
0
],
"mode": "Clamp safely",
"execution_horizon": 8,
"previous_action_json": "[0, 0, 0, 0, 0, 0, 0, 0]"
}
},
"6": {
"class_type": "VLATrajectoryPreview",
"inputs": {
"actions": [
"5",
0
],
"width": 960,
"height": 480,
"embodiment": [
"2",
0
]
}
},
"7": {
"class_type": "VLAActionInspect",
"inputs": {
"actions": [
"5",
0
],
"step_index": 0
}
},
"8": {
"class_type": "PreviewImage",
"inputs": {
"images": [
"6",
0
]
}
},
"9": {
"class_type": "ViewText",
"inputs": {
"text": [
"4",
1
]
}
},
"10": {
"class_type": "ViewText",
"inputs": {
"text": [
"5",
1
]
}
},
"11": {
"class_type": "ViewText",
"inputs": {
"text": [
"7",
0
]
}
}
}
+57
View File
@@ -0,0 +1,57 @@
{
"1": {
"class_type": "SimpleText",
"inputs": {
"input_text": "Model response:\n```json\n{\"scene\":{\"subject\":\"warehouse robot\",\"action\":\"moving a blue crate\"}}\n```"
}
},
"2": {
"class_type": "VLMJSONExtract",
"inputs": {
"text": [
"1",
0
],
"path": "$.scene.action",
"output_format": "Text",
"if_missing": "Error",
"default_value": ""
}
},
"3": {
"class_type": "VLMTextTemplate",
"inputs": {
"template": "{instruction}\n\nObserved action: {text1}",
"variables_json": "{\"instruction\":\"Write one concise video-generation prompt.\"}",
"missing_values": "Error",
"text1": [
"2",
0
]
}
},
"4": {
"class_type": "VLMTextClean",
"inputs": {
"text": [
"3",
0
],
"unicode_normalization": "NFC",
"whitespace": "Normalize line endings",
"trim_edges": true,
"remove_outer_markdown_fence": false,
"deduplicate_lines": false,
"max_characters": 0
}
},
"5": {
"class_type": "ViewText",
"inputs": {
"text": [
"4",
0
]
}
}
}
+32
View File
@@ -29,6 +29,38 @@ Runs Grounding DINO Tiny over `grounding_input.png`. Node 2 outputs:
`PreviewImage` displays output 2 and `ViewText` reports output 1.
### `vlm_performance_preflight_api.json`
Loads `vlm_api_people_birds.mp4` with Comfy core video nodes, applies the
`Fast video` performance profile, runs the track-aware adaptive sampler, and
then applies a 14-pixel-aligned image budget. The preview shows the exact batch
that can be connected to any local or hosted VLM. Three `ViewText` nodes report
the selected source indices/timestamps, pixel reduction, and active profile.
### `moondream3_preview_svg_segment_api.json`
Runs the official Moondream 3 Preview SVG segmentation skill over
`moondream_segment_input.png`. Read the linked model license and change
`license_accepted` to `true` before queueing. The graph previews the
black/white mask, isolated foreground cutout, and mask/box/polygon overlay;
`ViewText` receives the exact native SVG path plus its normalized bbox.
Moondream's path coordinates are normalized within the returned bbox. The
node preserves that path verbatim, safely flattens curves/arcs, applies an
even-odd fill for subpath holes, and supersamples the raster edge. The
canonical detection keeps both the primary polygon and the full in-process
mask.
### `moondream31_video_detect_api.json`
Loads `moondream_video_input.mp4`, passes the real frame batch and source FPS
to Moondream, and analyzes every frame with four concurrent requests. Photon
uses the Loader's `max_batch_size=4` scheduler capacity to form dynamic
batches. `ViewText` reports measured throughput and real-time factor. Increase
`frame_stride` to 2, 3, or more when full-frame analysis cannot keep up with
the source FPS; the canonical results preserve original frame indices and
timestamps.
### `sam2_video_tracking_api.json`
Runs this bounded pipeline:
@@ -0,0 +1,76 @@
{
"1": {
"class_type": "LoadVideo",
"inputs": {
"file": "moondream_video_input.mp4"
}
},
"2": {
"class_type": "GetVideoComponents",
"inputs": {
"video": [
"1",
0
]
}
},
"3": {
"class_type": "Moondream31Loader",
"inputs": {
"license_accepted": false,
"device": "Auto",
"max_batch_size": 4,
"kv_cache_profile": "Balanced (8K pages)",
"model_or_adapter": "moondream3.1-9B-A2B"
}
},
"4": {
"class_type": "Moondream31Detect",
"inputs": {
"model": [
"3",
0
],
"image": [
"2",
0
],
"object": "person",
"fps": [
"2",
2
],
"frame_stride": 1,
"parallel_requests": 4,
"max_objects": 100,
"unload_after": false
}
},
"5": {
"class_type": "PreviewImage",
"inputs": {
"images": [
"4",
2
]
}
},
"6": {
"class_type": "ViewText",
"inputs": {
"text": [
"4",
6
]
}
},
"7": {
"class_type": "ViewText",
"inputs": {
"text": [
"4",
1
]
}
}
}
@@ -0,0 +1,74 @@
{
"1": {
"class_type": "LoadImage",
"inputs": {
"image": "moondream_segment_input.png"
}
},
"2": {
"class_type": "Moondream31Loader",
"inputs": {
"license_accepted": false,
"device": "Auto",
"max_batch_size": 4,
"kv_cache_profile": "Balanced (8K pages)",
"model_or_adapter": "moondream3-preview"
}
},
"3": {
"class_type": "Moondream31Segment",
"inputs": {
"model": [
"2",
0
],
"image": [
"1",
0
],
"object": "main foreground object",
"fps": 1.0,
"frame_stride": 1,
"parallel_requests": 1,
"svg_supersample": 4,
"unload_after": false,
"spatial_refs_json": "[]"
}
},
"4": {
"class_type": "PreviewImage",
"inputs": {
"images": [
"3",
4
]
}
},
"5": {
"class_type": "PreviewImage",
"inputs": {
"images": [
"3",
5
]
}
},
"6": {
"class_type": "PreviewImage",
"inputs": {
"images": [
"3",
6
]
}
},
"7": {
"class_type": "ViewText",
"inputs": {
"text": [
"3",
2
]
}
}
}
@@ -0,0 +1,91 @@
{
"1": {
"class_type": "LoadVideo",
"inputs": {
"file": "video_understanding_input.mp4"
}
},
"2": {
"class_type": "GetVideoComponents",
"inputs": {
"video": [
"1",
0
]
}
},
"3": {
"class_type": "VLMVideoTemporalReasoner",
"inputs": {
"frames": [
"2",
0
],
"fps": [
"2",
2
],
"task": "Detailed temporal summary",
"question": "Describe what happens over time and identify the visible evidence.",
"model": "Qwen 3 VL 2B Instruct",
"custom_model_id": "",
"memory_mode": "ComfyUI managed (BF16)",
"max_frames": 16,
"max_events": 24,
"max_new_tokens": 768,
"strategy": "Hybrid: scene + motion + tracks",
"minimum_gap_seconds": 0.15,
"analysis_max_side": 448,
"attention_mode": "Auto (SDPA)",
"enable_thinking": false,
"strict_output": true,
"unload_after": false,
"stream_output": true
}
},
"4": {
"class_type": "ViewText",
"inputs": {
"text": [
"3",
0
]
}
},
"5": {
"class_type": "ViewText",
"inputs": {
"text": [
"3",
6
]
}
},
"6": {
"class_type": "ViewText",
"inputs": {
"text": [
"3",
7
]
}
},
"7": {
"class_type": "ViewText",
"inputs": {
"text": [
"3",
5
]
}
},
"8": {
"class_type": "PreviewImage",
"inputs": {
"images": [
"3",
3
]
}
}
}
@@ -0,0 +1,98 @@
{
"1": {
"class_type": "LoadVideo",
"inputs": {
"file": "vlm_api_people_birds.mp4"
}
},
"2": {
"class_type": "GetVideoComponents",
"inputs": {
"video": [
"1",
0
]
}
},
"3": {
"class_type": "VLMPerformanceProfile",
"inputs": {
"profile": "Fast video"
}
},
"4": {
"class_type": "VLMAdaptiveFrameSampler",
"inputs": {
"frames": [
"2",
0
],
"fps": [
"2",
2
],
"max_frames": [
"3",
0
],
"strategy": "Hybrid: scene + motion + tracks",
"minimum_gap_seconds": 0.15,
"thumbnail_size": 96
}
},
"5": {
"class_type": "VLMImagePixelBudget",
"inputs": {
"images": [
"4",
0
],
"max_megapixels": [
"3",
1
],
"max_edge": [
"3",
2
],
"multiple": "14",
"resize_quality": "Fast (area)"
}
},
"6": {
"class_type": "PreviewImage",
"inputs": {
"images": [
"5",
0
]
}
},
"7": {
"class_type": "ViewText",
"inputs": {
"text": [
"4",
3
]
}
},
"8": {
"class_type": "ViewText",
"inputs": {
"text": [
"5",
3
]
}
},
"9": {
"class_type": "ViewText",
"inputs": {
"text": [
"3",
5
]
}
}
}
+279
View File
@@ -0,0 +1,279 @@
"""Model-agnostic acceleration utilities for image and video VLM workflows.
These nodes reduce visual work *before* it reaches a model. They are therefore
portable across Transformers, llama.cpp, Photon, hosted APIs, CUDA, ROCm, MPS,
XPU, and CPU runtimes. No model is downloaded and no global PyTorch setting is
changed when this module is imported or executed.
"""
from __future__ import annotations
import json
import math
from typing import Any
import torch
import torch.nn.functional as functional
RESIZE_QUALITY = (
"Fast (area)",
"Quality (bicubic)",
)
PERFORMANCE_PROFILES = {
"Live / robotics": {
"max_frames": 24,
"max_megapixels": 0.5,
"max_edge": 896,
"batch_size": 8,
"unload_after": False,
},
"Fast video": {
"max_frames": 48,
"max_megapixels": 0.75,
"max_edge": 1024,
"batch_size": 8,
"unload_after": False,
},
"Balanced": {
"max_frames": 64,
"max_megapixels": 1.0,
"max_edge": 1344,
"batch_size": 4,
"unload_after": False,
},
"High detail": {
"max_frames": 96,
"max_megapixels": 2.0,
"max_edge": 2048,
"batch_size": 2,
"unload_after": False,
},
"Low VRAM handoff": {
"max_frames": 32,
"max_megapixels": 0.75,
"max_edge": 1024,
"batch_size": 1,
"unload_after": True,
},
}
def _json(value: Any) -> str:
return json.dumps(
value,
ensure_ascii=False,
allow_nan=False,
sort_keys=True,
indent=2,
)
def _validate_image_batch(images: torch.Tensor) -> tuple[torch.Tensor, bool]:
if not isinstance(images, torch.Tensor):
raise TypeError("images must be a ComfyUI IMAGE tensor.")
single = images.ndim == 3
value = images.unsqueeze(0) if single else images
if value.ndim != 4:
raise ValueError(
f"Expected an HWC/BHWC or CHW/BCHW IMAGE tensor, got {tuple(images.shape)}."
)
if value.shape[-1] in (1, 3, 4):
return value, single
if value.shape[1] in (1, 3, 4):
return value.permute(0, 2, 3, 1), single
raise ValueError(f"Unsupported image channel shape: {tuple(images.shape)}.")
def optimize_image_pixels(
images: torch.Tensor,
*,
max_megapixels: float,
max_edge: int,
multiple: int,
resize_quality: str,
) -> tuple[torch.Tensor, dict[str, Any]]:
"""Downscale a batch once to a bounded visual-token pixel budget."""
value, single = _validate_image_batch(images)
if not math.isfinite(float(max_megapixels)) or max_megapixels <= 0:
raise ValueError("max_megapixels must be finite and positive.")
if not isinstance(max_edge, int) or max_edge < 32:
raise ValueError("max_edge must be at least 32 pixels.")
if multiple not in {1, 14, 28, 32}:
raise ValueError("multiple must be one of 1, 14, 28, or 32.")
if resize_quality not in RESIZE_QUALITY:
raise ValueError(f"Unknown resize quality {resize_quality!r}.")
height, width = int(value.shape[1]), int(value.shape[2])
pixel_budget = float(max_megapixels) * 1_000_000
scale = min(
1.0,
float(max_edge) / max(width, height),
math.sqrt(pixel_budget / (width * height)),
)
def bounded_dimension(dimension: int) -> int:
target = max(1, math.floor(dimension * scale))
if multiple == 1 or target < multiple:
return target
return max(multiple, (target // multiple) * multiple)
output_width = bounded_dimension(width)
output_height = bounded_dimension(height)
output = value
resized_image = (output_height, output_width) != (height, width)
if resized_image:
nchw = value.permute(0, 3, 1, 2)
if resize_quality == "Fast (area)":
resized = functional.interpolate(
nchw,
size=(output_height, output_width),
mode="area",
)
else:
resized = functional.interpolate(
nchw,
size=(output_height, output_width),
mode="bicubic",
align_corners=False,
antialias=True,
)
output = resized.permute(0, 2, 3, 1).clamp(0.0, 1.0)
report = {
"frames": int(value.shape[0]),
"input_width": width,
"input_height": height,
"output_width": output_width,
"output_height": output_height,
"input_pixels_per_frame": width * height,
"output_pixels_per_frame": output_width * output_height,
"visual_work_reduction": (
(width * height) / max(1, output_width * output_height)
),
"resized": resized_image,
"multiple": multiple,
"quality": resize_quality,
}
if not resized_image:
return images, report
return (output[0] if single else output), report
class VLMPerformanceProfile:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"profile": (
tuple(PERFORMANCE_PROFILES),
{"default": "Balanced"},
)
}
}
RETURN_TYPES = ("INT", "FLOAT", "INT", "INT", "BOOLEAN", "STRING")
RETURN_NAMES = (
"max_frames",
"max_megapixels",
"max_edge",
"batch_size",
"unload_after",
"profile_json",
)
FUNCTION = "profile"
CATEGORY = "VLM Nodes/Performance"
DESCRIPTION = (
"Portable speed/quality presets for the sampler, pixel optimizer, "
"and VLM batch inputs. The profile never changes global runtime state."
)
def profile(self, profile):
values = dict(PERFORMANCE_PROFILES[profile])
values["profile"] = profile
return (
values["max_frames"],
values["max_megapixels"],
values["max_edge"],
values["batch_size"],
values["unload_after"],
_json(values),
)
class VLMImagePixelBudget:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"images": ("IMAGE",),
"max_megapixels": (
"FLOAT",
{"default": 1.0, "min": 0.01, "max": 64.0, "step": 0.05},
),
"max_edge": (
"INT",
{"default": 1344, "min": 32, "max": 16384, "step": 14},
),
"multiple": (
("1", "14", "28", "32"),
{
"default": "14",
"tooltip": (
"14/28 suit common VLM vision patches; 32 suits "
"many detector backbones. Use 1 for arbitrary sizes."
),
},
),
"resize_quality": (
RESIZE_QUALITY,
{"default": "Fast (area)"},
),
}
}
RETURN_TYPES = ("IMAGE", "INT", "INT", "STRING")
RETURN_NAMES = (
"optimized_images",
"width",
"height",
"optimization_report",
)
FUNCTION = "optimize"
CATEGORY = "VLM Nodes/Performance"
DESCRIPTION = (
"Apply one portable pixel budget before any VLM, avoiding repeated "
"high-resolution visual-token work while preserving aspect ratio."
)
def optimize(
self,
images,
max_megapixels,
max_edge,
multiple,
resize_quality,
):
output, report = optimize_image_pixels(
images,
max_megapixels=float(max_megapixels),
max_edge=int(max_edge),
multiple=int(multiple),
resize_quality=resize_quality,
)
return (
output,
report["output_width"],
report["output_height"],
_json(report),
)
NODE_CLASS_MAPPINGS = {
"VLMPerformanceProfile": VLMPerformanceProfile,
"VLMImagePixelBudget": VLMImagePixelBudget,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"VLMPerformanceProfile": "VLM Performance Profile",
"VLMImagePixelBudget": "VLM Image Pixel Budget",
}
+1 -2
View File
@@ -4,11 +4,10 @@ from __future__ import annotations
from pathlib import Path
import folder_paths
import numpy as np
import torch
import folder_paths
from .runtime import (
CachedModelNode,
execution_device,
+1 -1
View File
@@ -5,8 +5,8 @@ from __future__ import annotations
import colorsys
import hashlib
import math
from collections.abc import Iterable, Mapping
from dataclasses import dataclass
from typing import Iterable, Mapping
import numpy as np
import torch
+6 -1
View File
@@ -207,7 +207,12 @@ class OpenVocabularyDetector:
processor = transformers.AutoProcessor.from_pretrained(model_path)
model_class = transformers.AutoModelForZeroShotObjectDetection
dtype = torch_dtype(precision)
model = model_class.from_pretrained(model_path, dtype=dtype)
# Transformers 4.x consumes ``torch_dtype``; 5.x renamed it to
# ``dtype``. Passing the 5.x name to 4.x leaks into the model
# constructor and crashes Grounding DINO at runtime.
major = int(str(transformers.__version__).split(".", 1)[0])
dtype_kwargs = {"dtype": dtype} if major >= 5 else {"torch_dtype": dtype}
model = model_class.from_pretrained(model_path, **dtype_kwargs)
model.eval()
self.spec = spec
self.dtype = dtype
+1826
View File
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -118,7 +118,7 @@ class Joytag(CachedModelNode):
RETURN_TYPES = ("STRING",)
FUNCTION = "tags"
CATEGORY = "VLM Nodes/JoyTag"
CATEGORY = "VLM Nodes/Vision/Tagging"
def tags(
self,
+1 -1
View File
@@ -100,7 +100,7 @@ class Kosmos2model(CachedModelNode):
RETURN_TYPES = ("STRING",)
FUNCTION = "new_model_generate_predictions"
CATEGORY = "VLM Nodes/Kosmos-2"
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
def new_model_generate_predictions(
self,
+1 -1
View File
@@ -126,7 +126,7 @@ class MCLLaVAModel(CachedModelNode):
RETURN_TYPES = ("STRING",)
FUNCTION = "generate_image_description"
CATEGORY = "VLM Nodes/MC-LLaVA"
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
def generate_image_description(
self,
+1 -1
View File
@@ -160,7 +160,7 @@ class MiniCPMNode(CachedModelNode):
RETURN_TYPES = ("STRING",)
FUNCTION = "generate"
CATEGORY = "VLM Nodes/MiniCPM-V"
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
def generate(
self,
+478
View File
@@ -0,0 +1,478 @@
"""MiniMax music generation and cover support with fixed regional routing."""
from __future__ import annotations
import base64
import binascii
import io
import json
import os
from typing import Any
from urllib.parse import urlsplit
import numpy as np
import torch
from .audioldm2 import ANY
from .hosted_api import redact_sensitive
from .runtime import require_module
API_KEY_ENV = "MINIMAX_API_KEY"
REGION_ENDPOINTS = {
"global_en": "https://api.minimax.io/v1/music_generation",
"cn_zh": "https://api.minimaxi.com/v1/music_generation",
}
GENERATION_MODELS = (
"music-3.0",
"music-2.6",
"music-3.0-free",
"music-2.6-free",
)
COVER_MODELS = ("music-cover", "music-cover-free")
MUSIC_MODELS = GENERATION_MODELS + COVER_MODELS
DEFAULT_MODEL = "music-3.0"
REQUEST_FIELDS = frozenset(
{
"model",
"prompt",
"lyrics",
"stream",
"output_format",
"audio_setting",
"lyrics_optimizer",
"is_instrumental",
"audio_url",
"audio_base64",
"cover_feature_id",
}
)
OUTPUT_FORMATS = ("url", "hex")
STREAM_OUTPUT_FORMATS = ("hex",)
AUDIO_FORMATS = ("mp3", "wav", "pcm")
SAMPLE_RATES = (16000, 24000, 32000, 44100)
BITRATES = (32000, 64000, 128000, 256000)
REGIONAL_FIELDS = {"global_en": (), "cn_zh": ("aigc_watermark",)}
STATUS_IN_PROGRESS = 1
STATUS_COMPLETED = 2
MAX_COVER_BYTES = 50 * 1024 * 1024
MAX_AUDIO_BYTES = 128 * 1024 * 1024
def _clean_text(value: object) -> str:
return str(value or "").strip()
def _validate_cover_base64(value: str) -> None:
try:
decoded = base64.b64decode(value, validate=True)
except (binascii.Error, ValueError, TypeError):
raise ValueError("audio_base64 must contain valid base64 data.") from None
if len(decoded) > MAX_COVER_BYTES:
raise ValueError("audio_base64 exceeds the 50 MiB cover input limit.")
def build_music_request(
*,
region: str,
model: str,
prompt: str,
lyrics: str,
stream: bool,
output_format: str,
audio_format: str,
sample_rate: int,
bitrate: int,
lyrics_optimizer: bool,
is_instrumental: bool,
aigc_watermark: bool,
audio_url: str = "",
audio_base64: str = "",
cover_feature_id: str = "",
) -> dict[str, Any]:
"""Validate node inputs and build the documented JSON request body."""
if region not in REGION_ENDPOINTS:
raise ValueError(f"region must be one of {tuple(REGION_ENDPOINTS)}.")
if model not in MUSIC_MODELS:
raise ValueError(f"model must be one of {MUSIC_MODELS}.")
if output_format not in OUTPUT_FORMATS:
raise ValueError(f"output_format must be one of {OUTPUT_FORMATS}.")
if bool(stream) and output_format not in STREAM_OUTPUT_FORMATS:
raise ValueError("Streaming music responses require output_format='hex'.")
if audio_format not in AUDIO_FORMATS:
raise ValueError(f"audio_format must be one of {AUDIO_FORMATS}.")
if int(sample_rate) not in SAMPLE_RATES:
raise ValueError(f"sample_rate must be one of {SAMPLE_RATES}.")
if int(bitrate) not in BITRATES:
raise ValueError(f"bitrate must be one of {BITRATES}.")
clean_prompt = _clean_text(prompt)
clean_lyrics = _clean_text(lyrics)
clean_audio_url = _clean_text(audio_url)
clean_audio_base64 = _clean_text(audio_base64)
clean_cover_feature_id = _clean_text(cover_feature_id)
if len(clean_prompt) > 2000:
raise ValueError("prompt exceeds the 2,000-character music API limit.")
payload: dict[str, Any] = {
"model": model,
"stream": bool(stream),
"output_format": output_format,
"audio_setting": {
"sample_rate": int(sample_rate),
"bitrate": int(bitrate),
"format": audio_format,
},
}
if clean_prompt:
payload["prompt"] = clean_prompt
if clean_lyrics:
payload["lyrics"] = clean_lyrics
if model in COVER_MODELS:
if not 10 <= len(clean_prompt) <= 300:
raise ValueError("Cover generation requires a 10-300 character prompt.")
sources = (clean_audio_url, clean_audio_base64, clean_cover_feature_id)
if sum(bool(value) for value in sources) != 1:
raise ValueError(
"Cover generation requires exactly one of audio_url, "
"audio_base64, or cover_feature_id."
)
if clean_audio_base64:
_validate_cover_base64(clean_audio_base64)
payload["audio_base64"] = clean_audio_base64
elif clean_audio_url:
payload["audio_url"] = clean_audio_url
else:
if not 10 <= len(clean_lyrics) <= 1000:
raise ValueError(
"cover_feature_id requires lyrics between 10 and 1,000 characters."
)
payload["cover_feature_id"] = clean_cover_feature_id
if clean_lyrics and not 10 <= len(clean_lyrics) <= 1000:
raise ValueError("Cover lyrics must be between 10 and 1,000 characters.")
else:
if any((clean_audio_url, clean_audio_base64, clean_cover_feature_id)):
raise ValueError("Cover audio fields require a cover model.")
if len(clean_lyrics) > 3500:
raise ValueError("lyrics exceeds the 3,500-character music API limit.")
if bool(is_instrumental) and not clean_prompt:
raise ValueError("Instrumental generation requires a prompt.")
if not bool(is_instrumental) and not clean_lyrics and not bool(lyrics_optimizer):
raise ValueError(
"Non-instrumental generation requires lyrics or lyrics_optimizer."
)
payload["lyrics_optimizer"] = bool(lyrics_optimizer)
payload["is_instrumental"] = bool(is_instrumental)
if region == "cn_zh":
payload["aigc_watermark"] = bool(aigc_watermark)
return payload
def _response_parts(payload: object) -> tuple[str, int, dict[str, Any]]:
if not isinstance(payload, dict):
raise RuntimeError("MiniMax returned a non-object music response.")
base_response = payload.get("base_resp")
if not isinstance(base_response, dict):
raise RuntimeError("MiniMax returned no base_resp status.")
try:
success_code = int(base_response.get("status_code"))
except (TypeError, ValueError):
raise RuntimeError("MiniMax returned an invalid base_resp status code.") from None
if success_code != 0:
message = _clean_text(base_response.get("status_msg")) or "unknown API error"
raise RuntimeError(f"MiniMax music API error {success_code}: {message}")
data = payload.get("data")
if not isinstance(data, dict):
raise RuntimeError("MiniMax returned no music data object.")
try:
status = int(data.get("status"))
except (TypeError, ValueError):
raise RuntimeError("MiniMax returned an invalid music status.") from None
if status not in {STATUS_IN_PROGRESS, STATUS_COMPLETED}:
raise RuntimeError(f"MiniMax returned unsupported music status {status}.")
audio = data.get("audio", "")
if not isinstance(audio, str):
raise RuntimeError("MiniMax returned a non-string audio value.")
extra_info = payload.get("extra_info")
return audio.strip(), status, extra_info if isinstance(extra_info, dict) else {}
def _stream_audio(response: Any) -> tuple[str, dict[str, Any]]:
audio = ""
extra_info: dict[str, Any] = {}
completed = False
saw_payload = False
for line in response.iter_lines():
raw = line.decode("utf-8") if isinstance(line, bytes) else str(line)
raw = raw.strip()
if raw.startswith("data:"):
raw = raw[5:].strip()
if not raw or raw == "[DONE]" or raw.startswith("event:"):
continue
try:
payload = json.loads(raw)
except json.JSONDecodeError:
raise RuntimeError("MiniMax returned invalid streaming JSON.") from None
chunk, status, metadata = _response_parts(payload)
saw_payload = True
if chunk:
if chunk.startswith(audio):
audio = chunk
elif not audio.startswith(chunk):
audio += chunk
if metadata:
extra_info = metadata
completed = completed or status == STATUS_COMPLETED
if not saw_payload:
raise RuntimeError("MiniMax returned an empty streaming response.")
if not completed:
raise RuntimeError("MiniMax streaming ended before music generation completed.")
if not audio:
raise RuntimeError("MiniMax returned no audio data.")
return audio, extra_info
def _request_audio_value(
client: Any,
endpoint: str,
headers: dict[str, str],
payload: dict[str, Any],
) -> tuple[str, dict[str, Any]]:
if payload["stream"]:
with client.stream("POST", endpoint, headers=headers, json=payload) as response:
response.raise_for_status()
return _stream_audio(response)
response = client.post(endpoint, headers=headers, json=payload)
response.raise_for_status()
audio, status, extra_info = _response_parts(response.json())
if status != STATUS_COMPLETED:
raise RuntimeError(
"MiniMax music generation is still in progress and has no query endpoint."
)
if not audio:
raise RuntimeError("MiniMax returned no audio data.")
return audio, extra_info
def _download_audio(client: Any, url: str) -> bytes:
parsed = urlsplit(url)
if (
parsed.scheme != "https"
or not parsed.hostname
or parsed.username is not None
or parsed.password is not None
):
raise RuntimeError("MiniMax returned an invalid HTTPS audio URL.")
chunks: list[bytes] = []
total = 0
with client.stream("GET", url) as response:
response.raise_for_status()
for chunk in response.iter_bytes():
total += len(chunk)
if total > MAX_AUDIO_BYTES:
raise RuntimeError("MiniMax audio download exceeds 128 MiB.")
chunks.append(chunk)
return b"".join(chunks)
def _audio_bytes(client: Any, value: str, output_format: str) -> bytes:
if output_format == "url":
return _download_audio(client, value)
try:
return bytes.fromhex("".join(value.split()))
except ValueError:
raise RuntimeError("MiniMax returned invalid hexadecimal audio data.") from None
def _metadata_integer(metadata: dict[str, Any], name: str, default: int) -> int:
try:
value = int(metadata.get(name, default))
except (TypeError, ValueError):
return int(default)
return value if value > 0 else int(default)
def _decode_audio(
content: bytes,
audio_format: str,
requested_sample_rate: int,
metadata: dict[str, Any],
) -> tuple[np.ndarray, int]:
if not content:
raise RuntimeError("MiniMax returned an empty audio payload.")
if audio_format == "pcm":
if len(content) % 2:
raise RuntimeError("MiniMax returned an odd-length PCM payload.")
channels = _metadata_integer(metadata, "music_channel", 1)
raw = np.frombuffer(content, dtype="<i2")
if raw.size % channels:
raise RuntimeError("MiniMax PCM samples do not align with the channel count.")
samples = raw.astype(np.float32).reshape(-1, channels) / 32768.0
sample_rate = _metadata_integer(
metadata,
"music_sample_rate",
requested_sample_rate,
)
else:
soundfile = require_module("soundfile", "soundfile>=0.12")
try:
samples, sample_rate = soundfile.read(
io.BytesIO(content),
dtype="float32",
always_2d=True,
)
except Exception as exc:
detail = redact_sensitive(exc)
raise RuntimeError(f"Could not decode MiniMax {audio_format} audio: {detail}") from None
samples = np.asarray(samples, dtype=np.float32)
sample_rate = int(sample_rate)
if samples.ndim != 2 or not samples.size:
raise RuntimeError("MiniMax decoded to an empty audio array.")
if not np.isfinite(samples).all():
raise RuntimeError("MiniMax decoded audio contains non-finite samples.")
return np.ascontiguousarray(samples), int(sample_rate)
class MiniMaxMusicNode:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"region": (tuple(REGION_ENDPOINTS), {"default": "global_en"}),
"model": (MUSIC_MODELS, {"default": DEFAULT_MODEL}),
"prompt": ("STRING", {"default": "", "multiline": True}),
"lyrics": ("STRING", {"default": "", "multiline": True}),
"stream": ("BOOLEAN", {"default": False}),
"output_format": (OUTPUT_FORMATS, {"default": "hex"}),
"audio_format": (AUDIO_FORMATS, {"default": "mp3"}),
"sample_rate": (SAMPLE_RATES, {"default": 44100}),
"bitrate": (BITRATES, {"default": 256000}),
"lyrics_optimizer": ("BOOLEAN", {"default": False}),
"is_instrumental": ("BOOLEAN", {"default": False}),
"aigc_watermark": (
"BOOLEAN",
{
"default": False,
"tooltip": "Sent only to the cn_zh endpoint.",
},
),
},
"optional": {
"audio_url": ("STRING", {"default": ""}),
"audio_base64": ("STRING", {"default": "", "multiline": True}),
"cover_feature_id": ("STRING", {"default": ""}),
"timeout_seconds": (
"FLOAT",
{"default": 600.0, "min": 1.0, "max": 1800.0},
),
"use_system_proxy": ("BOOLEAN", {"default": False}),
},
}
RETURN_NAMES = ("wave_form", "sample_rate", "audio")
RETURN_TYPES = (ANY, "INT", "AUDIO")
OUTPUT_NODE = True
FUNCTION = "generate_music"
CATEGORY = "VLM Nodes/Audio"
DESCRIPTION = (
"Generate music or covers through fixed MiniMax regional endpoints. "
f"The API key is read only from {API_KEY_ENV}."
)
def generate_music(
self,
region,
model,
prompt,
lyrics,
stream,
output_format,
audio_format,
sample_rate,
bitrate,
lyrics_optimizer,
is_instrumental,
aigc_watermark,
audio_url="",
audio_base64="",
cover_feature_id="",
timeout_seconds=600.0,
use_system_proxy=False,
):
payload = build_music_request(
region=region,
model=model,
prompt=prompt,
lyrics=lyrics,
stream=stream,
output_format=output_format,
audio_format=audio_format,
sample_rate=sample_rate,
bitrate=bitrate,
lyrics_optimizer=lyrics_optimizer,
is_instrumental=is_instrumental,
aigc_watermark=aigc_watermark,
audio_url=audio_url,
audio_base64=audio_base64,
cover_feature_id=cover_feature_id,
)
api_key = os.getenv(API_KEY_ENV, "").strip()
if not api_key:
raise ValueError(
f"Set {API_KEY_ENV} in the environment that starts ComfyUI, "
"then restart the server."
)
httpx = require_module("httpx", "httpx>=0.27,<1")
client = httpx.Client(
timeout=max(1.0, min(1800.0, float(timeout_seconds))),
follow_redirects=False,
trust_env=bool(use_system_proxy),
)
headers = {
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
}
try:
value, metadata = _request_audio_value(
client,
REGION_ENDPOINTS[region],
headers,
payload,
)
content = _audio_bytes(client, value, output_format)
samples, actual_rate = _decode_audio(
content,
audio_format,
int(sample_rate),
metadata,
)
legacy = samples[:, 0] if samples.shape[1] == 1 else samples
audio = {
"waveform": torch.from_numpy(samples.T.copy()).unsqueeze(0),
"sample_rate": actual_rate,
}
return (legacy.tolist(), actual_rate, audio)
except Exception as exc:
detail = redact_sensitive(exc, (api_key,))
raise RuntimeError(f"MiniMax music request failed: {detail}") from None
finally:
try:
client.close()
except Exception:
pass
NODE_CLASS_MAPPINGS = {"MiniMaxMusicNode": MiniMaxMusicNode}
NODE_DISPLAY_NAME_MAPPINGS = {"MiniMaxMusicNode": "MiniMax Music"}
__all__ = [
"MiniMaxMusicNode",
"NODE_CLASS_MAPPINGS",
"NODE_DISPLAY_NAME_MAPPINGS",
]
+126 -29
View File
@@ -8,8 +8,9 @@ small and large VLM families while keeping downloads and VRAM allocation lazy.
from __future__ import annotations
import threading
from collections.abc import Callable
from dataclasses import dataclass
from typing import Any, Callable
from typing import Any
import torch
@@ -25,13 +26,14 @@ from .runtime import (
model_device,
move_inputs,
normalize_hf_model_id,
require_quantization_backend,
require_module,
require_quantization_backend,
reserve_external_vram,
snapshot_download,
tensor_batch_to_pil,
torch_dtype,
)
from .vision_types import VLM_VIDEO_SELECTION, VideoFrameSelection
@dataclass(frozen=True)
@@ -190,6 +192,24 @@ MODEL_CATALOG = {
),
}
RECOMMENDED_MODEL_LABELS = (
"Qwen 3.5 0.8B (fastest current)",
"Qwen 3.5 4B (recommended)",
"Qwen 3 VL 2B Instruct",
"Qwen 3 VL 4B Instruct",
"Qwen 3 VL 8B Instruct",
"SmolVLM2 500M Video (low VRAM)",
"SmolVLM2 2.2B Video",
"LFM2.5 VL 450M (edge)",
"InternVL 3.5 1B HF",
"Granite Vision 4.1 4B (structured documents)",
"Gemma 3 4B IT (license acceptance required)",
"Custom Hugging Face model",
)
LEGACY_MODEL_LABELS = tuple(
label for label in MODEL_CATALOG if label not in RECOMMENDED_MODEL_LABELS
)
MEMORY_MODES = (
"ComfyUI managed (BF16)",
"4-bit NF4 (bitsandbytes)",
@@ -445,6 +465,7 @@ class ModernVLMPredictor:
fps: float = 1.0,
enable_thinking: bool = False,
stream_callback: Callable[[str], None] | None = None,
video_selection: VideoFrameSelection | None = None,
) -> str:
primary_images = (
tensor_batch_to_pil(images) if images is not None else []
@@ -461,6 +482,27 @@ class ModernVLMPredictor:
f"{self.spec.family} does not advertise video support. "
"Disconnect video_frames or select Qwen/SmolVLM2."
)
if video_selection is not None:
if video is None:
raise ValueError(
"video_selection requires a connected video_frames batch."
)
if not isinstance(video_selection, VideoFrameSelection):
raise TypeError("video_selection must be a VLM Video Selection.")
if len(video_selection.frames) != len(video):
raise ValueError(
"video_selection frame count must match video_frames."
)
source_aspect = video_selection.width / video_selection.height
analysis_aspect = video[0].width / video[0].height
if abs(source_aspect - analysis_aspect) > max(
0.01,
source_aspect * 0.01,
):
raise ValueError(
"video_selection and video_frames must have the same "
"aspect ratio."
)
results = []
# A connected video is the primary visual input. Including ComfyUI's
@@ -483,25 +525,49 @@ class ModernVLMPredictor:
if video is not None
else [{"type": "image", "image": image}]
)
effective_prompt = (
f"The video frames are sampled at {float(fps):g} FPS.\n\n{prompt}"
if video is not None
else prompt
)
if video is not None and video_selection is not None:
timeline = ", ".join(
f"{position}=frame {frame.source_frame_index} "
f"at {frame.timestamp:.6f}s"
for position, frame in enumerate(video_selection.frames)
)
effective_prompt = (
"The supplied video images are irregular samples from one "
f"{video_selection.source_frame_count}-frame video at "
f"{video_selection.fps:g} FPS. Supplied-image mapping: "
f"{timeline}.\n\n{prompt}"
)
elif video is not None:
effective_prompt = (
f"The video frames are sampled at {float(fps):g} FPS.\n\n"
f"{prompt}"
)
else:
effective_prompt = prompt
content.append({"type": "text", "text": effective_prompt})
messages.append({"role": "user", "content": content})
metadata = None
if video is not None:
frame_rate = float(fps)
metadata = {
"total_num_frames": len(video),
"fps": frame_rate,
"duration": len(video) / frame_rate,
"frames_indices": list(range(len(video))),
"width": video[0].width,
"height": video[0].height,
}
if video_selection is not None:
metadata = {
"total_num_frames": video_selection.source_frame_count,
"fps": video_selection.fps,
"duration": video_selection.duration,
"frames_indices": list(video_selection.indices),
"width": video[0].width,
"height": video[0].height,
}
else:
frame_rate = float(fps)
metadata = {
"total_num_frames": len(video),
"fps": frame_rate,
"duration": len(video) / frame_rate,
"frames_indices": list(range(len(video))),
"width": video[0].width,
"height": video[0].height,
}
inputs = self._inputs(
messages,
enable_thinking,
@@ -607,7 +673,7 @@ class ModernVLM(CachedModelNode):
},
),
"model": (
list(MODEL_CATALOG),
list(RECOMMENDED_MODEL_LABELS),
{"default": "Qwen 3 VL 2B Instruct"},
),
"custom_model_id": ("STRING", {"default": ""}),
@@ -638,6 +704,7 @@ class ModernVLM(CachedModelNode):
},
),
"video_frames": ("IMAGE",),
"video_selection": (VLM_VIDEO_SELECTION,),
"fps": (
"FLOAT",
{"default": 1.0, "min": 0.1, "max": 60.0, "step": 0.1},
@@ -666,6 +733,15 @@ class ModernVLM(CachedModelNode):
FUNCTION = "run"
CATEGORY = "VLM Nodes/Modern"
@classmethod
def VALIDATE_INPUTS(cls, model):
# The visible combo is deliberately curated. Accepting every known
# catalog value here keeps workflows saved before the curation fully
# executable even when their model now lives under Legacy.
if model not in MODEL_CATALOG:
return f"Unsupported Modern VLM model {model!r}."
return True
def run(
self,
prompt,
@@ -678,6 +754,7 @@ class ModernVLM(CachedModelNode):
image=None,
system_prompt="You are an expert visual analyst.",
video_frames=None,
video_selection=None,
fps=1.0,
attention_mode="Auto (SDPA)",
enable_thinking=False,
@@ -705,25 +782,45 @@ class ModernVLM(CachedModelNode):
try:
return (
predictor.generate(
image,
prompt,
system_prompt,
max_new_tokens,
temperature,
top_p,
video_frames,
fps,
enable_thinking,
stream_callback,
images=image,
prompt=prompt,
system_prompt=system_prompt,
max_new_tokens=max_new_tokens,
temperature=temperature,
top_p=top_p,
video_frames=video_frames,
fps=fps,
video_selection=video_selection,
enable_thinking=enable_thinking,
stream_callback=stream_callback,
),
)
finally:
self.maybe_clear_model(unload_after)
NODE_CLASS_MAPPINGS = {"ModernVLM": ModernVLM}
class LegacyModernVLM(ModernVLM):
"""Compatibility surface for redundant, superseded, and very large tiers."""
@classmethod
def INPUT_TYPES(cls):
inputs = super().INPUT_TYPES()
inputs["required"]["model"] = (
list(LEGACY_MODEL_LABELS),
{"default": LEGACY_MODEL_LABELS[0]},
)
return inputs
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
NODE_CLASS_MAPPINGS = {
"ModernVLM": ModernVLM,
"LegacyModernVLM": LegacyModernVLM,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ModernVLM": (
"Modern VLM (Qwen / SmolVLM2 / LFM / InternVL / Granite / Gemma)"
)
),
"LegacyModernVLM": "[Legacy] Modern VLM Compatibility",
}
+2 -2
View File
@@ -14,8 +14,8 @@ from .runtime import (
external_device_map,
inference_context,
model_device,
require_quantization_backend,
require_module,
require_quantization_backend,
reserve_external_vram,
snapshot_download,
tensor_batch_to_pil,
@@ -155,7 +155,7 @@ class MolmoNode(CachedModelNode):
RETURN_TYPES = ("STRING",)
FUNCTION = "generate"
CATEGORY = "VLM Nodes/Molmo"
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
def generate(
self,
+69 -32
View File
@@ -1,7 +1,21 @@
"""Current Moondream 2 node using the model's supported query API."""
"""Current Moondream 2 node using the model's supported query API.
The pinned checkpoint was authored against Transformers 4.52.4. Loading it
through Transformers 5's ``from_pretrained`` compatibility path can silently
produce an all-EOS model even when every tensor is reported as loaded. The
checkpoint itself is a normal safetensors state dict, so instantiate its
official wrapper and load that state dict directly. This keeps Moondream in
ComfyUI's managed VRAM lifecycle without downgrading Transformers for the rest
of the node pack.
"""
from __future__ import annotations
import importlib
import sys
from pathlib import Path
from types import ModuleType
import torch
from .runtime import (
@@ -18,12 +32,59 @@ from .runtime import (
MODEL_ID = "vikhyatk/moondream2"
MODEL_REVISION = "2025-06-21"
_CHECKPOINT_PACKAGE = "_comfyui_vlm_moondream2_checkpoint"
def _checkpoint_module(model_path: str | Path):
"""Import the checkpoint's relative modules without HF's generated cache.
Hugging Face's dynamic-module cache can omit transitive relative imports
for a local snapshot. Giving the snapshot a private package namespace lets
Python resolve the checkpoint's own ``.config``, ``.vision``, and related
modules directly and deterministically.
"""
source = str(Path(model_path).resolve())
package = sys.modules.get(_CHECKPOINT_PACKAGE)
if package is None:
package = ModuleType(_CHECKPOINT_PACKAGE)
package.__path__ = [source]
package.__package__ = _CHECKPOINT_PACKAGE
sys.modules[_CHECKPOINT_PACKAGE] = package
elif list(getattr(package, "__path__", ())) != [source]:
raise RuntimeError(
"Moondream2 checkpoint source changed inside a running process. "
"Restart ComfyUI before loading a different snapshot."
)
return importlib.import_module(f"{_CHECKPOINT_PACKAGE}.hf_moondream")
def _load_native_checkpoint(model_path: str | Path):
checkpoint = _checkpoint_module(model_path)
safetensors = require_module("safetensors.torch")
config = checkpoint.HfConfig.from_pretrained(
model_path,
local_files_only=True,
)
model = checkpoint.HfMoondream(config)
weights = Path(model_path) / "model.safetensors"
if not weights.is_file():
raise FileNotFoundError(f"Moondream2 weights are missing: {weights}")
missing, unexpected = safetensors.load_model(
model,
str(weights),
strict=True,
)
if missing or unexpected:
raise RuntimeError(
"Moondream2 checkpoint did not load exactly: "
f"missing={sorted(missing)}, unexpected={sorted(unexpected)}"
)
return model.eval()
class Moondream2Predictor:
def __init__(self):
transformers = require_module("transformers")
dynamic_modules = require_module("transformers.dynamic_module_utils")
model_path = snapshot_download(
MODEL_ID,
"moondream2",
@@ -31,31 +92,7 @@ class Moondream2Predictor:
ignore_patterns=["*.bin", "*.gguf"],
)
self.dtype = torch_dtype("bfloat16")
config = transformers.AutoConfig.from_pretrained(
model_path,
revision=MODEL_REVISION,
trust_remote_code=True,
)
remote_class = dynamic_modules.get_class_from_dynamic_module(
"hf_moondream.HfMoondream",
model_path,
local_files_only=True,
)
class Transformers5Moondream(remote_class):
def __init__(self, model_config):
super().__init__(model_config)
# The pinned remote wrapper predates the Transformers 5 model
# loader and does not declare its tied-weight metadata. Calling
# the full post_init would reinitialize custom Moondream state.
self.all_tied_weights_keys = {}
model = Transformers5Moondream.from_pretrained(
model_path,
config=config,
dtype=self.dtype,
)
model.eval()
model = _load_native_checkpoint(model_path)
self.handle = ManagedTorchModel(model)
def close(self):
@@ -92,9 +129,9 @@ class Moondream2Predictor:
response = response.get("answer", response)
if not str(response).strip():
raise RuntimeError(
"Moondream2 returned an empty response on this "
"Torch/Transformers build. Use the Modern VLM node with "
"LFM2.5-VL 450M, InternVL 3.5 1B, or Qwen3-VL 2B."
"Moondream2 returned an empty response. Verify that the "
f"{MODEL_REVISION} snapshot is complete, then restart "
"ComfyUI so its checkpoint modules are reloaded."
)
results.append(str(response))
return batch_text(results)
@@ -134,7 +171,7 @@ class Moondream2model(CachedModelNode):
RETURN_TYPES = ("STRING",)
FUNCTION = "moondream2_generate_predictions"
CATEGORY = "VLM Nodes/Moondream2"
CATEGORY = "VLM Nodes/Modern/Edge"
def moondream2_generate_predictions(
self,
+1645
View File
File diff suppressed because it is too large Load Diff
+388
View File
@@ -0,0 +1,388 @@
"""Isolated Moondream 3.1 Photon worker.
This file is launched directly by the ComfyUI process with the dedicated
Moondream virtual environment. It intentionally has no imports from ComfyUI
or this package: Moondream pins a Pillow version that is incompatible with
current ComfyUI releases, so sharing one Python environment is unsafe.
"""
from __future__ import annotations
import argparse
import os
import platform
import sys
import time
import traceback
from collections.abc import Callable
from concurrent.futures import ThreadPoolExecutor
from dataclasses import replace
from importlib.metadata import PackageNotFoundError, version
from io import BytesIO
from multiprocessing.connection import Client
from typing import Any
from PIL import Image
def _honor_do_not_track() -> bool:
"""Disable anonymous Photon reporting when the sidecar requests privacy.
Kestrel 0.4.2 does not currently inspect the conventional DO_NOT_TRACK
environment variable. Base-model inference does not need its reporter, so
keep validation local, skip the telemetry loop, and still close the HTTP
client during engine shutdown. Finetune inference retains upstream auth
and reporting behavior because it explicitly receives an API key.
"""
if os.environ.get("DO_NOT_TRACK") != "1":
return False
if os.environ.get("MOONDREAM_API_KEY", "").strip():
return False
from kestrel.photon import PhotonReporter
async def validate_api_key(self) -> bool:
return False
def start(self) -> None:
return None
async def shutdown(self) -> None:
await self._client.aclose()
PhotonReporter.validate_api_key = validate_api_key
PhotonReporter.start = start
PhotonReporter.shutdown = shutdown
return True
def _register_moondream31_if_needed(model_name: str) -> bool:
"""Bridge the official model-card ID on runtimes released before the ID.
Moondream 3.1 uses the same MD3 Photon runtime/checkpoint format as the
preview. Stable moondream 1.3.0 / kestrel 0.4.2 shipped the safetensors
loader but omitted the new registry entry published by the later model
card. Prefer an upstream entry whenever present; otherwise clone only the
runtime metadata and point it at the official 3.1 weights.
"""
if model_name != "moondream3.1-9B-A2B":
return False
from kestrel.models import get_spec, register
try:
get_spec(model_name)
return False
except ValueError:
preview = get_spec("moondream3-preview")
register(
replace(
preview,
name=model_name,
repo_id="moondream/moondream3.1-9B-A2B",
filename="model.safetensors",
checkpoint_format="md3",
)
)
return True
def _base_model_name(value: str) -> str:
return str(value).split("/", 1)[0]
def _model_skills(model_name: str) -> frozenset[str]:
base_model = _base_model_name(model_name)
if base_model == "moondream3.1-9B-A2B":
# Source of truth: the final 3.1 model card. Segment remains a skill
# of the 3 Preview and cloud API, not the final local 3.1 checkpoint.
return frozenset(("caption", "query", "detect", "point"))
from kestrel.models import get_spec
spec = get_spec(base_model)
templates = spec.default_config.get("tokenizer", {}).get("templates", {})
return frozenset(
name for name, template in templates.items() if template is not None
)
def _image(value: bytes) -> Image.Image:
if not isinstance(value, bytes):
raise TypeError("Worker image payloads must be bytes.")
with Image.open(BytesIO(value)) as source:
return source.convert("RGB")
def _parallel(
images: list[bytes],
operation: Callable[[Image.Image], dict[str, Any]],
workers: int,
) -> list[dict[str, Any]]:
if not images:
return []
worker_count = max(1, min(int(workers), len(images)))
with ThreadPoolExecutor(max_workers=worker_count) as pool:
return list(pool.map(lambda value: operation(_image(value)), images))
def _private_shutdown(model: Any) -> None:
"""Best-effort graceful Photon shutdown before the process exits.
The public moondream package currently has no close method. Process
isolation remains the hard guarantee: the parent terminates this exact
process if this best-effort private cleanup ever changes or stalls.
"""
engine = getattr(model, "_engine", None)
loop = getattr(model, "_loop", None)
thread = getattr(model, "_thread", None)
if engine is not None and loop is not None:
try:
import asyncio
asyncio.run_coroutine_threadsafe(engine.shutdown(), loop).result(timeout=20)
except Exception: # noqa: BLE001,S110 - private best-effort cleanup.
pass
try:
loop.call_soon_threadsafe(loop.stop)
except Exception: # noqa: BLE001,S110 - private best-effort cleanup.
pass
if thread is not None:
try:
thread.join(timeout=5)
except Exception: # noqa: BLE001,S110 - private best-effort cleanup.
pass
def _request(
model: Any,
request: dict[str, Any],
send: Callable[[dict[str, Any]], None],
max_batch_size: int,
supported_skills: frozenset[str],
) -> bool:
request_id = request.get("id")
operation = request.get("operation")
if operation == "shutdown":
send({"id": request_id, "type": "result", "result": {"closed": True}})
return False
if operation not in supported_skills:
raise ValueError(
f"Model does not support the {operation!r} skill. "
f"Available skills: {', '.join(sorted(supported_skills))}."
)
started = time.perf_counter()
settings = {"max_tokens": int(request.get("max_tokens", 512))}
if operation in {"query", "caption"}:
image_payload = request.get("image")
image = _image(image_payload) if image_payload is not None else None
if operation == "query":
output = model.query(
image=image,
question=str(request["question"]),
stream=bool(request.get("stream", True)),
settings=settings,
reasoning=bool(request.get("reasoning", False)),
)
key = "answer"
else:
if image is None:
raise ValueError("Caption requires an image.")
output = model.caption(
image=image,
length=str(request.get("length", "normal")),
stream=bool(request.get("stream", True)),
settings=settings,
)
key = "caption"
value = output[key]
if isinstance(value, str):
text = value
else:
chunks = []
for chunk in value:
chunk_text = str(chunk)
chunks.append(chunk_text)
send(
{
"id": request_id,
"type": "chunk",
"text": chunk_text,
}
)
text = "".join(chunks)
result = {
key: text,
"elapsed_seconds": time.perf_counter() - started,
}
if operation == "query" and output.get("reasoning") is not None:
result["reasoning"] = output["reasoning"]
send({"id": request_id, "type": "result", "result": result})
return True
images = request.get("images")
if not isinstance(images, list):
raise TypeError(f"{operation} requires an image list.")
workers = min(
max_batch_size,
max(1, int(request.get("parallel_requests", max_batch_size))),
)
object_prompt = str(request.get("object", "")).strip()
if not object_prompt:
raise ValueError(f"{operation} requires a non-empty object prompt.")
if operation == "detect":
results = _parallel(
images,
lambda image: model.detect(image, object_prompt, settings=settings),
workers,
)
elif operation == "point":
results = _parallel(
images,
lambda image: model.point(image, object_prompt, settings=settings),
workers,
)
elif operation == "segment":
spatial_refs = request.get("spatial_refs") or None
results = _parallel(
images,
lambda image: model.segment(
image,
object_prompt,
spatial_refs=spatial_refs,
stream=False,
settings=settings,
),
workers,
)
else:
raise ValueError(f"Unknown worker operation {operation!r}.")
send(
{
"id": request_id,
"type": "result",
"result": {
"items": results,
"processed_frames": len(images),
"parallel_requests": workers,
"elapsed_seconds": time.perf_counter() - started,
},
}
)
return True
def main() -> int:
parser = argparse.ArgumentParser()
parser.add_argument("--host", default="127.0.0.1")
parser.add_argument("--port", type=int, required=True)
parser.add_argument("--auth-key")
parser.add_argument("--model", required=True)
parser.add_argument("--device", required=True)
parser.add_argument("--max-batch-size", type=int, required=True)
parser.add_argument("--kv-cache-pages", type=int, default=0)
args = parser.parse_args()
auth_key = args.auth_key or os.environ.pop("MOONDREAM_WORKER_AUTH", "")
if not auth_key:
parser.error("worker authentication is missing")
connection = Client(
(args.host, args.port),
authkey=bytes.fromhex(auth_key),
)
def send(value: dict[str, Any]) -> None:
connection.send(value)
send(
{
"type": "status",
"status": "loading",
"python": sys.version.split()[0],
"platform": platform.platform(),
"pid": os.getpid(),
}
)
model = None
try:
import moondream as md
base_model = _base_model_name(args.model)
compatibility_registration = _register_moondream31_if_needed(base_model)
telemetry_disabled = _honor_do_not_track()
supported_skills = _model_skills(args.model)
kwargs: dict[str, Any] = {
"local": True,
"model": args.model,
"device": args.device,
"max_batch_size": args.max_batch_size,
}
if args.kv_cache_pages > 0:
kwargs["kv_cache_pages"] = args.kv_cache_pages
model = md.vl(**kwargs)
try:
package_version = version("moondream")
except PackageNotFoundError:
package_version = "unknown"
send(
{
"type": "status",
"status": "ready",
"moondream_version": package_version,
"compatibility_registration": compatibility_registration,
"telemetry_disabled": telemetry_disabled,
"skills": sorted(supported_skills),
"pid": os.getpid(),
}
)
running = True
while running:
request = connection.recv()
request_id = request.get("id") if isinstance(request, dict) else None
try:
if not isinstance(request, dict):
raise TypeError("Worker requests must be dictionaries.")
running = _request(
model,
request,
send,
args.max_batch_size,
supported_skills,
)
except Exception as exc: # noqa: BLE001 - report request failures over IPC.
send(
{
"id": request_id,
"type": "error",
"error": f"{type(exc).__name__}: {exc}",
"traceback": traceback.format_exc(limit=12),
}
)
except Exception as exc: # noqa: BLE001 - report startup failures over IPC.
send(
{
"type": "fatal",
"error": f"{type(exc).__name__}: {exc}",
"traceback": traceback.format_exc(limit=20),
}
)
return 1
finally:
if model is not None:
_private_shutdown(model)
try:
connection.close()
except OSError:
pass
return 0
if __name__ == "__main__":
raise SystemExit(main())
+1 -1
View File
@@ -25,7 +25,7 @@ class MoonDream(CachedModelNode):
RETURN_TYPES = ("STRING",)
FUNCTION = "answer_questions"
CATEGORY = "VLM Nodes/MoonDream"
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
def answer_questions(self, image, question, unload_after=False):
predictor = self.get_or_create_model(
+2 -3
View File
@@ -24,15 +24,14 @@ from .runtime import (
normalize_hf_model_id,
pil_mask_to_tensor,
pil_to_tensor,
require_quantization_backend,
require_module,
require_quantization_backend,
reserve_external_vram,
snapshot_download,
tensor_batch_to_pil,
torch_dtype,
)
PALIGEMMA_MODELS = [
"gokaygokay/sd3-long-captioner-v2",
"google/paligemma-3b-ft-refcoco-seg-896",
@@ -263,7 +262,7 @@ class Paligemma(CachedModelNode):
RETURN_TYPES = ("STRING", "MASK", "IMAGE")
RETURN_NAMES = ("description", "mask", "visualization")
FUNCTION = "process_task"
CATEGORY = "VLM Nodes/Paligemma"
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
def process_task(
self,
+1 -1
View File
@@ -51,4 +51,4 @@ Optional: If asked to create a random prompt create one.
# Define the system message
system_msg_simple = """
You are an helpful asistant. Answer optional questions or help the user for their optional queries.
"""
"""
+2 -3
View File
@@ -17,15 +17,14 @@ from .runtime import (
inference_context,
model_device,
move_inputs,
require_quantization_backend,
require_module,
require_quantization_backend,
reserve_external_vram,
snapshot_download,
tensor_batch_to_pil,
torch_dtype,
)
QWEN2_VL_MODELS = {
"Qwen2-VL-2B": "Qwen/Qwen2-VL-2B-Instruct",
"Qwen2-VL-7B": "Qwen/Qwen2-VL-7B-Instruct",
@@ -347,7 +346,7 @@ class Qwen2VLNode(CachedModelNode):
RETURN_TYPES = ("STRING",)
FUNCTION = "generate"
CATEGORY = "VLM Nodes/Qwen2-VL"
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
def generate(
self,
+2457
View File
File diff suppressed because it is too large Load Diff
+93 -36
View File
@@ -18,11 +18,12 @@ import os
import platform
import re
import threading
from collections.abc import Callable, Iterable, Mapping
from contextlib import nullcontext
from dataclasses import dataclass
from importlib import metadata
from pathlib import Path
from typing import Any, Callable, Iterable, Mapping
from typing import Any
import folder_paths
import numpy as np
@@ -31,6 +32,7 @@ from PIL import Image
LOGGER = logging.getLogger("ComfyUI_VLM_nodes")
GGUF_EXTENSIONS = {".gguf"}
PIL_CONVERSION_CHUNK_BYTES = 128 * 1024**2
LLAMA_FLASH_ATTENTION_CHOICES = ("Auto", "Enabled", "Disabled")
LLAMA_SPLIT_MODE_CHOICES = ("Layer", "Row", "Single GPU")
LLAMA_VISION_HANDLER_CHOICES = (
@@ -157,39 +159,67 @@ def hf_download(repo_id: str, filename: str, subdirectory: str, **kwargs: Any) -
return Path(hub.hf_hub_download(**download_kwargs))
def _tensor_image_batch_to_uint8(images: torch.Tensor) -> np.ndarray:
"""Convert HWC/CHW/BHWC/BCHW image data in one vectorized transfer.
Video nodes previously moved and normalized every frame independently.
Converting a bounded batch at once reduces Python dispatch and host-device
transfer overhead while preserving the same clipping contract. The caller
chunks long videos to cap peak temporary memory. The returned array is
always contiguous BHWC RGB uint8.
"""
if not isinstance(images, torch.Tensor):
raise TypeError(f"Expected a torch.Tensor, got {type(images).__name__}.")
value = images.detach()
if value.ndim == 2:
value = value.unsqueeze(-1)
if value.ndim == 3:
value = value.unsqueeze(0)
if value.ndim != 4:
raise ValueError(
f"Expected an HWC/BHWC or CHW/BCHW image tensor, got {tuple(value.shape)}."
)
# ComfyUI uses BHWC. BCHW is accepted for compatibility with older nodes.
if value.shape[-1] not in (1, 3, 4) and value.shape[1] in (1, 3, 4):
value = value.permute(0, 2, 3, 1)
if value.shape[-1] not in (1, 3, 4):
raise ValueError(f"Unsupported image channel shape: {tuple(value.shape)}.")
value = torch.nan_to_num(
value.to(device="cpu", dtype=torch.float32),
nan=0.0,
posinf=1.0,
neginf=0.0,
)
if value.numel():
flat = value.reshape(value.shape[0], -1)
needs_byte_scale = (
(flat.amax(dim=1) > 1.0) | (flat.amin(dim=1) < 0.0)
).view(-1, 1, 1, 1)
value = torch.where(needs_byte_scale, value / 255.0, value)
value = value.clamp(0.0, 1.0).mul(255.0).round().to(torch.uint8)
if value.shape[-1] == 1:
value = value.expand(*value.shape[:-1], 3)
elif value.shape[-1] == 4:
value = value[..., :3]
return np.ascontiguousarray(value.numpy())
def tensor_to_pil(image: torch.Tensor, index: int = 0) -> Image.Image:
"""Convert a Comfy IMAGE tensor to an RGB PIL image without torchvision."""
if not isinstance(image, torch.Tensor):
raise TypeError(f"Expected a torch.Tensor, got {type(image).__name__}.")
value = image.detach()
value = image
if value.ndim == 4:
if not 0 <= index < value.shape[0]:
raise IndexError(f"Image batch index {index} is out of range.")
value = value[index]
if value.ndim == 2:
value = value.unsqueeze(-1)
if value.ndim != 3:
raise ValueError(
f"Expected an HWC/BHWC or CHW/BCHW image tensor, got {tuple(value.shape)}."
)
# ComfyUI uses HWC. CHW is accepted for compatibility with older callers.
if value.shape[-1] not in (1, 3, 4) and value.shape[0] in (1, 3, 4):
value = value.permute(1, 2, 0)
if value.shape[-1] not in (1, 3, 4):
raise ValueError(f"Unsupported image channel shape: {tuple(value.shape)}.")
value = torch.nan_to_num(
value.to(device="cpu", dtype=torch.float32), nan=0.0, posinf=1.0, neginf=0.0
)
if value.numel() and (value.max() > 1.0 or value.min() < 0.0):
value = value / 255.0
array = value.clamp(0.0, 1.0).mul(255.0).round().to(torch.uint8).numpy()
if array.shape[-1] == 1:
array = np.repeat(array, 3, axis=-1)
elif array.shape[-1] == 4:
array = array[..., :3]
elif index != 0:
raise IndexError("A single image only has batch index 0.")
array = _tensor_image_batch_to_uint8(value)[0]
return Image.fromarray(array, mode="RGB")
@@ -198,7 +228,24 @@ def tensor_batch_to_pil(images: torch.Tensor) -> list[Image.Image]:
return [tensor_to_pil(images)]
if images.ndim != 4:
raise ValueError(f"Expected an IMAGE batch, got {tuple(images.shape)}.")
return [tensor_to_pil(images, index) for index in range(images.shape[0])]
if images.shape[0] == 0:
return []
if images.device.type == "cpu":
# Per-frame conversion benchmarks faster for ordinary CPU-resident
# Comfy IMAGE batches and keeps the transient working set tiny.
return [tensor_to_pil(images, index) for index in range(images.shape[0])]
# Convert several frames per tensor operation without creating an
# unbounded full-video float32 temporary. On accelerator-resident batches,
# this amortizes device synchronization and transfers across many frames.
# The 128 MiB working-set ceiling keeps long HD/4K videos reliable.
frame_elements = max(1, int(images[0].numel()))
working_bytes = frame_elements * max(4, images.element_size())
chunk_frames = max(1, PIL_CONVERSION_CHUNK_BYTES // working_bytes)
output: list[Image.Image] = []
for start in range(0, int(images.shape[0]), chunk_frames):
frames = _tensor_image_batch_to_uint8(images[start : start + chunk_frames])
output.extend(Image.fromarray(frame, mode="RGB") for frame in frames)
return output
def pil_to_tensor(image: Image.Image) -> torch.Tensor:
@@ -540,18 +587,21 @@ class CachedModelNode:
def __init__(self) -> None:
self._model_handle = None
self._model_key = None
self._model_lock = threading.RLock()
def get_or_create_model(self, key: Any, factory: Callable[[], Any]):
if self._model_handle is None or self._model_key != key:
close_handle(self._model_handle)
self._model_handle = factory()
self._model_key = key
return self._model_handle
with self._model_lock:
if self._model_handle is None or self._model_key != key:
close_handle(self._model_handle)
self._model_handle = factory()
self._model_key = key
return self._model_handle
def clear_model(self) -> None:
close_handle(self._model_handle)
self._model_handle = None
self._model_key = None
with self._model_lock:
close_handle(self._model_handle)
self._model_handle = None
self._model_key = None
def maybe_clear_model(self, unload_after: bool) -> None:
if unload_after:
@@ -631,7 +681,10 @@ def llama_runtime_input_types() -> dict[str, tuple[Any, ...]]:
"min": 1,
"max": 8192,
"step": 1,
"tooltip": "Logical prompt batch. Lower this if context loading runs out of memory.",
"tooltip": (
"Logical prompt batch. Lower this if context loading runs "
"out of memory."
),
},
),
"n_ubatch": (
@@ -648,7 +701,11 @@ def llama_runtime_input_types() -> dict[str, tuple[Any, ...]]:
list(LLAMA_FLASH_ATTENTION_CHOICES),
{
"default": "Auto",
"tooltip": "Auto enables llama.cpp flash attention only with accelerator offload and safely retries without it when unsupported.",
"tooltip": (
"Auto enables llama.cpp flash attention only with "
"accelerator offload and safely retries without it when "
"unsupported."
),
},
),
"use_mmap": (
+1 -1
View File
@@ -258,7 +258,7 @@ class Sam2VideoPredictor:
):
raise ValueError("seed_mask must have shape [objects, height, width].")
object_ids = list(range(1, masks_for_seed.shape[0] + 1))
labels = {object_id: None for object_id in object_ids}
labels = dict.fromkeys(object_ids)
if not object_ids:
raise ValueError(
"Connect detections, a BOUNDING_BOX, or at least one seed mask."
+1074 -101
View File
File diff suppressed because it is too large Load Diff
+3 -222
View File
@@ -8,15 +8,14 @@ producing stricter output.
from __future__ import annotations
import json
import os
import re
from typing import Any, Literal, Optional
from typing import Any
import folder_paths
import torch
from pydantic import BaseModel, Field
from .prompts import system_msg_prompts, system_msg_simple
from .prompts import system_msg_prompts
from .runtime import (
LlamaHandle,
close_handle,
@@ -69,7 +68,7 @@ class ArtisticTechniques(BaseModel):
class ImageryTheme(BaseModel):
core_subject: str
additional_elements: Optional[list[str]] = None
additional_elements: list[str] | None = None
class VisualStyle(BaseModel):
@@ -157,222 +156,6 @@ def _structured_chat(
return raw, parsed
API_MODELS = [
"GPT-5.6 Terra",
"GPT-5.6 Sol",
"GPT-5.6 Luna",
"DeepSeek",
"Custom / OpenAI-compatible",
# Kept so saved workflows continue to deserialize without substitutions.
"ChatGPT-3.5",
"ChatGPT-4",
"gpt-3.5-turbo",
"gpt-3.5-turbo-0125",
"gpt-35-turbo",
"gpt-3.5-turbo-16k",
"gpt-3.5-turbo-16k-0613",
"gpt-4-0613",
"gpt-4-1106-preview",
"glm-4",
]
API_ROUTES = {
"GPT-5.6 Sol": ("gpt-5.6-sol", None, "Responses"),
"GPT-5.6 Terra": ("gpt-5.6-terra", None, "Responses"),
"GPT-5.6 Luna": ("gpt-5.6-luna", None, "Responses"),
"DeepSeek": ("deepseek-chat", "https://api.deepseek.com/v1", "Chat Completions"),
"ChatGPT-3.5": ("gpt-3.5-turbo", None, "Chat Completions"),
"ChatGPT-4": ("gpt-4", None, "Chat Completions"),
"gpt-35-turbo": ("gpt-35-turbo", None, "Chat Completions"),
"glm-4": ("glm-4", None, "Chat Completions"),
}
class PromptGenerateAPI:
def __init__(self):
self.session_history: list[dict[str, str]] = []
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model_name": (API_MODELS, {"default": "GPT-5.6 Terra"}),
"chat_type": (
"BOOLEAN",
{
"default": True,
"label_on": "Prompt Generator",
"label_off": "Simple Chat",
},
),
"api_key": (
"STRING",
{
"default": "",
"tooltip": (
"Leave blank to use OPENAI_API_KEY, DEEPSEEK_API_KEY, "
"or VLM_API_KEY."
),
},
),
"description": (
"STRING",
{"multiline": True, "default": ""},
),
"question": (
"STRING",
{"multiline": True, "default": ""},
),
"context_size": (
"INT",
{"default": 5, "min": 0, "max": 30, "step": 1},
),
"seed": (
"INT",
{
"default": 0,
"min": 0,
"max": 0xFFFFFFFFFFFFFFFF,
"step": 1,
},
),
},
"optional": {
"base_url": (
"STRING",
{
"default": "",
"tooltip": (
"OpenAI-compatible base URL, e.g. http://127.0.0.1:8000/v1."
),
},
),
"model_override": (
"STRING",
{
"default": "",
"tooltip": "Exact provider model ID. Overrides the picker.",
},
),
"api_mode": (
["Auto", "Responses", "Chat Completions"],
{"default": "Auto"},
),
"timeout_seconds": (
"FLOAT",
{"default": 120.0, "min": 1.0, "max": 1800.0},
),
"reasoning_effort": (
["none", "low", "medium", "high", "xhigh", "max"],
{"default": "none"},
),
},
}
RETURN_TYPES = ("STRING",)
FUNCTION = "generate_prompt"
CATEGORY = "VLM Nodes/LLM"
def _route(
self, model_name, model_override, base_url, api_mode
) -> tuple[str, str | None, str]:
route = API_ROUTES.get(model_name)
if route is None:
if model_name in API_MODELS and model_name not in {
"Custom / OpenAI-compatible"
}:
route = (model_name, None, "Chat Completions")
else:
route = ("", None, "Chat Completions")
model, route_url, route_mode = route
model = (model_override or model).strip()
if not model:
raise ValueError("A model ID is required for Custom / OpenAI-compatible.")
effective_url = (base_url or route_url or "").strip() or None
mode = route_mode if api_mode == "Auto" else api_mode
return model, effective_url, mode
def generate_prompt(
self,
model_name,
chat_type,
api_key,
description,
question,
context_size,
seed,
base_url="",
model_override="",
api_mode="Auto",
timeout_seconds=120.0,
reasoning_effort="none",
):
openai = require_module("openai", "openai")
model, effective_url, mode = self._route(
model_name, model_override, base_url, api_mode
)
key = (
api_key.strip()
or (os.getenv("DEEPSEEK_API_KEY", "") if model_name == "DeepSeek" else "")
or os.getenv("VLM_API_KEY", "")
or os.getenv("OPENAI_API_KEY", "")
)
if not key:
raise ValueError(
"No API key was supplied. Set OPENAI_API_KEY, "
"DEEPSEEK_API_KEY, or VLM_API_KEY, or enter the key in the node."
)
client_kwargs: dict[str, Any] = {
"api_key": key,
"timeout": float(timeout_seconds),
"max_retries": 2,
}
if effective_url:
client_kwargs["base_url"] = effective_url
client = openai.OpenAI(**client_kwargs)
system = system_msg_prompts if chat_type else system_msg_simple
user_message = (
f"Description:\n{description.strip()}\n\n"
f"Optional question:\n{question.strip()}"
).strip()
history_limit = max(0, int(context_size)) * 2
history = self.session_history[-history_limit:] if history_limit else []
if mode == "Responses":
response = client.responses.create(
model=model,
instructions=system,
input=history + [{"role": "user", "content": user_message}],
reasoning={"effort": reasoning_effort},
)
result = response.output_text
else:
messages = (
[{"role": "system", "content": system}]
+ history
+ [{"role": "user", "content": user_message}]
)
request: dict[str, Any] = {
"model": model,
"messages": messages,
"seed": int(seed),
}
if model.startswith("gpt-5.6"):
request["reasoning_effort"] = reasoning_effort
completion = client.chat.completions.create(**request)
result = completion.choices[0].message.content or ""
self.session_history.extend(
[
{"role": "user", "content": user_message},
{"role": "assistant", "content": result},
]
)
return (result,)
class LLMLoader:
@classmethod
def INPUT_TYPES(cls):
@@ -1206,7 +989,6 @@ NODE_CLASS_MAPPINGS = {
"KeywordExtraction": KeywordExtraction,
"LLavaPromptGenerator": LLavaPromptGenerator,
"Suggester": Suggester,
"PromptGenerateAPI": PromptGenerateAPI,
"CreativeArtPromptGenerator": CreativeArtPromptGenerator,
"ChatMusician": ChatMusician,
"StructuredOutput": StructuredOutput,
@@ -1221,7 +1003,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"KeywordExtraction": "Structured Keyword Extraction",
"LLavaPromptGenerator": "Structured Prompt Generator",
"Suggester": "Prompt Suggester",
"PromptGenerateAPI": "OpenAI-Compatible Prompt API",
"CreativeArtPromptGenerator": "Creative Art Prompt Generator",
"ChatMusician": "Chat Musician",
"StructuredOutput": "Structured Output",
+1 -1
View File
@@ -8,8 +8,8 @@ keeps the baseline portable across CUDA, ROCm, MPS, XPU, and CPU systems.
from __future__ import annotations
import math
from collections.abc import Iterable
from dataclasses import dataclass, field
from typing import Iterable
import numpy as np
from scipy.optimize import linear_sum_assignment
+1 -1
View File
@@ -105,7 +105,7 @@ class UformGen2QwenNode(CachedModelNode):
RETURN_TYPES = ("STRING",)
FUNCTION = "uform_gen2_qwen_chat"
CATEGORY = "VLM Nodes/UformGen2Qwen"
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
def uform_gen2_qwen_chat(
self, image, question, max_new_tokens=512, unload_after=False
File diff suppressed because it is too large Load Diff
+409
View File
@@ -19,12 +19,16 @@ VLM_DETECTIONS = "VLM_DETECTIONS"
VLM_TRACKS = "VLM_TRACKS"
VLM_POINTS = "VLM_POINTS"
VLM_EVENTS = "VLM_EVENTS"
VLM_VIDEO_SELECTION = "VLM_VIDEO_SELECTION"
VLM_SCENE_STATE = "VLM_SCENE_STATE"
SCHEMA_VERSION = 1
DETECTIONS_SCHEMA = "comfyui-vlm/detections"
TRACKS_SCHEMA = "comfyui-vlm/tracks"
POINTS_SCHEMA = "comfyui-vlm/points"
EVENTS_SCHEMA = "comfyui-vlm/events"
VIDEO_SELECTION_SCHEMA = "comfyui-vlm/video-selection"
SCENE_STATE_SCHEMA = "comfyui-vlm/scene-state"
PointXY = tuple[float, float]
BoxXYXY = tuple[float, float, float, float]
@@ -972,6 +976,403 @@ class EventSequence:
raise ValueError(f"Invalid event JSON: {exc.msg}.") from exc
@dataclass(frozen=True, slots=True)
class SelectedVideoFrame:
"""One source-frame reference preserved through adaptive sampling."""
source_frame_index: int
timestamp: float
score: float = 0.0
reasons: tuple[str, ...] = ()
def __post_init__(self) -> None:
if (
not isinstance(self.source_frame_index, int)
or self.source_frame_index < 0
):
raise ValueError("source_frame_index must be a non-negative integer.")
object.__setattr__(
self,
"timestamp",
_non_negative(self.timestamp, "timestamp"),
)
score = _finite(self.score, "selection score")
if not 0.0 <= score <= 1.0:
raise ValueError("selection score must be between 0 and 1.")
object.__setattr__(self, "score", score)
reasons = tuple(self.reasons)
if any(not isinstance(reason, str) or not reason.strip() for reason in reasons):
raise TypeError("selection reasons must be non-empty strings.")
object.__setattr__(self, "reasons", reasons)
def to_dict(self) -> dict[str, Any]:
result: dict[str, Any] = {
"source_frame_index": self.source_frame_index,
"timestamp": self.timestamp,
"score": self.score,
}
if self.reasons:
result["reasons"] = list(self.reasons)
return result
@classmethod
def from_dict(cls, value: Mapping[str, Any]) -> SelectedVideoFrame:
if not isinstance(value, Mapping):
raise TypeError("A selected video frame must be a JSON object.")
return cls(
source_frame_index=value["source_frame_index"],
timestamp=value["timestamp"],
score=value.get("score", 0.0),
reasons=tuple(value.get("reasons", ())),
)
@dataclass(frozen=True, slots=True)
class VideoFrameSelection:
"""Immutable map from a sampled IMAGE batch back to its source video."""
width: int
height: int
source_frame_count: int
fps: float
frames: tuple[SelectedVideoFrame, ...]
strategy: str = "adaptive"
source: str | None = None
metadata: FrozenDict = field(default_factory=FrozenDict)
version: int = SCHEMA_VERSION
def __post_init__(self) -> None:
if self.version != SCHEMA_VERSION:
raise ValueError(
f"Unsupported video selection schema version {self.version}."
)
if not isinstance(self.width, int) or self.width <= 0:
raise ValueError("width must be a positive integer.")
if not isinstance(self.height, int) or self.height <= 0:
raise ValueError("height must be a positive integer.")
if (
not isinstance(self.source_frame_count, int)
or self.source_frame_count <= 0
):
raise ValueError("source_frame_count must be a positive integer.")
fps = _finite(self.fps, "fps")
if fps <= 0:
raise ValueError("fps must be positive.")
frames = tuple(self.frames)
if not frames:
raise ValueError("A video selection requires at least one frame.")
if any(not isinstance(frame, SelectedVideoFrame) for frame in frames):
raise TypeError("frames must contain SelectedVideoFrame values.")
indices = [frame.source_frame_index for frame in frames]
if indices != sorted(set(indices)):
raise ValueError(
"Selected source frame indices must be unique and increasing."
)
if indices[-1] >= self.source_frame_count:
raise ValueError("A selected frame lies outside the source video.")
expected_timestamps = [index / fps for index in indices]
if any(
abs(frame.timestamp - expected) > max(1.0e-6, 0.51 / fps)
for frame, expected in zip(frames, expected_timestamps)
):
raise ValueError(
"Selected frame timestamps do not match source indices and fps."
)
strategy = str(self.strategy).strip()
if not strategy:
raise ValueError("strategy must not be empty.")
object.__setattr__(self, "fps", fps)
object.__setattr__(self, "frames", frames)
object.__setattr__(self, "strategy", strategy)
object.__setattr__(self, "source", _optional_text(self.source, "source"))
object.__setattr__(self, "metadata", _metadata(self.metadata))
@property
def duration(self) -> float:
return self.source_frame_count / self.fps
@property
def indices(self) -> tuple[int, ...]:
return tuple(frame.source_frame_index for frame in self.frames)
@property
def timestamps(self) -> tuple[float, ...]:
return tuple(frame.timestamp for frame in self.frames)
def to_dict(self) -> dict[str, Any]:
result: dict[str, Any] = {
"schema": VIDEO_SELECTION_SCHEMA,
"version": self.version,
"media": {
"width": self.width,
"height": self.height,
"source_frame_count": self.source_frame_count,
"fps": self.fps,
"duration": self.duration,
},
"strategy": self.strategy,
"frames": [frame.to_dict() for frame in self.frames],
}
if self.source is not None:
result["source"] = self.source
if self.metadata:
result["metadata"] = self.metadata.to_dict()
return result
def to_json(self, *, indent: int | None = None) -> str:
return json.dumps(
self.to_dict(),
ensure_ascii=False,
allow_nan=False,
sort_keys=True,
indent=indent,
)
@classmethod
def from_dict(cls, value: Mapping[str, Any]) -> VideoFrameSelection:
if not isinstance(value, Mapping):
raise TypeError("Video selection JSON must contain an object.")
if value.get("schema") != VIDEO_SELECTION_SCHEMA:
raise ValueError(f"Expected schema {VIDEO_SELECTION_SCHEMA!r}.")
if value.get("version") != SCHEMA_VERSION:
raise ValueError(
f"Unsupported video selection schema version "
f"{value.get('version')!r}."
)
media = value.get("media")
if not isinstance(media, Mapping):
raise ValueError("Video selection JSON requires a media object.")
return cls(
width=media["width"],
height=media["height"],
source_frame_count=media["source_frame_count"],
fps=media["fps"],
frames=tuple(
SelectedVideoFrame.from_dict(frame)
for frame in value.get("frames", ())
),
strategy=value.get("strategy", "adaptive"),
source=value.get("source"),
metadata=value.get("metadata"),
version=value["version"],
)
@classmethod
def from_json(cls, value: str) -> VideoFrameSelection:
try:
return cls.from_dict(json.loads(value))
except json.JSONDecodeError as exc:
raise ValueError(f"Invalid video selection JSON: {exc.msg}.") from exc
@dataclass(frozen=True, slots=True)
class SceneObjectState:
"""Compact latest state derived from a temporally consistent object track."""
track_id: int
first_seen: float
last_seen: float
last_bbox_xyxy: BoxXYXY
observation_count: int
label: str | None = None
state: str = "active"
mean_confidence: float | None = None
velocity_xy_px_s: PointXY = (0.0, 0.0)
metadata: FrozenDict = field(default_factory=FrozenDict)
def __post_init__(self) -> None:
if not isinstance(self.track_id, int) or self.track_id < 0:
raise ValueError("track_id must be a non-negative integer.")
first_seen = _non_negative(self.first_seen, "first_seen")
last_seen = _non_negative(self.last_seen, "last_seen")
if last_seen < first_seen:
raise ValueError("last_seen must be at or after first_seen.")
if (
not isinstance(self.observation_count, int)
or self.observation_count <= 0
):
raise ValueError("observation_count must be a positive integer.")
state = str(self.state).strip()
if not state:
raise ValueError("state must not be empty.")
velocity = tuple(_finite(value, "velocity") for value in self.velocity_xy_px_s)
if len(velocity) != 2:
raise ValueError("velocity_xy_px_s must contain exactly two values.")
object.__setattr__(self, "first_seen", first_seen)
object.__setattr__(self, "last_seen", last_seen)
object.__setattr__(self, "last_bbox_xyxy", _box(self.last_bbox_xyxy))
object.__setattr__(self, "label", _optional_text(self.label, "label"))
object.__setattr__(self, "state", state)
object.__setattr__(
self,
"mean_confidence",
_optional_score(self.mean_confidence),
)
object.__setattr__(self, "velocity_xy_px_s", velocity)
object.__setattr__(self, "metadata", _metadata(self.metadata))
def to_dict(self) -> dict[str, Any]:
result: dict[str, Any] = {
"track_id": self.track_id,
"first_seen": self.first_seen,
"last_seen": self.last_seen,
"last_bbox_xyxy": list(self.last_bbox_xyxy),
"observation_count": self.observation_count,
"state": self.state,
"velocity_xy_px_s": list(self.velocity_xy_px_s),
}
if self.label is not None:
result["label"] = self.label
if self.mean_confidence is not None:
result["mean_confidence"] = self.mean_confidence
if self.metadata:
result["metadata"] = self.metadata.to_dict()
return result
@classmethod
def from_dict(cls, value: Mapping[str, Any]) -> SceneObjectState:
if not isinstance(value, Mapping):
raise TypeError("A scene object must be a JSON object.")
return cls(
track_id=value["track_id"],
first_seen=value["first_seen"],
last_seen=value["last_seen"],
last_bbox_xyxy=value["last_bbox_xyxy"],
observation_count=value["observation_count"],
label=value.get("label"),
state=value.get("state", "active"),
mean_confidence=value.get("mean_confidence"),
velocity_xy_px_s=tuple(value.get("velocity_xy_px_s", (0.0, 0.0))),
metadata=value.get("metadata"),
)
@dataclass(frozen=True, slots=True)
class SceneState:
"""Persistent, serializable world-state summary for video reasoning."""
width: int
height: int
frame_count: int
fps: float | None
objects: tuple[SceneObjectState, ...] = ()
events: tuple[TemporalEvent, ...] = ()
source: str | None = None
metadata: FrozenDict = field(default_factory=FrozenDict)
version: int = SCHEMA_VERSION
def __post_init__(self) -> None:
if self.version != SCHEMA_VERSION:
raise ValueError(f"Unsupported scene state schema version {self.version}.")
if not isinstance(self.width, int) or self.width <= 0:
raise ValueError("width must be a positive integer.")
if not isinstance(self.height, int) or self.height <= 0:
raise ValueError("height must be a positive integer.")
if not isinstance(self.frame_count, int) or self.frame_count < 0:
raise ValueError("frame_count must be a non-negative integer.")
fps = None if self.fps is None else _finite(self.fps, "fps")
if fps is not None and fps <= 0:
raise ValueError("fps must be positive.")
objects = tuple(self.objects)
if any(not isinstance(item, SceneObjectState) for item in objects):
raise TypeError("objects must contain SceneObjectState values.")
ids = [item.track_id for item in objects]
if ids != sorted(set(ids)):
raise ValueError("Scene objects must have unique increasing track IDs.")
events = tuple(self.events)
if any(not isinstance(item, TemporalEvent) for item in events):
raise TypeError("events must contain TemporalEvent values.")
if list(events) != sorted(
events,
key=lambda event: (event.start_time, event.end_time),
):
raise ValueError("Scene events must be ordered by start_time.")
object.__setattr__(self, "fps", fps)
object.__setattr__(self, "objects", objects)
object.__setattr__(self, "events", events)
object.__setattr__(self, "source", _optional_text(self.source, "source"))
object.__setattr__(self, "metadata", _metadata(self.metadata))
@property
def duration(self) -> float | None:
return (
self.frame_count / self.fps
if self.fps is not None and self.frame_count
else None
)
def to_dict(self) -> dict[str, Any]:
media: dict[str, Any] = {
"width": self.width,
"height": self.height,
"frame_count": self.frame_count,
}
if self.fps is not None:
media["fps"] = self.fps
if self.duration is not None:
media["duration"] = self.duration
result: dict[str, Any] = {
"schema": SCENE_STATE_SCHEMA,
"version": self.version,
"media": media,
"objects": [item.to_dict() for item in self.objects],
"events": [item.to_dict() for item in self.events],
}
if self.source is not None:
result["source"] = self.source
if self.metadata:
result["metadata"] = self.metadata.to_dict()
return result
def to_json(self, *, indent: int | None = None) -> str:
return json.dumps(
self.to_dict(),
ensure_ascii=False,
allow_nan=False,
sort_keys=True,
indent=indent,
)
@classmethod
def from_dict(cls, value: Mapping[str, Any]) -> SceneState:
if not isinstance(value, Mapping):
raise TypeError("Scene state JSON must contain an object.")
if value.get("schema") != SCENE_STATE_SCHEMA:
raise ValueError(f"Expected schema {SCENE_STATE_SCHEMA!r}.")
if value.get("version") != SCHEMA_VERSION:
raise ValueError(
f"Unsupported scene state schema version "
f"{value.get('version')!r}."
)
media = value.get("media")
if not isinstance(media, Mapping):
raise ValueError("Scene state JSON requires a media object.")
return cls(
width=media["width"],
height=media["height"],
frame_count=media["frame_count"],
fps=media.get("fps"),
objects=tuple(
SceneObjectState.from_dict(item)
for item in value.get("objects", ())
),
events=tuple(
TemporalEvent.from_dict(item)
for item in value.get("events", ())
),
source=value.get("source"),
metadata=value.get("metadata"),
version=value["version"],
)
@classmethod
def from_json(cls, value: str) -> SceneState:
try:
return cls.from_dict(json.loads(value))
except json.JSONDecodeError as exc:
raise ValueError(f"Invalid scene state JSON: {exc.msg}.") from exc
__all__ = [
"BoxXYXY",
"DETECTIONS_SCHEMA",
@@ -985,7 +1386,11 @@ __all__ = [
"PointSequence",
"PointXY",
"Polygon",
"SCENE_STATE_SCHEMA",
"SCHEMA_VERSION",
"SceneObjectState",
"SceneState",
"SelectedVideoFrame",
"TRACKS_SCHEMA",
"TemporalEvent",
"Track",
@@ -993,6 +1398,10 @@ __all__ = [
"VLM_DETECTIONS",
"VLM_EVENTS",
"VLM_POINTS",
"VLM_SCENE_STATE",
"VLM_TRACKS",
"VLM_VIDEO_SELECTION",
"VIDEO_SELECTION_SCHEMA",
"VideoFrameSelection",
"VisionPoint",
]
+2 -1
View File
@@ -4,8 +4,9 @@ from __future__ import annotations
import json
import math
from collections.abc import Iterable
from dataclasses import replace
from typing import Any, Iterable
from typing import Any
import numpy as np
import torch
+55 -3
View File
@@ -1,10 +1,10 @@
[project]
name = "comfyui_vlm_nodes"
version = "3.0.0"
version = "3.5.0"
description = "Production-ready local and API vision-language nodes for ComfyUI"
readme = "README.md"
requires-python = ">=3.10"
license = "MIT"
license = "Apache-2.0"
license-files = ["LICENSE"]
dependencies = [
"accelerate>=1.1,<2",
@@ -12,13 +12,17 @@ dependencies = [
"diffusers>=0.34,<1",
"einops>=0.8,<1",
"huggingface-hub>=1.5,<2",
"openai>=1.30,<3",
"httpx>=0.27,<1",
"jsonschema>=4.22,<5",
"num2words>=0.5.14,<1",
"openai>=2,<3",
"pydantic>=2.7,<3",
"qwen-vl-utils>=0.0.14",
"safetensors>=0.4.3",
"scipy>=1.10,<2",
"soundfile>=0.12",
"symusic>=0.5",
"svgelements>=1.9.6,<2",
"transformers>=5.4,<6",
]
classifiers = [
@@ -43,11 +47,50 @@ quantization = [
gguf = [
"llama-cpp-python>=0.3.20,<1",
]
robotics-client = [
"msgpack>=1.0.8,<2",
"pyzmq>=26,<28",
"websockets>=14,<17",
]
[project.urls]
Repository = "https://github.com/gokayfem/ComfyUI_VLM_nodes"
Issues = "https://github.com/gokayfem/ComfyUI_VLM_nodes/issues"
[tool.pytest.ini_options]
testpaths = ["tests"]
# manual_*.py download multi-gigabyte checkpoints and are run by hand.
python_files = ["test_*.py"]
addopts = "-ra --strict-markers --strict-config"
filterwarnings = ["default"]
[tool.ruff]
line-length = 100
# nodes/joytagger is vendored upstream code kept byte-compatible with its
# source, including tab indentation. Reformatting it would break that.
extend-exclude = ["nodes/joytagger"]
[tool.ruff.lint]
select = ["E", "F", "W", "I", "UP", "C4", "B", "SIM"]
ignore = [
# zip(strict=) changes behaviour when lengths differ; enabling it needs a
# per-call-site audit rather than a blanket flag.
"B905",
# The streaming closures in modern_vlm are started and joined inside the
# same loop iteration, so the loop variable cannot change under them.
"B023",
# try/except/pass around optional backends stays readable as-is;
# contextlib.suppress would hide which dependency is being probed.
"SIM105",
]
[tool.ruff.lint.per-file-ignores]
# Prompt templates are data. Rewrapping them changes the model input.
"nodes/prompts.py" = ["E501"]
# The manual smoke scripts must bootstrap the package onto sys.path before
# they can import from it.
"tests/manual_*.py" = ["E402"]
[tool.comfy]
PublisherId = "gokayfem"
DisplayName = "ComfyUI VLM Nodes"
@@ -56,6 +99,9 @@ Icon = ""
[tool.setuptools]
packages = [
"comfyui_vlm_nodes",
"comfyui_vlm_nodes.examples",
"comfyui_vlm_nodes.examples.robotics",
"comfyui_vlm_nodes.examples.vision",
"comfyui_vlm_nodes.nodes",
"comfyui_vlm_nodes.nodes.joytagger",
"comfyui_vlm_nodes.web",
@@ -69,6 +115,12 @@ comfyui_vlm_nodes = "."
[tool.setuptools.package-data]
comfyui_vlm_nodes = [
"*.json",
"SECURITY.md",
"examples/*.json",
"examples/robotics/*.py",
"examples/robotics/*.md",
"examples/robotics/*.json",
"examples/vision/*.json",
"requirements*.txt",
]
"comfyui_vlm_nodes.web.js" = ["*.js"]
+8
View File
@@ -0,0 +1,8 @@
# Development and CI tooling. Not needed to run the nodes in ComfyUI.
# Install with ComfyUI's Python alongside requirements.txt:
# python -m pip install -r requirements.txt -r requirements-dev.txt
build>=1.2,<2
packaging>=24
pytest>=8,<9
pytest-cov>=5,<8
ruff>=0.14,<1
+11
View File
@@ -0,0 +1,11 @@
# Install this file only into the isolated Moondream sidecar environment.
# Do not install it into ComfyUI's main environment: moondream 1.3 pins
# Pillow <11 while current ComfyUI uses a newer Pillow release.
moondream==1.3.0
# moondream 1.3.0 expects this exact runtime API. 0.4.7+ renamed the
# prefix-mask kernel and is not source-compatible with kestrel 0.4.2.
kestrel-kernels==0.4.6
# Kestrel's CUDA 12 AOT kernels call cudaLibraryLoadData. PyTorch's cu126
# runtime (12.6.77) does not export it; 12.9.79 does and remains within the
# CUDA 12 ABI. Keep this inside the isolated Photon environment only.
nvidia-cuda-runtime-cu12==12.9.79; (sys_platform == "linux" and platform_machine == "x86_64") or (sys_platform == "win32" and platform_machine == "AMD64")
+5
View File
@@ -0,0 +1,5 @@
# Lightweight native clients only. Heavy VLA policy runtimes stay in a
# separate LeRobot, openpi, GR00T, OpenVLA/OFT, or Octo environment.
msgpack>=1.0.8,<2
pyzmq>=26,<28
websockets>=14,<17
+5 -1
View File
@@ -7,11 +7,15 @@ bitsandbytes>=0.50,<1; (sys_platform == "linux" and platform_machine == "x86_64"
diffusers>=0.34,<1
einops>=0.8,<1
huggingface-hub>=1.5,<2
openai>=1.30,<3
httpx>=0.27,<1
jsonschema>=4.22,<5
num2words>=0.5.14,<1
openai>=2,<3
pydantic>=2.7,<3
qwen-vl-utils>=0.0.14
safetensors>=0.4.3
scipy>=1.10,<2
soundfile>=0.12
symusic>=0.5
svgelements>=1.9.6,<2
transformers>=5.4,<6
+50
View File
@@ -0,0 +1,50 @@
"""Make this checkout importable from the manual smoke scripts.
The manual scripts run as `python tests/manual_*.py`, outside pytest, so they
do not get `conftest.py`. Without this they only import when the checkout
directory happens to be named `ComfyUI_VLM_nodes`, which is true in a normal
ComfyUI install but not in a git worktree named after a feature branch.
Usage, before importing anything from the package:
from _bootstrap import bootstrap
bootstrap()
"""
from __future__ import annotations
import importlib.util
import sys
from pathlib import Path
PACKAGE = "ComfyUI_VLM_nodes"
REPOSITORY = Path(__file__).resolve().parents[1]
def bootstrap() -> None:
"""Put the repository and ComfyUI on sys.path, then load this checkout."""
for candidate in (
REPOSITORY.parent,
REPOSITORY.parent / "ComfyUI",
REPOSITORY.parents[1],
):
if candidate.exists():
sys.path.insert(0, str(candidate))
if PACKAGE in sys.modules or REPOSITORY.name == PACKAGE:
return
# Load this checkout explicitly so the script can never pass by silently
# importing a sibling clone with the canonical directory name.
specification = importlib.util.spec_from_file_location(
PACKAGE,
REPOSITORY / "__init__.py",
submodule_search_locations=[str(REPOSITORY)],
)
if specification is None or specification.loader is None:
raise RuntimeError(f"Could not load package from {REPOSITORY}.")
package = importlib.util.module_from_spec(specification)
sys.modules[PACKAGE] = package
specification.loader.exec_module(package)
+5 -2
View File
@@ -11,9 +11,12 @@ from __future__ import annotations
import json
from transformers import AutoConfig, AutoProcessor
from _bootstrap import bootstrap
from ComfyUI_VLM_nodes.nodes.modern_vlm import MODEL_CATALOG
bootstrap()
from ComfyUI_VLM_nodes.nodes.modern_vlm import MODEL_CATALOG # noqa: E402
from transformers import AutoConfig, AutoProcessor # noqa: E402
def main() -> int:
+5 -1
View File
@@ -11,7 +11,11 @@ import json
import time
from pathlib import Path
from ComfyUI_VLM_nodes.nodes.runtime import (
from _bootstrap import bootstrap
bootstrap()
from ComfyUI_VLM_nodes.nodes.runtime import ( # noqa: E402
LlamaHandle,
default_llama_threads,
hf_download,
+178
View File
@@ -0,0 +1,178 @@
"""Opt-in real-weight smoke test for the GGUF *node classes*.
`manual_llama_cpp_smoke.py` proves the shared `LlamaHandle` runtime loads and
generates. This script goes one level up and drives the actual ComfyUI node
classes end to end against real weights, which covers the parts the offline
suite deliberately stubs:
* `LLMLoader` resolving a real file through ComfyUI's `folder_paths`
* `LLMSampler` producing real text from real sampling arguments
* `StructuredOutput` constraining a real model to a generated JSON Schema —
the llama.cpp grammar path, which cannot be verified with a stub
* `LLMOptionalMemoryFreeSimple` releasing a real llama.cpp allocation
Never run in CI: it downloads weights and needs `llama-cpp-python`.
Example:
python tests/manual_llm_node_smoke.py --download
python tests/manual_llm_node_smoke.py --model /models/qwen.gguf
"""
from __future__ import annotations
import argparse
import json
import shutil
import time
from pathlib import Path
from _bootstrap import bootstrap
bootstrap()
import folder_paths # noqa: E402
from ComfyUI_VLM_nodes.nodes.runtime import hf_download, model_root # noqa: E402
from ComfyUI_VLM_nodes.nodes.suggest import ( # noqa: E402
LLMLoader,
LLMOptionalMemoryFreeSimple,
LLMSampler,
StructuredOutput,
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--model", type=Path)
parser.add_argument("--download", action="store_true")
parser.add_argument("--repo", default="ggml-org/Qwen3.5-0.8B-GGUF")
parser.add_argument("--filename", default="Qwen3.5-0.8B-Q4_0.gguf")
parser.add_argument("--n-gpu-layers", type=int, default=-1)
return parser.parse_args()
def stage_model(args: argparse.Namespace) -> str:
"""Put the GGUF where ComfyUI's folder_paths can enumerate it."""
if args.model is None:
if not args.download:
raise SystemExit(
"Pass --model /path/to/model.gguf, or allow the small default "
"download with --download."
)
source = hf_download(args.repo, args.filename, "llm-node-smoke")
else:
source = args.model.resolve()
if not source.is_file():
raise SystemExit(f"{source} is not a file.")
destination = model_root() / source.name
if not destination.exists():
destination.parent.mkdir(parents=True, exist_ok=True)
shutil.copy2(source, destination)
# The loader nodes offer whatever folder_paths enumerates, so the staged
# file has to actually show up there. Staging happens before the first
# get_filename_list call in this process, so there is no cache to clear.
listed = folder_paths.get_filename_list("LLavacheckpoints")
if source.name not in listed:
raise SystemExit(
f"{source.name} is not enumerated in LLavacheckpoints: {listed}"
)
return source.name
def main() -> None:
args = parse_args()
checkpoint = stage_model(args)
results: dict[str, object] = {"checkpoint": checkpoint}
# 1. The loader must hand back a lazy handle that has not loaded yet.
started = time.perf_counter()
(model,) = LLMLoader().load_llm_checkpoint(
ckpt_name=checkpoint,
max_ctx=2048,
gpu_layers=args.n_gpu_layers,
n_threads=4,
)
results["loader_returned_without_loading"] = model._llm is None
results["loader_seconds"] = round(time.perf_counter() - started, 3)
# 2. Real generation through the real sampler node.
#
# Deliberately no assertion on what the model *says*: at 0.8B/Q4 the answer
# is often factually wrong, and that is model quality, not node
# correctness. What the node owns is that generation happens and that its
# sampling arguments actually reach llama.cpp — so assert determinism for a
# fixed seed at temperature 0 instead.
def sample(seed: int) -> tuple[str, float]:
started = time.perf_counter()
(text,) = LLMSampler().generate_text_advanced(
system_msg="You answer with a single short sentence.",
prompt="Name the largest planet in the solar system.",
model=model,
max_tokens=48,
temperature=0.0,
top_p=0.95,
top_k=40,
frequency_penalty=0.0,
presence_penalty=0.0,
repeat_penalty=1.1,
seed=seed,
)
return text, round(time.perf_counter() - started, 3)
text, elapsed = sample(42)
repeat, _ = sample(42)
results["sampler_seconds"] = elapsed
results["sampler_text"] = text
results["sampler_produced_text"] = bool(text.strip())
results["sampler_deterministic_for_fixed_seed"] = text == repeat
# 3. The grammar-constrained path. A stub cannot prove this works.
started = time.perf_counter()
(value,) = StructuredOutput().keyword_extract(
prompt="The photograph shows a calm, empty beach at sunrise.",
model=model,
temperature=0.0,
attribute_name="mood",
attribute_type="Category",
attribute_description="The overall mood of the described scene.",
categories="calm, tense, joyful, melancholy",
)
results["structured_seconds"] = round(time.perf_counter() - started, 3)
results["structured_value"] = value
# The whole point of the schema is that the model cannot answer off-menu.
results["structured_respected_enum"] = value in {
"calm",
"tense",
"joyful",
"melancholy",
}
model.close()
# 4. A managed-cache node must really release its allocation.
node = LLMOptionalMemoryFreeSimple()
(cached_text,) = node.generate_text(
ckpt_name=checkpoint,
max_ctx=2048,
gpu_layers=args.n_gpu_layers,
n_threads=4,
prompt="Say the word: ready",
temperature=0.0,
unload=True,
)
results["managed_cache_text"] = cached_text
results["managed_cache_released"] = node._handle is None and node._key is None
checks = {
key: value for key, value in results.items() if isinstance(value, bool)
}
results["ALL_CHECKS_PASSED"] = all(checks.values())
print(json.dumps(results, ensure_ascii=False, indent=2))
if not results["ALL_CHECKS_PASSED"]:
failed = [key for key, value in checks.items() if not value]
raise SystemExit(f"Failed checks: {failed}")
if __name__ == "__main__":
main()
+7 -1
View File
@@ -14,8 +14,14 @@ import json
import time
import torch
from _bootstrap import bootstrap
from ComfyUI_VLM_nodes.nodes.modern_vlm import MODEL_CATALOG, ModernVLMPredictor
bootstrap()
from ComfyUI_VLM_nodes.nodes.modern_vlm import ( # noqa: E402
MODEL_CATALOG,
ModernVLMPredictor,
)
def test_image() -> torch.Tensor:
+158
View File
@@ -0,0 +1,158 @@
"""Queue a real robotics policy workflow through ComfyUI's local API.
This is intentionally excluded from pytest: it requires a running ComfyUI
server, a running policy sidecar, a real image in ComfyUI's input directory,
and downloaded policy weights.
"""
from __future__ import annotations
import argparse
import json
import time
import urllib.request
import uuid
def _graph(image: str, policy_endpoint: str) -> dict:
camera_names = [
"observation.images.camera1",
"observation.images.camera2",
"observation.images.camera3",
]
return {
"1": {"class_type": "LoadImage", "inputs": {"image": image}},
"12": {
"class_type": "ImageScale",
"inputs": {
"image": ["1", 0],
"upscale_method": "lanczos",
"width": 256,
"height": 256,
"crop": "center",
},
},
"2": {
"class_type": "VLAEmbodimentProfile",
"inputs": {
"preset": "LeRobot SO-100 / SO-101 template",
"control_hz": 30.0,
"state_names_json": "",
"action_names_json": "",
"action_min_json": "",
"action_max_json": "",
"max_delta_json": "",
"camera_names_json": json.dumps(camera_names),
"action_mode_override": "",
},
},
"3": {
"class_type": "VLAObservationBuilder",
"inputs": {
"task": (
"Move the end effector toward the backpack and prepare to grasp it."
),
"state_json": "[0, 0, 0, 0, 0, 0]",
"primary_image": ["12", 0],
"primary_camera": camera_names[0],
"history_fps": 10.0,
"timestamp": 0.0,
"embodiment": ["2", 0],
"wrist_image": ["12", 0],
"wrist_camera": camera_names[1],
"secondary_image": ["12", 0],
"secondary_camera": camera_names[2],
},
},
"4": {
"class_type": "VLAHTTPPolicy",
"inputs": {
"observation": ["3", 0],
"endpoint": policy_endpoint,
"timeout_seconds": 120.0,
"include_history": True,
"allow_remote": False,
},
},
"5": {
"class_type": "VLAActionSafety",
"inputs": {
"actions": ["4", 0],
"embodiment": ["2", 0],
"mode": "Clamp safely",
"execution_horizon": 4,
"previous_action_json": "[0, 0, 0, 0, 0, 0]",
},
},
"6": {
"class_type": "VLATrajectoryPreview",
"inputs": {
"actions": ["5", 0],
"width": 960,
"height": 480,
"embodiment": ["2", 0],
},
},
"7": {
"class_type": "VLAActionInspect",
"inputs": {"actions": ["5", 0], "step_index": 0},
},
"8": {"class_type": "PreviewImage", "inputs": {"images": ["6", 0]}},
"9": {"class_type": "ViewText", "inputs": {"text": ["4", 1]}},
"10": {"class_type": "ViewText", "inputs": {"text": ["5", 1]}},
"11": {"class_type": "ViewText", "inputs": {"text": ["7", 0]}},
}
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--comfy-url", default="http://127.0.0.1:8188")
parser.add_argument("--policy-url", default="http://127.0.0.1:8787")
parser.add_argument(
"--image",
required=True,
help="Filename relative to the running ComfyUI instance's input directory.",
)
parser.add_argument("--timeout", type=float, default=180.0)
args = parser.parse_args()
base = args.comfy_url.rstrip("/")
body = json.dumps(
{
"prompt": _graph(args.image, args.policy_url),
"client_id": str(uuid.uuid4()),
}
).encode("utf-8")
request = urllib.request.Request(
f"{base}/prompt",
data=body,
headers={"Content-Type": "application/json"},
)
with urllib.request.urlopen(request, timeout=30) as response:
queued = json.load(response)
prompt_id = queued["prompt_id"]
deadline = time.monotonic() + args.timeout
while time.monotonic() < deadline:
with urllib.request.urlopen(
f"{base}/history/{prompt_id}",
timeout=10,
) as response:
history = json.load(response)
if prompt_id not in history:
time.sleep(0.5)
continue
entry = history[prompt_id]
result = {
"prompt_id": prompt_id,
"status": entry.get("status"),
"outputs": entry.get("outputs"),
}
print(json.dumps(result, indent=2))
if entry.get("status", {}).get("status_str") != "success":
raise SystemExit(1)
return
raise SystemExit(f"Timed out waiting for prompt {prompt_id}.")
if __name__ == "__main__":
main()
+2
View File
@@ -12,7 +12,9 @@ import json
import time
import torch
from _bootstrap import bootstrap
bootstrap()
BACKENDS = (
"florence-base",
+176
View File
@@ -0,0 +1,176 @@
"""Run adaptive temporal reasoning on a real local video and real VLM.
Example:
python tests/manual_video_intelligence_smoke.py \
/mnt/d/002.mp4 \
--model "Qwen 3 VL 2B Instruct" \
--output /mnt/d/comfyui-repair/video-intelligence-audit/result.json
"""
from __future__ import annotations
import argparse
import importlib.util
import json
import sys
import time
from pathlib import Path
import av
import torch
REPOSITORY = Path(__file__).resolve().parents[1]
if REPOSITORY.name != "ComfyUI_VLM_nodes":
specification = importlib.util.spec_from_file_location(
"ComfyUI_VLM_nodes",
REPOSITORY / "__init__.py",
submodule_search_locations=[str(REPOSITORY)],
)
if specification is None or specification.loader is None:
raise RuntimeError(f"Could not load package from {REPOSITORY}.")
package = importlib.util.module_from_spec(specification)
sys.modules["ComfyUI_VLM_nodes"] = package
specification.loader.exec_module(package)
from ComfyUI_VLM_nodes.nodes.modern_vlm import ModernVLMPredictor
from ComfyUI_VLM_nodes.nodes.video_intelligence import (
build_video_reasoning_prompt,
parse_video_reasoning_output,
resize_video_for_analysis,
sample_video_frames,
)
def load_video(path: Path) -> tuple[torch.Tensor, float]:
container = av.open(str(path))
try:
stream = container.streams.video[0]
rate = stream.average_rate or stream.guessed_rate
if rate is None:
raise RuntimeError("The video does not report a frame rate.")
frames = [
torch.from_numpy(frame.to_ndarray(format="rgb24")).to(torch.float32)
/ 255.0
for frame in container.decode(stream)
]
finally:
container.close()
if not frames:
raise RuntimeError("The video contains no decodable frames.")
return torch.stack(frames), float(rate)
def main() -> int:
parser = argparse.ArgumentParser()
parser.add_argument("video", type=Path)
parser.add_argument(
"--model",
default="Qwen 3 VL 2B Instruct",
)
parser.add_argument("--max-frames", type=int, default=12)
parser.add_argument("--analysis-max-side", type=int, default=448)
parser.add_argument("--max-new-tokens", type=int, default=512)
parser.add_argument("--output", type=Path)
args = parser.parse_args()
frames, fps = load_video(args.video)
sampled, selection, diagnostics = sample_video_frames(
frames,
fps=fps,
max_frames=args.max_frames,
strategy="Hybrid: scene + motion + tracks",
minimum_gap_seconds=0.2,
)
prompt = build_video_reasoning_prompt(
selection,
task="Detailed temporal summary",
question="What happens, and how do the people behave over time?",
max_events=12,
)
analysis_frames = resize_video_for_analysis(
sampled,
max_side=args.analysis_max_side,
)
predictor = ModernVLMPredictor(
args.model,
"",
"ComfyUI managed (BF16)",
"Auto (SDPA)",
)
started = time.perf_counter()
try:
raw = predictor.generate(
images=None,
prompt=prompt,
system_prompt=(
"You are a precise temporal video analyst. Return one JSON "
"object that obeys the supplied schema."
),
max_new_tokens=args.max_new_tokens,
temperature=0.0,
top_p=1.0,
video_frames=analysis_frames,
fps=fps,
video_selection=selection,
)
finally:
predictor.close()
reasoning_seconds = time.perf_counter() - started
result = {
"video": str(args.video),
"model": args.model,
"source_shape": list(frames.shape),
"fps": fps,
"selection": selection.to_dict(),
"sampling": diagnostics,
"analysis_shape": list(analysis_frames.shape),
"reasoning_seconds": reasoning_seconds,
"raw_response": raw,
"cuda_peak_gib": (
torch.cuda.max_memory_allocated() / 2**30
if torch.cuda.is_available()
else 0.0
),
}
try:
summary, events, normalized = parse_video_reasoning_output(raw, selection)
except (TypeError, ValueError) as exc:
result["structured_output_valid"] = False
result["structured_output_error"] = str(exc)
if args.output:
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(
json.dumps(
result,
ensure_ascii=False,
allow_nan=False,
indent=2,
sort_keys=True,
),
encoding="utf-8",
)
raise
result.update(
{
"structured_output_valid": True,
"summary": summary,
"events": events.to_dict(),
"normalized_response": json.loads(normalized),
}
)
encoded = json.dumps(
result,
ensure_ascii=False,
allow_nan=False,
indent=2,
sort_keys=True,
)
if args.output:
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(encoded, encoding="utf-8")
print(encoded)
return 0
if __name__ == "__main__":
raise SystemExit(main())
+111
View File
@@ -0,0 +1,111 @@
import json
import threading
import time
import pytest
import torch
from ComfyUI_VLM_nodes.nodes.acceleration import (
VLMImagePixelBudget,
VLMPerformanceProfile,
optimize_image_pixels,
)
from ComfyUI_VLM_nodes.nodes.runtime import (
CachedModelNode,
tensor_batch_to_pil,
tensor_to_pil,
)
def test_batch_conversion_matches_single_frame_contract():
images = torch.tensor(
[
[
[[float("nan"), 0.5, 2.0], [-1.0, 0.25, 1.0]],
[[0.0, 0.0, 0.0], [1.0, 1.0, 1.0]],
],
[
[[255.0, 128.0, 0.0], [0.0, 64.0, 255.0]],
[[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]],
],
]
)
batch = tensor_batch_to_pil(images)
assert len(batch) == 2
for index, converted in enumerate(batch):
assert converted.mode == "RGB"
assert converted.size == (2, 2)
assert converted.tobytes() == tensor_to_pil(images, index).tobytes()
with pytest.raises(IndexError, match="only has batch index 0"):
tensor_to_pil(images[0], 1)
def test_pixel_budget_preserves_aspect_and_patch_multiple():
images = torch.rand((3, 1080, 1920, 3), dtype=torch.float32)
output, report = optimize_image_pixels(
images,
max_megapixels=0.5,
max_edge=1024,
multiple=14,
resize_quality="Fast (area)",
)
assert output.ndim == 4
assert output.shape[0] == 3
assert output.shape[1] % 14 == 0
assert output.shape[2] % 14 == 0
assert output.shape[1] * output.shape[2] <= 500_000
assert output.shape[2] <= 1024
assert report["visual_work_reduction"] > 4
assert output.shape[2] / output.shape[1] == pytest.approx(16 / 9, rel=0.03)
def test_pixel_budget_never_upscales():
image = torch.rand((240, 320, 3), dtype=torch.float32)
output, report = optimize_image_pixels(
image,
max_megapixels=2.0,
max_edge=2048,
multiple=1,
resize_quality="Quality (bicubic)",
)
assert output is image
assert report["resized"] is False
def test_performance_nodes_return_standard_comfy_values():
profile = VLMPerformanceProfile().profile("Live / robotics")
assert profile[:5] == (24, 0.5, 896, 8, False)
assert json.loads(profile[5])["profile"] == "Live / robotics"
optimized = VLMImagePixelBudget().optimize(
torch.rand((1, 1000, 1600, 3)),
0.5,
1024,
"14",
"Fast (area)",
)
assert optimized[1] % 14 == 0
assert optimized[2] % 14 == 0
def test_cached_model_node_prevents_duplicate_concurrent_loads():
node = CachedModelNode()
factory_calls = []
handles = []
def factory():
factory_calls.append(1)
time.sleep(0.02)
return object()
def load():
handles.append(node.get_or_create_model("same-model", factory))
threads = [threading.Thread(target=load) for _ in range(8)]
for thread in threads:
thread.start()
for thread in threads:
thread.join()
assert len(factory_calls) == 1
assert len({id(handle) for handle in handles}) == 1
+110
View File
@@ -0,0 +1,110 @@
"""Keep the shipped documentation honest about what the pack actually contains.
The README node reference drifted to 47 undocumented nodes before these checks
existed, and the packaged license metadata disagreed with the LICENSE file.
Both are cheap to assert and expensive to notice by hand.
"""
from __future__ import annotations
import re
from pathlib import Path
import ComfyUI_VLM_nodes as package
REPOSITORY = Path(package.__file__).parent
def read(name: str) -> str:
return (REPOSITORY / name).read_text(encoding="utf-8")
def project_field(name: str) -> str:
"""Read a top-level [project] string field.
Deliberately regex-based rather than tomllib: this suite also runs on
Python 3.10, which has no tomllib in the standard library.
"""
match = re.search(rf'^{name}\s*=\s*"([^"]+)"', read("pyproject.toml"), re.M)
assert match is not None, f"pyproject.toml has no {name} field."
return match.group(1)
def test_every_registered_node_appears_in_the_readme():
readme = read("README.md")
documented = set(re.findall(r"`([^`]+)`", readme))
missing = sorted(set(package.NODE_CLASS_MAPPINGS) - documented)
assert not missing, (
"These nodes are registered but never named in README.md. "
f"Add them to the node reference: {missing}"
)
def test_node_reference_matches_registered_output_types():
row_pattern = re.compile(
r"^\|[^|]+\|\s*`(?P<node_id>[^`]+)`\s*\|(?P<outputs>[^|]*)\|$",
re.M,
)
documented = {
match.group("node_id"): tuple(
re.findall(r"`([^`]+)`", match.group("outputs"))
)
for match in row_pattern.finditer(read("README.md"))
}
mismatches = {}
for node_id, node_class in package.NODE_CLASS_MAPPINGS.items():
expected = tuple(
"*" if output is any else str(output)
for output in node_class.RETURN_TYPES
)
if documented.get(node_id) != expected:
mismatches[node_id] = {
"documented": documented.get(node_id),
"registered": expected,
}
assert not mismatches, (
"README.md output schemas do not match the registered RETURN_TYPES: "
f"{mismatches}"
)
def test_declared_license_matches_the_license_file():
declared = project_field("license")
license_text = read("LICENSE")
if "Apache License" in license_text:
expected = "Apache-2.0"
elif "MIT License" in license_text:
expected = "MIT"
else:
raise AssertionError("Could not identify the license in LICENSE.")
assert declared == expected, (
f"pyproject.toml declares {declared!r} but LICENSE is {expected}. "
"This metadata is embedded in built distribution artifacts."
)
def test_changelog_documents_the_current_version():
version = project_field("version")
changelog = read("CHANGELOG.md")
assert f"[{version}]" in changelog, (
f"pyproject version {version} has no CHANGELOG.md entry. The Comfy "
"Registry only publishes on a version change, so every release needs "
"one."
)
def test_contributor_and_security_docs_are_present():
for name in ("CONTRIBUTING.md", "SECURITY.md", "CHANGELOG.md", "LICENSE"):
assert (REPOSITORY / name).is_file(), f"{name} is missing."
def test_issue_templates_are_valid_and_request_diagnostics():
template_dir = REPOSITORY / ".github" / "ISSUE_TEMPLATE"
bug_report = (template_dir / "bug_report.yml").read_text(encoding="utf-8")
# Environment detail is what the historically unresolvable reports lacked.
assert "VLMRuntimeDiagnostics" in bug_report or "Diagnostics" in bug_report
assert "Node pack version" in bug_report
+904
View File
@@ -0,0 +1,904 @@
from __future__ import annotations
import json
from pathlib import Path
from types import SimpleNamespace
import pytest
import torch
from ComfyUI_VLM_nodes.nodes import hosted_api
class FakeHttpClient:
def __init__(self, **kwargs):
self.kwargs = kwargs
class FakeResponses:
def __init__(self, calls, failure=None, response_text="secure response"):
self.calls = calls
self.failure = failure
self.response_text = response_text
def create(self, **kwargs):
self.calls.append(("responses", kwargs))
if self.failure is not None:
raise self.failure
return SimpleNamespace(output_text=self.response_text)
class FakeChatCompletions:
def __init__(self, calls, failure=None, response_text="secure response"):
self.calls = calls
self.failure = failure
self.response_text = response_text
def create(self, **kwargs):
self.calls.append(("chat", kwargs))
if self.failure is not None:
raise self.failure
message = SimpleNamespace(content=self.response_text)
return SimpleNamespace(choices=[SimpleNamespace(message=message)])
def fake_openai_module(calls, failure=None, response_text="secure response"):
class FakeOpenAI:
def __init__(self, **kwargs):
calls.append(("client", kwargs))
self.responses = FakeResponses(
calls,
failure=failure,
response_text=response_text,
)
self.chat = SimpleNamespace(
completions=FakeChatCompletions(
calls,
failure=failure,
response_text=response_text,
)
)
def close(self):
calls.append(("close", {}))
return SimpleNamespace(
OpenAI=FakeOpenAI,
DefaultHttpxClient=FakeHttpClient,
)
def test_api_schemas_never_accept_plaintext_keys():
for node_class in (hosted_api.PromptGenerateAPI, hosted_api.HostedVLMAPI):
schema = node_class.INPUT_TYPES()
all_inputs = {
**schema.get("required", {}),
**schema.get("optional", {}),
**schema.get("hidden", {}),
}
assert "api_key" not in all_inputs
assert "credential_source" in all_inputs
assert "STRING" not in repr(all_inputs["credential_source"][0])
assert "web_search" in all_inputs
assert "output_format" in all_inputs
assert "json_schema" in all_inputs
assert "schema_api_style" in all_inputs
def test_json_schema_parser_blocks_remote_refs_and_bounds_input():
for keyword in ("$ref", "$dynamicRef", "$recursiveRef"):
with pytest.raises(ValueError, match="only local fragment"):
hosted_api.parse_json_schema(
"JSON Schema",
json.dumps(
{
"type": "object",
"properties": {
"payload": {
keyword: "https://attacker.example/schema.json"
}
},
}
),
)
with pytest.raises(ValueError, match="64,000"):
hosted_api.parse_json_schema("JSON Schema", "x" * 64_001)
def test_local_structured_output_validation_is_strict_and_normalized():
schema_text = json.dumps(
{
"type": "object",
"properties": {"count": {"type": "integer"}},
"required": ["count"],
"additionalProperties": False,
}
)
schema = hosted_api.parse_json_schema("JSON Schema", schema_text)
assert hosted_api.validate_structured_output(
'```json\n{"count": 2}\n```',
"JSON Schema",
schema,
) == '{\n "count": 2\n}'
with pytest.raises(RuntimeError, match=r"\$\.count \(type constraint\)"):
hosted_api.validate_structured_output(
'{"count": "two"}',
"JSON Schema",
schema,
)
with pytest.raises(RuntimeError, match="valid JSON"):
hosted_api.validate_structured_output(
'{"count":',
"JSON Schema",
schema,
)
def test_provider_catalog_uses_current_bound_credentials_and_endpoints():
assert len(hosted_api.PROVIDER_PROFILES) >= 18
expected = {
"OpenAI": "OPENAI_API_KEY",
"Google Gemini": "GEMINI_API_KEY",
"Anthropic": "ANTHROPIC_API_KEY",
"xAI": "XAI_API_KEY",
"DeepSeek": "DEEPSEEK_API_KEY",
"Groq": "GROQ_API_KEY",
"Mistral": "MISTRAL_API_KEY",
"Together AI": "TOGETHER_API_KEY",
"OpenRouter": "OPENROUTER_API_KEY",
"Custom / Local": "CUSTOM_API_KEY",
}
providers = {
profile.provider: profile.api_key_env
for profile in hosted_api.PROVIDER_PROFILES.values()
}
assert expected.items() <= providers.items()
for profile in hosted_api.PROVIDER_PROFILES.values():
if profile.base_url is not None:
assert profile.base_url.startswith("https://")
@pytest.mark.parametrize(
"url",
[
"http://example.com/v1",
"ftp://127.0.0.1/v1",
"https://user:secret@example.com/v1",
"https://example.com/v1?api_key=secret",
"not-a-url",
],
)
def test_custom_endpoint_rejects_unsafe_urls(url):
with pytest.raises(ValueError):
hosted_api.validate_custom_base_url(url)
@pytest.mark.parametrize(
"url",
[
"http://127.0.0.1:8000/v1",
"http://[::1]:11434/v1",
"http://localhost:1234/v1",
"https://example.com/v1/",
],
)
def test_custom_endpoint_accepts_https_or_loopback(url):
normalized, loopback = hosted_api.validate_custom_base_url(url)
assert normalized.startswith(("http://", "https://"))
assert loopback is (url.startswith("http://"))
def test_built_in_key_cannot_be_redirected(monkeypatch):
profile = hosted_api.provider_profile("OpenAI — GPT-5.6 Terra")
monkeypatch.setenv("OPENAI_API_KEY", "sk-real-secret-value")
with pytest.raises(ValueError, match="pinned to official hosts"):
hosted_api.resolve_endpoint(profile, "https://attacker.example/v1")
def test_legacy_plaintext_value_is_rejected_without_echo(monkeypatch):
profile = hosted_api.provider_profile("OpenAI — GPT-5.6 Terra")
secret = "sk-legacy-plaintext-that-must-not-appear"
with pytest.raises(ValueError) as captured:
hosted_api.resolve_api_key(profile, secret, loopback=False)
assert secret not in str(captured.value)
assert "legacy plaintext API key was removed" in str(captured.value)
def test_redaction_removes_exact_encoded_and_header_credentials():
secret = "sk-ant-example-SECRET_123456789"
message = (
f"Authorization: Bearer {secret}; api_key={secret}; "
f"url=https://user:{secret}@example.com; encoded={secret}"
)
redacted = hosted_api.redact_sensitive(message, (secret,))
assert secret not in redacted
assert "Bearer" not in redacted
assert "[REDACTED]" in redacted
def test_responses_call_is_stateless_private_and_provider_bound(monkeypatch):
calls = []
secret = "sk-openai-provider-bound-secret"
monkeypatch.setenv("OPENAI_API_KEY", secret)
monkeypatch.setattr(
hosted_api,
"require_module",
lambda *_args: fake_openai_module(calls),
)
node = hosted_api.PromptGenerateAPI()
assert not hasattr(node, "session_history")
result = node.generate_prompt(
"OpenAI — GPT-5.6 Terra",
False,
hosted_api.PROVIDER_CREDENTIAL,
"A scene",
"Improve it",
0,
0,
stream_output=False,
)
assert result == ("secure response",)
client_kwargs = next(payload for kind, payload in calls if kind == "client")
assert client_kwargs["api_key"] == secret
assert "base_url" not in client_kwargs
assert client_kwargs["http_client"].kwargs["follow_redirects"] is False
assert client_kwargs["http_client"].kwargs["trust_env"] is False
request = next(payload for kind, payload in calls if kind == "responses")
assert request["model"] == "gpt-5.6-terra"
assert request["store"] is False
assert "previous_response_id" not in request
assert "metadata" not in request
def test_openai_combines_web_search_structured_output_and_stream_contract(
monkeypatch,
):
calls = []
monkeypatch.setenv("OPENAI_API_KEY", "sk-test-structured-search")
monkeypatch.setattr(
hosted_api,
"require_module",
lambda import_name, *_args: (
fake_openai_module(calls, response_text='{"answer":"grounded"}')
if import_name == "openai"
else __import__(import_name)
),
)
schema = json.dumps(
{
"type": "object",
"properties": {"answer": {"type": "string"}},
"required": ["answer"],
"additionalProperties": False,
}
)
result = hosted_api.PromptGenerateAPI().generate_prompt(
"OpenAI — GPT-5.6 Sol",
False,
hosted_api.PROVIDER_CREDENTIAL,
"Find a current fact",
"",
0,
0,
web_search=True,
output_format="JSON Schema",
json_schema=schema,
stream_output=False,
)
assert result == ('{\n "answer": "grounded"\n}',)
request = next(payload for kind, payload in calls if kind == "responses")
assert request["tools"] == [{"type": "web_search"}]
assert request["text"]["format"]["type"] == "json_schema"
assert request["text"]["format"]["strict"] is True
assert request["text"]["format"]["schema"]["required"] == ["answer"]
assert "JSON Schema:" in request["instructions"]
def test_unsupported_web_search_fails_before_network(monkeypatch):
monkeypatch.setenv("DEEPSEEK_API_KEY", "deepseek-test-secret")
with pytest.raises(ValueError, match="does not expose native web search"):
hosted_api.PromptGenerateAPI().generate_prompt(
"DeepSeek — V4 Flash",
False,
hosted_api.PROVIDER_CREDENTIAL,
"Search now",
"",
0,
0,
web_search=True,
stream_output=False,
)
def test_responses_stream_collects_deltas_and_closes():
class Stream(list):
closed = False
def close(self):
self.closed = True
stream = Stream(
[
SimpleNamespace(type="response.created"),
SimpleNamespace(type="response.output_text.delta", delta="hello "),
SimpleNamespace(type="response.output_text.delta", delta="world"),
]
)
client = SimpleNamespace(
responses=SimpleNamespace(create=lambda **kwargs: stream)
)
assert hosted_api._stream_responses(client, {"model": "test"}, None) == (
"hello world"
)
assert stream.closed is True
def test_chat_stream_collects_deltas_and_closes():
class Stream(list):
closed = False
def close(self):
self.closed = True
stream = Stream(
[
SimpleNamespace(
choices=[
SimpleNamespace(delta=SimpleNamespace(content="frame "))
]
),
SimpleNamespace(
choices=[
SimpleNamespace(delta=SimpleNamespace(content="ready"))
]
),
]
)
client = SimpleNamespace(
chat=SimpleNamespace(
completions=SimpleNamespace(create=lambda **kwargs: stream)
)
)
assert hosted_api._stream_chat(client, {"model": "test"}, None) == (
"frame ready"
)
assert stream.closed is True
def test_provider_failure_never_echoes_api_key(monkeypatch):
calls = []
secret = "sk-secret-reflected-by-provider-123456"
failure = RuntimeError(f"Authorization: Bearer {secret} api_key={secret}")
monkeypatch.setenv("OPENAI_API_KEY", secret)
monkeypatch.setattr(
hosted_api,
"require_module",
lambda *_args: fake_openai_module(calls, failure=failure),
)
with pytest.raises(RuntimeError) as captured:
hosted_api.PromptGenerateAPI().generate_prompt(
"OpenAI — GPT-5.6 Terra",
False,
hosted_api.PROVIDER_CREDENTIAL,
"hello",
"",
0,
0,
stream_output=False,
)
assert secret not in str(captured.value)
assert "[REDACTED]" in str(captured.value)
def test_anthropic_uses_native_messages_and_keeps_key_out_of_body(monkeypatch):
calls = []
secret = "sk-ant-native-secret"
class FakeResponse:
def raise_for_status(self):
return None
def json(self):
return {
"content": [
{"type": "text", "text": "native Anthropic response"}
]
}
class FakeClient:
def __init__(self, **kwargs):
calls.append(("client", kwargs))
def post(self, url, **kwargs):
calls.append(("post", {"url": url, **kwargs}))
return FakeResponse()
def close(self):
calls.append(("close", {}))
monkeypatch.setenv("ANTHROPIC_API_KEY", secret)
monkeypatch.setattr(
hosted_api,
"require_module",
lambda import_name, *_args: (
SimpleNamespace(Client=FakeClient)
if import_name == "httpx"
else pytest.fail(f"Unexpected module request: {import_name}")
),
)
result = hosted_api.HostedVLMAPI().analyze(
"Anthropic — Claude Sonnet 5",
hosted_api.PROVIDER_CREDENTIAL,
"Read this image.",
"Be concise.",
1,
512,
80,
"auto",
images=torch.rand((1, 48, 64, 3)),
stream_output=False,
)
assert result == (
"native Anthropic response",
"claude-sonnet-5",
1,
)
request = next(payload for kind, payload in calls if kind == "post")
assert request["url"] == "https://api.anthropic.com/v1/messages"
assert request["headers"]["x-api-key"] == secret
assert secret not in repr(request["json"])
content = request["json"]["messages"][0]["content"]
assert content[1]["type"] == "image"
assert content[1]["source"]["type"] == "base64"
assert request["json"]["stream"] is False
def test_anthropic_native_stream_collects_text_deltas(monkeypatch):
class FakeStreamResponse:
def __enter__(self):
return self
def __exit__(self, *_args):
return False
def raise_for_status(self):
return None
def iter_lines(self):
yield 'event: content_block_delta'
yield (
'data: {"type":"content_block_delta","delta":'
'{"type":"text_delta","text":"hello "}}'
)
yield (
'data: {"type":"content_block_delta","delta":'
'{"type":"text_delta","text":"world"}}'
)
class FakeClient:
def __init__(self, **_kwargs):
pass
def stream(self, *_args, **_kwargs):
return FakeStreamResponse()
def close(self):
return None
monkeypatch.setattr(
hosted_api,
"require_module",
lambda import_name, *_args: (
SimpleNamespace(Client=FakeClient)
if import_name == "httpx"
else __import__(import_name)
),
)
profile = hosted_api.provider_profile("Anthropic — Claude Sonnet 5")
result = hosted_api._call_anthropic_api(
profile=profile,
model=profile.model,
endpoint=profile.base_url,
api_key="sk-ant-stream",
system_prompt="Be concise.",
prompt="Hello",
image_data=[],
timeout_seconds=30,
max_output_tokens=100,
stream_output=True,
use_system_proxy=False,
unique_id=None,
web_search=False,
output_format="Text",
output_schema=None,
)
assert result == "hello world"
def test_anthropic_native_search_and_structured_contracts(monkeypatch):
calls = []
class FakeResponse:
def raise_for_status(self):
return None
def json(self):
return {"content": [{"type": "text", "text": '{"answer":"yes"}'}]}
class FakeClient:
def __init__(self, **kwargs):
calls.append(("client", kwargs))
def post(self, url, **kwargs):
calls.append(("post", {"url": url, **kwargs}))
return FakeResponse()
def close(self):
return None
monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-ant-contract")
monkeypatch.setattr(
hosted_api,
"require_module",
lambda import_name, *_args: (
SimpleNamespace(Client=FakeClient)
if import_name == "httpx"
else __import__(import_name)
),
)
schema = (
'{"type":"object","properties":{"answer":{"type":"string"}},'
'"required":["answer"],"additionalProperties":false}'
)
result = hosted_api.PromptGenerateAPI().generate_prompt(
"Anthropic — Claude Sonnet 5",
False,
hosted_api.PROVIDER_CREDENTIAL,
"Return a value",
"",
0,
0,
output_format="JSON Schema",
json_schema=schema,
stream_output=False,
)
assert json.loads(result[0]) == {"answer": "yes"}
structured = next(payload for kind, payload in calls if kind == "post")
assert structured["json"]["output_config"]["format"]["type"] == "json_schema"
calls.clear()
hosted_api.PromptGenerateAPI().generate_prompt(
"Anthropic — Claude Sonnet 5",
False,
hosted_api.PROVIDER_CREDENTIAL,
"Search the web",
"",
0,
0,
web_search=True,
output_format="JSON Schema",
json_schema=schema,
stream_output=False,
)
searched = next(payload for kind, payload in calls if kind == "post")
assert searched["json"]["tools"][0]["type"] == "web_search_20260318"
assert searched["json"]["tools"][0]["allowed_callers"] == ["direct"]
assert "output_config" not in searched["json"]
def test_gemini_native_search_vision_and_schema_contract(monkeypatch):
calls = []
secret = "gemini-provider-secret"
class FakeResponse:
def raise_for_status(self):
return None
def json(self):
return {
"candidates": [
{
"content": {
"parts": [{"text": '{"objects":["tree"]}'}]
}
}
]
}
class FakeClient:
def __init__(self, **kwargs):
calls.append(("client", kwargs))
def post(self, url, **kwargs):
calls.append(("post", {"url": url, **kwargs}))
return FakeResponse()
def close(self):
return None
monkeypatch.setenv("GEMINI_API_KEY", secret)
monkeypatch.setattr(
hosted_api,
"require_module",
lambda import_name, *_args: (
SimpleNamespace(Client=FakeClient)
if import_name == "httpx"
else __import__(import_name)
),
)
schema = (
'{"type":"object","properties":{"objects":{"type":"array",'
'"items":{"type":"string"}}},"required":["objects"]}'
)
result = hosted_api.HostedVLMAPI().analyze(
"Google — Gemini 3.6 Flash",
hosted_api.PROVIDER_CREDENTIAL,
"Identify objects using current context.",
"Be precise.",
1,
512,
80,
"auto",
images=torch.rand((1, 32, 48, 3)),
web_search=True,
output_format="JSON Schema",
json_schema=schema,
stream_output=False,
)
assert json.loads(result[0]) == {"objects": ["tree"]}
request = next(payload for kind, payload in calls if kind == "post")
assert request["url"].endswith(
"/models/gemini-3.6-flash:generateContent"
)
assert request["headers"]["x-goog-api-key"] == secret
assert secret not in repr(request["json"])
assert request["json"]["tools"] == [{"google_search": {}}]
assert (
request["json"]["generationConfig"]["responseFormat"]["text"]["schema"][
"required"
]
== ["objects"]
)
inline = request["json"]["contents"][0]["parts"][1]["inlineData"]
assert inline["mimeType"] == "image/jpeg"
assert inline["data"]
def test_gemini_native_stream_collects_sse_deltas(monkeypatch):
class FakeStreamResponse:
def __enter__(self):
return self
def __exit__(self, *_args):
return False
def raise_for_status(self):
return None
def iter_lines(self):
yield (
'data: {"candidates":[{"content":{"parts":'
'[{"text":"frame "}]}}]}'
)
yield (
'data: {"candidates":[{"content":{"parts":'
'[{"text":"ready"}]}}]}'
)
class FakeClient:
def __init__(self, **_kwargs):
pass
def stream(self, *_args, **_kwargs):
return FakeStreamResponse()
def close(self):
return None
monkeypatch.setattr(
hosted_api,
"require_module",
lambda import_name, *_args: (
SimpleNamespace(Client=FakeClient)
if import_name == "httpx"
else __import__(import_name)
),
)
profile = hosted_api.provider_profile("Google — Gemini 3.6 Flash")
result = hosted_api._call_gemini_api(
profile=profile,
model=profile.model,
api_key="gemini-stream",
system_prompt="Be concise.",
prompt="Describe.",
image_data=[],
timeout_seconds=30,
max_output_tokens=100,
stream_output=True,
use_system_proxy=False,
unique_id=None,
web_search=True,
output_format="Text",
output_schema=None,
)
assert result == "frame ready"
def test_vlm_uniformly_samples_and_bounds_image_batch(monkeypatch):
calls = []
monkeypatch.setenv("OPENAI_API_KEY", "sk-test-only")
monkeypatch.setattr(
hosted_api,
"require_module",
lambda *_args: fake_openai_module(calls),
)
images = torch.rand((10, 96, 128, 3), dtype=torch.float32)
result = hosted_api.HostedVLMAPI().analyze(
"OpenAI — GPT-5.6 Terra",
hosted_api.PROVIDER_CREDENTIAL,
"Compare the sampled frames.",
"Be precise.",
4,
768,
82,
"low",
images=images,
stream_output=False,
)
assert result == ("secure response", "gpt-5.6-terra", 4)
request = next(payload for kind, payload in calls if kind == "responses")
content = request["input"][0]["content"]
image_parts = [part for part in content if part["type"] == "input_image"]
assert len(image_parts) == 4
assert all(part["image_url"].startswith("data:image/jpeg;base64,") for part in image_parts)
assert all(part["detail"] == "low" for part in image_parts)
def test_open_source_vlm_llama_cpp_schema_dialect_and_local_validation(
monkeypatch,
):
calls = []
monkeypatch.setattr(
hosted_api,
"require_module",
lambda import_name, *_args: (
fake_openai_module(calls, response_text='{"objects":["cat"]}')
if import_name == "openai"
else __import__(import_name)
),
)
schema = json.dumps(
{
"type": "object",
"properties": {
"objects": {
"type": "array",
"items": {"type": "string"},
}
},
"required": ["objects"],
"additionalProperties": False,
}
)
result = hosted_api.HostedVLMAPI().analyze(
"Custom / Local — OpenAI compatible",
hosted_api.LOCAL_NO_KEY,
"List visible objects.",
"Be precise.",
1,
512,
80,
"auto",
images=torch.rand((1, 32, 48, 3)),
base_url="http://127.0.0.1:8080/v1",
model_override="local-vlm",
output_format="JSON Schema",
json_schema=schema,
schema_api_style="llama.cpp JSON Schema",
stream_output=False,
)
assert result == (
'{\n "objects": [\n "cat"\n ]\n}',
"local-vlm",
1,
)
request = next(payload for kind, payload in calls if kind == "chat")
assert request["response_format"] == {
"type": "json_schema",
"schema": json.loads(schema),
}
image = request["messages"][1]["content"][1]
assert image["type"] == "image_url"
assert image["image_url"]["url"].startswith("data:image/jpeg;base64,")
def test_custom_openai_schema_style_uses_standard_wrapper(monkeypatch):
calls = []
monkeypatch.setattr(
hosted_api,
"require_module",
lambda import_name, *_args: (
fake_openai_module(calls, response_text='{"ok":true}')
if import_name == "openai"
else __import__(import_name)
),
)
schema = '{"type":"object","properties":{"ok":{"type":"boolean"}},"required":["ok"]}'
result, _, _ = hosted_api.execute_hosted(
model_name="Custom / Local — OpenAI compatible",
credential_source=hosted_api.LOCAL_NO_KEY,
prompt="Return status.",
system_prompt="Be exact.",
base_url="http://localhost:8000/v1",
model_override="local",
api_mode="Chat Completions",
timeout_seconds=30,
max_output_tokens=100,
reasoning_effort="none",
seed=0,
stream_output=False,
use_system_proxy=False,
unique_id=None,
output_format="JSON Schema",
json_schema=schema,
)
assert json.loads(result) == {"ok": True}
request = next(payload for kind, payload in calls if kind == "chat")
assert request["response_format"]["json_schema"]["strict"] is True
assert request["response_format"]["json_schema"]["schema"]["required"] == [
"ok"
]
def test_groq_auto_uses_documented_chat_route_for_structured_output(monkeypatch):
calls = []
monkeypatch.setenv("GROQ_API_KEY", "gsk-test-structured")
monkeypatch.setattr(
hosted_api,
"require_module",
lambda import_name, *_args: (
fake_openai_module(calls, response_text='{"ok":true}')
if import_name == "openai"
else __import__(import_name)
),
)
hosted_api.execute_hosted(
model_name="Groq — GPT-OSS 20B",
credential_source=hosted_api.PROVIDER_CREDENTIAL,
prompt="Return status.",
system_prompt="Be exact.",
base_url="",
model_override="",
api_mode="Auto",
timeout_seconds=30,
max_output_tokens=100,
reasoning_effort="none",
seed=0,
stream_output=False,
use_system_proxy=False,
unique_id=None,
output_format="JSON Schema",
json_schema=(
'{"type":"object","properties":{"ok":{"type":"boolean"}},'
'"required":["ok"],"additionalProperties":false}'
),
)
assert any(kind == "chat" for kind, _payload in calls)
assert not any(kind == "responses" for kind, _payload in calls)
def test_frontend_scrubs_legacy_key_before_graph_configuration():
web_root = Path(__file__).resolve().parents[1] / "web" / "js"
source = (
web_root / "apiSecurity.js"
).read_text("utf-8")
assert "beforeConfigureGraph" in source
assert "delete values.api_key" in source
assert "CREDENTIAL_WIDGET_INDEX = 2" in source
view_text = (web_root / "viewText.js").read_text("utf-8")
assert '"PromptGenerateAPI"' in view_text
assert '"HostedVLMAPI"' in view_text
+496
View File
@@ -0,0 +1,496 @@
"""Contract tests for the llama.cpp multimodal nodes in ``nodes/llavaloader.py``.
Covers batch handling, the vision message envelope, projector wiring, and the
cached-handle lifecycle. No llama.cpp wheel, mmproj, or GGUF weights required.
"""
from __future__ import annotations
import base64
from pathlib import Path
import pytest
import torch
from ComfyUI_VLM_nodes.nodes import llavaloader
from ComfyUI_VLM_nodes.nodes.runtime import LlamaHandle, LlavaClipConfig
MODEL_FILE = "llava.gguf"
CLIP_FILE = "mmproj.gguf"
class FakeLlama:
def __init__(self, contents: list[str] | None = None):
self.contents = contents or ["a description"]
self.calls: list[dict] = []
def create_chat_completion(self, **kwargs):
index = min(len(self.calls), len(self.contents) - 1)
self.calls.append(kwargs)
return {"choices": [{"message": {"content": self.contents[index]}}]}
class FakeHandle:
instances: list[FakeHandle] = []
def __init__(self, model_path, **kwargs):
self.model_path = model_path
self.kwargs = kwargs
self.closed = False
self.llama = FakeLlama()
FakeHandle.instances.append(self)
def ensure_loaded(self):
return self.llama
def close(self):
self.closed = True
@pytest.fixture
def resolved_paths(monkeypatch):
root = Path("/models/LLavacheckpoints")
monkeypatch.setattr(llavaloader, "resolve_model_path", lambda name: root / name)
return root
@pytest.fixture
def fake_handles(monkeypatch):
FakeHandle.instances = []
monkeypatch.setattr(llavaloader, "LlamaHandle", FakeHandle)
return FakeHandle
def image_batch(count: int = 1, size: int = 4) -> torch.Tensor:
"""A ComfyUI BHWC float image batch."""
return torch.rand(count, size, size, 3)
# --------------------------------------------------------------------------
# Widget ordering (see issue #156).
# --------------------------------------------------------------------------
def test_llava_sampler_simple_widget_order_is_frozen():
assert list(llavaloader.LLavaSamplerSimple.INPUT_TYPES()["required"]) == [
"image",
"prompt",
"model",
"temperature",
]
def test_llava_sampler_advanced_widget_order_is_frozen():
assert list(llavaloader.LLavaSamplerAdvanced.INPUT_TYPES()["required"]) == [
"image",
"system_msg",
"prompt",
"model",
"max_tokens",
"temperature",
"top_p",
"top_k",
"frequency_penalty",
"presence_penalty",
"repeat_penalty",
"seed",
]
def test_llava_loader_widget_order_is_frozen():
schema = llavaloader.LLavaLoader.INPUT_TYPES()
assert list(schema["required"]) == [
"ckpt_name",
"max_ctx",
"gpu_layers",
"n_threads",
"clip",
]
def test_optional_memory_free_simple_widget_order_is_frozen():
schema = llavaloader.LLavaOptionalMemoryFreeSimple.INPUT_TYPES()
assert list(schema["required"]) == [
"ckpt_name",
"clip_name",
"max_ctx",
"gpu_layers",
"n_threads",
"image",
"prompt",
"temperature",
"unload",
]
assert list(schema["optional"])[0] == "handler"
def test_every_llava_node_declares_a_callable_function_and_return_types():
for name, node_class in llavaloader.NODE_CLASS_MAPPINGS.items():
assert isinstance(node_class.RETURN_TYPES, tuple), name
assert node_class.RETURN_TYPES, name
assert callable(getattr(node_class, node_class.FUNCTION, None)), name
assert node_class.CATEGORY.startswith("VLM Nodes"), name
def test_display_names_cover_every_registered_node():
assert set(llavaloader.NODE_CLASS_MAPPINGS) == set(
llavaloader.NODE_DISPLAY_NAME_MAPPINGS
)
# --------------------------------------------------------------------------
# Vision message envelope.
# --------------------------------------------------------------------------
def test_vision_messages_place_the_image_before_the_text():
messages = llavaloader._vision_messages("sys", "what is this?", "data:image/png;b")
assert messages[0] == {"role": "system", "content": "sys"}
content = messages[1]["content"]
assert messages[1]["role"] == "user"
# llama.cpp vision handlers require the image part first.
assert content[0]["type"] == "image_url"
assert content[0]["image_url"]["url"] == "data:image/png;b"
assert content[1] == {"type": "text", "text": "what is this?"}
def test_run_batch_sends_a_png_data_uri_per_image():
llama = FakeLlama()
llavaloader._run_batch(
image_batch(1), llama, system_msg="sys", prompt="p", temperature=0.1
)
(call,) = llama.calls
url = call["messages"][1]["content"][0]["image_url"]["url"]
assert url.startswith("data:image/png;base64,")
# The payload must be real decodable PNG bytes.
decoded = base64.b64decode(url.split(",", 1)[1])
assert decoded.startswith(b"\x89PNG\r\n\x1a\n")
def test_run_batch_calls_the_model_once_per_batch_item():
llama = FakeLlama(["first", "second", "third"])
text = llavaloader._run_batch(
image_batch(3), llama, system_msg="sys", prompt="p", temperature=0.1
)
assert len(llama.calls) == 3
# Every batch item must survive into the response.
assert "first" in text
assert "second" in text
assert "third" in text
assert "--- Image 1 ---" in text
assert "--- Image 3 ---" in text
def test_run_batch_returns_bare_text_for_a_single_image():
llama = FakeLlama(["only one"])
text = llavaloader._run_batch(
image_batch(1), llama, system_msg="sys", prompt="p", temperature=0.1
)
assert text == "only one"
def test_run_batch_forwards_generation_kwargs_unchanged():
llama = FakeLlama()
llavaloader._run_batch(
image_batch(1),
llama,
system_msg="sys",
prompt="p",
max_tokens=32,
temperature=0.3,
top_p=0.7,
top_k=10,
seed=99,
)
(call,) = llama.calls
assert call["max_tokens"] == 32
assert call["temperature"] == 0.3
assert call["top_p"] == 0.7
assert call["top_k"] == 10
assert call["seed"] == 99
def test_sampler_simple_returns_a_single_string_output():
llama = FakeLlama(["a cat on a mat"])
result = llavaloader.LLavaSamplerSimple().generate_text(
image=image_batch(1), prompt="describe", model=llama, temperature=0.1
)
assert result == ("a cat on a mat",)
def test_sampler_advanced_uses_the_supplied_system_message():
llama = FakeLlama()
llavaloader.LLavaSamplerAdvanced().generate_text_advanced(
image=image_batch(1),
system_msg="answer in French",
prompt="describe",
model=llama,
max_tokens=16,
temperature=0.1,
top_p=0.9,
top_k=5,
frequency_penalty=0.0,
presence_penalty=0.0,
repeat_penalty=1.0,
seed=7,
)
(call,) = llama.calls
assert call["messages"][0] == {"role": "system", "content": "answer in French"}
# --------------------------------------------------------------------------
# Projector / clip wiring.
# --------------------------------------------------------------------------
def test_clip_factory_uses_the_config_create_hook():
config = LlavaClipConfig(Path("/models/mmproj.gguf"), "LLaVA 1.6")
assert llavaloader._clip_factory(config) == config.create
def test_clip_factory_accepts_a_precreated_handler():
sentinel = object()
factory = llavaloader._clip_factory(sentinel)
# Workflows saved before handler selection passed the handler itself.
assert factory() is sentinel
def test_make_handle_derives_the_projector_from_the_clip_config(resolved_paths):
config = LlavaClipConfig(resolved_paths / CLIP_FILE, "LLaVA 1.5")
handle = llavaloader._make_handle(MODEL_FILE, 4096, -1, 4, config)
assert isinstance(handle, LlamaHandle)
assert handle.projector_path == resolved_paths / CLIP_FILE
assert handle.chat_handler_factory == config.create
assert handle.n_ctx == 4096
# Still lazy: no llama.cpp object was constructed.
assert handle._llm is None
def test_make_handle_keeps_an_explicit_projector_override(resolved_paths):
config = LlavaClipConfig(resolved_paths / CLIP_FILE, "LLaVA 1.5")
override = Path("/models/other-mmproj.gguf")
handle = llavaloader._make_handle(
MODEL_FILE,
4096,
-1,
4,
config,
runtime_options={"projector_path": override},
)
assert handle.projector_path == override
def test_llava_loader_does_not_load_weights(resolved_paths):
config = LlavaClipConfig(resolved_paths / CLIP_FILE, "LLaVA 1.5")
(handle,) = llavaloader.LLavaLoader().load_llava_checkpoint(
ckpt_name=MODEL_FILE,
max_ctx=2048,
gpu_layers=10,
n_threads=8,
clip=config,
)
assert isinstance(handle, LlamaHandle)
assert handle._llm is None
assert handle.n_gpu_layers == 10
assert handle.n_threads == 8
def test_clip_loader_returns_a_frozen_config_with_the_chosen_handler(resolved_paths):
(config,) = llavaloader.LlavaClipLoader().load_clip_checkpoint(
CLIP_FILE, handler="MiniCPM-V 2.6"
)
assert isinstance(config, LlavaClipConfig)
assert config.model_path == resolved_paths / CLIP_FILE
assert config.handler == "MiniCPM-V 2.6"
def test_clip_config_rejects_an_unknown_handler(resolved_paths):
config = LlavaClipConfig(resolved_paths / CLIP_FILE, "Not A Handler")
with pytest.raises((ValueError, RuntimeError)) as error:
config.create()
# Either an unknown-handler rejection or a missing-wheel report is correct;
# a silent fallback to the wrong prompt format is not.
assert "handler" in str(error.value).lower() or "llama" in str(error.value).lower()
def test_clip_loader_defaults_to_the_embedded_gguf_chat_template():
handler = llavaloader.LlavaClipLoader.INPUT_TYPES()["optional"]["handler"]
choices, options = handler[0], handler[1]
assert options["default"] == "Auto (GGUF chat template)"
assert options["default"] in choices
assert "LLaVA 1.5" in choices
# --------------------------------------------------------------------------
# Cached-handle lifecycle (issue #137: "model never unloads").
# --------------------------------------------------------------------------
def test_cached_llava_reuses_one_handle_for_identical_settings(
fake_handles, resolved_paths
):
node = llavaloader.LLavaOptionalMemoryFreeSimple()
first = node._model(MODEL_FILE, CLIP_FILE, 4096, -1, 4)
second = node._model(MODEL_FILE, CLIP_FILE, 4096, -1, 4)
assert first is second
assert len(fake_handles.instances) == 1
def test_cached_llava_rebuilds_when_the_projector_changes(fake_handles, resolved_paths):
node = llavaloader.LLavaOptionalMemoryFreeSimple()
first = node._model(MODEL_FILE, CLIP_FILE, 4096, -1, 4)
second = node._model(MODEL_FILE, "other-mmproj.gguf", 4096, -1, 4)
assert first is not second
assert first.closed is True
assert len(fake_handles.instances) == 2
def test_cached_llava_rebuilds_when_the_handler_changes(fake_handles, resolved_paths):
node = llavaloader.LLavaOptionalMemoryFreeSimple()
node._model(MODEL_FILE, CLIP_FILE, 4096, -1, 4, handler="LLaVA 1.5")
node._model(MODEL_FILE, CLIP_FILE, 4096, -1, 4, handler="LLaVA 1.6")
assert len(fake_handles.instances) == 2
def test_cached_llava_unload_releases_the_handle(fake_handles, resolved_paths):
node = llavaloader.LLavaOptionalMemoryFreeSimple()
handle = node._model(MODEL_FILE, CLIP_FILE, 4096, -1, 4)
node._maybe_unload(False)
assert handle.closed is False
node._maybe_unload(True)
assert handle.closed is True
assert node._handle is None
assert node._key is None
def test_cached_llava_unload_is_safe_before_any_load(fake_handles, resolved_paths):
node = llavaloader.LLavaOptionalMemoryFreeSimple()
# Must not raise when nothing was ever loaded.
node._maybe_unload(True)
assert node._handle is None
def _memory_free_kwargs(**overrides):
kwargs = {
"ckpt_name": MODEL_FILE,
"clip_name": CLIP_FILE,
"max_ctx": 4096,
"gpu_layers": -1,
"n_threads": 4,
"image": image_batch(1),
"prompt": "describe this",
"temperature": 0.1,
"unload": False,
}
kwargs.update(overrides)
return kwargs
def test_memory_free_simple_generates_through_the_cached_handle(
fake_handles, resolved_paths
):
node = llavaloader.LLavaOptionalMemoryFreeSimple()
(text,) = node.generate_text(**_memory_free_kwargs())
assert text == "a description"
(handle,) = fake_handles.instances
assert handle.closed is False
(call,) = handle.llama.calls
assert call["temperature"] == 0.1
assert call["messages"][1]["content"][1]["text"] == "describe this"
def test_memory_free_simple_processes_every_image_in_the_batch(
fake_handles, resolved_paths
):
node = llavaloader.LLavaOptionalMemoryFreeSimple()
node.generate_text(**_memory_free_kwargs(image=image_batch(2)))
(handle,) = fake_handles.instances
assert len(handle.llama.calls) == 2
def test_memory_free_simple_unloads_after_generating_when_asked(
fake_handles, resolved_paths
):
node = llavaloader.LLavaOptionalMemoryFreeSimple()
node.generate_text(**_memory_free_kwargs(unload=True))
(handle,) = fake_handles.instances
assert handle.closed is True
assert node._handle is None
def test_memory_free_simple_unloads_even_when_generation_fails(
fake_handles, resolved_paths, monkeypatch
):
"""Issue #137: a failed generation must not strand the model in VRAM."""
def explode(*args, **kwargs):
raise RuntimeError("llama.cpp exploded")
monkeypatch.setattr(llavaloader, "_run_batch", explode)
node = llavaloader.LLavaOptionalMemoryFreeSimple()
with pytest.raises(RuntimeError, match="exploded"):
node.generate_text(**_memory_free_kwargs(unload=True))
(handle,) = fake_handles.instances
assert handle.closed is True
assert node._handle is None
def test_memory_free_advanced_forwards_the_system_message_and_sampling(
fake_handles, resolved_paths
):
node = llavaloader.LLavaOptionalMemoryFreeAdvanced()
(text,) = node.generate_text_advanced(
ckpt_name=MODEL_FILE,
clip_name=CLIP_FILE,
max_ctx=4096,
gpu_layers=-1,
n_threads=4,
image=image_batch(1),
system_msg="answer in German",
prompt="describe",
max_tokens=64,
temperature=0.4,
top_p=0.85,
top_k=25,
frequency_penalty=0.0,
presence_penalty=0.0,
repeat_penalty=1.05,
seed=5,
unload=False,
)
assert text == "a description"
(handle,) = fake_handles.instances
(call,) = handle.llama.calls
assert call["messages"][0] == {"role": "system", "content": "answer in German"}
assert call["max_tokens"] == 64
assert call["temperature"] == 0.4
assert call["top_p"] == 0.85
assert call["top_k"] == 25
assert call["seed"] == 5
def test_cached_llava_key_is_insensitive_to_runtime_option_ordering(
fake_handles, resolved_paths
):
node = llavaloader.LLavaOptionalMemoryFreeSimple()
node._model(MODEL_FILE, CLIP_FILE, 4096, -1, 4, n_batch=256, main_gpu=1)
node._model(MODEL_FILE, CLIP_FILE, 4096, -1, 4, main_gpu=1, n_batch=256)
assert len(fake_handles.instances) == 1
+291
View File
@@ -0,0 +1,291 @@
from types import SimpleNamespace
import numpy as np
import pytest
from ComfyUI_VLM_nodes.nodes import minimax_music
def generation_request(**overrides):
values = {
"region": "global_en",
"model": "music-3.0",
"prompt": "Reflective acoustic pop",
"lyrics": "[Verse]\nA quiet road under evening light",
"stream": False,
"output_format": "hex",
"audio_format": "mp3",
"sample_rate": 44100,
"bitrate": 256000,
"lyrics_optimizer": False,
"is_instrumental": False,
"aigc_watermark": False,
"audio_url": "",
"audio_base64": "",
"cover_feature_id": "",
}
values.update(overrides)
return minimax_music.build_music_request(**values)
def test_music_contract_matches_current_models_regions_and_formats():
assert minimax_music.REGION_ENDPOINTS == {
"global_en": "https://api.minimax.io/v1/music_generation",
"cn_zh": "https://api.minimaxi.com/v1/music_generation",
}
assert minimax_music.DEFAULT_MODEL == "music-3.0"
assert minimax_music.GENERATION_MODELS == (
"music-3.0",
"music-2.6",
"music-3.0-free",
"music-2.6-free",
)
assert minimax_music.COVER_MODELS == ("music-cover", "music-cover-free")
assert minimax_music.OUTPUT_FORMATS == ("url", "hex")
assert minimax_music.AUDIO_FORMATS == ("mp3", "wav", "pcm")
assert {
"model",
"prompt",
"lyrics",
"stream",
"output_format",
"audio_setting",
"lyrics_optimizer",
"is_instrumental",
"audio_url",
"audio_base64",
"cover_feature_id",
} == minimax_music.REQUEST_FIELDS
assert minimax_music.REGIONAL_FIELDS == {
"global_en": (),
"cn_zh": ("aigc_watermark",),
}
def test_generation_request_covers_generation_and_cn_fields():
request = generation_request(
region="cn_zh",
stream=True,
lyrics="",
lyrics_optimizer=True,
aigc_watermark=True,
audio_format="wav",
sample_rate=32000,
bitrate=128000,
)
assert request == {
"model": "music-3.0",
"prompt": "Reflective acoustic pop",
"stream": True,
"output_format": "hex",
"audio_setting": {
"sample_rate": 32000,
"bitrate": 128000,
"format": "wav",
},
"lyrics_optimizer": True,
"is_instrumental": False,
"aigc_watermark": True,
}
@pytest.mark.parametrize(
("source", "value"),
[
("audio_url", "https://media.example/reference.wav"),
("audio_base64", "dGVzdA=="),
("cover_feature_id", "feature-123"),
],
)
def test_cover_request_supports_each_documented_source(source, value):
overrides = {
"model": "music-cover",
"prompt": "Warm orchestral cover",
"lyrics": "Updated words for the cover",
source: value,
}
request = generation_request(**overrides)
assert request[source] == value
assert "lyrics_optimizer" not in request
assert "is_instrumental" not in request
def test_streaming_requires_hex_and_cover_sources_are_exclusive():
with pytest.raises(ValueError, match="output_format='hex'"):
generation_request(stream=True, output_format="url")
with pytest.raises(ValueError, match="exactly one"):
generation_request(
model="music-cover",
prompt="Warm orchestral cover",
audio_url="https://media.example/reference.wav",
audio_base64="dGVzdA==",
)
def test_stream_response_joins_hex_chunks_and_requires_completion():
class Response:
def iter_lines(self):
return iter(
[
'data: {"data":{"status":1,"audio":"0001"},'
'"base_resp":{"status_code":0}}',
'data: {"data":{"status":2,"audio":"0203"},'
'"extra_info":{"music_sample_rate":32000,"music_channel":2},'
'"base_resp":{"status_code":0}}',
"data: [DONE]",
]
)
audio, metadata = minimax_music._stream_audio(Response())
assert audio == "00010203"
assert metadata == {"music_sample_rate": 32000, "music_channel": 2}
def test_url_and_hex_response_decoding():
class DownloadResponse:
def __enter__(self):
return self
def __exit__(self, *_args):
return False
def raise_for_status(self):
return None
def iter_bytes(self):
return iter((b"ab", b"cd"))
class Client:
def stream(self, method, url):
assert method == "GET"
assert url == "https://media.example/music.wav"
return DownloadResponse()
client = Client()
assert minimax_music._audio_bytes(client, "61626364", "hex") == b"abcd"
assert (
minimax_music._audio_bytes(
client,
"https://media.example/music.wav",
"url",
)
== b"abcd"
)
def test_pcm_decoding_uses_response_sample_rate_and_channel_count():
content = np.array([0, 32767, -32768, 0], dtype="<i2").tobytes()
samples, sample_rate = minimax_music._decode_audio(
content,
"pcm",
44100,
{"music_sample_rate": 32000, "music_channel": 2},
)
assert samples.shape == (2, 2)
assert sample_rate == 32000
assert samples[0, 1] == pytest.approx(32767 / 32768)
def test_node_posts_to_fixed_region_and_returns_comfy_audio(monkeypatch):
captured = {}
class Response:
def raise_for_status(self):
return None
def json(self):
return {
"data": {"status": 2, "audio": "0102"},
"extra_info": {"music_sample_rate": 44100, "music_channel": 2},
"base_resp": {"status_code": 0},
}
class Client:
def __init__(self, **kwargs):
captured["client"] = kwargs
def post(self, endpoint, *, headers, json):
captured["endpoint"] = endpoint
captured["headers"] = headers
captured["request"] = json
return Response()
def close(self):
captured["closed"] = True
def fake_soundfile_read(buffer, **kwargs):
assert buffer.read() == b"\x01\x02"
assert kwargs == {"dtype": "float32", "always_2d": True}
return np.zeros((8, 2), dtype=np.float32), 44100
def fake_require_module(name, *_args):
if name == "httpx":
return SimpleNamespace(Client=Client)
if name == "soundfile":
return SimpleNamespace(read=fake_soundfile_read)
raise AssertionError(name)
monkeypatch.setenv(minimax_music.API_KEY_ENV, "test-key-not-for-production")
monkeypatch.setattr(minimax_music, "require_module", fake_require_module)
result = minimax_music.MiniMaxMusicNode().generate_music(
region="global_en",
model="music-3.0",
prompt="Reflective acoustic pop",
lyrics="[Verse]\nA quiet road under evening light",
stream=False,
output_format="hex",
audio_format="wav",
sample_rate=44100,
bitrate=256000,
lyrics_optimizer=False,
is_instrumental=False,
aigc_watermark=False,
)
assert captured["endpoint"] == minimax_music.REGION_ENDPOINTS["global_en"]
assert captured["headers"]["Authorization"].startswith("Bearer ")
assert captured["client"] == {
"timeout": 600.0,
"follow_redirects": False,
"trust_env": False,
}
assert captured["closed"] is True
assert len(result) == 3
assert result[1] == 44100
assert result[2]["waveform"].shape == (1, 2, 8)
def test_request_failures_redact_the_resolved_key(monkeypatch):
resolved_value = "unit-key"
class Client:
def __init__(self, **_kwargs):
pass
def post(self, *_args, **_kwargs):
raise RuntimeError(f"Authorization: Bearer {resolved_value}")
def close(self):
pass
monkeypatch.setenv(minimax_music.API_KEY_ENV, resolved_value)
monkeypatch.setattr(
minimax_music,
"require_module",
lambda *_args: SimpleNamespace(Client=Client),
)
with pytest.raises(RuntimeError) as captured:
minimax_music.MiniMaxMusicNode().generate_music(
region="global_en",
model="music-3.0",
prompt="Reflective acoustic pop",
lyrics="[Verse]\nA quiet road under evening light",
stream=False,
output_format="hex",
audio_format="mp3",
sample_rate=44100,
bitrate=256000,
lyrics_optimizer=False,
is_instrumental=False,
aigc_watermark=False,
)
assert resolved_value not in str(captured.value)
assert "[REDACTED]" in str(captured.value)
+77
View File
@@ -0,0 +1,77 @@
from __future__ import annotations
import importlib
import sys
from pathlib import Path
from types import ModuleType, SimpleNamespace
import torch
PACKAGE = __package__.split(".")[0] if __package__ else "ComfyUI_VLM_nodes"
module = importlib.import_module(f"{PACKAGE}.nodes.moondream2")
def test_native_checkpoint_loader_bypasses_transformers_from_pretrained(
tmp_path: Path,
monkeypatch,
):
package = ModuleType(module._CHECKPOINT_PACKAGE)
package.__path__ = [str(tmp_path.resolve())]
package.__package__ = module._CHECKPOINT_PACKAGE
checkpoint = ModuleType(f"{module._CHECKPOINT_PACKAGE}.hf_moondream")
calls = {}
class FakeConfig:
@classmethod
def from_pretrained(cls, model_path, **kwargs):
calls["config"] = (Path(model_path), kwargs)
return cls()
class FakeModel(torch.nn.Module):
def __init__(self, config):
super().__init__()
self.weight = torch.nn.Parameter(torch.zeros(1))
calls["model_config"] = config
checkpoint.HfConfig = FakeConfig
checkpoint.HfMoondream = FakeModel
monkeypatch.setitem(sys.modules, module._CHECKPOINT_PACKAGE, package)
monkeypatch.setitem(
sys.modules,
f"{module._CHECKPOINT_PACKAGE}.hf_moondream",
checkpoint,
)
weights = tmp_path / "model.safetensors"
weights.write_bytes(b"test")
def load_model(model, filename, *, strict):
calls["weights"] = (model, Path(filename), strict)
model.weight.data.fill_(1)
return set(), []
monkeypatch.setattr(
module,
"require_module",
lambda name: (
SimpleNamespace(load_model=load_model)
if name == "safetensors.torch"
else None
),
)
model = module._load_native_checkpoint(tmp_path)
assert isinstance(model, FakeModel)
assert not model.training
assert model.weight.item() == 1
assert calls["config"] == (tmp_path, {"local_files_only": True})
assert calls["weights"] == (model, weights, True)
def test_photon_requirements_pin_cuda_runtime_with_required_symbol():
requirements = (
Path(module.__file__).resolve().parents[1] / "requirements-moondream31.txt"
).read_text(encoding="utf-8")
assert "kestrel-kernels==0.4.6" in requirements
assert "nvidia-cuda-runtime-cu12==12.9.79" in requirements
+361
View File
@@ -0,0 +1,361 @@
import asyncio
import importlib
import inspect
import json
import sys
import types
from dataclasses import dataclass
import pytest
import torch
PACKAGE = __package__.split(".")[0] if __package__ else "ComfyUI_VLM_nodes"
module = importlib.import_module(f"{PACKAGE}.nodes.moondream31")
worker = importlib.import_module(f"{PACKAGE}.nodes.moondream31_worker")
Moondream31Detect = module.Moondream31Detect
Moondream31Loader = module.Moondream31Loader
Moondream31Model = module.Moondream31Model
Moondream31Segment = module.Moondream31Segment
svg_path_to_mask = module.svg_path_to_mask
def _fake_model(handler, model_name=module.MODEL_ID):
model = object.__new__(Moondream31Model)
model.config = module.Moondream31Config(
model=model_name,
device="cuda",
max_batch_size=4,
kv_cache_pages=8192,
)
model.request = handler
model.close = lambda: None
return model
def test_svg_path_is_transformed_from_bbox_space_to_image_pixels():
mask, polygon, contours = svg_path_to_mask(
"M 0 0 H 1 V 1 H 0 Z",
{"x_min": 0.25, "y_min": 0.25, "x_max": 0.75, "y_max": 0.75},
100,
80,
supersample=4,
)
assert mask.shape == (80, 100)
assert mask[40, 50] > 0.99
assert mask[5, 5] == 0
assert mask.sum().item() == pytest.approx(2000, rel=0.06)
assert len(polygon) >= 4
assert len(contours) == 1
xs = [point[0] for point in polygon]
ys = [point[1] for point in polygon]
assert min(xs) == pytest.approx(25)
assert max(xs) == pytest.approx(75)
assert min(ys) == pytest.approx(20)
assert max(ys) == pytest.approx(60)
def test_svg_curves_and_evenodd_holes_are_preserved():
path = "M 0 0 H 1 V 1 H 0 Z M .25 .25 C .4 .1 .6 .1 .75 .25 V .75 H .25 Z"
mask, polygon, contours = svg_path_to_mask(
path,
{"x_min": 0, "y_min": 0, "x_max": 1, "y_max": 1},
128,
128,
supersample=4,
precision_px=0.5,
)
assert len(contours) == 2
assert len(polygon) >= 4
assert mask[8, 8] > 0.99
assert mask[64, 64] < 0.01
@pytest.mark.parametrize(
("path", "bbox", "message"),
[
("", {"x_min": 0, "y_min": 0, "x_max": 1, "y_max": 1}, "empty"),
(
"M 0 0 L nan 1 Z",
{"x_min": 0, "y_min": 0, "x_max": 1, "y_max": 1},
"invalid",
),
(
"M 0 0 H 1 V 1 Z",
{"x_min": 0.7, "y_min": 0, "x_max": 0.2, "y_max": 1},
"positive",
),
],
)
def test_svg_rejects_malformed_or_unsafe_geometry(path, bbox, message):
with pytest.raises((TypeError, ValueError), match=message):
svg_path_to_mask(path, bbox, 64, 64)
def test_video_detect_uses_stride_parallelism_and_reports_measured_fps():
observed = {}
def request(operation, **payload):
observed["operation"] = operation
observed.update(payload)
return {
"items": [
{
"objects": [
{
"x_min": 0.1,
"y_min": 0.2,
"x_max": 0.4,
"y_max": 0.6,
}
]
},
{"objects": []},
],
"elapsed_seconds": 0.1,
"parallel_requests": 2,
}
images = torch.zeros((4, 48, 64, 3), dtype=torch.float32)
outputs = Moondream31Detect().detect(
_fake_model(request),
images,
"person",
30.0,
2,
2,
20,
False,
)
sequence = outputs[0]
performance = json.loads(outputs[-1])
assert observed["operation"] == "detect"
assert len(observed["images"]) == 2
assert observed["parallel_requests"] == 2
assert sequence.frame_count == 4
assert [frame.frame_index for frame in sequence.frames] == [0, 2]
assert sequence.frames[0].detections[0].bbox_xyxy == pytest.approx(
(6.4, 9.6, 25.6, 28.8)
)
assert outputs[2].shape == images.shape
assert outputs[3].shape == (4, 48, 64)
assert performance["processed_frames"] == 2
assert performance["worker_fps"] == pytest.approx(20)
assert performance["target_processed_fps"] == pytest.approx(15)
assert performance["parallel_requests"] == 2
def test_segment_exposes_svg_mask_cutout_overlay_and_structured_detection():
def request(operation, **payload):
assert operation == "segment"
assert payload["spatial_refs"] == [[0.5, 0.5]]
return {
"items": [
{
"path": "M 0 0 H 1 V 1 H 0 Z",
"bbox": {
"x_min": 0.25,
"y_min": 0.25,
"x_max": 0.75,
"y_max": 0.75,
},
}
],
"elapsed_seconds": 0.2,
"parallel_requests": 1,
}
image = torch.ones((1, 32, 40, 3), dtype=torch.float32)
outputs = Moondream31Segment().segment(
_fake_model(request, module.PREVIEW_MODEL_ID),
image,
"object",
1.0,
1,
1,
4,
False,
spatial_refs_json="[[0.5, 0.5]]",
)
sequence = outputs[0]
native = json.loads(outputs[2])
mask = outputs[3]
mask_image = outputs[4]
cutout = outputs[5]
overlay = outputs[6]
detection = sequence.frames[0].detections[0]
assert native[0]["path"].startswith("M 0 0")
assert mask.shape == (1, 32, 40)
assert mask_image.shape == (1, 32, 40, 3)
assert cutout.shape == image.shape
assert overlay.shape == image.shape
assert mask[0, 16, 20] > 0.99
assert mask[0, 2, 2] == 0
assert cutout[0, 16, 20].min() > 0.99
assert cutout[0, 2, 2].max() == 0
assert detection.mask is not None
assert detection.polygon is not None
assert detection.metadata["native_svg_path"].startswith("M 0 0")
def test_license_gate_and_node_registration():
with pytest.raises(ValueError, match="License"):
Moondream31Loader().load(
False,
"Auto",
4,
"Balanced (8K pages)",
)
assert set(module.NODE_CLASS_MAPPINGS) == {
"Moondream31Loader",
"Moondream31Query",
"Moondream31Caption",
"Moondream31Detect",
"Moondream31Point",
"Moondream31Segment",
}
assert all(
node.CATEGORY == "VLM Nodes/Moondream 3"
for node in module.NODE_CLASS_MAPPINGS.values()
)
def test_final_31_model_does_not_claim_preview_svg_segment():
with pytest.raises(ValueError, match="3 Preview"):
Moondream31Segment().segment(
_fake_model(lambda *_args, **_kwargs: {}),
torch.zeros((1, 16, 16, 3)),
"object",
1.0,
1,
1,
1,
False,
)
def test_worker_auth_is_not_exposed_in_process_arguments_and_logs_are_redacted(
tmp_path,
monkeypatch,
):
source = inspect.getsource(Moondream31Model.ensure_started)
assert '"--auth-key"' not in source
assert "MOONDREAM_WORKER_AUTH" in inspect.getsource(
module._worker_environment
)
log = tmp_path / "worker.log"
log.write_text(
"api_key=secret-value\nAuthorization: bearer-value\nCUDA error",
encoding="utf-8",
)
tail = module._safe_log_tail(log)
assert "secret-value" not in tail
assert "bearer-value" not in tail
assert "CUDA error" in tail
monkeypatch.setenv("PATH", "/runtime/bin")
monkeypatch.setenv("OPENAI_API_KEY", "must-not-cross")
monkeypatch.setenv("HF_TOKEN", "hf-server-side")
monkeypatch.setenv("MOONDREAM_API_KEY", "adapter-only")
monkeypatch.setenv("HTTPS_PROXY", "https://user:password@example.test")
monkeypatch.setenv("PYTORCH_ALLOC_CONF", "backend:cudaMallocAsync")
monkeypatch.setenv("PYTORCH_CUDA_ALLOC_CONF", "backend:cudaMallocAsync")
base_environment = module._worker_environment(
tmp_path,
b"\x01" * 32,
module.MODEL_ID,
)
assert base_environment["PATH"] == "/runtime/bin"
assert base_environment["HF_TOKEN"] == "hf-server-side"
assert "OPENAI_API_KEY" not in base_environment
assert "MOONDREAM_API_KEY" not in base_environment
assert "HTTPS_PROXY" not in base_environment
assert "PYTORCH_ALLOC_CONF" not in base_environment
assert "PYTORCH_CUDA_ALLOC_CONF" not in base_environment
assert base_environment["MOONDREAM_WORKER_AUTH"] == "01" * 32
adapter_environment = module._worker_environment(
tmp_path,
b"\x02" * 32,
f"{module.MODEL_ID}/adapter@step",
)
assert adapter_environment["MOONDREAM_API_KEY"] == "adapter-only"
def test_runtime_python_preserves_virtualenv_symlink(tmp_path, monkeypatch):
root = tmp_path / "runtime"
binary = tmp_path / "base-python"
binary.write_text("", encoding="utf-8")
venv_python = root / ".venv" / "bin" / "python"
venv_python.parent.mkdir(parents=True)
try:
venv_python.symlink_to(binary)
except OSError:
pytest.skip("This filesystem cannot create symlinks.")
monkeypatch.delenv("MOONDREAM_PYTHON", raising=False)
selected = module._runtime_python(root)
assert selected == venv_python.absolute()
assert selected != binary.resolve()
def test_worker_registers_official_31_id_only_when_upstream_is_missing(
monkeypatch,
):
@dataclass(frozen=True)
class Spec:
name: str
repo_id: str
filename: str
checkpoint_format: str
registry = {
"moondream3-preview": Spec(
"moondream3-preview",
"moondream/moondream3-preview",
"model_fp8.pt",
"md3",
)
}
fake = types.ModuleType("kestrel.models")
fake.get_spec = lambda name: (
registry[name] if name in registry else (_ for _ in ()).throw(ValueError(name))
)
fake.register = lambda spec: registry.__setitem__(spec.name, spec)
monkeypatch.setitem(sys.modules, "kestrel.models", fake)
assert worker._register_moondream31_if_needed("moondream3.1-9B-A2B")
registered = registry["moondream3.1-9B-A2B"]
assert registered.repo_id == "moondream/moondream3.1-9B-A2B"
assert registered.filename == "model.safetensors"
assert registered.checkpoint_format == "md3"
assert not worker._register_moondream31_if_needed("moondream3.1-9B-A2B")
assert not worker._register_moondream31_if_needed("custom-model")
def test_worker_honors_do_not_track_for_base_models(monkeypatch):
class SimpleClient:
def __init__(self):
self.closed = False
async def aclose(self):
self.closed = True
class Reporter:
def __init__(self):
self._client = SimpleClient()
fake = types.ModuleType("kestrel.photon")
fake.PhotonReporter = Reporter
monkeypatch.setitem(sys.modules, "kestrel.photon", fake)
monkeypatch.setenv("DO_NOT_TRACK", "1")
monkeypatch.delenv("MOONDREAM_API_KEY", raising=False)
assert worker._honor_do_not_track()
reporter = Reporter()
assert asyncio.run(reporter.validate_api_key()) is False
assert reporter.start() is None
asyncio.run(reporter.shutdown())
assert reporter._client.closed
monkeypatch.setenv("MOONDREAM_API_KEY", "finetune-key")
assert not worker._honor_do_not_track()
+41
View File
@@ -40,6 +40,7 @@ def test_every_module_imports_and_expected_nodes_exist():
assert package.IMPORT_ERRORS == {}
expected = {
"ModernVLM",
"LegacyModernVLM",
"VLMRuntimeDiagnostics",
"Florence2",
"Paligemma",
@@ -119,6 +120,10 @@ def test_dependency_metadata_matches_installer_requirements():
if line.strip() and not line.lstrip().startswith("#")
}
assert project_requirements == installer_requirements
assert any(
Requirement(value).name == "num2words"
for value in metadata["project"]["dependencies"]
)
bitsandbytes = next(
Requirement(value)
@@ -410,6 +415,37 @@ def test_modern_catalog_has_current_quality_and_low_vram_tiers():
assert "ibm-granite/granite-vision-4.1-4b" in repositories
def test_modern_picker_is_curated_and_legacy_models_remain_compatible():
visible = tuple(modern_vlm.ModernVLM.INPUT_TYPES()["required"]["model"][0])
legacy = tuple(
modern_vlm.LegacyModernVLM.INPUT_TYPES()["required"]["model"][0]
)
assert visible == modern_vlm.RECOMMENDED_MODEL_LABELS
assert legacy == modern_vlm.LEGACY_MODEL_LABELS
assert len(visible) == 12
assert set(visible).isdisjoint(legacy)
assert set(visible) | set(legacy) == set(modern_vlm.MODEL_CATALOG)
assert (
modern_vlm.ModernVLM.VALIDATE_INPUTS(
"Qwen 2.5 VL 3B Instruct (legacy workflows)"
)
is True
)
for node_name in (
"Kosmos2model",
"MCLLaVAModel",
"MiniCPMNode",
"MolmoNode",
"MoonDream",
"Paligemma",
"Qwen2VLNode",
"UformGen2QwenNode",
):
assert package.NODE_CLASS_MAPPINGS[node_name].CATEGORY.startswith(
"VLM Nodes/Legacy/"
)
def test_modern_video_is_primary_input_and_thinking_is_explicit():
assert "image" in modern_vlm.ModernVLM.INPUT_TYPES()["optional"]
assert "image" in qwen2vl.Qwen2VLNode.INPUT_TYPES()["optional"]
@@ -513,6 +549,11 @@ def test_view_text_frontend_rehydrates_and_uses_native_progress_channel():
assert 'api.addEventListener("progress_text"' in source
assert "onNodeOutputsUpdated(nodeOutputs)" in source
assert "connectedViewTextNodes(source)" in source
assert '"VLMVideoTemporalReasoner"' in source
assert 'makeButton("Save"' in source
assert 'makeButton("Wrap: on"' in source
assert 'makeButton("Follow: on"' in source
assert "isReroute(target)" in source
def test_internvl_video_uses_an_even_vision_patch_grid():
+753
View File
@@ -0,0 +1,753 @@
import json
import re
import threading
from http import HTTPStatus
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from types import SimpleNamespace
import numpy as np
import pytest
import torch
from ComfyUI_VLM_nodes.examples.robotics import lerobot_policy_server
from ComfyUI_VLM_nodes.nodes import robotics
from ComfyUI_VLM_nodes.nodes.robotics import (
EMBODIMENT_SCHEMA,
VLA_MODEL_CATALOG,
RobotActions,
actions_from_json,
actions_from_response,
blend_action_chunks,
build_observation,
make_embodiment,
observation_to_groot_payload,
observation_to_http_payload,
observation_to_openpi_payload,
render_action_preview,
validate_action_trajectory,
validate_policy_url,
validate_ws_url,
validate_zmq_host,
)
def _image(frames=1, height=24, width=32):
return torch.linspace(0, 1, frames * height * width * 3).reshape(
frames, height, width, 3
)
def _joint_profile():
return make_embodiment("Generic 7-DoF joint + gripper")
def _observation(profile=None):
profile = profile or _joint_profile()
return build_observation(
task="Pick up the blue cube.",
state_json=json.dumps([0.0] * profile.state_dim),
primary_image=_image(frames=3),
primary_camera=profile.camera_names[0],
history_fps=15,
timestamp=12.5,
embodiment=profile,
)
def test_embodiment_presets_are_explicit_roundtrippable_contracts():
for preset in robotics.EMBODIMENT_PRESETS:
profile = make_embodiment(preset)
encoded = profile.to_dict()
assert encoded["schema"] == EMBODIMENT_SCHEMA
assert len(encoded["action_names"]) == len(encoded["action_min"])
assert len(encoded["action_names"]) == len(encoded["action_max"])
assert len(encoded["action_names"]) == len(encoded["max_delta_per_step"])
assert robotics.RobotEmbodiment.from_dict(encoded) == profile
assert "Template limits" in profile.notes
def test_embodiment_rejects_invalid_bounds_and_mismatched_overrides():
with pytest.raises(ValueError, match="smaller"):
robotics.RobotEmbodiment(
name="bad",
state_names=("s",),
action_names=("a",),
action_min=(1,),
action_max=(0,),
max_delta=(0.1,),
control_hz=10,
action_mode="absolute",
camera_names=("front",),
)
with pytest.raises(ValueError, match="must contain 8"):
make_embodiment(
"Generic 7-DoF joint + gripper",
action_min=[-1],
)
def test_observation_preserves_history_and_validates_state_and_cameras():
profile = _joint_profile()
observation = _observation(profile)
summary = observation.summary()
assert summary["history_frames"] == 3
assert summary["cameras"][profile.camera_names[0]]["width"] == 32
assert summary["state"] == [0.0] * 8
assert summary["task"] == "Pick up the blue cube."
with pytest.raises(ValueError, match="State dimension"):
build_observation(
task="move",
state_json="[0]",
primary_image=_image(),
primary_camera=profile.camera_names[0],
history_fps=10,
timestamp=0,
embodiment=profile,
)
with pytest.raises(ValueError, match="not declared"):
build_observation(
task="move",
state_json=json.dumps([0] * profile.state_dim),
primary_image=_image(),
primary_camera="unknown",
history_fps=10,
timestamp=0,
embodiment=profile,
)
def test_http_payload_is_bounded_and_contains_no_tensor_details():
payload = observation_to_http_payload(_observation(), include_history=True)
assert payload["schema"] == "comfyui-vlm/robot-observation"
camera_frames = next(iter(payload["cameras"].values()))
assert len(camera_frames) == 3
assert all(item["encoding"] == "base64-jpeg" for item in camera_frames)
assert "device" not in json.dumps(payload)
def test_openpi_payload_supports_flat_and_aloha_shapes():
observation = _observation()
flat = observation_to_openpi_payload(
observation,
layout="Flat keys (DROID / LIBERO)",
state_key="observation/state",
prompt_key="prompt",
)
camera_key = observation.images[0][0]
assert flat[camera_key].shape == (24, 32, 3)
assert flat[camera_key].dtype == np.uint8
assert flat["observation/state"].dtype == np.float32
nested = observation_to_openpi_payload(
observation,
layout="Nested images (ALOHA)",
state_key="state",
prompt_key="prompt",
)
assert nested["images"][camera_key].shape == (3, 24, 32)
def test_openpi_array_codec_roundtrips_and_rejects_object_arrays():
source = np.arange(12, dtype=np.float32).reshape(3, 4)
encoded = robotics._openpi_pack_array(source)
decoded = robotics._openpi_unpack_array(encoded)
assert np.array_equal(decoded, source)
with pytest.raises(ValueError, match="does not support dtype"):
robotics._openpi_pack_array(np.array([object()], dtype=object))
forged = {
b"__ndarray__": True,
b"data": b"",
b"dtype": "|O",
b"shape": (0,),
}
with pytest.raises(ValueError, match="unsafe"):
robotics._openpi_unpack_array(forged)
def test_openpi_client_protocol_and_token_redaction(monkeypatch):
sent = []
class Connection:
responses = [b"metadata", b"actions"]
def __enter__(self):
return self
def __exit__(self, *_args):
pass
def recv(self, timeout):
assert timeout == 5
return self.responses.pop(0)
def send(self, value):
sent.append(value)
connect_calls = []
def connect(uri, **kwargs):
connect_calls.append((uri, kwargs))
return Connection()
fake_ws = SimpleNamespace(connect=connect)
class FakeMsgpack:
@staticmethod
def packb(value, default):
assert value["prompt"] == "Pick up the blue cube."
assert callable(default)
return b"encoded-request"
@staticmethod
def unpackb(value, object_hook):
assert callable(object_hook)
if value == b"metadata":
return {"model": "pi-test"}
return {
"actions": np.zeros((2, 8), dtype=np.float32),
"server_timing": {"infer_ms": 12.5},
}
def fake_require(name, *_args):
return fake_ws if name == "websockets.sync.client" else FakeMsgpack
monkeypatch.setattr(robotics, "require_module", fake_require)
monkeypatch.setenv("OPENPI_API_KEY", "openpi-secret")
actions, report = robotics.call_openpi_policy(
_observation(),
endpoint="ws://127.0.0.1:8000",
timeout_seconds=5,
allow_remote=False,
layout="Flat keys (DROID / LIBERO)",
state_key="observation/state",
prompt_key="prompt",
)
assert actions.values.shape == (2, 8)
assert sent == [b"encoded-request"]
assert connect_calls[0][1]["additional_headers"] == {
"Authorization": "Api-Key openpi-secret"
}
assert report["server_metadata"] == {"model": "pi-test"}
assert "openpi-secret" not in json.dumps(report)
def broken_connect(*_args, **_kwargs):
raise RuntimeError("token=openpi-secret")
fake_ws.connect = broken_connect
with pytest.raises(RuntimeError, match=r"token=\[REDACTED\]") as exc:
robotics.call_openpi_policy(
_observation(),
endpoint="ws://127.0.0.1:8000",
timeout_seconds=5,
allow_remote=False,
layout="Flat keys (DROID / LIBERO)",
state_key="observation/state",
prompt_key="prompt",
)
assert "openpi-secret" not in str(exc.value)
def test_groot_payload_uses_official_nested_batch_time_contract():
observation = _observation()
payload = observation_to_groot_payload(observation)
camera = next(iter(payload["video"].values()))
assert camera.shape == (1, 3, 24, 32, 3)
assert camera.dtype == np.uint8
assert payload["state"]["state"].shape == (1, 3, 8)
assert payload["state"]["state"].dtype == np.float32
assert payload["language"]["task"] == [["Pick up the blue cube."]]
def test_groot_array_codec_and_native_client(monkeypatch):
array = np.arange(6, dtype=np.float32).reshape(2, 3)
encoded = robotics._groot_encode(array)
decoded = robotics._groot_decode(encoded)
assert np.array_equal(decoded, array)
with pytest.raises(TypeError, match="object/void"):
robotics._groot_encode(np.array([object()], dtype=object))
with pytest.raises(ValueError, match="object/void"):
robotics._groot_decode({"nd": True, "kind": "O", "type": "|O"})
sockets = []
class Socket:
def __init__(self):
self.options = []
self.connected = None
self.request = None
self.closed = False
def setsockopt(self, key, value):
self.options.append((key, value))
def connect(self, value):
self.connected = value
def send(self, value):
self.request = value
def recv(self):
return b"response"
def close(self, linger):
assert linger == 0
self.closed = True
class Context:
terminated = False
def socket(self, kind):
assert kind == 1
socket = Socket()
sockets.append(socket)
return socket
def term(self):
self.terminated = True
fake_zmq = SimpleNamespace(
REQ=1,
RCVTIMEO=2,
SNDTIMEO=3,
LINGER=4,
Context=Context,
)
packed_requests = []
class FakeMsgpack:
@staticmethod
def packb(value, default):
packed_requests.append(value)
assert callable(default)
return b"request"
@staticmethod
def unpackb(value, object_hook, raw):
assert value == b"response"
assert callable(object_hook)
assert raw is False
return [
{
"arm": np.zeros((1, 3, 7), dtype=np.float32),
"gripper": np.ones((1, 3, 1), dtype=np.float32),
},
{"server": "ok"},
]
def fake_require(name, *_args):
return fake_zmq if name == "zmq" else FakeMsgpack
monkeypatch.setattr(robotics, "require_module", fake_require)
monkeypatch.setenv("GROOT_API_TOKEN", "groot-secret")
actions, report = robotics.call_groot_policy(
_observation(),
host="127.0.0.1",
port=5555,
timeout_seconds=3,
allow_remote=False,
)
assert actions.values.shape == (3, 8)
assert actions.stream_slices == (("arm", 0, 7), ("gripper", 7, 8))
assert packed_requests[0]["api_token"] == "groot-secret"
assert sockets[0].connected == "tcp://127.0.0.1:5555"
assert sockets[0].closed is True
assert report["policy_info"] == {"server": "ok"}
assert "groot-secret" not in json.dumps(report)
def test_action_response_parses_arrays_and_named_streams():
single = actions_from_response(
{"actions": [[[1, 2], [3, 4]]]},
source="test",
)
assert single.values.shape == (2, 2)
assert single.stream_slices == (("actions", 0, 2),)
streams = actions_from_response(
{
"arm": np.zeros((1, 4, 7), dtype=np.float32),
"gripper": np.ones((1, 4, 1), dtype=np.float32),
"info": {"ignored": True},
},
source="groot",
)
assert streams.values.shape == (4, 8)
assert streams.stream_slices == (("arm", 0, 7), ("gripper", 7, 8))
def test_actions_json_roundtrip_and_chunk_replanning():
profile = _joint_profile()
previous = RobotActions(
torch.zeros((5, 8)),
profile.action_names,
"previous",
)
new = RobotActions(
torch.ones((6, 8)),
profile.action_names,
"new",
)
parsed = actions_from_json(json.dumps(new.to_dict()))
assert torch.equal(parsed.values, new.values)
assert parsed.action_names == new.action_names
replanned, report = blend_action_chunks(
previous,
new,
executed_steps=2,
transition_steps=2,
max_horizon=4,
)
assert replanned.values.shape == (4, 8)
assert torch.allclose(replanned.values[0], torch.full((8,), 1 / 3))
assert torch.allclose(replanned.values[1], torch.full((8,), 2 / 3))
assert torch.equal(replanned.values[2:], torch.ones((2, 8)))
assert report["transition_steps_applied"] == 2
with pytest.raises(ValueError, match="outside"):
blend_action_chunks(
previous,
new,
executed_steps=99,
transition_steps=2,
max_horizon=4,
)
def test_action_safety_clamps_bounds_deltas_and_horizon():
profile = _joint_profile()
values = torch.tensor(
[
[4.0, 0, 0, 0, 0, 0, 0, 2.0],
[-4.0, 0, 0, 0, 0, 0, 0, -2.0],
[0.0, 0, 0, 0, 0, 0, 0, 0.5],
]
)
actions = RobotActions(
values=values,
action_names=profile.action_names,
source="unit",
)
safe, report = validate_action_trajectory(
actions,
profile,
mode="Clamp safely",
execution_horizon=2,
previous_action_json=json.dumps([0.0] * 8),
)
assert safe.horizon == 2
assert torch.all(safe.values <= torch.tensor(profile.action_max))
assert torch.all(safe.values >= torch.tensor(profile.action_min))
limits = torch.tensor(profile.max_delta)
previous = torch.zeros(8)
for step in safe.values:
assert torch.all((step - previous).abs() <= limits + 1.0e-6)
previous = step
assert report["violations"]["total"] > 0
assert report["changed"] is True
assert report["safe_for_handoff"] is True
def test_action_safety_blocks_or_holds_nonfinite_actions():
profile = _joint_profile()
values = torch.zeros((2, 8))
values[0, 2] = float("nan")
actions = RobotActions(values, profile.action_names, "unit")
with pytest.raises(ValueError, match="blocked"):
validate_action_trajectory(
actions,
profile,
mode="Block unsafe",
execution_horizon=2,
)
held, report = validate_action_trajectory(
actions,
profile,
mode="Hold position on unsafe",
execution_horizon=2,
previous_action_json=json.dumps([0.25] * 8),
)
assert torch.allclose(held.values, torch.full((2, 8), 0.25))
assert report["safe_for_handoff"] is True
with pytest.raises(ValueError, match="requires previous_action_json"):
validate_action_trajectory(
actions,
profile,
mode="Hold position on unsafe",
execution_horizon=2,
)
def test_action_json_rejects_nonfinite_before_serialization():
actions = RobotActions(
torch.tensor([[float("inf")]]),
("action",),
"unit",
)
with pytest.raises(ValueError, match="NaN or infinity"):
actions.to_dict()
def test_policy_endpoint_security_defaults():
assert (
validate_policy_url("http://127.0.0.1:8787", allow_remote=False)
== "http://127.0.0.1:8787/v1/infer"
)
assert validate_policy_url(
"https://policy.example/v1/infer",
allow_remote=True,
) == "https://policy.example/v1/infer"
with pytest.raises(ValueError, match="HTTPS"):
validate_policy_url("http://policy.example", allow_remote=True)
with pytest.raises(ValueError, match="disabled"):
validate_policy_url("https://policy.example", allow_remote=False)
with pytest.raises(ValueError, match="embedded credentials"):
validate_policy_url("https://secret@example.com", allow_remote=True)
assert validate_ws_url("127.0.0.1:8000", allow_remote=False) == "ws://127.0.0.1:8000"
with pytest.raises(ValueError, match="WSS"):
validate_ws_url("ws://policy.example", allow_remote=True)
assert validate_zmq_host("localhost", allow_remote=False) == "localhost"
with pytest.raises(ValueError, match="disabled"):
validate_zmq_host("policy.example", allow_remote=False)
class _PolicyHandler(BaseHTTPRequestHandler):
observed_auth = None
observed_payload = None
def log_message(self, *_args):
pass
def do_POST(self): # noqa: N802
type(self).observed_auth = self.headers.get("Authorization")
length = int(self.headers["Content-Length"])
type(self).observed_payload = json.loads(self.rfile.read(length))
body = json.dumps(
{
"actions": [[0.0] * 8, [0.1] * 8],
"action_names": [f"a{index}" for index in range(8)],
}
).encode()
self.send_response(HTTPStatus.OK)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(body)))
self.end_headers()
self.wfile.write(body)
def test_http_policy_real_loopback_request_uses_env_token_without_leaking(monkeypatch):
server = ThreadingHTTPServer(("127.0.0.1", 0), _PolicyHandler)
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
monkeypatch.setenv("VLA_POLICY_TOKEN", "top-secret-value")
try:
actions, report = robotics.call_http_policy(
_observation(),
endpoint=f"http://127.0.0.1:{server.server_port}",
timeout_seconds=5,
allow_remote=False,
include_history=False,
)
finally:
server.shutdown()
server.server_close()
thread.join(timeout=5)
assert actions.values.shape == (2, 8)
assert _PolicyHandler.observed_auth == "Bearer top-secret-value"
assert _PolicyHandler.observed_payload["task"] == "Pick up the blue cube."
assert report["authenticated"] is True
assert "top-secret-value" not in json.dumps(report)
assert "top-secret-value" not in repr(actions)
def test_http_policy_error_redacts_server_token(monkeypatch):
class BrokenClient:
def __init__(self, **_kwargs):
pass
def __enter__(self):
return self
def __exit__(self, *_args):
pass
def stream(self, *_args, **_kwargs):
raise RuntimeError("Authorization: Bearer secret-http-token")
monkeypatch.setenv("VLA_POLICY_TOKEN", "secret-http-token")
monkeypatch.setattr(
robotics,
"require_module",
lambda *_args: SimpleNamespace(Client=BrokenClient),
)
with pytest.raises(RuntimeError) as exc:
robotics.call_http_policy(
_observation(),
endpoint="http://127.0.0.1:8787",
timeout_seconds=5,
allow_remote=False,
include_history=False,
)
assert "secret-http-token" not in str(exc.value)
assert "[REDACTED]" in str(exc.value)
def test_http_policy_rejects_declared_oversized_response(monkeypatch):
class Response:
headers = {
"content-length": str(robotics.MAX_HTTP_RESPONSE_BYTES + 1),
}
def __enter__(self):
return self
def __exit__(self, *_args):
pass
def raise_for_status(self):
pass
def iter_bytes(self):
raise AssertionError("Oversized response body must not be read.")
class Client:
def __init__(self, **_kwargs):
pass
def __enter__(self):
return self
def __exit__(self, *_args):
pass
def stream(self, *_args, **_kwargs):
return Response()
monkeypatch.setattr(
robotics,
"require_module",
lambda *_args: SimpleNamespace(Client=Client),
)
with pytest.raises(RuntimeError, match="32 MiB"):
robotics.call_http_policy(
_observation(),
endpoint="http://127.0.0.1:8787",
timeout_seconds=5,
allow_remote=False,
include_history=False,
)
def test_catalog_has_unique_official_entries_and_clear_readiness():
assert 12 <= len(VLA_MODEL_CATALOG) <= 30
checkpoints = []
for label, info in VLA_MODEL_CATALOG.items():
assert label == info.label
assert info.official_url.startswith("https://")
assert info.backend
assert info.status
if info.checkpoint:
checkpoints.append(info.checkpoint)
assert len(checkpoints) == len(set(checkpoints))
assert any(info.family == "SmolVLA" for info in VLA_MODEL_CATALOG.values())
assert any(info.family == "Isaac GR00T N1.7" for info in VLA_MODEL_CATALOG.values())
assert any(info.family == "OpenVLA-OFT" for info in VLA_MODEL_CATALOG.values())
def test_trajectory_preview_is_a_comfy_image():
profile = _joint_profile()
actions = RobotActions(
torch.linspace(-0.5, 0.5, 5 * 8).reshape(5, 8),
profile.action_names,
"preview",
)
preview = render_action_preview(actions, embodiment=profile, width=640, height=320)
assert preview.shape == (1, 320, 640, 3)
assert preview.dtype == torch.float32
assert 0 <= float(preview.min()) <= float(preview.max()) <= 1
def test_robotics_nodes_are_registered_and_have_safe_categories():
expected = {
"VLAEmbodimentProfile",
"VLAObservationBuilder",
"VLAHTTPPolicy",
"VLAOpenPIWebSocketPolicy",
"VLAGr00tZMQPolicy",
"VLAActionSafety",
"VLAActionsFromJSON",
"VLAActionChunkReplan",
"VLAActionInspect",
"VLATrajectoryPreview",
"VLAModelCatalog",
}
assert expected == set(robotics.NODE_CLASS_MAPPINGS)
for node in robotics.NODE_CLASS_MAPPINGS.values():
assert node.CATEGORY.startswith("VLM Nodes/Robotics")
assert "forceInput" not in repr(node.INPUT_TYPES())
def test_lerobot_sidecar_is_packaged_and_avoids_pickle_transport():
root = Path(robotics.__file__).parents[1]
server = root / "examples" / "robotics" / "lerobot_policy_server.py"
source = server.read_text(encoding="utf-8")
compile(source, str(server), "exec")
assert "pickle.loads" not in source
assert "VLA_POLICY_TOKEN" in source
assert "ThreadingHTTPServer" in source
def test_lerobot_sidecar_exposes_checkpoint_feature_contract(monkeypatch):
visual = SimpleNamespace(
type=SimpleNamespace(value="VISUAL"),
shape=(3, 256, 256),
)
action = {"type": "ACTION", "shape": (6,)}
assert lerobot_policy_server._feature_metadata(
{"observation.images.camera1": visual, "action": action}
) == {
"observation.images.camera1": {
"type": "VISUAL",
"shape": [3, 256, 256],
},
"action": {"type": "ACTION", "shape": [6]},
}
assert lerobot_policy_server._optional_config_int(
SimpleNamespace(chunk_size=50), "chunk_size"
) == 50
monkeypatch.setattr(
lerobot_policy_server,
"_decode_image",
lambda _frame: np.zeros((1, 1, 3), dtype=np.uint8),
)
oversized_task = {
"schema": "comfyui-vlm/robot-observation",
"version": 1,
"cameras": {
"observation.images.front": [
{"encoding": "base64-jpeg", "data": "unused"}
]
},
"state": [0.0],
"task": "x" * (lerobot_policy_server.MAX_TASK_CHARS + 1),
}
with pytest.raises(ValueError, match="task must contain"):
lerobot_policy_server._decode_observation(oversized_task)
def test_robotics_client_requirements_match_optional_extra():
root = Path(robotics.__file__).parents[1]
pyproject = (root / "pyproject.toml").read_text(encoding="utf-8")
match = re.search(r"^robotics-client\s*=\s*\[(.*?)^\]", pyproject, re.M | re.S)
assert match is not None
optional_extra = set(re.findall(r'"([^"]+)"', match.group(1)))
requirement_file = {
line.strip()
for line in (root / "requirements-robotics-client.txt")
.read_text(encoding="utf-8")
.splitlines()
if line.strip() and not line.lstrip().startswith("#")
}
assert optional_extra == requirement_file
+211
View File
@@ -0,0 +1,211 @@
import json
from pathlib import Path
import ComfyUI_VLM_nodes as package
import pytest
from ComfyUI_VLM_nodes.nodes import simpletext
def test_simple_text_preserves_legacy_default_and_appends_metrics():
result = simpletext.SimpleText().simple_text(" one\r\ntwo ")
assert result == (" one\r\ntwo ", 12, 2, 2)
normalized = simpletext.SimpleText().simple_text(
" one\r\ntwo ",
trim_edges=True,
normalize_newlines=True,
)
assert normalized == ("one\ntwo", 7, 2, 2)
def test_json_to_text_keeps_legacy_smart_rendering_and_adds_canonical_output():
response = simpletext.JsonToText().json_to_text(
'{"prompt":"Create a red kite","suggestion1":"at sunset","tags":["red","sky"]}'
)
assert response["result"][0] == "a red kite\n\nat sunset\n\ntags: red, sky"
assert json.loads(response["result"][1])["tags"] == ["red", "sky"]
assert response["result"][2] == 3
def test_json_to_text_parses_fenced_model_response_and_json_paths():
response = simpletext.JsonToText().json_to_text(
'Model response:\n```json\n{"result":{"items":[{"name":"café"}]}}\n```',
format_mode="Pretty JSON",
json_path="$.result.items[0]",
)
assert json.loads(response["result"][0]) == {"name": "café"}
assert response["result"][1] == '{"name":"café"}'
pointer = simpletext.VLMJSONExtract().extract(
'{"a/b":{"~key":[10,20]}}',
"/a~1b/~0key/1",
"Text",
"Error",
"",
)
assert pointer == ("20", True, "integer", "20")
def test_json_extract_handles_negative_indexes_and_missing_policy():
node = simpletext.VLMJSONExtract()
assert node.extract(
'{"items":["first","last"]}',
"$.items[-1]",
"Text",
"Error",
"",
)[:3] == ("last", True, "string")
assert node.extract(
'{"items":[]}',
"$.missing",
"Text",
"Default value",
"fallback",
)[:3] == ("fallback", False, "string")
with pytest.raises(ValueError, match="not found"):
node.extract("{}", "$.missing", "Text", "Error", "")
def test_text_join_drops_empty_and_duplicate_parts():
result = simpletext.VLMTextJoin().join(
" first ",
"Blank line",
"|",
True,
True,
True,
text_b="second",
text_c="first",
)
assert result == (
"first\n\nsecond",
'["first","second"]',
2,
)
def test_text_template_is_safe_explicit_and_supports_literal_braces():
result = simpletext.VLMTextTemplate().render(
"{{schema}} {subject}: {text1}",
'{"subject":"robot"}',
"Error",
text1="moving a box",
)
assert result[0] == "{schema} robot: moving a box"
assert json.loads(result[1]) == {
"subject": "robot",
"text1": "moving a box",
}
assert result[2] == "[]"
with pytest.raises(ValueError, match="missing"):
simpletext.VLMTextTemplate().render(
"{known} {unknown}",
'{"known":"yes"}',
"Error",
)
def test_text_clean_normalizes_fences_duplicates_and_length():
result, diagnostics_json = simpletext.VLMTextClean().clean(
"```text\r\nA B\r\nA B\r\nC\r\n```",
"NFKC",
"Collapse horizontal",
True,
True,
True,
5,
)
assert result == "A B\nC"
diagnostics = json.loads(diagnostics_json)
assert diagnostics["changed"] is True
assert diagnostics["duplicate_lines_removed"] == 1
assert diagnostics["truncated"] is False
def test_text_replace_literal_regex_and_errors():
node = simpletext.VLMTextReplace()
assert node.replace(
"Cat cat cat",
"cat",
"dog",
"Literal",
False,
2,
"Keep text",
)[:2] == ("dog dog cat", 2)
assert node.replace(
"a1 b22",
r"\d+",
"#",
"Regular expression",
True,
0,
"Keep text",
)[:2] == ("a# b#", 2)
with pytest.raises(ValueError, match="not found"):
node.replace("hello", "x", "y", "Literal", True, 0, "Error")
def test_text_split_outputs_real_list_and_stable_json():
items, items_json, count = simpletext.VLMTextSplit().split(
'[" first ","second","first",""]',
"JSON array",
",",
True,
True,
True,
0,
)
assert items == ["first", "second"]
assert json.loads(items_json) == items
assert count == 2
assert simpletext.VLMTextSplit.OUTPUT_IS_LIST == (True, False, False)
def test_text_inspector_and_view_text_report_same_metrics():
inspected = simpletext.VLMTextInspect().inspect("hello\nworld")
assert inspected[1:6] == (11, 11, 2, 2, 3)
assert len(inspected[6]) == 64
assert json.loads(inspected[7])["words"] == 2
viewed = simpletext.ViewText().view_text("hello\nworld")
assert viewed["result"][:4] == ("hello\nworld", 11, 2, 2)
assert viewed["ui"]["text"] == ["hello\nworld"]
def test_text_node_categories_aliases_and_legacy_ids_are_stable():
assert simpletext.NODE_CLASS_MAPPINGS["SimpleText"] is simpletext.SimpleText
assert simpletext.NODE_CLASS_MAPPINGS["JsonToText"] is simpletext.JsonToText
assert simpletext.NODE_CLASS_MAPPINGS["ViewText"] is simpletext.ViewText
assert set(simpletext.NODE_CLASS_MAPPINGS) == {
"SimpleText",
"JsonToText",
"ViewText",
"VLMTextJoin",
"VLMTextTemplate",
"VLMTextClean",
"VLMTextReplace",
"VLMJSONExtract",
"VLMTextSplit",
"VLMTextInspect",
}
assert simpletext.SimpleText.CATEGORY == "VLM Nodes/Text/Create"
assert simpletext.ViewText.CATEGORY == "VLM Nodes/Text/Inspect"
def test_text_toolkit_api_example_uses_registered_inputs_and_output_indexes():
root = Path(package.__file__).parent
prompt = json.loads(
(root / "examples" / "text_toolkit_api.json").read_text("utf-8")
)
assert prompt["5"]["inputs"]["text"] == ["4", 0]
for node in prompt.values():
node_class = package.NODE_CLASS_MAPPINGS[node["class_type"]]
declared = {
name
for group in node_class.INPUT_TYPES().values()
if isinstance(group, dict)
for name in group
}
assert set(node["inputs"]) <= declared
+735
View File
@@ -0,0 +1,735 @@
"""Contract tests for the GGUF text/LLM nodes in ``nodes/suggest.py``.
These nodes carry the pack's longest bug history (widget-index drift, unexpected
sampling kwargs, JSON that never parsed), so the assertions below pin the
behaviours those reports depended on rather than the models themselves. No
llama.cpp wheel and no GGUF weights are required.
"""
from __future__ import annotations
import json
from pathlib import Path
from types import SimpleNamespace
import pytest
from ComfyUI_VLM_nodes.nodes import suggest
from ComfyUI_VLM_nodes.nodes.runtime import LlamaHandle
MODEL_FILE = "some-model.gguf"
class FakeLlama:
"""Records the kwargs llama.cpp would have received."""
def __init__(self, content: str = "generated text"):
self.content = content
self.calls: list[dict] = []
def create_chat_completion(self, **kwargs):
self.calls.append(kwargs)
return {"choices": [{"message": {"content": self.content}}]}
class FakeHandle:
"""Stands in for LlamaHandle so caching can be observed without weights."""
instances: list[FakeHandle] = []
def __init__(self, model_path, **kwargs):
self.model_path = model_path
self.kwargs = kwargs
self.closed = False
self.llama = FakeLlama()
FakeHandle.instances.append(self)
def ensure_loaded(self):
return self.llama
def close(self):
self.closed = True
@pytest.fixture
def resolved_model(monkeypatch):
"""Bypass folder_paths so no real GGUF has to exist on disk."""
path = Path("/models/LLavacheckpoints") / MODEL_FILE
monkeypatch.setattr(suggest, "resolve_model_path", lambda name: path)
return path
@pytest.fixture
def fake_handles(monkeypatch):
FakeHandle.instances = []
monkeypatch.setattr(suggest, "LlamaHandle", FakeHandle)
return FakeHandle
# --------------------------------------------------------------------------
# Widget ordering. Comfy serializes widget values by position, so a reordered
# INPUT_TYPES silently rebinds every saved workflow (issue #156).
# --------------------------------------------------------------------------
def test_llm_sampler_widget_order_is_frozen():
assert list(suggest.LLMSampler.INPUT_TYPES()["required"]) == [
"system_msg",
"prompt",
"model",
"max_tokens",
"temperature",
"top_p",
"top_k",
"frequency_penalty",
"presence_penalty",
"repeat_penalty",
"seed",
]
def test_llm_prompt_generator_widget_order_is_frozen():
assert list(suggest.LLMPromptGenerator.INPUT_TYPES()["required"]) == [
"prompt",
"model",
"max_tokens",
"temperature",
"top_p",
"top_k",
"frequency_penalty",
"presence_penalty",
"repeat_penalty",
]
def test_llm_loader_widget_order_is_frozen():
schema = suggest.LLMLoader.INPUT_TYPES()
assert list(schema["required"]) == [
"ckpt_name",
"max_ctx",
"gpu_layers",
"n_threads",
]
# chat_format must stay ahead of the shared runtime widgets.
assert list(schema["optional"])[0] == "chat_format"
def test_structured_output_widget_order_is_frozen():
assert list(suggest.StructuredOutput.INPUT_TYPES()["required"]) == [
"prompt",
"model",
"temperature",
"attribute_name",
"attribute_type",
"attribute_description",
"categories",
]
def test_every_suggest_node_declares_a_callable_function_and_return_types():
for name, node_class in suggest.NODE_CLASS_MAPPINGS.items():
assert isinstance(node_class.RETURN_TYPES, tuple), name
assert node_class.RETURN_TYPES, name
assert callable(getattr(node_class, node_class.FUNCTION, None)), name
assert node_class.CATEGORY.startswith("VLM Nodes"), name
def test_display_names_cover_every_registered_node():
assert set(suggest.NODE_CLASS_MAPPINGS) == set(suggest.NODE_DISPLAY_NAME_MAPPINGS)
# --------------------------------------------------------------------------
# Sampling kwargs. Issue #144 was an "unexpected keyword argument" crash, so
# the plumbing from node widget to create_chat_completion is asserted directly.
# --------------------------------------------------------------------------
def test_llm_sampler_forwards_every_sampling_argument():
llama = FakeLlama("a description")
result = suggest.LLMSampler().generate_text_advanced(
system_msg="be terse",
prompt="describe a cat",
model=llama,
max_tokens=64,
temperature=0.7,
top_p=0.8,
top_k=20,
frequency_penalty=0.1,
presence_penalty=0.2,
repeat_penalty=1.3,
seed=1234,
)
assert result == ("a description",)
(call,) = llama.calls
assert call["messages"] == [
{"role": "system", "content": "be terse"},
{"role": "user", "content": "describe a cat"},
]
assert call["max_tokens"] == 64
assert call["temperature"] == 0.7
assert call["top_p"] == 0.8
assert call["top_k"] == 20
assert call["frequency_penalty"] == 0.1
assert call["presence_penalty"] == 0.2
assert call["repeat_penalty"] == 1.3
assert call["seed"] == 1234
# No response_format unless a structured node asked for one.
assert "response_format" not in call
def test_chat_unwraps_a_lazy_handle_before_generating(fake_handles, resolved_model):
handle = FakeHandle(resolved_model)
text = suggest._chat(handle, prompt="hi", system="sys")
assert text == "generated text"
def test_llama_chat_content_rejects_an_empty_completion():
llama = FakeLlama(content=" ")
with pytest.raises(RuntimeError, match="empty response"):
suggest.LLMSampler().generate_text_advanced(
system_msg="s",
prompt="p",
model=llama,
max_tokens=8,
temperature=0.1,
top_p=0.9,
top_k=1,
frequency_penalty=0.0,
presence_penalty=0.0,
repeat_penalty=1.0,
seed=1,
)
# --------------------------------------------------------------------------
# Lazy loading. Building a loader node must not touch llama.cpp or the GGUF.
# --------------------------------------------------------------------------
def test_llm_loader_builds_a_lazy_handle_without_loading_weights(resolved_model):
(handle,) = suggest.LLMLoader().load_llm_checkpoint(
ckpt_name=MODEL_FILE,
max_ctx=8192,
gpu_layers=20,
n_threads=6,
)
assert isinstance(handle, LlamaHandle)
assert handle.model_path == resolved_model
assert handle.n_ctx == 8192
assert handle.n_gpu_layers == 20
assert handle.n_threads == 6
# Nothing was loaded: the llama.cpp object is still absent.
assert handle._llm is None
def test_llm_loader_treats_blank_chat_format_as_the_embedded_template(resolved_model):
(blank,) = suggest.LLMLoader().load_llm_checkpoint(
MODEL_FILE, 2048, -1, 4, chat_format=" "
)
assert blank.chat_format is None
(explicit,) = suggest.LLMLoader().load_llm_checkpoint(
MODEL_FILE, 2048, -1, 4, chat_format=" chatml "
)
assert explicit.chat_format == "chatml"
def test_llm_loader_forwards_advanced_runtime_options(resolved_model):
(handle,) = suggest.LLMLoader().load_llm_checkpoint(
MODEL_FILE,
2048,
-1,
4,
n_batch=256,
n_ubatch=128,
flash_attention="Disabled",
use_mmap=False,
split_mode="Single GPU",
main_gpu=2,
tensor_split="0.6,0.4",
)
assert handle.n_batch == 256
assert handle.n_ubatch == 128
assert handle.flash_attention == "Disabled"
assert handle.use_mmap is False
assert handle.split_mode == "Single GPU"
assert handle.main_gpu == 2
assert handle.tensor_split == [0.6, 0.4]
# --------------------------------------------------------------------------
# Structured output.
# --------------------------------------------------------------------------
def _capture_chat(monkeypatch, payload):
"""Replace _chat so the generated JSON Schema can be inspected."""
recorded: dict = {}
def fake_chat(model, **kwargs):
recorded.update(kwargs)
return payload
monkeypatch.setattr(suggest, "_chat", fake_chat)
return recorded
@pytest.mark.parametrize(
("declared", "expected"),
[
("str", "string"),
("int", "integer"),
("float", "number"),
("bool", "boolean"),
],
)
def test_structured_output_maps_scalar_types_to_json_schema(
monkeypatch, declared, expected
):
recorded = _capture_chat(monkeypatch, json.dumps({"result": "value"}))
suggest.StructuredOutput().keyword_extract(
prompt="p",
model=object(),
temperature=0.1,
attribute_name="result",
attribute_type=declared,
attribute_description="a description",
categories="",
)
schema = recorded["response_format"]["schema"]
assert schema["properties"]["result"]["type"] == expected
assert schema["properties"]["result"]["description"] == "a description"
assert schema["required"] == ["result"]
assert schema["additionalProperties"] is False
assert recorded["response_format"]["type"] == "json_object"
def test_structured_output_builds_an_enum_for_categories(monkeypatch):
recorded = _capture_chat(monkeypatch, json.dumps({"mood": "calm"}))
(value,) = suggest.StructuredOutput().keyword_extract(
prompt="p",
model=object(),
temperature=0.1,
attribute_name=" mood ",
attribute_type="Category",
attribute_description="",
categories=" calm , tense ,, bright ",
)
schema = recorded["response_format"]["schema"]
assert schema["properties"]["mood"]["enum"] == ["calm", "tense", "bright"]
assert schema["properties"]["mood"]["type"] == "string"
assert value == "calm"
def test_structured_output_serializes_non_string_values(monkeypatch):
_capture_chat(monkeypatch, json.dumps({"count": 7}))
(value,) = suggest.StructuredOutput().keyword_extract(
prompt="p",
model=object(),
temperature=0.1,
attribute_name="count",
attribute_type="int",
attribute_description="",
categories="",
)
assert value == "7"
def test_structured_output_rejects_an_empty_attribute_name():
with pytest.raises(ValueError, match="attribute_name cannot be empty"):
suggest.StructuredOutput().keyword_extract(
prompt="p",
model=object(),
temperature=0.1,
attribute_name=" ",
attribute_type="str",
attribute_description="",
categories="",
)
def test_structured_output_rejects_a_category_without_values():
with pytest.raises(ValueError, match="at least one comma-separated value"):
suggest.StructuredOutput().keyword_extract(
prompt="p",
model=object(),
temperature=0.1,
attribute_name="mood",
attribute_type="Category",
attribute_description="",
categories=" , ",
)
def test_structured_chat_reports_unparseable_json_with_a_bounded_excerpt(monkeypatch):
_capture_chat(monkeypatch, "x" * 900)
with pytest.raises(RuntimeError, match="did not return valid JSON") as error:
suggest.KeywordExtraction().keyword_extract(
prompt="p", model=object(), temperature=0.1
)
# The raw completion is truncated so a runaway response cannot flood the log.
assert len(str(error.value)) < 600
def test_keyword_extraction_returns_the_raw_json_document(monkeypatch):
payload = json.dumps(
{
"main_character": ["cat"],
"artform": ["photo"],
"photo_type": ["portrait"],
"color_with_objects": ["black cat"],
"digital_artform": [],
"background": ["studio"],
"lighting": ["soft"],
}
)
_capture_chat(monkeypatch, payload)
(raw,) = suggest.KeywordExtraction().keyword_extract(
prompt="a cat", model=object(), temperature=0.1
)
assert json.loads(raw)["main_character"] == ["cat"]
def test_llava_prompt_generator_returns_only_the_prompt_field(monkeypatch):
_capture_chat(monkeypatch, json.dumps({"prompt": "a moody portrait"}))
(text,) = suggest.LLavaPromptGenerator().generate_prompts(
prompt="p", model=object(), temperature=0.1
)
assert text == "a moody portrait"
def test_creative_art_prompt_generator_prefers_the_narrative(monkeypatch):
_capture_chat(
monkeypatch,
json.dumps(
{
"techniques": {"preferred": ["ink"], "avoided": []},
"theme": {"core_subject": "a harbour"},
"style": {"desired": ["muted"], "undesired": []},
"creative_descriptions": [{"description": "a quiet harbour at dawn"}],
}
),
)
(text,) = suggest.CreativeArtPromptGenerator().create_creative_art_prompts(
prompt="p", model=object(), temperature=0.1
)
assert text == "a quiet harbour at dawn"
def test_creative_art_prompt_generator_composes_a_fallback_without_narratives(
monkeypatch,
):
_capture_chat(
monkeypatch,
json.dumps(
{
"techniques": {"preferred": ["ink", "wash"], "avoided": []},
"theme": {"core_subject": "a harbour"},
"style": {"desired": ["muted", "grainy"], "undesired": []},
"creative_descriptions": [],
}
),
)
(text,) = suggest.CreativeArtPromptGenerator().create_creative_art_prompts(
prompt="p", model=object(), temperature=0.1
)
assert text == (
"a harbour. Techniques: ink, wash. Visual style: muted, grainy."
)
def test_suggester_switches_instruction_on_the_randomize_toggle(monkeypatch):
payload = json.dumps(
{f"suggestion{index}": f"idea {index}" for index in range(1, 6)}
)
similar = _capture_chat(monkeypatch, payload)
suggest.Suggester().generate_suggestions(
prompt="p", model=object(), temperature=0.1, randomize=True
)
assert "close, useful variations" in similar["system"]
different = _capture_chat(monkeypatch, payload)
suggest.Suggester().generate_suggestions(
prompt="p", model=object(), temperature=0.1, randomize=False
)
assert "deliberately different" in different["system"]
# --------------------------------------------------------------------------
# Handle caching. Issue #137 was "model never unloads"; these pin the reuse
# and teardown rules of the optional-memory-free nodes.
# --------------------------------------------------------------------------
def test_cached_llm_reuses_one_handle_for_identical_settings(
fake_handles, resolved_model
):
node = suggest.LLMOptionalMemoryFreeSimple()
first = node._model(MODEL_FILE, 2048, -1, 4)
second = node._model(MODEL_FILE, 2048, -1, 4)
assert first is second
assert len(fake_handles.instances) == 1
def test_cached_llm_closes_the_previous_handle_when_settings_change(
fake_handles, resolved_model
):
node = suggest.LLMOptionalMemoryFreeSimple()
first = node._model(MODEL_FILE, 2048, -1, 4)
second = node._model(MODEL_FILE, 4096, -1, 4)
assert first is not second
assert first.closed is True
assert second.closed is False
assert len(fake_handles.instances) == 2
def test_cached_llm_rebuilds_when_an_advanced_runtime_option_changes(
fake_handles, resolved_model
):
node = suggest.LLMOptionalMemoryFreeSimple()
node._model(MODEL_FILE, 2048, -1, 4, n_batch=512)
node._model(MODEL_FILE, 2048, -1, 4, n_batch=256)
assert len(fake_handles.instances) == 2
def test_cached_llm_unload_releases_the_handle(fake_handles, resolved_model):
node = suggest.LLMOptionalMemoryFreeSimple()
handle = node._model(MODEL_FILE, 2048, -1, 4)
node._maybe_unload(False)
assert handle.closed is False
assert node._handle is handle
node._maybe_unload(True)
assert handle.closed is True
assert node._handle is None
assert node._key is None
def test_any_type_never_reports_a_type_mismatch():
assert (suggest.ANY != "IMAGE") is False
assert (suggest.ANY != "STRING") is False
# --------------------------------------------------------------------------
# ChatMusician. Issue #149's workaround was for users to append "respond in
# ABC notation starting with X:1" themselves; the node now owns that.
# --------------------------------------------------------------------------
def _chat_musician_kwargs():
return {
"max_tokens": 256,
"temperature": 0.2,
"top_p": 0.9,
"top_k": 40,
"frequency_penalty": 0.0,
"presence_penalty": 0.0,
"repeat_penalty": 1.1,
"seed": 42,
"sample_rate": 44100,
}
ABC_TUNE = "X:1\nT:Test\nM:4/4\nK:C\nCDEF|GABc|"
def test_chat_musician_asks_for_abc_notation_without_user_help(monkeypatch):
recorded = _capture_chat(monkeypatch, ABC_TUNE)
monkeypatch.setattr(suggest, "require_module", lambda *a, **k: _fake_symusic())
suggest.ChatMusician().chat_musician(
prompt="a waltz", model=object(), **_chat_musician_kwargs()
)
assert "ABC notation" in recorded["prompt"]
assert "X:" in recorded["prompt"]
assert "a waltz" in recorded["prompt"]
assert "ABC notation" in recorded["system"]
def test_chat_musician_rejects_a_response_without_an_abc_header(monkeypatch):
_capture_chat(monkeypatch, "Sure! Here is a lovely tune for you.")
with pytest.raises(RuntimeError, match="did not contain ABC notation"):
suggest.ChatMusician().chat_musician(
prompt="p", model=object(), **_chat_musician_kwargs()
)
def test_chat_musician_strips_preamble_before_the_abc_header(monkeypatch):
_capture_chat(monkeypatch, "Here you go:\n\n" + ABC_TUNE)
monkeypatch.setattr(suggest, "require_module", lambda *a, **k: _fake_symusic())
abc, _legacy, _rate, _audio = suggest.ChatMusician().chat_musician(
prompt="p", model=object(), **_chat_musician_kwargs()
)
assert abc.startswith("X:1")
assert "Here you go" not in abc
def test_chat_musician_returns_comfy_audio_and_legacy_layouts(monkeypatch):
_capture_chat(monkeypatch, ABC_TUNE)
monkeypatch.setattr(suggest, "require_module", lambda *a, **k: _fake_symusic())
_abc, legacy, rate, audio = suggest.ChatMusician().chat_musician(
prompt="p", model=object(), **_chat_musician_kwargs()
)
# Comfy AUDIO is [batch, channels, samples].
assert audio["waveform"].shape == (1, 2, 100)
assert audio["sample_rate"] == 44100
assert rate == 44100
# soundfile-compatible legacy output is [samples, channels].
assert legacy.shape == (100, 2)
def _fake_symusic():
"""A symusic stand-in so the AUDIO contract is testable without the wheel."""
import numpy as np
class Synthesizer:
def __init__(self, sample_rate):
self.sample_rate = sample_rate
def render(self, score, stereo=True):
return np.zeros((2, 100), dtype=np.float32)
class Score:
@staticmethod
def from_abc(abc):
return SimpleNamespace(abc=abc)
return SimpleNamespace(Score=Score, Synthesizer=Synthesizer)
def _memory_free_kwargs(**overrides):
kwargs = {
"ckpt_name": MODEL_FILE,
"max_ctx": 2048,
"gpu_layers": -1,
"n_threads": 4,
"prompt": "write a haiku",
"temperature": 0.2,
"unload": False,
}
kwargs.update(overrides)
return kwargs
def test_memory_free_simple_generates_through_the_cached_handle(
fake_handles, resolved_model
):
node = suggest.LLMOptionalMemoryFreeSimple()
(text,) = node.generate_text(**_memory_free_kwargs())
assert text == "generated text"
(handle,) = fake_handles.instances
assert handle.closed is False
assert node._handle is handle
(call,) = handle.llama.calls
assert call["temperature"] == 0.2
assert call["messages"][1]["content"] == "write a haiku"
def test_memory_free_simple_unloads_after_generating_when_asked(
fake_handles, resolved_model
):
node = suggest.LLMOptionalMemoryFreeSimple()
node.generate_text(**_memory_free_kwargs(unload=True))
(handle,) = fake_handles.instances
assert handle.closed is True
assert node._handle is None
def test_memory_free_simple_unloads_even_when_generation_fails(
fake_handles, resolved_model, monkeypatch
):
"""Issue #137: a failed generation must not strand the model in VRAM."""
def explode(*args, **kwargs):
raise RuntimeError("llama.cpp exploded")
monkeypatch.setattr(suggest, "_chat", explode)
node = suggest.LLMOptionalMemoryFreeSimple()
with pytest.raises(RuntimeError, match="exploded"):
node.generate_text(**_memory_free_kwargs(unload=True))
(handle,) = fake_handles.instances
assert handle.closed is True
assert node._handle is None
def test_memory_free_advanced_forwards_every_sampling_argument(
fake_handles, resolved_model
):
node = suggest.LLMOptionalMemoryFreeAdvanced()
signature = suggest.LLMOptionalMemoryFreeAdvanced.INPUT_TYPES()["required"]
assert "system_msg" in signature
(text,) = node.generate_text_advanced(
ckpt_name=MODEL_FILE,
max_ctx=2048,
gpu_layers=-1,
n_threads=4,
system_msg="be brief",
prompt="a haiku",
max_tokens=48,
temperature=0.5,
top_p=0.8,
top_k=15,
frequency_penalty=0.1,
presence_penalty=0.2,
repeat_penalty=1.2,
seed=11,
unload=False,
)
assert text == "generated text"
(handle,) = fake_handles.instances
(call,) = handle.llama.calls
assert call["messages"][0] == {"role": "system", "content": "be brief"}
assert call["max_tokens"] == 48
assert call["temperature"] == 0.5
assert call["top_p"] == 0.8
assert call["top_k"] == 15
assert call["seed"] == 11
def test_cached_llm_key_is_insensitive_to_runtime_option_ordering(
fake_handles, resolved_model
):
node = suggest.LLMOptionalMemoryFreeSimple()
node._model(MODEL_FILE, 2048, -1, 4, n_batch=256, main_gpu=1)
node._model(MODEL_FILE, 2048, -1, 4, main_gpu=1, n_batch=256)
# Keyword order must not invalidate the cache and reload the GGUF.
assert len(fake_handles.instances) == 1
def test_schema_helper_emits_a_json_schema_for_a_pydantic_model():
schema = suggest._schema(suggest.PromptGen)
assert schema["properties"]["prompt"]["type"] == "string"
assert schema["required"] == ["prompt"]
def test_response_content_extraction_matches_the_runtime_helper():
response = {"choices": [{"message": {"content": " text "}}]}
assert suggest._response_content(response) == "text"
def test_stub_handle_matches_the_real_handle_api():
"""Guard the stub: LlamaHandle must keep the API these tests rely on."""
assert callable(getattr(LlamaHandle, "ensure_loaded", None))
assert callable(getattr(LlamaHandle, "close", None))
+377
View File
@@ -0,0 +1,377 @@
from __future__ import annotations
import json
from pathlib import Path
import pytest
import torch
from ComfyUI_VLM_nodes.nodes.video_intelligence import (
NODE_CLASS_MAPPINGS,
VLMAdaptiveFrameSampler,
build_scene_state,
build_video_reasoning_prompt,
parse_video_reasoning_output,
resize_video_for_analysis,
sample_video_frames,
scene_state_summary,
track_aware_crops,
)
from ComfyUI_VLM_nodes.nodes.vision_types import (
Detection,
EventSequence,
SceneState,
SelectedVideoFrame,
Track,
TrackSequence,
VideoFrameSelection,
)
def _moving_video(frame_count=20, height=48, width=64):
frames = torch.zeros((frame_count, height, width, 3), dtype=torch.float32)
for frame_index in range(frame_count):
x = min(width - 9, 2 + frame_index * 2)
frames[frame_index, 16:28, x : x + 8, 0] = 1.0
if frame_index >= frame_count // 2:
frames[frame_index, :, :, 2] += 0.55
return frames.clamp(0, 1)
def _tracks(width=64, height=48, frame_count=20, fps=10.0):
detections = []
for frame_index in (0, 5, 10, 15, 19):
x = min(width - 12, 2 + frame_index * 2)
detections.append(
Detection(
bbox_xyxy=(x, 14, x + 10, 30),
label="red object",
score=0.9 - frame_index * 0.005,
frame_index=frame_index,
timestamp=frame_index / fps,
track_id=3,
metadata={"track_state": "active"},
)
)
return TrackSequence(
width=width,
height=height,
frame_count=frame_count,
fps=fps,
tracks=(
Track(
track_id=3,
detections=tuple(detections),
label="red object",
score=0.85,
),
),
source="unit-test-tracker",
)
def _selection():
return VideoFrameSelection(
width=64,
height=48,
source_frame_count=20,
fps=10.0,
strategy="Hybrid: scene + motion + tracks",
frames=(
SelectedVideoFrame(0, 0.0, 1.0, ("first-frame",)),
SelectedVideoFrame(5, 0.5, 0.7, ("motion",)),
SelectedVideoFrame(10, 1.0, 0.9, ("scene-change",)),
SelectedVideoFrame(19, 1.9, 1.0, ("last-frame",)),
),
)
def test_uniform_sampling_is_deterministic_and_preserves_timestamps():
frames = _moving_video()
sampled, selection, diagnostics = sample_video_frames(
frames,
fps=10.0,
max_frames=5,
strategy="Uniform coverage",
)
assert sampled.shape == (5, 48, 64, 3)
assert selection.indices == (0, 5, 10, 14, 19)
assert selection.timestamps == pytest.approx((0.0, 0.5, 1.0, 1.4, 1.9))
assert diagnostics["visual_reduction_ratio"] == pytest.approx(0.75)
assert torch.equal(sampled[2], frames[10])
def test_hybrid_sampling_captures_boundaries_scene_change_and_motion():
frames = _moving_video()
sampled, selection, diagnostics = sample_video_frames(
frames,
fps=10.0,
max_frames=7,
strategy="Hybrid: scene + motion + tracks",
minimum_gap_seconds=0.2,
)
assert sampled.shape[0] == 7
assert selection.indices[0] == 0
assert selection.indices[-1] == 19
assert any(9 <= index <= 11 for index in selection.indices)
assert diagnostics["motion_peak"] > 0
assert diagnostics["scene_peak"] > 0
assert selection.to_json() == VideoFrameSelection.from_json(
selection.to_json()
).to_json()
def test_track_priority_uses_track_changes_and_validates_dimensions():
frames = _moving_video()
tracks = _tracks()
_sampled, selection, diagnostics = sample_video_frames(
frames,
fps=10.0,
max_frames=6,
strategy="Track-change priority",
tracks=tracks,
)
assert diagnostics["track_peak"] == pytest.approx(1.0)
assert any(
"track-change" in frame.reasons for frame in selection.frames
)
bad_tracks = TrackSequence(
width=65,
height=48,
frame_count=20,
fps=10,
tracks=(),
)
with pytest.raises(ValueError, match="dimensions"):
sample_video_frames(
frames,
fps=10,
max_frames=4,
tracks=bad_tracks,
)
@pytest.mark.parametrize(
"frames",
[
torch.zeros(4, 16, 16),
torch.zeros(4, 16, 16, 2),
torch.zeros(0, 16, 16, 3),
torch.zeros(4, 16, 16, 3, dtype=torch.uint8),
],
)
def test_sampling_rejects_invalid_video_tensors(frames):
with pytest.raises((TypeError, ValueError)):
sample_video_frames(frames, fps=24, max_frames=4)
def test_track_aware_crops_preserve_identity_and_source_frame_mapping():
frames = _moving_video()
crops, manifest = track_aware_crops(
frames,
_tracks(),
crops_per_track=3,
max_crops=8,
output_size=96,
context_scale=1.4,
)
assert crops.shape == (3, 96, 96, 3)
assert [item["track_id"] for item in manifest] == [3, 3, 3]
assert [item["source_frame_index"] for item in manifest] == [0, 10, 19]
assert crops.max().item() > 0.5
def test_analysis_resize_reduces_pixels_without_changing_batch_or_aspect():
frames = _moving_video(height=128, width=256)
resized = resize_video_for_analysis(frames, max_side=128)
assert resized.shape == (20, 64, 128, 3)
assert resized.min().item() >= 0
assert resized.max().item() <= 1
assert resize_video_for_analysis(frames, max_side=0) is frames
def test_scene_state_ignores_predicted_track_samples_and_computes_velocity():
tracks = _tracks()
predicted = Detection(
bbox_xyxy=(48, 14, 58, 30),
label="red object",
frame_index=18,
timestamp=1.8,
track_id=3,
metadata={"track_state": "predicted"},
)
track = tracks.tracks[0]
with_prediction = TrackSequence(
width=tracks.width,
height=tracks.height,
frame_count=tracks.frame_count,
fps=tracks.fps,
tracks=(
Track(
track_id=3,
detections=tuple(
sorted(
(*track.detections, predicted),
key=lambda item: item.frame_index,
)
),
label=track.label,
),
),
)
scene = build_scene_state(with_prediction)
assert len(scene.objects) == 1
item = scene.objects[0]
assert item.observation_count == 5
assert item.velocity_xy_px_s[0] > 0
assert "#3 red object" in scene_state_summary(scene)
assert SceneState.from_json(scene.to_json()).to_json() == scene.to_json()
def test_reasoning_prompt_explains_irregular_source_timeline():
prompt = build_video_reasoning_prompt(
_selection(),
task="Robotics scene understanding",
question="",
max_events=12,
)
assert "irregularly spaced" in prompt
assert "supplied image 2: source frame 10, timestamp 1.000000s" in prompt
assert "Do not propose motor commands" in prompt
assert "evidence_frame_indices" in prompt
def test_structured_video_output_parses_fenced_json_and_preserves_evidence():
response = """Result:
```json
{
"summary": "A red object moves to the right.",
"events": [
{
"start_time": 0.0,
"end_time": 1.9,
"label": "object motion",
"text": "The red object moves from left to right.",
"score": 0.94,
"evidence_frame_indices": [0, 10, 19]
}
]
}
```
"""
summary, events, normalized = parse_video_reasoning_output(
response,
_selection(),
)
assert summary == "A red object moves to the right."
assert len(events.events) == 1
assert events.events[0].metadata["evidence_frame_indices"] == (0, 10, 19)
assert json.loads(normalized)["events"][0]["label"] == "object motion"
def test_structured_output_normalizes_supplied_image_positions_to_source_frames():
response = json.dumps(
{
"summary": "A transition occurs.",
"events": [
{
"start_time": 0.5,
"end_time": 1.9,
"label": "transition",
"text": "The scene changes.",
"score": 0.8,
# Positions 1 and 3 in the supplied image batch.
"evidence_frame_indices": [1, 3],
}
],
}
)
_summary, events, _normalized = parse_video_reasoning_output(
response,
_selection(),
)
event = events.events[0]
assert event.metadata["evidence_frame_indices"] == (5, 19)
assert event.metadata["evidence_index_mode"] == "supplied-image-position"
@pytest.mark.parametrize(
("event_patch", "error"),
[
({"end_time": 2.1}, "outside"),
({"score": 1.2}, "between"),
({"evidence_frame_indices": [0, 7]}, "not supplied"),
({"evidence_frame_indices": [0, 0]}, "duplicate"),
({"label": "", "text": ""}, "requires"),
],
)
def test_structured_video_output_rejects_unverifiable_events(event_patch, error):
event = {
"start_time": 0.0,
"end_time": 1.0,
"label": "motion",
"text": "Object moves.",
"score": 0.8,
"evidence_frame_indices": [0, 10],
}
event.update(event_patch)
with pytest.raises((TypeError, ValueError), match=error):
parse_video_reasoning_output(
json.dumps({"summary": "test", "events": [event]}),
_selection(),
)
def test_scene_state_accepts_validated_events():
_summary, events, _normalized = parse_video_reasoning_output(
json.dumps(
{
"summary": "motion",
"events": [
{
"start_time": 0.0,
"end_time": 1.9,
"label": "motion",
"text": "Object moves.",
"score": 0.9,
"evidence_frame_indices": [0, 19],
}
],
}
),
_selection(),
)
scene = build_scene_state(_tracks(), events)
assert isinstance(events, EventSequence)
assert len(scene.events) == 1
assert "Event 0.000s–1.900s" in scene_state_summary(scene)
def test_node_surface_registers_all_video_intelligence_nodes():
assert set(NODE_CLASS_MAPPINGS) == {
"VLMAdaptiveFrameSampler",
"VLMTrackAwareCrops",
"VLMBuildSceneState",
"VLMVideoReasoningPrompt",
"VLMEventsFromVideoJSON",
"VLMVideoTemporalReasoner",
}
inputs = VLMAdaptiveFrameSampler.INPUT_TYPES()
assert inputs["required"]["frames"][0] == "IMAGE"
assert inputs["optional"]["tracks"][0] == "VLM_TRACKS"
reasoner = NODE_CLASS_MAPPINGS["VLMVideoTemporalReasoner"]
assert reasoner.RETURN_NAMES[-2:] == ("events_json", "selection_json")
def test_api_example_uses_direct_json_outputs_and_preview():
example = json.loads(
(
Path(__file__).resolve().parents[1]
/ "examples"
/ "vision"
/ "video_temporal_reasoning_api.json"
).read_text(encoding="utf-8")
)
assert example["3"]["class_type"] == "VLMVideoTemporalReasoner"
assert example["5"]["inputs"]["text"] == ["3", 6]
assert example["6"]["inputs"]["text"] == ["3", 7]
assert example["8"]["inputs"]["images"] == ["3", 3]
+8
View File
@@ -13,11 +13,15 @@ DETECTIONS_SCHEMA = vision_types.DETECTIONS_SCHEMA
EVENTS_SCHEMA = vision_types.EVENTS_SCHEMA
POINTS_SCHEMA = vision_types.POINTS_SCHEMA
SCHEMA_VERSION = vision_types.SCHEMA_VERSION
SCENE_STATE_SCHEMA = vision_types.SCENE_STATE_SCHEMA
TRACKS_SCHEMA = vision_types.TRACKS_SCHEMA
VIDEO_SELECTION_SCHEMA = vision_types.VIDEO_SELECTION_SCHEMA
VLM_DETECTIONS = vision_types.VLM_DETECTIONS
VLM_EVENTS = vision_types.VLM_EVENTS
VLM_POINTS = vision_types.VLM_POINTS
VLM_SCENE_STATE = vision_types.VLM_SCENE_STATE
VLM_TRACKS = vision_types.VLM_TRACKS
VLM_VIDEO_SELECTION = vision_types.VLM_VIDEO_SELECTION
Detection = vision_types.Detection
DetectionSequence = vision_types.DetectionSequence
EventSequence = vision_types.EventSequence
@@ -86,11 +90,15 @@ def test_public_socket_and_schema_names_are_stable():
assert VLM_TRACKS == "VLM_TRACKS"
assert VLM_POINTS == "VLM_POINTS"
assert VLM_EVENTS == "VLM_EVENTS"
assert VLM_VIDEO_SELECTION == "VLM_VIDEO_SELECTION"
assert VLM_SCENE_STATE == "VLM_SCENE_STATE"
assert SCHEMA_VERSION == 1
assert DETECTIONS_SCHEMA == "comfyui-vlm/detections"
assert TRACKS_SCHEMA == "comfyui-vlm/tracks"
assert POINTS_SCHEMA == "comfyui-vlm/points"
assert EVENTS_SCHEMA == "comfyui-vlm/events"
assert VIDEO_SELECTION_SCHEMA == "comfyui-vlm/video-selection"
assert SCENE_STATE_SCHEMA == "comfyui-vlm/scene-state"
def test_detection_payload_is_validated_immutable_and_mask_safe():
+62
View File
@@ -0,0 +1,62 @@
import { app } from "../../../scripts/app.js";
const LLM_NODE = "PromptGenerateAPI";
const SAFE_SOURCE = "Provider environment variable";
const NO_KEY_SOURCE = "No key (loopback custom endpoint only)";
const SAFE_SOURCES = new Set([SAFE_SOURCE, NO_KEY_SOURCE]);
const CREDENTIAL_WIDGET_INDEX = 2;
function visitGraphNodes(graphData, callback) {
for (const node of graphData?.nodes ?? []) {
callback(node);
}
for (const subgraph of graphData?.definitions?.subgraphs ?? []) {
visitGraphNodes(subgraph, callback);
}
}
function scrubSerializedNode(node) {
if (node?.type !== LLM_NODE) {
return;
}
const values = node.widgets_values;
if (Array.isArray(values)) {
const saved = values[CREDENTIAL_WIDGET_INDEX];
if (!SAFE_SOURCES.has(saved)) {
values[CREDENTIAL_WIDGET_INDEX] = SAFE_SOURCE;
}
return;
}
if (values && typeof values === "object") {
// Some frontend versions serialize widgets by name.
delete values.api_key;
if (!SAFE_SOURCES.has(values.credential_source)) {
values.credential_source = SAFE_SOURCE;
}
}
}
function enforceLiveWidget(node) {
if (node?.type !== LLM_NODE) {
return;
}
const widget = node.widgets?.find(
(item) => item.name === "credential_source",
);
if (widget && !SAFE_SOURCES.has(widget.value)) {
widget.value = SAFE_SOURCE;
widget.callback?.(SAFE_SOURCE);
}
}
app.registerExtension({
name: "gokayfem.vlm.api-credential-security",
async beforeConfigureGraph(graphData) {
// Runs on the cloned workflow before LiteGraph creates any widgets, so a
// legacy key never reaches a DOM input or the active graph.
visitGraphNodes(graphData, scrubSerializedNode);
},
loadedGraphNode(node) {
enforceLiveWidget(node);
},
});
+179 -83
View File
@@ -3,83 +3,151 @@ import { api } from "../../../scripts/api.js";
const OUTPUT_NAME = "output_text";
const VIEW_TEXT_NODE = "ViewText";
const MODERN_VLM_NODE = "ModernVLM";
const STREAMING_SOURCE_NODES = new Set([
"ModernVLM",
"Moondream31Query",
"Moondream31Caption",
"PromptGenerateAPI",
"HostedVLMAPI",
"VLMVideoTemporalReasoner",
]);
function textMetrics(value) {
const text = String(value ?? "");
const words = text.trim() ? text.trim().split(/\s+/u).length : 0;
const lines = text ? text.split("\n").length : 0;
return `${text.length.toLocaleString()} chars · ${words.toLocaleString()} words · ${lines.toLocaleString()} lines`;
}
function makeButton(label, title, handler) {
const button = document.createElement("button");
button.textContent = label;
button.type = "button";
button.title = title;
button.addEventListener("click", handler);
Object.assign(button.style, {
color: "var(--input-text, #ddd)",
background: "var(--comfy-input-bg, #202020)",
border: "1px solid var(--border-color, #555)",
borderRadius: "5px",
padding: "3px 8px",
cursor: "pointer",
whiteSpace: "nowrap",
});
return button;
}
function ensureOutputWidget(node) {
let widget = node.widgets?.find((item) => item.name === OUTPUT_NAME);
if (!widget) {
const container = document.createElement("div");
const header = document.createElement("div");
const status = document.createElement("span");
const copy = document.createElement("button");
const output = document.createElement("textarea");
status.textContent = "Ready";
copy.textContent = "Copy";
copy.type = "button";
copy.title = "Copy the complete VLM response";
copy.addEventListener("click", async () => {
const previous = copy.textContent;
try {
await navigator.clipboard.writeText(output.value);
copy.textContent = "Copied";
} catch {
copy.textContent = "Copy failed";
}
window.setTimeout(() => {
copy.textContent = previous;
}, 1200);
});
output.readOnly = true;
output.setAttribute("aria-label", "VLM text output");
header.append(status, copy);
container.append(header, output);
Object.assign(container.style, {
display: "flex",
flexDirection: "column",
width: "100%",
height: "100%",
minHeight: "150px",
gap: "6px",
});
Object.assign(header.style, {
display: "flex",
alignItems: "center",
justifyContent: "space-between",
color: "var(--descrip-text, #aaa)",
fontSize: "12px",
});
Object.assign(copy.style, {
color: "var(--input-text, #ddd)",
background: "var(--comfy-input-bg, #202020)",
border: "1px solid var(--border-color, #555)",
borderRadius: "5px",
padding: "3px 9px",
cursor: "pointer",
});
Object.assign(output.style, {
width: "100%",
flex: "1",
minHeight: "120px",
resize: "vertical",
boxSizing: "border-box",
color: "var(--input-text, #ddd)",
background: "var(--comfy-input-bg, #202020)",
border: "1px solid var(--border-color, #555)",
borderRadius: "6px",
padding: "8px",
lineHeight: "1.45",
whiteSpace: "pre-wrap",
});
widget = node.addDOMWidget(OUTPUT_NAME, "STRING", container, {
serialize: false,
hideOnZoom: false,
});
widget.serialize = false;
widget.inputEl = output;
widget.statusEl = status;
if (widget) {
return widget;
}
const container = document.createElement("div");
const header = document.createElement("div");
const status = document.createElement("span");
const meta = document.createElement("span");
const actions = document.createElement("div");
const output = document.createElement("textarea");
status.textContent = "Ready";
meta.textContent = textMetrics("");
output.readOnly = true;
output.wrap = "soft";
output.spellcheck = false;
output.setAttribute("aria-label", "VLM text output");
const copy = makeButton("Copy", "Copy complete text", async () => {
const previous = copy.textContent;
try {
await navigator.clipboard.writeText(output.value);
copy.textContent = "Copied";
} catch {
copy.textContent = "Copy failed";
}
window.setTimeout(() => {
copy.textContent = previous;
}, 1200);
});
const download = makeButton("Save", "Download output as a UTF-8 text file", () => {
const blob = new Blob([output.value], {
type: "text/plain;charset=utf-8",
});
const url = URL.createObjectURL(blob);
const anchor = document.createElement("a");
anchor.href = url;
anchor.download = `vlm-output-${new Date().toISOString().replaceAll(":", "-")}.txt`;
anchor.click();
URL.revokeObjectURL(url);
});
const wrap = makeButton("Wrap: on", "Toggle long-line wrapping", () => {
const enabled = output.wrap !== "off";
output.wrap = enabled ? "off" : "soft";
output.style.whiteSpace = enabled ? "pre" : "pre-wrap";
output.style.overflowX = enabled ? "auto" : "hidden";
wrap.textContent = enabled ? "Wrap: off" : "Wrap: on";
});
const follow = makeButton("Follow: on", "Follow streaming output", () => {
widget.followOutput = !widget.followOutput;
follow.textContent = widget.followOutput ? "Follow: on" : "Follow: off";
});
actions.append(wrap, follow, copy, download);
header.append(status, meta, actions);
container.append(header, output);
Object.assign(container.style, {
display: "flex",
flexDirection: "column",
width: "100%",
height: "100%",
minHeight: "190px",
gap: "6px",
});
Object.assign(header.style, {
display: "grid",
gridTemplateColumns: "auto minmax(0, 1fr) auto",
alignItems: "center",
gap: "9px",
color: "var(--descrip-text, #aaa)",
fontSize: "11px",
});
Object.assign(meta.style, {
overflow: "hidden",
textOverflow: "ellipsis",
whiteSpace: "nowrap",
});
Object.assign(actions.style, {
display: "flex",
gap: "4px",
justifyContent: "flex-end",
});
Object.assign(output.style, {
width: "100%",
flex: "1",
minHeight: "160px",
resize: "vertical",
boxSizing: "border-box",
color: "var(--input-text, #ddd)",
background: "var(--comfy-input-bg, #202020)",
border: "1px solid var(--border-color, #555)",
borderRadius: "6px",
padding: "9px",
lineHeight: "1.45",
whiteSpace: "pre-wrap",
overflowWrap: "anywhere",
tabSize: "4",
});
widget = node.addDOMWidget(OUTPUT_NAME, "STRING", container, {
serialize: false,
hideOnZoom: false,
});
widget.serialize = false;
widget.inputEl = output;
widget.statusEl = status;
widget.metaEl = meta;
widget.followOutput = true;
return widget;
}
@@ -88,8 +156,10 @@ function setOutput(node, text, state = "Complete") {
const value = Array.isArray(text) ? text.join("\n\n") : String(text ?? "");
widget.value = value;
widget.inputEl.value = value;
if (widget.statusEl) {
widget.statusEl.textContent = state;
widget.statusEl.textContent = state;
widget.metaEl.textContent = textMetrics(value);
if (widget.followOutput && state === "Streaming…") {
widget.inputEl.scrollTop = widget.inputEl.scrollHeight;
}
node.setDirtyCanvas?.(true, true);
}
@@ -104,18 +174,38 @@ function findNode(graph, id) {
?? null;
}
function linkFor(graph, linkId) {
return graph?.links?.get?.(linkId)
?? graph?._links?.get?.(linkId)
?? null;
}
function isReroute(node) {
return String(node?.type ?? "").toLowerCase().includes("reroute");
}
function connectedViewTextNodes(source) {
if (!source?.graph) {
return [];
}
const found = new Set();
for (const output of source.outputs ?? []) {
for (const linkId of output.links ?? []) {
const link = source.graph.links?.get?.(linkId)
?? source.graph._links?.get?.(linkId);
const target = findNode(source.graph, link?.target_id);
if (target?.type === VIEW_TEXT_NODE) {
found.add(target);
const visited = new Set([source.id]);
const queue = [source];
while (queue.length) {
const current = queue.shift();
for (const output of current.outputs ?? []) {
for (const linkId of output.links ?? []) {
const link = linkFor(source.graph, linkId);
const target = findNode(source.graph, link?.target_id);
if (!target || visited.has(target.id)) {
continue;
}
visited.add(target.id);
if (target.type === VIEW_TEXT_NODE) {
found.add(target);
} else if (isReroute(target)) {
queue.push(target);
}
}
}
}
@@ -131,7 +221,7 @@ function updateFromProgress({ nodeId, text }) {
setOutput(source, text, "Streaming…");
return;
}
if (source.type !== MODERN_VLM_NODE) {
if (!STREAMING_SOURCE_NODES.has(source.type)) {
return;
}
for (const target of connectedViewTextNodes(source)) {
@@ -154,6 +244,12 @@ app.registerExtension({
nodeType.prototype.onNodeCreated = function (...args) {
const result = onCreated?.apply(this, args);
ensureOutputWidget(this);
if (Array.isArray(this.size)) {
this.setSize?.([
Math.max(this.size[0], 430),
Math.max(this.size[1], 290),
]);
}
return result;
};
const onExecuted = nodeType.prototype.onExecuted;