Compare commits
23
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3f9612774e | ||
|
|
7a74f5a079 | ||
|
|
67344abe6a | ||
|
|
e79f316908 | ||
|
|
f4bc8b9eef | ||
|
|
5779c50b20 | ||
|
|
48101541a3 | ||
|
|
fcfdf7b210 | ||
|
|
56cdd25aa9 | ||
|
|
9aeca11c35 | ||
|
|
79a929c1ca | ||
|
|
e2b20cde13 | ||
|
|
04275b57cb | ||
|
|
2e41de6ac2 | ||
|
|
8b4226474a | ||
|
|
8a432184d3 | ||
|
|
da5f4d5787 | ||
|
|
45d21d0642 | ||
|
|
8e55c81b34 | ||
|
|
f06a2a3e6c | ||
|
|
c13ee2364e | ||
|
|
b8ae298abf | ||
|
|
bb7d51777c |
@@ -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
|
||||
@@ -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.
|
||||
@@ -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
|
||||
@@ -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"
|
||||
-->
|
||||
@@ -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: ["*"]
|
||||
@@ -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
@@ -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
|
||||
@@ -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"
|
||||
@@ -137,6 +137,38 @@ Authoritative references:
|
||||
- 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
|
||||
@@ -173,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
@@ -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.
|
||||
@@ -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
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# ComfyUI VLM Nodes
|
||||
|
||||
Production-oriented vision-language, structured prompting, audio, and utility
|
||||
nodes for ComfyUI. Version 3.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
|
||||
@@ -103,9 +103,10 @@ useful model capabilities:
|
||||
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.
|
||||
- **Legacy Molmo, Kosmos-2, UForm, MCLLaVA, and MiniCPM-V 2.6 GGUF**, plus
|
||||
maintained JoyTag.
|
||||
@@ -339,11 +340,231 @@ API-format examples are in [`examples/vision`](examples/vision):
|
||||
|
||||
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:
|
||||
@@ -358,6 +579,35 @@ 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
|
||||
@@ -397,6 +647,12 @@ 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:
|
||||
|
||||
@@ -447,10 +703,13 @@ Gemma 3 and PaLI-Gemma require accepting their model licenses on Hugging Face.
|
||||
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. Its random IPC secret is not placed
|
||||
on the process command line.
|
||||
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.
|
||||
@@ -594,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>
|
||||
|
||||
+36
@@ -39,7 +39,11 @@ restart ComfyUI:
|
||||
| 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:
|
||||
@@ -59,6 +63,38 @@ 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.
|
||||
|
||||
@@ -18,6 +18,7 @@ node_list = [
|
||||
"llavaloader",
|
||||
"mcllava",
|
||||
"minicpm",
|
||||
"minimax_music",
|
||||
"modern_vlm",
|
||||
"molmo",
|
||||
"moondream31",
|
||||
@@ -26,6 +27,7 @@ node_list = [
|
||||
"paligemma",
|
||||
"playmusic",
|
||||
"qwen2vl",
|
||||
"robotics",
|
||||
"sam2",
|
||||
"sam3_adapter",
|
||||
"simpletext",
|
||||
|
||||
@@ -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)
|
||||
@@ -0,0 +1 @@
|
||||
"""Runnable, dependency-isolated robotics policy bridge examples."""
|
||||
@@ -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
|
||||
]
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -15,7 +15,6 @@ from typing import Any
|
||||
import torch
|
||||
import torch.nn.functional as functional
|
||||
|
||||
|
||||
RESIZE_QUALITY = (
|
||||
"Fast (area)",
|
||||
"Quality (bicubic)",
|
||||
|
||||
+1
-2
@@ -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
@@ -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
@@ -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
|
||||
|
||||
+3
-2
@@ -9,13 +9,14 @@ Custom endpoints can read only ``CUSTOM_API_KEY``.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ipaddress
|
||||
import io
|
||||
import ipaddress
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, replace
|
||||
from typing import Any, Callable
|
||||
from typing import Any
|
||||
from urllib.parse import quote, quote_plus, urlsplit, urlunsplit
|
||||
|
||||
import torch
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
+3
-2
@@ -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,8 +26,8 @@ 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,
|
||||
|
||||
+1
-1
@@ -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,
|
||||
|
||||
+68
-31
@@ -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)
|
||||
|
||||
+11
-2
@@ -214,6 +214,12 @@ def _worker_environment(
|
||||
"HTTP_PROXY",
|
||||
"HTTPS_PROXY",
|
||||
"NO_PROXY",
|
||||
# ComfyUI may select cudaMallocAsync for its own allocator. Photon
|
||||
# captures CUDA graphs in a separate process, where inheriting that
|
||||
# override can abort during warmup on an uncaptured/captured free.
|
||||
# Let the sidecar use PyTorch's native allocator instead.
|
||||
"PYTORCH_ALLOC_CONF",
|
||||
"PYTORCH_CUDA_ALLOC_CONF",
|
||||
}
|
||||
secret_name = re.compile(
|
||||
r"(?:API[_-]?KEY|AUTHORIZATION|CREDENTIAL|PASSWORD|SECRET|TOKEN)",
|
||||
@@ -306,8 +312,11 @@ class Moondream31Model:
|
||||
)
|
||||
if "cudaLibraryLoadData" in tail:
|
||||
detail += (
|
||||
" The installed Photon CUDA kernel requires a newer compatible "
|
||||
"NVIDIA driver/runtime combination than this machine exposes."
|
||||
" Photon's CUDA 12 kernels require libcudart 12.9 or newer; "
|
||||
"CUDA runtime 12.6 does not export cudaLibraryLoadData. Re-run "
|
||||
"the README isolated-runtime install command so "
|
||||
"requirements-moondream31.txt upgrades only this sidecar to "
|
||||
"nvidia-cuda-runtime-cu12 12.9.79."
|
||||
)
|
||||
if tail:
|
||||
detail += f"\n\nSanitized worker log tail:\n{tail}"
|
||||
|
||||
@@ -25,6 +25,38 @@ 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.
|
||||
|
||||
@@ -283,6 +315,7 @@ def main() -> int:
|
||||
|
||||
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,
|
||||
@@ -303,6 +336,7 @@ def main() -> int:
|
||||
"status": "ready",
|
||||
"moondream_version": package_version,
|
||||
"compatibility_registration": compatibility_registration,
|
||||
"telemetry_disabled": telemetry_disabled,
|
||||
"skills": sorted(supported_skills),
|
||||
"pid": os.getpid(),
|
||||
}
|
||||
|
||||
+1
-2
@@ -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",
|
||||
|
||||
+1
-1
@@ -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.
|
||||
"""
|
||||
"""
|
||||
|
||||
+1
-2
@@ -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",
|
||||
|
||||
+2457
File diff suppressed because it is too large
Load Diff
+11
-3
@@ -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
|
||||
@@ -680,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": (
|
||||
@@ -697,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
@@ -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."
|
||||
|
||||
@@ -14,7 +14,6 @@ import re
|
||||
import unicodedata
|
||||
from typing import Any
|
||||
|
||||
|
||||
TEXT_CATEGORY = "VLM Nodes/Text"
|
||||
CREATE_CATEGORY = f"{TEXT_CATEGORY}/Create"
|
||||
TRANSFORM_CATEGORY = f"{TEXT_CATEGORY}/Transform"
|
||||
|
||||
+2
-2
@@ -9,7 +9,7 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
from typing import Any, Literal, Optional
|
||||
from typing import Any
|
||||
|
||||
import folder_paths
|
||||
import torch
|
||||
@@ -68,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):
|
||||
|
||||
+1
-1
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
+46
-2
@@ -1,10 +1,10 @@
|
||||
[project]
|
||||
name = "comfyui_vlm_nodes"
|
||||
version = "3.3.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",
|
||||
@@ -14,6 +14,7 @@ dependencies = [
|
||||
"huggingface-hub>=1.5,<2",
|
||||
"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",
|
||||
@@ -46,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"
|
||||
@@ -60,6 +100,7 @@ Icon = ""
|
||||
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",
|
||||
@@ -76,6 +117,9 @@ comfyui_vlm_nodes = [
|
||||
"*.json",
|
||||
"SECURITY.md",
|
||||
"examples/*.json",
|
||||
"examples/robotics/*.py",
|
||||
"examples/robotics/*.md",
|
||||
"examples/robotics/*.json",
|
||||
"examples/vision/*.json",
|
||||
"requirements*.txt",
|
||||
]
|
||||
|
||||
@@ -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
|
||||
@@ -5,3 +5,7 @@ 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")
|
||||
|
||||
@@ -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
|
||||
@@ -9,6 +9,7 @@ einops>=0.8,<1
|
||||
huggingface-hub>=1.5,<2
|
||||
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
|
||||
|
||||
@@ -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)
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
@@ -12,7 +12,9 @@ import json
|
||||
import time
|
||||
|
||||
import torch
|
||||
from _bootstrap import bootstrap
|
||||
|
||||
bootstrap()
|
||||
|
||||
BACKENDS = (
|
||||
"florence-base",
|
||||
|
||||
@@ -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
|
||||
@@ -6,7 +6,6 @@ from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from ComfyUI_VLM_nodes.nodes import hosted_api
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -1,3 +1,4 @@
|
||||
import asyncio
|
||||
import importlib
|
||||
import inspect
|
||||
import json
|
||||
@@ -257,6 +258,8 @@ def test_worker_auth_is_not_exposed_in_process_arguments_and_logs_are_redacted(
|
||||
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,
|
||||
@@ -267,6 +270,8 @@ def test_worker_auth_is_not_exposed_in_process_arguments_and_logs_are_redacted(
|
||||
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(
|
||||
@@ -325,3 +330,32 @@ def test_worker_registers_official_31_id_only_when_upstream_is_missing(
|
||||
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()
|
||||
|
||||
@@ -120,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)
|
||||
|
||||
@@ -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
|
||||
@@ -1,9 +1,8 @@
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
import ComfyUI_VLM_nodes as package
|
||||
import pytest
|
||||
from ComfyUI_VLM_nodes.nodes import simpletext
|
||||
|
||||
|
||||
|
||||
@@ -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))
|
||||
@@ -5,7 +5,6 @@ from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from ComfyUI_VLM_nodes.nodes.video_intelligence import (
|
||||
NODE_CLASS_MAPPINGS,
|
||||
VLMAdaptiveFrameSampler,
|
||||
|
||||
Reference in New Issue
Block a user