Compare commits
27
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 | ||
|
|
102f1662ac | ||
|
|
44fefcb57a | ||
|
|
505b324f66 | ||
|
|
39fc116341 |
@@ -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,9 +72,25 @@ 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
|
||||
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
|
||||
run: python -m build
|
||||
|
||||
+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"
|
||||
@@ -26,6 +26,102 @@ optional optimization, not an import requirement. DirectML/private-use devices
|
||||
receive a safe FP32 fallback, but are best-effort because current ComfyUI itself
|
||||
does not treat DirectML as a primary performance backend.
|
||||
|
||||
## Detection and segmentation backends
|
||||
|
||||
The structured vision nodes do not install a second PyTorch build. Grounding
|
||||
DINO, OWLv2, OmDet Turbo, Florence-2, and SAM2.1 use the device selected by
|
||||
ComfyUI and participate in its model loading/offloading lifecycle. The core
|
||||
SAM3.1 adapter performs schema validation and report generation on the compact
|
||||
core payload; ComfyUI itself owns SAM3 inference and mask packing.
|
||||
|
||||
| Backend | Detection / Florence | SAM2.1 video | Comfy core SAM3.1 | Practical limitation |
|
||||
| --- | --- | --- | --- | --- |
|
||||
| NVIDIA CUDA | Managed BF16 when supported, otherwise FP16 | Preferred accelerated path; CPU state/storage is the default | Supported when the installed ComfyUI version recognizes the checkpoint | Resolution, frame count, and object count still dominate VRAM/RAM |
|
||||
| AMD ROCm on Linux | Uses PyTorch's `cuda` device and BF16/FP16 capability checks | Same managed path; keep inference state on CPU unless measured otherwise | Follows ComfyUI core ROCm support | Individual Transformers kernels may fall back or differ in performance |
|
||||
| AMD ROCm on Windows | Uses the device exposed by the selected ComfyUI PyTorch build | Same API contract | Follows that ComfyUI build | Treat as hardware-validation pending, not equivalent to a Linux ROCm pass |
|
||||
| Apple Metal / MPS | FP16, or BF16 only when macOS/PyTorch report support | Supported contract with CPU video storage; use Tiny and short slices first | Follows ComfyUI core MPS support | Unified memory is shared with the OS; unsupported operators may fall back to CPU |
|
||||
| Intel XPU | BF16/FP16 capability-selected managed path | Supported contract; use CPU state for portability | Follows ComfyUI core XPU support | Model-specific operator coverage and real throughput require hardware validation |
|
||||
| CPU | FP32 portable path | Functionally supported but slow; use Tiny, low resolution, and short slices | Adapter/report works; core SAM3 inference is memory intensive | No half-precision speed assumption and no accelerator kernel |
|
||||
|
||||
`precision=auto` is the safe default for open-vocabulary detection and SAM2.1.
|
||||
Explicit BF16 silently falls back to FP16 or FP32 when the selected backend
|
||||
cannot execute BF16. This is a portability fallback, not proof that every
|
||||
model family has been run on every vendor device. See
|
||||
[MODEL_VALIDATION.md](MODEL_VALIDATION.md) for real-hardware evidence.
|
||||
|
||||
### Moondream 3 / 3.1 Photon
|
||||
|
||||
Moondream Photon is deliberately isolated from ComfyUI's main Python environment
|
||||
because `moondream==1.3.0` requires Pillow 10 while current ComfyUI uses a
|
||||
newer Pillow. Its worker cache, virtual environment, and logs live under
|
||||
`models/LLavacheckpoints/moondream31-runtime`; it never replaces ComfyUI's
|
||||
PyTorch or Pillow.
|
||||
|
||||
| Platform | Official local Photon support | This integration |
|
||||
| --- | --- | --- |
|
||||
| Linux/WSL + NVIDIA Ampere or newer | Supported | 3.1 query/caption/detection/pointing; 3 Preview SVG segmentation |
|
||||
| Windows + NVIDIA Ampere or newer | Supported | Same isolated worker contract |
|
||||
| Apple Silicon macOS 13+ | Supported with MPS | Same contract; use a conservative KV-cache profile on low-memory systems |
|
||||
| AMD ROCm, Intel GPU, CPU | Not currently provided upstream | Node stays importable and fails before model work with an actionable support message |
|
||||
|
||||
The final Moondream 3.1 model card lists query, caption, detect, and point; it
|
||||
does not list segment. Native SVG segment uses `moondream3-preview`, and the
|
||||
loader rejects a 3.1/segment mismatch before inference.
|
||||
|
||||
`max_batch_size` controls Photon's scheduler capacity. The detection, point,
|
||||
and preview-segmentation nodes issue `parallel_requests` frame requests concurrently,
|
||||
allowing Photon to build GPU batches. `frame_stride` bounds work for high-frame
|
||||
rate sources. Performance JSON records warm worker time, end-to-end time,
|
||||
processed/skipped frames, worker/sustained FPS, target sampled FPS, and
|
||||
real-time factor; it is a measurement from the current run, not a universal
|
||||
benchmark claim.
|
||||
|
||||
### Video memory and chunking
|
||||
|
||||
- Core `Video Slice` should bound work before `GetVideoComponents` materializes
|
||||
frames. Scale the resulting `IMAGE` batch before running detection or
|
||||
segmentation.
|
||||
- Open-vocabulary detection runs frame by frame. SAM2.1 keeps source frames on
|
||||
CPU, defaults its inference state to CPU, and caches at most one vision
|
||||
feature in the video session.
|
||||
- SAM2.1 output masks and previews are CPU tensors. Core SAM3 keeps its track
|
||||
masks bit-packed; `VLMSAM3TrackAdapter` does not unpack the complete volume.
|
||||
- `unload_after=true` releases the node's owned detector/SAM2 model after a
|
||||
run. Leave it false for repeated work with one model; set it true before a
|
||||
different large family must load on a constrained accelerator.
|
||||
- Each slice or queue run starts a new propagation/tracking session. Carrying
|
||||
an ID across independent chunks requires an explicit application-level
|
||||
overlap/reconciliation step; the nodes never claim cross-run identity.
|
||||
|
||||
### Model licenses and access
|
||||
|
||||
Model licenses are independent from this repository's code license. Check the
|
||||
model card before redistributing weights or outputs.
|
||||
|
||||
- The `facebook/sam2.1-hiera-*` Transformers checkpoints are published under
|
||||
Apache-2.0.
|
||||
- Meta SAM3 uses the SAM License. The upstream `facebook/sam3` repository is
|
||||
access-gated and asks the Hugging Face account holder to accept its terms and
|
||||
share the requested contact information.
|
||||
- ComfyUI's `Comfy-Org/sam3.1` checkpoint is marked `sam-license`; the example
|
||||
expects `sam3.1_multiplex_fp16.safetensors` under
|
||||
`ComfyUI/models/checkpoints`.
|
||||
- `HF_TOKEN` is used when Hugging Face requires authenticated access. Tokens
|
||||
must be supplied by the environment and must not be embedded in workflows.
|
||||
- Moondream 3.1 uses the Moondream Model License 1.0. The Loader requires an
|
||||
explicit workflow acknowledgement. The license permits local product use
|
||||
but restricts offering general-purpose hosted Moondream access; review the
|
||||
current upstream terms for the intended deployment.
|
||||
|
||||
Authoritative references:
|
||||
|
||||
- [Meta SAM3 model and access terms](https://huggingface.co/facebook/sam3)
|
||||
- [Meta SAM3 license](https://huggingface.co/facebook/sam3/blob/main/LICENSE)
|
||||
- [ComfyUI SAM3.1 checkpoint](https://huggingface.co/Comfy-Org/sam3.1)
|
||||
- [SAM2.1 Hiera Tiny model card](https://huggingface.co/facebook/sam2.1-hiera-tiny)
|
||||
- [Moondream 3.1 model card](https://huggingface.co/moondream/moondream3.1-9B-A2B)
|
||||
- [Moondream Model License 1.0](https://moondream.ai/licenses/model/1.0)
|
||||
|
||||
## Dependency behavior
|
||||
|
||||
- Python 3.10 through 3.13 is covered by CI.
|
||||
@@ -36,9 +132,43 @@ does not treat DirectML as a primary performance backend.
|
||||
from blocking the whole node pack.
|
||||
- `requirements-quantization.txt` is available for an explicit quantization
|
||||
install or source-build environment.
|
||||
- `requirements-moondream31.txt` belongs only in the isolated Photon sidecar;
|
||||
installing it into ComfyUI's environment would create a Pillow conflict.
|
||||
- Model downloads, imports, and package compilation never occur during node
|
||||
discovery.
|
||||
|
||||
## Robotics / VLA policy compatibility
|
||||
|
||||
ComfyUI's robotics schemas, safety gate, trajectory tools, and universal HTTP
|
||||
client run wherever this node pack runs. Policy runtime compatibility is
|
||||
separate:
|
||||
|
||||
| Policy route | ComfyUI client | Policy environment | Practical boundary |
|
||||
| --- | --- | --- | --- |
|
||||
| Universal VLA HTTP | Windows, Linux, macOS; CUDA, ROCm, Metal, XPU, CPU | Any host that implements `comfyui-vla-http-v1` | Loopback HTTP or trusted HTTPS; no pickle |
|
||||
| LeRobot sidecar | Same universal client | Current LeRobot supports Linux, Windows, and macOS; individual policy extras/operators vary | Python/PyTorch live outside ComfyUI; fine-tuned checkpoint required for the target embodiment |
|
||||
| openpi WebSocket | Lightweight optional client on every ComfyUI platform | Upstream currently tests Ubuntu 22.04 + NVIDIA, inference above 8 GB VRAM | Use WSL/Docker/Linux server; remote transport must be WSS |
|
||||
| Isaac-GR00T N1.7 ZMQ | Lightweight optional client on every ComfyUI platform | NVIDIA CUDA/Jetson Linux according to upstream deployment matrix | ZMQ has no transport encryption; use a private network/tunnel |
|
||||
| OpenVLA-OFT | Universal client with a project-specific bridge | Upstream PyTorch/CUDA environment | OFT is the preferred high-frequency multi-image OpenVLA route |
|
||||
| Octo | Universal client with a project-specific bridge | Isolated JAX environment | Kept as a lightweight research baseline, not the default maintained runtime |
|
||||
|
||||
Install only the native client protocols into ComfyUI:
|
||||
|
||||
```bash
|
||||
python -m pip install -r requirements-robotics-client.txt
|
||||
```
|
||||
|
||||
Do not install `lerobot[all]`, openpi, Isaac-GR00T, OpenVLA, or JAX into
|
||||
ComfyUI's Python. The included LeRobot HTTP sidecar belongs in its own
|
||||
environment and optionally moves its owned policy to CPU after an idle
|
||||
interval. It does not flush ComfyUI's accelerator cache.
|
||||
|
||||
An embodiment profile is a workflow contract, not a hardware certification.
|
||||
The supplied profiles are visibly labeled templates. Before real deployment,
|
||||
replace action bounds/deltas with the trained dataset's semantics and the
|
||||
manufacturer/controller limits. ComfyUI never opens ROS, serial, CAN, or robot
|
||||
SDK transports.
|
||||
|
||||
Install manually:
|
||||
|
||||
```bash
|
||||
@@ -75,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
|
||||
@@ -50,6 +115,65 @@ SmolVLM2 256M/500M/2.2B, LFM2.5 VL 450M/1.6B, InternVL 3.5 1B/2B, and Granite
|
||||
Vision 3.3 2B/4.1 4B. Gemma 3 4B is the sixteenth entry and correctly requires
|
||||
license acceptance plus `HF_TOKEN`.
|
||||
|
||||
## Structured vision validation
|
||||
|
||||
The versioned detection/track/point/event payloads, geometry and mask
|
||||
conversion, strict spatial parser, Grounding-family adapters, SAM2.1 session
|
||||
plumbing, SAM3 bit-packed payload adapter, and ByteTrack-style association pass
|
||||
the local WSL contract suite. Those tests validate schemas, shapes, output
|
||||
ordering, bounds, timestamps, deterministic IDs, and error handling.
|
||||
|
||||
Representative real-weight checks were then submitted through ComfyUI's local
|
||||
`POST /prompt` API and verified from `/history/{prompt_id}`. The test machine
|
||||
used ComfyUI 0.28.0, Python 3.12.12, PyTorch 2.13.0+cu126, Transformers 5.14.1,
|
||||
and an NVIDIA RTX 3090. Input media, checkpoints, model caches, ComfyUI, and
|
||||
this checkout all remained on the D drive under WSL.
|
||||
|
||||
| Family | Representative checkpoint policy | Real-weight status |
|
||||
| --- | --- | --- |
|
||||
| Grounding DINO | Tiny; Base uses the same loader/processor contract | **Passed**: FP16, four real 640x360 video frames in two-frame micro-batches; person and bird boxes/labels were visually checked, serialized, timestamped, and in bounds |
|
||||
| OWLv2 | Base Ensemble | Pending |
|
||||
| OmDet Turbo | Swin Tiny | Pending |
|
||||
| SAM2.1 video | Hiera Tiny; sibling sizes use the same session adapter | **Passed**: FP16, real 12-frame 640x360 clip at 24 FPS with CPU preprocessing/state. Grounding's core `BOUNDING_BOX` output connected directly: the forward union-only run kept one person ID on frames 0-11; a last-frame reverse run kept two IDs for 24 observations and emitted 24 frame-major object masks. All geometry was in bounds and first/last overlays and masks were visually checked |
|
||||
| Comfy core SAM3.1 | `sam3.1_multiplex_fp16.safetensors`, only after license/access is available | Pending |
|
||||
| SAM3 adapter/report | Synthetic core payload contract | Passed without weights; real core handoff pending |
|
||||
| ByteTrack-style tracker | Deterministic synthetic crossing, missed-frame, and expiry cases | Passed; no model weights exist |
|
||||
| Florence-2 multitask | Base FT; Large uses the same native Transformers contract | **Passed**: real object-detection API run produced bounded woman, face, and clothing boxes plus a visually checked overlay |
|
||||
|
||||
The SAM2 API check initially exposed a real session-lifecycle defect that unit
|
||||
fixtures did not: prompt insertion must be followed by inference on the seeded
|
||||
frame before propagation. The implementation now performs that seed pass and
|
||||
also propagates in reverse when `seed_frame` is greater than zero. Later live
|
||||
checks exercised nested multi-object core boxes, CPU preprocessing/state,
|
||||
union-only low-memory output, optional object-mask output, disabled preview
|
||||
rendering, reverse propagation, and `unload_after=true` for both models. The
|
||||
final unload run returned total reported GPU memory use to within 4 MiB of the
|
||||
pre-run `nvidia-smi` baseline.
|
||||
|
||||
Grounding DINO and SAM2 sibling sizes are catalog-available but were not
|
||||
downloaded or executed. OWLv2, OmDet Turbo, and gated SAM3 remain explicitly
|
||||
unverified; the UI never presents them as locally tested simply because their
|
||||
schemas import.
|
||||
|
||||
The acceptance run for each model family must record:
|
||||
|
||||
1. Exact checkpoint revision, ComfyUI/Python/PyTorch/Transformers versions,
|
||||
device, dtype, peak accelerator allocation, and wall time.
|
||||
2. A real image or short bounded video with manually verified boxes, labels,
|
||||
masks, timestamps, and stable IDs.
|
||||
3. The canonical JSON schema/version and every advertised output socket,
|
||||
including preview/report output through ComfyUI's local `/prompt` API.
|
||||
4. A second queue using the cached model, followed by an `unload_after=true`
|
||||
run where that option exists.
|
||||
5. Failure behavior for an absent checkpoint or gated access without exposing
|
||||
a token.
|
||||
|
||||
One checkpoint per distinct implementation family is enough for sibling model
|
||||
sizes that share the same code path. Validation prioritizes the smallest useful
|
||||
checkpoint and will not download or execute a 30B model. A larger variant is
|
||||
tested only when it has a different loader, processor, postprocessor, or
|
||||
quantization path.
|
||||
|
||||
## Not marked passed
|
||||
|
||||
- Qwen 3 VL 30B-A3B: weights are available locally, but inference validation
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# ComfyUI VLM Nodes
|
||||
|
||||
Production-oriented vision-language, structured prompting, audio, and utility
|
||||
nodes for ComfyUI. Version 2.3 supports ComfyUI's selected NVIDIA CUDA, AMD
|
||||
nodes for ComfyUI. Version 3.4 supports ComfyUI's selected NVIDIA CUDA, AMD
|
||||
ROCm, Apple Metal, Intel XPU, and CPU device without replacing its PyTorch
|
||||
build. It removes startup installers and global accelerator cache flushes,
|
||||
adds real image/video batches and live token streaming, and uses ComfyUI model
|
||||
@@ -9,19 +9,40 @@ residency and offloading.
|
||||
|
||||
## Modern model coverage
|
||||
|
||||
The **Modern VLM** node provides one stable interface for:
|
||||
The **Modern VLM** node provides one stable interface with a deliberately
|
||||
small, 12-choice production picker:
|
||||
|
||||
- Qwen 3.5 0.8B, 2B, 4B, 9B, 27B, and 35B-A3B
|
||||
- Qwen 3.6 27B
|
||||
- Qwen 3 VL 2B, 4B, 8B, and 30B-A3B Instruct
|
||||
- Qwen 2.5 VL 3B and 7B for existing workflows
|
||||
- Gemma 3 4B, 12B, and 27B IT
|
||||
- SmolVLM2 256M, 500M, and 2.2B video models
|
||||
- Liquid LFM2.5-VL 450M and 1.6B edge models
|
||||
- InternVL 3.5 1B and 2B standard Hugging Face checkpoints
|
||||
- Granite Vision 3.3 2B and 4.1 4B for documents, charts, and OCR
|
||||
- Qwen 3.5 0.8B and 4B
|
||||
- Qwen 3 VL 2B, 4B, and 8B Instruct
|
||||
- SmolVLM2 500M and 2.2B Video
|
||||
- Liquid LFM2.5-VL 450M
|
||||
- InternVL 3.5 1B
|
||||
- Granite Vision 4.1 4B
|
||||
- Gemma 3 4B IT
|
||||
- a compatible custom Hugging Face image-to-text repository
|
||||
|
||||
The separate **[Legacy] Modern VLM Compatibility** node contains redundant,
|
||||
superseded, experimental, and very large tiers:
|
||||
|
||||
- Qwen 3.5 2B, 9B, 27B, and 35B-A3B
|
||||
- Qwen 3.6 27B
|
||||
- Qwen 3 VL 30B-A3B Instruct
|
||||
- Qwen 2.5 VL 3B and 7B for existing workflows
|
||||
- Gemma 3 12B and 27B IT
|
||||
- SmolVLM2 256M Video
|
||||
- Liquid LFM2.5-VL 1.6B
|
||||
- InternVL 3.5 2B
|
||||
- Granite Vision 3.3 2B
|
||||
|
||||
Previously saved `ModernVLM` workflows remain valid even when their selected
|
||||
model moved to Legacy. The server accepts every known catalog value for
|
||||
backward compatibility; only the visible new-workflow picker is curated.
|
||||
Dedicated Molmo, PaLI-Gemma, Qwen2-VL, MiniCPM-V, Kosmos-2, MC-LLaVA, UForm,
|
||||
and script-style MoonDream nodes are also collected under
|
||||
`VLM Nodes/Legacy/Model Loaders`. Maintained creator-facing Florence-2,
|
||||
Moondream2, JoyTag, llama.cpp/GGUF, detection, segmentation, tracking, API,
|
||||
and video-intelligence nodes stay in their functional categories.
|
||||
|
||||
Sixteen curated sub-4B/low-VRAM choices are marked internally as the
|
||||
small-and-fast tier. The default is Qwen 3 VL 2B: it is much quicker to load
|
||||
than larger checkpoints while retaining broad image and video understanding.
|
||||
@@ -41,21 +62,509 @@ when ComfyUI rehydrates workflow output history. Disable `stream_output` for
|
||||
API-only or headless runs that do not need incremental UI updates. Streaming is
|
||||
best-effort and never changes the final `STRING` output or makes inference fail.
|
||||
|
||||
## Text workflow toolkit
|
||||
|
||||
The original `SimpleText`, `JsonToText`, and `ViewText` node IDs and their
|
||||
first `STRING` outputs remain stable for saved workflows. They now live in
|
||||
organized `VLM Nodes/Text` subcategories and expose descriptive names, search
|
||||
aliases, tooltips, appended metrics, and strict error messages:
|
||||
|
||||
| Node | Purpose |
|
||||
| --- | --- |
|
||||
| `Text` (`SimpleText`) | Multiline/dynamic prompt source with optional edge/newline normalization and character, word, and line outputs |
|
||||
| `View Text (Streaming)` | Read-only live output with counts, copy, UTF-8 download, line wrapping, stream following, reroute traversal, and history rehydration |
|
||||
| `JSON to Text` | Plain or fenced JSON parsing with readable, values-only, key/value, pretty, and compact render modes |
|
||||
| `Text Join` | Join up to eight prompt/context values with empty-value removal and stable deduplication |
|
||||
| `Text Template` | Safe named placeholders from a JSON object plus four convenient live text sockets, with explicit missing-key policy |
|
||||
| `Text Clean` | Unicode NFC/NFKC, newline/whitespace cleanup, enclosing Markdown-fence removal, line deduplication, and deterministic length caps |
|
||||
| `Text Replace` | Literal or regex substitution with case, count, and missing-pattern controls |
|
||||
| `JSON Extract` | JSONPath-lite (`$.items[0]`) and RFC 6901 JSON Pointer extraction from plain or fenced model responses |
|
||||
| `Text Split / Batch` | Lines, paragraphs, delimiters, regex, CSV, or JSON arrays converted to a real mapped Comfy `STRING` list |
|
||||
| `Text Inspector` | Pass-through text plus characters, UTF-8 bytes, words, lines, rough token budget, SHA-256, and JSON metadata |
|
||||
|
||||
The JSON utilities never evaluate code, follow references, access files, or
|
||||
make network requests. Template fields are direct names rather than Python
|
||||
attribute/index expressions. `approx_tokens` is deliberately labeled as a
|
||||
rough UTF-8 budget estimate; use the target model tokenizer when exact billing
|
||||
or context accounting matters.
|
||||
|
||||
Specialized nodes remain available where a generic chat node would discard
|
||||
useful model capabilities:
|
||||
|
||||
- **Moondream 3.1 9B-A2B**: official 2B-active Photon runtime with query,
|
||||
caption, and high-throughput image/video detection and pointing.
|
||||
- **Moondream 3 Preview segment**: native SVG segmentation through the same
|
||||
isolated Photon loader. The SVG is preserved and also converted into antialiased
|
||||
`MASK`, black/white previews, foreground cutouts, overlays, polygons,
|
||||
canonical `VLM_DETECTIONS`, and core bounding boxes. Detection/pointing
|
||||
submit frames concurrently so Photon can dynamically batch them; every run
|
||||
reports measured worker FPS, end-to-end FPS, and real-time factor.
|
||||
- **Florence-2**: captioning, OCR, detection, region captioning, and referring
|
||||
expression segmentation, with structured JSON, mask, and overlay outputs.
|
||||
- **PaLI-Gemma**: caption/VQA plus the official 16-token VQ-VAE segmentation
|
||||
decoder; segmentation tokens are no longer misinterpreted as polygon points.
|
||||
- **Moondream2**: pinned query API with explicit decoding controls. Its current
|
||||
checkpoint is not marked passed on the tested Torch/Transformers stack; use a
|
||||
small Modern VLM preset for production.
|
||||
- **Moondream2**: pinned query API with explicit decoding controls. The official
|
||||
checkpoint is loaded through its native safetensors state dict, avoiding the
|
||||
silent empty-output regression in Transformers 5 while retaining ComfyUI
|
||||
managed loading and unloading.
|
||||
- **Qwen2-VL**: image batches and real video-frame batches.
|
||||
- **Molmo, Kosmos-2, UForm, MCLLaVA, JoyTag, and MiniCPM-V 2.6 GGUF**.
|
||||
- **Legacy Molmo, Kosmos-2, UForm, MCLLaVA, and MiniCPM-V 2.6 GGUF**, plus
|
||||
maintained JoyTag.
|
||||
- **llama.cpp LLaVA/GGUF**, structured prompt suggestions, OpenAI-compatible
|
||||
prompting, and AudioLDM2.
|
||||
|
||||
## Structured detection, segmentation, and tracking
|
||||
|
||||
The vision nodes use stable, typed sockets instead of passing model-specific
|
||||
lists between nodes:
|
||||
|
||||
| Socket | JSON schema | Purpose |
|
||||
| --- | --- | --- |
|
||||
| `VLM_DETECTIONS` | `comfyui-vlm/detections`, version 1 | Per-frame boxes, labels, scores, optional polygons/quads, and in-process masks |
|
||||
| `VLM_TRACKS` | `comfyui-vlm/tracks`, version 1 | Durable object IDs with ordered observations over time |
|
||||
| `VLM_POINTS` | `comfyui-vlm/points`, version 1 | Pixel-coordinate points, including detection centers |
|
||||
| `VLM_EVENTS` | `comfyui-vlm/events`, version 1 | Ordered temporal events for downstream video analysis |
|
||||
| `VLM_VIDEO_SELECTION` | `comfyui-vlm/video-selection`, version 1 | Exact mapping from sampled images to source frame indices and timestamps |
|
||||
| `VLM_SCENE_STATE` | `comfyui-vlm/scene-state`, version 1 | Compact persistent objects, motion, visibility, and validated events |
|
||||
|
||||
All spatial coordinates are source-image pixels. Bounding boxes are
|
||||
`[x1, y1, x2, y2]` with an exclusive right/bottom edge; polygons contain at
|
||||
least three points and quads exactly four. JSON roots contain `schema`,
|
||||
`version`, media dimensions/frame count/FPS, and their ordered records. Dense
|
||||
mask tensors remain in-process and are deliberately omitted from JSON so API
|
||||
results do not unexpectedly grow by hundreds of megabytes.
|
||||
|
||||
The utility layer converts without model-specific glue:
|
||||
|
||||
- `VLMStructuredSpatialParser` strictly parses pixel, normalized 0–1, or
|
||||
normalized 0–1000 JSON from any VLM into `VLM_DETECTIONS` and `VLM_POINTS`.
|
||||
`VLMSpatialPromptBuilder` creates the matching constrained prompt.
|
||||
- `VLMDetectionsToBoundingBoxes`, `VLMDetectionsToPoints`, and
|
||||
`VLMDetectionsToMasks` emit Comfy core boxes, center points, combined and
|
||||
individual binary masks, inverse masks, ready-to-preview black-and-white
|
||||
images, and stable-color instance maps. Polygon/quad masks are rasterized
|
||||
when present, otherwise the bounding box is used. Existing output indexes
|
||||
remain stable; the creator-facing mask images and instance map are appended.
|
||||
- `VLMFilterDetections`, `VLMSelectDetection`, `VLMCropDetections`, and
|
||||
`VLMRenderDetections` provide label/score/area/frame selection, padded crops,
|
||||
and deterministic overlays.
|
||||
- `VLMMaskProcessor` accepts any Comfy `MASK`, including SAM2/SAM3 masks, and
|
||||
returns a feathered matte, strict binary mask, inverse mask, and
|
||||
black-and-white image. Its grow/shrink and Gaussian feathering run in Torch
|
||||
without OpenCV or SciPy.
|
||||
- `VLMMaskComposite` applies still-image or video mask batches to a source and
|
||||
returns the replacement composite, isolated foreground, original
|
||||
background-only plate, and black-and-white mask image. A single mask or
|
||||
background broadcasts safely across a video batch.
|
||||
- `VLMDetectionsFromJSON` and `VLMDetectionsToJSON` are the explicit API and
|
||||
persistence boundary for the versioned detection schema.
|
||||
|
||||
### Universal VLM performance utilities
|
||||
|
||||
The performance nodes sit before any local or hosted VLM, so their savings do
|
||||
not depend on CUDA, ROCm, MPS, XPU, CPU, Transformers, llama.cpp, or Photon:
|
||||
|
||||
- `VLM Performance Profile` emits coherent `max_frames`, pixel budget,
|
||||
longest-edge, batch-size, and `unload_after` values. `Live / robotics`,
|
||||
`Fast video`, `Balanced`, `High detail`, and `Low VRAM handoff` are explicit
|
||||
starting points rather than hidden global flags.
|
||||
- `VLM Adaptive Frame Sampler` is the existing track-aware temporal gate. It
|
||||
combines uniform coverage, scene changes, motion, and optional track changes
|
||||
while preserving source frame indices and timestamps.
|
||||
- `VLM Image Pixel Budget` downsizes the selected analysis copy once, preserves
|
||||
aspect ratio, never upscales, and can align dimensions to 14/28-pixel VLM
|
||||
patches or 32-pixel detector backbones. Fast area and antialiased bicubic
|
||||
modes are available.
|
||||
|
||||
The recommended order is `Video Slice` → `VLM Adaptive Frame Sampler` →
|
||||
`VLM Image Pixel Budget` → any VLM. A model's own official processor still
|
||||
performs its required normalization/crop; the pixel-budget node simply prevents
|
||||
every downstream model from repeatedly receiving unnecessary source pixels.
|
||||
Local torch models remain registered with ComfyUI's smart model manager, while
|
||||
external allocators reserve space before loading and close only the handle they
|
||||
own.
|
||||
|
||||
On the real `vlm_api_people_birds.mp4` input in this repository's D-drive test
|
||||
environment, the utilities selected 10 of 60 1280×720 frames and resized them
|
||||
to 938×518 in about 0.44 seconds on a cold WSL run. That reduced the
|
||||
frame×pixel analysis workload by 11.38× before model inference. This is an
|
||||
input-work reduction measurement, not a claim that every model runs 11.38×
|
||||
faster; token generation and model-specific vision encoders still determine
|
||||
end-to-end speed.
|
||||
|
||||
### Adaptive video intelligence
|
||||
|
||||
The video-intelligence layer keeps generative VLM inference out of the
|
||||
per-frame loop:
|
||||
|
||||
- `VLMAdaptiveFrameSampler` combines scene-change, motion, track-change, and
|
||||
uniform-coverage signals. It always preserves the real source frame index
|
||||
and timestamp, enforces a frame budget, and returns selection/diagnostic
|
||||
JSON. `Uniform coverage`, motion, scene, and track-priority modes remain
|
||||
available for deterministic experiments.
|
||||
- `VLMVideoTemporalReasoner` is the one-node path. It adaptively samples the
|
||||
input, downsizes only the VLM analysis copy (448-pixel longest side by
|
||||
default), runs a recommended video-capable model, parses the result into
|
||||
validated `VLM_EVENTS`, and returns summary, events, selection, sampled
|
||||
previews, raw response, diagnostics, event JSON, and selection JSON.
|
||||
- `VLMVideoReasoningPrompt` and `VLMEventsFromVideoJSON` expose the same strict
|
||||
timestamp/evidence contract for custom local or hosted VLM workflows.
|
||||
- `VLMTrackAwareCrops` chooses representative observations for each durable
|
||||
track, adds configurable context, and letterboxes crops to one batch size.
|
||||
This lets a VLM label identities without rereading every full frame.
|
||||
- `VLMBuildSceneState` converts tracks plus optional events into a compact
|
||||
persistent world-state summary with first/last observation, current box,
|
||||
confidence, state, and pixel velocity.
|
||||
|
||||
Small VLMs commonly return evidence as positions in the supplied image batch
|
||||
even when asked for source indices. The parser accepts that form only when
|
||||
every value is an unambiguous valid supplied-image position, maps it back to
|
||||
the immutable source selection, and records the normalization mode. Arbitrary
|
||||
or unsupplied evidence frames, out-of-range timestamps, invalid confidence,
|
||||
duplicate evidence, malformed JSON, and non-finite values fail validation.
|
||||
|
||||
On the repository's real-data smoke test (RTX 3090, Qwen3-VL 2B, 157-frame
|
||||
896x448 H.264 clip), hybrid sampling selected 12 frames in 0.30 seconds,
|
||||
reduced temporal inputs by 92.36%, reduced analysis pixels by 75%, used
|
||||
4.24 GiB peak allocated VRAM in the standalone runner, and produced a valid
|
||||
timestamped result in 35.17 seconds. The equivalent live ComfyUI `/prompt`
|
||||
graph completed in 37.45 seconds. These are one-machine measurements, not
|
||||
portable performance guarantees.
|
||||
|
||||
### Open-vocabulary image and video detection
|
||||
|
||||
`VLMOpenVocabularyDetection` exposes one interface for:
|
||||
|
||||
- Grounding DINO Tiny and Base
|
||||
- OWLv2 Base Ensemble
|
||||
- OmDet Turbo Swin Tiny
|
||||
|
||||
It accepts a still image or an `IMAGE` batch of video frames and processes the
|
||||
batch frame by frame. Outputs, in socket order, are `detections`, `json`,
|
||||
`preview`, `box_mask`, and Comfy core `bounding_boxes`. Connect the FPS output
|
||||
of `GetVideoComponents` when the input is video so every timestamp is correct.
|
||||
For tracking-by-detection, run detection over the complete bounded batch and
|
||||
connect it to `VLMTrackDetections`.
|
||||
|
||||
`VLMTrackDetections` uses a ByteTrack-style two-stage high/low-confidence
|
||||
association, motion prediction, label-aware matching, and time-based expiry.
|
||||
IDs are durable within the supplied sequence and survive short missed
|
||||
detections when `emit_predictions` is enabled. Independent Comfy queue runs or
|
||||
independently sliced chunks are separate tracking sessions; they do not
|
||||
silently reuse IDs.
|
||||
|
||||
### SAM2.1 and Comfy core SAM3.1
|
||||
|
||||
`VLMSAM2VideoSegmentation` propagates first-frame detections, one core
|
||||
`BOUNDING_BOX`, or seed masks through an `IMAGE` batch using SAM2.1 Hiera Tiny,
|
||||
Small, Base+, or Large. It returns `VLM_TRACKS`, report JSON, per-frame union
|
||||
masks, frame-major individual object masks, and an overlay batch. The object
|
||||
IDs assigned at the seed frame remain stable for that video session.
|
||||
|
||||
`VLMSAM3TrackAdapter` is intentionally an adapter, not a second SAM3 loader. It
|
||||
validates ComfyUI core `SAM3_TRACK_DATA`, preserves the core bit-packed mask
|
||||
payload unchanged, and exposes lightweight `VLM_TRACKS` metadata with mask
|
||||
references. Connect its passthrough output to core `SAM3_TrackPreview` or
|
||||
`SAM3_TrackToMask`, and connect `tracks` to `VLMTrackReport`. This avoids
|
||||
duplicating dense masks in memory or JSON.
|
||||
|
||||
SAM3 weights use Meta's SAM License. The upstream `facebook/sam3` repository
|
||||
requires accepting access terms and sharing the requested account information;
|
||||
the ComfyUI checkpoint is also marked `sam-license`. Review and accept the
|
||||
license before downloading. The example names ComfyUI's
|
||||
`sam3.1_multiplex_fp16.safetensors`; if it is unavailable, use the SAM2.1
|
||||
workflow rather than substituting an unrelated checkpoint.
|
||||
|
||||
### Florence-2 task coverage
|
||||
|
||||
`Florence2` exposes all 15 supported task contracts:
|
||||
|
||||
| Task | Extra input | Structured result |
|
||||
| --- | --- | --- |
|
||||
| Caption | none | text |
|
||||
| Detailed caption | none | text |
|
||||
| More detailed caption | none | text |
|
||||
| OCR | none | text |
|
||||
| OCR with regions | none | text plus quadrilateral regions |
|
||||
| Object detection | none | labeled boxes |
|
||||
| Dense region caption | none | captions with boxes |
|
||||
| Caption to phrase grounding | `text_input` | phrase boxes |
|
||||
| Referring expression segmentation | `text_input` | polygons and mask |
|
||||
| Region to segmentation | one `BOUNDING_BOX` per image | polygons and mask |
|
||||
| Open vocabulary detection | `text_input` | model-provided spatial records |
|
||||
| Region to category | one `BOUNDING_BOX` per image | text |
|
||||
| Region to description | one `BOUNDING_BOX` per image | text |
|
||||
| Region to OCR | one `BOUNDING_BOX` per image | text |
|
||||
| Region proposals | none | boxes |
|
||||
|
||||
Every task returns `text`, `structured_json`, `mask`, and `visualization`.
|
||||
Tasks that do not produce a spatial result return an empty mask and the source
|
||||
image visualization. Region tasks reject ambiguous multi-box input; use
|
||||
`VLMSelectDetection` to isolate the record, then supply exactly one core
|
||||
`BOUNDING_BOX` with the same pixel coordinates.
|
||||
|
||||
### Video memory strategy
|
||||
|
||||
- Trim long media with core `Video Slice`, then use `GetVideoComponents`.
|
||||
Downscale the complete frame batch before detection or segmentation and keep
|
||||
every frame at identical dimensions.
|
||||
- Grounding detection supports configurable micro-batches; keep `batch_size=1`
|
||||
for minimum VRAM or increase it when memory allows. It returns both nested
|
||||
per-frame core `BOUNDING_BOX` values and flat metadata-rich
|
||||
`BOUNDING_BOXES`.
|
||||
- SAM2.1 stores source video frames on CPU, keeps its inference state on CPU by
|
||||
default, and limits the vision-feature cache to one frame. Union masks and
|
||||
previews return on CPU. Full per-object mask volumes are opt-in with
|
||||
`mask_output=union_and_objects`; disable `render_preview` to avoid another
|
||||
full-resolution overlay copy on long clips.
|
||||
- Start with Grounding DINO Tiny plus SAM2.1 Hiera Tiny. Increase detector or
|
||||
segmenter size only after the pipeline is correct. `unload_after=false`
|
||||
caches one model per node instance; use `true` when another large model must
|
||||
run immediately afterward.
|
||||
- A `Video Slice` is an independent propagation session. For very long media,
|
||||
use bounded slices, reseed each slice, and keep the overlap/output mapping in
|
||||
the caller. The pack does not pretend IDs are globally stable across separate
|
||||
queues.
|
||||
- The SAM3 adapter never unpacks the complete mask volume for its report. Use
|
||||
core `SAM3_TrackToMask` only when a dense selected mask is actually needed.
|
||||
|
||||
API-format examples are in [`examples/vision`](examples/vision):
|
||||
|
||||
- [`grounding_dino_image_api.json`](examples/vision/grounding_dino_image_api.json)
|
||||
- [`moondream3_preview_svg_segment_api.json`](examples/vision/moondream3_preview_svg_segment_api.json)
|
||||
- [`moondream31_video_detect_api.json`](examples/vision/moondream31_video_detect_api.json)
|
||||
- [`sam2_video_tracking_api.json`](examples/vision/sam2_video_tracking_api.json)
|
||||
- [`sam3_core_adapter_blueprint_api.json`](examples/vision/sam3_core_adapter_blueprint_api.json)
|
||||
- [`video_temporal_reasoning_api.json`](examples/vision/video_temporal_reasoning_api.json)
|
||||
- [`vlm_performance_preflight_api.json`](examples/vision/vlm_performance_preflight_api.json)
|
||||
|
||||
The dependency-free text-toolkit example is
|
||||
[`examples/text_toolkit_api.json`](examples/text_toolkit_api.json).
|
||||
Robotics policy, safety, and sidecar examples are in
|
||||
[`examples/robotics`](examples/robotics), including a complete universal
|
||||
HTTP policy graph.
|
||||
|
||||
Upload the named media to ComfyUI's input directory, adjust the filenames and
|
||||
labels, then submit the JSON object as the `prompt` value to `/prompt`. These
|
||||
are API graphs, not frontend workflow-export JSON.
|
||||
|
||||
## Node reference
|
||||
|
||||
All 89 registered nodes, grouped by their menu category. The **Node ID** is the
|
||||
`class_type` written into workflow and API JSON — search for that string when
|
||||
you need to find a node you saw on a canvas.
|
||||
|
||||
### Modern VLM
|
||||
|
||||
The main entry point for current vision-language models.
|
||||
|
||||
| Node | Node ID | Outputs |
|
||||
| --- | --- | --- |
|
||||
| Modern VLM (Qwen / SmolVLM2 / LFM / InternVL / Granite / Gemma) | `ModernVLM` | `STRING` |
|
||||
| Moondream 2 | `Moondream2model` | `STRING` |
|
||||
|
||||
### Moondream 3
|
||||
|
||||
Moondream 3 / 3.1 in an isolated Photon runtime. Load once, then reuse the
|
||||
`MOONDREAM31_MODEL` output across the task nodes.
|
||||
|
||||
| Node | Node ID | Outputs |
|
||||
| --- | --- | --- |
|
||||
| Moondream 3 / 3.1 Loader (Isolated Photon) | `Moondream31Loader` | `MOONDREAM31_MODEL`, `STRING` |
|
||||
| Moondream 3 / 3.1 Caption | `Moondream31Caption` | `STRING`, `STRING` |
|
||||
| Moondream 3 / 3.1 Query | `Moondream31Query` | `STRING`, `STRING`, `STRING` |
|
||||
| Moondream 3 / 3.1 Detect (Image / Video) | `Moondream31Detect` | `VLM_DETECTIONS`, `STRING`, `IMAGE`, `MASK`, `BOUNDING_BOX`, `BOUNDING_BOXES`, `STRING` |
|
||||
| Moondream 3 / 3.1 Point (Image / Video) | `Moondream31Point` | `VLM_POINTS`, `STRING`, `IMAGE`, `STRING` |
|
||||
| Moondream 3 Preview SVG Segment (Image / Video) | `Moondream31Segment` | `VLM_DETECTIONS`, `STRING`, `STRING`, `MASK`, `IMAGE`, `IMAGE`, `IMAGE`, `BOUNDING_BOX`, `BOUNDING_BOXES`, `STRING` |
|
||||
|
||||
### Florence-2
|
||||
|
||||
| Node | Node ID | Outputs |
|
||||
| --- | --- | --- |
|
||||
| Florence-2 Multitask Vision | `Florence2` | `STRING`, `STRING`, `MASK`, `IMAGE` |
|
||||
|
||||
### Vision: detection, segmentation, tracking
|
||||
|
||||
Open-vocabulary detection and video segmentation. These emit the structured
|
||||
`VLM_DETECTIONS` / `VLM_POINTS` / `VLM_TRACKS` types rather than loose strings.
|
||||
|
||||
| Node | Node ID | Outputs |
|
||||
| --- | --- | --- |
|
||||
| VLM Open-Vocabulary Detection | `VLMOpenVocabularyDetection` | `VLM_DETECTIONS`, `STRING`, `IMAGE`, `MASK`, `BOUNDING_BOX`, `BOUNDING_BOXES` |
|
||||
| VLM SAM2.1 Video Segmentation | `VLMSAM2VideoSegmentation` | `VLM_TRACKS`, `STRING`, `MASK`, `MASK`, `IMAGE` |
|
||||
| VLM SAM3 Track Adapter | `VLMSAM3TrackAdapter` | `VLM_TRACKS`, `SAM3_TRACK_DATA` |
|
||||
| VLM Track Detections | `VLMTrackDetections` | `VLM_TRACKS` |
|
||||
| VLM Track Report | `VLMTrackReport` | `STRING`, `STRING` |
|
||||
| JoyTag | `Joytag` | `STRING` |
|
||||
|
||||
### Vision: spatial reasoning
|
||||
|
||||
| Node | Node ID | Outputs |
|
||||
| --- | --- | --- |
|
||||
| VLM Spatial Prompt Builder | `VLMSpatialPromptBuilder` | `STRING` |
|
||||
| VLM Structured Spatial Parser | `VLMStructuredSpatialParser` | `VLM_DETECTIONS`, `VLM_POINTS`, `STRING` |
|
||||
|
||||
### Vision: detection utilities
|
||||
|
||||
Converters and filters between structured detections and ordinary Comfy types.
|
||||
|
||||
| Node | Node ID | Outputs |
|
||||
| --- | --- | --- |
|
||||
| Filter VLM Detections | `VLMFilterDetections` | `VLM_DETECTIONS` |
|
||||
| Select VLM Detection | `VLMSelectDetection` | `VLM_DETECTIONS` |
|
||||
| Crop VLM Detections | `VLMCropDetections` | `IMAGE`, `STRING` |
|
||||
| Render VLM Detections | `VLMRenderDetections` | `IMAGE` |
|
||||
| VLM Detection Centers | `VLMDetectionsToPoints` | `VLM_POINTS`, `STRING` |
|
||||
| VLM Detections from JSON | `VLMDetectionsFromJSON` | `VLM_DETECTIONS` |
|
||||
| VLM Detections to JSON | `VLMDetectionsToJSON` | `STRING` |
|
||||
| VLM Detections to Bounding Boxes | `VLMDetectionsToBoundingBoxes` | `BOUNDING_BOXES`, `STRING` |
|
||||
| VLM Detections to Masks | `VLMDetectionsToMasks` | `MASK`, `MASK`, `STRING`, `MASK`, `IMAGE`, `IMAGE`, `IMAGE` |
|
||||
|
||||
### Vision: mask tools
|
||||
|
||||
| Node | Node ID | Outputs |
|
||||
| --- | --- | --- |
|
||||
| VLM Mask Processor | `VLMMaskProcessor` | `MASK`, `MASK`, `MASK`, `IMAGE` |
|
||||
| VLM Mask Composite | `VLMMaskComposite` | `IMAGE`, `IMAGE`, `IMAGE`, `IMAGE` |
|
||||
|
||||
### Video intelligence
|
||||
|
||||
Adaptive frame selection and temporal reasoning for long videos.
|
||||
|
||||
| Node | Node ID | Outputs |
|
||||
| --- | --- | --- |
|
||||
| VLM Adaptive Frame Sampler | `VLMAdaptiveFrameSampler` | `IMAGE`, `VLM_VIDEO_SELECTION`, `STRING`, `STRING` |
|
||||
| VLM Video Reasoning Prompt | `VLMVideoReasoningPrompt` | `STRING`, `STRING` |
|
||||
| VLM Video Temporal Reasoner | `VLMVideoTemporalReasoner` | `STRING`, `VLM_EVENTS`, `VLM_VIDEO_SELECTION`, `IMAGE`, `STRING`, `STRING`, `STRING`, `STRING` |
|
||||
| VLM Temporal Events From JSON | `VLMEventsFromVideoJSON` | `VLM_EVENTS`, `STRING`, `STRING` |
|
||||
| VLM Persistent Scene State | `VLMBuildSceneState` | `VLM_SCENE_STATE`, `STRING`, `STRING` |
|
||||
| VLM Track-Aware Semantic Crops | `VLMTrackAwareCrops` | `IMAGE`, `STRING` |
|
||||
|
||||
### LLM (local GGUF)
|
||||
|
||||
llama.cpp text models. `LLM Loader (GGUF)` produces the `CUSTOM` model handle
|
||||
the samplers consume; the *Managed Cache* variants own their own handle and can
|
||||
release it after each run.
|
||||
|
||||
| Node | Node ID | Outputs |
|
||||
| --- | --- | --- |
|
||||
| LLM Loader (GGUF) | `LLMLoader` | `CUSTOM` |
|
||||
| LLM Sampler | `LLMSampler` | `STRING` |
|
||||
| LLM Prompt Generator | `LLMPromptGenerator` | `STRING` |
|
||||
| LLM (Managed Cache) | `LLMOptionalMemoryFreeSimple` | `STRING` |
|
||||
| LLM (Managed Cache, Advanced) | `LLMOptionalMemoryFreeAdvanced` | `STRING` |
|
||||
| Structured Output | `StructuredOutput` | `STRING` |
|
||||
| Structured Keyword Extraction | `KeywordExtraction` | `STRING` |
|
||||
| Structured Prompt Generator | `LLavaPromptGenerator` | `STRING` |
|
||||
| Creative Art Prompt Generator | `CreativeArtPromptGenerator` | `STRING` |
|
||||
| Prompt Suggester | `Suggester` | `STRING` |
|
||||
|
||||
### LLaVA (local GGUF multimodal)
|
||||
|
||||
Vision models through llama.cpp. These need both a GGUF and its vision
|
||||
projector (mmproj).
|
||||
|
||||
| Node | Node ID | Outputs |
|
||||
| --- | --- | --- |
|
||||
| LLaVA Loader | `LLava Loader Simple` | `CUSTOM` |
|
||||
| LLaVA Vision Projector Loader | `LlavaClipLoader` | `CUSTOM` |
|
||||
| LLaVA Sampler | `LLavaSamplerSimple` | `STRING` |
|
||||
| LLaVA Sampler (Advanced) | `LLavaSamplerAdvanced` | `STRING` |
|
||||
| LLaVA (Managed Cache) | `LLavaOptionalMemoryFreeSimple` | `STRING` |
|
||||
| LLaVA (Managed Cache, Advanced) | `LLavaOptionalMemoryFreeAdvanced` | `STRING` |
|
||||
|
||||
### Hosted APIs
|
||||
|
||||
| Node | Node ID | Outputs |
|
||||
| --- | --- | --- |
|
||||
| Hosted VLM API (Secure) | `HostedVLMAPI` | `STRING`, `STRING`, `INT` |
|
||||
| Hosted LLM API (Secure) | `PromptGenerateAPI` | `STRING` |
|
||||
|
||||
### Robotics / VLA policies
|
||||
|
||||
These nodes build and inspect policy observations/actions. They never send
|
||||
commands to robot hardware. Heavy policy runtimes stay in isolated LeRobot,
|
||||
openpi, GR00T, OpenVLA/OFT, or JAX environments.
|
||||
|
||||
| Node | Node ID | Outputs |
|
||||
| --- | --- | --- |
|
||||
| VLA Embodiment Profile | `VLAEmbodimentProfile` | `VLA_EMBODIMENT`, `STRING`, `INT`, `INT` |
|
||||
| VLA Observation Builder | `VLAObservationBuilder` | `VLA_OBSERVATION`, `STRING`, `INT` |
|
||||
| VLA Policy — Universal HTTP | `VLAHTTPPolicy` | `VLA_ACTIONS`, `STRING` |
|
||||
| VLA Policy — OpenPI WebSocket | `VLAOpenPIWebSocketPolicy` | `VLA_ACTIONS`, `STRING` |
|
||||
| VLA Policy — GR00T N1.7 ZMQ | `VLAGr00tZMQPolicy` | `VLA_ACTIONS`, `STRING` |
|
||||
| VLA Action Safety Gate | `VLAActionSafety` | `VLA_ACTIONS`, `STRING`, `BOOLEAN` |
|
||||
| VLA Actions From JSON | `VLAActionsFromJSON` | `VLA_ACTIONS`, `STRING` |
|
||||
| VLA Action Chunk Replan | `VLAActionChunkReplan` | `VLA_ACTIONS`, `STRING` |
|
||||
| VLA Action Inspect | `VLAActionInspect` | `STRING`, `STRING`, `INT`, `INT` |
|
||||
| VLA Trajectory Preview | `VLATrajectoryPreview` | `IMAGE` |
|
||||
| VLA Model Catalog | `VLAModelCatalog` | `STRING`, `STRING`, `STRING`, `STRING` |
|
||||
|
||||
### Text toolkit
|
||||
|
||||
Dependency-free string handling, so a VLM response can be shaped without an
|
||||
extra node pack.
|
||||
|
||||
| Node | Node ID | Outputs |
|
||||
| --- | --- | --- |
|
||||
| Text | `SimpleText` | `STRING`, `INT`, `INT`, `INT` |
|
||||
| Text Join | `VLMTextJoin` | `STRING`, `STRING`, `INT` |
|
||||
| Text Template | `VLMTextTemplate` | `STRING`, `STRING`, `STRING` |
|
||||
| Text Clean | `VLMTextClean` | `STRING`, `STRING` |
|
||||
| Text Replace | `VLMTextReplace` | `STRING`, `INT`, `STRING` |
|
||||
| Text Split / Batch | `VLMTextSplit` | `STRING`, `STRING`, `INT` |
|
||||
| Text Inspector | `VLMTextInspect` | `STRING`, `INT`, `INT`, `INT`, `INT`, `INT`, `STRING`, `STRING` |
|
||||
| View Text (Streaming) | `ViewText` | `STRING`, `INT`, `INT`, `INT`, `STRING` |
|
||||
| JSON Extract | `VLMJSONExtract` | `STRING`, `BOOLEAN`, `STRING`, `STRING` |
|
||||
| JSON to Text | `JsonToText` | `STRING`, `STRING`, `INT` |
|
||||
|
||||
### Performance and diagnostics
|
||||
|
||||
Run **VLM Runtime Diagnostics** before reporting a bug — it reports your
|
||||
device, backend, and which optional packages are installed.
|
||||
|
||||
| Node | Node ID | Outputs |
|
||||
| --- | --- | --- |
|
||||
| VLM Runtime Diagnostics | `VLMRuntimeDiagnostics` | `STRING` |
|
||||
| VLM Performance Profile | `VLMPerformanceProfile` | `INT`, `FLOAT`, `INT`, `INT`, `BOOLEAN`, `STRING` |
|
||||
| VLM Image Pixel Budget | `VLMImagePixelBudget` | `IMAGE`, `INT`, `INT`, `STRING` |
|
||||
|
||||
### Audio
|
||||
|
||||
| Node | Node ID | Outputs |
|
||||
| --- | --- | --- |
|
||||
| AudioLDM2 | `AudioLDM2Node` | `*`, `INT`, `AUDIO` |
|
||||
| Chat Musician | `ChatMusician` | `STRING`, `*`, `INT`, `AUDIO` |
|
||||
| MiniMax Music | `MiniMaxMusicNode` | `*`, `INT`, `AUDIO` |
|
||||
| PlayMusic Node | `PlayMusic` | `*` |
|
||||
| Save Audio | `SaveAudioNode` | — |
|
||||
|
||||
MiniMax Music reads `MINIMAX_API_KEY` only from the ComfyUI server
|
||||
environment. It uses fixed `global_en` and `cn_zh` endpoints, supports music
|
||||
generation and cover models, decodes URL or hexadecimal responses, and emits
|
||||
MP3, WAV, or PCM results through the existing waveform and `AUDIO` sockets.
|
||||
The `aigc_watermark` field is sent only for `cn_zh` requests. See the official
|
||||
[global](https://platform.minimax.io/docs/api-reference/music-generation) or
|
||||
[China](https://platform.minimaxi.com/docs/api-reference/music-generation)
|
||||
music API reference for account and content requirements.
|
||||
|
||||
### Legacy model loaders
|
||||
|
||||
Kept for existing workflows. New graphs should prefer **Modern VLM**, which
|
||||
covers most of these architectures through one interface.
|
||||
|
||||
| Node | Node ID | Outputs |
|
||||
| --- | --- | --- |
|
||||
| Qwen2-VL | `Qwen2VLNode` | `STRING` |
|
||||
| MiniCPM-V 2.6 (GGUF) | `MiniCPMNode` | `STRING` |
|
||||
| Molmo Vision-Language Model | `MolmoNode` | `STRING` |
|
||||
| PaLI-Gemma (Official Segmentation) | `Paligemma` | `STRING`, `MASK`, `IMAGE` |
|
||||
| Kosmos-2 | `Kosmos2model` | `STRING` |
|
||||
| MC-LLaVA | `MCLLaVAModel` | `STRING` |
|
||||
| UForm Gen2 Qwen | `UformGen2QwenNode` | `STRING` |
|
||||
| MoonDream (Moondream 2) | `MoonDream` | `STRING` |
|
||||
| [Legacy] Modern VLM Compatibility | `LegacyModernVLM` | `STRING` |
|
||||
|
||||
## Install
|
||||
|
||||
Install through ComfyUI Manager, or clone into `ComfyUI/custom_nodes` and run:
|
||||
@@ -70,6 +579,80 @@ Current official bitsandbytes wheels are installed automatically only on their
|
||||
supported OS/architecture combinations. Unsupported machines retain all
|
||||
non-quantized nodes.
|
||||
|
||||
### Robotics / VLA isolated runtimes
|
||||
|
||||
The robotics nodes keep policy dependencies outside ComfyUI. The universal
|
||||
HTTP client works without another package. Native openpi WebSocket and
|
||||
GR00T ZeroMQ clients use the lightweight optional extra:
|
||||
|
||||
```bash
|
||||
python -m pip install \
|
||||
-r ComfyUI/custom_nodes/ComfyUI_VLM_nodes/requirements-robotics-client.txt
|
||||
```
|
||||
|
||||
`VLA Model Catalog` covers current SmolVLA, X-VLA, π0/π0-FAST/π0.5,
|
||||
GR00T N1.7, WALL-OSS, MolmoAct2, VLA-JEPA, LingBot-VA, FastWAM, EO-1,
|
||||
EVO-1, OpenVLA-OFT, and Octo routes. “Available” means a supported isolated
|
||||
runtime/checkpoint path; base and architecture-only entries still require
|
||||
embodiment-specific training and transforms.
|
||||
|
||||
Start with SmolVLA for small consumer hardware. The included authenticated
|
||||
LeRobot sidecar loads one chosen policy, uses its serialized processors,
|
||||
returns action chunks over bounded JSON/JPEG, keeps it resident for speed,
|
||||
and can offload it to CPU after an idle timeout. Remote policy URLs require
|
||||
encrypted transport and explicit opt-in. Tokens are fixed environment
|
||||
variables (`VLA_POLICY_TOKEN`, `OPENPI_API_KEY`, or `GROOT_API_TOKEN`) and are
|
||||
never workflow inputs.
|
||||
|
||||
See [`examples/robotics/README.md`](examples/robotics/README.md) for D-drive
|
||||
WSL setup, platform boundaries, current model readiness, observation schemas,
|
||||
action safety semantics, and the runnable API example.
|
||||
|
||||
### Moondream 3 / 3.1 isolated runtime
|
||||
|
||||
Moondream's official Photon package pins Pillow below version 11 while
|
||||
current ComfyUI uses a newer Pillow. It therefore runs in a dedicated sidecar
|
||||
environment and never changes ComfyUI's Python packages. Read and accept the
|
||||
[Moondream Model License 1.0](https://moondream.ai/licenses/model/1.0), then
|
||||
create the environment under the registered `LLavacheckpoints` model folder.
|
||||
|
||||
Linux/WSL/macOS:
|
||||
|
||||
```bash
|
||||
runtime="ComfyUI/models/LLavacheckpoints/moondream31-runtime"
|
||||
uv venv "$runtime/.venv" --python 3.12
|
||||
uv pip install --python "$runtime/.venv/bin/python" \
|
||||
-r ComfyUI/custom_nodes/ComfyUI_VLM_nodes/requirements-moondream31.txt
|
||||
```
|
||||
|
||||
Windows PowerShell:
|
||||
|
||||
```powershell
|
||||
$runtime = "ComfyUI\models\LLavacheckpoints\moondream31-runtime"
|
||||
uv venv "$runtime\.venv" --python 3.12
|
||||
uv pip install --python "$runtime\.venv\Scripts\python.exe" `
|
||||
-r "ComfyUI\custom_nodes\ComfyUI_VLM_nodes\requirements-moondream31.txt"
|
||||
```
|
||||
|
||||
The first Loader execution downloads the selected official model below that
|
||||
runtime's `cache` directory. Use `moondream3.1-9B-A2B` for query, caption,
|
||||
detection, and pointing. Use `moondream3-preview` only for the SVG segment
|
||||
skill; the final 3.1 model card does not list segment. Set the server-side
|
||||
`MOONDREAM_PYTHON` environment variable
|
||||
when using a different isolated environment. Do not put this path or any
|
||||
credential in a workflow.
|
||||
|
||||
Official Photon local inference currently supports NVIDIA Ampere-or-newer on
|
||||
Linux/Windows and Apple Silicon on macOS 13 or newer. It does not currently
|
||||
provide local ROCm, Intel GPU, or CPU execution. Those platforms retain every
|
||||
portable Transformers, GGUF, API, and vision utility node in this pack.
|
||||
|
||||
On CUDA 12 x86-64 systems the isolated requirements deliberately install
|
||||
`nvidia-cuda-runtime-cu12==12.9.79`. Kestrel 0.4.6's AOT kernels require the
|
||||
`cudaLibraryLoadData` entry point, which is absent from the CUDA 12.6 runtime
|
||||
bundled by cu126 PyTorch. This pin updates only Photon's private runtime; it
|
||||
does not replace ComfyUI's PyTorch build or the host NVIDIA driver.
|
||||
|
||||
GGUF nodes use optional `llama-cpp-python`. Install a wheel built for the
|
||||
desired CUDA, ROCm/HIP, Metal, Vulkan, SYCL, or CPU backend:
|
||||
|
||||
@@ -113,7 +696,20 @@ Gemma 3 and PaLI-Gemma require accepting their model licenses on Hugging Face.
|
||||
The runtime reports llama.cpp's own compiled backend, GPU-offload, mmap, and
|
||||
mlock capabilities in **VLM Runtime Diagnostics**.
|
||||
- `unload_after=false` caches one model per node instance for fast repeated
|
||||
queues. Turn it on for maximum reclamation between prompts.
|
||||
queues. Cache creation is serialized, so concurrent API work cannot make the
|
||||
same node allocate duplicate model handles. Turn it on for maximum
|
||||
reclamation between prompts.
|
||||
- Moondream Photon asks ComfyUI to make room before it starts, then owns one
|
||||
exact isolated process. `unload_after=true` gracefully shuts it down and
|
||||
terminates that process if necessary, which releases Photon model, KV-cache,
|
||||
and CUDA-graph allocations without flushing unrelated ComfyUI models. The
|
||||
sidecar intentionally does not inherit ComfyUI's PyTorch allocator override;
|
||||
Photon's CUDA-graph capture uses the native allocator in its own process. The
|
||||
worker does not inherit unrelated provider keys or proxy credentials; only
|
||||
`HF_TOKEN`, and `MOONDREAM_API_KEY` for an explicitly selected adapter, may
|
||||
cross into its server-side environment. Base-model sidecars honor
|
||||
`DO_NOT_TRACK` locally and do not start Kestrel's anonymous telemetry task.
|
||||
Its random IPC secret is not placed on the process command line.
|
||||
- A connected `video_frames` batch becomes the primary visual input. The
|
||||
optional still-image socket is ignored for video inference so smaller models
|
||||
cannot silently answer from the wrong media.
|
||||
@@ -134,10 +730,96 @@ are not available for the installed PyTorch/backend combination.
|
||||
|
||||
## API nodes
|
||||
|
||||
`PromptGenerateAPI` supports the current OpenAI Responses API, the legacy Chat
|
||||
Completions API, and compatible base URLs. API keys can be supplied by node or
|
||||
environment (`OPENAI_API_KEY`, `ANTHROPIC_API_KEY`, `GEMINI_API_KEY`,
|
||||
`GROQ_API_KEY`). Keys are never persisted by this repository.
|
||||
**Hosted LLM API (Secure)** and **Hosted VLM API (Secure)** share a provider
|
||||
layer built around the current OpenAI Responses and Chat Completions request
|
||||
shapes, with Anthropic using its native Messages/vision contract and Gemini
|
||||
switching to its native multimodal contract for grounded or structured calls.
|
||||
The VLM node
|
||||
accepts a still image or a video-frame batch, samples
|
||||
frames uniformly, resizes and JPEG-compresses them, and enforces per-image and
|
||||
total request limits before upload. Both nodes can stream text into a connected
|
||||
`ViewText` node.
|
||||
|
||||
Both API nodes also expose:
|
||||
|
||||
- **Native web search** for OpenAI, Gemini, Anthropic, xAI, and any compatible
|
||||
model routed through OpenRouter. Unsupported presets fail clearly before a
|
||||
model request instead of silently pretending to search. Search can add
|
||||
provider cost and has provider-specific data terms, so it is off by default.
|
||||
- **JSON object** and **JSON Schema** output. Completed JSON is always parsed
|
||||
locally, JSON Schema results are validated locally, and invalid results fail
|
||||
the node instead of flowing into downstream automation.
|
||||
- **Open-source structured VLM output** through Custom / Local endpoints.
|
||||
OpenAI-standard mode supports vLLM, Ollama, and compatible servers;
|
||||
`llama.cpp JSON Schema` emits llama.cpp's direct schema dialect; and
|
||||
`JSON object + local validation` is a portable fallback for servers that
|
||||
implement only JSON mode.
|
||||
|
||||
User-provided schemas are capped at 64,000 characters, bounded by depth/node
|
||||
count, checked against their declared JSON Schema draft, and may use only local
|
||||
fragment `$ref` values. Remote/file references are rejected so validation can
|
||||
never turn into an unexpected network or filesystem lookup.
|
||||
|
||||
Curated production profiles include:
|
||||
|
||||
| Provider | Presets | Server environment variable |
|
||||
| --- | --- | --- |
|
||||
| OpenAI | GPT-5.6 Terra, Sol, Luna | `OPENAI_API_KEY` |
|
||||
| Google | Gemini 3.6 Flash, 3.5 Flash, 3.5 Flash-Lite | `GEMINI_API_KEY` |
|
||||
| Anthropic | Claude Fable 5, Opus 5, Sonnet 5, Haiku 4.5 | `ANTHROPIC_API_KEY` |
|
||||
| xAI | Grok 4.5 | `XAI_API_KEY` |
|
||||
| DeepSeek | V4 Flash, V4 Pro | `DEEPSEEK_API_KEY` |
|
||||
| Groq | Qwen 3.6 27B Vision, GPT-OSS 20B | `GROQ_API_KEY` |
|
||||
| Mistral | Mistral Large, Mistral Small, Ministral 14B | `MISTRAL_API_KEY` |
|
||||
| Together AI | Kimi K2.5, Qwen 3.5 9B | `TOGETHER_API_KEY` |
|
||||
| OpenRouter | Any compatible model ID | `OPENROUTER_API_KEY` |
|
||||
| Custom/local | OpenAI-compatible endpoint | `CUSTOM_API_KEY` |
|
||||
|
||||
Preset IDs were reviewed on 2026-07-29 against the official
|
||||
[OpenAI](https://developers.openai.com/api/docs/models),
|
||||
[Gemini](https://ai.google.dev/gemini-api/docs/models),
|
||||
[Claude](https://platform.claude.com/docs/en/about-claude/models/overview),
|
||||
[xAI](https://docs.x.ai/developers/models),
|
||||
[DeepSeek](https://api-docs.deepseek.com/updates/),
|
||||
[Groq](https://console.groq.com/docs/models),
|
||||
[Mistral](https://docs.mistral.ai/models/), and
|
||||
[Together](https://docs.together.ai/docs/inference/recommended-models), plus
|
||||
[OpenRouter's multimodal compatibility](https://openrouter.ai/docs/guides/overview/multimodal/overview)
|
||||
catalogs. Use `model_override` when a provider exposes a newer compatible model
|
||||
before the next node-pack release.
|
||||
|
||||
The capability routing follows the current official
|
||||
[OpenAI web-search](https://developers.openai.com/api/docs/guides/tools-web-search)
|
||||
and [structured-output](https://developers.openai.com/api/docs/guides/structured-outputs)
|
||||
contracts,
|
||||
[Gemini grounding](https://ai.google.dev/gemini-api/docs/google-search) and
|
||||
[structured output](https://ai.google.dev/gemini-api/docs/structured-output),
|
||||
[Claude web-search](https://platform.claude.com/docs/en/agents-and-tools/tool-use/web-search-tool)
|
||||
and [structured-output](https://platform.claude.com/docs/en/build-with-claude/structured-outputs)
|
||||
contracts, [xAI web search](https://docs.x.ai/developers/tools/web-search) and
|
||||
[structured outputs](https://docs.x.ai/developers/model-capabilities/text/structured-outputs),
|
||||
and [OpenRouter server-side search](https://openrouter.ai/docs/guides/features/server-tools/web-search).
|
||||
The local dialect is based on the
|
||||
[llama.cpp server API](https://github.com/ggml-org/llama.cpp/blob/master/tools/server/README.md).
|
||||
|
||||
API keys are not node inputs. A workflow contains only the provider selection,
|
||||
and the server resolves that provider's fixed environment variable at execution
|
||||
time. Built-in credentials are pinned to the provider's official HTTPS host;
|
||||
only the custom profile accepts a URL, and it can read only `CUSTOM_API_KEY`.
|
||||
Remote custom URLs require HTTPS, while keyless HTTP is restricted to
|
||||
`localhost`/loopback. Redirect following and environment proxies are disabled
|
||||
by default, API calls are stateless, OpenAI Responses explicitly use
|
||||
`store=false`, and provider exceptions are redacted before ComfyUI receives
|
||||
them.
|
||||
|
||||
Web search sends the prompt (and, where supported, the same multimodal request)
|
||||
to the selected provider's server-side search system. Do not enable it for
|
||||
content that must not be processed under that provider's search terms.
|
||||
|
||||
Opening an older `PromptGenerateAPI` workflow automatically clears its former
|
||||
plaintext key widget before the graph is configured. Save the migrated workflow
|
||||
to overwrite the old file, and rotate any key that was previously saved or
|
||||
shared. See [SECURITY.md](SECURITY.md) for setup and the exact threat model.
|
||||
|
||||
## Reliability guarantees
|
||||
|
||||
@@ -171,3 +853,23 @@ catalog-only evidence matrix.
|
||||
|
||||
Please report reproducible bugs at the
|
||||
[issue tracker](https://github.com/gokayfem/ComfyUI_VLM_nodes/issues).
|
||||
|
||||
<details>
|
||||
<summary><strong>Cite this project</strong></summary>
|
||||
|
||||
If ComfyUI VLM Nodes supports your work, please cite the software. GitHub also
|
||||
provides ready-to-copy APA and BibTeX entries via **Cite this repository**.
|
||||
|
||||
```bibtex
|
||||
@software{Aydogan_ComfyUI_VLM_Nodes_2026,
|
||||
author = {Aydoğan, Gökay},
|
||||
title = {ComfyUI VLM Nodes},
|
||||
version = {3.5.0},
|
||||
year = {2026},
|
||||
url = {https://github.com/gokayfem/ComfyUI_VLM_nodes}
|
||||
}
|
||||
```
|
||||
|
||||
[ORCID](https://orcid.org/0000-0002-2343-9433) · [Citation metadata](CITATION.cff)
|
||||
|
||||
</details>
|
||||
|
||||
+120
@@ -0,0 +1,120 @@
|
||||
# API credential security
|
||||
|
||||
## Guarantees
|
||||
|
||||
- API keys are never accepted as node inputs, widget values, workflow fields,
|
||||
outputs, metadata, or log messages.
|
||||
- Each built-in provider reads only its standard server-side environment
|
||||
variable and sends it only to that provider's fixed official HTTPS endpoint.
|
||||
- A built-in provider key cannot be combined with a workflow-supplied URL.
|
||||
- The custom endpoint reads only `CUSTOM_API_KEY`. Remote custom endpoints must
|
||||
use HTTPS; unencrypted and keyless requests are limited to loopback.
|
||||
- HTTP redirects and environment proxies are disabled by default. Proxy use is
|
||||
an explicit non-secret node option for installations that require it.
|
||||
- Hosted calls are stateless. No Python node-instance conversation history is
|
||||
retained, and OpenAI Responses requests set `store=false`.
|
||||
- Exceptions are bounded and redact the resolved key, URL-encoded variants,
|
||||
bearer tokens, common provider-key formats, authorization fields, and URL
|
||||
user-info before the message reaches ComfyUI.
|
||||
- Local image/video-frame uploads are uniformly sampled, resized,
|
||||
JPEG-compressed, limited to 4 MiB per image, and limited to 24 MiB total.
|
||||
- User JSON Schemas are size/depth/node bounded and may contain only local
|
||||
fragment `$ref` values. Remote URLs and file references are rejected before
|
||||
validation, preventing schema resolution from becoming an SSRF or local-file
|
||||
access path.
|
||||
|
||||
## Configure credentials
|
||||
|
||||
Set the matching variable in the environment that launches ComfyUI, then
|
||||
restart ComfyUI:
|
||||
|
||||
| Provider | Variable |
|
||||
| --- | --- |
|
||||
| OpenAI | `OPENAI_API_KEY` |
|
||||
| Google Gemini | `GEMINI_API_KEY` |
|
||||
| Anthropic | `ANTHROPIC_API_KEY` |
|
||||
| xAI | `XAI_API_KEY` |
|
||||
| DeepSeek | `DEEPSEEK_API_KEY` |
|
||||
| Groq | `GROQ_API_KEY` |
|
||||
| Mistral | `MISTRAL_API_KEY` |
|
||||
| Together AI | `TOGETHER_API_KEY` |
|
||||
| OpenRouter | `OPENROUTER_API_KEY` |
|
||||
| MiniMax | `MINIMAX_API_KEY` |
|
||||
| Custom remote endpoint | `CUSTOM_API_KEY` |
|
||||
| Universal VLA policy server | `VLA_POLICY_TOKEN` |
|
||||
| openpi WebSocket server | `OPENPI_API_KEY` |
|
||||
| Isaac-GR00T ZMQ server | `GROOT_API_TOKEN` |
|
||||
|
||||
For an interactive POSIX/WSL session, this avoids putting the value in shell
|
||||
history:
|
||||
|
||||
```bash
|
||||
read -rsp "Provider API key: " OPENAI_API_KEY
|
||||
export OPENAI_API_KEY
|
||||
python main.py
|
||||
```
|
||||
|
||||
Use the equivalent secret manager or service environment mechanism for a
|
||||
persistent installation. Do not commit a `.env` file, workflow containing an
|
||||
old key, shell script containing a key, or copied ComfyUI log.
|
||||
|
||||
Web search is disabled by default. Enabling it sends the request content to the
|
||||
selected provider's server-side search system and may have separate retention,
|
||||
regional-availability, and billing terms. Treat it as an explicit data-sharing
|
||||
choice; do not enable it for content that is outside those terms.
|
||||
|
||||
## Robotics policy endpoints
|
||||
|
||||
Robotics tokens are also server-side only. Workflow nodes select an endpoint,
|
||||
but cannot select an arbitrary environment variable or contain the secret
|
||||
value.
|
||||
|
||||
- The universal policy client permits unencrypted HTTP only on loopback.
|
||||
Remote use requires HTTPS plus `allow_remote=true`; redirects and
|
||||
environment proxies are disabled.
|
||||
- The openpi client permits unencrypted WebSocket only on loopback. Remote use
|
||||
requires WSS plus `allow_remote=true`.
|
||||
- GR00T's official ZeroMQ protocol has token authentication but no built-in
|
||||
transport encryption. Keep it on loopback/private infrastructure or place it
|
||||
inside an authenticated encrypted tunnel. Never expose its port directly to
|
||||
the public internet.
|
||||
- Camera payloads are JPEG-compressed and bounded per frame and per request.
|
||||
Response sizes, camera count, observation history, state/action dimensions,
|
||||
and action horizons are bounded before use.
|
||||
- MessagePack ndarray decoders reject object/void dtypes and never deserialize
|
||||
pickle. The included HTTP sidecar uses bounded JSON instead of LeRobot's
|
||||
pickle-based asynchronous transport.
|
||||
- Errors redact the resolved token and authorization-like values. Reports
|
||||
include only endpoint scheme/host/port, not request headers, full camera
|
||||
payloads, or state data.
|
||||
|
||||
Robot observations may expose people, homes, workplaces, proprietary tasks,
|
||||
and physical state. Treat them as sensitive even when no API key is present.
|
||||
The safety node is a data validation gate, not a certified control system.
|
||||
This package intentionally contains no ROS, serial, CAN, motor, or robot SDK
|
||||
transport; a separate controller must enforce emergency stop, deadman,
|
||||
watchdog, collision/workspace, command-age, and manufacturer limits.
|
||||
|
||||
## Legacy workflows
|
||||
|
||||
Versions before this security update exposed an `api_key` text widget.
|
||||
The frontend migration clears position 3 of every serialized
|
||||
`PromptGenerateAPI` node before LiteGraph creates the active node, including
|
||||
nodes inside saved subgraph definitions. The backend independently rejects any
|
||||
value that is not one of the two safe credential-source choices.
|
||||
|
||||
The source workflow file is not rewritten merely by opening it. Save the
|
||||
migrated workflow, securely remove old copies, and rotate any credential that
|
||||
was ever saved, shared, committed, backed up, or placed in an exported PNG.
|
||||
|
||||
## Threat boundary
|
||||
|
||||
ComfyUI custom nodes execute Python code with the permissions of the ComfyUI
|
||||
process. Another untrusted custom-node package can read the same process
|
||||
environment regardless of protections in this repository. Install only trusted
|
||||
node packs, keep ComfyUI authenticated and bound to a trusted interface, and do
|
||||
not expose an unauthenticated server to the public internet.
|
||||
|
||||
If a key may have been exposed, revoke it with the provider immediately, review
|
||||
usage, create a replacement with the minimum needed project permissions and
|
||||
spend limit, and restart ComfyUI with the replacement.
|
||||
+12
@@ -7,24 +7,36 @@ LOGGER = logging.getLogger("ComfyUI_VLM_nodes")
|
||||
register_model_folder()
|
||||
|
||||
node_list = [
|
||||
"acceleration",
|
||||
"audioldm2",
|
||||
"diagnostics",
|
||||
"florence2",
|
||||
"grounding",
|
||||
"hosted_api",
|
||||
"joytag",
|
||||
"kosmos2",
|
||||
"llavaloader",
|
||||
"mcllava",
|
||||
"minicpm",
|
||||
"minimax_music",
|
||||
"modern_vlm",
|
||||
"molmo",
|
||||
"moondream31",
|
||||
"moondream2",
|
||||
"moondream_script",
|
||||
"paligemma",
|
||||
"playmusic",
|
||||
"qwen2vl",
|
||||
"robotics",
|
||||
"sam2",
|
||||
"sam3_adapter",
|
||||
"simpletext",
|
||||
"spatial_parser",
|
||||
"suggest",
|
||||
"tracking",
|
||||
"uform",
|
||||
"video_intelligence",
|
||||
"vision_utils",
|
||||
]
|
||||
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
|
||||
@@ -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
|
||||
]
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,57 @@
|
||||
{
|
||||
"1": {
|
||||
"class_type": "SimpleText",
|
||||
"inputs": {
|
||||
"input_text": "Model response:\n```json\n{\"scene\":{\"subject\":\"warehouse robot\",\"action\":\"moving a blue crate\"}}\n```"
|
||||
}
|
||||
},
|
||||
"2": {
|
||||
"class_type": "VLMJSONExtract",
|
||||
"inputs": {
|
||||
"text": [
|
||||
"1",
|
||||
0
|
||||
],
|
||||
"path": "$.scene.action",
|
||||
"output_format": "Text",
|
||||
"if_missing": "Error",
|
||||
"default_value": ""
|
||||
}
|
||||
},
|
||||
"3": {
|
||||
"class_type": "VLMTextTemplate",
|
||||
"inputs": {
|
||||
"template": "{instruction}\n\nObserved action: {text1}",
|
||||
"variables_json": "{\"instruction\":\"Write one concise video-generation prompt.\"}",
|
||||
"missing_values": "Error",
|
||||
"text1": [
|
||||
"2",
|
||||
0
|
||||
]
|
||||
}
|
||||
},
|
||||
"4": {
|
||||
"class_type": "VLMTextClean",
|
||||
"inputs": {
|
||||
"text": [
|
||||
"3",
|
||||
0
|
||||
],
|
||||
"unicode_normalization": "NFC",
|
||||
"whitespace": "Normalize line endings",
|
||||
"trim_edges": true,
|
||||
"remove_outer_markdown_fence": false,
|
||||
"deduplicate_lines": false,
|
||||
"max_characters": 0
|
||||
}
|
||||
},
|
||||
"5": {
|
||||
"class_type": "ViewText",
|
||||
"inputs": {
|
||||
"text": [
|
||||
"4",
|
||||
0
|
||||
]
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,124 @@
|
||||
# Vision API examples
|
||||
|
||||
These files contain ComfyUI API prompt graphs: the object that belongs under
|
||||
the `prompt` key in a `POST /prompt` request. They are not frontend workflow
|
||||
exports and are not intended for drag-and-drop import into the canvas.
|
||||
|
||||
Before queueing:
|
||||
|
||||
1. Copy the named image/video into `ComfyUI/input`, or change the `image`/`file`
|
||||
widget value to an existing input filename.
|
||||
2. Restart ComfyUI after installing or updating this node pack.
|
||||
3. Confirm every `class_type` is present in `/object_info`.
|
||||
4. Wrap the loaded JSON as `{"prompt": graph}` in the API request.
|
||||
|
||||
## Examples
|
||||
|
||||
### `grounding_dino_image_api.json`
|
||||
|
||||
Runs Grounding DINO Tiny over `grounding_input.png`. Node 2 outputs:
|
||||
|
||||
| Index | Output |
|
||||
| ---: | --- |
|
||||
| 0 | `VLM_DETECTIONS` |
|
||||
| 1 | Structured detection JSON |
|
||||
| 2 | Detection overlay |
|
||||
| 3 | Box mask |
|
||||
| 4 | Core nested per-frame `BOUNDING_BOX` |
|
||||
| 5 | Flat metadata-rich `BOUNDING_BOXES` |
|
||||
|
||||
`PreviewImage` displays output 2 and `ViewText` reports output 1.
|
||||
|
||||
### `vlm_performance_preflight_api.json`
|
||||
|
||||
Loads `vlm_api_people_birds.mp4` with Comfy core video nodes, applies the
|
||||
`Fast video` performance profile, runs the track-aware adaptive sampler, and
|
||||
then applies a 14-pixel-aligned image budget. The preview shows the exact batch
|
||||
that can be connected to any local or hosted VLM. Three `ViewText` nodes report
|
||||
the selected source indices/timestamps, pixel reduction, and active profile.
|
||||
|
||||
### `moondream3_preview_svg_segment_api.json`
|
||||
|
||||
Runs the official Moondream 3 Preview SVG segmentation skill over
|
||||
`moondream_segment_input.png`. Read the linked model license and change
|
||||
`license_accepted` to `true` before queueing. The graph previews the
|
||||
black/white mask, isolated foreground cutout, and mask/box/polygon overlay;
|
||||
`ViewText` receives the exact native SVG path plus its normalized bbox.
|
||||
|
||||
Moondream's path coordinates are normalized within the returned bbox. The
|
||||
node preserves that path verbatim, safely flattens curves/arcs, applies an
|
||||
even-odd fill for subpath holes, and supersamples the raster edge. The
|
||||
canonical detection keeps both the primary polygon and the full in-process
|
||||
mask.
|
||||
|
||||
### `moondream31_video_detect_api.json`
|
||||
|
||||
Loads `moondream_video_input.mp4`, passes the real frame batch and source FPS
|
||||
to Moondream, and analyzes every frame with four concurrent requests. Photon
|
||||
uses the Loader's `max_batch_size=4` scheduler capacity to form dynamic
|
||||
batches. `ViewText` reports measured throughput and real-time factor. Increase
|
||||
`frame_stride` to 2, 3, or more when full-frame analysis cannot keep up with
|
||||
the source FPS; the canonical results preserve original frame indices and
|
||||
timestamps.
|
||||
|
||||
### `sam2_video_tracking_api.json`
|
||||
|
||||
Runs this bounded pipeline:
|
||||
|
||||
`LoadVideo` → `Video Slice` → `GetVideoComponents` → `ImageScale` →
|
||||
`ImageFromBatch` → Grounding DINO first-frame detection → SAM2.1 propagation.
|
||||
|
||||
The example limits the source to two seconds, scales its largest dimension to
|
||||
768 pixels while preserving aspect ratio, unloads Grounding DINO after
|
||||
seeding, and keeps SAM2.1 video state on CPU. The example requests only the
|
||||
union mask volume; change `mask_output` to `union_and_objects` only when every
|
||||
per-object mask is required. `VLMTrackReport` is an output node and the final
|
||||
`PreviewImage` displays SAM2.1 output index 4.
|
||||
|
||||
For a longer source, change `start_time` and keep a bounded `duration`.
|
||||
Independent slices create independent object-ID sessions.
|
||||
|
||||
### `sam3_core_adapter_blueprint_api.json`
|
||||
|
||||
Uses ComfyUI core nodes to load and run SAM3.1, then passes core
|
||||
`SAM3_TRACK_DATA` through `VLMSAM3TrackAdapter`. The adapter's output 1 is the
|
||||
unchanged core payload consumed by `SAM3_TrackPreview`; output 0 is canonical
|
||||
`VLM_TRACKS` consumed by `VLMTrackReport`.
|
||||
|
||||
The graph intentionally names:
|
||||
|
||||
`ComfyUI/models/checkpoints/sam3.1_multiplex_fp16.safetensors`
|
||||
|
||||
The checkpoint is not bundled. Review the SAM License before downloading
|
||||
[Comfy-Org/sam3.1](https://huggingface.co/Comfy-Org/sam3.1). ComfyUI rejects
|
||||
the graph at prompt validation when the named checkpoint is absent. Use the
|
||||
SAM2.1 example when SAM3.1 access or compatible core support is unavailable.
|
||||
|
||||
## Output history
|
||||
|
||||
ComfyUI returns image/video previews in the execution history and text reports
|
||||
in the output-node UI payload. Canonical JSON is also available on the linked
|
||||
string outputs. Dense masks intentionally stay as tensors rather than being
|
||||
embedded in the JSON report.
|
||||
|
||||
## Creator mask outputs
|
||||
|
||||
`VLM Detections to Masks` preserves its original first three outputs and
|
||||
appends creator-ready derivatives:
|
||||
|
||||
| Index | Output |
|
||||
| ---: | --- |
|
||||
| 0 | Per-frame combined/union `MASK` |
|
||||
| 1 | Flattened per-object `MASK` batch |
|
||||
| 2 | JSON mapping each object mask to its frame/detection/track |
|
||||
| 3 | Per-frame inverse/background `MASK` |
|
||||
| 4 | Combined masks as black-and-white `IMAGE` batches |
|
||||
| 5 | Individual masks as black-and-white `IMAGE` batches |
|
||||
| 6 | Stable-color per-frame instance maps |
|
||||
|
||||
All binary mask values are exactly zero or one. `VLM Mask Processor` can grow,
|
||||
shrink, and feather any of these masks and returns processed, binary, inverse,
|
||||
and black-and-white image outputs. `VLM Mask Composite` accepts the resulting
|
||||
mask plus still-image or video frames and returns a composite, isolated
|
||||
foreground, background-only plate, and mask image. Connect an optional
|
||||
background image/video batch to replace the solid background color.
|
||||
@@ -0,0 +1,45 @@
|
||||
{
|
||||
"1": {
|
||||
"class_type": "LoadImage",
|
||||
"inputs": {
|
||||
"image": "grounding_input.png"
|
||||
}
|
||||
},
|
||||
"2": {
|
||||
"class_type": "VLMOpenVocabularyDetection",
|
||||
"inputs": {
|
||||
"image": [
|
||||
"1",
|
||||
0
|
||||
],
|
||||
"model": "Grounding DINO Tiny (fast)",
|
||||
"labels": "person, dog, bicycle",
|
||||
"box_threshold": 0.3,
|
||||
"text_threshold": 0.25,
|
||||
"max_detections": 100,
|
||||
"fps": 1.0,
|
||||
"nms_threshold": 0.5,
|
||||
"precision": "auto",
|
||||
"batch_size": 1,
|
||||
"unload_after": false
|
||||
}
|
||||
},
|
||||
"3": {
|
||||
"class_type": "PreviewImage",
|
||||
"inputs": {
|
||||
"images": [
|
||||
"2",
|
||||
2
|
||||
]
|
||||
}
|
||||
},
|
||||
"4": {
|
||||
"class_type": "ViewText",
|
||||
"inputs": {
|
||||
"text": [
|
||||
"2",
|
||||
1
|
||||
]
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,76 @@
|
||||
{
|
||||
"1": {
|
||||
"class_type": "LoadVideo",
|
||||
"inputs": {
|
||||
"file": "moondream_video_input.mp4"
|
||||
}
|
||||
},
|
||||
"2": {
|
||||
"class_type": "GetVideoComponents",
|
||||
"inputs": {
|
||||
"video": [
|
||||
"1",
|
||||
0
|
||||
]
|
||||
}
|
||||
},
|
||||
"3": {
|
||||
"class_type": "Moondream31Loader",
|
||||
"inputs": {
|
||||
"license_accepted": false,
|
||||
"device": "Auto",
|
||||
"max_batch_size": 4,
|
||||
"kv_cache_profile": "Balanced (8K pages)",
|
||||
"model_or_adapter": "moondream3.1-9B-A2B"
|
||||
}
|
||||
},
|
||||
"4": {
|
||||
"class_type": "Moondream31Detect",
|
||||
"inputs": {
|
||||
"model": [
|
||||
"3",
|
||||
0
|
||||
],
|
||||
"image": [
|
||||
"2",
|
||||
0
|
||||
],
|
||||
"object": "person",
|
||||
"fps": [
|
||||
"2",
|
||||
2
|
||||
],
|
||||
"frame_stride": 1,
|
||||
"parallel_requests": 4,
|
||||
"max_objects": 100,
|
||||
"unload_after": false
|
||||
}
|
||||
},
|
||||
"5": {
|
||||
"class_type": "PreviewImage",
|
||||
"inputs": {
|
||||
"images": [
|
||||
"4",
|
||||
2
|
||||
]
|
||||
}
|
||||
},
|
||||
"6": {
|
||||
"class_type": "ViewText",
|
||||
"inputs": {
|
||||
"text": [
|
||||
"4",
|
||||
6
|
||||
]
|
||||
}
|
||||
},
|
||||
"7": {
|
||||
"class_type": "ViewText",
|
||||
"inputs": {
|
||||
"text": [
|
||||
"4",
|
||||
1
|
||||
]
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,74 @@
|
||||
{
|
||||
"1": {
|
||||
"class_type": "LoadImage",
|
||||
"inputs": {
|
||||
"image": "moondream_segment_input.png"
|
||||
}
|
||||
},
|
||||
"2": {
|
||||
"class_type": "Moondream31Loader",
|
||||
"inputs": {
|
||||
"license_accepted": false,
|
||||
"device": "Auto",
|
||||
"max_batch_size": 4,
|
||||
"kv_cache_profile": "Balanced (8K pages)",
|
||||
"model_or_adapter": "moondream3-preview"
|
||||
}
|
||||
},
|
||||
"3": {
|
||||
"class_type": "Moondream31Segment",
|
||||
"inputs": {
|
||||
"model": [
|
||||
"2",
|
||||
0
|
||||
],
|
||||
"image": [
|
||||
"1",
|
||||
0
|
||||
],
|
||||
"object": "main foreground object",
|
||||
"fps": 1.0,
|
||||
"frame_stride": 1,
|
||||
"parallel_requests": 1,
|
||||
"svg_supersample": 4,
|
||||
"unload_after": false,
|
||||
"spatial_refs_json": "[]"
|
||||
}
|
||||
},
|
||||
"4": {
|
||||
"class_type": "PreviewImage",
|
||||
"inputs": {
|
||||
"images": [
|
||||
"3",
|
||||
4
|
||||
]
|
||||
}
|
||||
},
|
||||
"5": {
|
||||
"class_type": "PreviewImage",
|
||||
"inputs": {
|
||||
"images": [
|
||||
"3",
|
||||
5
|
||||
]
|
||||
}
|
||||
},
|
||||
"6": {
|
||||
"class_type": "PreviewImage",
|
||||
"inputs": {
|
||||
"images": [
|
||||
"3",
|
||||
6
|
||||
]
|
||||
}
|
||||
},
|
||||
"7": {
|
||||
"class_type": "ViewText",
|
||||
"inputs": {
|
||||
"text": [
|
||||
"3",
|
||||
2
|
||||
]
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,116 @@
|
||||
{
|
||||
"1": {
|
||||
"class_type": "LoadVideo",
|
||||
"inputs": {
|
||||
"file": "tracking_input.mp4"
|
||||
}
|
||||
},
|
||||
"2": {
|
||||
"class_type": "Video Slice",
|
||||
"inputs": {
|
||||
"video": [
|
||||
"1",
|
||||
0
|
||||
],
|
||||
"start_time": 0.0,
|
||||
"duration": 2.0,
|
||||
"strict_duration": false
|
||||
}
|
||||
},
|
||||
"3": {
|
||||
"class_type": "GetVideoComponents",
|
||||
"inputs": {
|
||||
"video": [
|
||||
"2",
|
||||
0
|
||||
]
|
||||
}
|
||||
},
|
||||
"4": {
|
||||
"class_type": "ImageScaleToMaxDimension",
|
||||
"inputs": {
|
||||
"image": [
|
||||
"3",
|
||||
0
|
||||
],
|
||||
"upscale_method": "area",
|
||||
"largest_size": 768
|
||||
}
|
||||
},
|
||||
"5": {
|
||||
"class_type": "ImageFromBatch",
|
||||
"inputs": {
|
||||
"image": [
|
||||
"4",
|
||||
0
|
||||
],
|
||||
"batch_index": 0,
|
||||
"length": 1
|
||||
}
|
||||
},
|
||||
"6": {
|
||||
"class_type": "VLMOpenVocabularyDetection",
|
||||
"inputs": {
|
||||
"image": [
|
||||
"5",
|
||||
0
|
||||
],
|
||||
"model": "Grounding DINO Tiny (fast)",
|
||||
"labels": "person, dog, vehicle",
|
||||
"box_threshold": 0.3,
|
||||
"text_threshold": 0.25,
|
||||
"max_detections": 16,
|
||||
"fps": [
|
||||
"3",
|
||||
2
|
||||
],
|
||||
"nms_threshold": 0.5,
|
||||
"precision": "auto",
|
||||
"batch_size": 1,
|
||||
"unload_after": true
|
||||
}
|
||||
},
|
||||
"7": {
|
||||
"class_type": "VLMSAM2VideoSegmentation",
|
||||
"inputs": {
|
||||
"images": [
|
||||
"4",
|
||||
0
|
||||
],
|
||||
"model": "SAM2.1 Hiera Tiny (fast)",
|
||||
"seed_frame": 0,
|
||||
"fps": [
|
||||
"3",
|
||||
2
|
||||
],
|
||||
"detections": [
|
||||
"6",
|
||||
0
|
||||
],
|
||||
"mask_threshold": 0.0,
|
||||
"precision": "auto",
|
||||
"keep_video_on_cpu": true,
|
||||
"mask_output": "union_only",
|
||||
"render_preview": true,
|
||||
"unload_after": false
|
||||
}
|
||||
},
|
||||
"8": {
|
||||
"class_type": "VLMTrackReport",
|
||||
"inputs": {
|
||||
"tracks": [
|
||||
"7",
|
||||
0
|
||||
]
|
||||
}
|
||||
},
|
||||
"9": {
|
||||
"class_type": "PreviewImage",
|
||||
"inputs": {
|
||||
"images": [
|
||||
"7",
|
||||
4
|
||||
]
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,116 @@
|
||||
{
|
||||
"1": {
|
||||
"class_type": "LoadVideo",
|
||||
"inputs": {
|
||||
"file": "tracking_input.mp4"
|
||||
}
|
||||
},
|
||||
"2": {
|
||||
"class_type": "Video Slice",
|
||||
"inputs": {
|
||||
"video": [
|
||||
"1",
|
||||
0
|
||||
],
|
||||
"start_time": 0.0,
|
||||
"duration": 2.0,
|
||||
"strict_duration": false
|
||||
}
|
||||
},
|
||||
"3": {
|
||||
"class_type": "GetVideoComponents",
|
||||
"inputs": {
|
||||
"video": [
|
||||
"2",
|
||||
0
|
||||
]
|
||||
}
|
||||
},
|
||||
"4": {
|
||||
"class_type": "ImageScaleToMaxDimension",
|
||||
"inputs": {
|
||||
"image": [
|
||||
"3",
|
||||
0
|
||||
],
|
||||
"upscale_method": "area",
|
||||
"largest_size": 768
|
||||
}
|
||||
},
|
||||
"5": {
|
||||
"class_type": "CheckpointLoaderSimple",
|
||||
"inputs": {
|
||||
"ckpt_name": "sam3.1_multiplex_fp16.safetensors"
|
||||
}
|
||||
},
|
||||
"6": {
|
||||
"class_type": "CLIPTextEncode",
|
||||
"inputs": {
|
||||
"text": "person, dog, vehicle",
|
||||
"clip": [
|
||||
"5",
|
||||
1
|
||||
]
|
||||
}
|
||||
},
|
||||
"7": {
|
||||
"class_type": "SAM3_VideoTrack",
|
||||
"inputs": {
|
||||
"images": [
|
||||
"4",
|
||||
0
|
||||
],
|
||||
"model": [
|
||||
"5",
|
||||
0
|
||||
],
|
||||
"conditioning": [
|
||||
"6",
|
||||
0
|
||||
],
|
||||
"detection_threshold": 0.5,
|
||||
"max_objects": 8,
|
||||
"detect_interval": 1
|
||||
}
|
||||
},
|
||||
"8": {
|
||||
"class_type": "VLMSAM3TrackAdapter",
|
||||
"inputs": {
|
||||
"track_data": [
|
||||
"7",
|
||||
0
|
||||
],
|
||||
"fps": [
|
||||
"3",
|
||||
2
|
||||
]
|
||||
}
|
||||
},
|
||||
"9": {
|
||||
"class_type": "VLMTrackReport",
|
||||
"inputs": {
|
||||
"tracks": [
|
||||
"8",
|
||||
0
|
||||
]
|
||||
}
|
||||
},
|
||||
"10": {
|
||||
"class_type": "SAM3_TrackPreview",
|
||||
"inputs": {
|
||||
"track_data": [
|
||||
"8",
|
||||
1
|
||||
],
|
||||
"images": [
|
||||
"4",
|
||||
0
|
||||
],
|
||||
"opacity": 0.5,
|
||||
"fps": [
|
||||
"3",
|
||||
2
|
||||
]
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
{
|
||||
"1": {
|
||||
"class_type": "LoadVideo",
|
||||
"inputs": {
|
||||
"file": "video_understanding_input.mp4"
|
||||
}
|
||||
},
|
||||
"2": {
|
||||
"class_type": "GetVideoComponents",
|
||||
"inputs": {
|
||||
"video": [
|
||||
"1",
|
||||
0
|
||||
]
|
||||
}
|
||||
},
|
||||
"3": {
|
||||
"class_type": "VLMVideoTemporalReasoner",
|
||||
"inputs": {
|
||||
"frames": [
|
||||
"2",
|
||||
0
|
||||
],
|
||||
"fps": [
|
||||
"2",
|
||||
2
|
||||
],
|
||||
"task": "Detailed temporal summary",
|
||||
"question": "Describe what happens over time and identify the visible evidence.",
|
||||
"model": "Qwen 3 VL 2B Instruct",
|
||||
"custom_model_id": "",
|
||||
"memory_mode": "ComfyUI managed (BF16)",
|
||||
"max_frames": 16,
|
||||
"max_events": 24,
|
||||
"max_new_tokens": 768,
|
||||
"strategy": "Hybrid: scene + motion + tracks",
|
||||
"minimum_gap_seconds": 0.15,
|
||||
"analysis_max_side": 448,
|
||||
"attention_mode": "Auto (SDPA)",
|
||||
"enable_thinking": false,
|
||||
"strict_output": true,
|
||||
"unload_after": false,
|
||||
"stream_output": true
|
||||
}
|
||||
},
|
||||
"4": {
|
||||
"class_type": "ViewText",
|
||||
"inputs": {
|
||||
"text": [
|
||||
"3",
|
||||
0
|
||||
]
|
||||
}
|
||||
},
|
||||
"5": {
|
||||
"class_type": "ViewText",
|
||||
"inputs": {
|
||||
"text": [
|
||||
"3",
|
||||
6
|
||||
]
|
||||
}
|
||||
},
|
||||
"6": {
|
||||
"class_type": "ViewText",
|
||||
"inputs": {
|
||||
"text": [
|
||||
"3",
|
||||
7
|
||||
]
|
||||
}
|
||||
},
|
||||
"7": {
|
||||
"class_type": "ViewText",
|
||||
"inputs": {
|
||||
"text": [
|
||||
"3",
|
||||
5
|
||||
]
|
||||
}
|
||||
},
|
||||
"8": {
|
||||
"class_type": "PreviewImage",
|
||||
"inputs": {
|
||||
"images": [
|
||||
"3",
|
||||
3
|
||||
]
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,98 @@
|
||||
{
|
||||
"1": {
|
||||
"class_type": "LoadVideo",
|
||||
"inputs": {
|
||||
"file": "vlm_api_people_birds.mp4"
|
||||
}
|
||||
},
|
||||
"2": {
|
||||
"class_type": "GetVideoComponents",
|
||||
"inputs": {
|
||||
"video": [
|
||||
"1",
|
||||
0
|
||||
]
|
||||
}
|
||||
},
|
||||
"3": {
|
||||
"class_type": "VLMPerformanceProfile",
|
||||
"inputs": {
|
||||
"profile": "Fast video"
|
||||
}
|
||||
},
|
||||
"4": {
|
||||
"class_type": "VLMAdaptiveFrameSampler",
|
||||
"inputs": {
|
||||
"frames": [
|
||||
"2",
|
||||
0
|
||||
],
|
||||
"fps": [
|
||||
"2",
|
||||
2
|
||||
],
|
||||
"max_frames": [
|
||||
"3",
|
||||
0
|
||||
],
|
||||
"strategy": "Hybrid: scene + motion + tracks",
|
||||
"minimum_gap_seconds": 0.15,
|
||||
"thumbnail_size": 96
|
||||
}
|
||||
},
|
||||
"5": {
|
||||
"class_type": "VLMImagePixelBudget",
|
||||
"inputs": {
|
||||
"images": [
|
||||
"4",
|
||||
0
|
||||
],
|
||||
"max_megapixels": [
|
||||
"3",
|
||||
1
|
||||
],
|
||||
"max_edge": [
|
||||
"3",
|
||||
2
|
||||
],
|
||||
"multiple": "14",
|
||||
"resize_quality": "Fast (area)"
|
||||
}
|
||||
},
|
||||
"6": {
|
||||
"class_type": "PreviewImage",
|
||||
"inputs": {
|
||||
"images": [
|
||||
"5",
|
||||
0
|
||||
]
|
||||
}
|
||||
},
|
||||
"7": {
|
||||
"class_type": "ViewText",
|
||||
"inputs": {
|
||||
"text": [
|
||||
"4",
|
||||
3
|
||||
]
|
||||
}
|
||||
},
|
||||
"8": {
|
||||
"class_type": "ViewText",
|
||||
"inputs": {
|
||||
"text": [
|
||||
"5",
|
||||
3
|
||||
]
|
||||
}
|
||||
},
|
||||
"9": {
|
||||
"class_type": "ViewText",
|
||||
"inputs": {
|
||||
"text": [
|
||||
"3",
|
||||
5
|
||||
]
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,279 @@
|
||||
"""Model-agnostic acceleration utilities for image and video VLM workflows.
|
||||
|
||||
These nodes reduce visual work *before* it reaches a model. They are therefore
|
||||
portable across Transformers, llama.cpp, Photon, hosted APIs, CUDA, ROCm, MPS,
|
||||
XPU, and CPU runtimes. No model is downloaded and no global PyTorch setting is
|
||||
changed when this module is imported or executed.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import math
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as functional
|
||||
|
||||
RESIZE_QUALITY = (
|
||||
"Fast (area)",
|
||||
"Quality (bicubic)",
|
||||
)
|
||||
PERFORMANCE_PROFILES = {
|
||||
"Live / robotics": {
|
||||
"max_frames": 24,
|
||||
"max_megapixels": 0.5,
|
||||
"max_edge": 896,
|
||||
"batch_size": 8,
|
||||
"unload_after": False,
|
||||
},
|
||||
"Fast video": {
|
||||
"max_frames": 48,
|
||||
"max_megapixels": 0.75,
|
||||
"max_edge": 1024,
|
||||
"batch_size": 8,
|
||||
"unload_after": False,
|
||||
},
|
||||
"Balanced": {
|
||||
"max_frames": 64,
|
||||
"max_megapixels": 1.0,
|
||||
"max_edge": 1344,
|
||||
"batch_size": 4,
|
||||
"unload_after": False,
|
||||
},
|
||||
"High detail": {
|
||||
"max_frames": 96,
|
||||
"max_megapixels": 2.0,
|
||||
"max_edge": 2048,
|
||||
"batch_size": 2,
|
||||
"unload_after": False,
|
||||
},
|
||||
"Low VRAM handoff": {
|
||||
"max_frames": 32,
|
||||
"max_megapixels": 0.75,
|
||||
"max_edge": 1024,
|
||||
"batch_size": 1,
|
||||
"unload_after": True,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _json(value: Any) -> str:
|
||||
return json.dumps(
|
||||
value,
|
||||
ensure_ascii=False,
|
||||
allow_nan=False,
|
||||
sort_keys=True,
|
||||
indent=2,
|
||||
)
|
||||
|
||||
|
||||
def _validate_image_batch(images: torch.Tensor) -> tuple[torch.Tensor, bool]:
|
||||
if not isinstance(images, torch.Tensor):
|
||||
raise TypeError("images must be a ComfyUI IMAGE tensor.")
|
||||
single = images.ndim == 3
|
||||
value = images.unsqueeze(0) if single else images
|
||||
if value.ndim != 4:
|
||||
raise ValueError(
|
||||
f"Expected an HWC/BHWC or CHW/BCHW IMAGE tensor, got {tuple(images.shape)}."
|
||||
)
|
||||
if value.shape[-1] in (1, 3, 4):
|
||||
return value, single
|
||||
if value.shape[1] in (1, 3, 4):
|
||||
return value.permute(0, 2, 3, 1), single
|
||||
raise ValueError(f"Unsupported image channel shape: {tuple(images.shape)}.")
|
||||
|
||||
|
||||
def optimize_image_pixels(
|
||||
images: torch.Tensor,
|
||||
*,
|
||||
max_megapixels: float,
|
||||
max_edge: int,
|
||||
multiple: int,
|
||||
resize_quality: str,
|
||||
) -> tuple[torch.Tensor, dict[str, Any]]:
|
||||
"""Downscale a batch once to a bounded visual-token pixel budget."""
|
||||
|
||||
value, single = _validate_image_batch(images)
|
||||
if not math.isfinite(float(max_megapixels)) or max_megapixels <= 0:
|
||||
raise ValueError("max_megapixels must be finite and positive.")
|
||||
if not isinstance(max_edge, int) or max_edge < 32:
|
||||
raise ValueError("max_edge must be at least 32 pixels.")
|
||||
if multiple not in {1, 14, 28, 32}:
|
||||
raise ValueError("multiple must be one of 1, 14, 28, or 32.")
|
||||
if resize_quality not in RESIZE_QUALITY:
|
||||
raise ValueError(f"Unknown resize quality {resize_quality!r}.")
|
||||
|
||||
height, width = int(value.shape[1]), int(value.shape[2])
|
||||
pixel_budget = float(max_megapixels) * 1_000_000
|
||||
scale = min(
|
||||
1.0,
|
||||
float(max_edge) / max(width, height),
|
||||
math.sqrt(pixel_budget / (width * height)),
|
||||
)
|
||||
|
||||
def bounded_dimension(dimension: int) -> int:
|
||||
target = max(1, math.floor(dimension * scale))
|
||||
if multiple == 1 or target < multiple:
|
||||
return target
|
||||
return max(multiple, (target // multiple) * multiple)
|
||||
|
||||
output_width = bounded_dimension(width)
|
||||
output_height = bounded_dimension(height)
|
||||
output = value
|
||||
resized_image = (output_height, output_width) != (height, width)
|
||||
if resized_image:
|
||||
nchw = value.permute(0, 3, 1, 2)
|
||||
if resize_quality == "Fast (area)":
|
||||
resized = functional.interpolate(
|
||||
nchw,
|
||||
size=(output_height, output_width),
|
||||
mode="area",
|
||||
)
|
||||
else:
|
||||
resized = functional.interpolate(
|
||||
nchw,
|
||||
size=(output_height, output_width),
|
||||
mode="bicubic",
|
||||
align_corners=False,
|
||||
antialias=True,
|
||||
)
|
||||
output = resized.permute(0, 2, 3, 1).clamp(0.0, 1.0)
|
||||
report = {
|
||||
"frames": int(value.shape[0]),
|
||||
"input_width": width,
|
||||
"input_height": height,
|
||||
"output_width": output_width,
|
||||
"output_height": output_height,
|
||||
"input_pixels_per_frame": width * height,
|
||||
"output_pixels_per_frame": output_width * output_height,
|
||||
"visual_work_reduction": (
|
||||
(width * height) / max(1, output_width * output_height)
|
||||
),
|
||||
"resized": resized_image,
|
||||
"multiple": multiple,
|
||||
"quality": resize_quality,
|
||||
}
|
||||
if not resized_image:
|
||||
return images, report
|
||||
return (output[0] if single else output), report
|
||||
|
||||
|
||||
class VLMPerformanceProfile:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"profile": (
|
||||
tuple(PERFORMANCE_PROFILES),
|
||||
{"default": "Balanced"},
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("INT", "FLOAT", "INT", "INT", "BOOLEAN", "STRING")
|
||||
RETURN_NAMES = (
|
||||
"max_frames",
|
||||
"max_megapixels",
|
||||
"max_edge",
|
||||
"batch_size",
|
||||
"unload_after",
|
||||
"profile_json",
|
||||
)
|
||||
FUNCTION = "profile"
|
||||
CATEGORY = "VLM Nodes/Performance"
|
||||
DESCRIPTION = (
|
||||
"Portable speed/quality presets for the sampler, pixel optimizer, "
|
||||
"and VLM batch inputs. The profile never changes global runtime state."
|
||||
)
|
||||
|
||||
def profile(self, profile):
|
||||
values = dict(PERFORMANCE_PROFILES[profile])
|
||||
values["profile"] = profile
|
||||
return (
|
||||
values["max_frames"],
|
||||
values["max_megapixels"],
|
||||
values["max_edge"],
|
||||
values["batch_size"],
|
||||
values["unload_after"],
|
||||
_json(values),
|
||||
)
|
||||
|
||||
|
||||
class VLMImagePixelBudget:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
"max_megapixels": (
|
||||
"FLOAT",
|
||||
{"default": 1.0, "min": 0.01, "max": 64.0, "step": 0.05},
|
||||
),
|
||||
"max_edge": (
|
||||
"INT",
|
||||
{"default": 1344, "min": 32, "max": 16384, "step": 14},
|
||||
),
|
||||
"multiple": (
|
||||
("1", "14", "28", "32"),
|
||||
{
|
||||
"default": "14",
|
||||
"tooltip": (
|
||||
"14/28 suit common VLM vision patches; 32 suits "
|
||||
"many detector backbones. Use 1 for arbitrary sizes."
|
||||
),
|
||||
},
|
||||
),
|
||||
"resize_quality": (
|
||||
RESIZE_QUALITY,
|
||||
{"default": "Fast (area)"},
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "INT", "INT", "STRING")
|
||||
RETURN_NAMES = (
|
||||
"optimized_images",
|
||||
"width",
|
||||
"height",
|
||||
"optimization_report",
|
||||
)
|
||||
FUNCTION = "optimize"
|
||||
CATEGORY = "VLM Nodes/Performance"
|
||||
DESCRIPTION = (
|
||||
"Apply one portable pixel budget before any VLM, avoiding repeated "
|
||||
"high-resolution visual-token work while preserving aspect ratio."
|
||||
)
|
||||
|
||||
def optimize(
|
||||
self,
|
||||
images,
|
||||
max_megapixels,
|
||||
max_edge,
|
||||
multiple,
|
||||
resize_quality,
|
||||
):
|
||||
output, report = optimize_image_pixels(
|
||||
images,
|
||||
max_megapixels=float(max_megapixels),
|
||||
max_edge=int(max_edge),
|
||||
multiple=int(multiple),
|
||||
resize_quality=resize_quality,
|
||||
)
|
||||
return (
|
||||
output,
|
||||
report["output_width"],
|
||||
report["output_height"],
|
||||
_json(report),
|
||||
)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"VLMPerformanceProfile": VLMPerformanceProfile,
|
||||
"VLMImagePixelBudget": VLMImagePixelBudget,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"VLMPerformanceProfile": "VLM Performance Profile",
|
||||
"VLMImagePixelBudget": "VLM Image Pixel Budget",
|
||||
}
|
||||
+1
-2
@@ -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,
|
||||
|
||||
+355
-53
@@ -2,7 +2,11 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from numbers import Real
|
||||
|
||||
import torch
|
||||
from PIL import Image, ImageDraw
|
||||
@@ -24,24 +28,173 @@ from .runtime import (
|
||||
|
||||
MODELS = {
|
||||
"Florence-2 base FT (fast)": "florence-community/Florence-2-base-ft",
|
||||
"Florence-2 large FT (recommended)": (
|
||||
"florence-community/Florence-2-large-ft"
|
||||
),
|
||||
"Florence-2 large FT (recommended)": ("florence-community/Florence-2-large-ft"),
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class FlorenceTaskSpec:
|
||||
"""Declarative contract for one official Florence-2 task."""
|
||||
|
||||
token: str
|
||||
input_kind: str
|
||||
output_kind: str
|
||||
|
||||
|
||||
TASKS = {
|
||||
"Caption": "<CAPTION>",
|
||||
"Detailed caption": "<DETAILED_CAPTION>",
|
||||
"More detailed caption": "<MORE_DETAILED_CAPTION>",
|
||||
"OCR": "<OCR>",
|
||||
"OCR with regions": "<OCR_WITH_REGION>",
|
||||
"Object detection": "<OD>",
|
||||
"Dense region caption": "<DENSE_REGION_CAPTION>",
|
||||
"Region proposals": "<REGION_PROPOSAL>",
|
||||
"Referring expression segmentation": "<REFERRING_EXPRESSION_SEGMENTATION>",
|
||||
"Open vocabulary detection": "<OPEN_VOCABULARY_DETECTION>",
|
||||
"Caption": FlorenceTaskSpec("<CAPTION>", "none", "text"),
|
||||
"Detailed caption": FlorenceTaskSpec("<DETAILED_CAPTION>", "none", "text"),
|
||||
"More detailed caption": FlorenceTaskSpec(
|
||||
"<MORE_DETAILED_CAPTION>", "none", "text"
|
||||
),
|
||||
"OCR": FlorenceTaskSpec("<OCR>", "none", "text"),
|
||||
"OCR with regions": FlorenceTaskSpec("<OCR_WITH_REGION>", "none", "quad_boxes"),
|
||||
"Object detection": FlorenceTaskSpec("<OD>", "none", "boxes"),
|
||||
"Dense region caption": FlorenceTaskSpec("<DENSE_REGION_CAPTION>", "none", "boxes"),
|
||||
"Caption to phrase grounding": FlorenceTaskSpec(
|
||||
"<CAPTION_TO_PHRASE_GROUNDING>", "text", "boxes"
|
||||
),
|
||||
"Referring expression segmentation": FlorenceTaskSpec(
|
||||
"<REFERRING_EXPRESSION_SEGMENTATION>", "text", "polygons"
|
||||
),
|
||||
"Region to segmentation": FlorenceTaskSpec(
|
||||
"<REGION_TO_SEGMENTATION>", "region", "polygons"
|
||||
),
|
||||
"Open vocabulary detection": FlorenceTaskSpec(
|
||||
"<OPEN_VOCABULARY_DETECTION>", "text", "mixed"
|
||||
),
|
||||
"Region to category": FlorenceTaskSpec("<REGION_TO_CATEGORY>", "region", "text"),
|
||||
"Region to description": FlorenceTaskSpec(
|
||||
"<REGION_TO_DESCRIPTION>", "region", "text"
|
||||
),
|
||||
"Region to OCR": FlorenceTaskSpec("<REGION_TO_OCR>", "region", "text"),
|
||||
"Region proposals": FlorenceTaskSpec("<REGION_PROPOSAL>", "none", "boxes"),
|
||||
}
|
||||
|
||||
|
||||
def _clean_decoded_text(value):
|
||||
"""Remove generation wrappers without discarding Florence location tokens."""
|
||||
|
||||
text = str(value)
|
||||
for token in ("<s>", "</s>", "<pad>"):
|
||||
text = text.replace(token, "")
|
||||
return text.strip()
|
||||
|
||||
|
||||
def _select_region(region, image_index, batch_size):
|
||||
"""Select one core BOUNDING_BOX for the current image.
|
||||
|
||||
Core primitive boxes are dictionaries. Detection nodes may emit either a
|
||||
flat per-image list or a nested batch list, so both common shapes are
|
||||
accepted while ambiguous multi-region inputs fail explicitly.
|
||||
"""
|
||||
|
||||
if region is None or isinstance(region, dict):
|
||||
return region
|
||||
if not isinstance(region, (list, tuple)):
|
||||
raise TypeError("region must be a core BOUNDING_BOX dictionary.")
|
||||
if not region:
|
||||
return None
|
||||
|
||||
if all(isinstance(item, dict) for item in region):
|
||||
if len(region) == 1:
|
||||
return region[0]
|
||||
if len(region) == batch_size:
|
||||
return region[image_index]
|
||||
raise ValueError("Region tasks require exactly one BOUNDING_BOX per image.")
|
||||
|
||||
if len(region) != batch_size:
|
||||
raise ValueError("Batched BOUNDING_BOX input must contain one entry per image.")
|
||||
frame_regions = region[image_index]
|
||||
if isinstance(frame_regions, dict):
|
||||
return frame_regions
|
||||
if not isinstance(frame_regions, (list, tuple)) or len(frame_regions) != 1:
|
||||
raise ValueError(
|
||||
"Region tasks require exactly one BOUNDING_BOX per image; "
|
||||
"select a detection before connecting it."
|
||||
)
|
||||
if not isinstance(frame_regions[0], dict):
|
||||
raise TypeError("Each BOUNDING_BOX entry must be a dictionary.")
|
||||
return frame_regions[0]
|
||||
|
||||
|
||||
def _encode_region(region, image_size):
|
||||
"""Encode an absolute-pixel core BOUNDING_BOX as Florence location tokens."""
|
||||
|
||||
if not isinstance(region, dict):
|
||||
raise TypeError("region must be a core BOUNDING_BOX dictionary.")
|
||||
|
||||
try:
|
||||
x = float(region["x"])
|
||||
y = float(region["y"])
|
||||
box_width = float(region["width"])
|
||||
box_height = float(region["height"])
|
||||
except KeyError as exc:
|
||||
raise ValueError("region must contain x, y, width, and height.") from exc
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise ValueError("region coordinates must be numeric.") from exc
|
||||
|
||||
values = (x, y, box_width, box_height)
|
||||
if not all(math.isfinite(value) for value in values):
|
||||
raise ValueError("region coordinates must be finite.")
|
||||
if box_width <= 0 or box_height <= 0:
|
||||
raise ValueError("region width and height must be greater than zero.")
|
||||
|
||||
image_width, image_height = image_size
|
||||
if image_width <= 0 or image_height <= 0:
|
||||
raise ValueError("image dimensions must be greater than zero.")
|
||||
|
||||
x0 = max(0.0, min(float(image_width), x))
|
||||
y0 = max(0.0, min(float(image_height), y))
|
||||
x1 = max(0.0, min(float(image_width), x + box_width))
|
||||
y1 = max(0.0, min(float(image_height), y + box_height))
|
||||
if x1 <= x0 or y1 <= y0:
|
||||
raise ValueError("region does not overlap the input image.")
|
||||
|
||||
coordinates = (
|
||||
x0 / image_width,
|
||||
y0 / image_height,
|
||||
x1 / image_width,
|
||||
y1 / image_height,
|
||||
)
|
||||
bins = [
|
||||
max(0, min(999, math.floor(coordinate * 1000))) for coordinate in coordinates
|
||||
]
|
||||
return "".join(f"<loc_{value}>" for value in bins)
|
||||
|
||||
|
||||
def _task_extra_input(task_name, text_input, region, image_size):
|
||||
"""Validate and prepare the optional suffix for a Florence task prompt."""
|
||||
|
||||
try:
|
||||
spec = TASKS[task_name]
|
||||
except KeyError as exc:
|
||||
raise ValueError(f"Unsupported Florence-2 task: {task_name}") from exc
|
||||
|
||||
text = (text_input or "").strip()
|
||||
if spec.input_kind == "none":
|
||||
if text:
|
||||
raise ValueError(f"{task_name} does not accept text input.")
|
||||
if region is not None:
|
||||
raise ValueError(f"{task_name} does not accept a region input.")
|
||||
return ""
|
||||
if spec.input_kind == "text":
|
||||
if not text:
|
||||
raise ValueError(f"{task_name} requires text input.")
|
||||
if region is not None:
|
||||
raise ValueError(f"{task_name} does not accept a region input.")
|
||||
return text
|
||||
if spec.input_kind == "region":
|
||||
if text:
|
||||
raise ValueError(
|
||||
f"{task_name} uses the region input and does not accept text."
|
||||
)
|
||||
if region is None:
|
||||
raise ValueError(f"{task_name} requires a connected BOUNDING_BOX region.")
|
||||
return _encode_region(region, image_size)
|
||||
raise RuntimeError(f"Unknown Florence task input kind: {spec.input_kind}")
|
||||
|
||||
|
||||
class FlorencePredictor:
|
||||
def __init__(self, model_label):
|
||||
transformers = require_module("transformers")
|
||||
@@ -66,9 +219,7 @@ class FlorencePredictor:
|
||||
|
||||
def run(self, image, task_token, text, max_new_tokens, beams):
|
||||
prompt = task_token + (text.strip() if text.strip() else "")
|
||||
inputs = self.processor(
|
||||
text=prompt, images=image, return_tensors="pt"
|
||||
)
|
||||
inputs = self.processor(text=prompt, images=image, return_tensors="pt")
|
||||
model = self.handle.ensure_loaded()
|
||||
device = model_device(model)
|
||||
inputs = move_inputs(inputs, device, floating_dtype=self.dtype)
|
||||
@@ -80,9 +231,7 @@ class FlorencePredictor:
|
||||
do_sample=False,
|
||||
early_stopping=int(beams) > 1,
|
||||
)
|
||||
raw = self.processor.batch_decode(
|
||||
generated, skip_special_tokens=False
|
||||
)[0]
|
||||
raw = self.processor.batch_decode(generated, skip_special_tokens=False)[0]
|
||||
parsed = self.processor.post_process_generation(
|
||||
raw, task=task_token, image_size=image.size
|
||||
)
|
||||
@@ -95,39 +244,156 @@ def _json_default(value):
|
||||
return str(value)
|
||||
|
||||
|
||||
_SPATIAL_KEYS = frozenset(
|
||||
{
|
||||
"bboxes",
|
||||
"quad_boxes",
|
||||
"polygons",
|
||||
"labels",
|
||||
"bboxes_labels",
|
||||
"polygons_labels",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _spatial_result(parsed):
|
||||
if not isinstance(parsed, dict):
|
||||
return {}
|
||||
if _SPATIAL_KEYS.intersection(parsed):
|
||||
return parsed
|
||||
result = next(iter(parsed.values()), {})
|
||||
return result if isinstance(result, dict) else {}
|
||||
|
||||
|
||||
def _stable_color(kind, index, label):
|
||||
key = f"{kind}:{index}:{label}".encode("utf-8", errors="replace")
|
||||
digest = hashlib.blake2b(key, digest_size=3).digest()
|
||||
return tuple(64 + channel % 192 for channel in digest)
|
||||
|
||||
|
||||
def _points(values, image_size):
|
||||
if not isinstance(values, (list, tuple)) or len(values) < 6:
|
||||
return []
|
||||
width, height = image_size
|
||||
points = []
|
||||
for index in range(0, len(values) - 1, 2):
|
||||
x, y = values[index], values[index + 1]
|
||||
if not isinstance(x, Real) or not isinstance(y, Real):
|
||||
return []
|
||||
if not math.isfinite(float(x)) or not math.isfinite(float(y)):
|
||||
return []
|
||||
points.append(
|
||||
(
|
||||
max(0, min(width - 1, round(float(x)))),
|
||||
max(0, min(height - 1, round(float(y)))),
|
||||
)
|
||||
)
|
||||
return points
|
||||
|
||||
|
||||
def _box(values, image_size):
|
||||
if not isinstance(values, (list, tuple)) or len(values) < 4:
|
||||
return None
|
||||
if not all(isinstance(value, Real) for value in values[:4]):
|
||||
return None
|
||||
coordinates = [float(value) for value in values[:4]]
|
||||
if not all(math.isfinite(value) for value in coordinates):
|
||||
return None
|
||||
x0, y0, x1, y1 = coordinates
|
||||
x0, x1 = sorted((x0, x1))
|
||||
y0, y1 = sorted((y0, y1))
|
||||
width, height = image_size
|
||||
x0 = max(0, min(width - 1, round(x0)))
|
||||
x1 = max(0, min(width - 1, round(x1)))
|
||||
y0 = max(0, min(height - 1, round(y0)))
|
||||
y1 = max(0, min(height - 1, round(y1)))
|
||||
if x1 <= x0 or y1 <= y0:
|
||||
return None
|
||||
return x0, y0, x1, y1
|
||||
|
||||
|
||||
def _polygon_list(group):
|
||||
if not isinstance(group, (list, tuple)) or not group:
|
||||
return []
|
||||
if isinstance(group[0], Real):
|
||||
return [group]
|
||||
return [item for item in group if isinstance(item, (list, tuple))]
|
||||
|
||||
|
||||
def _label_with_score(labels, scores, index):
|
||||
label = str(labels[index]) if index < len(labels) else ""
|
||||
if index < len(scores) and isinstance(scores[index], Real):
|
||||
score = f"{float(scores[index]):.3f}"
|
||||
return f"{label} {score}".strip()
|
||||
return label
|
||||
|
||||
|
||||
def _draw_label(draw, position, text, color, image_size):
|
||||
if not text:
|
||||
return
|
||||
x, y = position
|
||||
try:
|
||||
left, top, right, bottom = draw.textbbox((0, 0), text)
|
||||
text_width, text_height = right - left, bottom - top
|
||||
except AttributeError:
|
||||
text_width, text_height = draw.textlength(text), 11
|
||||
width, height = image_size
|
||||
x = max(0, min(width - text_width - 4, x))
|
||||
y = max(0, min(height - text_height - 4, y))
|
||||
background = (0, 0, 0) if sum(color) > 360 else (255, 255, 255)
|
||||
foreground = (255, 255, 255) if background == (0, 0, 0) else (0, 0, 0)
|
||||
draw.rectangle(
|
||||
(x, y, x + text_width + 4, y + text_height + 4),
|
||||
fill=background,
|
||||
)
|
||||
draw.text((x + 2, y + 2), text, fill=foreground)
|
||||
|
||||
|
||||
def _visualize(image, parsed):
|
||||
result = next(iter(parsed.values()), parsed) if isinstance(parsed, dict) else {}
|
||||
result = _spatial_result(parsed)
|
||||
mask = Image.new("L", image.size, 0)
|
||||
visual = image.copy().convert("RGB")
|
||||
mask_draw = ImageDraw.Draw(mask)
|
||||
draw = ImageDraw.Draw(visual)
|
||||
labels = result.get("labels", []) if isinstance(result, dict) else []
|
||||
width = max(2, min(8, round(min(image.size) / 256 * 3)))
|
||||
labels = result.get("labels", [])
|
||||
scores = result.get("scores", [])
|
||||
|
||||
for index, box in enumerate(result.get("bboxes", [])):
|
||||
box = [float(value) for value in box]
|
||||
draw.rectangle(box, outline="#00ff88", width=3)
|
||||
if index < len(labels):
|
||||
draw.text((box[0] + 3, box[1] + 3), str(labels[index]), fill="#00ff88")
|
||||
box_labels = result.get("bboxes_labels", labels)
|
||||
for index, values in enumerate(result.get("bboxes", [])):
|
||||
box = _box(values, image.size)
|
||||
if box is None:
|
||||
continue
|
||||
label = _label_with_score(box_labels, scores, index)
|
||||
color = _stable_color("box", index, label)
|
||||
mask_draw.rectangle(box, fill=255)
|
||||
draw.rectangle(box, outline=color, width=width)
|
||||
_draw_label(draw, (box[0], box[1]), label, color, image.size)
|
||||
|
||||
for quad in result.get("quad_boxes", []):
|
||||
points = [
|
||||
(float(quad[index]), float(quad[index + 1]))
|
||||
for index in range(0, len(quad), 2)
|
||||
]
|
||||
draw.line(points + [points[0]], fill="#00c8ff", width=3)
|
||||
for index, values in enumerate(result.get("quad_boxes", [])):
|
||||
points = _points(values, image.size)
|
||||
if len(points) < 3:
|
||||
continue
|
||||
label = _label_with_score(labels, scores, index)
|
||||
color = _stable_color("quad", index, label)
|
||||
mask_draw.polygon(points, fill=255)
|
||||
draw.line(points + [points[0]], fill=color, width=width)
|
||||
_draw_label(draw, points[0], label, color, image.size)
|
||||
|
||||
polygons = result.get("polygons", [])
|
||||
for group in polygons:
|
||||
# Florence may return either one flat polygon or a list of polygons.
|
||||
groups = [group] if group and isinstance(group[0], (int, float)) else group
|
||||
for polygon in groups:
|
||||
points = [
|
||||
(float(polygon[index]), float(polygon[index + 1]))
|
||||
for index in range(0, len(polygon), 2)
|
||||
]
|
||||
if len(points) >= 3:
|
||||
mask_draw.polygon(points, fill=255)
|
||||
draw.line(points + [points[0]], fill="#ff4da6", width=3)
|
||||
polygon_labels = result.get("polygons_labels", labels)
|
||||
for index, group in enumerate(result.get("polygons", [])):
|
||||
label = _label_with_score(polygon_labels, scores, index)
|
||||
color = _stable_color("polygon", index, label)
|
||||
label_drawn = False
|
||||
for polygon in _polygon_list(group):
|
||||
points = _points(polygon, image.size)
|
||||
if len(points) < 3:
|
||||
continue
|
||||
mask_draw.polygon(points, fill=255)
|
||||
draw.line(points + [points[0]], fill=color, width=width)
|
||||
if not label_drawn:
|
||||
_draw_label(draw, points[0], label, color, image.size)
|
||||
label_drawn = True
|
||||
return mask, visual
|
||||
|
||||
|
||||
@@ -143,7 +409,10 @@ class Florence2(CachedModelNode):
|
||||
{
|
||||
"default": "",
|
||||
"multiline": True,
|
||||
"tooltip": "Required for referring-expression and open-vocabulary tasks.",
|
||||
"tooltip": (
|
||||
"Required only for phrase grounding, referring-expression "
|
||||
"segmentation, and open-vocabulary detection."
|
||||
),
|
||||
},
|
||||
),
|
||||
"model": (
|
||||
@@ -158,6 +427,15 @@ class Florence2(CachedModelNode):
|
||||
},
|
||||
"optional": {
|
||||
"unload_after": ("BOOLEAN", {"default": False}),
|
||||
"region": (
|
||||
"BOUNDING_BOX",
|
||||
{
|
||||
"tooltip": (
|
||||
"Core bounding box input required by Region to "
|
||||
"Segmentation/Category/Description/OCR."
|
||||
)
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -175,28 +453,52 @@ class Florence2(CachedModelNode):
|
||||
max_new_tokens,
|
||||
beams,
|
||||
unload_after=False,
|
||||
region=None,
|
||||
):
|
||||
predictor = self.get_or_create_model(
|
||||
model, lambda: FlorencePredictor(model)
|
||||
)
|
||||
images = tensor_batch_to_pil(image)
|
||||
if not images:
|
||||
raise ValueError("Florence-2 requires at least one input image.")
|
||||
try:
|
||||
spec = TASKS[task]
|
||||
except KeyError as exc:
|
||||
raise ValueError(f"Unsupported Florence-2 task: {task}") from exc
|
||||
|
||||
extra_inputs = []
|
||||
for index, pil_image in enumerate(images):
|
||||
selected_region = _select_region(region, index, len(images))
|
||||
extra_inputs.append(
|
||||
_task_extra_input(
|
||||
task,
|
||||
text_input,
|
||||
selected_region,
|
||||
pil_image.size,
|
||||
)
|
||||
)
|
||||
|
||||
predictor = self.get_or_create_model(model, lambda: FlorencePredictor(model))
|
||||
texts, records, masks, visuals = [], [], [], []
|
||||
try:
|
||||
for pil_image in tensor_batch_to_pil(image):
|
||||
for pil_image, extra_input in zip(images, extra_inputs):
|
||||
raw, parsed = predictor.run(
|
||||
pil_image,
|
||||
TASKS[task],
|
||||
text_input,
|
||||
spec.token,
|
||||
extra_input,
|
||||
max_new_tokens,
|
||||
beams,
|
||||
)
|
||||
texts.append(raw)
|
||||
texts.append(_clean_decoded_text(raw))
|
||||
records.append(parsed)
|
||||
mask, visual = _visualize(pil_image, parsed)
|
||||
masks.append(pil_mask_to_tensor(mask))
|
||||
visuals.append(pil_to_tensor(visual))
|
||||
return (
|
||||
batch_text(texts),
|
||||
json.dumps(records, ensure_ascii=False, default=_json_default),
|
||||
json.dumps(
|
||||
records,
|
||||
ensure_ascii=False,
|
||||
default=_json_default,
|
||||
sort_keys=True,
|
||||
),
|
||||
torch.cat(masks),
|
||||
torch.cat(visuals),
|
||||
)
|
||||
|
||||
@@ -0,0 +1,413 @@
|
||||
"""Dependency-light geometry, mask, color, and association primitives."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import colorsys
|
||||
import hashlib
|
||||
import math
|
||||
from collections.abc import Iterable, Mapping
|
||||
from dataclasses import dataclass
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image, ImageDraw
|
||||
|
||||
from .vision_types import BoxXYXY, Detection, PointXY, Polygon
|
||||
|
||||
|
||||
def _dimensions(width: int, height: int) -> tuple[int, int]:
|
||||
if not isinstance(width, int) or width <= 0:
|
||||
raise ValueError("width must be a positive integer.")
|
||||
if not isinstance(height, int) or height <= 0:
|
||||
raise ValueError("height must be a positive integer.")
|
||||
return width, height
|
||||
|
||||
|
||||
def _ordered_box(box: Iterable[float]) -> BoxXYXY:
|
||||
values = tuple(float(value) for value in box)
|
||||
if len(values) != 4 or not all(math.isfinite(value) for value in values):
|
||||
raise ValueError("A box must contain four finite xyxy values.")
|
||||
x1, y1, x2, y2 = values
|
||||
if x2 < x1 or y2 < y1:
|
||||
raise ValueError("A box must satisfy x2 >= x1 and y2 >= y1.")
|
||||
return x1, y1, x2, y2
|
||||
|
||||
|
||||
def clip_box(box: Iterable[float], width: int, height: int) -> BoxXYXY:
|
||||
"""Clamp a pixel xyxy box to an image, preserving exclusive x2/y2."""
|
||||
|
||||
width, height = _dimensions(width, height)
|
||||
x1, y1, x2, y2 = _ordered_box(box)
|
||||
return (
|
||||
min(max(x1, 0.0), float(width)),
|
||||
min(max(y1, 0.0), float(height)),
|
||||
min(max(x2, 0.0), float(width)),
|
||||
min(max(y2, 0.0), float(height)),
|
||||
)
|
||||
|
||||
|
||||
def clip_polygon(
|
||||
polygon: Iterable[Iterable[float]],
|
||||
width: int,
|
||||
height: int,
|
||||
) -> Polygon:
|
||||
width, height = _dimensions(width, height)
|
||||
points = []
|
||||
for point in polygon:
|
||||
values = tuple(float(value) for value in point)
|
||||
if len(values) != 2 or not all(math.isfinite(value) for value in values):
|
||||
raise ValueError("Polygon points must contain two finite values.")
|
||||
points.append(
|
||||
(
|
||||
min(max(values[0], 0.0), float(width)),
|
||||
min(max(values[1], 0.0), float(height)),
|
||||
)
|
||||
)
|
||||
if len(points) < 3:
|
||||
raise ValueError("A polygon requires at least three points.")
|
||||
return tuple(points)
|
||||
|
||||
|
||||
def normalize_box(
|
||||
box: Iterable[float],
|
||||
width: int,
|
||||
height: int,
|
||||
) -> BoxXYXY:
|
||||
width, height = _dimensions(width, height)
|
||||
x1, y1, x2, y2 = clip_box(box, width, height)
|
||||
return x1 / width, y1 / height, x2 / width, y2 / height
|
||||
|
||||
|
||||
def denormalize_box(
|
||||
box: Iterable[float],
|
||||
width: int,
|
||||
height: int,
|
||||
) -> BoxXYXY:
|
||||
width, height = _dimensions(width, height)
|
||||
x1, y1, x2, y2 = _ordered_box(box)
|
||||
if any(value < 0.0 or value > 1.0 for value in (x1, y1, x2, y2)):
|
||||
raise ValueError("Normalized box coordinates must be between 0 and 1.")
|
||||
return x1 * width, y1 * height, x2 * width, y2 * height
|
||||
|
||||
|
||||
def box_area(box: Iterable[float]) -> float:
|
||||
x1, y1, x2, y2 = _ordered_box(box)
|
||||
return (x2 - x1) * (y2 - y1)
|
||||
|
||||
|
||||
def box_center(box: Iterable[float]) -> PointXY:
|
||||
x1, y1, x2, y2 = _ordered_box(box)
|
||||
return (x1 + x2) * 0.5, (y1 + y2) * 0.5
|
||||
|
||||
|
||||
def polygon_area(polygon: Iterable[Iterable[float]]) -> float:
|
||||
points = [tuple(float(value) for value in point) for point in polygon]
|
||||
if len(points) < 3 or any(len(point) != 2 for point in points):
|
||||
raise ValueError("A polygon requires at least three xy points.")
|
||||
if any(not math.isfinite(value) for point in points for value in point):
|
||||
raise ValueError("Polygon coordinates must be finite.")
|
||||
twice_area = sum(
|
||||
x1 * y2 - x2 * y1 for (x1, y1), (x2, y2) in zip(points, points[1:] + points[:1])
|
||||
)
|
||||
return abs(twice_area) * 0.5
|
||||
|
||||
|
||||
def bbox_iou(first: Iterable[float], second: Iterable[float]) -> float:
|
||||
ax1, ay1, ax2, ay2 = _ordered_box(first)
|
||||
bx1, by1, bx2, by2 = _ordered_box(second)
|
||||
intersection = max(0.0, min(ax2, bx2) - max(ax1, bx1)) * max(
|
||||
0.0, min(ay2, by2) - max(ay1, by1)
|
||||
)
|
||||
union = box_area(first) + box_area(second) - intersection
|
||||
return intersection / union if union > 0 else 0.0
|
||||
|
||||
|
||||
def mask_iou(
|
||||
first: torch.Tensor | np.ndarray,
|
||||
second: torch.Tensor | np.ndarray,
|
||||
*,
|
||||
threshold: float = 0.5,
|
||||
) -> float:
|
||||
first_tensor = torch.as_tensor(first)
|
||||
second_tensor = torch.as_tensor(second)
|
||||
if first_tensor.ndim != 2 or second_tensor.ndim != 2:
|
||||
raise ValueError("Masks must have shape [height, width].")
|
||||
if first_tensor.shape != second_tensor.shape:
|
||||
raise ValueError("Masks must have the same shape.")
|
||||
first_bool = first_tensor > float(threshold)
|
||||
second_bool = second_tensor > float(threshold)
|
||||
intersection = torch.logical_and(first_bool, second_bool).sum().item()
|
||||
union = torch.logical_or(first_bool, second_bool).sum().item()
|
||||
return float(intersection / union) if union else 0.0
|
||||
|
||||
|
||||
def deterministic_color(value: object) -> tuple[int, int, int]:
|
||||
"""Return a readable RGB color that is stable across Python processes."""
|
||||
|
||||
digest = hashlib.sha256(str(value).encode("utf-8")).digest()
|
||||
hue = int.from_bytes(digest[:2], "big") / 65535.0
|
||||
saturation = 0.62 + digest[2] / 255.0 * 0.22
|
||||
brightness = 0.78 + digest[3] / 255.0 * 0.17
|
||||
return tuple(
|
||||
round(channel * 255)
|
||||
for channel in colorsys.hsv_to_rgb(hue, saturation, brightness)
|
||||
)
|
||||
|
||||
|
||||
def box_to_mask(
|
||||
box: Iterable[float],
|
||||
width: int,
|
||||
height: int,
|
||||
) -> torch.Tensor:
|
||||
width, height = _dimensions(width, height)
|
||||
x1, y1, x2, y2 = clip_box(box, width, height)
|
||||
left = max(0, min(width, math.floor(x1)))
|
||||
top = max(0, min(height, math.floor(y1)))
|
||||
right = max(left, min(width, math.ceil(x2)))
|
||||
bottom = max(top, min(height, math.ceil(y2)))
|
||||
mask = torch.zeros((height, width), dtype=torch.float32)
|
||||
mask[top:bottom, left:right] = 1.0
|
||||
return mask
|
||||
|
||||
|
||||
def polygon_to_mask(
|
||||
polygon: Iterable[Iterable[float]],
|
||||
width: int,
|
||||
height: int,
|
||||
) -> torch.Tensor:
|
||||
width, height = _dimensions(width, height)
|
||||
points = clip_polygon(polygon, width, height)
|
||||
canvas = Image.new("L", (width, height), 0)
|
||||
ImageDraw.Draw(canvas).polygon(points, fill=255)
|
||||
array = np.asarray(canvas, dtype=np.float32) / 255.0
|
||||
return torch.from_numpy(array.copy())
|
||||
|
||||
|
||||
def quad_to_mask(
|
||||
quad: Iterable[Iterable[float]],
|
||||
width: int,
|
||||
height: int,
|
||||
) -> torch.Tensor:
|
||||
points = tuple(tuple(point) for point in quad)
|
||||
if len(points) != 4:
|
||||
raise ValueError("A quad must contain exactly four points.")
|
||||
return polygon_to_mask(points, width, height)
|
||||
|
||||
|
||||
def detection_to_mask(
|
||||
detection: Detection,
|
||||
width: int,
|
||||
height: int,
|
||||
) -> torch.Tensor:
|
||||
"""Rasterize the most precise geometry available on a detection."""
|
||||
|
||||
width, height = _dimensions(width, height)
|
||||
if not isinstance(detection, Detection):
|
||||
raise TypeError("detection must be a Detection.")
|
||||
if detection.mask is not None:
|
||||
if tuple(detection.mask.shape) != (height, width):
|
||||
raise ValueError("Detection mask shape does not match the image.")
|
||||
return detection.mask.detach().to(dtype=torch.float32).clamp(0, 1).clone()
|
||||
if detection.polygon is not None:
|
||||
return polygon_to_mask(detection.polygon, width, height)
|
||||
if detection.quad is not None:
|
||||
return quad_to_mask(detection.quad, width, height)
|
||||
return box_to_mask(detection.bbox_xyxy, width, height)
|
||||
|
||||
|
||||
def individual_detection_masks(
|
||||
detections: Iterable[Detection],
|
||||
width: int,
|
||||
height: int,
|
||||
) -> torch.Tensor:
|
||||
width, height = _dimensions(width, height)
|
||||
masks = [detection_to_mask(detection, width, height) for detection in detections]
|
||||
if not masks:
|
||||
return torch.zeros((0, height, width), dtype=torch.float32)
|
||||
return torch.stack(masks).to(dtype=torch.float32)
|
||||
|
||||
|
||||
def union_detection_mask(
|
||||
detections: Iterable[Detection],
|
||||
width: int,
|
||||
height: int,
|
||||
) -> torch.Tensor:
|
||||
masks = individual_detection_masks(detections, width, height)
|
||||
if masks.shape[0] == 0:
|
||||
return torch.zeros((height, width), dtype=torch.float32)
|
||||
return masks.amax(dim=0).clamp(0, 1)
|
||||
|
||||
|
||||
def bbox_from_mask(
|
||||
mask: torch.Tensor | np.ndarray,
|
||||
*,
|
||||
threshold: float = 0.5,
|
||||
) -> BoxXYXY | None:
|
||||
value = torch.as_tensor(mask)
|
||||
if value.ndim != 2:
|
||||
raise ValueError("mask must have shape [height, width].")
|
||||
locations = torch.nonzero(value > float(threshold), as_tuple=False)
|
||||
if locations.numel() == 0:
|
||||
return None
|
||||
y1, x1 = locations.amin(dim=0).tolist()
|
||||
y2, x2 = locations.amax(dim=0).tolist()
|
||||
return float(x1), float(y1), float(x2 + 1), float(y2 + 1)
|
||||
|
||||
|
||||
def translate_box(
|
||||
box: Iterable[float],
|
||||
dx: float,
|
||||
dy: float,
|
||||
) -> BoxXYXY:
|
||||
x1, y1, x2, y2 = _ordered_box(box)
|
||||
dx = float(dx)
|
||||
dy = float(dy)
|
||||
if not math.isfinite(dx) or not math.isfinite(dy):
|
||||
raise ValueError("Box motion must be finite.")
|
||||
return x1 + dx, y1 + dy, x2 + dx, y2 + dy
|
||||
|
||||
|
||||
def expand_box(
|
||||
box: Iterable[float],
|
||||
width: int,
|
||||
height: int,
|
||||
*,
|
||||
padding: float = 0.0,
|
||||
square: bool = False,
|
||||
) -> BoxXYXY:
|
||||
"""Pad and optionally square a box around its center, then clip it."""
|
||||
|
||||
width, height = _dimensions(width, height)
|
||||
if not math.isfinite(float(padding)) or padding < 0:
|
||||
raise ValueError("padding must be finite and non-negative.")
|
||||
x1, y1, x2, y2 = _ordered_box(box)
|
||||
x1 -= padding
|
||||
y1 -= padding
|
||||
x2 += padding
|
||||
y2 += padding
|
||||
if square:
|
||||
center_x, center_y = (x1 + x2) * 0.5, (y1 + y2) * 0.5
|
||||
half = max(x2 - x1, y2 - y1) * 0.5
|
||||
x1, y1, x2, y2 = (
|
||||
center_x - half,
|
||||
center_y - half,
|
||||
center_x + half,
|
||||
center_y + half,
|
||||
)
|
||||
side = x2 - x1
|
||||
if side <= width:
|
||||
if x1 < 0:
|
||||
x2 -= x1
|
||||
x1 = 0.0
|
||||
elif x2 > width:
|
||||
x1 -= x2 - width
|
||||
x2 = float(width)
|
||||
if side <= height:
|
||||
if y1 < 0:
|
||||
y2 -= y1
|
||||
y1 = 0.0
|
||||
elif y2 > height:
|
||||
y1 -= y2 - height
|
||||
y2 = float(height)
|
||||
return clip_box((x1, y1, x2, y2), width, height)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AssociationResult:
|
||||
"""Stable one-to-one detection assignment by descending overlap."""
|
||||
|
||||
matches: tuple[tuple[int, int, float], ...]
|
||||
unmatched_previous: tuple[int, ...]
|
||||
unmatched_current: tuple[int, ...]
|
||||
|
||||
|
||||
def associate_detections(
|
||||
previous: Iterable[Detection],
|
||||
current: Iterable[Detection],
|
||||
*,
|
||||
minimum_iou: float = 0.3,
|
||||
label_aware: bool = True,
|
||||
motion_by_track: Mapping[int, tuple[float, float]] | None = None,
|
||||
) -> AssociationResult:
|
||||
"""Associate detections without SciPy or backend-specific operators.
|
||||
|
||||
Candidates are greedily selected by descending IoU with deterministic
|
||||
index tie-breaks. Optional per-track motion offsets predict the previous
|
||||
box before overlap is measured.
|
||||
"""
|
||||
|
||||
previous_items = tuple(previous)
|
||||
current_items = tuple(current)
|
||||
if not 0.0 <= float(minimum_iou) <= 1.0:
|
||||
raise ValueError("minimum_iou must be between 0 and 1.")
|
||||
if any(not isinstance(item, Detection) for item in previous_items):
|
||||
raise TypeError("previous must contain Detection values.")
|
||||
if any(not isinstance(item, Detection) for item in current_items):
|
||||
raise TypeError("current must contain Detection values.")
|
||||
|
||||
candidates = []
|
||||
for previous_index, old in enumerate(previous_items):
|
||||
old_box = old.bbox_xyxy
|
||||
if old.track_id is not None and motion_by_track:
|
||||
motion = motion_by_track.get(old.track_id)
|
||||
if motion is not None:
|
||||
old_box = translate_box(old_box, motion[0], motion[1])
|
||||
for current_index, new in enumerate(current_items):
|
||||
if (
|
||||
label_aware
|
||||
and old.label is not None
|
||||
and new.label is not None
|
||||
and " ".join(old.label.casefold().split())
|
||||
!= " ".join(new.label.casefold().split())
|
||||
):
|
||||
continue
|
||||
overlap = bbox_iou(old_box, new.bbox_xyxy)
|
||||
if overlap >= float(minimum_iou):
|
||||
candidates.append((-overlap, previous_index, current_index, overlap))
|
||||
|
||||
matched_previous: set[int] = set()
|
||||
matched_current: set[int] = set()
|
||||
matches = []
|
||||
for _negative, previous_index, current_index, overlap in sorted(candidates):
|
||||
if previous_index in matched_previous or current_index in matched_current:
|
||||
continue
|
||||
matched_previous.add(previous_index)
|
||||
matched_current.add(current_index)
|
||||
matches.append((previous_index, current_index, overlap))
|
||||
|
||||
return AssociationResult(
|
||||
matches=tuple(matches),
|
||||
unmatched_previous=tuple(
|
||||
index
|
||||
for index in range(len(previous_items))
|
||||
if index not in matched_previous
|
||||
),
|
||||
unmatched_current=tuple(
|
||||
index for index in range(len(current_items)) if index not in matched_current
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"AssociationResult",
|
||||
"associate_detections",
|
||||
"bbox_from_mask",
|
||||
"bbox_iou",
|
||||
"box_area",
|
||||
"box_center",
|
||||
"box_to_mask",
|
||||
"clip_box",
|
||||
"clip_polygon",
|
||||
"denormalize_box",
|
||||
"detection_to_mask",
|
||||
"deterministic_color",
|
||||
"expand_box",
|
||||
"individual_detection_masks",
|
||||
"mask_iou",
|
||||
"normalize_box",
|
||||
"polygon_area",
|
||||
"polygon_to_mask",
|
||||
"quad_to_mask",
|
||||
"translate_box",
|
||||
"union_detection_mask",
|
||||
]
|
||||
@@ -0,0 +1,537 @@
|
||||
"""Fast open-vocabulary object detection with maintained Transformers models.
|
||||
|
||||
The node deliberately presents one stable ComfyUI interface while keeping
|
||||
model-specific preprocessing and postprocessing behind a small adapter. Model
|
||||
downloads are lazy, inference participates in ComfyUI's VRAM management, and
|
||||
all spatial output uses the pack's versioned pixel-coordinate contract.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import ImageDraw
|
||||
|
||||
from .geometry import deterministic_color
|
||||
from .runtime import (
|
||||
CachedModelNode,
|
||||
ManagedTorchModel,
|
||||
inference_context,
|
||||
model_device,
|
||||
move_inputs,
|
||||
require_module,
|
||||
snapshot_download,
|
||||
tensor_batch_to_pil,
|
||||
torch_dtype,
|
||||
)
|
||||
from .vision_types import (
|
||||
VLM_DETECTIONS,
|
||||
Detection,
|
||||
DetectionSequence,
|
||||
FrameDetections,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class DetectorSpec:
|
||||
model_id: str
|
||||
cache_name: str
|
||||
family: str
|
||||
description: str
|
||||
|
||||
|
||||
MODEL_SPECS = {
|
||||
"Grounding DINO Tiny (fast)": DetectorSpec(
|
||||
"IDEA-Research/grounding-dino-tiny",
|
||||
"grounding-dino-tiny",
|
||||
"grounding_dino",
|
||||
"Fast, accurate open-vocabulary grounding.",
|
||||
),
|
||||
"Grounding DINO Base": DetectorSpec(
|
||||
"IDEA-Research/grounding-dino-base",
|
||||
"grounding-dino-base",
|
||||
"grounding_dino",
|
||||
"Higher-quality open-vocabulary grounding.",
|
||||
),
|
||||
"OWLv2 Base Ensemble": DetectorSpec(
|
||||
"google/owlv2-base-patch16-ensemble",
|
||||
"owlv2-base-patch16-ensemble",
|
||||
"owlv2",
|
||||
"Strong zero-shot detector for lists of visual concepts.",
|
||||
),
|
||||
"OmDet Turbo Swin Tiny (fast)": DetectorSpec(
|
||||
"omlab/omdet-turbo-swin-tiny-hf",
|
||||
"omdet-turbo-swin-tiny",
|
||||
"omdet",
|
||||
"Efficient real-time-oriented open-vocabulary detector.",
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def parse_labels(value: str) -> list[str]:
|
||||
"""Parse user concepts without splitting meaningful multi-word labels."""
|
||||
|
||||
labels: list[str] = []
|
||||
for line in str(value or "").replace(";", "\n").splitlines():
|
||||
for candidate in line.split(","):
|
||||
label = " ".join(candidate.strip().split())
|
||||
if label and label not in labels:
|
||||
labels.append(label)
|
||||
if not labels:
|
||||
raise ValueError("Enter at least one object label or referring phrase.")
|
||||
return labels
|
||||
|
||||
|
||||
def _safe_score(value: Any) -> float:
|
||||
score = float(value.item() if hasattr(value, "item") else value)
|
||||
return min(1.0, max(0.0, score))
|
||||
|
||||
|
||||
def _result_labels(result: dict[str, Any], labels: list[str]) -> list[str]:
|
||||
text_labels = result.get("text_labels")
|
||||
if text_labels is not None:
|
||||
return [str(label) for label in text_labels]
|
||||
|
||||
raw_labels = result.get("labels", result.get("classes", []))
|
||||
resolved = []
|
||||
for value in raw_labels:
|
||||
if isinstance(value, str):
|
||||
resolved.append(value)
|
||||
continue
|
||||
index = int(value.item() if hasattr(value, "item") else value)
|
||||
resolved.append(labels[index] if 0 <= index < len(labels) else str(index))
|
||||
return resolved
|
||||
|
||||
|
||||
def result_to_detections(
|
||||
result: dict[str, Any],
|
||||
*,
|
||||
labels: list[str],
|
||||
width: int,
|
||||
height: int,
|
||||
frame_index: int,
|
||||
timestamp: float,
|
||||
source: str,
|
||||
max_detections: int,
|
||||
) -> tuple[Detection, ...]:
|
||||
"""Normalize a Transformers detector result into immutable detections."""
|
||||
|
||||
boxes = result.get("boxes", ())
|
||||
scores = result.get("scores", ())
|
||||
resolved_labels = _result_labels(result, labels)
|
||||
count = min(len(boxes), len(scores), len(resolved_labels))
|
||||
records = []
|
||||
for index in range(count):
|
||||
box_value = boxes[index]
|
||||
if hasattr(box_value, "detach"):
|
||||
box_value = box_value.detach().to(device="cpu").tolist()
|
||||
x1, y1, x2, y2 = (float(value) for value in box_value)
|
||||
x1 = min(float(width), max(0.0, x1))
|
||||
y1 = min(float(height), max(0.0, y1))
|
||||
x2 = min(float(width), max(x1, x2))
|
||||
y2 = min(float(height), max(y1, y2))
|
||||
if x2 <= x1 or y2 <= y1:
|
||||
continue
|
||||
records.append(
|
||||
Detection(
|
||||
bbox_xyxy=(x1, y1, x2, y2),
|
||||
label=resolved_labels[index].strip() or None,
|
||||
score=_safe_score(scores[index]),
|
||||
frame_index=frame_index,
|
||||
timestamp=timestamp,
|
||||
source=source,
|
||||
metadata={"model_id": source},
|
||||
)
|
||||
)
|
||||
records.sort(
|
||||
key=lambda item: (
|
||||
-(item.score or 0.0),
|
||||
item.label or "",
|
||||
item.bbox_xyxy,
|
||||
)
|
||||
)
|
||||
return tuple(records[:max_detections])
|
||||
|
||||
|
||||
def _post_process(
|
||||
processor: Any,
|
||||
spec: DetectorSpec,
|
||||
outputs: Any,
|
||||
inputs: dict[str, Any],
|
||||
labels: list[str],
|
||||
sizes: list[tuple[int, int]],
|
||||
box_threshold: float,
|
||||
text_threshold: float,
|
||||
nms_threshold: float,
|
||||
max_detections: int,
|
||||
) -> list[dict[str, Any]]:
|
||||
if spec.family == "grounding_dino":
|
||||
kwargs = {
|
||||
"threshold": float(box_threshold),
|
||||
"text_threshold": float(text_threshold),
|
||||
"target_sizes": sizes,
|
||||
}
|
||||
input_ids = inputs.get("input_ids")
|
||||
if input_ids is not None:
|
||||
kwargs["input_ids"] = input_ids
|
||||
return processor.post_process_grounded_object_detection(outputs, **kwargs)
|
||||
if spec.family == "omdet":
|
||||
return processor.post_process_grounded_object_detection(
|
||||
outputs,
|
||||
text_labels=[labels] * len(sizes),
|
||||
threshold=float(box_threshold),
|
||||
nms_threshold=float(nms_threshold),
|
||||
target_sizes=sizes,
|
||||
max_num_det=int(max_detections),
|
||||
)
|
||||
return processor.post_process_grounded_object_detection(
|
||||
outputs,
|
||||
threshold=float(box_threshold),
|
||||
target_sizes=sizes,
|
||||
text_labels=[labels] * len(sizes),
|
||||
)
|
||||
|
||||
|
||||
class OpenVocabularyDetector:
|
||||
def __init__(self, spec: DetectorSpec, precision: str = "auto"):
|
||||
transformers = require_module("transformers")
|
||||
model_path = snapshot_download(
|
||||
spec.model_id,
|
||||
spec.cache_name,
|
||||
ignore_patterns=["*.bin", "*.gguf", "*.onnx", "*.tflite"],
|
||||
)
|
||||
processor = transformers.AutoProcessor.from_pretrained(model_path)
|
||||
model_class = transformers.AutoModelForZeroShotObjectDetection
|
||||
dtype = torch_dtype(precision)
|
||||
# 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
|
||||
self.processor = processor
|
||||
self.handle = ManagedTorchModel(model, processor=processor)
|
||||
|
||||
def close(self):
|
||||
self.handle.close()
|
||||
|
||||
def detect(
|
||||
self,
|
||||
images: torch.Tensor,
|
||||
labels: list[str],
|
||||
*,
|
||||
box_threshold: float,
|
||||
text_threshold: float,
|
||||
nms_threshold: float,
|
||||
max_detections: int,
|
||||
fps: float,
|
||||
batch_size: int,
|
||||
) -> DetectionSequence:
|
||||
if not math.isfinite(fps) or fps <= 0:
|
||||
raise ValueError("fps must be finite and positive.")
|
||||
if not isinstance(batch_size, int) or batch_size < 1:
|
||||
raise ValueError("batch_size must be a positive integer.")
|
||||
frames = []
|
||||
pil_images = tensor_batch_to_pil(images)
|
||||
model = self.handle.ensure_loaded()
|
||||
device = model_device(model)
|
||||
for start in range(0, len(pil_images), batch_size):
|
||||
image_batch = pil_images[start : start + batch_size]
|
||||
text = [labels] * len(image_batch)
|
||||
inputs = self.processor(
|
||||
images=image_batch,
|
||||
text=text,
|
||||
return_tensors="pt",
|
||||
)
|
||||
inputs = move_inputs(inputs, device, floating_dtype=self.dtype)
|
||||
with torch.inference_mode(), inference_context(device, self.dtype):
|
||||
outputs = model(**inputs)
|
||||
results = _post_process(
|
||||
self.processor,
|
||||
self.spec,
|
||||
outputs,
|
||||
inputs,
|
||||
labels,
|
||||
[(image.height, image.width) for image in image_batch],
|
||||
box_threshold,
|
||||
text_threshold,
|
||||
nms_threshold,
|
||||
max_detections,
|
||||
)
|
||||
if len(results) != len(image_batch):
|
||||
raise RuntimeError(
|
||||
f"{self.spec.model_id} returned {len(results)} result sets "
|
||||
f"for a batch of {len(image_batch)} images."
|
||||
)
|
||||
for offset, (image, result) in enumerate(
|
||||
zip(image_batch, results, strict=True)
|
||||
):
|
||||
frame_index = start + offset
|
||||
detections = result_to_detections(
|
||||
result,
|
||||
labels=labels,
|
||||
width=image.width,
|
||||
height=image.height,
|
||||
frame_index=frame_index,
|
||||
timestamp=frame_index / fps,
|
||||
source=self.spec.model_id,
|
||||
max_detections=max_detections,
|
||||
)
|
||||
frames.append(
|
||||
FrameDetections(
|
||||
frame_index=frame_index,
|
||||
timestamp=frame_index / fps,
|
||||
width=image.width,
|
||||
height=image.height,
|
||||
detections=detections,
|
||||
)
|
||||
)
|
||||
first = pil_images[0]
|
||||
return DetectionSequence(
|
||||
width=first.width,
|
||||
height=first.height,
|
||||
frames=tuple(frames),
|
||||
frame_count=len(frames),
|
||||
fps=fps,
|
||||
source=self.spec.model_id,
|
||||
metadata={"labels": labels, "model_family": self.spec.family},
|
||||
)
|
||||
|
||||
|
||||
def render_detections(
|
||||
images: torch.Tensor, detections: DetectionSequence
|
||||
) -> torch.Tensor:
|
||||
rendered = []
|
||||
for index, image in enumerate(tensor_batch_to_pil(images)):
|
||||
canvas = image.copy()
|
||||
draw = ImageDraw.Draw(canvas)
|
||||
frame = detections.frame(index)
|
||||
for detection in frame.detections if frame else ():
|
||||
color = deterministic_color(
|
||||
detection.track_id
|
||||
if detection.track_id is not None
|
||||
else detection.label or "object"
|
||||
)
|
||||
color = tuple(int(component) for component in color)
|
||||
x1, y1, x2, y2 = detection.bbox_xyxy
|
||||
draw.rectangle(
|
||||
(x1, y1, max(x1, x2 - 1), max(y1, y2 - 1)),
|
||||
outline=color,
|
||||
width=max(2, round(min(image.size) / 256)),
|
||||
)
|
||||
label = detection.label or "object"
|
||||
if detection.score is not None:
|
||||
label += f" {detection.score:.2f}"
|
||||
text_box = draw.textbbox((x1, y1), label)
|
||||
draw.rectangle(text_box, fill=color)
|
||||
draw.text((x1, y1), label, fill=(0, 0, 0))
|
||||
array = torch.from_numpy(np.asarray(canvas, dtype=np.float32).copy())
|
||||
rendered.append(array / 255.0)
|
||||
return torch.stack(rendered)
|
||||
|
||||
|
||||
def detection_box_masks(
|
||||
detections: DetectionSequence,
|
||||
) -> torch.Tensor:
|
||||
masks = torch.zeros(
|
||||
(detections.frame_count, detections.height, detections.width),
|
||||
dtype=torch.float32,
|
||||
)
|
||||
for frame in detections.frames:
|
||||
for detection in frame.detections:
|
||||
x1, y1, x2, y2 = detection.bbox_xyxy
|
||||
ix1, iy1 = int(x1), int(y1)
|
||||
ix2, iy2 = int(math.ceil(x2)), int(math.ceil(y2))
|
||||
masks[frame.frame_index, iy1:iy2, ix1:ix2] = 1.0
|
||||
return masks
|
||||
|
||||
|
||||
def _core_box(detection: Detection) -> dict[str, Any]:
|
||||
x1, y1, x2, y2 = detection.bbox_xyxy
|
||||
left, top = math.floor(x1), math.floor(y1)
|
||||
right, bottom = math.ceil(x2), math.ceil(y2)
|
||||
return {
|
||||
"x": left,
|
||||
"y": top,
|
||||
"width": right - left,
|
||||
"height": bottom - top,
|
||||
"label": detection.label,
|
||||
"score": detection.score,
|
||||
"metadata": {
|
||||
"frame_index": detection.frame_index,
|
||||
"label": detection.label,
|
||||
"score": detection.score,
|
||||
"source": detection.source,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def core_bounding_box_frames(
|
||||
detections: DetectionSequence,
|
||||
) -> list[list[dict[str, Any]]]:
|
||||
"""Return the nested per-frame convention used by core BOUNDING_BOX."""
|
||||
|
||||
frames = [[] for _index in range(detections.frame_count)]
|
||||
for frame in detections.frames:
|
||||
frames[frame.frame_index] = [
|
||||
_core_box(detection) for detection in frame.detections
|
||||
]
|
||||
return frames
|
||||
|
||||
|
||||
def core_bounding_boxes(detections: DetectionSequence) -> list[dict[str, Any]]:
|
||||
"""Return the flat metadata-rich BOUNDING_BOXES contract."""
|
||||
|
||||
result = []
|
||||
for frame in detections.frames:
|
||||
for detection in frame.detections:
|
||||
result.append(_core_box(detection))
|
||||
return result
|
||||
|
||||
|
||||
class VLMOpenVocabularyDetection(CachedModelNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"model": (tuple(MODEL_SPECS),),
|
||||
"labels": (
|
||||
"STRING",
|
||||
{
|
||||
"multiline": True,
|
||||
"default": "person, animal, vehicle",
|
||||
"tooltip": "Comma, semicolon, or newline-separated concepts.",
|
||||
},
|
||||
),
|
||||
"box_threshold": (
|
||||
"FLOAT",
|
||||
{"default": 0.3, "min": 0.0, "max": 1.0, "step": 0.01},
|
||||
),
|
||||
"text_threshold": (
|
||||
"FLOAT",
|
||||
{"default": 0.25, "min": 0.0, "max": 1.0, "step": 0.01},
|
||||
),
|
||||
"max_detections": (
|
||||
"INT",
|
||||
{"default": 100, "min": 1, "max": 1000},
|
||||
),
|
||||
"fps": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": 0.001,
|
||||
"max": 1000.0,
|
||||
"step": 0.001,
|
||||
"tooltip": (
|
||||
"Connect Get Video Components fps for video batches."
|
||||
),
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"nms_threshold": (
|
||||
"FLOAT",
|
||||
{"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01},
|
||||
),
|
||||
"precision": (("auto", "bfloat16", "float16", "float32"),),
|
||||
"batch_size": (
|
||||
"INT",
|
||||
{
|
||||
"default": 1,
|
||||
"min": 1,
|
||||
"max": 16,
|
||||
"tooltip": (
|
||||
"Frames per model call. Increase only when VRAM allows."
|
||||
),
|
||||
},
|
||||
),
|
||||
"unload_after": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = (
|
||||
VLM_DETECTIONS,
|
||||
"STRING",
|
||||
"IMAGE",
|
||||
"MASK",
|
||||
"BOUNDING_BOX",
|
||||
"BOUNDING_BOXES",
|
||||
)
|
||||
RETURN_NAMES = (
|
||||
"detections",
|
||||
"json",
|
||||
"preview",
|
||||
"box_mask",
|
||||
"bounding_boxes",
|
||||
"bounding_boxes_with_metadata",
|
||||
)
|
||||
FUNCTION = "detect"
|
||||
CATEGORY = "VLM Nodes/Vision/Detection"
|
||||
DESCRIPTION = (
|
||||
"Detect text-specified objects with one portable interface. Outputs "
|
||||
"versioned detections, JSON, preview, box masks, and core boxes."
|
||||
)
|
||||
|
||||
def detect(
|
||||
self,
|
||||
image,
|
||||
model,
|
||||
labels,
|
||||
box_threshold,
|
||||
text_threshold,
|
||||
max_detections,
|
||||
fps,
|
||||
nms_threshold=0.5,
|
||||
precision="auto",
|
||||
batch_size=1,
|
||||
unload_after=False,
|
||||
):
|
||||
concepts = parse_labels(labels)
|
||||
fps_value = float(fps)
|
||||
batch_size_value = int(batch_size)
|
||||
if not math.isfinite(fps_value) or fps_value <= 0:
|
||||
raise ValueError("fps must be finite and positive.")
|
||||
if batch_size_value < 1:
|
||||
raise ValueError("batch_size must be a positive integer.")
|
||||
spec = MODEL_SPECS[model]
|
||||
predictor = self.get_or_create_model(
|
||||
(spec.model_id, precision),
|
||||
lambda: OpenVocabularyDetector(spec, precision),
|
||||
)
|
||||
try:
|
||||
detections = predictor.detect(
|
||||
image,
|
||||
concepts,
|
||||
box_threshold=box_threshold,
|
||||
text_threshold=text_threshold,
|
||||
nms_threshold=nms_threshold,
|
||||
max_detections=max_detections,
|
||||
fps=fps_value,
|
||||
batch_size=batch_size_value,
|
||||
)
|
||||
return (
|
||||
detections,
|
||||
detections.to_json(indent=2),
|
||||
render_detections(image, detections),
|
||||
detection_box_masks(detections),
|
||||
core_bounding_box_frames(detections),
|
||||
core_bounding_boxes(detections),
|
||||
)
|
||||
finally:
|
||||
self.maybe_clear_model(unload_after)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"VLMOpenVocabularyDetection": VLMOpenVocabularyDetection,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"VLMOpenVocabularyDetection": "VLM Open-Vocabulary Detection",
|
||||
}
|
||||
+1826
File diff suppressed because it is too large
Load Diff
+1
-1
@@ -118,7 +118,7 @@ class Joytag(CachedModelNode):
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "tags"
|
||||
CATEGORY = "VLM Nodes/JoyTag"
|
||||
CATEGORY = "VLM Nodes/Vision/Tagging"
|
||||
|
||||
def tags(
|
||||
self,
|
||||
|
||||
+1
-1
@@ -100,7 +100,7 @@ class Kosmos2model(CachedModelNode):
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "new_model_generate_predictions"
|
||||
CATEGORY = "VLM Nodes/Kosmos-2"
|
||||
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
|
||||
|
||||
def new_model_generate_predictions(
|
||||
self,
|
||||
|
||||
+1
-1
@@ -126,7 +126,7 @@ class MCLLaVAModel(CachedModelNode):
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "generate_image_description"
|
||||
CATEGORY = "VLM Nodes/MC-LLaVA"
|
||||
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
|
||||
|
||||
def generate_image_description(
|
||||
self,
|
||||
|
||||
+1
-1
@@ -160,7 +160,7 @@ class MiniCPMNode(CachedModelNode):
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "generate"
|
||||
CATEGORY = "VLM Nodes/MiniCPM-V"
|
||||
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
|
||||
|
||||
def generate(
|
||||
self,
|
||||
|
||||
@@ -0,0 +1,478 @@
|
||||
"""MiniMax music generation and cover support with fixed regional routing."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import binascii
|
||||
import io
|
||||
import json
|
||||
import os
|
||||
from typing import Any
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from .audioldm2 import ANY
|
||||
from .hosted_api import redact_sensitive
|
||||
from .runtime import require_module
|
||||
|
||||
API_KEY_ENV = "MINIMAX_API_KEY"
|
||||
REGION_ENDPOINTS = {
|
||||
"global_en": "https://api.minimax.io/v1/music_generation",
|
||||
"cn_zh": "https://api.minimaxi.com/v1/music_generation",
|
||||
}
|
||||
GENERATION_MODELS = (
|
||||
"music-3.0",
|
||||
"music-2.6",
|
||||
"music-3.0-free",
|
||||
"music-2.6-free",
|
||||
)
|
||||
COVER_MODELS = ("music-cover", "music-cover-free")
|
||||
MUSIC_MODELS = GENERATION_MODELS + COVER_MODELS
|
||||
DEFAULT_MODEL = "music-3.0"
|
||||
REQUEST_FIELDS = frozenset(
|
||||
{
|
||||
"model",
|
||||
"prompt",
|
||||
"lyrics",
|
||||
"stream",
|
||||
"output_format",
|
||||
"audio_setting",
|
||||
"lyrics_optimizer",
|
||||
"is_instrumental",
|
||||
"audio_url",
|
||||
"audio_base64",
|
||||
"cover_feature_id",
|
||||
}
|
||||
)
|
||||
OUTPUT_FORMATS = ("url", "hex")
|
||||
STREAM_OUTPUT_FORMATS = ("hex",)
|
||||
AUDIO_FORMATS = ("mp3", "wav", "pcm")
|
||||
SAMPLE_RATES = (16000, 24000, 32000, 44100)
|
||||
BITRATES = (32000, 64000, 128000, 256000)
|
||||
REGIONAL_FIELDS = {"global_en": (), "cn_zh": ("aigc_watermark",)}
|
||||
STATUS_IN_PROGRESS = 1
|
||||
STATUS_COMPLETED = 2
|
||||
MAX_COVER_BYTES = 50 * 1024 * 1024
|
||||
MAX_AUDIO_BYTES = 128 * 1024 * 1024
|
||||
|
||||
|
||||
def _clean_text(value: object) -> str:
|
||||
return str(value or "").strip()
|
||||
|
||||
|
||||
def _validate_cover_base64(value: str) -> None:
|
||||
try:
|
||||
decoded = base64.b64decode(value, validate=True)
|
||||
except (binascii.Error, ValueError, TypeError):
|
||||
raise ValueError("audio_base64 must contain valid base64 data.") from None
|
||||
if len(decoded) > MAX_COVER_BYTES:
|
||||
raise ValueError("audio_base64 exceeds the 50 MiB cover input limit.")
|
||||
|
||||
|
||||
def build_music_request(
|
||||
*,
|
||||
region: str,
|
||||
model: str,
|
||||
prompt: str,
|
||||
lyrics: str,
|
||||
stream: bool,
|
||||
output_format: str,
|
||||
audio_format: str,
|
||||
sample_rate: int,
|
||||
bitrate: int,
|
||||
lyrics_optimizer: bool,
|
||||
is_instrumental: bool,
|
||||
aigc_watermark: bool,
|
||||
audio_url: str = "",
|
||||
audio_base64: str = "",
|
||||
cover_feature_id: str = "",
|
||||
) -> dict[str, Any]:
|
||||
"""Validate node inputs and build the documented JSON request body."""
|
||||
|
||||
if region not in REGION_ENDPOINTS:
|
||||
raise ValueError(f"region must be one of {tuple(REGION_ENDPOINTS)}.")
|
||||
if model not in MUSIC_MODELS:
|
||||
raise ValueError(f"model must be one of {MUSIC_MODELS}.")
|
||||
if output_format not in OUTPUT_FORMATS:
|
||||
raise ValueError(f"output_format must be one of {OUTPUT_FORMATS}.")
|
||||
if bool(stream) and output_format not in STREAM_OUTPUT_FORMATS:
|
||||
raise ValueError("Streaming music responses require output_format='hex'.")
|
||||
if audio_format not in AUDIO_FORMATS:
|
||||
raise ValueError(f"audio_format must be one of {AUDIO_FORMATS}.")
|
||||
if int(sample_rate) not in SAMPLE_RATES:
|
||||
raise ValueError(f"sample_rate must be one of {SAMPLE_RATES}.")
|
||||
if int(bitrate) not in BITRATES:
|
||||
raise ValueError(f"bitrate must be one of {BITRATES}.")
|
||||
|
||||
clean_prompt = _clean_text(prompt)
|
||||
clean_lyrics = _clean_text(lyrics)
|
||||
clean_audio_url = _clean_text(audio_url)
|
||||
clean_audio_base64 = _clean_text(audio_base64)
|
||||
clean_cover_feature_id = _clean_text(cover_feature_id)
|
||||
if len(clean_prompt) > 2000:
|
||||
raise ValueError("prompt exceeds the 2,000-character music API limit.")
|
||||
|
||||
payload: dict[str, Any] = {
|
||||
"model": model,
|
||||
"stream": bool(stream),
|
||||
"output_format": output_format,
|
||||
"audio_setting": {
|
||||
"sample_rate": int(sample_rate),
|
||||
"bitrate": int(bitrate),
|
||||
"format": audio_format,
|
||||
},
|
||||
}
|
||||
if clean_prompt:
|
||||
payload["prompt"] = clean_prompt
|
||||
if clean_lyrics:
|
||||
payload["lyrics"] = clean_lyrics
|
||||
|
||||
if model in COVER_MODELS:
|
||||
if not 10 <= len(clean_prompt) <= 300:
|
||||
raise ValueError("Cover generation requires a 10-300 character prompt.")
|
||||
sources = (clean_audio_url, clean_audio_base64, clean_cover_feature_id)
|
||||
if sum(bool(value) for value in sources) != 1:
|
||||
raise ValueError(
|
||||
"Cover generation requires exactly one of audio_url, "
|
||||
"audio_base64, or cover_feature_id."
|
||||
)
|
||||
if clean_audio_base64:
|
||||
_validate_cover_base64(clean_audio_base64)
|
||||
payload["audio_base64"] = clean_audio_base64
|
||||
elif clean_audio_url:
|
||||
payload["audio_url"] = clean_audio_url
|
||||
else:
|
||||
if not 10 <= len(clean_lyrics) <= 1000:
|
||||
raise ValueError(
|
||||
"cover_feature_id requires lyrics between 10 and 1,000 characters."
|
||||
)
|
||||
payload["cover_feature_id"] = clean_cover_feature_id
|
||||
if clean_lyrics and not 10 <= len(clean_lyrics) <= 1000:
|
||||
raise ValueError("Cover lyrics must be between 10 and 1,000 characters.")
|
||||
else:
|
||||
if any((clean_audio_url, clean_audio_base64, clean_cover_feature_id)):
|
||||
raise ValueError("Cover audio fields require a cover model.")
|
||||
if len(clean_lyrics) > 3500:
|
||||
raise ValueError("lyrics exceeds the 3,500-character music API limit.")
|
||||
if bool(is_instrumental) and not clean_prompt:
|
||||
raise ValueError("Instrumental generation requires a prompt.")
|
||||
if not bool(is_instrumental) and not clean_lyrics and not bool(lyrics_optimizer):
|
||||
raise ValueError(
|
||||
"Non-instrumental generation requires lyrics or lyrics_optimizer."
|
||||
)
|
||||
payload["lyrics_optimizer"] = bool(lyrics_optimizer)
|
||||
payload["is_instrumental"] = bool(is_instrumental)
|
||||
|
||||
if region == "cn_zh":
|
||||
payload["aigc_watermark"] = bool(aigc_watermark)
|
||||
return payload
|
||||
|
||||
|
||||
def _response_parts(payload: object) -> tuple[str, int, dict[str, Any]]:
|
||||
if not isinstance(payload, dict):
|
||||
raise RuntimeError("MiniMax returned a non-object music response.")
|
||||
base_response = payload.get("base_resp")
|
||||
if not isinstance(base_response, dict):
|
||||
raise RuntimeError("MiniMax returned no base_resp status.")
|
||||
try:
|
||||
success_code = int(base_response.get("status_code"))
|
||||
except (TypeError, ValueError):
|
||||
raise RuntimeError("MiniMax returned an invalid base_resp status code.") from None
|
||||
if success_code != 0:
|
||||
message = _clean_text(base_response.get("status_msg")) or "unknown API error"
|
||||
raise RuntimeError(f"MiniMax music API error {success_code}: {message}")
|
||||
|
||||
data = payload.get("data")
|
||||
if not isinstance(data, dict):
|
||||
raise RuntimeError("MiniMax returned no music data object.")
|
||||
try:
|
||||
status = int(data.get("status"))
|
||||
except (TypeError, ValueError):
|
||||
raise RuntimeError("MiniMax returned an invalid music status.") from None
|
||||
if status not in {STATUS_IN_PROGRESS, STATUS_COMPLETED}:
|
||||
raise RuntimeError(f"MiniMax returned unsupported music status {status}.")
|
||||
audio = data.get("audio", "")
|
||||
if not isinstance(audio, str):
|
||||
raise RuntimeError("MiniMax returned a non-string audio value.")
|
||||
extra_info = payload.get("extra_info")
|
||||
return audio.strip(), status, extra_info if isinstance(extra_info, dict) else {}
|
||||
|
||||
|
||||
def _stream_audio(response: Any) -> tuple[str, dict[str, Any]]:
|
||||
audio = ""
|
||||
extra_info: dict[str, Any] = {}
|
||||
completed = False
|
||||
saw_payload = False
|
||||
for line in response.iter_lines():
|
||||
raw = line.decode("utf-8") if isinstance(line, bytes) else str(line)
|
||||
raw = raw.strip()
|
||||
if raw.startswith("data:"):
|
||||
raw = raw[5:].strip()
|
||||
if not raw or raw == "[DONE]" or raw.startswith("event:"):
|
||||
continue
|
||||
try:
|
||||
payload = json.loads(raw)
|
||||
except json.JSONDecodeError:
|
||||
raise RuntimeError("MiniMax returned invalid streaming JSON.") from None
|
||||
chunk, status, metadata = _response_parts(payload)
|
||||
saw_payload = True
|
||||
if chunk:
|
||||
if chunk.startswith(audio):
|
||||
audio = chunk
|
||||
elif not audio.startswith(chunk):
|
||||
audio += chunk
|
||||
if metadata:
|
||||
extra_info = metadata
|
||||
completed = completed or status == STATUS_COMPLETED
|
||||
if not saw_payload:
|
||||
raise RuntimeError("MiniMax returned an empty streaming response.")
|
||||
if not completed:
|
||||
raise RuntimeError("MiniMax streaming ended before music generation completed.")
|
||||
if not audio:
|
||||
raise RuntimeError("MiniMax returned no audio data.")
|
||||
return audio, extra_info
|
||||
|
||||
|
||||
def _request_audio_value(
|
||||
client: Any,
|
||||
endpoint: str,
|
||||
headers: dict[str, str],
|
||||
payload: dict[str, Any],
|
||||
) -> tuple[str, dict[str, Any]]:
|
||||
if payload["stream"]:
|
||||
with client.stream("POST", endpoint, headers=headers, json=payload) as response:
|
||||
response.raise_for_status()
|
||||
return _stream_audio(response)
|
||||
|
||||
response = client.post(endpoint, headers=headers, json=payload)
|
||||
response.raise_for_status()
|
||||
audio, status, extra_info = _response_parts(response.json())
|
||||
if status != STATUS_COMPLETED:
|
||||
raise RuntimeError(
|
||||
"MiniMax music generation is still in progress and has no query endpoint."
|
||||
)
|
||||
if not audio:
|
||||
raise RuntimeError("MiniMax returned no audio data.")
|
||||
return audio, extra_info
|
||||
|
||||
|
||||
def _download_audio(client: Any, url: str) -> bytes:
|
||||
parsed = urlsplit(url)
|
||||
if (
|
||||
parsed.scheme != "https"
|
||||
or not parsed.hostname
|
||||
or parsed.username is not None
|
||||
or parsed.password is not None
|
||||
):
|
||||
raise RuntimeError("MiniMax returned an invalid HTTPS audio URL.")
|
||||
chunks: list[bytes] = []
|
||||
total = 0
|
||||
with client.stream("GET", url) as response:
|
||||
response.raise_for_status()
|
||||
for chunk in response.iter_bytes():
|
||||
total += len(chunk)
|
||||
if total > MAX_AUDIO_BYTES:
|
||||
raise RuntimeError("MiniMax audio download exceeds 128 MiB.")
|
||||
chunks.append(chunk)
|
||||
return b"".join(chunks)
|
||||
|
||||
|
||||
def _audio_bytes(client: Any, value: str, output_format: str) -> bytes:
|
||||
if output_format == "url":
|
||||
return _download_audio(client, value)
|
||||
try:
|
||||
return bytes.fromhex("".join(value.split()))
|
||||
except ValueError:
|
||||
raise RuntimeError("MiniMax returned invalid hexadecimal audio data.") from None
|
||||
|
||||
|
||||
def _metadata_integer(metadata: dict[str, Any], name: str, default: int) -> int:
|
||||
try:
|
||||
value = int(metadata.get(name, default))
|
||||
except (TypeError, ValueError):
|
||||
return int(default)
|
||||
return value if value > 0 else int(default)
|
||||
|
||||
|
||||
def _decode_audio(
|
||||
content: bytes,
|
||||
audio_format: str,
|
||||
requested_sample_rate: int,
|
||||
metadata: dict[str, Any],
|
||||
) -> tuple[np.ndarray, int]:
|
||||
if not content:
|
||||
raise RuntimeError("MiniMax returned an empty audio payload.")
|
||||
if audio_format == "pcm":
|
||||
if len(content) % 2:
|
||||
raise RuntimeError("MiniMax returned an odd-length PCM payload.")
|
||||
channels = _metadata_integer(metadata, "music_channel", 1)
|
||||
raw = np.frombuffer(content, dtype="<i2")
|
||||
if raw.size % channels:
|
||||
raise RuntimeError("MiniMax PCM samples do not align with the channel count.")
|
||||
samples = raw.astype(np.float32).reshape(-1, channels) / 32768.0
|
||||
sample_rate = _metadata_integer(
|
||||
metadata,
|
||||
"music_sample_rate",
|
||||
requested_sample_rate,
|
||||
)
|
||||
else:
|
||||
soundfile = require_module("soundfile", "soundfile>=0.12")
|
||||
try:
|
||||
samples, sample_rate = soundfile.read(
|
||||
io.BytesIO(content),
|
||||
dtype="float32",
|
||||
always_2d=True,
|
||||
)
|
||||
except Exception as exc:
|
||||
detail = redact_sensitive(exc)
|
||||
raise RuntimeError(f"Could not decode MiniMax {audio_format} audio: {detail}") from None
|
||||
samples = np.asarray(samples, dtype=np.float32)
|
||||
sample_rate = int(sample_rate)
|
||||
if samples.ndim != 2 or not samples.size:
|
||||
raise RuntimeError("MiniMax decoded to an empty audio array.")
|
||||
if not np.isfinite(samples).all():
|
||||
raise RuntimeError("MiniMax decoded audio contains non-finite samples.")
|
||||
return np.ascontiguousarray(samples), int(sample_rate)
|
||||
|
||||
|
||||
class MiniMaxMusicNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"region": (tuple(REGION_ENDPOINTS), {"default": "global_en"}),
|
||||
"model": (MUSIC_MODELS, {"default": DEFAULT_MODEL}),
|
||||
"prompt": ("STRING", {"default": "", "multiline": True}),
|
||||
"lyrics": ("STRING", {"default": "", "multiline": True}),
|
||||
"stream": ("BOOLEAN", {"default": False}),
|
||||
"output_format": (OUTPUT_FORMATS, {"default": "hex"}),
|
||||
"audio_format": (AUDIO_FORMATS, {"default": "mp3"}),
|
||||
"sample_rate": (SAMPLE_RATES, {"default": 44100}),
|
||||
"bitrate": (BITRATES, {"default": 256000}),
|
||||
"lyrics_optimizer": ("BOOLEAN", {"default": False}),
|
||||
"is_instrumental": ("BOOLEAN", {"default": False}),
|
||||
"aigc_watermark": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": "Sent only to the cn_zh endpoint.",
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"audio_url": ("STRING", {"default": ""}),
|
||||
"audio_base64": ("STRING", {"default": "", "multiline": True}),
|
||||
"cover_feature_id": ("STRING", {"default": ""}),
|
||||
"timeout_seconds": (
|
||||
"FLOAT",
|
||||
{"default": 600.0, "min": 1.0, "max": 1800.0},
|
||||
),
|
||||
"use_system_proxy": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_NAMES = ("wave_form", "sample_rate", "audio")
|
||||
RETURN_TYPES = (ANY, "INT", "AUDIO")
|
||||
OUTPUT_NODE = True
|
||||
FUNCTION = "generate_music"
|
||||
CATEGORY = "VLM Nodes/Audio"
|
||||
DESCRIPTION = (
|
||||
"Generate music or covers through fixed MiniMax regional endpoints. "
|
||||
f"The API key is read only from {API_KEY_ENV}."
|
||||
)
|
||||
|
||||
def generate_music(
|
||||
self,
|
||||
region,
|
||||
model,
|
||||
prompt,
|
||||
lyrics,
|
||||
stream,
|
||||
output_format,
|
||||
audio_format,
|
||||
sample_rate,
|
||||
bitrate,
|
||||
lyrics_optimizer,
|
||||
is_instrumental,
|
||||
aigc_watermark,
|
||||
audio_url="",
|
||||
audio_base64="",
|
||||
cover_feature_id="",
|
||||
timeout_seconds=600.0,
|
||||
use_system_proxy=False,
|
||||
):
|
||||
payload = build_music_request(
|
||||
region=region,
|
||||
model=model,
|
||||
prompt=prompt,
|
||||
lyrics=lyrics,
|
||||
stream=stream,
|
||||
output_format=output_format,
|
||||
audio_format=audio_format,
|
||||
sample_rate=sample_rate,
|
||||
bitrate=bitrate,
|
||||
lyrics_optimizer=lyrics_optimizer,
|
||||
is_instrumental=is_instrumental,
|
||||
aigc_watermark=aigc_watermark,
|
||||
audio_url=audio_url,
|
||||
audio_base64=audio_base64,
|
||||
cover_feature_id=cover_feature_id,
|
||||
)
|
||||
api_key = os.getenv(API_KEY_ENV, "").strip()
|
||||
if not api_key:
|
||||
raise ValueError(
|
||||
f"Set {API_KEY_ENV} in the environment that starts ComfyUI, "
|
||||
"then restart the server."
|
||||
)
|
||||
|
||||
httpx = require_module("httpx", "httpx>=0.27,<1")
|
||||
client = httpx.Client(
|
||||
timeout=max(1.0, min(1800.0, float(timeout_seconds))),
|
||||
follow_redirects=False,
|
||||
trust_env=bool(use_system_proxy),
|
||||
)
|
||||
headers = {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
try:
|
||||
value, metadata = _request_audio_value(
|
||||
client,
|
||||
REGION_ENDPOINTS[region],
|
||||
headers,
|
||||
payload,
|
||||
)
|
||||
content = _audio_bytes(client, value, output_format)
|
||||
samples, actual_rate = _decode_audio(
|
||||
content,
|
||||
audio_format,
|
||||
int(sample_rate),
|
||||
metadata,
|
||||
)
|
||||
legacy = samples[:, 0] if samples.shape[1] == 1 else samples
|
||||
audio = {
|
||||
"waveform": torch.from_numpy(samples.T.copy()).unsqueeze(0),
|
||||
"sample_rate": actual_rate,
|
||||
}
|
||||
return (legacy.tolist(), actual_rate, audio)
|
||||
except Exception as exc:
|
||||
detail = redact_sensitive(exc, (api_key,))
|
||||
raise RuntimeError(f"MiniMax music request failed: {detail}") from None
|
||||
finally:
|
||||
try:
|
||||
client.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {"MiniMaxMusicNode": MiniMaxMusicNode}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {"MiniMaxMusicNode": "MiniMax Music"}
|
||||
|
||||
|
||||
__all__ = [
|
||||
"MiniMaxMusicNode",
|
||||
"NODE_CLASS_MAPPINGS",
|
||||
"NODE_DISPLAY_NAME_MAPPINGS",
|
||||
]
|
||||
+126
-29
@@ -8,8 +8,9 @@ small and large VLM families while keeping downloads and VRAM allocation lazy.
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Callable
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
@@ -25,13 +26,14 @@ from .runtime import (
|
||||
model_device,
|
||||
move_inputs,
|
||||
normalize_hf_model_id,
|
||||
require_quantization_backend,
|
||||
require_module,
|
||||
require_quantization_backend,
|
||||
reserve_external_vram,
|
||||
snapshot_download,
|
||||
tensor_batch_to_pil,
|
||||
torch_dtype,
|
||||
)
|
||||
from .vision_types import VLM_VIDEO_SELECTION, VideoFrameSelection
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@@ -190,6 +192,24 @@ MODEL_CATALOG = {
|
||||
),
|
||||
}
|
||||
|
||||
RECOMMENDED_MODEL_LABELS = (
|
||||
"Qwen 3.5 0.8B (fastest current)",
|
||||
"Qwen 3.5 4B (recommended)",
|
||||
"Qwen 3 VL 2B Instruct",
|
||||
"Qwen 3 VL 4B Instruct",
|
||||
"Qwen 3 VL 8B Instruct",
|
||||
"SmolVLM2 500M Video (low VRAM)",
|
||||
"SmolVLM2 2.2B Video",
|
||||
"LFM2.5 VL 450M (edge)",
|
||||
"InternVL 3.5 1B HF",
|
||||
"Granite Vision 4.1 4B (structured documents)",
|
||||
"Gemma 3 4B IT (license acceptance required)",
|
||||
"Custom Hugging Face model",
|
||||
)
|
||||
LEGACY_MODEL_LABELS = tuple(
|
||||
label for label in MODEL_CATALOG if label not in RECOMMENDED_MODEL_LABELS
|
||||
)
|
||||
|
||||
MEMORY_MODES = (
|
||||
"ComfyUI managed (BF16)",
|
||||
"4-bit NF4 (bitsandbytes)",
|
||||
@@ -445,6 +465,7 @@ class ModernVLMPredictor:
|
||||
fps: float = 1.0,
|
||||
enable_thinking: bool = False,
|
||||
stream_callback: Callable[[str], None] | None = None,
|
||||
video_selection: VideoFrameSelection | None = None,
|
||||
) -> str:
|
||||
primary_images = (
|
||||
tensor_batch_to_pil(images) if images is not None else []
|
||||
@@ -461,6 +482,27 @@ class ModernVLMPredictor:
|
||||
f"{self.spec.family} does not advertise video support. "
|
||||
"Disconnect video_frames or select Qwen/SmolVLM2."
|
||||
)
|
||||
if video_selection is not None:
|
||||
if video is None:
|
||||
raise ValueError(
|
||||
"video_selection requires a connected video_frames batch."
|
||||
)
|
||||
if not isinstance(video_selection, VideoFrameSelection):
|
||||
raise TypeError("video_selection must be a VLM Video Selection.")
|
||||
if len(video_selection.frames) != len(video):
|
||||
raise ValueError(
|
||||
"video_selection frame count must match video_frames."
|
||||
)
|
||||
source_aspect = video_selection.width / video_selection.height
|
||||
analysis_aspect = video[0].width / video[0].height
|
||||
if abs(source_aspect - analysis_aspect) > max(
|
||||
0.01,
|
||||
source_aspect * 0.01,
|
||||
):
|
||||
raise ValueError(
|
||||
"video_selection and video_frames must have the same "
|
||||
"aspect ratio."
|
||||
)
|
||||
|
||||
results = []
|
||||
# A connected video is the primary visual input. Including ComfyUI's
|
||||
@@ -483,25 +525,49 @@ class ModernVLMPredictor:
|
||||
if video is not None
|
||||
else [{"type": "image", "image": image}]
|
||||
)
|
||||
effective_prompt = (
|
||||
f"The video frames are sampled at {float(fps):g} FPS.\n\n{prompt}"
|
||||
if video is not None
|
||||
else prompt
|
||||
)
|
||||
if video is not None and video_selection is not None:
|
||||
timeline = ", ".join(
|
||||
f"{position}=frame {frame.source_frame_index} "
|
||||
f"at {frame.timestamp:.6f}s"
|
||||
for position, frame in enumerate(video_selection.frames)
|
||||
)
|
||||
effective_prompt = (
|
||||
"The supplied video images are irregular samples from one "
|
||||
f"{video_selection.source_frame_count}-frame video at "
|
||||
f"{video_selection.fps:g} FPS. Supplied-image mapping: "
|
||||
f"{timeline}.\n\n{prompt}"
|
||||
)
|
||||
elif video is not None:
|
||||
effective_prompt = (
|
||||
f"The video frames are sampled at {float(fps):g} FPS.\n\n"
|
||||
f"{prompt}"
|
||||
)
|
||||
else:
|
||||
effective_prompt = prompt
|
||||
content.append({"type": "text", "text": effective_prompt})
|
||||
messages.append({"role": "user", "content": content})
|
||||
|
||||
metadata = None
|
||||
if video is not None:
|
||||
frame_rate = float(fps)
|
||||
metadata = {
|
||||
"total_num_frames": len(video),
|
||||
"fps": frame_rate,
|
||||
"duration": len(video) / frame_rate,
|
||||
"frames_indices": list(range(len(video))),
|
||||
"width": video[0].width,
|
||||
"height": video[0].height,
|
||||
}
|
||||
if video_selection is not None:
|
||||
metadata = {
|
||||
"total_num_frames": video_selection.source_frame_count,
|
||||
"fps": video_selection.fps,
|
||||
"duration": video_selection.duration,
|
||||
"frames_indices": list(video_selection.indices),
|
||||
"width": video[0].width,
|
||||
"height": video[0].height,
|
||||
}
|
||||
else:
|
||||
frame_rate = float(fps)
|
||||
metadata = {
|
||||
"total_num_frames": len(video),
|
||||
"fps": frame_rate,
|
||||
"duration": len(video) / frame_rate,
|
||||
"frames_indices": list(range(len(video))),
|
||||
"width": video[0].width,
|
||||
"height": video[0].height,
|
||||
}
|
||||
inputs = self._inputs(
|
||||
messages,
|
||||
enable_thinking,
|
||||
@@ -607,7 +673,7 @@ class ModernVLM(CachedModelNode):
|
||||
},
|
||||
),
|
||||
"model": (
|
||||
list(MODEL_CATALOG),
|
||||
list(RECOMMENDED_MODEL_LABELS),
|
||||
{"default": "Qwen 3 VL 2B Instruct"},
|
||||
),
|
||||
"custom_model_id": ("STRING", {"default": ""}),
|
||||
@@ -638,6 +704,7 @@ class ModernVLM(CachedModelNode):
|
||||
},
|
||||
),
|
||||
"video_frames": ("IMAGE",),
|
||||
"video_selection": (VLM_VIDEO_SELECTION,),
|
||||
"fps": (
|
||||
"FLOAT",
|
||||
{"default": 1.0, "min": 0.1, "max": 60.0, "step": 0.1},
|
||||
@@ -666,6 +733,15 @@ class ModernVLM(CachedModelNode):
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "VLM Nodes/Modern"
|
||||
|
||||
@classmethod
|
||||
def VALIDATE_INPUTS(cls, model):
|
||||
# The visible combo is deliberately curated. Accepting every known
|
||||
# catalog value here keeps workflows saved before the curation fully
|
||||
# executable even when their model now lives under Legacy.
|
||||
if model not in MODEL_CATALOG:
|
||||
return f"Unsupported Modern VLM model {model!r}."
|
||||
return True
|
||||
|
||||
def run(
|
||||
self,
|
||||
prompt,
|
||||
@@ -678,6 +754,7 @@ class ModernVLM(CachedModelNode):
|
||||
image=None,
|
||||
system_prompt="You are an expert visual analyst.",
|
||||
video_frames=None,
|
||||
video_selection=None,
|
||||
fps=1.0,
|
||||
attention_mode="Auto (SDPA)",
|
||||
enable_thinking=False,
|
||||
@@ -705,25 +782,45 @@ class ModernVLM(CachedModelNode):
|
||||
try:
|
||||
return (
|
||||
predictor.generate(
|
||||
image,
|
||||
prompt,
|
||||
system_prompt,
|
||||
max_new_tokens,
|
||||
temperature,
|
||||
top_p,
|
||||
video_frames,
|
||||
fps,
|
||||
enable_thinking,
|
||||
stream_callback,
|
||||
images=image,
|
||||
prompt=prompt,
|
||||
system_prompt=system_prompt,
|
||||
max_new_tokens=max_new_tokens,
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
video_frames=video_frames,
|
||||
fps=fps,
|
||||
video_selection=video_selection,
|
||||
enable_thinking=enable_thinking,
|
||||
stream_callback=stream_callback,
|
||||
),
|
||||
)
|
||||
finally:
|
||||
self.maybe_clear_model(unload_after)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {"ModernVLM": ModernVLM}
|
||||
class LegacyModernVLM(ModernVLM):
|
||||
"""Compatibility surface for redundant, superseded, and very large tiers."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
inputs = super().INPUT_TYPES()
|
||||
inputs["required"]["model"] = (
|
||||
list(LEGACY_MODEL_LABELS),
|
||||
{"default": LEGACY_MODEL_LABELS[0]},
|
||||
)
|
||||
return inputs
|
||||
|
||||
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"ModernVLM": ModernVLM,
|
||||
"LegacyModernVLM": LegacyModernVLM,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"ModernVLM": (
|
||||
"Modern VLM (Qwen / SmolVLM2 / LFM / InternVL / Granite / Gemma)"
|
||||
)
|
||||
),
|
||||
"LegacyModernVLM": "[Legacy] Modern VLM Compatibility",
|
||||
}
|
||||
|
||||
+2
-2
@@ -14,8 +14,8 @@ from .runtime import (
|
||||
external_device_map,
|
||||
inference_context,
|
||||
model_device,
|
||||
require_quantization_backend,
|
||||
require_module,
|
||||
require_quantization_backend,
|
||||
reserve_external_vram,
|
||||
snapshot_download,
|
||||
tensor_batch_to_pil,
|
||||
@@ -155,7 +155,7 @@ class MolmoNode(CachedModelNode):
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "generate"
|
||||
CATEGORY = "VLM Nodes/Molmo"
|
||||
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
|
||||
|
||||
def generate(
|
||||
self,
|
||||
|
||||
+69
-32
@@ -1,7 +1,21 @@
|
||||
"""Current Moondream 2 node using the model's supported query API."""
|
||||
"""Current Moondream 2 node using the model's supported query API.
|
||||
|
||||
The pinned checkpoint was authored against Transformers 4.52.4. Loading it
|
||||
through Transformers 5's ``from_pretrained`` compatibility path can silently
|
||||
produce an all-EOS model even when every tensor is reported as loaded. The
|
||||
checkpoint itself is a normal safetensors state dict, so instantiate its
|
||||
official wrapper and load that state dict directly. This keeps Moondream in
|
||||
ComfyUI's managed VRAM lifecycle without downgrading Transformers for the rest
|
||||
of the node pack.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from types import ModuleType
|
||||
|
||||
import torch
|
||||
|
||||
from .runtime import (
|
||||
@@ -18,12 +32,59 @@ from .runtime import (
|
||||
|
||||
MODEL_ID = "vikhyatk/moondream2"
|
||||
MODEL_REVISION = "2025-06-21"
|
||||
_CHECKPOINT_PACKAGE = "_comfyui_vlm_moondream2_checkpoint"
|
||||
|
||||
|
||||
def _checkpoint_module(model_path: str | Path):
|
||||
"""Import the checkpoint's relative modules without HF's generated cache.
|
||||
|
||||
Hugging Face's dynamic-module cache can omit transitive relative imports
|
||||
for a local snapshot. Giving the snapshot a private package namespace lets
|
||||
Python resolve the checkpoint's own ``.config``, ``.vision``, and related
|
||||
modules directly and deterministically.
|
||||
"""
|
||||
|
||||
source = str(Path(model_path).resolve())
|
||||
package = sys.modules.get(_CHECKPOINT_PACKAGE)
|
||||
if package is None:
|
||||
package = ModuleType(_CHECKPOINT_PACKAGE)
|
||||
package.__path__ = [source]
|
||||
package.__package__ = _CHECKPOINT_PACKAGE
|
||||
sys.modules[_CHECKPOINT_PACKAGE] = package
|
||||
elif list(getattr(package, "__path__", ())) != [source]:
|
||||
raise RuntimeError(
|
||||
"Moondream2 checkpoint source changed inside a running process. "
|
||||
"Restart ComfyUI before loading a different snapshot."
|
||||
)
|
||||
return importlib.import_module(f"{_CHECKPOINT_PACKAGE}.hf_moondream")
|
||||
|
||||
|
||||
def _load_native_checkpoint(model_path: str | Path):
|
||||
checkpoint = _checkpoint_module(model_path)
|
||||
safetensors = require_module("safetensors.torch")
|
||||
config = checkpoint.HfConfig.from_pretrained(
|
||||
model_path,
|
||||
local_files_only=True,
|
||||
)
|
||||
model = checkpoint.HfMoondream(config)
|
||||
weights = Path(model_path) / "model.safetensors"
|
||||
if not weights.is_file():
|
||||
raise FileNotFoundError(f"Moondream2 weights are missing: {weights}")
|
||||
missing, unexpected = safetensors.load_model(
|
||||
model,
|
||||
str(weights),
|
||||
strict=True,
|
||||
)
|
||||
if missing or unexpected:
|
||||
raise RuntimeError(
|
||||
"Moondream2 checkpoint did not load exactly: "
|
||||
f"missing={sorted(missing)}, unexpected={sorted(unexpected)}"
|
||||
)
|
||||
return model.eval()
|
||||
|
||||
|
||||
class Moondream2Predictor:
|
||||
def __init__(self):
|
||||
transformers = require_module("transformers")
|
||||
dynamic_modules = require_module("transformers.dynamic_module_utils")
|
||||
model_path = snapshot_download(
|
||||
MODEL_ID,
|
||||
"moondream2",
|
||||
@@ -31,31 +92,7 @@ class Moondream2Predictor:
|
||||
ignore_patterns=["*.bin", "*.gguf"],
|
||||
)
|
||||
self.dtype = torch_dtype("bfloat16")
|
||||
config = transformers.AutoConfig.from_pretrained(
|
||||
model_path,
|
||||
revision=MODEL_REVISION,
|
||||
trust_remote_code=True,
|
||||
)
|
||||
remote_class = dynamic_modules.get_class_from_dynamic_module(
|
||||
"hf_moondream.HfMoondream",
|
||||
model_path,
|
||||
local_files_only=True,
|
||||
)
|
||||
|
||||
class Transformers5Moondream(remote_class):
|
||||
def __init__(self, model_config):
|
||||
super().__init__(model_config)
|
||||
# The pinned remote wrapper predates the Transformers 5 model
|
||||
# loader and does not declare its tied-weight metadata. Calling
|
||||
# the full post_init would reinitialize custom Moondream state.
|
||||
self.all_tied_weights_keys = {}
|
||||
|
||||
model = Transformers5Moondream.from_pretrained(
|
||||
model_path,
|
||||
config=config,
|
||||
dtype=self.dtype,
|
||||
)
|
||||
model.eval()
|
||||
model = _load_native_checkpoint(model_path)
|
||||
self.handle = ManagedTorchModel(model)
|
||||
|
||||
def close(self):
|
||||
@@ -92,9 +129,9 @@ class Moondream2Predictor:
|
||||
response = response.get("answer", response)
|
||||
if not str(response).strip():
|
||||
raise RuntimeError(
|
||||
"Moondream2 returned an empty response on this "
|
||||
"Torch/Transformers build. Use the Modern VLM node with "
|
||||
"LFM2.5-VL 450M, InternVL 3.5 1B, or Qwen3-VL 2B."
|
||||
"Moondream2 returned an empty response. Verify that the "
|
||||
f"{MODEL_REVISION} snapshot is complete, then restart "
|
||||
"ComfyUI so its checkpoint modules are reloaded."
|
||||
)
|
||||
results.append(str(response))
|
||||
return batch_text(results)
|
||||
@@ -134,7 +171,7 @@ class Moondream2model(CachedModelNode):
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "moondream2_generate_predictions"
|
||||
CATEGORY = "VLM Nodes/Moondream2"
|
||||
CATEGORY = "VLM Nodes/Modern/Edge"
|
||||
|
||||
def moondream2_generate_predictions(
|
||||
self,
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,388 @@
|
||||
"""Isolated Moondream 3.1 Photon worker.
|
||||
|
||||
This file is launched directly by the ComfyUI process with the dedicated
|
||||
Moondream virtual environment. It intentionally has no imports from ComfyUI
|
||||
or this package: Moondream pins a Pillow version that is incompatible with
|
||||
current ComfyUI releases, so sharing one Python environment is unsafe.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import platform
|
||||
import sys
|
||||
import time
|
||||
import traceback
|
||||
from collections.abc import Callable
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from dataclasses import replace
|
||||
from importlib.metadata import PackageNotFoundError, version
|
||||
from io import BytesIO
|
||||
from multiprocessing.connection import Client
|
||||
from typing import Any
|
||||
|
||||
from PIL import Image
|
||||
|
||||
|
||||
def _honor_do_not_track() -> bool:
|
||||
"""Disable anonymous Photon reporting when the sidecar requests privacy.
|
||||
|
||||
Kestrel 0.4.2 does not currently inspect the conventional DO_NOT_TRACK
|
||||
environment variable. Base-model inference does not need its reporter, so
|
||||
keep validation local, skip the telemetry loop, and still close the HTTP
|
||||
client during engine shutdown. Finetune inference retains upstream auth
|
||||
and reporting behavior because it explicitly receives an API key.
|
||||
"""
|
||||
|
||||
if os.environ.get("DO_NOT_TRACK") != "1":
|
||||
return False
|
||||
if os.environ.get("MOONDREAM_API_KEY", "").strip():
|
||||
return False
|
||||
|
||||
from kestrel.photon import PhotonReporter
|
||||
|
||||
async def validate_api_key(self) -> bool:
|
||||
return False
|
||||
|
||||
def start(self) -> None:
|
||||
return None
|
||||
|
||||
async def shutdown(self) -> None:
|
||||
await self._client.aclose()
|
||||
|
||||
PhotonReporter.validate_api_key = validate_api_key
|
||||
PhotonReporter.start = start
|
||||
PhotonReporter.shutdown = shutdown
|
||||
return True
|
||||
|
||||
|
||||
def _register_moondream31_if_needed(model_name: str) -> bool:
|
||||
"""Bridge the official model-card ID on runtimes released before the ID.
|
||||
|
||||
Moondream 3.1 uses the same MD3 Photon runtime/checkpoint format as the
|
||||
preview. Stable moondream 1.3.0 / kestrel 0.4.2 shipped the safetensors
|
||||
loader but omitted the new registry entry published by the later model
|
||||
card. Prefer an upstream entry whenever present; otherwise clone only the
|
||||
runtime metadata and point it at the official 3.1 weights.
|
||||
"""
|
||||
|
||||
if model_name != "moondream3.1-9B-A2B":
|
||||
return False
|
||||
from kestrel.models import get_spec, register
|
||||
|
||||
try:
|
||||
get_spec(model_name)
|
||||
return False
|
||||
except ValueError:
|
||||
preview = get_spec("moondream3-preview")
|
||||
register(
|
||||
replace(
|
||||
preview,
|
||||
name=model_name,
|
||||
repo_id="moondream/moondream3.1-9B-A2B",
|
||||
filename="model.safetensors",
|
||||
checkpoint_format="md3",
|
||||
)
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
def _base_model_name(value: str) -> str:
|
||||
return str(value).split("/", 1)[0]
|
||||
|
||||
|
||||
def _model_skills(model_name: str) -> frozenset[str]:
|
||||
base_model = _base_model_name(model_name)
|
||||
if base_model == "moondream3.1-9B-A2B":
|
||||
# Source of truth: the final 3.1 model card. Segment remains a skill
|
||||
# of the 3 Preview and cloud API, not the final local 3.1 checkpoint.
|
||||
return frozenset(("caption", "query", "detect", "point"))
|
||||
from kestrel.models import get_spec
|
||||
|
||||
spec = get_spec(base_model)
|
||||
templates = spec.default_config.get("tokenizer", {}).get("templates", {})
|
||||
return frozenset(
|
||||
name for name, template in templates.items() if template is not None
|
||||
)
|
||||
|
||||
|
||||
def _image(value: bytes) -> Image.Image:
|
||||
if not isinstance(value, bytes):
|
||||
raise TypeError("Worker image payloads must be bytes.")
|
||||
with Image.open(BytesIO(value)) as source:
|
||||
return source.convert("RGB")
|
||||
|
||||
|
||||
def _parallel(
|
||||
images: list[bytes],
|
||||
operation: Callable[[Image.Image], dict[str, Any]],
|
||||
workers: int,
|
||||
) -> list[dict[str, Any]]:
|
||||
if not images:
|
||||
return []
|
||||
worker_count = max(1, min(int(workers), len(images)))
|
||||
with ThreadPoolExecutor(max_workers=worker_count) as pool:
|
||||
return list(pool.map(lambda value: operation(_image(value)), images))
|
||||
|
||||
|
||||
def _private_shutdown(model: Any) -> None:
|
||||
"""Best-effort graceful Photon shutdown before the process exits.
|
||||
|
||||
The public moondream package currently has no close method. Process
|
||||
isolation remains the hard guarantee: the parent terminates this exact
|
||||
process if this best-effort private cleanup ever changes or stalls.
|
||||
"""
|
||||
|
||||
engine = getattr(model, "_engine", None)
|
||||
loop = getattr(model, "_loop", None)
|
||||
thread = getattr(model, "_thread", None)
|
||||
if engine is not None and loop is not None:
|
||||
try:
|
||||
import asyncio
|
||||
|
||||
asyncio.run_coroutine_threadsafe(engine.shutdown(), loop).result(timeout=20)
|
||||
except Exception: # noqa: BLE001,S110 - private best-effort cleanup.
|
||||
pass
|
||||
try:
|
||||
loop.call_soon_threadsafe(loop.stop)
|
||||
except Exception: # noqa: BLE001,S110 - private best-effort cleanup.
|
||||
pass
|
||||
if thread is not None:
|
||||
try:
|
||||
thread.join(timeout=5)
|
||||
except Exception: # noqa: BLE001,S110 - private best-effort cleanup.
|
||||
pass
|
||||
|
||||
|
||||
def _request(
|
||||
model: Any,
|
||||
request: dict[str, Any],
|
||||
send: Callable[[dict[str, Any]], None],
|
||||
max_batch_size: int,
|
||||
supported_skills: frozenset[str],
|
||||
) -> bool:
|
||||
request_id = request.get("id")
|
||||
operation = request.get("operation")
|
||||
if operation == "shutdown":
|
||||
send({"id": request_id, "type": "result", "result": {"closed": True}})
|
||||
return False
|
||||
if operation not in supported_skills:
|
||||
raise ValueError(
|
||||
f"Model does not support the {operation!r} skill. "
|
||||
f"Available skills: {', '.join(sorted(supported_skills))}."
|
||||
)
|
||||
|
||||
started = time.perf_counter()
|
||||
settings = {"max_tokens": int(request.get("max_tokens", 512))}
|
||||
if operation in {"query", "caption"}:
|
||||
image_payload = request.get("image")
|
||||
image = _image(image_payload) if image_payload is not None else None
|
||||
if operation == "query":
|
||||
output = model.query(
|
||||
image=image,
|
||||
question=str(request["question"]),
|
||||
stream=bool(request.get("stream", True)),
|
||||
settings=settings,
|
||||
reasoning=bool(request.get("reasoning", False)),
|
||||
)
|
||||
key = "answer"
|
||||
else:
|
||||
if image is None:
|
||||
raise ValueError("Caption requires an image.")
|
||||
output = model.caption(
|
||||
image=image,
|
||||
length=str(request.get("length", "normal")),
|
||||
stream=bool(request.get("stream", True)),
|
||||
settings=settings,
|
||||
)
|
||||
key = "caption"
|
||||
|
||||
value = output[key]
|
||||
if isinstance(value, str):
|
||||
text = value
|
||||
else:
|
||||
chunks = []
|
||||
for chunk in value:
|
||||
chunk_text = str(chunk)
|
||||
chunks.append(chunk_text)
|
||||
send(
|
||||
{
|
||||
"id": request_id,
|
||||
"type": "chunk",
|
||||
"text": chunk_text,
|
||||
}
|
||||
)
|
||||
text = "".join(chunks)
|
||||
result = {
|
||||
key: text,
|
||||
"elapsed_seconds": time.perf_counter() - started,
|
||||
}
|
||||
if operation == "query" and output.get("reasoning") is not None:
|
||||
result["reasoning"] = output["reasoning"]
|
||||
send({"id": request_id, "type": "result", "result": result})
|
||||
return True
|
||||
|
||||
images = request.get("images")
|
||||
if not isinstance(images, list):
|
||||
raise TypeError(f"{operation} requires an image list.")
|
||||
workers = min(
|
||||
max_batch_size,
|
||||
max(1, int(request.get("parallel_requests", max_batch_size))),
|
||||
)
|
||||
object_prompt = str(request.get("object", "")).strip()
|
||||
if not object_prompt:
|
||||
raise ValueError(f"{operation} requires a non-empty object prompt.")
|
||||
|
||||
if operation == "detect":
|
||||
results = _parallel(
|
||||
images,
|
||||
lambda image: model.detect(image, object_prompt, settings=settings),
|
||||
workers,
|
||||
)
|
||||
elif operation == "point":
|
||||
results = _parallel(
|
||||
images,
|
||||
lambda image: model.point(image, object_prompt, settings=settings),
|
||||
workers,
|
||||
)
|
||||
elif operation == "segment":
|
||||
spatial_refs = request.get("spatial_refs") or None
|
||||
results = _parallel(
|
||||
images,
|
||||
lambda image: model.segment(
|
||||
image,
|
||||
object_prompt,
|
||||
spatial_refs=spatial_refs,
|
||||
stream=False,
|
||||
settings=settings,
|
||||
),
|
||||
workers,
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Unknown worker operation {operation!r}.")
|
||||
|
||||
send(
|
||||
{
|
||||
"id": request_id,
|
||||
"type": "result",
|
||||
"result": {
|
||||
"items": results,
|
||||
"processed_frames": len(images),
|
||||
"parallel_requests": workers,
|
||||
"elapsed_seconds": time.perf_counter() - started,
|
||||
},
|
||||
}
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
def main() -> int:
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--host", default="127.0.0.1")
|
||||
parser.add_argument("--port", type=int, required=True)
|
||||
parser.add_argument("--auth-key")
|
||||
parser.add_argument("--model", required=True)
|
||||
parser.add_argument("--device", required=True)
|
||||
parser.add_argument("--max-batch-size", type=int, required=True)
|
||||
parser.add_argument("--kv-cache-pages", type=int, default=0)
|
||||
args = parser.parse_args()
|
||||
|
||||
auth_key = args.auth_key or os.environ.pop("MOONDREAM_WORKER_AUTH", "")
|
||||
if not auth_key:
|
||||
parser.error("worker authentication is missing")
|
||||
connection = Client(
|
||||
(args.host, args.port),
|
||||
authkey=bytes.fromhex(auth_key),
|
||||
)
|
||||
|
||||
def send(value: dict[str, Any]) -> None:
|
||||
connection.send(value)
|
||||
|
||||
send(
|
||||
{
|
||||
"type": "status",
|
||||
"status": "loading",
|
||||
"python": sys.version.split()[0],
|
||||
"platform": platform.platform(),
|
||||
"pid": os.getpid(),
|
||||
}
|
||||
)
|
||||
|
||||
model = None
|
||||
try:
|
||||
import moondream as md
|
||||
|
||||
base_model = _base_model_name(args.model)
|
||||
compatibility_registration = _register_moondream31_if_needed(base_model)
|
||||
telemetry_disabled = _honor_do_not_track()
|
||||
supported_skills = _model_skills(args.model)
|
||||
kwargs: dict[str, Any] = {
|
||||
"local": True,
|
||||
"model": args.model,
|
||||
"device": args.device,
|
||||
"max_batch_size": args.max_batch_size,
|
||||
}
|
||||
if args.kv_cache_pages > 0:
|
||||
kwargs["kv_cache_pages"] = args.kv_cache_pages
|
||||
model = md.vl(**kwargs)
|
||||
try:
|
||||
package_version = version("moondream")
|
||||
except PackageNotFoundError:
|
||||
package_version = "unknown"
|
||||
send(
|
||||
{
|
||||
"type": "status",
|
||||
"status": "ready",
|
||||
"moondream_version": package_version,
|
||||
"compatibility_registration": compatibility_registration,
|
||||
"telemetry_disabled": telemetry_disabled,
|
||||
"skills": sorted(supported_skills),
|
||||
"pid": os.getpid(),
|
||||
}
|
||||
)
|
||||
|
||||
running = True
|
||||
while running:
|
||||
request = connection.recv()
|
||||
request_id = request.get("id") if isinstance(request, dict) else None
|
||||
try:
|
||||
if not isinstance(request, dict):
|
||||
raise TypeError("Worker requests must be dictionaries.")
|
||||
running = _request(
|
||||
model,
|
||||
request,
|
||||
send,
|
||||
args.max_batch_size,
|
||||
supported_skills,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 - report request failures over IPC.
|
||||
send(
|
||||
{
|
||||
"id": request_id,
|
||||
"type": "error",
|
||||
"error": f"{type(exc).__name__}: {exc}",
|
||||
"traceback": traceback.format_exc(limit=12),
|
||||
}
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 - report startup failures over IPC.
|
||||
send(
|
||||
{
|
||||
"type": "fatal",
|
||||
"error": f"{type(exc).__name__}: {exc}",
|
||||
"traceback": traceback.format_exc(limit=20),
|
||||
}
|
||||
)
|
||||
return 1
|
||||
finally:
|
||||
if model is not None:
|
||||
_private_shutdown(model)
|
||||
try:
|
||||
connection.close()
|
||||
except OSError:
|
||||
pass
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -25,7 +25,7 @@ class MoonDream(CachedModelNode):
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "answer_questions"
|
||||
CATEGORY = "VLM Nodes/MoonDream"
|
||||
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
|
||||
|
||||
def answer_questions(self, image, question, unload_after=False):
|
||||
predictor = self.get_or_create_model(
|
||||
|
||||
+2
-3
@@ -24,15 +24,14 @@ from .runtime import (
|
||||
normalize_hf_model_id,
|
||||
pil_mask_to_tensor,
|
||||
pil_to_tensor,
|
||||
require_quantization_backend,
|
||||
require_module,
|
||||
require_quantization_backend,
|
||||
reserve_external_vram,
|
||||
snapshot_download,
|
||||
tensor_batch_to_pil,
|
||||
torch_dtype,
|
||||
)
|
||||
|
||||
|
||||
PALIGEMMA_MODELS = [
|
||||
"gokaygokay/sd3-long-captioner-v2",
|
||||
"google/paligemma-3b-ft-refcoco-seg-896",
|
||||
@@ -263,7 +262,7 @@ class Paligemma(CachedModelNode):
|
||||
RETURN_TYPES = ("STRING", "MASK", "IMAGE")
|
||||
RETURN_NAMES = ("description", "mask", "visualization")
|
||||
FUNCTION = "process_task"
|
||||
CATEGORY = "VLM Nodes/Paligemma"
|
||||
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
|
||||
|
||||
def process_task(
|
||||
self,
|
||||
|
||||
+1
-1
@@ -51,4 +51,4 @@ Optional: If asked to create a random prompt create one.
|
||||
# Define the system message
|
||||
system_msg_simple = """
|
||||
You are an helpful asistant. Answer optional questions or help the user for their optional queries.
|
||||
"""
|
||||
"""
|
||||
|
||||
+2
-3
@@ -17,15 +17,14 @@ from .runtime import (
|
||||
inference_context,
|
||||
model_device,
|
||||
move_inputs,
|
||||
require_quantization_backend,
|
||||
require_module,
|
||||
require_quantization_backend,
|
||||
reserve_external_vram,
|
||||
snapshot_download,
|
||||
tensor_batch_to_pil,
|
||||
torch_dtype,
|
||||
)
|
||||
|
||||
|
||||
QWEN2_VL_MODELS = {
|
||||
"Qwen2-VL-2B": "Qwen/Qwen2-VL-2B-Instruct",
|
||||
"Qwen2-VL-7B": "Qwen/Qwen2-VL-7B-Instruct",
|
||||
@@ -347,7 +346,7 @@ class Qwen2VLNode(CachedModelNode):
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "generate"
|
||||
CATEGORY = "VLM Nodes/Qwen2-VL"
|
||||
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
|
||||
|
||||
def generate(
|
||||
self,
|
||||
|
||||
+2457
File diff suppressed because it is too large
Load Diff
+93
-36
@@ -18,11 +18,12 @@ import os
|
||||
import platform
|
||||
import re
|
||||
import threading
|
||||
from collections.abc import Callable, Iterable, Mapping
|
||||
from contextlib import nullcontext
|
||||
from dataclasses import dataclass
|
||||
from importlib import metadata
|
||||
from pathlib import Path
|
||||
from typing import Any, Callable, Iterable, Mapping
|
||||
from typing import Any
|
||||
|
||||
import folder_paths
|
||||
import numpy as np
|
||||
@@ -31,6 +32,7 @@ from PIL import Image
|
||||
|
||||
LOGGER = logging.getLogger("ComfyUI_VLM_nodes")
|
||||
GGUF_EXTENSIONS = {".gguf"}
|
||||
PIL_CONVERSION_CHUNK_BYTES = 128 * 1024**2
|
||||
LLAMA_FLASH_ATTENTION_CHOICES = ("Auto", "Enabled", "Disabled")
|
||||
LLAMA_SPLIT_MODE_CHOICES = ("Layer", "Row", "Single GPU")
|
||||
LLAMA_VISION_HANDLER_CHOICES = (
|
||||
@@ -157,39 +159,67 @@ def hf_download(repo_id: str, filename: str, subdirectory: str, **kwargs: Any) -
|
||||
return Path(hub.hf_hub_download(**download_kwargs))
|
||||
|
||||
|
||||
def _tensor_image_batch_to_uint8(images: torch.Tensor) -> np.ndarray:
|
||||
"""Convert HWC/CHW/BHWC/BCHW image data in one vectorized transfer.
|
||||
|
||||
Video nodes previously moved and normalized every frame independently.
|
||||
Converting a bounded batch at once reduces Python dispatch and host-device
|
||||
transfer overhead while preserving the same clipping contract. The caller
|
||||
chunks long videos to cap peak temporary memory. The returned array is
|
||||
always contiguous BHWC RGB uint8.
|
||||
"""
|
||||
|
||||
if not isinstance(images, torch.Tensor):
|
||||
raise TypeError(f"Expected a torch.Tensor, got {type(images).__name__}.")
|
||||
value = images.detach()
|
||||
if value.ndim == 2:
|
||||
value = value.unsqueeze(-1)
|
||||
if value.ndim == 3:
|
||||
value = value.unsqueeze(0)
|
||||
if value.ndim != 4:
|
||||
raise ValueError(
|
||||
f"Expected an HWC/BHWC or CHW/BCHW image tensor, got {tuple(value.shape)}."
|
||||
)
|
||||
|
||||
# ComfyUI uses BHWC. BCHW is accepted for compatibility with older nodes.
|
||||
if value.shape[-1] not in (1, 3, 4) and value.shape[1] in (1, 3, 4):
|
||||
value = value.permute(0, 2, 3, 1)
|
||||
if value.shape[-1] not in (1, 3, 4):
|
||||
raise ValueError(f"Unsupported image channel shape: {tuple(value.shape)}.")
|
||||
|
||||
value = torch.nan_to_num(
|
||||
value.to(device="cpu", dtype=torch.float32),
|
||||
nan=0.0,
|
||||
posinf=1.0,
|
||||
neginf=0.0,
|
||||
)
|
||||
if value.numel():
|
||||
flat = value.reshape(value.shape[0], -1)
|
||||
needs_byte_scale = (
|
||||
(flat.amax(dim=1) > 1.0) | (flat.amin(dim=1) < 0.0)
|
||||
).view(-1, 1, 1, 1)
|
||||
value = torch.where(needs_byte_scale, value / 255.0, value)
|
||||
value = value.clamp(0.0, 1.0).mul(255.0).round().to(torch.uint8)
|
||||
if value.shape[-1] == 1:
|
||||
value = value.expand(*value.shape[:-1], 3)
|
||||
elif value.shape[-1] == 4:
|
||||
value = value[..., :3]
|
||||
return np.ascontiguousarray(value.numpy())
|
||||
|
||||
|
||||
def tensor_to_pil(image: torch.Tensor, index: int = 0) -> Image.Image:
|
||||
"""Convert a Comfy IMAGE tensor to an RGB PIL image without torchvision."""
|
||||
|
||||
if not isinstance(image, torch.Tensor):
|
||||
raise TypeError(f"Expected a torch.Tensor, got {type(image).__name__}.")
|
||||
value = image.detach()
|
||||
value = image
|
||||
if value.ndim == 4:
|
||||
if not 0 <= index < value.shape[0]:
|
||||
raise IndexError(f"Image batch index {index} is out of range.")
|
||||
value = value[index]
|
||||
if value.ndim == 2:
|
||||
value = value.unsqueeze(-1)
|
||||
if value.ndim != 3:
|
||||
raise ValueError(
|
||||
f"Expected an HWC/BHWC or CHW/BCHW image tensor, got {tuple(value.shape)}."
|
||||
)
|
||||
|
||||
# ComfyUI uses HWC. CHW is accepted for compatibility with older callers.
|
||||
if value.shape[-1] not in (1, 3, 4) and value.shape[0] in (1, 3, 4):
|
||||
value = value.permute(1, 2, 0)
|
||||
if value.shape[-1] not in (1, 3, 4):
|
||||
raise ValueError(f"Unsupported image channel shape: {tuple(value.shape)}.")
|
||||
|
||||
value = torch.nan_to_num(
|
||||
value.to(device="cpu", dtype=torch.float32), nan=0.0, posinf=1.0, neginf=0.0
|
||||
)
|
||||
if value.numel() and (value.max() > 1.0 or value.min() < 0.0):
|
||||
value = value / 255.0
|
||||
array = value.clamp(0.0, 1.0).mul(255.0).round().to(torch.uint8).numpy()
|
||||
if array.shape[-1] == 1:
|
||||
array = np.repeat(array, 3, axis=-1)
|
||||
elif array.shape[-1] == 4:
|
||||
array = array[..., :3]
|
||||
elif index != 0:
|
||||
raise IndexError("A single image only has batch index 0.")
|
||||
array = _tensor_image_batch_to_uint8(value)[0]
|
||||
return Image.fromarray(array, mode="RGB")
|
||||
|
||||
|
||||
@@ -198,7 +228,24 @@ def tensor_batch_to_pil(images: torch.Tensor) -> list[Image.Image]:
|
||||
return [tensor_to_pil(images)]
|
||||
if images.ndim != 4:
|
||||
raise ValueError(f"Expected an IMAGE batch, got {tuple(images.shape)}.")
|
||||
return [tensor_to_pil(images, index) for index in range(images.shape[0])]
|
||||
if images.shape[0] == 0:
|
||||
return []
|
||||
if images.device.type == "cpu":
|
||||
# Per-frame conversion benchmarks faster for ordinary CPU-resident
|
||||
# Comfy IMAGE batches and keeps the transient working set tiny.
|
||||
return [tensor_to_pil(images, index) for index in range(images.shape[0])]
|
||||
# Convert several frames per tensor operation without creating an
|
||||
# unbounded full-video float32 temporary. On accelerator-resident batches,
|
||||
# this amortizes device synchronization and transfers across many frames.
|
||||
# The 128 MiB working-set ceiling keeps long HD/4K videos reliable.
|
||||
frame_elements = max(1, int(images[0].numel()))
|
||||
working_bytes = frame_elements * max(4, images.element_size())
|
||||
chunk_frames = max(1, PIL_CONVERSION_CHUNK_BYTES // working_bytes)
|
||||
output: list[Image.Image] = []
|
||||
for start in range(0, int(images.shape[0]), chunk_frames):
|
||||
frames = _tensor_image_batch_to_uint8(images[start : start + chunk_frames])
|
||||
output.extend(Image.fromarray(frame, mode="RGB") for frame in frames)
|
||||
return output
|
||||
|
||||
|
||||
def pil_to_tensor(image: Image.Image) -> torch.Tensor:
|
||||
@@ -540,18 +587,21 @@ class CachedModelNode:
|
||||
def __init__(self) -> None:
|
||||
self._model_handle = None
|
||||
self._model_key = None
|
||||
self._model_lock = threading.RLock()
|
||||
|
||||
def get_or_create_model(self, key: Any, factory: Callable[[], Any]):
|
||||
if self._model_handle is None or self._model_key != key:
|
||||
close_handle(self._model_handle)
|
||||
self._model_handle = factory()
|
||||
self._model_key = key
|
||||
return self._model_handle
|
||||
with self._model_lock:
|
||||
if self._model_handle is None or self._model_key != key:
|
||||
close_handle(self._model_handle)
|
||||
self._model_handle = factory()
|
||||
self._model_key = key
|
||||
return self._model_handle
|
||||
|
||||
def clear_model(self) -> None:
|
||||
close_handle(self._model_handle)
|
||||
self._model_handle = None
|
||||
self._model_key = None
|
||||
with self._model_lock:
|
||||
close_handle(self._model_handle)
|
||||
self._model_handle = None
|
||||
self._model_key = None
|
||||
|
||||
def maybe_clear_model(self, unload_after: bool) -> None:
|
||||
if unload_after:
|
||||
@@ -631,7 +681,10 @@ def llama_runtime_input_types() -> dict[str, tuple[Any, ...]]:
|
||||
"min": 1,
|
||||
"max": 8192,
|
||||
"step": 1,
|
||||
"tooltip": "Logical prompt batch. Lower this if context loading runs out of memory.",
|
||||
"tooltip": (
|
||||
"Logical prompt batch. Lower this if context loading runs "
|
||||
"out of memory."
|
||||
),
|
||||
},
|
||||
),
|
||||
"n_ubatch": (
|
||||
@@ -648,7 +701,11 @@ def llama_runtime_input_types() -> dict[str, tuple[Any, ...]]:
|
||||
list(LLAMA_FLASH_ATTENTION_CHOICES),
|
||||
{
|
||||
"default": "Auto",
|
||||
"tooltip": "Auto enables llama.cpp flash attention only with accelerator offload and safely retries without it when unsupported.",
|
||||
"tooltip": (
|
||||
"Auto enables llama.cpp flash attention only with "
|
||||
"accelerator offload and safely retries without it when "
|
||||
"unsupported."
|
||||
),
|
||||
},
|
||||
),
|
||||
"use_mmap": (
|
||||
|
||||
+539
@@ -0,0 +1,539 @@
|
||||
"""Prompt-seeded SAM2.1 image-batch/video segmentation and tracking."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from .geometry import bbox_from_mask, deterministic_color
|
||||
from .runtime import (
|
||||
CachedModelNode,
|
||||
ManagedTorchModel,
|
||||
inference_context,
|
||||
model_device,
|
||||
require_module,
|
||||
snapshot_download,
|
||||
tensor_batch_to_pil,
|
||||
torch_dtype,
|
||||
)
|
||||
from .vision_types import (
|
||||
VLM_DETECTIONS,
|
||||
VLM_TRACKS,
|
||||
Detection,
|
||||
DetectionSequence,
|
||||
Track,
|
||||
TrackSequence,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Sam2Spec:
|
||||
model_id: str
|
||||
cache_name: str
|
||||
|
||||
|
||||
SAM2_MODELS = {
|
||||
"SAM2.1 Hiera Tiny (fast)": Sam2Spec(
|
||||
"facebook/sam2.1-hiera-tiny", "sam2.1-hiera-tiny"
|
||||
),
|
||||
"SAM2.1 Hiera Small": Sam2Spec("facebook/sam2.1-hiera-small", "sam2.1-hiera-small"),
|
||||
"SAM2.1 Hiera Base+": Sam2Spec(
|
||||
"facebook/sam2.1-hiera-base-plus", "sam2.1-hiera-base-plus"
|
||||
),
|
||||
"SAM2.1 Hiera Large": Sam2Spec("facebook/sam2.1-hiera-large", "sam2.1-hiera-large"),
|
||||
}
|
||||
|
||||
|
||||
def _core_box(value: dict[str, Any]) -> tuple[float, float, float, float]:
|
||||
try:
|
||||
x, y = float(value["x"]), float(value["y"])
|
||||
width, height = float(value["width"]), float(value["height"])
|
||||
except (KeyError, TypeError, ValueError) as exc:
|
||||
raise ValueError(
|
||||
"BOUNDING_BOX must contain numeric x, y, width, and height."
|
||||
) from exc
|
||||
if width <= 0 or height <= 0:
|
||||
raise ValueError("BOUNDING_BOX width and height must be positive.")
|
||||
return x, y, x + width, y + height
|
||||
|
||||
|
||||
def _core_box_frames(value: Any) -> list[list[dict[str, Any]]]:
|
||||
"""Normalize core dict/flat/nested BOUNDING_BOX values to frame lists."""
|
||||
|
||||
if value is None:
|
||||
return []
|
||||
if isinstance(value, dict):
|
||||
return [[value]]
|
||||
if not isinstance(value, list):
|
||||
raise TypeError("BOUNDING_BOX must be a dict, list of dicts, or frame list.")
|
||||
if not value:
|
||||
return []
|
||||
if all(isinstance(item, dict) for item in value):
|
||||
return [value]
|
||||
if all(
|
||||
isinstance(frame, list) and all(isinstance(item, dict) for item in frame)
|
||||
for frame in value
|
||||
):
|
||||
return value
|
||||
raise TypeError("BOUNDING_BOX contains an unsupported nested value.")
|
||||
|
||||
|
||||
def _box_label(value: dict[str, Any]) -> str | None:
|
||||
label = value.get("label")
|
||||
if isinstance(label, str) and label.strip():
|
||||
return label.strip()
|
||||
metadata = value.get("metadata")
|
||||
if isinstance(metadata, dict):
|
||||
label = metadata.get("label")
|
||||
if isinstance(label, str) and label.strip():
|
||||
return label.strip()
|
||||
return None
|
||||
|
||||
|
||||
def seed_boxes(
|
||||
*,
|
||||
width: int,
|
||||
height: int,
|
||||
frame_index: int,
|
||||
detections: DetectionSequence | None,
|
||||
bounding_box: Any,
|
||||
) -> tuple[list[list[float]], list[int], dict[int, str | None]]:
|
||||
boxes: list[list[float]] = []
|
||||
object_ids: list[int] = []
|
||||
labels: dict[int, str | None] = {}
|
||||
if detections is not None:
|
||||
if not isinstance(detections, DetectionSequence):
|
||||
raise TypeError("detections must be a VLM Detection Sequence.")
|
||||
if detections.width != width or detections.height != height:
|
||||
raise ValueError(
|
||||
"Detection dimensions must exactly match the SAM2 video frames."
|
||||
)
|
||||
frame = detections.frame(frame_index)
|
||||
if frame is None and len(detections.frames) == 1:
|
||||
# A detector run over a selected single image is an explicit seed
|
||||
# annotation and may be applied to any chosen video frame.
|
||||
frame = detections.frames[0]
|
||||
elif frame is None and detections.frames:
|
||||
raise ValueError(
|
||||
f"Detections do not contain the requested seed frame {frame_index}."
|
||||
)
|
||||
for index, detection in enumerate(frame.detections if frame else (), 1):
|
||||
track_id = detection.track_id
|
||||
object_id = int(track_id if track_id is not None else index)
|
||||
while object_id in object_ids:
|
||||
object_id += 1
|
||||
x1, y1, x2, y2 = detection.bbox_xyxy
|
||||
boxes.append(
|
||||
[
|
||||
min(width, max(0.0, x1)),
|
||||
min(height, max(0.0, y1)),
|
||||
min(width, max(0.0, x2)),
|
||||
min(height, max(0.0, y2)),
|
||||
]
|
||||
)
|
||||
object_ids.append(object_id)
|
||||
labels[object_id] = detection.label
|
||||
box_frames = _core_box_frames(bounding_box)
|
||||
if len(box_frames) == 1:
|
||||
selected_boxes = box_frames[0]
|
||||
elif box_frames and frame_index < len(box_frames):
|
||||
selected_boxes = box_frames[frame_index]
|
||||
elif box_frames:
|
||||
raise ValueError(
|
||||
f"BOUNDING_BOX has {len(box_frames)} frames but seed_frame is "
|
||||
f"{frame_index}."
|
||||
)
|
||||
else:
|
||||
selected_boxes = []
|
||||
for value in selected_boxes:
|
||||
x1, y1, x2, y2 = _core_box(value)
|
||||
object_id = max(object_ids, default=0) + 1
|
||||
boxes.append(
|
||||
[
|
||||
min(width, max(0.0, x1)),
|
||||
min(height, max(0.0, y1)),
|
||||
min(width, max(0.0, x2)),
|
||||
min(height, max(0.0, y2)),
|
||||
]
|
||||
)
|
||||
object_ids.append(object_id)
|
||||
labels[object_id] = _box_label(value)
|
||||
valid = []
|
||||
for box, object_id in zip(boxes, object_ids):
|
||||
if box[2] > box[0] and box[3] > box[1]:
|
||||
valid.append((box, object_id))
|
||||
return (
|
||||
[box for box, _object_id in valid],
|
||||
[object_id for _box, object_id in valid],
|
||||
labels,
|
||||
)
|
||||
|
||||
|
||||
def _normalize_processed_masks(value: torch.Tensor) -> torch.Tensor:
|
||||
masks = value.detach().to(device="cpu")
|
||||
if masks.ndim == 4 and masks.shape[1] == 1:
|
||||
masks = masks[:, 0]
|
||||
elif masks.ndim == 4 and masks.shape[0] == 1:
|
||||
masks = masks[0]
|
||||
if masks.ndim == 2:
|
||||
masks = masks.unsqueeze(0)
|
||||
if masks.ndim != 3:
|
||||
raise RuntimeError(
|
||||
f"SAM2 returned an unsupported mask shape {tuple(masks.shape)}."
|
||||
)
|
||||
return masks if masks.dtype == torch.bool else masks > 0.5
|
||||
|
||||
|
||||
class Sam2VideoPredictor:
|
||||
def __init__(self, spec: Sam2Spec, precision: str):
|
||||
transformers = require_module("transformers")
|
||||
if not hasattr(transformers, "Sam2VideoModel"):
|
||||
raise RuntimeError(
|
||||
"SAM2 video requires Transformers with Sam2VideoModel support."
|
||||
)
|
||||
model_path = snapshot_download(
|
||||
spec.model_id,
|
||||
spec.cache_name,
|
||||
ignore_patterns=["*.pt", "*.bin", "*.onnx", "*.tflite"],
|
||||
)
|
||||
self.processor = transformers.Sam2VideoProcessor.from_pretrained(model_path)
|
||||
self.dtype = torch_dtype(precision)
|
||||
model = transformers.Sam2VideoModel.from_pretrained(
|
||||
model_path, dtype=self.dtype
|
||||
)
|
||||
model.eval()
|
||||
self.handle = ManagedTorchModel(model, processor=self.processor)
|
||||
self.spec = spec
|
||||
|
||||
def close(self):
|
||||
self.handle.close()
|
||||
|
||||
def propagate(
|
||||
self,
|
||||
images: torch.Tensor,
|
||||
*,
|
||||
seed_frame: int,
|
||||
fps: float,
|
||||
detections: DetectionSequence | None,
|
||||
bounding_box: Any,
|
||||
seed_mask: torch.Tensor | None,
|
||||
mask_threshold: float,
|
||||
keep_video_on_cpu: bool,
|
||||
mask_output: str,
|
||||
render_preview: bool,
|
||||
) -> tuple[TrackSequence, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
if not math.isfinite(fps) or fps <= 0:
|
||||
raise ValueError("fps must be finite and positive.")
|
||||
if mask_output not in {"union_only", "union_and_objects"}:
|
||||
raise ValueError(f"Unsupported mask_output mode {mask_output!r}.")
|
||||
pil_images = tensor_batch_to_pil(images)
|
||||
if not pil_images:
|
||||
raise ValueError("SAM2 requires at least one image.")
|
||||
if not 0 <= seed_frame < len(pil_images):
|
||||
raise ValueError(
|
||||
f"seed_frame {seed_frame} is outside the {len(pil_images)}-frame batch."
|
||||
)
|
||||
width, height = pil_images[0].size
|
||||
if any(image.size != (width, height) for image in pil_images):
|
||||
raise ValueError("Every video frame must have identical dimensions.")
|
||||
|
||||
boxes, object_ids, labels = seed_boxes(
|
||||
width=width,
|
||||
height=height,
|
||||
frame_index=seed_frame,
|
||||
detections=detections,
|
||||
bounding_box=bounding_box,
|
||||
)
|
||||
masks_for_seed = None
|
||||
if not boxes and seed_mask is not None:
|
||||
masks_for_seed = seed_mask.detach().to(device="cpu", dtype=torch.float32)
|
||||
if masks_for_seed.ndim == 2:
|
||||
masks_for_seed = masks_for_seed.unsqueeze(0)
|
||||
if masks_for_seed.ndim != 3 or tuple(masks_for_seed.shape[-2:]) != (
|
||||
height,
|
||||
width,
|
||||
):
|
||||
raise ValueError("seed_mask must have shape [objects, height, width].")
|
||||
object_ids = list(range(1, masks_for_seed.shape[0] + 1))
|
||||
labels = dict.fromkeys(object_ids)
|
||||
if not object_ids:
|
||||
raise ValueError(
|
||||
"Connect detections, a BOUNDING_BOX, or at least one seed mask."
|
||||
)
|
||||
|
||||
model = self.handle.ensure_loaded()
|
||||
device = model_device(model)
|
||||
state_device = torch.device("cpu") if keep_video_on_cpu else device
|
||||
processing_device = state_device if keep_video_on_cpu else device
|
||||
object_count = len(object_ids)
|
||||
union = torch.zeros((len(pil_images), height, width), dtype=torch.float32)
|
||||
if mask_output == "union_and_objects":
|
||||
individual = torch.zeros(
|
||||
(len(pil_images) * object_count, height, width),
|
||||
dtype=torch.float32,
|
||||
)
|
||||
else:
|
||||
individual = torch.zeros((0, height, width), dtype=torch.float32)
|
||||
source_preview = images.detach().to(device="cpu", dtype=torch.float32)
|
||||
preview = source_preview.clone() if render_preview else source_preview
|
||||
with torch.inference_mode(), inference_context(device, self.dtype):
|
||||
session = self.processor.init_video_session(
|
||||
video=pil_images,
|
||||
inference_device=device,
|
||||
inference_state_device=state_device,
|
||||
processing_device=processing_device,
|
||||
video_storage_device=state_device,
|
||||
max_vision_features_cache_size=1,
|
||||
dtype=self.dtype,
|
||||
)
|
||||
seed_kwargs: dict[str, Any] = {
|
||||
"inference_session": session,
|
||||
"frame_idx": int(seed_frame),
|
||||
"obj_ids": object_ids,
|
||||
}
|
||||
if boxes:
|
||||
seed_kwargs["input_boxes"] = [boxes]
|
||||
else:
|
||||
seed_kwargs["input_masks"] = [
|
||||
masks_for_seed[index] for index in range(len(object_ids))
|
||||
]
|
||||
self.processor.add_inputs_to_inference_session(**seed_kwargs)
|
||||
|
||||
per_object: dict[int, dict[int, Detection]] = {
|
||||
object_id: {} for object_id in object_ids
|
||||
}
|
||||
|
||||
def record_output(output):
|
||||
frame_index = int(output.frame_idx)
|
||||
processed = self.processor.post_process_masks(
|
||||
[output.pred_masks],
|
||||
original_sizes=[[height, width]],
|
||||
mask_threshold=float(mask_threshold),
|
||||
binarize=True,
|
||||
)[0]
|
||||
processed = _normalize_processed_masks(processed)
|
||||
current_ids = list(getattr(session, "obj_ids", object_ids))
|
||||
if render_preview:
|
||||
# The Transformers iterator may revisit the seed frame in
|
||||
# both directions. Rebuild that frame from the immutable
|
||||
# source so opacity is never accumulated across visits.
|
||||
preview[frame_index].copy_(source_preview[frame_index])
|
||||
for object_index, object_id in enumerate(current_ids):
|
||||
if object_index >= processed.shape[0]:
|
||||
continue
|
||||
mask = processed[object_index]
|
||||
float_mask = mask.to(dtype=torch.float32)
|
||||
union[frame_index] = torch.maximum(union[frame_index], float_mask)
|
||||
if mask_output == "union_and_objects":
|
||||
individual[frame_index * object_count + object_index] = (
|
||||
float_mask
|
||||
)
|
||||
if render_preview:
|
||||
color = torch.tensor(
|
||||
deterministic_color(object_id),
|
||||
dtype=preview.dtype,
|
||||
).div(255.0)
|
||||
alpha = float_mask.unsqueeze(-1) * 0.45
|
||||
preview[frame_index] = (
|
||||
preview[frame_index] * (1.0 - alpha) + color * alpha
|
||||
)
|
||||
bbox = bbox_from_mask(mask)
|
||||
if bbox is None:
|
||||
continue
|
||||
per_object.setdefault(int(object_id), {})[frame_index] = Detection(
|
||||
bbox_xyxy=bbox,
|
||||
label=labels.get(int(object_id)),
|
||||
frame_index=frame_index,
|
||||
timestamp=frame_index / fps,
|
||||
track_id=int(object_id),
|
||||
source=self.spec.model_id,
|
||||
metadata={
|
||||
"observation": (
|
||||
"detected"
|
||||
if frame_index == seed_frame
|
||||
else "propagated"
|
||||
),
|
||||
**(
|
||||
{
|
||||
"mask_batch_index": (
|
||||
frame_index * object_count + object_index
|
||||
)
|
||||
}
|
||||
if mask_output == "union_and_objects"
|
||||
else {"object_mask_output": "disabled"}
|
||||
),
|
||||
},
|
||||
)
|
||||
|
||||
# SAM2 does not consider prompt insertion itself an inference pass.
|
||||
# Running the conditioned frame establishes the track start before
|
||||
# either propagation direction is requested.
|
||||
record_output(model(inference_session=session, frame_idx=seed_frame))
|
||||
for output in model.propagate_in_video_iterator(
|
||||
inference_session=session,
|
||||
start_frame_idx=seed_frame,
|
||||
show_progress_bar=False,
|
||||
):
|
||||
record_output(output)
|
||||
if seed_frame > 0:
|
||||
for output in model.propagate_in_video_iterator(
|
||||
inference_session=session,
|
||||
start_frame_idx=seed_frame,
|
||||
reverse=True,
|
||||
show_progress_bar=False,
|
||||
):
|
||||
record_output(output)
|
||||
|
||||
tracks = tuple(
|
||||
Track(
|
||||
track_id=object_id,
|
||||
detections=tuple(
|
||||
records[frame_index] for frame_index in sorted(records)
|
||||
),
|
||||
label=labels.get(object_id),
|
||||
source=self.spec.model_id,
|
||||
metadata={"backend": "transformers-sam2-video"},
|
||||
)
|
||||
for object_id, records in sorted(per_object.items())
|
||||
if records
|
||||
)
|
||||
track_sequence = TrackSequence(
|
||||
width=width,
|
||||
height=height,
|
||||
tracks=tracks,
|
||||
frame_count=len(pil_images),
|
||||
fps=fps,
|
||||
source=self.spec.model_id,
|
||||
metadata={
|
||||
"seed_frame": seed_frame,
|
||||
"object_ids": object_ids,
|
||||
"mask_output": mask_output,
|
||||
"mask_order": (
|
||||
"frame_major_object_minor"
|
||||
if mask_output == "union_and_objects"
|
||||
else None
|
||||
),
|
||||
},
|
||||
)
|
||||
if render_preview:
|
||||
preview.clamp_(0, 1)
|
||||
return track_sequence, union, individual, preview
|
||||
|
||||
|
||||
class VLMSAM2VideoSegmentation(CachedModelNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
"model": (tuple(SAM2_MODELS),),
|
||||
"seed_frame": ("INT", {"default": 0, "min": 0, "max": 1000000}),
|
||||
"fps": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 24.0,
|
||||
"min": 0.001,
|
||||
"max": 1000.0,
|
||||
"step": 0.001,
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"detections": (VLM_DETECTIONS,),
|
||||
"bounding_box": ("BOUNDING_BOX",),
|
||||
"seed_mask": ("MASK",),
|
||||
"mask_threshold": (
|
||||
"FLOAT",
|
||||
{"default": 0.0, "min": -10.0, "max": 10.0, "step": 0.05},
|
||||
),
|
||||
"precision": (("auto", "bfloat16", "float16", "float32"),),
|
||||
"keep_video_on_cpu": ("BOOLEAN", {"default": True}),
|
||||
"mask_output": (
|
||||
("union_only", "union_and_objects"),
|
||||
{
|
||||
"default": "union_only",
|
||||
"tooltip": (
|
||||
"Per-object full-resolution masks can be very large."
|
||||
),
|
||||
},
|
||||
),
|
||||
"render_preview": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": True,
|
||||
"tooltip": (
|
||||
"Disable to return the input batch without another "
|
||||
"full-size overlay copy."
|
||||
),
|
||||
},
|
||||
),
|
||||
"unload_after": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = (VLM_TRACKS, "STRING", "MASK", "MASK", "IMAGE")
|
||||
RETURN_NAMES = (
|
||||
"tracks",
|
||||
"json",
|
||||
"union_masks",
|
||||
"object_masks",
|
||||
"preview",
|
||||
)
|
||||
FUNCTION = "segment"
|
||||
CATEGORY = "VLM Nodes/Vision/Segmentation"
|
||||
DESCRIPTION = (
|
||||
"Track detection boxes or masks through an IMAGE batch with SAM2.1. "
|
||||
"Connect the fps output of Get Video Components for correct timestamps."
|
||||
)
|
||||
|
||||
def segment(
|
||||
self,
|
||||
images,
|
||||
model,
|
||||
seed_frame,
|
||||
fps,
|
||||
detections=None,
|
||||
bounding_box=None,
|
||||
seed_mask=None,
|
||||
mask_threshold=0.0,
|
||||
precision="auto",
|
||||
keep_video_on_cpu=True,
|
||||
mask_output="union_only",
|
||||
render_preview=True,
|
||||
unload_after=False,
|
||||
):
|
||||
fps_value = float(fps)
|
||||
if not math.isfinite(fps_value) or fps_value <= 0:
|
||||
raise ValueError("fps must be finite and positive.")
|
||||
spec = SAM2_MODELS[model]
|
||||
predictor = self.get_or_create_model(
|
||||
(spec.model_id, precision),
|
||||
lambda: Sam2VideoPredictor(spec, precision),
|
||||
)
|
||||
try:
|
||||
tracks, union, individual, preview = predictor.propagate(
|
||||
images,
|
||||
seed_frame=int(seed_frame),
|
||||
fps=fps_value,
|
||||
detections=detections,
|
||||
bounding_box=bounding_box,
|
||||
seed_mask=seed_mask,
|
||||
mask_threshold=float(mask_threshold),
|
||||
keep_video_on_cpu=bool(keep_video_on_cpu),
|
||||
mask_output=mask_output,
|
||||
render_preview=bool(render_preview),
|
||||
)
|
||||
return tracks, tracks.to_json(indent=2), union, individual, preview
|
||||
finally:
|
||||
self.maybe_clear_model(unload_after)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"VLMSAM2VideoSegmentation": VLMSAM2VideoSegmentation,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"VLMSAM2VideoSegmentation": "VLM SAM2.1 Video Segmentation",
|
||||
}
|
||||
@@ -0,0 +1,560 @@
|
||||
"""Guarded adapters for ComfyUI core ``SAM3_TRACK_DATA`` payloads.
|
||||
|
||||
Core SAM3 keeps masks bit-packed for memory efficiency. This module preserves
|
||||
that payload untouched and emits small canonical ``VLM_TRACKS`` metadata with
|
||||
mask references instead of embedding dense masks in JSON.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import math
|
||||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from .geometry import bbox_from_mask, clip_box
|
||||
from .vision_types import (
|
||||
VLM_DETECTIONS,
|
||||
VLM_TRACKS,
|
||||
Detection,
|
||||
DetectionSequence,
|
||||
Track,
|
||||
TrackSequence,
|
||||
)
|
||||
|
||||
SAM3_TRACK_DATA = "SAM3_TRACK_DATA"
|
||||
SAM3_ADAPTER_SOURCE = "comfyui-core-sam3"
|
||||
_REQUIRED_KEYS = frozenset({"packed_masks", "n_frames", "scores", "orig_size"})
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SAM3TrackLayout:
|
||||
n_frames: int
|
||||
n_objects: int
|
||||
mask_height: int
|
||||
mask_width: int
|
||||
orig_height: int
|
||||
orig_width: int
|
||||
scores: tuple[float | None, ...]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _SeedIdentity:
|
||||
track_id: int | None
|
||||
label: str | None
|
||||
text: str | None
|
||||
score: float | None
|
||||
source: str | None
|
||||
|
||||
|
||||
def _integer(value: Any, name: str, *, minimum: int = 0) -> int:
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise TypeError(f"{name} must be an integer.")
|
||||
if value < minimum:
|
||||
raise ValueError(f"{name} must be at least {minimum}.")
|
||||
return value
|
||||
|
||||
|
||||
def _scores(
|
||||
values: Any,
|
||||
*,
|
||||
n_objects: int,
|
||||
) -> tuple[float | None, ...]:
|
||||
if isinstance(values, torch.Tensor):
|
||||
if values.ndim != 1:
|
||||
raise ValueError("SAM3 scores tensor must have shape [objects].")
|
||||
items = values.detach().cpu().tolist()
|
||||
elif isinstance(values, Sequence) and not isinstance(values, (str, bytes)):
|
||||
items = list(values)
|
||||
else:
|
||||
raise TypeError("SAM3 scores must be a one-dimensional sequence.")
|
||||
if len(items) > n_objects:
|
||||
raise ValueError("SAM3 scores contain more entries than mask objects.")
|
||||
parsed: list[float | None] = []
|
||||
for value in items:
|
||||
if value is None:
|
||||
parsed.append(None)
|
||||
continue
|
||||
score = float(value)
|
||||
if not math.isfinite(score) or not 0.0 <= score <= 1.0:
|
||||
raise ValueError("SAM3 scores must be finite values from 0 to 1.")
|
||||
parsed.append(score)
|
||||
parsed.extend([None] * (n_objects - len(parsed)))
|
||||
return tuple(parsed)
|
||||
|
||||
|
||||
def validate_sam3_track_data(track_data: Any) -> SAM3TrackLayout:
|
||||
"""Validate the private core payload before interpreting its bit layout."""
|
||||
|
||||
if not isinstance(track_data, Mapping):
|
||||
raise TypeError("SAM3_TRACK_DATA must be a mapping.")
|
||||
missing = sorted(_REQUIRED_KEYS - set(track_data))
|
||||
if missing:
|
||||
raise ValueError(
|
||||
"SAM3_TRACK_DATA is missing required keys: " + ", ".join(missing)
|
||||
)
|
||||
|
||||
n_frames = _integer(track_data["n_frames"], "n_frames")
|
||||
orig_size = track_data["orig_size"]
|
||||
if not isinstance(orig_size, (tuple, list)) or len(orig_size) != 2:
|
||||
raise TypeError("SAM3 orig_size must be (height, width).")
|
||||
orig_height = _integer(orig_size[0], "orig_height", minimum=1)
|
||||
orig_width = _integer(orig_size[1], "orig_width", minimum=1)
|
||||
|
||||
packed = track_data["packed_masks"]
|
||||
if packed is None:
|
||||
scores = _scores(track_data["scores"], n_objects=0)
|
||||
return SAM3TrackLayout(
|
||||
n_frames=n_frames,
|
||||
n_objects=0,
|
||||
mask_height=0,
|
||||
mask_width=0,
|
||||
orig_height=orig_height,
|
||||
orig_width=orig_width,
|
||||
scores=scores,
|
||||
)
|
||||
if not isinstance(packed, torch.Tensor):
|
||||
raise TypeError("SAM3 packed_masks must be a torch.Tensor or None.")
|
||||
if packed.dtype != torch.uint8:
|
||||
raise TypeError("SAM3 packed_masks must use torch.uint8.")
|
||||
if packed.ndim != 4:
|
||||
raise ValueError(
|
||||
"SAM3 packed_masks must have shape [frames, objects, height, packed_width]."
|
||||
)
|
||||
if packed.shape[0] != n_frames:
|
||||
raise ValueError("SAM3 n_frames does not match packed_masks.")
|
||||
n_objects = int(packed.shape[1])
|
||||
mask_height = int(packed.shape[2])
|
||||
packed_width = int(packed.shape[3])
|
||||
if n_objects < 1 or mask_height < 1 or packed_width < 1:
|
||||
raise ValueError("SAM3 packed_masks dimensions must be positive.")
|
||||
scores = _scores(track_data["scores"], n_objects=n_objects)
|
||||
return SAM3TrackLayout(
|
||||
n_frames=n_frames,
|
||||
n_objects=n_objects,
|
||||
mask_height=mask_height,
|
||||
mask_width=packed_width * 8,
|
||||
orig_height=orig_height,
|
||||
orig_width=orig_width,
|
||||
scores=scores,
|
||||
)
|
||||
|
||||
|
||||
def unpack_sam3_mask(packed_mask: torch.Tensor) -> torch.Tensor:
|
||||
"""Unpack exactly one object/frame mask, avoiding full-video expansion."""
|
||||
|
||||
if not isinstance(packed_mask, torch.Tensor):
|
||||
raise TypeError("packed_mask must be a torch.Tensor.")
|
||||
if packed_mask.dtype != torch.uint8 or packed_mask.ndim != 2:
|
||||
raise ValueError("packed_mask must be uint8 with shape [height, packed_width].")
|
||||
bits = torch.tensor(
|
||||
(1, 2, 4, 8, 16, 32, 64, 128),
|
||||
dtype=torch.uint8,
|
||||
device=packed_mask.device,
|
||||
)
|
||||
return (
|
||||
torch.bitwise_and(packed_mask.unsqueeze(-1), bits)
|
||||
.ne(0)
|
||||
.reshape(packed_mask.shape[0], packed_mask.shape[1] * 8)
|
||||
)
|
||||
|
||||
|
||||
def iter_sam3_masks(
|
||||
track_data: Mapping[str, Any],
|
||||
*,
|
||||
present_only: bool = False,
|
||||
) -> Iterator[tuple[int, int, torch.Tensor]]:
|
||||
"""Yield one unpacked mask at a time as ``(frame, object, mask)``."""
|
||||
|
||||
layout = validate_sam3_track_data(track_data)
|
||||
packed = track_data["packed_masks"]
|
||||
if packed is None:
|
||||
return
|
||||
for frame_index in range(layout.n_frames):
|
||||
for object_index in range(layout.n_objects):
|
||||
mask = unpack_sam3_mask(packed[frame_index, object_index])
|
||||
if present_only and not bool(mask.any().item()):
|
||||
continue
|
||||
yield frame_index, object_index, mask
|
||||
|
||||
|
||||
def _seeds_from_detections(
|
||||
sequence: DetectionSequence,
|
||||
) -> list[_SeedIdentity]:
|
||||
for frame in sequence.frames:
|
||||
if frame.detections:
|
||||
return [
|
||||
_SeedIdentity(
|
||||
track_id=detection.track_id,
|
||||
label=detection.label,
|
||||
text=detection.text,
|
||||
score=detection.score,
|
||||
source=detection.source,
|
||||
)
|
||||
for detection in frame.detections
|
||||
]
|
||||
return []
|
||||
|
||||
|
||||
def _seeds_from_tracks(sequence: TrackSequence) -> list[_SeedIdentity]:
|
||||
return [
|
||||
_SeedIdentity(
|
||||
track_id=track.track_id,
|
||||
label=track.label or track.detections[0].label,
|
||||
text=track.detections[0].text,
|
||||
score=track.score,
|
||||
source=track.source,
|
||||
)
|
||||
for track in sequence.tracks
|
||||
]
|
||||
|
||||
|
||||
def _seed_identities(
|
||||
*,
|
||||
seed_detections: DetectionSequence | None,
|
||||
seed_tracks: TrackSequence | None,
|
||||
n_objects: int,
|
||||
) -> tuple[tuple[_SeedIdentity, ...], int]:
|
||||
if seed_detections is not None and seed_tracks is not None:
|
||||
raise ValueError("Connect seed_detections or seed_tracks, not both.")
|
||||
if seed_detections is not None:
|
||||
if not isinstance(seed_detections, DetectionSequence):
|
||||
raise TypeError("seed_detections must be a DetectionSequence.")
|
||||
seeds = _seeds_from_detections(seed_detections)
|
||||
elif seed_tracks is not None:
|
||||
if not isinstance(seed_tracks, TrackSequence):
|
||||
raise TypeError("seed_tracks must be a TrackSequence.")
|
||||
seeds = _seeds_from_tracks(seed_tracks)
|
||||
else:
|
||||
seeds = []
|
||||
|
||||
used_ids: set[int] = set()
|
||||
next_id = 0
|
||||
identities = []
|
||||
for object_index in range(n_objects):
|
||||
seed = seeds[object_index] if object_index < len(seeds) else None
|
||||
preferred = None if seed is None else seed.track_id
|
||||
if preferred is not None and preferred not in used_ids:
|
||||
track_id = preferred
|
||||
else:
|
||||
while next_id in used_ids:
|
||||
next_id += 1
|
||||
track_id = next_id
|
||||
next_id += 1
|
||||
used_ids.add(track_id)
|
||||
identities.append(
|
||||
_SeedIdentity(
|
||||
track_id=track_id,
|
||||
label=None if seed is None else seed.label,
|
||||
text=None if seed is None else seed.text,
|
||||
score=None if seed is None else seed.score,
|
||||
source=None if seed is None else seed.source,
|
||||
)
|
||||
)
|
||||
return tuple(identities), min(len(seeds), n_objects)
|
||||
|
||||
|
||||
def _scaled_bbox(
|
||||
bbox: tuple[float, float, float, float],
|
||||
layout: SAM3TrackLayout,
|
||||
) -> tuple[float, float, float, float]:
|
||||
scale_x = layout.orig_width / layout.mask_width
|
||||
scale_y = layout.orig_height / layout.mask_height
|
||||
x1, y1, x2, y2 = bbox
|
||||
return clip_box(
|
||||
(
|
||||
x1 * scale_x,
|
||||
y1 * scale_y,
|
||||
x2 * scale_x,
|
||||
y2 * scale_y,
|
||||
),
|
||||
layout.orig_width,
|
||||
layout.orig_height,
|
||||
)
|
||||
|
||||
|
||||
def sam3_track_data_to_tracks(
|
||||
track_data: Mapping[str, Any],
|
||||
*,
|
||||
seed_detections: DetectionSequence | None = None,
|
||||
seed_tracks: TrackSequence | None = None,
|
||||
fps: float | None = None,
|
||||
source: str = SAM3_ADAPTER_SOURCE,
|
||||
) -> TrackSequence:
|
||||
"""Create canonical sparse metadata while retaining packed masks separately."""
|
||||
|
||||
layout = validate_sam3_track_data(track_data)
|
||||
if fps is not None:
|
||||
fps = float(fps)
|
||||
if not math.isfinite(fps) or fps <= 0:
|
||||
raise ValueError("fps must be finite and positive or None.")
|
||||
identities, seed_count = _seed_identities(
|
||||
seed_detections=seed_detections,
|
||||
seed_tracks=seed_tracks,
|
||||
n_objects=layout.n_objects,
|
||||
)
|
||||
detections_by_object: list[list[Detection]] = [
|
||||
[] for _index in range(layout.n_objects)
|
||||
]
|
||||
|
||||
for frame_index, object_index, mask in iter_sam3_masks(
|
||||
track_data, present_only=True
|
||||
):
|
||||
bbox = bbox_from_mask(mask)
|
||||
if bbox is None:
|
||||
continue
|
||||
identity = identities[object_index]
|
||||
timestamp = frame_index / fps if fps is not None else 0.0
|
||||
score = layout.scores[object_index]
|
||||
if score is None:
|
||||
score = identity.score
|
||||
detections_by_object[object_index].append(
|
||||
Detection(
|
||||
bbox_xyxy=_scaled_bbox(bbox, layout),
|
||||
label=identity.label,
|
||||
text=identity.text,
|
||||
score=score,
|
||||
frame_index=frame_index,
|
||||
timestamp=timestamp,
|
||||
track_id=identity.track_id,
|
||||
source=source,
|
||||
metadata={
|
||||
"observation": "propagated",
|
||||
"visibility": "visible",
|
||||
"sam3_object_index": object_index,
|
||||
"mask_ref": {
|
||||
"type": SAM3_TRACK_DATA,
|
||||
"frame_index": frame_index,
|
||||
"object_index": object_index,
|
||||
},
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
tracks = []
|
||||
for object_index, detections in enumerate(detections_by_object):
|
||||
if not detections:
|
||||
continue
|
||||
identity = identities[object_index]
|
||||
present_frames = len(detections)
|
||||
final_state = (
|
||||
"active" if detections[-1].frame_index == layout.n_frames - 1 else "lost"
|
||||
)
|
||||
score = layout.scores[object_index]
|
||||
if score is None:
|
||||
score = identity.score
|
||||
tracks.append(
|
||||
Track(
|
||||
track_id=identity.track_id,
|
||||
detections=tuple(detections),
|
||||
label=identity.label,
|
||||
score=score,
|
||||
source=source,
|
||||
metadata={
|
||||
"state": final_state,
|
||||
"sam3_object_index": object_index,
|
||||
"first_frame": detections[0].frame_index,
|
||||
"last_observed_frame": detections[-1].frame_index,
|
||||
"present_frames": present_frames,
|
||||
"presence_ratio": (
|
||||
present_frames / layout.n_frames if layout.n_frames else 0.0
|
||||
),
|
||||
"seeded": object_index < seed_count,
|
||||
},
|
||||
)
|
||||
)
|
||||
return TrackSequence(
|
||||
width=layout.orig_width,
|
||||
height=layout.orig_height,
|
||||
tracks=tuple(sorted(tracks, key=lambda item: item.track_id)),
|
||||
frame_count=layout.n_frames,
|
||||
fps=fps,
|
||||
source=source,
|
||||
metadata={
|
||||
"adapter": "sam3-track-data/v1",
|
||||
"mask_payload": {
|
||||
"type": SAM3_TRACK_DATA,
|
||||
"encoding": "little-endian-bitpack",
|
||||
"mask_width": layout.mask_width,
|
||||
"mask_height": layout.mask_height,
|
||||
},
|
||||
"object_slots": layout.n_objects,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def track_report_payload(tracks: TrackSequence) -> dict[str, Any]:
|
||||
"""Return a compact, history-safe report with no tensor content."""
|
||||
|
||||
if not isinstance(tracks, TrackSequence):
|
||||
raise TypeError("tracks must be a TrackSequence.")
|
||||
records = []
|
||||
state_counts: dict[str, int] = {}
|
||||
total_observations = 0
|
||||
for track in tracks.tracks:
|
||||
observations: dict[str, int] = {}
|
||||
for detection in track.detections:
|
||||
kind = str(detection.metadata.to_dict().get("observation", "detected"))
|
||||
observations[kind] = observations.get(kind, 0) + 1
|
||||
total_observations += len(track.detections)
|
||||
state = str(track.metadata.to_dict().get("state", "unknown"))
|
||||
state_counts[state] = state_counts.get(state, 0) + 1
|
||||
records.append(
|
||||
{
|
||||
"track_id": track.track_id,
|
||||
"label": track.label,
|
||||
"state": state,
|
||||
"score": track.score,
|
||||
"first_frame": track.detections[0].frame_index,
|
||||
"last_frame": track.detections[-1].frame_index,
|
||||
"observation_count": len(track.detections),
|
||||
"observations": observations,
|
||||
}
|
||||
)
|
||||
media: dict[str, Any] = {
|
||||
"width": tracks.width,
|
||||
"height": tracks.height,
|
||||
"frame_count": tracks.frame_count,
|
||||
}
|
||||
if tracks.fps is not None:
|
||||
media["fps"] = tracks.fps
|
||||
return {
|
||||
"schema": "comfyui-vlm/track-report",
|
||||
"version": 1,
|
||||
"media": media,
|
||||
"track_count": len(tracks.tracks),
|
||||
"observation_count": total_observations,
|
||||
"state_counts": dict(sorted(state_counts.items())),
|
||||
"tracks": records,
|
||||
}
|
||||
|
||||
|
||||
def track_report_json(
|
||||
tracks: TrackSequence,
|
||||
*,
|
||||
indent: int | None = 2,
|
||||
) -> str:
|
||||
return json.dumps(
|
||||
track_report_payload(tracks),
|
||||
ensure_ascii=False,
|
||||
allow_nan=False,
|
||||
sort_keys=True,
|
||||
indent=indent,
|
||||
)
|
||||
|
||||
|
||||
def track_report_text(tracks: TrackSequence) -> str:
|
||||
report = track_report_payload(tracks)
|
||||
media = report["media"]
|
||||
lines = [
|
||||
(
|
||||
f"Tracks: {report['track_count']} | "
|
||||
f"Observations: {report['observation_count']} | "
|
||||
f"Frames: {media['frame_count']} | "
|
||||
f"Size: {media['width']}x{media['height']}"
|
||||
)
|
||||
]
|
||||
if "fps" in media:
|
||||
lines[0] += f" | FPS: {media['fps']:g}"
|
||||
for track in report["tracks"]:
|
||||
label = track["label"] or "(unlabeled)"
|
||||
lines.append(
|
||||
f"#{track['track_id']} {label}: {track['state']}, "
|
||||
f"frames {track['first_frame']}-{track['last_frame']}, "
|
||||
f"{track['observation_count']} observations"
|
||||
)
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
class VLMSAM3TrackAdapter:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"track_data": (SAM3_TRACK_DATA,),
|
||||
"fps": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.0,
|
||||
"min": 0.0,
|
||||
"max": 1000.0,
|
||||
"step": 0.01,
|
||||
"tooltip": "0 keeps timestamps unknown.",
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"seed_detections": (VLM_DETECTIONS,),
|
||||
"seed_tracks": (VLM_TRACKS,),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = (VLM_TRACKS, SAM3_TRACK_DATA)
|
||||
RETURN_NAMES = ("tracks", "track_data")
|
||||
FUNCTION = "adapt"
|
||||
CATEGORY = "VLM Nodes/Vision/Tracking"
|
||||
|
||||
def adapt(
|
||||
self,
|
||||
track_data,
|
||||
fps,
|
||||
seed_detections=None,
|
||||
seed_tracks=None,
|
||||
):
|
||||
tracks = sam3_track_data_to_tracks(
|
||||
track_data,
|
||||
seed_detections=seed_detections,
|
||||
seed_tracks=seed_tracks,
|
||||
fps=None if fps <= 0 else fps,
|
||||
)
|
||||
return tracks, track_data
|
||||
|
||||
|
||||
class VLMTrackReport:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {"required": {"tracks": (VLM_TRACKS,)}}
|
||||
|
||||
RETURN_TYPES = ("STRING", "STRING")
|
||||
RETURN_NAMES = ("report_json", "report_text")
|
||||
FUNCTION = "report"
|
||||
CATEGORY = "VLM Nodes/Vision/Tracking"
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def report(self, tracks):
|
||||
report_json = track_report_json(tracks)
|
||||
report_text = track_report_text(tracks)
|
||||
return {
|
||||
"ui": {"text": [report_text]},
|
||||
"result": (report_json, report_text),
|
||||
}
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"VLMSAM3TrackAdapter": VLMSAM3TrackAdapter,
|
||||
"VLMTrackReport": VLMTrackReport,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"VLMSAM3TrackAdapter": "VLM SAM3 Track Adapter",
|
||||
"VLMTrackReport": "VLM Track Report",
|
||||
}
|
||||
|
||||
__all__ = [
|
||||
"NODE_CLASS_MAPPINGS",
|
||||
"NODE_DISPLAY_NAME_MAPPINGS",
|
||||
"SAM3TrackLayout",
|
||||
"SAM3_ADAPTER_SOURCE",
|
||||
"SAM3_TRACK_DATA",
|
||||
"VLMSAM3TrackAdapter",
|
||||
"VLMTrackReport",
|
||||
"iter_sam3_masks",
|
||||
"sam3_track_data_to_tracks",
|
||||
"track_report_json",
|
||||
"track_report_payload",
|
||||
"track_report_text",
|
||||
"unpack_sam3_mask",
|
||||
"validate_sam3_track_data",
|
||||
]
|
||||
+1074
-101
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
+3
-222
@@ -8,15 +8,14 @@ producing stricter output.
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
from typing import Any, Literal, Optional
|
||||
from typing import Any
|
||||
|
||||
import folder_paths
|
||||
import torch
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from .prompts import system_msg_prompts, system_msg_simple
|
||||
from .prompts import system_msg_prompts
|
||||
from .runtime import (
|
||||
LlamaHandle,
|
||||
close_handle,
|
||||
@@ -69,7 +68,7 @@ class ArtisticTechniques(BaseModel):
|
||||
|
||||
class ImageryTheme(BaseModel):
|
||||
core_subject: str
|
||||
additional_elements: Optional[list[str]] = None
|
||||
additional_elements: list[str] | None = None
|
||||
|
||||
|
||||
class VisualStyle(BaseModel):
|
||||
@@ -157,222 +156,6 @@ def _structured_chat(
|
||||
return raw, parsed
|
||||
|
||||
|
||||
API_MODELS = [
|
||||
"GPT-5.6 Terra",
|
||||
"GPT-5.6 Sol",
|
||||
"GPT-5.6 Luna",
|
||||
"DeepSeek",
|
||||
"Custom / OpenAI-compatible",
|
||||
# Kept so saved workflows continue to deserialize without substitutions.
|
||||
"ChatGPT-3.5",
|
||||
"ChatGPT-4",
|
||||
"gpt-3.5-turbo",
|
||||
"gpt-3.5-turbo-0125",
|
||||
"gpt-35-turbo",
|
||||
"gpt-3.5-turbo-16k",
|
||||
"gpt-3.5-turbo-16k-0613",
|
||||
"gpt-4-0613",
|
||||
"gpt-4-1106-preview",
|
||||
"glm-4",
|
||||
]
|
||||
|
||||
API_ROUTES = {
|
||||
"GPT-5.6 Sol": ("gpt-5.6-sol", None, "Responses"),
|
||||
"GPT-5.6 Terra": ("gpt-5.6-terra", None, "Responses"),
|
||||
"GPT-5.6 Luna": ("gpt-5.6-luna", None, "Responses"),
|
||||
"DeepSeek": ("deepseek-chat", "https://api.deepseek.com/v1", "Chat Completions"),
|
||||
"ChatGPT-3.5": ("gpt-3.5-turbo", None, "Chat Completions"),
|
||||
"ChatGPT-4": ("gpt-4", None, "Chat Completions"),
|
||||
"gpt-35-turbo": ("gpt-35-turbo", None, "Chat Completions"),
|
||||
"glm-4": ("glm-4", None, "Chat Completions"),
|
||||
}
|
||||
|
||||
|
||||
class PromptGenerateAPI:
|
||||
def __init__(self):
|
||||
self.session_history: list[dict[str, str]] = []
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model_name": (API_MODELS, {"default": "GPT-5.6 Terra"}),
|
||||
"chat_type": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": True,
|
||||
"label_on": "Prompt Generator",
|
||||
"label_off": "Simple Chat",
|
||||
},
|
||||
),
|
||||
"api_key": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": (
|
||||
"Leave blank to use OPENAI_API_KEY, DEEPSEEK_API_KEY, "
|
||||
"or VLM_API_KEY."
|
||||
),
|
||||
},
|
||||
),
|
||||
"description": (
|
||||
"STRING",
|
||||
{"multiline": True, "default": ""},
|
||||
),
|
||||
"question": (
|
||||
"STRING",
|
||||
{"multiline": True, "default": ""},
|
||||
),
|
||||
"context_size": (
|
||||
"INT",
|
||||
{"default": 5, "min": 0, "max": 30, "step": 1},
|
||||
),
|
||||
"seed": (
|
||||
"INT",
|
||||
{
|
||||
"default": 0,
|
||||
"min": 0,
|
||||
"max": 0xFFFFFFFFFFFFFFFF,
|
||||
"step": 1,
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"base_url": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": (
|
||||
"OpenAI-compatible base URL, e.g. http://127.0.0.1:8000/v1."
|
||||
),
|
||||
},
|
||||
),
|
||||
"model_override": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "Exact provider model ID. Overrides the picker.",
|
||||
},
|
||||
),
|
||||
"api_mode": (
|
||||
["Auto", "Responses", "Chat Completions"],
|
||||
{"default": "Auto"},
|
||||
),
|
||||
"timeout_seconds": (
|
||||
"FLOAT",
|
||||
{"default": 120.0, "min": 1.0, "max": 1800.0},
|
||||
),
|
||||
"reasoning_effort": (
|
||||
["none", "low", "medium", "high", "xhigh", "max"],
|
||||
{"default": "none"},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "generate_prompt"
|
||||
CATEGORY = "VLM Nodes/LLM"
|
||||
|
||||
def _route(
|
||||
self, model_name, model_override, base_url, api_mode
|
||||
) -> tuple[str, str | None, str]:
|
||||
route = API_ROUTES.get(model_name)
|
||||
if route is None:
|
||||
if model_name in API_MODELS and model_name not in {
|
||||
"Custom / OpenAI-compatible"
|
||||
}:
|
||||
route = (model_name, None, "Chat Completions")
|
||||
else:
|
||||
route = ("", None, "Chat Completions")
|
||||
model, route_url, route_mode = route
|
||||
model = (model_override or model).strip()
|
||||
if not model:
|
||||
raise ValueError("A model ID is required for Custom / OpenAI-compatible.")
|
||||
effective_url = (base_url or route_url or "").strip() or None
|
||||
mode = route_mode if api_mode == "Auto" else api_mode
|
||||
return model, effective_url, mode
|
||||
|
||||
def generate_prompt(
|
||||
self,
|
||||
model_name,
|
||||
chat_type,
|
||||
api_key,
|
||||
description,
|
||||
question,
|
||||
context_size,
|
||||
seed,
|
||||
base_url="",
|
||||
model_override="",
|
||||
api_mode="Auto",
|
||||
timeout_seconds=120.0,
|
||||
reasoning_effort="none",
|
||||
):
|
||||
openai = require_module("openai", "openai")
|
||||
model, effective_url, mode = self._route(
|
||||
model_name, model_override, base_url, api_mode
|
||||
)
|
||||
key = (
|
||||
api_key.strip()
|
||||
or (os.getenv("DEEPSEEK_API_KEY", "") if model_name == "DeepSeek" else "")
|
||||
or os.getenv("VLM_API_KEY", "")
|
||||
or os.getenv("OPENAI_API_KEY", "")
|
||||
)
|
||||
if not key:
|
||||
raise ValueError(
|
||||
"No API key was supplied. Set OPENAI_API_KEY, "
|
||||
"DEEPSEEK_API_KEY, or VLM_API_KEY, or enter the key in the node."
|
||||
)
|
||||
|
||||
client_kwargs: dict[str, Any] = {
|
||||
"api_key": key,
|
||||
"timeout": float(timeout_seconds),
|
||||
"max_retries": 2,
|
||||
}
|
||||
if effective_url:
|
||||
client_kwargs["base_url"] = effective_url
|
||||
client = openai.OpenAI(**client_kwargs)
|
||||
|
||||
system = system_msg_prompts if chat_type else system_msg_simple
|
||||
user_message = (
|
||||
f"Description:\n{description.strip()}\n\n"
|
||||
f"Optional question:\n{question.strip()}"
|
||||
).strip()
|
||||
history_limit = max(0, int(context_size)) * 2
|
||||
history = self.session_history[-history_limit:] if history_limit else []
|
||||
|
||||
if mode == "Responses":
|
||||
response = client.responses.create(
|
||||
model=model,
|
||||
instructions=system,
|
||||
input=history + [{"role": "user", "content": user_message}],
|
||||
reasoning={"effort": reasoning_effort},
|
||||
)
|
||||
result = response.output_text
|
||||
else:
|
||||
messages = (
|
||||
[{"role": "system", "content": system}]
|
||||
+ history
|
||||
+ [{"role": "user", "content": user_message}]
|
||||
)
|
||||
request: dict[str, Any] = {
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
"seed": int(seed),
|
||||
}
|
||||
if model.startswith("gpt-5.6"):
|
||||
request["reasoning_effort"] = reasoning_effort
|
||||
completion = client.chat.completions.create(**request)
|
||||
result = completion.choices[0].message.content or ""
|
||||
|
||||
self.session_history.extend(
|
||||
[
|
||||
{"role": "user", "content": user_message},
|
||||
{"role": "assistant", "content": result},
|
||||
]
|
||||
)
|
||||
return (result,)
|
||||
|
||||
|
||||
class LLMLoader:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
@@ -1206,7 +989,6 @@ NODE_CLASS_MAPPINGS = {
|
||||
"KeywordExtraction": KeywordExtraction,
|
||||
"LLavaPromptGenerator": LLavaPromptGenerator,
|
||||
"Suggester": Suggester,
|
||||
"PromptGenerateAPI": PromptGenerateAPI,
|
||||
"CreativeArtPromptGenerator": CreativeArtPromptGenerator,
|
||||
"ChatMusician": ChatMusician,
|
||||
"StructuredOutput": StructuredOutput,
|
||||
@@ -1221,7 +1003,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"KeywordExtraction": "Structured Keyword Extraction",
|
||||
"LLavaPromptGenerator": "Structured Prompt Generator",
|
||||
"Suggester": "Prompt Suggester",
|
||||
"PromptGenerateAPI": "OpenAI-Compatible Prompt API",
|
||||
"CreativeArtPromptGenerator": "Creative Art Prompt Generator",
|
||||
"ChatMusician": "Chat Musician",
|
||||
"StructuredOutput": "Structured Output",
|
||||
|
||||
@@ -0,0 +1,766 @@
|
||||
"""Deterministic tracking-by-detection for canonical VLM vision payloads.
|
||||
|
||||
The tracker intentionally owns only temporal association and identity. Dense
|
||||
mask propagation remains the responsibility of SAM-style video models. This
|
||||
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
|
||||
|
||||
import numpy as np
|
||||
from scipy.optimize import linear_sum_assignment
|
||||
|
||||
from .geometry import bbox_iou, clip_box, mask_iou
|
||||
from .vision_types import (
|
||||
VLM_DETECTIONS,
|
||||
VLM_TRACKS,
|
||||
Detection,
|
||||
DetectionSequence,
|
||||
FrozenDict,
|
||||
Track,
|
||||
TrackSequence,
|
||||
)
|
||||
|
||||
TRACKER_SOURCE = "vlm-bytetrack"
|
||||
_CHI_SQUARE_FOUR_DOF_99 = 13.2767
|
||||
_MIN_SIZE = 1.0e-3
|
||||
|
||||
|
||||
def _normalized_label(value: str | None) -> str | None:
|
||||
if value is None:
|
||||
return None
|
||||
normalized = " ".join(value.casefold().split())
|
||||
return normalized or None
|
||||
|
||||
|
||||
def _score_or_one(detection: Detection) -> float:
|
||||
return 1.0 if detection.score is None else detection.score
|
||||
|
||||
|
||||
def _box_to_measurement(box: Iterable[float]) -> np.ndarray:
|
||||
x1, y1, x2, y2 = (float(value) for value in box)
|
||||
return np.asarray(
|
||||
(
|
||||
(x1 + x2) * 0.5,
|
||||
(y1 + y2) * 0.5,
|
||||
max(x2 - x1, _MIN_SIZE),
|
||||
max(y2 - y1, _MIN_SIZE),
|
||||
),
|
||||
dtype=np.float64,
|
||||
)
|
||||
|
||||
|
||||
def _measurement_to_box(measurement: np.ndarray) -> tuple[float, ...]:
|
||||
center_x, center_y, width, height = measurement[:4]
|
||||
width = max(float(width), _MIN_SIZE)
|
||||
height = max(float(height), _MIN_SIZE)
|
||||
return (
|
||||
float(center_x - width * 0.5),
|
||||
float(center_y - height * 0.5),
|
||||
float(center_x + width * 0.5),
|
||||
float(center_y + height * 0.5),
|
||||
)
|
||||
|
||||
|
||||
class _BoxKalmanFilter:
|
||||
"""Small constant-velocity Kalman filter with no optional dependencies."""
|
||||
|
||||
_observation = np.concatenate(
|
||||
(np.eye(4, dtype=np.float64), np.zeros((4, 4), dtype=np.float64)),
|
||||
axis=1,
|
||||
)
|
||||
|
||||
def __init__(self, box: Iterable[float]):
|
||||
measurement = _box_to_measurement(box)
|
||||
self.mean = np.concatenate((measurement, np.zeros(4, dtype=np.float64)))
|
||||
scale = max(measurement[2], measurement[3], 1.0)
|
||||
self.covariance = np.diag(
|
||||
(
|
||||
scale * scale * 0.01,
|
||||
scale * scale * 0.01,
|
||||
scale * scale * 0.04,
|
||||
scale * scale * 0.04,
|
||||
scale * scale,
|
||||
scale * scale,
|
||||
scale * scale * 0.25,
|
||||
scale * scale * 0.25,
|
||||
)
|
||||
)
|
||||
|
||||
@property
|
||||
def box(self) -> tuple[float, ...]:
|
||||
return _measurement_to_box(self.mean)
|
||||
|
||||
def predict(self, delta_seconds: float) -> None:
|
||||
delta = max(float(delta_seconds), 1.0e-6)
|
||||
transition = np.eye(8, dtype=np.float64)
|
||||
transition[:4, 4:] = np.eye(4, dtype=np.float64) * delta
|
||||
scale = max(self.mean[2], self.mean[3], 1.0)
|
||||
position_noise = max(scale * 0.02 * delta, 1.0e-3)
|
||||
velocity_noise = max(scale * 0.01 * math.sqrt(delta), 1.0e-3)
|
||||
process_noise = np.diag((position_noise,) * 4 + (velocity_noise,) * 4) ** 2
|
||||
self.mean = transition @ self.mean
|
||||
self.covariance = transition @ self.covariance @ transition.T + process_noise
|
||||
self.mean[2:4] = np.maximum(self.mean[2:4], _MIN_SIZE)
|
||||
|
||||
def projected(self) -> tuple[np.ndarray, np.ndarray]:
|
||||
scale = max(self.mean[2], self.mean[3], 1.0)
|
||||
measurement_noise = (
|
||||
np.diag(
|
||||
(
|
||||
max(scale * 0.025, 1.0e-3),
|
||||
max(scale * 0.025, 1.0e-3),
|
||||
max(scale * 0.05, 1.0e-3),
|
||||
max(scale * 0.05, 1.0e-3),
|
||||
)
|
||||
)
|
||||
** 2
|
||||
)
|
||||
projected_mean = self._observation @ self.mean
|
||||
projected_covariance = (
|
||||
self._observation @ self.covariance @ self._observation.T
|
||||
+ measurement_noise
|
||||
)
|
||||
return projected_mean, projected_covariance
|
||||
|
||||
def gating_distance(self, box: Iterable[float]) -> float:
|
||||
measurement = _box_to_measurement(box)
|
||||
projected_mean, projected_covariance = self.projected()
|
||||
residual = measurement - projected_mean
|
||||
try:
|
||||
solved = np.linalg.solve(projected_covariance, residual)
|
||||
except np.linalg.LinAlgError:
|
||||
solved = np.linalg.pinv(projected_covariance) @ residual
|
||||
return float(residual @ solved)
|
||||
|
||||
def update(self, box: Iterable[float]) -> None:
|
||||
measurement = _box_to_measurement(box)
|
||||
projected_mean, projected_covariance = self.projected()
|
||||
cross_covariance = self.covariance @ self._observation.T
|
||||
try:
|
||||
gain = np.linalg.solve(projected_covariance, cross_covariance.T).T
|
||||
except np.linalg.LinAlgError:
|
||||
gain = cross_covariance @ np.linalg.pinv(projected_covariance)
|
||||
innovation = measurement - projected_mean
|
||||
self.mean = self.mean + gain @ innovation
|
||||
identity = np.eye(8, dtype=np.float64)
|
||||
residual_projection = identity - gain @ self._observation
|
||||
self.covariance = residual_projection @ self.covariance @ residual_projection.T
|
||||
self.mean[2:4] = np.maximum(self.mean[2:4], _MIN_SIZE)
|
||||
|
||||
|
||||
def _merged_metadata(
|
||||
detection: Detection,
|
||||
*,
|
||||
observation: str,
|
||||
track_state: str,
|
||||
association_stage: str,
|
||||
association_score: float | None,
|
||||
) -> FrozenDict:
|
||||
metadata = detection.metadata.to_dict()
|
||||
if detection.track_id is not None:
|
||||
metadata.setdefault("source_track_id", detection.track_id)
|
||||
metadata.update(
|
||||
{
|
||||
"observation": observation,
|
||||
"track_state": track_state,
|
||||
"association_stage": association_stage,
|
||||
"association_score": association_score,
|
||||
}
|
||||
)
|
||||
return FrozenDict(metadata)
|
||||
|
||||
|
||||
def _tracked_detection(
|
||||
detection: Detection,
|
||||
*,
|
||||
track_id: int,
|
||||
timestamp: float,
|
||||
track_state: str,
|
||||
association_stage: str,
|
||||
association_score: float | None,
|
||||
) -> Detection:
|
||||
return Detection(
|
||||
bbox_xyxy=detection.bbox_xyxy,
|
||||
label=detection.label,
|
||||
text=detection.text,
|
||||
score=detection.score,
|
||||
polygon=detection.polygon,
|
||||
quad=detection.quad,
|
||||
frame_index=detection.frame_index,
|
||||
timestamp=timestamp,
|
||||
track_id=track_id,
|
||||
source=detection.source,
|
||||
metadata=_merged_metadata(
|
||||
detection,
|
||||
observation="detected",
|
||||
track_state=track_state,
|
||||
association_stage=association_stage,
|
||||
association_score=association_score,
|
||||
),
|
||||
mask=detection.mask,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class _TrackState:
|
||||
track_id: int
|
||||
filter: _BoxKalmanFilter
|
||||
detections: list[Detection]
|
||||
label: str | None
|
||||
text: str | None
|
||||
state: str
|
||||
hits: int
|
||||
first_frame: int
|
||||
last_observed_frame: int
|
||||
last_observed_timestamp: float
|
||||
last_timestamp: float
|
||||
last_mask: object | None = None
|
||||
misses: int = 0
|
||||
removed_frame: int | None = None
|
||||
observed_scores: list[float] = field(default_factory=list)
|
||||
|
||||
@property
|
||||
def predicted_box(self) -> tuple[float, ...]:
|
||||
return self.filter.box
|
||||
|
||||
|
||||
def _labels_compatible(
|
||||
track: _TrackState,
|
||||
detection: Detection,
|
||||
*,
|
||||
label_aware: bool,
|
||||
) -> bool:
|
||||
if not label_aware:
|
||||
return True
|
||||
old_label = _normalized_label(track.label)
|
||||
new_label = _normalized_label(detection.label)
|
||||
return old_label is None or new_label is None or old_label == new_label
|
||||
|
||||
|
||||
def _overlap(track: _TrackState, detection: Detection) -> float:
|
||||
overlap = bbox_iou(track.predicted_box, detection.bbox_xyxy)
|
||||
if track.last_mask is not None and detection.mask is not None:
|
||||
try:
|
||||
overlap = max(
|
||||
overlap,
|
||||
mask_iou(track.last_mask, detection.mask),
|
||||
)
|
||||
except (TypeError, ValueError):
|
||||
# Boxes remain a valid association primitive when mask resolutions
|
||||
# differ across detector backends.
|
||||
pass
|
||||
return overlap
|
||||
|
||||
|
||||
def _hungarian_matches(
|
||||
tracks: list[_TrackState],
|
||||
detections: list[Detection],
|
||||
*,
|
||||
minimum_iou: float,
|
||||
label_aware: bool,
|
||||
motion_gate: float,
|
||||
) -> tuple[
|
||||
list[tuple[int, int, float]],
|
||||
list[int],
|
||||
list[int],
|
||||
]:
|
||||
if not tracks or not detections:
|
||||
return (
|
||||
[],
|
||||
list(range(len(tracks))),
|
||||
list(range(len(detections))),
|
||||
)
|
||||
|
||||
cost = np.full((len(tracks), len(detections)), np.inf, dtype=np.float64)
|
||||
overlaps = np.zeros_like(cost)
|
||||
for track_index, track in enumerate(tracks):
|
||||
for detection_index, detection in enumerate(detections):
|
||||
if not _labels_compatible(track, detection, label_aware=label_aware):
|
||||
continue
|
||||
if track.filter.gating_distance(detection.bbox_xyxy) > motion_gate:
|
||||
continue
|
||||
overlap = _overlap(track, detection)
|
||||
if overlap < minimum_iou:
|
||||
continue
|
||||
overlaps[track_index, detection_index] = overlap
|
||||
cost[track_index, detection_index] = 1.0 - overlap
|
||||
|
||||
finite = np.isfinite(cost)
|
||||
if not finite.any():
|
||||
return (
|
||||
[],
|
||||
list(range(len(tracks))),
|
||||
list(range(len(detections))),
|
||||
)
|
||||
safe_cost = np.where(finite, cost, 1.0e6)
|
||||
row_indices, column_indices = linear_sum_assignment(safe_cost)
|
||||
matches = sorted(
|
||||
(
|
||||
(int(row), int(column), float(overlaps[row, column]))
|
||||
for row, column in zip(row_indices, column_indices)
|
||||
if finite[row, column]
|
||||
),
|
||||
key=lambda item: (tracks[item[0]].track_id, item[1]),
|
||||
)
|
||||
matched_tracks = {track_index for track_index, _index, _score in matches}
|
||||
matched_detections = {
|
||||
detection_index for _index, detection_index, _score in matches
|
||||
}
|
||||
return (
|
||||
matches,
|
||||
[index for index in range(len(tracks)) if index not in matched_tracks],
|
||||
[index for index in range(len(detections)) if index not in matched_detections],
|
||||
)
|
||||
|
||||
|
||||
class VLMByteTracker:
|
||||
"""ByteTrack-style high/low confidence association over a whole sequence."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
high_threshold: float = 0.6,
|
||||
low_threshold: float = 0.1,
|
||||
match_iou_threshold: float = 0.3,
|
||||
low_match_iou_threshold: float = 0.2,
|
||||
max_age_seconds: float = 1.0,
|
||||
min_hits: int = 2,
|
||||
label_aware: bool = True,
|
||||
emit_predictions: bool = True,
|
||||
motion_gate: float = _CHI_SQUARE_FOUR_DOF_99,
|
||||
fps_fallback: float = 30.0,
|
||||
):
|
||||
values = (
|
||||
high_threshold,
|
||||
low_threshold,
|
||||
match_iou_threshold,
|
||||
low_match_iou_threshold,
|
||||
)
|
||||
if any(not 0.0 <= float(value) <= 1.0 for value in values):
|
||||
raise ValueError("Thresholds must be between 0 and 1.")
|
||||
if low_threshold > high_threshold:
|
||||
raise ValueError(
|
||||
"low_threshold must be less than or equal to high_threshold."
|
||||
)
|
||||
if not math.isfinite(float(max_age_seconds)) or max_age_seconds < 0:
|
||||
raise ValueError("max_age_seconds must be finite and non-negative.")
|
||||
if not isinstance(min_hits, int) or min_hits < 1:
|
||||
raise ValueError("min_hits must be a positive integer.")
|
||||
if not math.isfinite(float(motion_gate)) or motion_gate <= 0:
|
||||
raise ValueError("motion_gate must be finite and positive.")
|
||||
if not math.isfinite(float(fps_fallback)) or fps_fallback <= 0:
|
||||
raise ValueError("fps_fallback must be finite and positive.")
|
||||
self.high_threshold = float(high_threshold)
|
||||
self.low_threshold = float(low_threshold)
|
||||
self.match_iou_threshold = float(match_iou_threshold)
|
||||
self.low_match_iou_threshold = float(low_match_iou_threshold)
|
||||
self.max_age_seconds = float(max_age_seconds)
|
||||
self.min_hits = min_hits
|
||||
self.label_aware = bool(label_aware)
|
||||
self.emit_predictions = bool(emit_predictions)
|
||||
self.motion_gate = float(motion_gate)
|
||||
self.fps_fallback = float(fps_fallback)
|
||||
self._tracks: list[_TrackState] = []
|
||||
self._next_track_id = 0
|
||||
|
||||
def _timestamp(
|
||||
self,
|
||||
frame_index: int,
|
||||
frame_timestamp: float | None,
|
||||
fps: float,
|
||||
) -> float:
|
||||
expected = frame_index / fps
|
||||
if frame_timestamp is None or (frame_index > 0 and frame_timestamp <= 0.0):
|
||||
return expected
|
||||
return max(float(frame_timestamp), expected)
|
||||
|
||||
def _spawn(
|
||||
self,
|
||||
detection: Detection,
|
||||
*,
|
||||
timestamp: float,
|
||||
) -> None:
|
||||
state = "active" if self.min_hits == 1 else "tentative"
|
||||
track_id = self._next_track_id
|
||||
self._next_track_id += 1
|
||||
tracked = _tracked_detection(
|
||||
detection,
|
||||
track_id=track_id,
|
||||
timestamp=timestamp,
|
||||
track_state=state,
|
||||
association_stage="new",
|
||||
association_score=None,
|
||||
)
|
||||
scores = [] if detection.score is None else [detection.score]
|
||||
self._tracks.append(
|
||||
_TrackState(
|
||||
track_id=track_id,
|
||||
filter=_BoxKalmanFilter(detection.bbox_xyxy),
|
||||
detections=[tracked],
|
||||
label=detection.label,
|
||||
text=detection.text,
|
||||
state=state,
|
||||
hits=1,
|
||||
first_frame=detection.frame_index,
|
||||
last_observed_frame=detection.frame_index,
|
||||
last_observed_timestamp=timestamp,
|
||||
last_timestamp=timestamp,
|
||||
last_mask=detection.mask,
|
||||
observed_scores=scores,
|
||||
)
|
||||
)
|
||||
|
||||
def _update_track(
|
||||
self,
|
||||
track: _TrackState,
|
||||
detection: Detection,
|
||||
*,
|
||||
timestamp: float,
|
||||
stage: str,
|
||||
association_score: float,
|
||||
) -> None:
|
||||
track.filter.update(detection.bbox_xyxy)
|
||||
track.hits += 1
|
||||
track.misses = 0
|
||||
track.state = "active" if track.hits >= self.min_hits else "tentative"
|
||||
if track.label is None:
|
||||
track.label = detection.label
|
||||
if track.text is None:
|
||||
track.text = detection.text
|
||||
track.last_observed_frame = detection.frame_index
|
||||
track.last_observed_timestamp = timestamp
|
||||
track.last_timestamp = timestamp
|
||||
track.last_mask = detection.mask
|
||||
if detection.score is not None:
|
||||
track.observed_scores.append(detection.score)
|
||||
track.detections.append(
|
||||
_tracked_detection(
|
||||
detection,
|
||||
track_id=track.track_id,
|
||||
timestamp=timestamp,
|
||||
track_state=track.state,
|
||||
association_stage=stage,
|
||||
association_score=association_score,
|
||||
)
|
||||
)
|
||||
|
||||
def _mark_missed(
|
||||
self,
|
||||
track: _TrackState,
|
||||
*,
|
||||
frame_index: int,
|
||||
timestamp: float,
|
||||
width: int,
|
||||
height: int,
|
||||
) -> None:
|
||||
track.misses += 1
|
||||
elapsed = max(0.0, timestamp - track.last_observed_timestamp)
|
||||
if track.state == "tentative" or elapsed > self.max_age_seconds:
|
||||
track.state = "removed"
|
||||
track.removed_frame = frame_index
|
||||
return
|
||||
track.state = "lost"
|
||||
if not self.emit_predictions:
|
||||
return
|
||||
box = clip_box(track.predicted_box, width, height)
|
||||
if box[2] <= box[0] or box[3] <= box[1]:
|
||||
return
|
||||
track.detections.append(
|
||||
Detection(
|
||||
bbox_xyxy=box,
|
||||
label=track.label,
|
||||
text=track.text,
|
||||
score=None,
|
||||
frame_index=frame_index,
|
||||
timestamp=timestamp,
|
||||
track_id=track.track_id,
|
||||
source=TRACKER_SOURCE,
|
||||
metadata={
|
||||
"observation": "predicted",
|
||||
"track_state": "lost",
|
||||
"association_stage": "unmatched",
|
||||
"association_score": None,
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
def _predict(
|
||||
self,
|
||||
*,
|
||||
timestamp: float,
|
||||
) -> list[_TrackState]:
|
||||
candidates = [track for track in self._tracks if track.state != "removed"]
|
||||
for track in candidates:
|
||||
delta = max(timestamp - track.last_timestamp, 1.0e-6)
|
||||
track.filter.predict(delta)
|
||||
track.last_timestamp = timestamp
|
||||
return candidates
|
||||
|
||||
def _process_frame(
|
||||
self,
|
||||
detections: list[Detection],
|
||||
*,
|
||||
frame_index: int,
|
||||
timestamp: float,
|
||||
width: int,
|
||||
height: int,
|
||||
) -> None:
|
||||
candidates = self._predict(timestamp=timestamp)
|
||||
high = [
|
||||
detection
|
||||
for detection in detections
|
||||
if _score_or_one(detection) >= self.high_threshold
|
||||
]
|
||||
low = [
|
||||
detection
|
||||
for detection in detections
|
||||
if self.low_threshold <= _score_or_one(detection) < self.high_threshold
|
||||
]
|
||||
|
||||
high_matches, unmatched_candidate_indices, unmatched_high_indices = (
|
||||
_hungarian_matches(
|
||||
candidates,
|
||||
high,
|
||||
minimum_iou=self.match_iou_threshold,
|
||||
label_aware=self.label_aware,
|
||||
motion_gate=self.motion_gate,
|
||||
)
|
||||
)
|
||||
matched_track_ids = set()
|
||||
for track_index, detection_index, overlap in high_matches:
|
||||
track = candidates[track_index]
|
||||
self._update_track(
|
||||
track,
|
||||
high[detection_index],
|
||||
timestamp=timestamp,
|
||||
stage="high",
|
||||
association_score=overlap,
|
||||
)
|
||||
matched_track_ids.add(track.track_id)
|
||||
|
||||
low_candidates = [
|
||||
candidates[index]
|
||||
for index in unmatched_candidate_indices
|
||||
if candidates[index].state in {"active", "lost"}
|
||||
]
|
||||
low_matches, _unmatched_low_track_indices, _unmatched_low_indices = (
|
||||
_hungarian_matches(
|
||||
low_candidates,
|
||||
low,
|
||||
minimum_iou=self.low_match_iou_threshold,
|
||||
label_aware=self.label_aware,
|
||||
motion_gate=self.motion_gate,
|
||||
)
|
||||
)
|
||||
for track_index, detection_index, overlap in low_matches:
|
||||
track = low_candidates[track_index]
|
||||
self._update_track(
|
||||
track,
|
||||
low[detection_index],
|
||||
timestamp=timestamp,
|
||||
stage="low",
|
||||
association_score=overlap,
|
||||
)
|
||||
matched_track_ids.add(track.track_id)
|
||||
|
||||
for track in candidates:
|
||||
if track.track_id not in matched_track_ids:
|
||||
self._mark_missed(
|
||||
track,
|
||||
frame_index=frame_index,
|
||||
timestamp=timestamp,
|
||||
width=width,
|
||||
height=height,
|
||||
)
|
||||
|
||||
for detection_index in unmatched_high_indices:
|
||||
self._spawn(high[detection_index], timestamp=timestamp)
|
||||
|
||||
def track(self, sequence: DetectionSequence) -> TrackSequence:
|
||||
if not isinstance(sequence, DetectionSequence):
|
||||
raise TypeError("sequence must be a DetectionSequence.")
|
||||
self._tracks = []
|
||||
self._next_track_id = 0
|
||||
fps = sequence.fps or self.fps_fallback
|
||||
frames = {frame.frame_index: frame for frame in sequence.frames}
|
||||
for frame_index in range(sequence.frame_count):
|
||||
frame = frames.get(frame_index)
|
||||
timestamp = self._timestamp(
|
||||
frame_index,
|
||||
None if frame is None else frame.timestamp,
|
||||
fps,
|
||||
)
|
||||
detections = (
|
||||
[]
|
||||
if frame is None
|
||||
else [
|
||||
Detection(
|
||||
bbox_xyxy=detection.bbox_xyxy,
|
||||
label=detection.label,
|
||||
text=detection.text,
|
||||
score=detection.score,
|
||||
polygon=detection.polygon,
|
||||
quad=detection.quad,
|
||||
frame_index=frame_index,
|
||||
timestamp=timestamp,
|
||||
track_id=detection.track_id,
|
||||
source=detection.source,
|
||||
metadata=detection.metadata,
|
||||
mask=detection.mask,
|
||||
)
|
||||
for detection in frame.detections
|
||||
]
|
||||
)
|
||||
self._process_frame(
|
||||
detections,
|
||||
frame_index=frame_index,
|
||||
timestamp=timestamp,
|
||||
width=sequence.width,
|
||||
height=sequence.height,
|
||||
)
|
||||
|
||||
tracks = []
|
||||
for track in sorted(self._tracks, key=lambda item: item.track_id):
|
||||
score = (
|
||||
sum(track.observed_scores) / len(track.observed_scores)
|
||||
if track.observed_scores
|
||||
else None
|
||||
)
|
||||
tracks.append(
|
||||
Track(
|
||||
track_id=track.track_id,
|
||||
detections=tuple(track.detections),
|
||||
label=track.label,
|
||||
score=score,
|
||||
source=TRACKER_SOURCE,
|
||||
metadata={
|
||||
"state": track.state,
|
||||
"hits": track.hits,
|
||||
"misses": track.misses,
|
||||
"first_frame": track.first_frame,
|
||||
"last_observed_frame": track.last_observed_frame,
|
||||
"removed_frame": track.removed_frame,
|
||||
},
|
||||
)
|
||||
)
|
||||
metadata = sequence.metadata.to_dict()
|
||||
metadata["tracker"] = {
|
||||
"algorithm": "bytetrack-style-hungarian",
|
||||
"high_threshold": self.high_threshold,
|
||||
"low_threshold": self.low_threshold,
|
||||
"match_iou_threshold": self.match_iou_threshold,
|
||||
"low_match_iou_threshold": self.low_match_iou_threshold,
|
||||
"max_age_seconds": self.max_age_seconds,
|
||||
"min_hits": self.min_hits,
|
||||
"label_aware": self.label_aware,
|
||||
"emit_predictions": self.emit_predictions,
|
||||
}
|
||||
return TrackSequence(
|
||||
width=sequence.width,
|
||||
height=sequence.height,
|
||||
tracks=tuple(tracks),
|
||||
frame_count=sequence.frame_count,
|
||||
fps=sequence.fps,
|
||||
source=TRACKER_SOURCE,
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
|
||||
def associate_detection_sequence(
|
||||
sequence: DetectionSequence,
|
||||
**tracker_options,
|
||||
) -> TrackSequence:
|
||||
"""Convenience function for callers that do not need a reusable tracker."""
|
||||
|
||||
return VLMByteTracker(**tracker_options).track(sequence)
|
||||
|
||||
|
||||
class VLMTrackDetections:
|
||||
"""ComfyUI node wrapper for deterministic tracking-by-detection."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"detections": (VLM_DETECTIONS,),
|
||||
"high_threshold": (
|
||||
"FLOAT",
|
||||
{"default": 0.6, "min": 0.0, "max": 1.0, "step": 0.01},
|
||||
),
|
||||
"low_threshold": (
|
||||
"FLOAT",
|
||||
{"default": 0.1, "min": 0.0, "max": 1.0, "step": 0.01},
|
||||
),
|
||||
"match_iou_threshold": (
|
||||
"FLOAT",
|
||||
{"default": 0.3, "min": 0.0, "max": 1.0, "step": 0.01},
|
||||
),
|
||||
"low_match_iou_threshold": (
|
||||
"FLOAT",
|
||||
{"default": 0.2, "min": 0.0, "max": 1.0, "step": 0.01},
|
||||
),
|
||||
"max_age_seconds": (
|
||||
"FLOAT",
|
||||
{"default": 1.0, "min": 0.0, "max": 60.0, "step": 0.05},
|
||||
),
|
||||
"min_hits": (
|
||||
"INT",
|
||||
{"default": 2, "min": 1, "max": 100},
|
||||
),
|
||||
"label_aware": ("BOOLEAN", {"default": True}),
|
||||
"emit_predictions": ("BOOLEAN", {"default": True}),
|
||||
"fps_fallback": (
|
||||
"FLOAT",
|
||||
{"default": 30.0, "min": 0.01, "max": 1000.0},
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = (VLM_TRACKS,)
|
||||
RETURN_NAMES = ("tracks",)
|
||||
FUNCTION = "track"
|
||||
CATEGORY = "VLM Nodes/Vision/Tracking"
|
||||
|
||||
def track(
|
||||
self,
|
||||
detections,
|
||||
high_threshold,
|
||||
low_threshold,
|
||||
match_iou_threshold,
|
||||
low_match_iou_threshold,
|
||||
max_age_seconds,
|
||||
min_hits,
|
||||
label_aware,
|
||||
emit_predictions,
|
||||
fps_fallback,
|
||||
):
|
||||
tracker = VLMByteTracker(
|
||||
high_threshold=high_threshold,
|
||||
low_threshold=low_threshold,
|
||||
match_iou_threshold=match_iou_threshold,
|
||||
low_match_iou_threshold=low_match_iou_threshold,
|
||||
max_age_seconds=max_age_seconds,
|
||||
min_hits=min_hits,
|
||||
label_aware=label_aware,
|
||||
emit_predictions=emit_predictions,
|
||||
fps_fallback=fps_fallback,
|
||||
)
|
||||
return (tracker.track(detections),)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {"VLMTrackDetections": VLMTrackDetections}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {"VLMTrackDetections": "VLM Track Detections"}
|
||||
|
||||
__all__ = [
|
||||
"NODE_CLASS_MAPPINGS",
|
||||
"NODE_DISPLAY_NAME_MAPPINGS",
|
||||
"TRACKER_SOURCE",
|
||||
"VLMByteTracker",
|
||||
"VLMTrackDetections",
|
||||
"associate_detection_sequence",
|
||||
]
|
||||
+1
-1
@@ -105,7 +105,7 @@ class UformGen2QwenNode(CachedModelNode):
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "uform_gen2_qwen_chat"
|
||||
CATEGORY = "VLM Nodes/UformGen2Qwen"
|
||||
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
|
||||
|
||||
def uform_gen2_qwen_chat(
|
||||
self, image, question, max_new_tokens=512, unload_after=False
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
+77
-3
@@ -1,22 +1,28 @@
|
||||
[project]
|
||||
name = "comfyui_vlm_nodes"
|
||||
version = "2.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 = { file = "LICENSE" }
|
||||
license = "Apache-2.0"
|
||||
license-files = ["LICENSE"]
|
||||
dependencies = [
|
||||
"accelerate>=1.1,<2",
|
||||
"bitsandbytes>=0.50,<1; (sys_platform == 'linux' and platform_machine == 'x86_64') or (sys_platform == 'linux' and platform_machine == 'aarch64') or (sys_platform == 'win32' and platform_machine == 'AMD64') or (sys_platform == 'win32' and platform_machine == 'ARM64') or (sys_platform == 'darwin' and platform_machine == 'arm64')",
|
||||
"diffusers>=0.34,<1",
|
||||
"einops>=0.8,<1",
|
||||
"huggingface-hub>=1.5,<2",
|
||||
"openai>=1.30,<3",
|
||||
"httpx>=0.27,<1",
|
||||
"jsonschema>=4.22,<5",
|
||||
"num2words>=0.5.14,<1",
|
||||
"openai>=2,<3",
|
||||
"pydantic>=2.7,<3",
|
||||
"qwen-vl-utils>=0.0.14",
|
||||
"safetensors>=0.4.3",
|
||||
"scipy>=1.10,<2",
|
||||
"soundfile>=0.12",
|
||||
"symusic>=0.5",
|
||||
"svgelements>=1.9.6,<2",
|
||||
"transformers>=5.4,<6",
|
||||
]
|
||||
classifiers = [
|
||||
@@ -41,12 +47,80 @@ 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"
|
||||
Icon = ""
|
||||
|
||||
[tool.setuptools]
|
||||
packages = [
|
||||
"comfyui_vlm_nodes",
|
||||
"comfyui_vlm_nodes.examples",
|
||||
"comfyui_vlm_nodes.examples.robotics",
|
||||
"comfyui_vlm_nodes.examples.vision",
|
||||
"comfyui_vlm_nodes.nodes",
|
||||
"comfyui_vlm_nodes.nodes.joytagger",
|
||||
"comfyui_vlm_nodes.web",
|
||||
"comfyui_vlm_nodes.web.js",
|
||||
]
|
||||
include-package-data = true
|
||||
|
||||
[tool.setuptools.package-dir]
|
||||
comfyui_vlm_nodes = "."
|
||||
|
||||
[tool.setuptools.package-data]
|
||||
comfyui_vlm_nodes = [
|
||||
"*.json",
|
||||
"SECURITY.md",
|
||||
"examples/*.json",
|
||||
"examples/robotics/*.py",
|
||||
"examples/robotics/*.md",
|
||||
"examples/robotics/*.json",
|
||||
"examples/vision/*.json",
|
||||
"requirements*.txt",
|
||||
]
|
||||
"comfyui_vlm_nodes.web.js" = ["*.js"]
|
||||
|
||||
@@ -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
|
||||
@@ -0,0 +1,11 @@
|
||||
# Install this file only into the isolated Moondream sidecar environment.
|
||||
# Do not install it into ComfyUI's main environment: moondream 1.3 pins
|
||||
# Pillow <11 while current ComfyUI uses a newer Pillow release.
|
||||
moondream==1.3.0
|
||||
# moondream 1.3.0 expects this exact runtime API. 0.4.7+ renamed the
|
||||
# prefix-mask kernel and is not source-compatible with kestrel 0.4.2.
|
||||
kestrel-kernels==0.4.6
|
||||
# Kestrel's CUDA 12 AOT kernels call cudaLibraryLoadData. PyTorch's cu126
|
||||
# runtime (12.6.77) does not export it; 12.9.79 does and remains within the
|
||||
# CUDA 12 ABI. Keep this inside the isolated Photon environment only.
|
||||
nvidia-cuda-runtime-cu12==12.9.79; (sys_platform == "linux" and platform_machine == "x86_64") or (sys_platform == "win32" and platform_machine == "AMD64")
|
||||
@@ -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
|
||||
+6
-1
@@ -7,10 +7,15 @@ bitsandbytes>=0.50,<1; (sys_platform == "linux" and platform_machine == "x86_64"
|
||||
diffusers>=0.34,<1
|
||||
einops>=0.8,<1
|
||||
huggingface-hub>=1.5,<2
|
||||
openai>=1.30,<3
|
||||
httpx>=0.27,<1
|
||||
jsonschema>=4.22,<5
|
||||
num2words>=0.5.14,<1
|
||||
openai>=2,<3
|
||||
pydantic>=2.7,<3
|
||||
qwen-vl-utils>=0.0.14
|
||||
safetensors>=0.4.3
|
||||
scipy>=1.10,<2
|
||||
soundfile>=0.12
|
||||
symusic>=0.5
|
||||
svgelements>=1.9.6,<2
|
||||
transformers>=5.4,<6
|
||||
|
||||
@@ -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)
|
||||
+16
-1
@@ -2,10 +2,10 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib.util
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
REPOSITORY = Path(__file__).resolve().parents[1]
|
||||
for candidate in (
|
||||
REPOSITORY.parent,
|
||||
@@ -14,3 +14,18 @@ for candidate in (
|
||||
):
|
||||
if candidate.exists():
|
||||
sys.path.insert(0, str(candidate))
|
||||
|
||||
# Git worktrees are often intentionally named after a feature branch rather
|
||||
# than the import package. Load this checkout explicitly so tests can never
|
||||
# pass by silently importing a sibling clone with the canonical directory name.
|
||||
if REPOSITORY.name != "ComfyUI_VLM_nodes":
|
||||
specification = importlib.util.spec_from_file_location(
|
||||
"ComfyUI_VLM_nodes",
|
||||
REPOSITORY / "__init__.py",
|
||||
submodule_search_locations=[str(REPOSITORY)],
|
||||
)
|
||||
if specification is None or specification.loader is None:
|
||||
raise RuntimeError(f"Could not load package from {REPOSITORY}.")
|
||||
package = importlib.util.module_from_spec(specification)
|
||||
sys.modules["ComfyUI_VLM_nodes"] = package
|
||||
specification.loader.exec_module(package)
|
||||
|
||||
@@ -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,176 @@
|
||||
"""Run adaptive temporal reasoning on a real local video and real VLM.
|
||||
|
||||
Example:
|
||||
python tests/manual_video_intelligence_smoke.py \
|
||||
/mnt/d/002.mp4 \
|
||||
--model "Qwen 3 VL 2B Instruct" \
|
||||
--output /mnt/d/comfyui-repair/video-intelligence-audit/result.json
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import importlib.util
|
||||
import json
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import av
|
||||
import torch
|
||||
|
||||
REPOSITORY = Path(__file__).resolve().parents[1]
|
||||
if REPOSITORY.name != "ComfyUI_VLM_nodes":
|
||||
specification = importlib.util.spec_from_file_location(
|
||||
"ComfyUI_VLM_nodes",
|
||||
REPOSITORY / "__init__.py",
|
||||
submodule_search_locations=[str(REPOSITORY)],
|
||||
)
|
||||
if specification is None or specification.loader is None:
|
||||
raise RuntimeError(f"Could not load package from {REPOSITORY}.")
|
||||
package = importlib.util.module_from_spec(specification)
|
||||
sys.modules["ComfyUI_VLM_nodes"] = package
|
||||
specification.loader.exec_module(package)
|
||||
|
||||
from ComfyUI_VLM_nodes.nodes.modern_vlm import ModernVLMPredictor
|
||||
from ComfyUI_VLM_nodes.nodes.video_intelligence import (
|
||||
build_video_reasoning_prompt,
|
||||
parse_video_reasoning_output,
|
||||
resize_video_for_analysis,
|
||||
sample_video_frames,
|
||||
)
|
||||
|
||||
|
||||
def load_video(path: Path) -> tuple[torch.Tensor, float]:
|
||||
container = av.open(str(path))
|
||||
try:
|
||||
stream = container.streams.video[0]
|
||||
rate = stream.average_rate or stream.guessed_rate
|
||||
if rate is None:
|
||||
raise RuntimeError("The video does not report a frame rate.")
|
||||
frames = [
|
||||
torch.from_numpy(frame.to_ndarray(format="rgb24")).to(torch.float32)
|
||||
/ 255.0
|
||||
for frame in container.decode(stream)
|
||||
]
|
||||
finally:
|
||||
container.close()
|
||||
if not frames:
|
||||
raise RuntimeError("The video contains no decodable frames.")
|
||||
return torch.stack(frames), float(rate)
|
||||
|
||||
|
||||
def main() -> int:
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("video", type=Path)
|
||||
parser.add_argument(
|
||||
"--model",
|
||||
default="Qwen 3 VL 2B Instruct",
|
||||
)
|
||||
parser.add_argument("--max-frames", type=int, default=12)
|
||||
parser.add_argument("--analysis-max-side", type=int, default=448)
|
||||
parser.add_argument("--max-new-tokens", type=int, default=512)
|
||||
parser.add_argument("--output", type=Path)
|
||||
args = parser.parse_args()
|
||||
|
||||
frames, fps = load_video(args.video)
|
||||
sampled, selection, diagnostics = sample_video_frames(
|
||||
frames,
|
||||
fps=fps,
|
||||
max_frames=args.max_frames,
|
||||
strategy="Hybrid: scene + motion + tracks",
|
||||
minimum_gap_seconds=0.2,
|
||||
)
|
||||
prompt = build_video_reasoning_prompt(
|
||||
selection,
|
||||
task="Detailed temporal summary",
|
||||
question="What happens, and how do the people behave over time?",
|
||||
max_events=12,
|
||||
)
|
||||
analysis_frames = resize_video_for_analysis(
|
||||
sampled,
|
||||
max_side=args.analysis_max_side,
|
||||
)
|
||||
predictor = ModernVLMPredictor(
|
||||
args.model,
|
||||
"",
|
||||
"ComfyUI managed (BF16)",
|
||||
"Auto (SDPA)",
|
||||
)
|
||||
started = time.perf_counter()
|
||||
try:
|
||||
raw = predictor.generate(
|
||||
images=None,
|
||||
prompt=prompt,
|
||||
system_prompt=(
|
||||
"You are a precise temporal video analyst. Return one JSON "
|
||||
"object that obeys the supplied schema."
|
||||
),
|
||||
max_new_tokens=args.max_new_tokens,
|
||||
temperature=0.0,
|
||||
top_p=1.0,
|
||||
video_frames=analysis_frames,
|
||||
fps=fps,
|
||||
video_selection=selection,
|
||||
)
|
||||
finally:
|
||||
predictor.close()
|
||||
reasoning_seconds = time.perf_counter() - started
|
||||
result = {
|
||||
"video": str(args.video),
|
||||
"model": args.model,
|
||||
"source_shape": list(frames.shape),
|
||||
"fps": fps,
|
||||
"selection": selection.to_dict(),
|
||||
"sampling": diagnostics,
|
||||
"analysis_shape": list(analysis_frames.shape),
|
||||
"reasoning_seconds": reasoning_seconds,
|
||||
"raw_response": raw,
|
||||
"cuda_peak_gib": (
|
||||
torch.cuda.max_memory_allocated() / 2**30
|
||||
if torch.cuda.is_available()
|
||||
else 0.0
|
||||
),
|
||||
}
|
||||
try:
|
||||
summary, events, normalized = parse_video_reasoning_output(raw, selection)
|
||||
except (TypeError, ValueError) as exc:
|
||||
result["structured_output_valid"] = False
|
||||
result["structured_output_error"] = str(exc)
|
||||
if args.output:
|
||||
args.output.parent.mkdir(parents=True, exist_ok=True)
|
||||
args.output.write_text(
|
||||
json.dumps(
|
||||
result,
|
||||
ensure_ascii=False,
|
||||
allow_nan=False,
|
||||
indent=2,
|
||||
sort_keys=True,
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
raise
|
||||
result.update(
|
||||
{
|
||||
"structured_output_valid": True,
|
||||
"summary": summary,
|
||||
"events": events.to_dict(),
|
||||
"normalized_response": json.loads(normalized),
|
||||
}
|
||||
)
|
||||
encoded = json.dumps(
|
||||
result,
|
||||
ensure_ascii=False,
|
||||
allow_nan=False,
|
||||
indent=2,
|
||||
sort_keys=True,
|
||||
)
|
||||
if args.output:
|
||||
args.output.parent.mkdir(parents=True, exist_ok=True)
|
||||
args.output.write_text(encoded, encoding="utf-8")
|
||||
print(encoded)
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -0,0 +1,111 @@
|
||||
import json
|
||||
import threading
|
||||
import time
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from ComfyUI_VLM_nodes.nodes.acceleration import (
|
||||
VLMImagePixelBudget,
|
||||
VLMPerformanceProfile,
|
||||
optimize_image_pixels,
|
||||
)
|
||||
from ComfyUI_VLM_nodes.nodes.runtime import (
|
||||
CachedModelNode,
|
||||
tensor_batch_to_pil,
|
||||
tensor_to_pil,
|
||||
)
|
||||
|
||||
|
||||
def test_batch_conversion_matches_single_frame_contract():
|
||||
images = torch.tensor(
|
||||
[
|
||||
[
|
||||
[[float("nan"), 0.5, 2.0], [-1.0, 0.25, 1.0]],
|
||||
[[0.0, 0.0, 0.0], [1.0, 1.0, 1.0]],
|
||||
],
|
||||
[
|
||||
[[255.0, 128.0, 0.0], [0.0, 64.0, 255.0]],
|
||||
[[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]],
|
||||
],
|
||||
]
|
||||
)
|
||||
batch = tensor_batch_to_pil(images)
|
||||
assert len(batch) == 2
|
||||
for index, converted in enumerate(batch):
|
||||
assert converted.mode == "RGB"
|
||||
assert converted.size == (2, 2)
|
||||
assert converted.tobytes() == tensor_to_pil(images, index).tobytes()
|
||||
|
||||
with pytest.raises(IndexError, match="only has batch index 0"):
|
||||
tensor_to_pil(images[0], 1)
|
||||
|
||||
|
||||
def test_pixel_budget_preserves_aspect_and_patch_multiple():
|
||||
images = torch.rand((3, 1080, 1920, 3), dtype=torch.float32)
|
||||
output, report = optimize_image_pixels(
|
||||
images,
|
||||
max_megapixels=0.5,
|
||||
max_edge=1024,
|
||||
multiple=14,
|
||||
resize_quality="Fast (area)",
|
||||
)
|
||||
assert output.ndim == 4
|
||||
assert output.shape[0] == 3
|
||||
assert output.shape[1] % 14 == 0
|
||||
assert output.shape[2] % 14 == 0
|
||||
assert output.shape[1] * output.shape[2] <= 500_000
|
||||
assert output.shape[2] <= 1024
|
||||
assert report["visual_work_reduction"] > 4
|
||||
assert output.shape[2] / output.shape[1] == pytest.approx(16 / 9, rel=0.03)
|
||||
|
||||
|
||||
def test_pixel_budget_never_upscales():
|
||||
image = torch.rand((240, 320, 3), dtype=torch.float32)
|
||||
output, report = optimize_image_pixels(
|
||||
image,
|
||||
max_megapixels=2.0,
|
||||
max_edge=2048,
|
||||
multiple=1,
|
||||
resize_quality="Quality (bicubic)",
|
||||
)
|
||||
assert output is image
|
||||
assert report["resized"] is False
|
||||
|
||||
|
||||
def test_performance_nodes_return_standard_comfy_values():
|
||||
profile = VLMPerformanceProfile().profile("Live / robotics")
|
||||
assert profile[:5] == (24, 0.5, 896, 8, False)
|
||||
assert json.loads(profile[5])["profile"] == "Live / robotics"
|
||||
|
||||
optimized = VLMImagePixelBudget().optimize(
|
||||
torch.rand((1, 1000, 1600, 3)),
|
||||
0.5,
|
||||
1024,
|
||||
"14",
|
||||
"Fast (area)",
|
||||
)
|
||||
assert optimized[1] % 14 == 0
|
||||
assert optimized[2] % 14 == 0
|
||||
|
||||
|
||||
def test_cached_model_node_prevents_duplicate_concurrent_loads():
|
||||
node = CachedModelNode()
|
||||
factory_calls = []
|
||||
handles = []
|
||||
|
||||
def factory():
|
||||
factory_calls.append(1)
|
||||
time.sleep(0.02)
|
||||
return object()
|
||||
|
||||
def load():
|
||||
handles.append(node.get_or_create_model("same-model", factory))
|
||||
|
||||
threads = [threading.Thread(target=load) for _ in range(8)]
|
||||
for thread in threads:
|
||||
thread.start()
|
||||
for thread in threads:
|
||||
thread.join()
|
||||
|
||||
assert len(factory_calls) == 1
|
||||
assert len({id(handle) for handle in handles}) == 1
|
||||
@@ -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
|
||||
@@ -0,0 +1,250 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from ComfyUI_VLM_nodes.nodes import florence2
|
||||
from PIL import Image
|
||||
|
||||
EXPECTED_TASKS = {
|
||||
"Caption": ("<CAPTION>", "none"),
|
||||
"Detailed caption": ("<DETAILED_CAPTION>", "none"),
|
||||
"More detailed caption": ("<MORE_DETAILED_CAPTION>", "none"),
|
||||
"OCR": ("<OCR>", "none"),
|
||||
"OCR with regions": ("<OCR_WITH_REGION>", "none"),
|
||||
"Object detection": ("<OD>", "none"),
|
||||
"Dense region caption": ("<DENSE_REGION_CAPTION>", "none"),
|
||||
"Caption to phrase grounding": ("<CAPTION_TO_PHRASE_GROUNDING>", "text"),
|
||||
"Referring expression segmentation": (
|
||||
"<REFERRING_EXPRESSION_SEGMENTATION>",
|
||||
"text",
|
||||
),
|
||||
"Region to segmentation": ("<REGION_TO_SEGMENTATION>", "region"),
|
||||
"Open vocabulary detection": ("<OPEN_VOCABULARY_DETECTION>", "text"),
|
||||
"Region to category": ("<REGION_TO_CATEGORY>", "region"),
|
||||
"Region to description": ("<REGION_TO_DESCRIPTION>", "region"),
|
||||
"Region to OCR": ("<REGION_TO_OCR>", "region"),
|
||||
"Region proposals": ("<REGION_PROPOSAL>", "none"),
|
||||
}
|
||||
|
||||
|
||||
def test_registry_covers_all_official_transformers_tasks():
|
||||
assert len(florence2.TASKS) == 15
|
||||
assert {
|
||||
name: (spec.token, spec.input_kind) for name, spec in florence2.TASKS.items()
|
||||
} == EXPECTED_TASKS
|
||||
assert {spec.output_kind for spec in florence2.TASKS.values()} == {
|
||||
"text",
|
||||
"boxes",
|
||||
"quad_boxes",
|
||||
"polygons",
|
||||
"mixed",
|
||||
}
|
||||
|
||||
|
||||
def test_node_contract_preserves_outputs_and_adds_core_region_input():
|
||||
schema = florence2.Florence2.INPUT_TYPES()
|
||||
assert florence2.NODE_CLASS_MAPPINGS["Florence2"] is florence2.Florence2
|
||||
assert florence2.Florence2.RETURN_TYPES[:4] == (
|
||||
"STRING",
|
||||
"STRING",
|
||||
"MASK",
|
||||
"IMAGE",
|
||||
)
|
||||
assert florence2.Florence2.RETURN_NAMES[:4] == (
|
||||
"text",
|
||||
"structured_json",
|
||||
"mask",
|
||||
"visualization",
|
||||
)
|
||||
assert schema["optional"]["region"][0] == "BOUNDING_BOX"
|
||||
assert "forceInput" not in repr(schema["optional"]["region"])
|
||||
|
||||
|
||||
def test_region_encoding_uses_core_xywh_and_florence_location_bins():
|
||||
region = {"x": 10, "y": 20, "width": 40, "height": 100}
|
||||
assert florence2._encode_region(region, (100, 200)) == (
|
||||
"<loc_100><loc_100><loc_500><loc_600>"
|
||||
)
|
||||
|
||||
clamped = {"x": -10, "y": -20, "width": 200, "height": 300}
|
||||
assert florence2._encode_region(clamped, (100, 200)) == (
|
||||
"<loc_0><loc_0><loc_999><loc_999>"
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="greater than zero"):
|
||||
florence2._encode_region(
|
||||
{"x": 0, "y": 0, "width": 0, "height": 10},
|
||||
(100, 100),
|
||||
)
|
||||
with pytest.raises(ValueError, match="does not overlap"):
|
||||
florence2._encode_region(
|
||||
{"x": 200, "y": 200, "width": 10, "height": 10},
|
||||
(100, 100),
|
||||
)
|
||||
|
||||
|
||||
def test_task_inputs_are_validated_before_inference():
|
||||
image_size = (100, 100)
|
||||
region = {"x": 10, "y": 10, "width": 20, "height": 20}
|
||||
|
||||
assert florence2._task_extra_input("Caption", "", None, image_size) == ""
|
||||
with pytest.raises(ValueError, match="does not accept text"):
|
||||
florence2._task_extra_input("Caption", "unexpected", None, image_size)
|
||||
with pytest.raises(ValueError, match="requires text"):
|
||||
florence2._task_extra_input("Open vocabulary detection", "", None, image_size)
|
||||
assert (
|
||||
florence2._task_extra_input(
|
||||
"Open vocabulary detection", "red car", None, image_size
|
||||
)
|
||||
== "red car"
|
||||
)
|
||||
with pytest.raises(ValueError, match="requires a connected BOUNDING_BOX"):
|
||||
florence2._task_extra_input("Region to OCR", "", None, image_size)
|
||||
assert (
|
||||
florence2._task_extra_input("Region to OCR", "", region, image_size)
|
||||
== "<loc_100><loc_100><loc_300><loc_300>"
|
||||
)
|
||||
with pytest.raises(ValueError, match="does not accept text"):
|
||||
florence2._task_extra_input("Region to OCR", "also text", region, image_size)
|
||||
|
||||
|
||||
def test_region_selection_accepts_core_and_batched_detector_shapes():
|
||||
first = {"x": 1, "y": 2, "width": 3, "height": 4}
|
||||
second = {"x": 5, "y": 6, "width": 7, "height": 8}
|
||||
assert florence2._select_region(first, 0, 2) is first
|
||||
assert florence2._select_region([first, second], 1, 2) is second
|
||||
assert florence2._select_region([[first], [second]], 0, 2) is first
|
||||
|
||||
with pytest.raises(ValueError, match="exactly one"):
|
||||
florence2._select_region([[first, second]], 0, 1)
|
||||
|
||||
|
||||
def test_visualization_is_deterministic_and_masks_every_spatial_shape():
|
||||
image = Image.new("RGB", (48, 36), "black")
|
||||
parsed = {
|
||||
"<OPEN_VOCABULARY_DETECTION>": {
|
||||
"bboxes": [[1, 1, 10, 10]],
|
||||
"bboxes_labels": ["box"],
|
||||
"quad_boxes": [[14, 1, 22, 1, 22, 10, 14, 10]],
|
||||
"labels": ["ocr"],
|
||||
"polygons": [[[26, 1, 40, 1, 40, 12, 26, 12]]],
|
||||
"polygons_labels": ["polygon"],
|
||||
}
|
||||
}
|
||||
|
||||
mask_a, visual_a = florence2._visualize(image, parsed)
|
||||
mask_b, visual_b = florence2._visualize(image, parsed)
|
||||
mask = np.asarray(mask_a)
|
||||
|
||||
assert mask[5, 5] == 255
|
||||
assert mask[5, 18] == 255
|
||||
assert mask[5, 30] == 255
|
||||
assert mask_a.tobytes() == mask_b.tobytes()
|
||||
assert visual_a.tobytes() == visual_b.tobytes()
|
||||
|
||||
|
||||
def test_predictor_generation_is_deterministic_without_downloads():
|
||||
calls = {}
|
||||
|
||||
class FakeProcessor:
|
||||
def __call__(self, text, images, return_tensors):
|
||||
calls["prompt"] = text
|
||||
assert images.size == (8, 8)
|
||||
assert return_tensors == "pt"
|
||||
return {
|
||||
"input_ids": torch.tensor([[1]], dtype=torch.long),
|
||||
"pixel_values": torch.zeros((1, 3, 8, 8)),
|
||||
}
|
||||
|
||||
def batch_decode(self, generated, skip_special_tokens):
|
||||
assert skip_special_tokens is False
|
||||
return ["<s>answer</s>"]
|
||||
|
||||
def post_process_generation(self, raw, task, image_size):
|
||||
return {task: "answer"}
|
||||
|
||||
class FakeModel(torch.nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.anchor = torch.nn.Parameter(torch.zeros(()))
|
||||
|
||||
def generate(self, **kwargs):
|
||||
calls["generation"] = kwargs
|
||||
return torch.tensor([[2]], dtype=torch.long)
|
||||
|
||||
class FakeHandle:
|
||||
def __init__(self):
|
||||
self.model = FakeModel()
|
||||
|
||||
def ensure_loaded(self):
|
||||
return self.model
|
||||
|
||||
predictor = object.__new__(florence2.FlorencePredictor)
|
||||
predictor.dtype = torch.float32
|
||||
predictor.processor = FakeProcessor()
|
||||
predictor.handle = FakeHandle()
|
||||
|
||||
raw, parsed = predictor.run(
|
||||
Image.new("RGB", (8, 8)),
|
||||
"<CAPTION>",
|
||||
"",
|
||||
32,
|
||||
3,
|
||||
)
|
||||
assert raw == "<s>answer</s>"
|
||||
assert parsed == {"<CAPTION>": "answer"}
|
||||
assert calls["prompt"] == "<CAPTION>"
|
||||
assert calls["generation"]["do_sample"] is False
|
||||
assert calls["generation"]["num_beams"] == 3
|
||||
assert calls["generation"]["early_stopping"] is True
|
||||
|
||||
|
||||
def test_node_cleans_text_preserves_structured_data_and_unloads_target():
|
||||
parsed = {
|
||||
"<OD>": {
|
||||
"bboxes": [[2, 2, 12, 12]],
|
||||
"labels": ["person"],
|
||||
"quad_boxes": [[14, 2, 22, 2, 22, 12, 14, 12]],
|
||||
"polygons": [[[24, 2, 30, 2, 30, 12, 24, 12]]],
|
||||
}
|
||||
}
|
||||
calls = []
|
||||
|
||||
class FakePredictor:
|
||||
def run(
|
||||
self,
|
||||
image,
|
||||
task_token,
|
||||
extra_input,
|
||||
max_new_tokens,
|
||||
beams,
|
||||
):
|
||||
calls.append((image.size, task_token, extra_input, max_new_tokens, beams))
|
||||
return "<s>person<loc_1><loc_2></s><pad>", parsed
|
||||
|
||||
node = florence2.Florence2()
|
||||
node.get_or_create_model = lambda key, factory: FakePredictor()
|
||||
unloads = []
|
||||
node.maybe_clear_model = unloads.append
|
||||
|
||||
output = node.run(
|
||||
torch.zeros((1, 32, 32, 3)),
|
||||
"Object detection",
|
||||
"",
|
||||
"Florence-2 base FT (fast)",
|
||||
64,
|
||||
1,
|
||||
unload_after=True,
|
||||
)
|
||||
|
||||
assert len(output) == 4
|
||||
assert output[0] == "person<loc_1><loc_2>"
|
||||
assert json.loads(output[1]) == [parsed]
|
||||
assert output[2].shape == (1, 32, 32)
|
||||
assert output[2].max().item() == 1.0
|
||||
assert output[3].shape == (1, 32, 32, 3)
|
||||
assert calls == [((32, 32), "<OD>", "", 64, 1)]
|
||||
assert unloads == [True]
|
||||
@@ -0,0 +1,166 @@
|
||||
import importlib
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
PACKAGE = Path(__file__).resolve().parents[1].name
|
||||
geometry = importlib.import_module(f"{PACKAGE}.nodes.geometry")
|
||||
vision_types = importlib.import_module(f"{PACKAGE}.nodes.vision_types")
|
||||
|
||||
associate_detections = geometry.associate_detections
|
||||
bbox_from_mask = geometry.bbox_from_mask
|
||||
bbox_iou = geometry.bbox_iou
|
||||
box_area = geometry.box_area
|
||||
box_center = geometry.box_center
|
||||
box_to_mask = geometry.box_to_mask
|
||||
clip_box = geometry.clip_box
|
||||
clip_polygon = geometry.clip_polygon
|
||||
denormalize_box = geometry.denormalize_box
|
||||
detection_to_mask = geometry.detection_to_mask
|
||||
deterministic_color = geometry.deterministic_color
|
||||
expand_box = geometry.expand_box
|
||||
individual_detection_masks = geometry.individual_detection_masks
|
||||
mask_iou = geometry.mask_iou
|
||||
normalize_box = geometry.normalize_box
|
||||
polygon_area = geometry.polygon_area
|
||||
polygon_to_mask = geometry.polygon_to_mask
|
||||
quad_to_mask = geometry.quad_to_mask
|
||||
translate_box = geometry.translate_box
|
||||
union_detection_mask = geometry.union_detection_mask
|
||||
Detection = vision_types.Detection
|
||||
|
||||
|
||||
def test_box_clipping_normalization_area_and_center():
|
||||
assert clip_box((0, 2, 25, 22), 20, 10) == (0, 2, 20, 10)
|
||||
normalized = normalize_box((5, 2, 15, 8), 20, 10)
|
||||
assert normalized == pytest.approx((0.25, 0.2, 0.75, 0.8))
|
||||
assert denormalize_box(normalized, 20, 10) == pytest.approx((5, 2, 15, 8))
|
||||
assert box_area((5, 2, 15, 8)) == 60
|
||||
assert box_center((5, 2, 15, 8)) == (10, 5)
|
||||
with pytest.raises(ValueError, match="Normalized"):
|
||||
denormalize_box((0, 0, 2, 1), 20, 10)
|
||||
with pytest.raises(ValueError, match="x2"):
|
||||
clip_box((2, 0, 1, 1), 20, 10)
|
||||
|
||||
|
||||
def test_polygon_clipping_and_area():
|
||||
polygon = ((-2, -3), (8, 0), (8, 5), (0, 5))
|
||||
assert clip_polygon(polygon, 6, 4) == (
|
||||
(0, 0),
|
||||
(6, 0),
|
||||
(6, 4),
|
||||
(0, 4),
|
||||
)
|
||||
assert polygon_area(((0, 0), (5, 0), (5, 4), (0, 4))) == 20
|
||||
assert polygon_area(((0, 0), (0, 4), (5, 4), (5, 0))) == 20
|
||||
with pytest.raises(ValueError, match="at least three"):
|
||||
polygon_area(((0, 0), (1, 1)))
|
||||
|
||||
|
||||
def test_bbox_and_mask_iou():
|
||||
assert bbox_iou((0, 0, 10, 10), (5, 0, 15, 10)) == pytest.approx(1 / 3)
|
||||
assert bbox_iou((0, 0, 1, 1), (2, 2, 3, 3)) == 0
|
||||
first = torch.zeros((4, 4))
|
||||
second = torch.zeros((4, 4))
|
||||
first[:2, :2] = 1
|
||||
second[1:3, :2] = 1
|
||||
assert mask_iou(first, second) == pytest.approx(1 / 3)
|
||||
assert mask_iou(torch.zeros((2, 2)), torch.zeros((2, 2))) == 0
|
||||
with pytest.raises(ValueError, match="same shape"):
|
||||
mask_iou(torch.zeros((2, 2)), torch.zeros((3, 2)))
|
||||
|
||||
|
||||
def test_box_polygon_and_quad_rasterization():
|
||||
box = box_to_mask((1.2, 2.1, 4.1, 5.2), 8, 7)
|
||||
assert box.shape == (7, 8)
|
||||
assert box.sum().item() == 16
|
||||
polygon = polygon_to_mask(((1, 1), (5, 1), (5, 5), (1, 5)), 8, 8)
|
||||
quad = quad_to_mask(((1, 1), (5, 1), (5, 5), (1, 5)), 8, 8)
|
||||
assert torch.equal(polygon, quad)
|
||||
assert polygon.sum() > 0
|
||||
with pytest.raises(ValueError, match="exactly four"):
|
||||
quad_to_mask(((0, 0), (1, 0), (1, 1)), 4, 4)
|
||||
|
||||
|
||||
def test_detection_mask_priority_union_individual_and_bbox():
|
||||
explicit = torch.zeros((8, 8))
|
||||
explicit[3:6, 2:5] = 1
|
||||
with_mask = Detection(
|
||||
bbox_xyxy=(0, 0, 8, 8),
|
||||
polygon=((0, 0), (8, 0), (8, 8), (0, 8)),
|
||||
mask=explicit,
|
||||
)
|
||||
polygon_only = Detection(
|
||||
bbox_xyxy=(1, 1, 5, 5),
|
||||
polygon=((1, 1), (5, 1), (5, 5), (1, 5)),
|
||||
)
|
||||
assert torch.equal(detection_to_mask(with_mask, 8, 8), explicit)
|
||||
masks = individual_detection_masks((with_mask, polygon_only), 8, 8)
|
||||
assert masks.shape == (2, 8, 8)
|
||||
union = union_detection_mask((with_mask, polygon_only), 8, 8)
|
||||
assert union.shape == (8, 8)
|
||||
assert torch.all(union >= masks[0])
|
||||
assert individual_detection_masks((), 8, 8).shape == (0, 8, 8)
|
||||
assert union_detection_mask((), 8, 8).sum() == 0
|
||||
assert bbox_from_mask(explicit) == (2, 3, 5, 6)
|
||||
assert bbox_from_mask(torch.zeros((2, 2))) is None
|
||||
|
||||
|
||||
def test_deterministic_color_and_box_expansion():
|
||||
assert deterministic_color("track-1") == deterministic_color("track-1")
|
||||
assert deterministic_color("track-1") != deterministic_color("track-2")
|
||||
assert all(0 <= channel <= 255 for channel in deterministic_color("object"))
|
||||
assert expand_box((4, 4, 8, 6), 12, 12, padding=1) == (3, 3, 9, 7)
|
||||
squared = expand_box((4, 4, 8, 6), 12, 12, square=True)
|
||||
assert squared == (4, 3, 8, 7)
|
||||
assert translate_box((1, 2, 3, 4), 2, 1) == (3, 3, 5, 5)
|
||||
|
||||
|
||||
def test_label_aware_stable_association_with_motion():
|
||||
previous = (
|
||||
Detection(
|
||||
bbox_xyxy=(0, 0, 10, 10),
|
||||
label="cat",
|
||||
track_id=4,
|
||||
),
|
||||
Detection(
|
||||
bbox_xyxy=(20, 0, 30, 10),
|
||||
label="dog",
|
||||
track_id=9,
|
||||
),
|
||||
)
|
||||
current = (
|
||||
Detection(bbox_xyxy=(5, 0, 15, 10), label="cat"),
|
||||
Detection(bbox_xyxy=(20, 0, 30, 10), label="bird"),
|
||||
Detection(bbox_xyxy=(40, 0, 50, 10), label="dog"),
|
||||
)
|
||||
without_motion = associate_detections(
|
||||
previous,
|
||||
current,
|
||||
minimum_iou=0.3,
|
||||
)
|
||||
assert without_motion.matches == ((0, 0, pytest.approx(1 / 3)),)
|
||||
assert without_motion.unmatched_previous == (1,)
|
||||
assert without_motion.unmatched_current == (1, 2)
|
||||
|
||||
with_motion = associate_detections(
|
||||
previous,
|
||||
current,
|
||||
minimum_iou=0.9,
|
||||
motion_by_track={4: (5, 0), 9: (20, 0)},
|
||||
)
|
||||
assert with_motion.matches == (
|
||||
(0, 0, 1.0),
|
||||
(1, 2, 1.0),
|
||||
)
|
||||
assert with_motion.unmatched_previous == ()
|
||||
assert with_motion.unmatched_current == (1,)
|
||||
|
||||
label_agnostic = associate_detections(
|
||||
previous,
|
||||
current,
|
||||
minimum_iou=0.9,
|
||||
label_aware=False,
|
||||
)
|
||||
assert label_agnostic.matches == ((1, 1, 1.0),)
|
||||
@@ -0,0 +1,141 @@
|
||||
from types import SimpleNamespace
|
||||
|
||||
import torch
|
||||
from ComfyUI_VLM_nodes.nodes.grounding import (
|
||||
MODEL_SPECS,
|
||||
VLMOpenVocabularyDetection,
|
||||
core_bounding_box_frames,
|
||||
core_bounding_boxes,
|
||||
detection_box_masks,
|
||||
parse_labels,
|
||||
result_to_detections,
|
||||
)
|
||||
from ComfyUI_VLM_nodes.nodes.vision_types import (
|
||||
DetectionSequence,
|
||||
FrameDetections,
|
||||
)
|
||||
|
||||
|
||||
def test_detector_catalog_is_small_fast_and_portable():
|
||||
assert "Grounding DINO Tiny (fast)" in MODEL_SPECS
|
||||
assert "OmDet Turbo Swin Tiny (fast)" in MODEL_SPECS
|
||||
assert all("/" in spec.model_id for spec in MODEL_SPECS.values())
|
||||
schema = VLMOpenVocabularyDetection.INPUT_TYPES()
|
||||
assert tuple(MODEL_SPECS) == schema["required"]["model"][0]
|
||||
|
||||
|
||||
def test_label_parser_preserves_phrases_and_removes_duplicates():
|
||||
assert parse_labels("red car, person\nsmall dog;person") == [
|
||||
"red car",
|
||||
"person",
|
||||
"small dog",
|
||||
]
|
||||
|
||||
|
||||
def test_transformers_results_are_clipped_sorted_and_normalized():
|
||||
result = {
|
||||
"boxes": torch.tensor([[-5.0, 2.0, 20.0, 12.0], [5.0, 5.0, 9.0, 9.0]]),
|
||||
"scores": torch.tensor([0.25, 0.9]),
|
||||
"text_labels": ["cat", "dog"],
|
||||
}
|
||||
detections = result_to_detections(
|
||||
result,
|
||||
labels=["cat", "dog"],
|
||||
width=16,
|
||||
height=10,
|
||||
frame_index=0,
|
||||
timestamp=0.0,
|
||||
source="test/model",
|
||||
max_detections=20,
|
||||
)
|
||||
assert [item.label for item in detections] == ["dog", "cat"]
|
||||
assert detections[1].bbox_xyxy == (0.0, 2.0, 16.0, 10.0)
|
||||
assert detections[0].metadata["model_id"] == "test/model"
|
||||
|
||||
|
||||
def test_max_detections_is_applied_after_confidence_sorting():
|
||||
detections = result_to_detections(
|
||||
{
|
||||
"boxes": [[0, 0, 1, 1], [1, 1, 2, 2], [2, 2, 3, 3]],
|
||||
"scores": [0.1, 0.9, 0.8],
|
||||
"text_labels": ["low", "best", "second"],
|
||||
},
|
||||
labels=["low", "best", "second"],
|
||||
width=4,
|
||||
height=4,
|
||||
frame_index=0,
|
||||
timestamp=0.0,
|
||||
source="test/model",
|
||||
max_detections=2,
|
||||
)
|
||||
assert [item.label for item in detections] == ["best", "second"]
|
||||
|
||||
|
||||
def test_box_masks_and_core_boxes_keep_geometry_and_metadata():
|
||||
detections = result_to_detections(
|
||||
{
|
||||
"boxes": torch.tensor([[1.0, 2.0, 4.0, 5.0]]),
|
||||
"scores": torch.tensor([0.8]),
|
||||
"labels": torch.tensor([0]),
|
||||
},
|
||||
labels=["cat"],
|
||||
width=8,
|
||||
height=6,
|
||||
frame_index=0,
|
||||
timestamp=0.0,
|
||||
source="test/model",
|
||||
max_detections=5,
|
||||
)
|
||||
sequence = DetectionSequence(
|
||||
width=8,
|
||||
height=6,
|
||||
frames=(
|
||||
FrameDetections(
|
||||
frame_index=0,
|
||||
timestamp=0.0,
|
||||
width=8,
|
||||
height=6,
|
||||
detections=detections,
|
||||
),
|
||||
),
|
||||
frame_count=1,
|
||||
)
|
||||
masks = detection_box_masks(sequence)
|
||||
assert masks.shape == (1, 6, 8)
|
||||
assert masks.sum().item() == 9
|
||||
boxes = core_bounding_boxes(sequence)
|
||||
assert boxes == [
|
||||
{
|
||||
"x": 1,
|
||||
"y": 2,
|
||||
"width": 3,
|
||||
"height": 3,
|
||||
"label": "cat",
|
||||
"score": detections[0].score,
|
||||
"metadata": {
|
||||
"frame_index": 0,
|
||||
"label": "cat",
|
||||
"score": detections[0].score,
|
||||
"source": "test/model",
|
||||
},
|
||||
}
|
||||
]
|
||||
assert core_bounding_box_frames(sequence) == [boxes]
|
||||
|
||||
|
||||
def test_result_label_indices_are_resolved():
|
||||
detections = result_to_detections(
|
||||
{
|
||||
"boxes": [[0, 0, 4, 4]],
|
||||
"scores": [SimpleNamespace(item=lambda: 0.5)],
|
||||
"classes": [SimpleNamespace(item=lambda: 1)],
|
||||
},
|
||||
labels=["cat", "dog"],
|
||||
width=4,
|
||||
height=4,
|
||||
frame_index=0,
|
||||
timestamp=0.0,
|
||||
source="omdet",
|
||||
max_detections=1,
|
||||
)
|
||||
assert detections[0].label == "dog"
|
||||
@@ -0,0 +1,904 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from ComfyUI_VLM_nodes.nodes import hosted_api
|
||||
|
||||
|
||||
class FakeHttpClient:
|
||||
def __init__(self, **kwargs):
|
||||
self.kwargs = kwargs
|
||||
|
||||
|
||||
class FakeResponses:
|
||||
def __init__(self, calls, failure=None, response_text="secure response"):
|
||||
self.calls = calls
|
||||
self.failure = failure
|
||||
self.response_text = response_text
|
||||
|
||||
def create(self, **kwargs):
|
||||
self.calls.append(("responses", kwargs))
|
||||
if self.failure is not None:
|
||||
raise self.failure
|
||||
return SimpleNamespace(output_text=self.response_text)
|
||||
|
||||
|
||||
class FakeChatCompletions:
|
||||
def __init__(self, calls, failure=None, response_text="secure response"):
|
||||
self.calls = calls
|
||||
self.failure = failure
|
||||
self.response_text = response_text
|
||||
|
||||
def create(self, **kwargs):
|
||||
self.calls.append(("chat", kwargs))
|
||||
if self.failure is not None:
|
||||
raise self.failure
|
||||
message = SimpleNamespace(content=self.response_text)
|
||||
return SimpleNamespace(choices=[SimpleNamespace(message=message)])
|
||||
|
||||
|
||||
def fake_openai_module(calls, failure=None, response_text="secure response"):
|
||||
class FakeOpenAI:
|
||||
def __init__(self, **kwargs):
|
||||
calls.append(("client", kwargs))
|
||||
self.responses = FakeResponses(
|
||||
calls,
|
||||
failure=failure,
|
||||
response_text=response_text,
|
||||
)
|
||||
self.chat = SimpleNamespace(
|
||||
completions=FakeChatCompletions(
|
||||
calls,
|
||||
failure=failure,
|
||||
response_text=response_text,
|
||||
)
|
||||
)
|
||||
|
||||
def close(self):
|
||||
calls.append(("close", {}))
|
||||
|
||||
return SimpleNamespace(
|
||||
OpenAI=FakeOpenAI,
|
||||
DefaultHttpxClient=FakeHttpClient,
|
||||
)
|
||||
|
||||
|
||||
def test_api_schemas_never_accept_plaintext_keys():
|
||||
for node_class in (hosted_api.PromptGenerateAPI, hosted_api.HostedVLMAPI):
|
||||
schema = node_class.INPUT_TYPES()
|
||||
all_inputs = {
|
||||
**schema.get("required", {}),
|
||||
**schema.get("optional", {}),
|
||||
**schema.get("hidden", {}),
|
||||
}
|
||||
assert "api_key" not in all_inputs
|
||||
assert "credential_source" in all_inputs
|
||||
assert "STRING" not in repr(all_inputs["credential_source"][0])
|
||||
assert "web_search" in all_inputs
|
||||
assert "output_format" in all_inputs
|
||||
assert "json_schema" in all_inputs
|
||||
assert "schema_api_style" in all_inputs
|
||||
|
||||
|
||||
def test_json_schema_parser_blocks_remote_refs_and_bounds_input():
|
||||
for keyword in ("$ref", "$dynamicRef", "$recursiveRef"):
|
||||
with pytest.raises(ValueError, match="only local fragment"):
|
||||
hosted_api.parse_json_schema(
|
||||
"JSON Schema",
|
||||
json.dumps(
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"payload": {
|
||||
keyword: "https://attacker.example/schema.json"
|
||||
}
|
||||
},
|
||||
}
|
||||
),
|
||||
)
|
||||
with pytest.raises(ValueError, match="64,000"):
|
||||
hosted_api.parse_json_schema("JSON Schema", "x" * 64_001)
|
||||
|
||||
|
||||
def test_local_structured_output_validation_is_strict_and_normalized():
|
||||
schema_text = json.dumps(
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {"count": {"type": "integer"}},
|
||||
"required": ["count"],
|
||||
"additionalProperties": False,
|
||||
}
|
||||
)
|
||||
schema = hosted_api.parse_json_schema("JSON Schema", schema_text)
|
||||
assert hosted_api.validate_structured_output(
|
||||
'```json\n{"count": 2}\n```',
|
||||
"JSON Schema",
|
||||
schema,
|
||||
) == '{\n "count": 2\n}'
|
||||
with pytest.raises(RuntimeError, match=r"\$\.count \(type constraint\)"):
|
||||
hosted_api.validate_structured_output(
|
||||
'{"count": "two"}',
|
||||
"JSON Schema",
|
||||
schema,
|
||||
)
|
||||
with pytest.raises(RuntimeError, match="valid JSON"):
|
||||
hosted_api.validate_structured_output(
|
||||
'{"count":',
|
||||
"JSON Schema",
|
||||
schema,
|
||||
)
|
||||
|
||||
|
||||
def test_provider_catalog_uses_current_bound_credentials_and_endpoints():
|
||||
assert len(hosted_api.PROVIDER_PROFILES) >= 18
|
||||
expected = {
|
||||
"OpenAI": "OPENAI_API_KEY",
|
||||
"Google Gemini": "GEMINI_API_KEY",
|
||||
"Anthropic": "ANTHROPIC_API_KEY",
|
||||
"xAI": "XAI_API_KEY",
|
||||
"DeepSeek": "DEEPSEEK_API_KEY",
|
||||
"Groq": "GROQ_API_KEY",
|
||||
"Mistral": "MISTRAL_API_KEY",
|
||||
"Together AI": "TOGETHER_API_KEY",
|
||||
"OpenRouter": "OPENROUTER_API_KEY",
|
||||
"Custom / Local": "CUSTOM_API_KEY",
|
||||
}
|
||||
providers = {
|
||||
profile.provider: profile.api_key_env
|
||||
for profile in hosted_api.PROVIDER_PROFILES.values()
|
||||
}
|
||||
assert expected.items() <= providers.items()
|
||||
for profile in hosted_api.PROVIDER_PROFILES.values():
|
||||
if profile.base_url is not None:
|
||||
assert profile.base_url.startswith("https://")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"url",
|
||||
[
|
||||
"http://example.com/v1",
|
||||
"ftp://127.0.0.1/v1",
|
||||
"https://user:secret@example.com/v1",
|
||||
"https://example.com/v1?api_key=secret",
|
||||
"not-a-url",
|
||||
],
|
||||
)
|
||||
def test_custom_endpoint_rejects_unsafe_urls(url):
|
||||
with pytest.raises(ValueError):
|
||||
hosted_api.validate_custom_base_url(url)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"url",
|
||||
[
|
||||
"http://127.0.0.1:8000/v1",
|
||||
"http://[::1]:11434/v1",
|
||||
"http://localhost:1234/v1",
|
||||
"https://example.com/v1/",
|
||||
],
|
||||
)
|
||||
def test_custom_endpoint_accepts_https_or_loopback(url):
|
||||
normalized, loopback = hosted_api.validate_custom_base_url(url)
|
||||
assert normalized.startswith(("http://", "https://"))
|
||||
assert loopback is (url.startswith("http://"))
|
||||
|
||||
|
||||
def test_built_in_key_cannot_be_redirected(monkeypatch):
|
||||
profile = hosted_api.provider_profile("OpenAI — GPT-5.6 Terra")
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "sk-real-secret-value")
|
||||
with pytest.raises(ValueError, match="pinned to official hosts"):
|
||||
hosted_api.resolve_endpoint(profile, "https://attacker.example/v1")
|
||||
|
||||
|
||||
def test_legacy_plaintext_value_is_rejected_without_echo(monkeypatch):
|
||||
profile = hosted_api.provider_profile("OpenAI — GPT-5.6 Terra")
|
||||
secret = "sk-legacy-plaintext-that-must-not-appear"
|
||||
with pytest.raises(ValueError) as captured:
|
||||
hosted_api.resolve_api_key(profile, secret, loopback=False)
|
||||
assert secret not in str(captured.value)
|
||||
assert "legacy plaintext API key was removed" in str(captured.value)
|
||||
|
||||
|
||||
def test_redaction_removes_exact_encoded_and_header_credentials():
|
||||
secret = "sk-ant-example-SECRET_123456789"
|
||||
message = (
|
||||
f"Authorization: Bearer {secret}; api_key={secret}; "
|
||||
f"url=https://user:{secret}@example.com; encoded={secret}"
|
||||
)
|
||||
redacted = hosted_api.redact_sensitive(message, (secret,))
|
||||
assert secret not in redacted
|
||||
assert "Bearer" not in redacted
|
||||
assert "[REDACTED]" in redacted
|
||||
|
||||
|
||||
def test_responses_call_is_stateless_private_and_provider_bound(monkeypatch):
|
||||
calls = []
|
||||
secret = "sk-openai-provider-bound-secret"
|
||||
monkeypatch.setenv("OPENAI_API_KEY", secret)
|
||||
monkeypatch.setattr(
|
||||
hosted_api,
|
||||
"require_module",
|
||||
lambda *_args: fake_openai_module(calls),
|
||||
)
|
||||
|
||||
node = hosted_api.PromptGenerateAPI()
|
||||
assert not hasattr(node, "session_history")
|
||||
result = node.generate_prompt(
|
||||
"OpenAI — GPT-5.6 Terra",
|
||||
False,
|
||||
hosted_api.PROVIDER_CREDENTIAL,
|
||||
"A scene",
|
||||
"Improve it",
|
||||
0,
|
||||
0,
|
||||
stream_output=False,
|
||||
)
|
||||
assert result == ("secure response",)
|
||||
|
||||
client_kwargs = next(payload for kind, payload in calls if kind == "client")
|
||||
assert client_kwargs["api_key"] == secret
|
||||
assert "base_url" not in client_kwargs
|
||||
assert client_kwargs["http_client"].kwargs["follow_redirects"] is False
|
||||
assert client_kwargs["http_client"].kwargs["trust_env"] is False
|
||||
|
||||
request = next(payload for kind, payload in calls if kind == "responses")
|
||||
assert request["model"] == "gpt-5.6-terra"
|
||||
assert request["store"] is False
|
||||
assert "previous_response_id" not in request
|
||||
assert "metadata" not in request
|
||||
|
||||
|
||||
def test_openai_combines_web_search_structured_output_and_stream_contract(
|
||||
monkeypatch,
|
||||
):
|
||||
calls = []
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "sk-test-structured-search")
|
||||
monkeypatch.setattr(
|
||||
hosted_api,
|
||||
"require_module",
|
||||
lambda import_name, *_args: (
|
||||
fake_openai_module(calls, response_text='{"answer":"grounded"}')
|
||||
if import_name == "openai"
|
||||
else __import__(import_name)
|
||||
),
|
||||
)
|
||||
schema = json.dumps(
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {"answer": {"type": "string"}},
|
||||
"required": ["answer"],
|
||||
"additionalProperties": False,
|
||||
}
|
||||
)
|
||||
result = hosted_api.PromptGenerateAPI().generate_prompt(
|
||||
"OpenAI — GPT-5.6 Sol",
|
||||
False,
|
||||
hosted_api.PROVIDER_CREDENTIAL,
|
||||
"Find a current fact",
|
||||
"",
|
||||
0,
|
||||
0,
|
||||
web_search=True,
|
||||
output_format="JSON Schema",
|
||||
json_schema=schema,
|
||||
stream_output=False,
|
||||
)
|
||||
assert result == ('{\n "answer": "grounded"\n}',)
|
||||
request = next(payload for kind, payload in calls if kind == "responses")
|
||||
assert request["tools"] == [{"type": "web_search"}]
|
||||
assert request["text"]["format"]["type"] == "json_schema"
|
||||
assert request["text"]["format"]["strict"] is True
|
||||
assert request["text"]["format"]["schema"]["required"] == ["answer"]
|
||||
assert "JSON Schema:" in request["instructions"]
|
||||
|
||||
|
||||
def test_unsupported_web_search_fails_before_network(monkeypatch):
|
||||
monkeypatch.setenv("DEEPSEEK_API_KEY", "deepseek-test-secret")
|
||||
with pytest.raises(ValueError, match="does not expose native web search"):
|
||||
hosted_api.PromptGenerateAPI().generate_prompt(
|
||||
"DeepSeek — V4 Flash",
|
||||
False,
|
||||
hosted_api.PROVIDER_CREDENTIAL,
|
||||
"Search now",
|
||||
"",
|
||||
0,
|
||||
0,
|
||||
web_search=True,
|
||||
stream_output=False,
|
||||
)
|
||||
|
||||
|
||||
def test_responses_stream_collects_deltas_and_closes():
|
||||
class Stream(list):
|
||||
closed = False
|
||||
|
||||
def close(self):
|
||||
self.closed = True
|
||||
|
||||
stream = Stream(
|
||||
[
|
||||
SimpleNamespace(type="response.created"),
|
||||
SimpleNamespace(type="response.output_text.delta", delta="hello "),
|
||||
SimpleNamespace(type="response.output_text.delta", delta="world"),
|
||||
]
|
||||
)
|
||||
client = SimpleNamespace(
|
||||
responses=SimpleNamespace(create=lambda **kwargs: stream)
|
||||
)
|
||||
assert hosted_api._stream_responses(client, {"model": "test"}, None) == (
|
||||
"hello world"
|
||||
)
|
||||
assert stream.closed is True
|
||||
|
||||
|
||||
def test_chat_stream_collects_deltas_and_closes():
|
||||
class Stream(list):
|
||||
closed = False
|
||||
|
||||
def close(self):
|
||||
self.closed = True
|
||||
|
||||
stream = Stream(
|
||||
[
|
||||
SimpleNamespace(
|
||||
choices=[
|
||||
SimpleNamespace(delta=SimpleNamespace(content="frame "))
|
||||
]
|
||||
),
|
||||
SimpleNamespace(
|
||||
choices=[
|
||||
SimpleNamespace(delta=SimpleNamespace(content="ready"))
|
||||
]
|
||||
),
|
||||
]
|
||||
)
|
||||
client = SimpleNamespace(
|
||||
chat=SimpleNamespace(
|
||||
completions=SimpleNamespace(create=lambda **kwargs: stream)
|
||||
)
|
||||
)
|
||||
assert hosted_api._stream_chat(client, {"model": "test"}, None) == (
|
||||
"frame ready"
|
||||
)
|
||||
assert stream.closed is True
|
||||
|
||||
|
||||
def test_provider_failure_never_echoes_api_key(monkeypatch):
|
||||
calls = []
|
||||
secret = "sk-secret-reflected-by-provider-123456"
|
||||
failure = RuntimeError(f"Authorization: Bearer {secret} api_key={secret}")
|
||||
monkeypatch.setenv("OPENAI_API_KEY", secret)
|
||||
monkeypatch.setattr(
|
||||
hosted_api,
|
||||
"require_module",
|
||||
lambda *_args: fake_openai_module(calls, failure=failure),
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError) as captured:
|
||||
hosted_api.PromptGenerateAPI().generate_prompt(
|
||||
"OpenAI — GPT-5.6 Terra",
|
||||
False,
|
||||
hosted_api.PROVIDER_CREDENTIAL,
|
||||
"hello",
|
||||
"",
|
||||
0,
|
||||
0,
|
||||
stream_output=False,
|
||||
)
|
||||
assert secret not in str(captured.value)
|
||||
assert "[REDACTED]" in str(captured.value)
|
||||
|
||||
|
||||
def test_anthropic_uses_native_messages_and_keeps_key_out_of_body(monkeypatch):
|
||||
calls = []
|
||||
secret = "sk-ant-native-secret"
|
||||
|
||||
class FakeResponse:
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
def json(self):
|
||||
return {
|
||||
"content": [
|
||||
{"type": "text", "text": "native Anthropic response"}
|
||||
]
|
||||
}
|
||||
|
||||
class FakeClient:
|
||||
def __init__(self, **kwargs):
|
||||
calls.append(("client", kwargs))
|
||||
|
||||
def post(self, url, **kwargs):
|
||||
calls.append(("post", {"url": url, **kwargs}))
|
||||
return FakeResponse()
|
||||
|
||||
def close(self):
|
||||
calls.append(("close", {}))
|
||||
|
||||
monkeypatch.setenv("ANTHROPIC_API_KEY", secret)
|
||||
monkeypatch.setattr(
|
||||
hosted_api,
|
||||
"require_module",
|
||||
lambda import_name, *_args: (
|
||||
SimpleNamespace(Client=FakeClient)
|
||||
if import_name == "httpx"
|
||||
else pytest.fail(f"Unexpected module request: {import_name}")
|
||||
),
|
||||
)
|
||||
result = hosted_api.HostedVLMAPI().analyze(
|
||||
"Anthropic — Claude Sonnet 5",
|
||||
hosted_api.PROVIDER_CREDENTIAL,
|
||||
"Read this image.",
|
||||
"Be concise.",
|
||||
1,
|
||||
512,
|
||||
80,
|
||||
"auto",
|
||||
images=torch.rand((1, 48, 64, 3)),
|
||||
stream_output=False,
|
||||
)
|
||||
assert result == (
|
||||
"native Anthropic response",
|
||||
"claude-sonnet-5",
|
||||
1,
|
||||
)
|
||||
request = next(payload for kind, payload in calls if kind == "post")
|
||||
assert request["url"] == "https://api.anthropic.com/v1/messages"
|
||||
assert request["headers"]["x-api-key"] == secret
|
||||
assert secret not in repr(request["json"])
|
||||
content = request["json"]["messages"][0]["content"]
|
||||
assert content[1]["type"] == "image"
|
||||
assert content[1]["source"]["type"] == "base64"
|
||||
assert request["json"]["stream"] is False
|
||||
|
||||
|
||||
def test_anthropic_native_stream_collects_text_deltas(monkeypatch):
|
||||
class FakeStreamResponse:
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *_args):
|
||||
return False
|
||||
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
def iter_lines(self):
|
||||
yield 'event: content_block_delta'
|
||||
yield (
|
||||
'data: {"type":"content_block_delta","delta":'
|
||||
'{"type":"text_delta","text":"hello "}}'
|
||||
)
|
||||
yield (
|
||||
'data: {"type":"content_block_delta","delta":'
|
||||
'{"type":"text_delta","text":"world"}}'
|
||||
)
|
||||
|
||||
class FakeClient:
|
||||
def __init__(self, **_kwargs):
|
||||
pass
|
||||
|
||||
def stream(self, *_args, **_kwargs):
|
||||
return FakeStreamResponse()
|
||||
|
||||
def close(self):
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(
|
||||
hosted_api,
|
||||
"require_module",
|
||||
lambda import_name, *_args: (
|
||||
SimpleNamespace(Client=FakeClient)
|
||||
if import_name == "httpx"
|
||||
else __import__(import_name)
|
||||
),
|
||||
)
|
||||
profile = hosted_api.provider_profile("Anthropic — Claude Sonnet 5")
|
||||
result = hosted_api._call_anthropic_api(
|
||||
profile=profile,
|
||||
model=profile.model,
|
||||
endpoint=profile.base_url,
|
||||
api_key="sk-ant-stream",
|
||||
system_prompt="Be concise.",
|
||||
prompt="Hello",
|
||||
image_data=[],
|
||||
timeout_seconds=30,
|
||||
max_output_tokens=100,
|
||||
stream_output=True,
|
||||
use_system_proxy=False,
|
||||
unique_id=None,
|
||||
web_search=False,
|
||||
output_format="Text",
|
||||
output_schema=None,
|
||||
)
|
||||
assert result == "hello world"
|
||||
|
||||
|
||||
def test_anthropic_native_search_and_structured_contracts(monkeypatch):
|
||||
calls = []
|
||||
|
||||
class FakeResponse:
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
def json(self):
|
||||
return {"content": [{"type": "text", "text": '{"answer":"yes"}'}]}
|
||||
|
||||
class FakeClient:
|
||||
def __init__(self, **kwargs):
|
||||
calls.append(("client", kwargs))
|
||||
|
||||
def post(self, url, **kwargs):
|
||||
calls.append(("post", {"url": url, **kwargs}))
|
||||
return FakeResponse()
|
||||
|
||||
def close(self):
|
||||
return None
|
||||
|
||||
monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-ant-contract")
|
||||
monkeypatch.setattr(
|
||||
hosted_api,
|
||||
"require_module",
|
||||
lambda import_name, *_args: (
|
||||
SimpleNamespace(Client=FakeClient)
|
||||
if import_name == "httpx"
|
||||
else __import__(import_name)
|
||||
),
|
||||
)
|
||||
schema = (
|
||||
'{"type":"object","properties":{"answer":{"type":"string"}},'
|
||||
'"required":["answer"],"additionalProperties":false}'
|
||||
)
|
||||
|
||||
result = hosted_api.PromptGenerateAPI().generate_prompt(
|
||||
"Anthropic — Claude Sonnet 5",
|
||||
False,
|
||||
hosted_api.PROVIDER_CREDENTIAL,
|
||||
"Return a value",
|
||||
"",
|
||||
0,
|
||||
0,
|
||||
output_format="JSON Schema",
|
||||
json_schema=schema,
|
||||
stream_output=False,
|
||||
)
|
||||
assert json.loads(result[0]) == {"answer": "yes"}
|
||||
structured = next(payload for kind, payload in calls if kind == "post")
|
||||
assert structured["json"]["output_config"]["format"]["type"] == "json_schema"
|
||||
|
||||
calls.clear()
|
||||
hosted_api.PromptGenerateAPI().generate_prompt(
|
||||
"Anthropic — Claude Sonnet 5",
|
||||
False,
|
||||
hosted_api.PROVIDER_CREDENTIAL,
|
||||
"Search the web",
|
||||
"",
|
||||
0,
|
||||
0,
|
||||
web_search=True,
|
||||
output_format="JSON Schema",
|
||||
json_schema=schema,
|
||||
stream_output=False,
|
||||
)
|
||||
searched = next(payload for kind, payload in calls if kind == "post")
|
||||
assert searched["json"]["tools"][0]["type"] == "web_search_20260318"
|
||||
assert searched["json"]["tools"][0]["allowed_callers"] == ["direct"]
|
||||
assert "output_config" not in searched["json"]
|
||||
|
||||
|
||||
def test_gemini_native_search_vision_and_schema_contract(monkeypatch):
|
||||
calls = []
|
||||
secret = "gemini-provider-secret"
|
||||
|
||||
class FakeResponse:
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
def json(self):
|
||||
return {
|
||||
"candidates": [
|
||||
{
|
||||
"content": {
|
||||
"parts": [{"text": '{"objects":["tree"]}'}]
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
class FakeClient:
|
||||
def __init__(self, **kwargs):
|
||||
calls.append(("client", kwargs))
|
||||
|
||||
def post(self, url, **kwargs):
|
||||
calls.append(("post", {"url": url, **kwargs}))
|
||||
return FakeResponse()
|
||||
|
||||
def close(self):
|
||||
return None
|
||||
|
||||
monkeypatch.setenv("GEMINI_API_KEY", secret)
|
||||
monkeypatch.setattr(
|
||||
hosted_api,
|
||||
"require_module",
|
||||
lambda import_name, *_args: (
|
||||
SimpleNamespace(Client=FakeClient)
|
||||
if import_name == "httpx"
|
||||
else __import__(import_name)
|
||||
),
|
||||
)
|
||||
schema = (
|
||||
'{"type":"object","properties":{"objects":{"type":"array",'
|
||||
'"items":{"type":"string"}}},"required":["objects"]}'
|
||||
)
|
||||
result = hosted_api.HostedVLMAPI().analyze(
|
||||
"Google — Gemini 3.6 Flash",
|
||||
hosted_api.PROVIDER_CREDENTIAL,
|
||||
"Identify objects using current context.",
|
||||
"Be precise.",
|
||||
1,
|
||||
512,
|
||||
80,
|
||||
"auto",
|
||||
images=torch.rand((1, 32, 48, 3)),
|
||||
web_search=True,
|
||||
output_format="JSON Schema",
|
||||
json_schema=schema,
|
||||
stream_output=False,
|
||||
)
|
||||
assert json.loads(result[0]) == {"objects": ["tree"]}
|
||||
request = next(payload for kind, payload in calls if kind == "post")
|
||||
assert request["url"].endswith(
|
||||
"/models/gemini-3.6-flash:generateContent"
|
||||
)
|
||||
assert request["headers"]["x-goog-api-key"] == secret
|
||||
assert secret not in repr(request["json"])
|
||||
assert request["json"]["tools"] == [{"google_search": {}}]
|
||||
assert (
|
||||
request["json"]["generationConfig"]["responseFormat"]["text"]["schema"][
|
||||
"required"
|
||||
]
|
||||
== ["objects"]
|
||||
)
|
||||
inline = request["json"]["contents"][0]["parts"][1]["inlineData"]
|
||||
assert inline["mimeType"] == "image/jpeg"
|
||||
assert inline["data"]
|
||||
|
||||
|
||||
def test_gemini_native_stream_collects_sse_deltas(monkeypatch):
|
||||
class FakeStreamResponse:
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *_args):
|
||||
return False
|
||||
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
def iter_lines(self):
|
||||
yield (
|
||||
'data: {"candidates":[{"content":{"parts":'
|
||||
'[{"text":"frame "}]}}]}'
|
||||
)
|
||||
yield (
|
||||
'data: {"candidates":[{"content":{"parts":'
|
||||
'[{"text":"ready"}]}}]}'
|
||||
)
|
||||
|
||||
class FakeClient:
|
||||
def __init__(self, **_kwargs):
|
||||
pass
|
||||
|
||||
def stream(self, *_args, **_kwargs):
|
||||
return FakeStreamResponse()
|
||||
|
||||
def close(self):
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(
|
||||
hosted_api,
|
||||
"require_module",
|
||||
lambda import_name, *_args: (
|
||||
SimpleNamespace(Client=FakeClient)
|
||||
if import_name == "httpx"
|
||||
else __import__(import_name)
|
||||
),
|
||||
)
|
||||
profile = hosted_api.provider_profile("Google — Gemini 3.6 Flash")
|
||||
result = hosted_api._call_gemini_api(
|
||||
profile=profile,
|
||||
model=profile.model,
|
||||
api_key="gemini-stream",
|
||||
system_prompt="Be concise.",
|
||||
prompt="Describe.",
|
||||
image_data=[],
|
||||
timeout_seconds=30,
|
||||
max_output_tokens=100,
|
||||
stream_output=True,
|
||||
use_system_proxy=False,
|
||||
unique_id=None,
|
||||
web_search=True,
|
||||
output_format="Text",
|
||||
output_schema=None,
|
||||
)
|
||||
assert result == "frame ready"
|
||||
|
||||
|
||||
def test_vlm_uniformly_samples_and_bounds_image_batch(monkeypatch):
|
||||
calls = []
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "sk-test-only")
|
||||
monkeypatch.setattr(
|
||||
hosted_api,
|
||||
"require_module",
|
||||
lambda *_args: fake_openai_module(calls),
|
||||
)
|
||||
images = torch.rand((10, 96, 128, 3), dtype=torch.float32)
|
||||
result = hosted_api.HostedVLMAPI().analyze(
|
||||
"OpenAI — GPT-5.6 Terra",
|
||||
hosted_api.PROVIDER_CREDENTIAL,
|
||||
"Compare the sampled frames.",
|
||||
"Be precise.",
|
||||
4,
|
||||
768,
|
||||
82,
|
||||
"low",
|
||||
images=images,
|
||||
stream_output=False,
|
||||
)
|
||||
assert result == ("secure response", "gpt-5.6-terra", 4)
|
||||
request = next(payload for kind, payload in calls if kind == "responses")
|
||||
content = request["input"][0]["content"]
|
||||
image_parts = [part for part in content if part["type"] == "input_image"]
|
||||
assert len(image_parts) == 4
|
||||
assert all(part["image_url"].startswith("data:image/jpeg;base64,") for part in image_parts)
|
||||
assert all(part["detail"] == "low" for part in image_parts)
|
||||
|
||||
|
||||
def test_open_source_vlm_llama_cpp_schema_dialect_and_local_validation(
|
||||
monkeypatch,
|
||||
):
|
||||
calls = []
|
||||
monkeypatch.setattr(
|
||||
hosted_api,
|
||||
"require_module",
|
||||
lambda import_name, *_args: (
|
||||
fake_openai_module(calls, response_text='{"objects":["cat"]}')
|
||||
if import_name == "openai"
|
||||
else __import__(import_name)
|
||||
),
|
||||
)
|
||||
schema = json.dumps(
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"objects": {
|
||||
"type": "array",
|
||||
"items": {"type": "string"},
|
||||
}
|
||||
},
|
||||
"required": ["objects"],
|
||||
"additionalProperties": False,
|
||||
}
|
||||
)
|
||||
result = hosted_api.HostedVLMAPI().analyze(
|
||||
"Custom / Local — OpenAI compatible",
|
||||
hosted_api.LOCAL_NO_KEY,
|
||||
"List visible objects.",
|
||||
"Be precise.",
|
||||
1,
|
||||
512,
|
||||
80,
|
||||
"auto",
|
||||
images=torch.rand((1, 32, 48, 3)),
|
||||
base_url="http://127.0.0.1:8080/v1",
|
||||
model_override="local-vlm",
|
||||
output_format="JSON Schema",
|
||||
json_schema=schema,
|
||||
schema_api_style="llama.cpp JSON Schema",
|
||||
stream_output=False,
|
||||
)
|
||||
assert result == (
|
||||
'{\n "objects": [\n "cat"\n ]\n}',
|
||||
"local-vlm",
|
||||
1,
|
||||
)
|
||||
request = next(payload for kind, payload in calls if kind == "chat")
|
||||
assert request["response_format"] == {
|
||||
"type": "json_schema",
|
||||
"schema": json.loads(schema),
|
||||
}
|
||||
image = request["messages"][1]["content"][1]
|
||||
assert image["type"] == "image_url"
|
||||
assert image["image_url"]["url"].startswith("data:image/jpeg;base64,")
|
||||
|
||||
|
||||
def test_custom_openai_schema_style_uses_standard_wrapper(monkeypatch):
|
||||
calls = []
|
||||
monkeypatch.setattr(
|
||||
hosted_api,
|
||||
"require_module",
|
||||
lambda import_name, *_args: (
|
||||
fake_openai_module(calls, response_text='{"ok":true}')
|
||||
if import_name == "openai"
|
||||
else __import__(import_name)
|
||||
),
|
||||
)
|
||||
schema = '{"type":"object","properties":{"ok":{"type":"boolean"}},"required":["ok"]}'
|
||||
result, _, _ = hosted_api.execute_hosted(
|
||||
model_name="Custom / Local — OpenAI compatible",
|
||||
credential_source=hosted_api.LOCAL_NO_KEY,
|
||||
prompt="Return status.",
|
||||
system_prompt="Be exact.",
|
||||
base_url="http://localhost:8000/v1",
|
||||
model_override="local",
|
||||
api_mode="Chat Completions",
|
||||
timeout_seconds=30,
|
||||
max_output_tokens=100,
|
||||
reasoning_effort="none",
|
||||
seed=0,
|
||||
stream_output=False,
|
||||
use_system_proxy=False,
|
||||
unique_id=None,
|
||||
output_format="JSON Schema",
|
||||
json_schema=schema,
|
||||
)
|
||||
assert json.loads(result) == {"ok": True}
|
||||
request = next(payload for kind, payload in calls if kind == "chat")
|
||||
assert request["response_format"]["json_schema"]["strict"] is True
|
||||
assert request["response_format"]["json_schema"]["schema"]["required"] == [
|
||||
"ok"
|
||||
]
|
||||
|
||||
|
||||
def test_groq_auto_uses_documented_chat_route_for_structured_output(monkeypatch):
|
||||
calls = []
|
||||
monkeypatch.setenv("GROQ_API_KEY", "gsk-test-structured")
|
||||
monkeypatch.setattr(
|
||||
hosted_api,
|
||||
"require_module",
|
||||
lambda import_name, *_args: (
|
||||
fake_openai_module(calls, response_text='{"ok":true}')
|
||||
if import_name == "openai"
|
||||
else __import__(import_name)
|
||||
),
|
||||
)
|
||||
hosted_api.execute_hosted(
|
||||
model_name="Groq — GPT-OSS 20B",
|
||||
credential_source=hosted_api.PROVIDER_CREDENTIAL,
|
||||
prompt="Return status.",
|
||||
system_prompt="Be exact.",
|
||||
base_url="",
|
||||
model_override="",
|
||||
api_mode="Auto",
|
||||
timeout_seconds=30,
|
||||
max_output_tokens=100,
|
||||
reasoning_effort="none",
|
||||
seed=0,
|
||||
stream_output=False,
|
||||
use_system_proxy=False,
|
||||
unique_id=None,
|
||||
output_format="JSON Schema",
|
||||
json_schema=(
|
||||
'{"type":"object","properties":{"ok":{"type":"boolean"}},'
|
||||
'"required":["ok"],"additionalProperties":false}'
|
||||
),
|
||||
)
|
||||
assert any(kind == "chat" for kind, _payload in calls)
|
||||
assert not any(kind == "responses" for kind, _payload in calls)
|
||||
|
||||
|
||||
def test_frontend_scrubs_legacy_key_before_graph_configuration():
|
||||
web_root = Path(__file__).resolve().parents[1] / "web" / "js"
|
||||
source = (
|
||||
web_root / "apiSecurity.js"
|
||||
).read_text("utf-8")
|
||||
assert "beforeConfigureGraph" in source
|
||||
assert "delete values.api_key" in source
|
||||
assert "CREDENTIAL_WIDGET_INDEX = 2" in source
|
||||
view_text = (web_root / "viewText.js").read_text("utf-8")
|
||||
assert '"PromptGenerateAPI"' in view_text
|
||||
assert '"HostedVLMAPI"' in view_text
|
||||
@@ -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
|
||||
@@ -0,0 +1,361 @@
|
||||
import asyncio
|
||||
import importlib
|
||||
import inspect
|
||||
import json
|
||||
import sys
|
||||
import types
|
||||
from dataclasses import dataclass
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
PACKAGE = __package__.split(".")[0] if __package__ else "ComfyUI_VLM_nodes"
|
||||
module = importlib.import_module(f"{PACKAGE}.nodes.moondream31")
|
||||
worker = importlib.import_module(f"{PACKAGE}.nodes.moondream31_worker")
|
||||
|
||||
Moondream31Detect = module.Moondream31Detect
|
||||
Moondream31Loader = module.Moondream31Loader
|
||||
Moondream31Model = module.Moondream31Model
|
||||
Moondream31Segment = module.Moondream31Segment
|
||||
svg_path_to_mask = module.svg_path_to_mask
|
||||
|
||||
|
||||
def _fake_model(handler, model_name=module.MODEL_ID):
|
||||
model = object.__new__(Moondream31Model)
|
||||
model.config = module.Moondream31Config(
|
||||
model=model_name,
|
||||
device="cuda",
|
||||
max_batch_size=4,
|
||||
kv_cache_pages=8192,
|
||||
)
|
||||
model.request = handler
|
||||
model.close = lambda: None
|
||||
return model
|
||||
|
||||
|
||||
def test_svg_path_is_transformed_from_bbox_space_to_image_pixels():
|
||||
mask, polygon, contours = svg_path_to_mask(
|
||||
"M 0 0 H 1 V 1 H 0 Z",
|
||||
{"x_min": 0.25, "y_min": 0.25, "x_max": 0.75, "y_max": 0.75},
|
||||
100,
|
||||
80,
|
||||
supersample=4,
|
||||
)
|
||||
assert mask.shape == (80, 100)
|
||||
assert mask[40, 50] > 0.99
|
||||
assert mask[5, 5] == 0
|
||||
assert mask.sum().item() == pytest.approx(2000, rel=0.06)
|
||||
assert len(polygon) >= 4
|
||||
assert len(contours) == 1
|
||||
xs = [point[0] for point in polygon]
|
||||
ys = [point[1] for point in polygon]
|
||||
assert min(xs) == pytest.approx(25)
|
||||
assert max(xs) == pytest.approx(75)
|
||||
assert min(ys) == pytest.approx(20)
|
||||
assert max(ys) == pytest.approx(60)
|
||||
|
||||
|
||||
def test_svg_curves_and_evenodd_holes_are_preserved():
|
||||
path = "M 0 0 H 1 V 1 H 0 Z M .25 .25 C .4 .1 .6 .1 .75 .25 V .75 H .25 Z"
|
||||
mask, polygon, contours = svg_path_to_mask(
|
||||
path,
|
||||
{"x_min": 0, "y_min": 0, "x_max": 1, "y_max": 1},
|
||||
128,
|
||||
128,
|
||||
supersample=4,
|
||||
precision_px=0.5,
|
||||
)
|
||||
assert len(contours) == 2
|
||||
assert len(polygon) >= 4
|
||||
assert mask[8, 8] > 0.99
|
||||
assert mask[64, 64] < 0.01
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("path", "bbox", "message"),
|
||||
[
|
||||
("", {"x_min": 0, "y_min": 0, "x_max": 1, "y_max": 1}, "empty"),
|
||||
(
|
||||
"M 0 0 L nan 1 Z",
|
||||
{"x_min": 0, "y_min": 0, "x_max": 1, "y_max": 1},
|
||||
"invalid",
|
||||
),
|
||||
(
|
||||
"M 0 0 H 1 V 1 Z",
|
||||
{"x_min": 0.7, "y_min": 0, "x_max": 0.2, "y_max": 1},
|
||||
"positive",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_svg_rejects_malformed_or_unsafe_geometry(path, bbox, message):
|
||||
with pytest.raises((TypeError, ValueError), match=message):
|
||||
svg_path_to_mask(path, bbox, 64, 64)
|
||||
|
||||
|
||||
def test_video_detect_uses_stride_parallelism_and_reports_measured_fps():
|
||||
observed = {}
|
||||
|
||||
def request(operation, **payload):
|
||||
observed["operation"] = operation
|
||||
observed.update(payload)
|
||||
return {
|
||||
"items": [
|
||||
{
|
||||
"objects": [
|
||||
{
|
||||
"x_min": 0.1,
|
||||
"y_min": 0.2,
|
||||
"x_max": 0.4,
|
||||
"y_max": 0.6,
|
||||
}
|
||||
]
|
||||
},
|
||||
{"objects": []},
|
||||
],
|
||||
"elapsed_seconds": 0.1,
|
||||
"parallel_requests": 2,
|
||||
}
|
||||
|
||||
images = torch.zeros((4, 48, 64, 3), dtype=torch.float32)
|
||||
outputs = Moondream31Detect().detect(
|
||||
_fake_model(request),
|
||||
images,
|
||||
"person",
|
||||
30.0,
|
||||
2,
|
||||
2,
|
||||
20,
|
||||
False,
|
||||
)
|
||||
sequence = outputs[0]
|
||||
performance = json.loads(outputs[-1])
|
||||
assert observed["operation"] == "detect"
|
||||
assert len(observed["images"]) == 2
|
||||
assert observed["parallel_requests"] == 2
|
||||
assert sequence.frame_count == 4
|
||||
assert [frame.frame_index for frame in sequence.frames] == [0, 2]
|
||||
assert sequence.frames[0].detections[0].bbox_xyxy == pytest.approx(
|
||||
(6.4, 9.6, 25.6, 28.8)
|
||||
)
|
||||
assert outputs[2].shape == images.shape
|
||||
assert outputs[3].shape == (4, 48, 64)
|
||||
assert performance["processed_frames"] == 2
|
||||
assert performance["worker_fps"] == pytest.approx(20)
|
||||
assert performance["target_processed_fps"] == pytest.approx(15)
|
||||
assert performance["parallel_requests"] == 2
|
||||
|
||||
|
||||
def test_segment_exposes_svg_mask_cutout_overlay_and_structured_detection():
|
||||
def request(operation, **payload):
|
||||
assert operation == "segment"
|
||||
assert payload["spatial_refs"] == [[0.5, 0.5]]
|
||||
return {
|
||||
"items": [
|
||||
{
|
||||
"path": "M 0 0 H 1 V 1 H 0 Z",
|
||||
"bbox": {
|
||||
"x_min": 0.25,
|
||||
"y_min": 0.25,
|
||||
"x_max": 0.75,
|
||||
"y_max": 0.75,
|
||||
},
|
||||
}
|
||||
],
|
||||
"elapsed_seconds": 0.2,
|
||||
"parallel_requests": 1,
|
||||
}
|
||||
|
||||
image = torch.ones((1, 32, 40, 3), dtype=torch.float32)
|
||||
outputs = Moondream31Segment().segment(
|
||||
_fake_model(request, module.PREVIEW_MODEL_ID),
|
||||
image,
|
||||
"object",
|
||||
1.0,
|
||||
1,
|
||||
1,
|
||||
4,
|
||||
False,
|
||||
spatial_refs_json="[[0.5, 0.5]]",
|
||||
)
|
||||
sequence = outputs[0]
|
||||
native = json.loads(outputs[2])
|
||||
mask = outputs[3]
|
||||
mask_image = outputs[4]
|
||||
cutout = outputs[5]
|
||||
overlay = outputs[6]
|
||||
detection = sequence.frames[0].detections[0]
|
||||
assert native[0]["path"].startswith("M 0 0")
|
||||
assert mask.shape == (1, 32, 40)
|
||||
assert mask_image.shape == (1, 32, 40, 3)
|
||||
assert cutout.shape == image.shape
|
||||
assert overlay.shape == image.shape
|
||||
assert mask[0, 16, 20] > 0.99
|
||||
assert mask[0, 2, 2] == 0
|
||||
assert cutout[0, 16, 20].min() > 0.99
|
||||
assert cutout[0, 2, 2].max() == 0
|
||||
assert detection.mask is not None
|
||||
assert detection.polygon is not None
|
||||
assert detection.metadata["native_svg_path"].startswith("M 0 0")
|
||||
|
||||
|
||||
def test_license_gate_and_node_registration():
|
||||
with pytest.raises(ValueError, match="License"):
|
||||
Moondream31Loader().load(
|
||||
False,
|
||||
"Auto",
|
||||
4,
|
||||
"Balanced (8K pages)",
|
||||
)
|
||||
assert set(module.NODE_CLASS_MAPPINGS) == {
|
||||
"Moondream31Loader",
|
||||
"Moondream31Query",
|
||||
"Moondream31Caption",
|
||||
"Moondream31Detect",
|
||||
"Moondream31Point",
|
||||
"Moondream31Segment",
|
||||
}
|
||||
assert all(
|
||||
node.CATEGORY == "VLM Nodes/Moondream 3"
|
||||
for node in module.NODE_CLASS_MAPPINGS.values()
|
||||
)
|
||||
|
||||
|
||||
def test_final_31_model_does_not_claim_preview_svg_segment():
|
||||
with pytest.raises(ValueError, match="3 Preview"):
|
||||
Moondream31Segment().segment(
|
||||
_fake_model(lambda *_args, **_kwargs: {}),
|
||||
torch.zeros((1, 16, 16, 3)),
|
||||
"object",
|
||||
1.0,
|
||||
1,
|
||||
1,
|
||||
1,
|
||||
False,
|
||||
)
|
||||
|
||||
|
||||
def test_worker_auth_is_not_exposed_in_process_arguments_and_logs_are_redacted(
|
||||
tmp_path,
|
||||
monkeypatch,
|
||||
):
|
||||
source = inspect.getsource(Moondream31Model.ensure_started)
|
||||
assert '"--auth-key"' not in source
|
||||
assert "MOONDREAM_WORKER_AUTH" in inspect.getsource(
|
||||
module._worker_environment
|
||||
)
|
||||
log = tmp_path / "worker.log"
|
||||
log.write_text(
|
||||
"api_key=secret-value\nAuthorization: bearer-value\nCUDA error",
|
||||
encoding="utf-8",
|
||||
)
|
||||
tail = module._safe_log_tail(log)
|
||||
assert "secret-value" not in tail
|
||||
assert "bearer-value" not in tail
|
||||
assert "CUDA error" in tail
|
||||
|
||||
monkeypatch.setenv("PATH", "/runtime/bin")
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "must-not-cross")
|
||||
monkeypatch.setenv("HF_TOKEN", "hf-server-side")
|
||||
monkeypatch.setenv("MOONDREAM_API_KEY", "adapter-only")
|
||||
monkeypatch.setenv("HTTPS_PROXY", "https://user:password@example.test")
|
||||
monkeypatch.setenv("PYTORCH_ALLOC_CONF", "backend:cudaMallocAsync")
|
||||
monkeypatch.setenv("PYTORCH_CUDA_ALLOC_CONF", "backend:cudaMallocAsync")
|
||||
base_environment = module._worker_environment(
|
||||
tmp_path,
|
||||
b"\x01" * 32,
|
||||
module.MODEL_ID,
|
||||
)
|
||||
assert base_environment["PATH"] == "/runtime/bin"
|
||||
assert base_environment["HF_TOKEN"] == "hf-server-side"
|
||||
assert "OPENAI_API_KEY" not in base_environment
|
||||
assert "MOONDREAM_API_KEY" not in base_environment
|
||||
assert "HTTPS_PROXY" not in base_environment
|
||||
assert "PYTORCH_ALLOC_CONF" not in base_environment
|
||||
assert "PYTORCH_CUDA_ALLOC_CONF" not in base_environment
|
||||
assert base_environment["MOONDREAM_WORKER_AUTH"] == "01" * 32
|
||||
|
||||
adapter_environment = module._worker_environment(
|
||||
tmp_path,
|
||||
b"\x02" * 32,
|
||||
f"{module.MODEL_ID}/adapter@step",
|
||||
)
|
||||
assert adapter_environment["MOONDREAM_API_KEY"] == "adapter-only"
|
||||
|
||||
|
||||
def test_runtime_python_preserves_virtualenv_symlink(tmp_path, monkeypatch):
|
||||
root = tmp_path / "runtime"
|
||||
binary = tmp_path / "base-python"
|
||||
binary.write_text("", encoding="utf-8")
|
||||
venv_python = root / ".venv" / "bin" / "python"
|
||||
venv_python.parent.mkdir(parents=True)
|
||||
try:
|
||||
venv_python.symlink_to(binary)
|
||||
except OSError:
|
||||
pytest.skip("This filesystem cannot create symlinks.")
|
||||
monkeypatch.delenv("MOONDREAM_PYTHON", raising=False)
|
||||
selected = module._runtime_python(root)
|
||||
assert selected == venv_python.absolute()
|
||||
assert selected != binary.resolve()
|
||||
|
||||
|
||||
def test_worker_registers_official_31_id_only_when_upstream_is_missing(
|
||||
monkeypatch,
|
||||
):
|
||||
@dataclass(frozen=True)
|
||||
class Spec:
|
||||
name: str
|
||||
repo_id: str
|
||||
filename: str
|
||||
checkpoint_format: str
|
||||
|
||||
registry = {
|
||||
"moondream3-preview": Spec(
|
||||
"moondream3-preview",
|
||||
"moondream/moondream3-preview",
|
||||
"model_fp8.pt",
|
||||
"md3",
|
||||
)
|
||||
}
|
||||
fake = types.ModuleType("kestrel.models")
|
||||
fake.get_spec = lambda name: (
|
||||
registry[name] if name in registry else (_ for _ in ()).throw(ValueError(name))
|
||||
)
|
||||
fake.register = lambda spec: registry.__setitem__(spec.name, spec)
|
||||
monkeypatch.setitem(sys.modules, "kestrel.models", fake)
|
||||
|
||||
assert worker._register_moondream31_if_needed("moondream3.1-9B-A2B")
|
||||
registered = registry["moondream3.1-9B-A2B"]
|
||||
assert registered.repo_id == "moondream/moondream3.1-9B-A2B"
|
||||
assert registered.filename == "model.safetensors"
|
||||
assert registered.checkpoint_format == "md3"
|
||||
assert not worker._register_moondream31_if_needed("moondream3.1-9B-A2B")
|
||||
assert not worker._register_moondream31_if_needed("custom-model")
|
||||
|
||||
|
||||
def test_worker_honors_do_not_track_for_base_models(monkeypatch):
|
||||
class SimpleClient:
|
||||
def __init__(self):
|
||||
self.closed = False
|
||||
|
||||
async def aclose(self):
|
||||
self.closed = True
|
||||
|
||||
class Reporter:
|
||||
def __init__(self):
|
||||
self._client = SimpleClient()
|
||||
|
||||
fake = types.ModuleType("kestrel.photon")
|
||||
fake.PhotonReporter = Reporter
|
||||
monkeypatch.setitem(sys.modules, "kestrel.photon", fake)
|
||||
monkeypatch.setenv("DO_NOT_TRACK", "1")
|
||||
monkeypatch.delenv("MOONDREAM_API_KEY", raising=False)
|
||||
|
||||
assert worker._honor_do_not_track()
|
||||
reporter = Reporter()
|
||||
assert asyncio.run(reporter.validate_api_key()) is False
|
||||
assert reporter.start() is None
|
||||
asyncio.run(reporter.shutdown())
|
||||
assert reporter._client.closed
|
||||
|
||||
monkeypatch.setenv("MOONDREAM_API_KEY", "finetune-key")
|
||||
assert not worker._honor_do_not_track()
|
||||
@@ -40,6 +40,7 @@ def test_every_module_imports_and_expected_nodes_exist():
|
||||
assert package.IMPORT_ERRORS == {}
|
||||
expected = {
|
||||
"ModernVLM",
|
||||
"LegacyModernVLM",
|
||||
"VLMRuntimeDiagnostics",
|
||||
"Florence2",
|
||||
"Paligemma",
|
||||
@@ -119,6 +120,10 @@ def test_dependency_metadata_matches_installer_requirements():
|
||||
if line.strip() and not line.lstrip().startswith("#")
|
||||
}
|
||||
assert project_requirements == installer_requirements
|
||||
assert any(
|
||||
Requirement(value).name == "num2words"
|
||||
for value in metadata["project"]["dependencies"]
|
||||
)
|
||||
|
||||
bitsandbytes = next(
|
||||
Requirement(value)
|
||||
@@ -410,6 +415,37 @@ def test_modern_catalog_has_current_quality_and_low_vram_tiers():
|
||||
assert "ibm-granite/granite-vision-4.1-4b" in repositories
|
||||
|
||||
|
||||
def test_modern_picker_is_curated_and_legacy_models_remain_compatible():
|
||||
visible = tuple(modern_vlm.ModernVLM.INPUT_TYPES()["required"]["model"][0])
|
||||
legacy = tuple(
|
||||
modern_vlm.LegacyModernVLM.INPUT_TYPES()["required"]["model"][0]
|
||||
)
|
||||
assert visible == modern_vlm.RECOMMENDED_MODEL_LABELS
|
||||
assert legacy == modern_vlm.LEGACY_MODEL_LABELS
|
||||
assert len(visible) == 12
|
||||
assert set(visible).isdisjoint(legacy)
|
||||
assert set(visible) | set(legacy) == set(modern_vlm.MODEL_CATALOG)
|
||||
assert (
|
||||
modern_vlm.ModernVLM.VALIDATE_INPUTS(
|
||||
"Qwen 2.5 VL 3B Instruct (legacy workflows)"
|
||||
)
|
||||
is True
|
||||
)
|
||||
for node_name in (
|
||||
"Kosmos2model",
|
||||
"MCLLaVAModel",
|
||||
"MiniCPMNode",
|
||||
"MolmoNode",
|
||||
"MoonDream",
|
||||
"Paligemma",
|
||||
"Qwen2VLNode",
|
||||
"UformGen2QwenNode",
|
||||
):
|
||||
assert package.NODE_CLASS_MAPPINGS[node_name].CATEGORY.startswith(
|
||||
"VLM Nodes/Legacy/"
|
||||
)
|
||||
|
||||
|
||||
def test_modern_video_is_primary_input_and_thinking_is_explicit():
|
||||
assert "image" in modern_vlm.ModernVLM.INPUT_TYPES()["optional"]
|
||||
assert "image" in qwen2vl.Qwen2VLNode.INPUT_TYPES()["optional"]
|
||||
@@ -513,6 +549,11 @@ def test_view_text_frontend_rehydrates_and_uses_native_progress_channel():
|
||||
assert 'api.addEventListener("progress_text"' in source
|
||||
assert "onNodeOutputsUpdated(nodeOutputs)" in source
|
||||
assert "connectedViewTextNodes(source)" in source
|
||||
assert '"VLMVideoTemporalReasoner"' in source
|
||||
assert 'makeButton("Save"' in source
|
||||
assert 'makeButton("Wrap: on"' in source
|
||||
assert 'makeButton("Follow: on"' in source
|
||||
assert "isReroute(target)" in source
|
||||
|
||||
|
||||
def test_internvl_video_uses_an_even_vision_patch_grid():
|
||||
|
||||
@@ -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
|
||||
@@ -0,0 +1,251 @@
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from ComfyUI_VLM_nodes.nodes.sam2 import (
|
||||
SAM2_MODELS,
|
||||
Sam2Spec,
|
||||
Sam2VideoPredictor,
|
||||
VLMSAM2VideoSegmentation,
|
||||
_core_box,
|
||||
_normalize_processed_masks,
|
||||
seed_boxes,
|
||||
)
|
||||
from ComfyUI_VLM_nodes.nodes.vision_types import (
|
||||
Detection,
|
||||
DetectionSequence,
|
||||
FrameDetections,
|
||||
)
|
||||
|
||||
|
||||
def test_sam2_catalog_leads_with_tiny_and_has_no_30b_models():
|
||||
assert next(iter(SAM2_MODELS)) == "SAM2.1 Hiera Tiny (fast)"
|
||||
assert all("30b" not in spec.model_id.lower() for spec in SAM2_MODELS.values())
|
||||
assert VLMSAM2VideoSegmentation.RETURN_NAMES[0:2] == ("tracks", "json")
|
||||
|
||||
|
||||
def test_core_box_is_converted_from_xywh():
|
||||
assert _core_box({"x": 2, "y": 3, "width": 5, "height": 7}) == (
|
||||
2.0,
|
||||
3.0,
|
||||
7.0,
|
||||
10.0,
|
||||
)
|
||||
with pytest.raises(ValueError, match="positive"):
|
||||
_core_box({"x": 2, "y": 3, "width": 0, "height": 7})
|
||||
|
||||
|
||||
def test_detection_seeds_keep_labels_and_ids():
|
||||
detection = Detection(
|
||||
bbox_xyxy=(1, 2, 9, 10),
|
||||
label="cat",
|
||||
frame_index=2,
|
||||
timestamp=0.2,
|
||||
track_id=7,
|
||||
)
|
||||
sequence = DetectionSequence(
|
||||
width=10,
|
||||
height=10,
|
||||
frames=(
|
||||
FrameDetections(
|
||||
frame_index=2,
|
||||
timestamp=0.2,
|
||||
width=10,
|
||||
height=10,
|
||||
detections=(detection,),
|
||||
),
|
||||
),
|
||||
frame_count=3,
|
||||
fps=10,
|
||||
)
|
||||
boxes, ids, labels = seed_boxes(
|
||||
width=10,
|
||||
height=10,
|
||||
frame_index=2,
|
||||
detections=sequence,
|
||||
bounding_box=None,
|
||||
)
|
||||
assert boxes == [[1.0, 2.0, 9.0, 10.0]]
|
||||
assert ids == [7]
|
||||
assert labels == {7: "cat"}
|
||||
|
||||
|
||||
def test_processed_mask_shapes_are_normalized():
|
||||
assert _normalize_processed_masks(torch.zeros(2, 1, 4, 5)).shape == (
|
||||
2,
|
||||
4,
|
||||
5,
|
||||
)
|
||||
assert _normalize_processed_masks(torch.zeros(4, 5)).shape == (1, 4, 5)
|
||||
with pytest.raises(RuntimeError, match="unsupported"):
|
||||
_normalize_processed_masks(torch.zeros(1, 2, 3, 4, 5))
|
||||
|
||||
|
||||
def test_video_session_runs_seed_frame_before_both_propagation_directions():
|
||||
class FakeProcessor:
|
||||
def init_video_session(self, **_kwargs):
|
||||
return SimpleNamespace(obj_ids=[1], seed_inferred=False)
|
||||
|
||||
def add_inputs_to_inference_session(self, **kwargs):
|
||||
kwargs["inference_session"].seed_frame = kwargs["frame_idx"]
|
||||
|
||||
def post_process_masks(self, masks, **_kwargs):
|
||||
return masks
|
||||
|
||||
class FakeModel(torch.nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.anchor = torch.nn.Parameter(torch.zeros(()))
|
||||
self.seed_calls = []
|
||||
self.propagation_directions = []
|
||||
|
||||
def forward(self, inference_session, frame_idx):
|
||||
inference_session.seed_inferred = True
|
||||
self.seed_calls.append(frame_idx)
|
||||
return SimpleNamespace(
|
||||
frame_idx=frame_idx,
|
||||
pred_masks=torch.ones(1, 1, 4, 4),
|
||||
)
|
||||
|
||||
def propagate_in_video_iterator(
|
||||
self,
|
||||
inference_session,
|
||||
start_frame_idx,
|
||||
reverse=False,
|
||||
**_kwargs,
|
||||
):
|
||||
assert inference_session.seed_inferred
|
||||
self.propagation_directions.append(reverse)
|
||||
indices = (
|
||||
range(start_frame_idx, 3)
|
||||
if not reverse
|
||||
else range(start_frame_idx, -1, -1)
|
||||
)
|
||||
for frame_index in indices:
|
||||
yield SimpleNamespace(
|
||||
frame_idx=frame_index,
|
||||
pred_masks=torch.ones(1, 1, 4, 4),
|
||||
)
|
||||
|
||||
model = FakeModel()
|
||||
predictor = Sam2VideoPredictor.__new__(Sam2VideoPredictor)
|
||||
predictor.processor = FakeProcessor()
|
||||
predictor.dtype = torch.float32
|
||||
predictor.spec = Sam2Spec("test/sam2", "test-sam2")
|
||||
predictor.handle = SimpleNamespace(ensure_loaded=lambda: model)
|
||||
detections = DetectionSequence(
|
||||
width=4,
|
||||
height=4,
|
||||
frames=(
|
||||
FrameDetections(
|
||||
frame_index=0,
|
||||
timestamp=0.0,
|
||||
width=4,
|
||||
height=4,
|
||||
detections=(
|
||||
Detection(
|
||||
bbox_xyxy=(0, 0, 4, 4),
|
||||
label="object",
|
||||
frame_index=0,
|
||||
timestamp=0.0,
|
||||
),
|
||||
),
|
||||
),
|
||||
),
|
||||
frame_count=1,
|
||||
)
|
||||
tracks, union, individual, preview = predictor.propagate(
|
||||
torch.zeros(3, 4, 4, 3),
|
||||
seed_frame=1,
|
||||
fps=10.0,
|
||||
detections=detections,
|
||||
bounding_box=None,
|
||||
seed_mask=None,
|
||||
mask_threshold=0.0,
|
||||
keep_video_on_cpu=True,
|
||||
mask_output="union_and_objects",
|
||||
render_preview=True,
|
||||
)
|
||||
assert model.seed_calls == [1]
|
||||
assert model.propagation_directions == [False, True]
|
||||
assert [item.frame_index for item in tracks.tracks[0].detections] == [0, 1, 2]
|
||||
assert union.shape == (3, 4, 4)
|
||||
assert individual.shape == (3, 4, 4)
|
||||
assert preview.shape == (3, 4, 4, 3)
|
||||
assert torch.equal(preview[0], preview[1])
|
||||
assert torch.equal(preview[1], preview[2])
|
||||
|
||||
|
||||
def test_multi_object_mask_seeds_are_passed_as_one_mask_per_object():
|
||||
class FakeProcessor:
|
||||
def __init__(self):
|
||||
self.received_masks = None
|
||||
|
||||
def init_video_session(self, **kwargs):
|
||||
assert str(kwargs["processing_device"]) == "cpu"
|
||||
return SimpleNamespace(obj_ids=[1, 2], seed_inferred=False)
|
||||
|
||||
def add_inputs_to_inference_session(self, **kwargs):
|
||||
self.received_masks = kwargs["input_masks"]
|
||||
|
||||
def post_process_masks(self, masks, **_kwargs):
|
||||
return masks
|
||||
|
||||
class FakeModel(torch.nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.anchor = torch.nn.Parameter(torch.zeros(()))
|
||||
|
||||
def forward(self, inference_session, frame_idx):
|
||||
inference_session.seed_inferred = True
|
||||
return SimpleNamespace(
|
||||
frame_idx=frame_idx,
|
||||
pred_masks=torch.ones(2, 1, 4, 4),
|
||||
)
|
||||
|
||||
def propagate_in_video_iterator(self, **_kwargs):
|
||||
return iter(())
|
||||
|
||||
processor = FakeProcessor()
|
||||
predictor = Sam2VideoPredictor.__new__(Sam2VideoPredictor)
|
||||
predictor.processor = processor
|
||||
predictor.dtype = torch.float32
|
||||
predictor.spec = Sam2Spec("test/sam2", "test-sam2")
|
||||
predictor.handle = SimpleNamespace(ensure_loaded=FakeModel)
|
||||
images = torch.zeros(1, 4, 4, 3)
|
||||
tracks, _union, individual, preview = predictor.propagate(
|
||||
images,
|
||||
seed_frame=0,
|
||||
fps=24.0,
|
||||
detections=None,
|
||||
bounding_box=None,
|
||||
seed_mask=torch.ones(2, 4, 4),
|
||||
mask_threshold=0.0,
|
||||
keep_video_on_cpu=True,
|
||||
mask_output="union_only",
|
||||
render_preview=False,
|
||||
)
|
||||
assert isinstance(processor.received_masks, list)
|
||||
assert len(processor.received_masks) == 2
|
||||
assert len(tracks.tracks) == 2
|
||||
assert individual.shape == (0, 4, 4)
|
||||
assert preview.data_ptr() == images.data_ptr()
|
||||
|
||||
|
||||
def test_nested_core_boxes_select_seed_frame_and_keep_top_level_labels():
|
||||
boxes, ids, labels = seed_boxes(
|
||||
width=20,
|
||||
height=20,
|
||||
frame_index=1,
|
||||
detections=None,
|
||||
bounding_box=[
|
||||
[{"x": 0, "y": 0, "width": 2, "height": 2, "label": "old"}],
|
||||
[
|
||||
{"x": 3, "y": 4, "width": 5, "height": 6, "label": "person"},
|
||||
{"x": 10, "y": 11, "width": 4, "height": 3},
|
||||
],
|
||||
],
|
||||
)
|
||||
assert boxes == [[3.0, 4.0, 8.0, 10.0], [10.0, 11.0, 14.0, 14.0]]
|
||||
assert ids == [1, 2]
|
||||
assert labels == {1: "person", 2: None}
|
||||
@@ -0,0 +1,213 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from ComfyUI_VLM_nodes.nodes.sam3_adapter import (
|
||||
VLMTrackReport,
|
||||
iter_sam3_masks,
|
||||
sam3_track_data_to_tracks,
|
||||
track_report_json,
|
||||
track_report_payload,
|
||||
track_report_text,
|
||||
unpack_sam3_mask,
|
||||
validate_sam3_track_data,
|
||||
)
|
||||
from ComfyUI_VLM_nodes.nodes.vision_types import (
|
||||
Detection,
|
||||
DetectionSequence,
|
||||
FrameDetections,
|
||||
)
|
||||
|
||||
|
||||
def _pack_masks(masks):
|
||||
masks = masks.to(torch.uint8)
|
||||
width = masks.shape[-1]
|
||||
assert width % 8 == 0
|
||||
bits = 1 << torch.arange(8, dtype=torch.int64)
|
||||
grouped = masks.reshape(*masks.shape[:-1], width // 8, 8)
|
||||
return (grouped * bits).sum(dim=-1).to(torch.uint8)
|
||||
|
||||
|
||||
def _sample_track_data():
|
||||
masks = torch.zeros(2, 2, 4, 8, dtype=torch.bool)
|
||||
masks[0, 0, 1:3, 2:5] = True
|
||||
masks[1, 0, 1:4, 3:6] = True
|
||||
masks[1, 1, 0:2, 0:2] = True
|
||||
return {
|
||||
"packed_masks": _pack_masks(masks),
|
||||
"n_frames": 2,
|
||||
"scores": [0.9, 0.75],
|
||||
"orig_size": (40, 80),
|
||||
}
|
||||
|
||||
|
||||
def _seed_detections():
|
||||
detections = (
|
||||
Detection(
|
||||
bbox_xyxy=(20, 10, 50, 30),
|
||||
label="cat",
|
||||
text="the cat",
|
||||
score=0.95,
|
||||
frame_index=0,
|
||||
timestamp=0.0,
|
||||
track_id=7,
|
||||
source="seed",
|
||||
),
|
||||
Detection(
|
||||
bbox_xyxy=(0, 0, 20, 20),
|
||||
label="fish",
|
||||
score=0.8,
|
||||
frame_index=0,
|
||||
timestamp=0.0,
|
||||
track_id=9,
|
||||
source="seed",
|
||||
),
|
||||
)
|
||||
return DetectionSequence(
|
||||
width=80,
|
||||
height=40,
|
||||
frames=(
|
||||
FrameDetections(
|
||||
frame_index=0,
|
||||
timestamp=0.0,
|
||||
width=80,
|
||||
height=40,
|
||||
detections=detections,
|
||||
),
|
||||
),
|
||||
frame_count=2,
|
||||
fps=10.0,
|
||||
)
|
||||
|
||||
|
||||
def test_validate_and_stream_unpack_without_expanding_the_video():
|
||||
track_data = _sample_track_data()
|
||||
layout = validate_sam3_track_data(track_data)
|
||||
assert (
|
||||
layout.n_frames,
|
||||
layout.n_objects,
|
||||
layout.mask_height,
|
||||
layout.mask_width,
|
||||
) == (2, 2, 4, 8)
|
||||
yielded = list(iter_sam3_masks(track_data, present_only=True))
|
||||
assert [(frame, obj) for frame, obj, _mask in yielded] == [
|
||||
(0, 0),
|
||||
(1, 0),
|
||||
(1, 1),
|
||||
]
|
||||
assert all(mask.shape == (4, 8) for _frame, _obj, mask in yielded)
|
||||
assert yielded[0][2].sum().item() == 6
|
||||
one = unpack_sam3_mask(track_data["packed_masks"][1, 1])
|
||||
assert one.dtype == torch.bool
|
||||
assert one.sum().item() == 4
|
||||
|
||||
|
||||
def test_adapter_derives_scaled_boxes_and_preserves_seed_identity():
|
||||
tracks = sam3_track_data_to_tracks(
|
||||
_sample_track_data(),
|
||||
seed_detections=_seed_detections(),
|
||||
fps=10.0,
|
||||
)
|
||||
assert (tracks.width, tracks.height, tracks.frame_count, tracks.fps) == (
|
||||
80,
|
||||
40,
|
||||
2,
|
||||
10.0,
|
||||
)
|
||||
assert [track.track_id for track in tracks.tracks] == [7, 9]
|
||||
assert [track.label for track in tracks.tracks] == ["cat", "fish"]
|
||||
cat, fish = tracks.tracks
|
||||
assert cat.detections[0].bbox_xyxy == (20.0, 10.0, 50.0, 30.0)
|
||||
assert cat.detections[1].bbox_xyxy == (30.0, 10.0, 60.0, 40.0)
|
||||
assert fish.detections[0].frame_index == 1
|
||||
assert fish.detections[0].bbox_xyxy == (0.0, 0.0, 20.0, 20.0)
|
||||
assert cat.detections[0].metadata["mask_ref"]["object_index"] == 0
|
||||
assert fish.detections[0].metadata["mask_ref"]["object_index"] == 1
|
||||
assert cat.metadata["seeded"] is True
|
||||
assert fish.metadata["seeded"] is True
|
||||
|
||||
serialized = tracks.to_json()
|
||||
assert "mask_ref" in serialized
|
||||
assert "packed_masks" not in serialized
|
||||
assert "tensor" not in serialized.casefold()
|
||||
|
||||
|
||||
def test_adapter_handles_an_empty_core_result():
|
||||
tracks = sam3_track_data_to_tracks(
|
||||
{
|
||||
"packed_masks": None,
|
||||
"n_frames": 3,
|
||||
"scores": [],
|
||||
"orig_size": (48, 64),
|
||||
},
|
||||
fps=24.0,
|
||||
)
|
||||
assert tracks.tracks == ()
|
||||
assert tracks.frame_count == 3
|
||||
assert tracks.metadata["object_slots"] == 0
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("track_data", "message"),
|
||||
(
|
||||
(
|
||||
{"packed_masks": None, "n_frames": 1, "scores": []},
|
||||
"orig_size",
|
||||
),
|
||||
(
|
||||
{
|
||||
"packed_masks": torch.zeros(1, 1, 2, 1),
|
||||
"n_frames": 1,
|
||||
"scores": [0.5],
|
||||
"orig_size": (2, 8),
|
||||
},
|
||||
"uint8",
|
||||
),
|
||||
(
|
||||
{
|
||||
"packed_masks": torch.zeros(2, 1, 2, 1, dtype=torch.uint8),
|
||||
"n_frames": 1,
|
||||
"scores": [0.5],
|
||||
"orig_size": (2, 8),
|
||||
},
|
||||
"n_frames",
|
||||
),
|
||||
(
|
||||
{
|
||||
"packed_masks": torch.zeros(1, 1, 2, 1, dtype=torch.uint8),
|
||||
"n_frames": 1,
|
||||
"scores": [1.5],
|
||||
"orig_size": (2, 8),
|
||||
},
|
||||
"0 to 1",
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_adapter_rejects_incompatible_private_payloads(track_data, message):
|
||||
with pytest.raises((TypeError, ValueError), match=message):
|
||||
validate_sam3_track_data(track_data)
|
||||
|
||||
|
||||
def test_track_report_is_small_deterministic_and_history_safe():
|
||||
tracks = sam3_track_data_to_tracks(
|
||||
_sample_track_data(),
|
||||
seed_detections=_seed_detections(),
|
||||
fps=10.0,
|
||||
)
|
||||
payload = track_report_payload(tracks)
|
||||
assert payload["track_count"] == 2
|
||||
assert payload["observation_count"] == 3
|
||||
assert payload["state_counts"] == {"active": 2}
|
||||
encoded = track_report_json(tracks)
|
||||
assert json.loads(encoded) == payload
|
||||
assert "packed_masks" not in encoded
|
||||
text = track_report_text(tracks)
|
||||
assert "Tracks: 2" in text
|
||||
assert "#7 cat" in text
|
||||
assert "#9 fish" in text
|
||||
|
||||
node_result = VLMTrackReport().report(tracks)
|
||||
assert node_result["result"] == (encoded, text)
|
||||
assert node_result["ui"]["text"] == [text]
|
||||
@@ -0,0 +1,211 @@
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import ComfyUI_VLM_nodes as package
|
||||
import pytest
|
||||
from ComfyUI_VLM_nodes.nodes import simpletext
|
||||
|
||||
|
||||
def test_simple_text_preserves_legacy_default_and_appends_metrics():
|
||||
result = simpletext.SimpleText().simple_text(" one\r\ntwo ")
|
||||
assert result == (" one\r\ntwo ", 12, 2, 2)
|
||||
|
||||
normalized = simpletext.SimpleText().simple_text(
|
||||
" one\r\ntwo ",
|
||||
trim_edges=True,
|
||||
normalize_newlines=True,
|
||||
)
|
||||
assert normalized == ("one\ntwo", 7, 2, 2)
|
||||
|
||||
|
||||
def test_json_to_text_keeps_legacy_smart_rendering_and_adds_canonical_output():
|
||||
response = simpletext.JsonToText().json_to_text(
|
||||
'{"prompt":"Create a red kite","suggestion1":"at sunset","tags":["red","sky"]}'
|
||||
)
|
||||
assert response["result"][0] == "a red kite\n\nat sunset\n\ntags: red, sky"
|
||||
assert json.loads(response["result"][1])["tags"] == ["red", "sky"]
|
||||
assert response["result"][2] == 3
|
||||
|
||||
|
||||
def test_json_to_text_parses_fenced_model_response_and_json_paths():
|
||||
response = simpletext.JsonToText().json_to_text(
|
||||
'Model response:\n```json\n{"result":{"items":[{"name":"café"}]}}\n```',
|
||||
format_mode="Pretty JSON",
|
||||
json_path="$.result.items[0]",
|
||||
)
|
||||
assert json.loads(response["result"][0]) == {"name": "café"}
|
||||
assert response["result"][1] == '{"name":"café"}'
|
||||
|
||||
pointer = simpletext.VLMJSONExtract().extract(
|
||||
'{"a/b":{"~key":[10,20]}}',
|
||||
"/a~1b/~0key/1",
|
||||
"Text",
|
||||
"Error",
|
||||
"",
|
||||
)
|
||||
assert pointer == ("20", True, "integer", "20")
|
||||
|
||||
|
||||
def test_json_extract_handles_negative_indexes_and_missing_policy():
|
||||
node = simpletext.VLMJSONExtract()
|
||||
assert node.extract(
|
||||
'{"items":["first","last"]}',
|
||||
"$.items[-1]",
|
||||
"Text",
|
||||
"Error",
|
||||
"",
|
||||
)[:3] == ("last", True, "string")
|
||||
assert node.extract(
|
||||
'{"items":[]}',
|
||||
"$.missing",
|
||||
"Text",
|
||||
"Default value",
|
||||
"fallback",
|
||||
)[:3] == ("fallback", False, "string")
|
||||
with pytest.raises(ValueError, match="not found"):
|
||||
node.extract("{}", "$.missing", "Text", "Error", "")
|
||||
|
||||
|
||||
def test_text_join_drops_empty_and_duplicate_parts():
|
||||
result = simpletext.VLMTextJoin().join(
|
||||
" first ",
|
||||
"Blank line",
|
||||
"|",
|
||||
True,
|
||||
True,
|
||||
True,
|
||||
text_b="second",
|
||||
text_c="first",
|
||||
)
|
||||
assert result == (
|
||||
"first\n\nsecond",
|
||||
'["first","second"]',
|
||||
2,
|
||||
)
|
||||
|
||||
|
||||
def test_text_template_is_safe_explicit_and_supports_literal_braces():
|
||||
result = simpletext.VLMTextTemplate().render(
|
||||
"{{schema}} {subject}: {text1}",
|
||||
'{"subject":"robot"}',
|
||||
"Error",
|
||||
text1="moving a box",
|
||||
)
|
||||
assert result[0] == "{schema} robot: moving a box"
|
||||
assert json.loads(result[1]) == {
|
||||
"subject": "robot",
|
||||
"text1": "moving a box",
|
||||
}
|
||||
assert result[2] == "[]"
|
||||
|
||||
with pytest.raises(ValueError, match="missing"):
|
||||
simpletext.VLMTextTemplate().render(
|
||||
"{known} {unknown}",
|
||||
'{"known":"yes"}',
|
||||
"Error",
|
||||
)
|
||||
|
||||
|
||||
def test_text_clean_normalizes_fences_duplicates_and_length():
|
||||
result, diagnostics_json = simpletext.VLMTextClean().clean(
|
||||
"```text\r\nA B\r\nA B\r\nC\r\n```",
|
||||
"NFKC",
|
||||
"Collapse horizontal",
|
||||
True,
|
||||
True,
|
||||
True,
|
||||
5,
|
||||
)
|
||||
assert result == "A B\nC"
|
||||
diagnostics = json.loads(diagnostics_json)
|
||||
assert diagnostics["changed"] is True
|
||||
assert diagnostics["duplicate_lines_removed"] == 1
|
||||
assert diagnostics["truncated"] is False
|
||||
|
||||
|
||||
def test_text_replace_literal_regex_and_errors():
|
||||
node = simpletext.VLMTextReplace()
|
||||
assert node.replace(
|
||||
"Cat cat cat",
|
||||
"cat",
|
||||
"dog",
|
||||
"Literal",
|
||||
False,
|
||||
2,
|
||||
"Keep text",
|
||||
)[:2] == ("dog dog cat", 2)
|
||||
assert node.replace(
|
||||
"a1 b22",
|
||||
r"\d+",
|
||||
"#",
|
||||
"Regular expression",
|
||||
True,
|
||||
0,
|
||||
"Keep text",
|
||||
)[:2] == ("a# b#", 2)
|
||||
with pytest.raises(ValueError, match="not found"):
|
||||
node.replace("hello", "x", "y", "Literal", True, 0, "Error")
|
||||
|
||||
|
||||
def test_text_split_outputs_real_list_and_stable_json():
|
||||
items, items_json, count = simpletext.VLMTextSplit().split(
|
||||
'[" first ","second","first",""]',
|
||||
"JSON array",
|
||||
",",
|
||||
True,
|
||||
True,
|
||||
True,
|
||||
0,
|
||||
)
|
||||
assert items == ["first", "second"]
|
||||
assert json.loads(items_json) == items
|
||||
assert count == 2
|
||||
assert simpletext.VLMTextSplit.OUTPUT_IS_LIST == (True, False, False)
|
||||
|
||||
|
||||
def test_text_inspector_and_view_text_report_same_metrics():
|
||||
inspected = simpletext.VLMTextInspect().inspect("hello\nworld")
|
||||
assert inspected[1:6] == (11, 11, 2, 2, 3)
|
||||
assert len(inspected[6]) == 64
|
||||
assert json.loads(inspected[7])["words"] == 2
|
||||
|
||||
viewed = simpletext.ViewText().view_text("hello\nworld")
|
||||
assert viewed["result"][:4] == ("hello\nworld", 11, 2, 2)
|
||||
assert viewed["ui"]["text"] == ["hello\nworld"]
|
||||
|
||||
|
||||
def test_text_node_categories_aliases_and_legacy_ids_are_stable():
|
||||
assert simpletext.NODE_CLASS_MAPPINGS["SimpleText"] is simpletext.SimpleText
|
||||
assert simpletext.NODE_CLASS_MAPPINGS["JsonToText"] is simpletext.JsonToText
|
||||
assert simpletext.NODE_CLASS_MAPPINGS["ViewText"] is simpletext.ViewText
|
||||
assert set(simpletext.NODE_CLASS_MAPPINGS) == {
|
||||
"SimpleText",
|
||||
"JsonToText",
|
||||
"ViewText",
|
||||
"VLMTextJoin",
|
||||
"VLMTextTemplate",
|
||||
"VLMTextClean",
|
||||
"VLMTextReplace",
|
||||
"VLMJSONExtract",
|
||||
"VLMTextSplit",
|
||||
"VLMTextInspect",
|
||||
}
|
||||
assert simpletext.SimpleText.CATEGORY == "VLM Nodes/Text/Create"
|
||||
assert simpletext.ViewText.CATEGORY == "VLM Nodes/Text/Inspect"
|
||||
|
||||
|
||||
def test_text_toolkit_api_example_uses_registered_inputs_and_output_indexes():
|
||||
root = Path(package.__file__).parent
|
||||
prompt = json.loads(
|
||||
(root / "examples" / "text_toolkit_api.json").read_text("utf-8")
|
||||
)
|
||||
assert prompt["5"]["inputs"]["text"] == ["4", 0]
|
||||
for node in prompt.values():
|
||||
node_class = package.NODE_CLASS_MAPPINGS[node["class_type"]]
|
||||
declared = {
|
||||
name
|
||||
for group in node_class.INPUT_TYPES().values()
|
||||
if isinstance(group, dict)
|
||||
for name in group
|
||||
}
|
||||
assert set(node["inputs"]) <= declared
|
||||
@@ -0,0 +1,376 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
import pytest
|
||||
from ComfyUI_VLM_nodes.nodes.spatial_parser import (
|
||||
COORDINATE_MODES,
|
||||
VLMSpatialPromptBuilder,
|
||||
VLMStructuredSpatialParser,
|
||||
build_spatial_prompt,
|
||||
load_json_document,
|
||||
parse_spatial_response,
|
||||
)
|
||||
from ComfyUI_VLM_nodes.nodes.vision_types import (
|
||||
VLM_DETECTIONS,
|
||||
VLM_POINTS,
|
||||
DetectionSequence,
|
||||
PointSequence,
|
||||
)
|
||||
|
||||
|
||||
def test_prompt_builder_is_explicit_and_provider_neutral():
|
||||
prompt = build_spatial_prompt(
|
||||
"Find every vehicle.",
|
||||
coordinate_mode="normalized_0_1000",
|
||||
width=1920,
|
||||
height=1080,
|
||||
frame_count=12,
|
||||
fps=24.0,
|
||||
)
|
||||
|
||||
assert prompt.startswith("Perform this visual analysis task:")
|
||||
assert "Find every vehicle." in prompt
|
||||
assert "Return only one valid JSON object" in prompt
|
||||
assert "normalized_0_1000" in prompt
|
||||
assert '"frame_count":12' in prompt
|
||||
assert '"fps":24.0' in prompt
|
||||
assert "zero-based frame_index" in prompt
|
||||
assert "bbox_xyxy" in prompt
|
||||
assert "polygon" in prompt
|
||||
assert '"point"' in prompt
|
||||
assert "score" in prompt
|
||||
example = json.loads(prompt.split("Required JSON shape:\n", 1)[1])
|
||||
assert example["coordinate_mode"] == "normalized_0_1000"
|
||||
|
||||
|
||||
def test_node_contracts_use_canonical_spatial_types():
|
||||
assert tuple(COORDINATE_MODES) == (
|
||||
"pixel",
|
||||
"normalized_0_1",
|
||||
"normalized_0_1000",
|
||||
)
|
||||
assert VLMSpatialPromptBuilder.RETURN_TYPES == ("STRING",)
|
||||
assert VLMStructuredSpatialParser.RETURN_TYPES == (
|
||||
VLM_DETECTIONS,
|
||||
VLM_POINTS,
|
||||
"STRING",
|
||||
)
|
||||
schema = VLMStructuredSpatialParser.INPUT_TYPES()
|
||||
assert tuple(schema["required"]["coordinate_mode"][0]) == COORDINATE_MODES
|
||||
|
||||
|
||||
def test_parser_accepts_only_complete_plain_or_fenced_json():
|
||||
assert load_json_document(" ") == {}
|
||||
assert load_json_document('```json\n{"frames":[]}\n```') == {"frames": []}
|
||||
|
||||
for invalid in (
|
||||
'Here is the result: {"frames":[]}',
|
||||
'Result:\n```json\n{"frames":[]}\n```',
|
||||
'{"frames":[]} trailing',
|
||||
"```python\n{}\n```",
|
||||
'{"x": 1, "x": 2}',
|
||||
'{"score": NaN}',
|
||||
):
|
||||
with pytest.raises(ValueError):
|
||||
load_json_document(invalid)
|
||||
|
||||
|
||||
def test_normalized_video_parse_clips_and_preserves_metadata():
|
||||
response = json.dumps(
|
||||
{
|
||||
"coordinate_mode": "normalized_0_1",
|
||||
"media": {
|
||||
"width": 200,
|
||||
"height": 100,
|
||||
"frame_count": 3,
|
||||
"fps": 2,
|
||||
"codec": "test-codec",
|
||||
},
|
||||
"source": "unit-vlm",
|
||||
"metadata": {"request_id": "abc"},
|
||||
"vendor": {"latency_ms": 12},
|
||||
"frames": [
|
||||
{
|
||||
"frame_index": 0,
|
||||
"metadata": {"scene": "start"},
|
||||
"detections": [
|
||||
{
|
||||
"class": "cat",
|
||||
"confidence": 0.75,
|
||||
"bbox": [-0.1, 0.2, 1.2, 0.8],
|
||||
"polygon": [
|
||||
[-0.5, 0.2],
|
||||
[0.5, 0.2],
|
||||
[0.5, 1.5],
|
||||
],
|
||||
"instance_id": "cat-1",
|
||||
"metadata": {"occluded": False},
|
||||
}
|
||||
],
|
||||
"points": [
|
||||
{
|
||||
"name": "nose",
|
||||
"point": [0.25, 0.5],
|
||||
"confidence": 0.9,
|
||||
"landmark_id": 7,
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"frame_index": 2,
|
||||
"detections": [
|
||||
{
|
||||
"label": "sign",
|
||||
"quad": [
|
||||
[0.1, 0.1],
|
||||
[0.9, 0.1],
|
||||
[0.9, 0.9],
|
||||
[0.1, 0.9],
|
||||
],
|
||||
"text": "STOP",
|
||||
}
|
||||
],
|
||||
},
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
detections, points, normalized_json = parse_spatial_response(
|
||||
f"```json\n{response}\n```",
|
||||
width=200,
|
||||
height=100,
|
||||
coordinate_mode="normalized_0_1",
|
||||
)
|
||||
|
||||
assert isinstance(detections, DetectionSequence)
|
||||
assert isinstance(points, PointSequence)
|
||||
assert detections.frame_count == points.frame_count == 3
|
||||
assert detections.fps == points.fps == 2.0
|
||||
assert [frame.frame_index for frame in detections.frames] == [0, 2]
|
||||
cat, sign = detections.all_detections()
|
||||
assert cat.bbox_xyxy == (0.0, 20.0, 200.0, 80.0)
|
||||
assert cat.polygon == ((0.0, 20.0), (100.0, 20.0), (100.0, 100.0))
|
||||
assert cat.label == "cat"
|
||||
assert cat.score == 0.75
|
||||
assert cat.metadata.to_dict() == {
|
||||
"instance_id": "cat-1",
|
||||
"occluded": False,
|
||||
}
|
||||
assert sign.bbox_xyxy == (20.0, 10.0, 180.0, 90.0)
|
||||
assert sign.quad is not None and len(sign.quad) == 4
|
||||
assert sign.text == "STOP"
|
||||
assert points.points[0].x == 50.0
|
||||
assert points.points[0].y == 50.0
|
||||
assert points.points[0].label == "nose"
|
||||
assert points.points[0].metadata["landmark_id"] == 7
|
||||
assert detections.frames[0].metadata["scene"] == "start"
|
||||
assert detections.metadata.to_dict() == {
|
||||
"coordinate_mode": "normalized_0_1",
|
||||
"media_metadata": {"codec": "test-codec"},
|
||||
"request_id": "abc",
|
||||
"vendor": {"latency_ms": 12},
|
||||
}
|
||||
|
||||
normalized = json.loads(normalized_json)
|
||||
assert normalized["schema"] == "comfyui-vlm/spatial"
|
||||
assert normalized["detections"] == detections.to_dict()
|
||||
assert normalized["points"] == points.to_dict()
|
||||
|
||||
|
||||
def test_pixel_aliases_xywh_flat_segmentation_and_multiple_points():
|
||||
response = json.dumps(
|
||||
{
|
||||
"objects": [
|
||||
{
|
||||
"name": "panel",
|
||||
"box": {"x": -2, "y": 5, "width": 15, "height": 30},
|
||||
"segmentation": [0, 5, 13, 5, 13, 35, 0, 35],
|
||||
"points": [[2, 7], [12, 30]],
|
||||
"score": 1,
|
||||
}
|
||||
]
|
||||
}
|
||||
)
|
||||
detections, points, _json = parse_spatial_response(
|
||||
response,
|
||||
width=10,
|
||||
height=20,
|
||||
coordinate_mode="pixel",
|
||||
frame_count=1,
|
||||
)
|
||||
|
||||
detection = detections.all_detections()[0]
|
||||
assert detection.bbox_xyxy == (0.0, 5.0, 10.0, 20.0)
|
||||
assert detection.polygon == (
|
||||
(0.0, 5.0),
|
||||
(10.0, 5.0),
|
||||
(10.0, 20.0),
|
||||
(0.0, 20.0),
|
||||
)
|
||||
assert [(point.x, point.y) for point in points.points] == [
|
||||
(2.0, 7.0),
|
||||
(10.0, 20.0),
|
||||
]
|
||||
|
||||
|
||||
def test_direct_multi_point_record_keeps_shared_fields_without_duplicates():
|
||||
detections, points, _json = parse_spatial_response(
|
||||
'{"label":"hand","confidence":0.8,"points":[[1,2],[3,4]],'
|
||||
'"metadata":{"side":"left"}}',
|
||||
width=10,
|
||||
height=10,
|
||||
)
|
||||
|
||||
assert detections.all_detections() == ()
|
||||
assert [(point.x, point.y) for point in points.points] == [
|
||||
(1.0, 2.0),
|
||||
(3.0, 4.0),
|
||||
]
|
||||
assert {point.label for point in points.points} == {"hand"}
|
||||
assert {point.score for point in points.points} == {0.8}
|
||||
assert points.points[0].metadata["side"] == "left"
|
||||
assert detections.frames[0].metadata.to_dict() == {}
|
||||
|
||||
|
||||
def test_top_level_record_batch_groups_video_frame_indices_and_timestamps():
|
||||
response = json.dumps(
|
||||
[
|
||||
{"frame_index": 2, "bbox_xyxy": [1, 2, 3, 4], "label": "late"},
|
||||
{"frame_index": 0, "point": [5, 6], "label": "early"},
|
||||
{"frame_index": 0, "box": [0, 0, 4, 5], "label": "first"},
|
||||
]
|
||||
)
|
||||
detections, points, _json = parse_spatial_response(
|
||||
response,
|
||||
width=20,
|
||||
height=10,
|
||||
fps=4,
|
||||
coordinate_mode="pixel",
|
||||
)
|
||||
|
||||
assert [frame.frame_index for frame in detections.frames] == [0, 2]
|
||||
assert detections.frames[1].timestamp == 0.5
|
||||
assert detections.frame_count == 3
|
||||
assert [item.label for item in detections.all_detections()] == [
|
||||
"first",
|
||||
"late",
|
||||
]
|
||||
assert points.points[0].frame_index == 0
|
||||
|
||||
|
||||
def test_normalized_1000_polygon_without_box_derives_clipped_bbox():
|
||||
detections, points, _json = parse_spatial_response(
|
||||
'{"polygon":[[-100,100],[500,100],[1200,900]],"label":"shape"}',
|
||||
width=300,
|
||||
height=200,
|
||||
coordinate_mode="normalized_0_1000",
|
||||
)
|
||||
|
||||
detection = detections.all_detections()[0]
|
||||
assert detection.bbox_xyxy == (0.0, 20.0, 300.0, 180.0)
|
||||
assert not points.points
|
||||
|
||||
|
||||
def test_empty_payloads_are_valid_and_predictable():
|
||||
for response in ("", "{}", "[]", '{"frames":[]}'):
|
||||
detections, points, normalized_json = parse_spatial_response(
|
||||
response,
|
||||
width=640,
|
||||
height=480,
|
||||
coordinate_mode="pixel",
|
||||
frame_count=5,
|
||||
fps=25,
|
||||
source="empty-test",
|
||||
)
|
||||
assert detections.frames == ()
|
||||
assert detections.frame_count == 5
|
||||
assert detections.fps == 25
|
||||
assert points.points == ()
|
||||
assert points.frame_count == 5
|
||||
normalized = json.loads(normalized_json)
|
||||
assert normalized["detections"]["frames"] == []
|
||||
assert normalized["points"]["points"] == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("response", "message"),
|
||||
[
|
||||
('{"bbox":[4,4,2,5]}', "x2 >= x1"),
|
||||
('{"polygon":[[1,1],[2,2]]}', "at least three"),
|
||||
('{"quad":[[0,0],[1,0],[1,1]]}', "exactly 4"),
|
||||
('{"point":[1]}', "two coordinates"),
|
||||
('{"score":2,"point":[1,1]}', "between 0 and 1"),
|
||||
(
|
||||
'{"bbox":[0,0,1,1],"polygon":[[0,0],[1,0],[0,1]],'
|
||||
'"quad":[[0,0],[1,0],[1,1],[0,1]]}',
|
||||
"either polygon",
|
||||
),
|
||||
(
|
||||
'{"bbox":[0,0,1,1],"coordinate_mode":"normalized_0_1"}',
|
||||
"parser is set",
|
||||
),
|
||||
('{"label":"no geometry"}', "requires bbox"),
|
||||
],
|
||||
)
|
||||
def test_strict_validation(response, message):
|
||||
with pytest.raises((TypeError, ValueError), match=message):
|
||||
parse_spatial_response(
|
||||
response,
|
||||
width=10,
|
||||
height=10,
|
||||
coordinate_mode="pixel",
|
||||
)
|
||||
|
||||
|
||||
def test_dimensions_timing_and_alias_conflicts_are_rejected():
|
||||
cases = [
|
||||
('{"media":{"width":20}}', {"width": 10}, "does not match"),
|
||||
('{"media":{"fps":30}}', {"fps": 24}, "does not match"),
|
||||
(
|
||||
'{"label":"a","class":"b","point":[1,1]}',
|
||||
{},
|
||||
"Conflicting aliases",
|
||||
),
|
||||
(
|
||||
'{"frames":[{"frame_index":0},{"frame_index":0}]}',
|
||||
{},
|
||||
"Duplicate frame_index",
|
||||
),
|
||||
(
|
||||
'{"detections":[],"objects":[]}',
|
||||
{},
|
||||
"multiple detection collection",
|
||||
),
|
||||
]
|
||||
for response, overrides, message in cases:
|
||||
arguments = {
|
||||
"width": 10,
|
||||
"height": 10,
|
||||
"coordinate_mode": "pixel",
|
||||
**overrides,
|
||||
}
|
||||
with pytest.raises(ValueError, match=message):
|
||||
parse_spatial_response(response, **arguments)
|
||||
|
||||
|
||||
def test_node_methods_return_direct_canonical_payloads():
|
||||
prompt = VLMSpatialPromptBuilder().build(
|
||||
"Locate the subject.",
|
||||
"pixel",
|
||||
100,
|
||||
80,
|
||||
1,
|
||||
0,
|
||||
)[0]
|
||||
detections, points, normalized = VLMStructuredSpatialParser().parse(
|
||||
'{"bbox_xyxy":[1,2,30,40],"point":[5,6]}',
|
||||
"pixel",
|
||||
100,
|
||||
80,
|
||||
)
|
||||
|
||||
assert isinstance(prompt, str)
|
||||
assert isinstance(detections, DetectionSequence)
|
||||
assert isinstance(points, PointSequence)
|
||||
assert json.loads(normalized)["version"] == 1
|
||||
@@ -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))
|
||||
@@ -0,0 +1,234 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
from ComfyUI_VLM_nodes.nodes.tracking import (
|
||||
VLMByteTracker,
|
||||
associate_detection_sequence,
|
||||
)
|
||||
from ComfyUI_VLM_nodes.nodes.vision_types import (
|
||||
Detection,
|
||||
DetectionSequence,
|
||||
FrameDetections,
|
||||
)
|
||||
|
||||
|
||||
def _frame(
|
||||
frame_index,
|
||||
detections,
|
||||
*,
|
||||
width=100,
|
||||
height=100,
|
||||
fps=10.0,
|
||||
):
|
||||
timestamp = frame_index / fps
|
||||
records = tuple(
|
||||
Detection(
|
||||
bbox_xyxy=record["box"],
|
||||
label=record.get("label"),
|
||||
score=record.get("score"),
|
||||
frame_index=frame_index,
|
||||
timestamp=timestamp,
|
||||
mask=record.get("mask"),
|
||||
)
|
||||
for record in detections
|
||||
)
|
||||
return FrameDetections(
|
||||
frame_index=frame_index,
|
||||
timestamp=timestamp,
|
||||
width=width,
|
||||
height=height,
|
||||
detections=records,
|
||||
)
|
||||
|
||||
|
||||
def _sequence(frames, *, frame_count=None, fps=10.0, width=100, height=100):
|
||||
return DetectionSequence(
|
||||
width=width,
|
||||
height=height,
|
||||
frames=tuple(frames),
|
||||
frame_count=frame_count or 0,
|
||||
fps=fps,
|
||||
source="test-detector",
|
||||
)
|
||||
|
||||
|
||||
def test_low_confidence_second_stage_keeps_the_track_id():
|
||||
sequence = _sequence(
|
||||
(
|
||||
_frame(
|
||||
0,
|
||||
[{"box": (10, 10, 30, 30), "score": 0.9, "label": "cat"}],
|
||||
),
|
||||
_frame(
|
||||
1,
|
||||
[{"box": (11, 10, 31, 30), "score": 0.25, "label": "cat"}],
|
||||
),
|
||||
)
|
||||
)
|
||||
tracks = associate_detection_sequence(
|
||||
sequence,
|
||||
high_threshold=0.6,
|
||||
low_threshold=0.1,
|
||||
min_hits=1,
|
||||
)
|
||||
assert len(tracks.tracks) == 1
|
||||
track = tracks.tracks[0]
|
||||
assert track.track_id == 0
|
||||
assert [item.track_id for item in track.detections] == [0, 0]
|
||||
assert track.detections[1].metadata["association_stage"] == "low"
|
||||
assert track.metadata["state"] == "active"
|
||||
|
||||
|
||||
def test_low_confidence_detection_cannot_start_a_track():
|
||||
sequence = _sequence(
|
||||
(
|
||||
_frame(
|
||||
0,
|
||||
[{"box": (10, 10, 30, 30), "score": 0.2, "label": "cat"}],
|
||||
),
|
||||
)
|
||||
)
|
||||
tracks = associate_detection_sequence(
|
||||
sequence,
|
||||
high_threshold=0.6,
|
||||
low_threshold=0.1,
|
||||
min_hits=1,
|
||||
)
|
||||
assert tracks.tracks == ()
|
||||
|
||||
|
||||
def test_label_aware_matching_prevents_cross_class_identity_reuse():
|
||||
frames = (
|
||||
_frame(
|
||||
0,
|
||||
[{"box": (10, 10, 30, 30), "score": 0.9, "label": "cat"}],
|
||||
),
|
||||
_frame(
|
||||
1,
|
||||
[{"box": (10, 10, 30, 30), "score": 0.9, "label": "dog"}],
|
||||
),
|
||||
)
|
||||
label_aware = associate_detection_sequence(
|
||||
_sequence(frames),
|
||||
min_hits=1,
|
||||
label_aware=True,
|
||||
emit_predictions=False,
|
||||
)
|
||||
class_agnostic = associate_detection_sequence(
|
||||
_sequence(frames),
|
||||
min_hits=1,
|
||||
label_aware=False,
|
||||
emit_predictions=False,
|
||||
)
|
||||
assert len(label_aware.tracks) == 2
|
||||
assert [track.label for track in label_aware.tracks] == ["cat", "dog"]
|
||||
assert len(class_agnostic.tracks) == 1
|
||||
assert [
|
||||
detection.frame_index for detection in class_agnostic.tracks[0].detections
|
||||
] == [0, 1]
|
||||
|
||||
|
||||
def test_max_age_seconds_uses_fps_and_marks_removed_deterministically():
|
||||
sequence = _sequence(
|
||||
(
|
||||
_frame(
|
||||
0,
|
||||
[{"box": (10, 10, 30, 30), "score": 0.9}],
|
||||
fps=2.0,
|
||||
),
|
||||
),
|
||||
frame_count=4,
|
||||
fps=2.0,
|
||||
)
|
||||
tracks = associate_detection_sequence(
|
||||
sequence,
|
||||
min_hits=1,
|
||||
max_age_seconds=1.0,
|
||||
emit_predictions=True,
|
||||
)
|
||||
track = tracks.tracks[0]
|
||||
assert [item.frame_index for item in track.detections] == [0, 1, 2]
|
||||
assert [item.metadata["observation"] for item in track.detections] == [
|
||||
"detected",
|
||||
"predicted",
|
||||
"predicted",
|
||||
]
|
||||
assert track.metadata["state"] == "removed"
|
||||
assert track.metadata["removed_frame"] == 3
|
||||
assert track.metadata["last_observed_frame"] == 0
|
||||
|
||||
|
||||
def test_mask_iou_can_rescue_a_zero_bbox_iou_match():
|
||||
mask = torch.zeros(100, 100)
|
||||
mask[40:60, 40:60] = 1
|
||||
sequence = _sequence(
|
||||
(
|
||||
_frame(
|
||||
0,
|
||||
[
|
||||
{
|
||||
"box": (0, 0, 10, 10),
|
||||
"score": 0.9,
|
||||
"mask": mask,
|
||||
}
|
||||
],
|
||||
),
|
||||
_frame(
|
||||
1,
|
||||
[
|
||||
{
|
||||
"box": (20, 0, 30, 10),
|
||||
"score": 0.9,
|
||||
"mask": mask,
|
||||
}
|
||||
],
|
||||
),
|
||||
)
|
||||
)
|
||||
tracks = VLMByteTracker(
|
||||
min_hits=1,
|
||||
motion_gate=1.0e12,
|
||||
).track(sequence)
|
||||
assert len(tracks.tracks) == 1
|
||||
assert len(tracks.tracks[0].detections) == 2
|
||||
|
||||
|
||||
def test_hungarian_results_and_serialization_are_repeatable():
|
||||
sequence = _sequence(
|
||||
(
|
||||
_frame(
|
||||
0,
|
||||
[
|
||||
{"box": (5, 5, 20, 20), "score": 0.9},
|
||||
{"box": (40, 5, 55, 20), "score": 0.9},
|
||||
],
|
||||
),
|
||||
_frame(
|
||||
1,
|
||||
[
|
||||
{"box": (41, 5, 56, 20), "score": 0.9},
|
||||
{"box": (6, 5, 21, 20), "score": 0.9},
|
||||
],
|
||||
),
|
||||
)
|
||||
)
|
||||
options = {"min_hits": 1, "emit_predictions": False}
|
||||
first = associate_detection_sequence(sequence, **options)
|
||||
second = associate_detection_sequence(sequence, **options)
|
||||
assert first.to_json() == second.to_json()
|
||||
assert [
|
||||
[detection.bbox_xyxy for detection in track.detections]
|
||||
for track in first.tracks
|
||||
] == [
|
||||
[(5.0, 5.0, 20.0, 20.0), (6.0, 5.0, 21.0, 20.0)],
|
||||
[(40.0, 5.0, 55.0, 20.0), (41.0, 5.0, 56.0, 20.0)],
|
||||
]
|
||||
|
||||
|
||||
def test_tracker_rejects_invalid_threshold_order():
|
||||
try:
|
||||
VLMByteTracker(high_threshold=0.2, low_threshold=0.3)
|
||||
except ValueError as error:
|
||||
assert "low_threshold" in str(error)
|
||||
else:
|
||||
raise AssertionError("Expected invalid threshold order to fail.")
|
||||
@@ -0,0 +1,377 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from ComfyUI_VLM_nodes.nodes.video_intelligence import (
|
||||
NODE_CLASS_MAPPINGS,
|
||||
VLMAdaptiveFrameSampler,
|
||||
build_scene_state,
|
||||
build_video_reasoning_prompt,
|
||||
parse_video_reasoning_output,
|
||||
resize_video_for_analysis,
|
||||
sample_video_frames,
|
||||
scene_state_summary,
|
||||
track_aware_crops,
|
||||
)
|
||||
from ComfyUI_VLM_nodes.nodes.vision_types import (
|
||||
Detection,
|
||||
EventSequence,
|
||||
SceneState,
|
||||
SelectedVideoFrame,
|
||||
Track,
|
||||
TrackSequence,
|
||||
VideoFrameSelection,
|
||||
)
|
||||
|
||||
|
||||
def _moving_video(frame_count=20, height=48, width=64):
|
||||
frames = torch.zeros((frame_count, height, width, 3), dtype=torch.float32)
|
||||
for frame_index in range(frame_count):
|
||||
x = min(width - 9, 2 + frame_index * 2)
|
||||
frames[frame_index, 16:28, x : x + 8, 0] = 1.0
|
||||
if frame_index >= frame_count // 2:
|
||||
frames[frame_index, :, :, 2] += 0.55
|
||||
return frames.clamp(0, 1)
|
||||
|
||||
|
||||
def _tracks(width=64, height=48, frame_count=20, fps=10.0):
|
||||
detections = []
|
||||
for frame_index in (0, 5, 10, 15, 19):
|
||||
x = min(width - 12, 2 + frame_index * 2)
|
||||
detections.append(
|
||||
Detection(
|
||||
bbox_xyxy=(x, 14, x + 10, 30),
|
||||
label="red object",
|
||||
score=0.9 - frame_index * 0.005,
|
||||
frame_index=frame_index,
|
||||
timestamp=frame_index / fps,
|
||||
track_id=3,
|
||||
metadata={"track_state": "active"},
|
||||
)
|
||||
)
|
||||
return TrackSequence(
|
||||
width=width,
|
||||
height=height,
|
||||
frame_count=frame_count,
|
||||
fps=fps,
|
||||
tracks=(
|
||||
Track(
|
||||
track_id=3,
|
||||
detections=tuple(detections),
|
||||
label="red object",
|
||||
score=0.85,
|
||||
),
|
||||
),
|
||||
source="unit-test-tracker",
|
||||
)
|
||||
|
||||
|
||||
def _selection():
|
||||
return VideoFrameSelection(
|
||||
width=64,
|
||||
height=48,
|
||||
source_frame_count=20,
|
||||
fps=10.0,
|
||||
strategy="Hybrid: scene + motion + tracks",
|
||||
frames=(
|
||||
SelectedVideoFrame(0, 0.0, 1.0, ("first-frame",)),
|
||||
SelectedVideoFrame(5, 0.5, 0.7, ("motion",)),
|
||||
SelectedVideoFrame(10, 1.0, 0.9, ("scene-change",)),
|
||||
SelectedVideoFrame(19, 1.9, 1.0, ("last-frame",)),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_uniform_sampling_is_deterministic_and_preserves_timestamps():
|
||||
frames = _moving_video()
|
||||
sampled, selection, diagnostics = sample_video_frames(
|
||||
frames,
|
||||
fps=10.0,
|
||||
max_frames=5,
|
||||
strategy="Uniform coverage",
|
||||
)
|
||||
assert sampled.shape == (5, 48, 64, 3)
|
||||
assert selection.indices == (0, 5, 10, 14, 19)
|
||||
assert selection.timestamps == pytest.approx((0.0, 0.5, 1.0, 1.4, 1.9))
|
||||
assert diagnostics["visual_reduction_ratio"] == pytest.approx(0.75)
|
||||
assert torch.equal(sampled[2], frames[10])
|
||||
|
||||
|
||||
def test_hybrid_sampling_captures_boundaries_scene_change_and_motion():
|
||||
frames = _moving_video()
|
||||
sampled, selection, diagnostics = sample_video_frames(
|
||||
frames,
|
||||
fps=10.0,
|
||||
max_frames=7,
|
||||
strategy="Hybrid: scene + motion + tracks",
|
||||
minimum_gap_seconds=0.2,
|
||||
)
|
||||
assert sampled.shape[0] == 7
|
||||
assert selection.indices[0] == 0
|
||||
assert selection.indices[-1] == 19
|
||||
assert any(9 <= index <= 11 for index in selection.indices)
|
||||
assert diagnostics["motion_peak"] > 0
|
||||
assert diagnostics["scene_peak"] > 0
|
||||
assert selection.to_json() == VideoFrameSelection.from_json(
|
||||
selection.to_json()
|
||||
).to_json()
|
||||
|
||||
|
||||
def test_track_priority_uses_track_changes_and_validates_dimensions():
|
||||
frames = _moving_video()
|
||||
tracks = _tracks()
|
||||
_sampled, selection, diagnostics = sample_video_frames(
|
||||
frames,
|
||||
fps=10.0,
|
||||
max_frames=6,
|
||||
strategy="Track-change priority",
|
||||
tracks=tracks,
|
||||
)
|
||||
assert diagnostics["track_peak"] == pytest.approx(1.0)
|
||||
assert any(
|
||||
"track-change" in frame.reasons for frame in selection.frames
|
||||
)
|
||||
bad_tracks = TrackSequence(
|
||||
width=65,
|
||||
height=48,
|
||||
frame_count=20,
|
||||
fps=10,
|
||||
tracks=(),
|
||||
)
|
||||
with pytest.raises(ValueError, match="dimensions"):
|
||||
sample_video_frames(
|
||||
frames,
|
||||
fps=10,
|
||||
max_frames=4,
|
||||
tracks=bad_tracks,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"frames",
|
||||
[
|
||||
torch.zeros(4, 16, 16),
|
||||
torch.zeros(4, 16, 16, 2),
|
||||
torch.zeros(0, 16, 16, 3),
|
||||
torch.zeros(4, 16, 16, 3, dtype=torch.uint8),
|
||||
],
|
||||
)
|
||||
def test_sampling_rejects_invalid_video_tensors(frames):
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
sample_video_frames(frames, fps=24, max_frames=4)
|
||||
|
||||
|
||||
def test_track_aware_crops_preserve_identity_and_source_frame_mapping():
|
||||
frames = _moving_video()
|
||||
crops, manifest = track_aware_crops(
|
||||
frames,
|
||||
_tracks(),
|
||||
crops_per_track=3,
|
||||
max_crops=8,
|
||||
output_size=96,
|
||||
context_scale=1.4,
|
||||
)
|
||||
assert crops.shape == (3, 96, 96, 3)
|
||||
assert [item["track_id"] for item in manifest] == [3, 3, 3]
|
||||
assert [item["source_frame_index"] for item in manifest] == [0, 10, 19]
|
||||
assert crops.max().item() > 0.5
|
||||
|
||||
|
||||
def test_analysis_resize_reduces_pixels_without_changing_batch_or_aspect():
|
||||
frames = _moving_video(height=128, width=256)
|
||||
resized = resize_video_for_analysis(frames, max_side=128)
|
||||
assert resized.shape == (20, 64, 128, 3)
|
||||
assert resized.min().item() >= 0
|
||||
assert resized.max().item() <= 1
|
||||
assert resize_video_for_analysis(frames, max_side=0) is frames
|
||||
|
||||
|
||||
def test_scene_state_ignores_predicted_track_samples_and_computes_velocity():
|
||||
tracks = _tracks()
|
||||
predicted = Detection(
|
||||
bbox_xyxy=(48, 14, 58, 30),
|
||||
label="red object",
|
||||
frame_index=18,
|
||||
timestamp=1.8,
|
||||
track_id=3,
|
||||
metadata={"track_state": "predicted"},
|
||||
)
|
||||
track = tracks.tracks[0]
|
||||
with_prediction = TrackSequence(
|
||||
width=tracks.width,
|
||||
height=tracks.height,
|
||||
frame_count=tracks.frame_count,
|
||||
fps=tracks.fps,
|
||||
tracks=(
|
||||
Track(
|
||||
track_id=3,
|
||||
detections=tuple(
|
||||
sorted(
|
||||
(*track.detections, predicted),
|
||||
key=lambda item: item.frame_index,
|
||||
)
|
||||
),
|
||||
label=track.label,
|
||||
),
|
||||
),
|
||||
)
|
||||
scene = build_scene_state(with_prediction)
|
||||
assert len(scene.objects) == 1
|
||||
item = scene.objects[0]
|
||||
assert item.observation_count == 5
|
||||
assert item.velocity_xy_px_s[0] > 0
|
||||
assert "#3 red object" in scene_state_summary(scene)
|
||||
assert SceneState.from_json(scene.to_json()).to_json() == scene.to_json()
|
||||
|
||||
|
||||
def test_reasoning_prompt_explains_irregular_source_timeline():
|
||||
prompt = build_video_reasoning_prompt(
|
||||
_selection(),
|
||||
task="Robotics scene understanding",
|
||||
question="",
|
||||
max_events=12,
|
||||
)
|
||||
assert "irregularly spaced" in prompt
|
||||
assert "supplied image 2: source frame 10, timestamp 1.000000s" in prompt
|
||||
assert "Do not propose motor commands" in prompt
|
||||
assert "evidence_frame_indices" in prompt
|
||||
|
||||
|
||||
def test_structured_video_output_parses_fenced_json_and_preserves_evidence():
|
||||
response = """Result:
|
||||
```json
|
||||
{
|
||||
"summary": "A red object moves to the right.",
|
||||
"events": [
|
||||
{
|
||||
"start_time": 0.0,
|
||||
"end_time": 1.9,
|
||||
"label": "object motion",
|
||||
"text": "The red object moves from left to right.",
|
||||
"score": 0.94,
|
||||
"evidence_frame_indices": [0, 10, 19]
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
"""
|
||||
summary, events, normalized = parse_video_reasoning_output(
|
||||
response,
|
||||
_selection(),
|
||||
)
|
||||
assert summary == "A red object moves to the right."
|
||||
assert len(events.events) == 1
|
||||
assert events.events[0].metadata["evidence_frame_indices"] == (0, 10, 19)
|
||||
assert json.loads(normalized)["events"][0]["label"] == "object motion"
|
||||
|
||||
|
||||
def test_structured_output_normalizes_supplied_image_positions_to_source_frames():
|
||||
response = json.dumps(
|
||||
{
|
||||
"summary": "A transition occurs.",
|
||||
"events": [
|
||||
{
|
||||
"start_time": 0.5,
|
||||
"end_time": 1.9,
|
||||
"label": "transition",
|
||||
"text": "The scene changes.",
|
||||
"score": 0.8,
|
||||
# Positions 1 and 3 in the supplied image batch.
|
||||
"evidence_frame_indices": [1, 3],
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
_summary, events, _normalized = parse_video_reasoning_output(
|
||||
response,
|
||||
_selection(),
|
||||
)
|
||||
event = events.events[0]
|
||||
assert event.metadata["evidence_frame_indices"] == (5, 19)
|
||||
assert event.metadata["evidence_index_mode"] == "supplied-image-position"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("event_patch", "error"),
|
||||
[
|
||||
({"end_time": 2.1}, "outside"),
|
||||
({"score": 1.2}, "between"),
|
||||
({"evidence_frame_indices": [0, 7]}, "not supplied"),
|
||||
({"evidence_frame_indices": [0, 0]}, "duplicate"),
|
||||
({"label": "", "text": ""}, "requires"),
|
||||
],
|
||||
)
|
||||
def test_structured_video_output_rejects_unverifiable_events(event_patch, error):
|
||||
event = {
|
||||
"start_time": 0.0,
|
||||
"end_time": 1.0,
|
||||
"label": "motion",
|
||||
"text": "Object moves.",
|
||||
"score": 0.8,
|
||||
"evidence_frame_indices": [0, 10],
|
||||
}
|
||||
event.update(event_patch)
|
||||
with pytest.raises((TypeError, ValueError), match=error):
|
||||
parse_video_reasoning_output(
|
||||
json.dumps({"summary": "test", "events": [event]}),
|
||||
_selection(),
|
||||
)
|
||||
|
||||
|
||||
def test_scene_state_accepts_validated_events():
|
||||
_summary, events, _normalized = parse_video_reasoning_output(
|
||||
json.dumps(
|
||||
{
|
||||
"summary": "motion",
|
||||
"events": [
|
||||
{
|
||||
"start_time": 0.0,
|
||||
"end_time": 1.9,
|
||||
"label": "motion",
|
||||
"text": "Object moves.",
|
||||
"score": 0.9,
|
||||
"evidence_frame_indices": [0, 19],
|
||||
}
|
||||
],
|
||||
}
|
||||
),
|
||||
_selection(),
|
||||
)
|
||||
scene = build_scene_state(_tracks(), events)
|
||||
assert isinstance(events, EventSequence)
|
||||
assert len(scene.events) == 1
|
||||
assert "Event 0.000s–1.900s" in scene_state_summary(scene)
|
||||
|
||||
|
||||
def test_node_surface_registers_all_video_intelligence_nodes():
|
||||
assert set(NODE_CLASS_MAPPINGS) == {
|
||||
"VLMAdaptiveFrameSampler",
|
||||
"VLMTrackAwareCrops",
|
||||
"VLMBuildSceneState",
|
||||
"VLMVideoReasoningPrompt",
|
||||
"VLMEventsFromVideoJSON",
|
||||
"VLMVideoTemporalReasoner",
|
||||
}
|
||||
inputs = VLMAdaptiveFrameSampler.INPUT_TYPES()
|
||||
assert inputs["required"]["frames"][0] == "IMAGE"
|
||||
assert inputs["optional"]["tracks"][0] == "VLM_TRACKS"
|
||||
reasoner = NODE_CLASS_MAPPINGS["VLMVideoTemporalReasoner"]
|
||||
assert reasoner.RETURN_NAMES[-2:] == ("events_json", "selection_json")
|
||||
|
||||
|
||||
def test_api_example_uses_direct_json_outputs_and_preview():
|
||||
example = json.loads(
|
||||
(
|
||||
Path(__file__).resolve().parents[1]
|
||||
/ "examples"
|
||||
/ "vision"
|
||||
/ "video_temporal_reasoning_api.json"
|
||||
).read_text(encoding="utf-8")
|
||||
)
|
||||
assert example["3"]["class_type"] == "VLMVideoTemporalReasoner"
|
||||
assert example["5"]["inputs"]["text"] == ["3", 6]
|
||||
assert example["6"]["inputs"]["text"] == ["3", 7]
|
||||
assert example["8"]["inputs"]["images"] == ["3", 3]
|
||||
@@ -0,0 +1,273 @@
|
||||
import importlib
|
||||
import json
|
||||
from dataclasses import FrozenInstanceError
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
PACKAGE = Path(__file__).resolve().parents[1].name
|
||||
vision_types = importlib.import_module(f"{PACKAGE}.nodes.vision_types")
|
||||
|
||||
DETECTIONS_SCHEMA = vision_types.DETECTIONS_SCHEMA
|
||||
EVENTS_SCHEMA = vision_types.EVENTS_SCHEMA
|
||||
POINTS_SCHEMA = vision_types.POINTS_SCHEMA
|
||||
SCHEMA_VERSION = vision_types.SCHEMA_VERSION
|
||||
SCENE_STATE_SCHEMA = vision_types.SCENE_STATE_SCHEMA
|
||||
TRACKS_SCHEMA = vision_types.TRACKS_SCHEMA
|
||||
VIDEO_SELECTION_SCHEMA = vision_types.VIDEO_SELECTION_SCHEMA
|
||||
VLM_DETECTIONS = vision_types.VLM_DETECTIONS
|
||||
VLM_EVENTS = vision_types.VLM_EVENTS
|
||||
VLM_POINTS = vision_types.VLM_POINTS
|
||||
VLM_SCENE_STATE = vision_types.VLM_SCENE_STATE
|
||||
VLM_TRACKS = vision_types.VLM_TRACKS
|
||||
VLM_VIDEO_SELECTION = vision_types.VLM_VIDEO_SELECTION
|
||||
Detection = vision_types.Detection
|
||||
DetectionSequence = vision_types.DetectionSequence
|
||||
EventSequence = vision_types.EventSequence
|
||||
FrameDetections = vision_types.FrameDetections
|
||||
FrozenDict = vision_types.FrozenDict
|
||||
PointSequence = vision_types.PointSequence
|
||||
TemporalEvent = vision_types.TemporalEvent
|
||||
Track = vision_types.Track
|
||||
TrackSequence = vision_types.TrackSequence
|
||||
VisionPoint = vision_types.VisionPoint
|
||||
|
||||
|
||||
def sample_sequence() -> DetectionSequence:
|
||||
mask = torch.zeros((24, 32), dtype=torch.float32)
|
||||
mask[3:12, 4:18] = 1
|
||||
first = Detection(
|
||||
bbox_xyxy=(4, 3, 18, 12),
|
||||
label="cat",
|
||||
text="sleeping cat",
|
||||
score=0.875,
|
||||
polygon=((4, 3), (18, 3), (18, 12), (4, 12)),
|
||||
quad=((4, 3), (18, 3), (18, 12), (4, 12)),
|
||||
frame_index=0,
|
||||
timestamp=0.0,
|
||||
track_id=7,
|
||||
source="unit-test",
|
||||
metadata={"attributes": ["small", "red"], "visible": True},
|
||||
mask=mask,
|
||||
)
|
||||
second = Detection(
|
||||
bbox_xyxy=(6.5, 5.0, 20.25, 15.0),
|
||||
label="cat",
|
||||
score=0.8,
|
||||
frame_index=1,
|
||||
timestamp=0.5,
|
||||
track_id=7,
|
||||
)
|
||||
return DetectionSequence(
|
||||
width=32,
|
||||
height=24,
|
||||
frame_count=2,
|
||||
fps=2.0,
|
||||
source="synthetic",
|
||||
metadata={"nested": {"value": 3}},
|
||||
frames=(
|
||||
FrameDetections(
|
||||
frame_index=0,
|
||||
timestamp=0.0,
|
||||
width=32,
|
||||
height=24,
|
||||
detections=(first,),
|
||||
),
|
||||
FrameDetections(
|
||||
frame_index=1,
|
||||
timestamp=0.5,
|
||||
width=32,
|
||||
height=24,
|
||||
detections=(second,),
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_public_socket_and_schema_names_are_stable():
|
||||
assert VLM_DETECTIONS == "VLM_DETECTIONS"
|
||||
assert VLM_TRACKS == "VLM_TRACKS"
|
||||
assert VLM_POINTS == "VLM_POINTS"
|
||||
assert VLM_EVENTS == "VLM_EVENTS"
|
||||
assert VLM_VIDEO_SELECTION == "VLM_VIDEO_SELECTION"
|
||||
assert VLM_SCENE_STATE == "VLM_SCENE_STATE"
|
||||
assert SCHEMA_VERSION == 1
|
||||
assert DETECTIONS_SCHEMA == "comfyui-vlm/detections"
|
||||
assert TRACKS_SCHEMA == "comfyui-vlm/tracks"
|
||||
assert POINTS_SCHEMA == "comfyui-vlm/points"
|
||||
assert EVENTS_SCHEMA == "comfyui-vlm/events"
|
||||
assert VIDEO_SELECTION_SCHEMA == "comfyui-vlm/video-selection"
|
||||
assert SCENE_STATE_SCHEMA == "comfyui-vlm/scene-state"
|
||||
|
||||
|
||||
def test_detection_payload_is_validated_immutable_and_mask_safe():
|
||||
original = torch.ones((4, 5))
|
||||
detection = Detection(
|
||||
bbox_xyxy=(0, 0, 5, 4),
|
||||
label="object",
|
||||
score=1.0,
|
||||
metadata={"items": [1, {"ready": True}]},
|
||||
mask=original,
|
||||
)
|
||||
original.zero_()
|
||||
assert detection.mask.sum().item() == 20
|
||||
assert isinstance(detection.metadata, FrozenDict)
|
||||
assert detection.metadata["items"][1]["ready"] is True
|
||||
|
||||
with pytest.raises(FrozenInstanceError):
|
||||
detection.label = "changed"
|
||||
with pytest.raises(TypeError):
|
||||
detection.metadata["new"] = "value"
|
||||
with pytest.raises(TypeError, match="Metadata"):
|
||||
Detection(
|
||||
bbox_xyxy=(0, 0, 1, 1),
|
||||
metadata={"tensor": torch.ones(1)},
|
||||
)
|
||||
|
||||
record = detection.to_dict()
|
||||
assert "mask" not in record
|
||||
assert "mask" not in json.dumps(record)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("kwargs", "error"),
|
||||
[
|
||||
({"bbox_xyxy": (2, 0, 1, 2)}, "x2"),
|
||||
({"bbox_xyxy": (-1, 0, 1, 2)}, "non-negative"),
|
||||
({"bbox_xyxy": (0, 0, 1, 2), "score": 1.1}, "between 0 and 1"),
|
||||
({"bbox_xyxy": (0, 0, 1, 2), "frame_index": -1}, "frame_index"),
|
||||
({"bbox_xyxy": (0, 0, 1, 2), "timestamp": -0.1}, "timestamp"),
|
||||
(
|
||||
{"bbox_xyxy": (0, 0, 1, 2), "quad": ((0, 0), (1, 0), (1, 1))},
|
||||
"exactly 4",
|
||||
),
|
||||
(
|
||||
{"bbox_xyxy": (0, 0, 1, 2), "mask": torch.ones(1, 2, 3)},
|
||||
"shape",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_detection_rejects_invalid_values(kwargs, error):
|
||||
with pytest.raises((TypeError, ValueError), match=error):
|
||||
Detection(**kwargs)
|
||||
|
||||
|
||||
def test_detection_sequence_json_round_trip_is_versioned_and_tensor_free():
|
||||
sequence = sample_sequence()
|
||||
encoded = sequence.to_json(indent=2)
|
||||
decoded_json = json.loads(encoded)
|
||||
assert decoded_json["schema"] == DETECTIONS_SCHEMA
|
||||
assert decoded_json["version"] == SCHEMA_VERSION
|
||||
assert decoded_json["media"] == {
|
||||
"fps": 2.0,
|
||||
"frame_count": 2,
|
||||
"height": 24,
|
||||
"width": 32,
|
||||
}
|
||||
assert "mask" not in encoded
|
||||
|
||||
restored = DetectionSequence.from_json(encoded)
|
||||
assert restored.to_dict() == sequence.to_dict()
|
||||
assert restored.frames[0].detections[0].mask is None
|
||||
assert restored.all_detections()[1].center == pytest.approx((13.375, 10.0))
|
||||
assert restored.frame(99) is None
|
||||
|
||||
decoded_json["version"] = 99
|
||||
with pytest.raises(ValueError, match="Unsupported"):
|
||||
DetectionSequence.from_dict(decoded_json)
|
||||
decoded_json["version"] = SCHEMA_VERSION
|
||||
decoded_json["schema"] = "other"
|
||||
with pytest.raises(ValueError, match="Expected schema"):
|
||||
DetectionSequence.from_dict(decoded_json)
|
||||
|
||||
|
||||
def test_frame_and_sequence_enforce_dimensions_order_and_timestamps():
|
||||
detection = Detection(
|
||||
bbox_xyxy=(0, 0, 11, 5),
|
||||
frame_index=0,
|
||||
timestamp=0,
|
||||
)
|
||||
with pytest.raises(ValueError, match="width"):
|
||||
FrameDetections(
|
||||
frame_index=0,
|
||||
timestamp=0,
|
||||
width=10,
|
||||
height=10,
|
||||
detections=(detection,),
|
||||
)
|
||||
mismatch = Detection(
|
||||
bbox_xyxy=(0, 0, 1, 1),
|
||||
frame_index=1,
|
||||
timestamp=0,
|
||||
)
|
||||
with pytest.raises(ValueError, match="frame_index"):
|
||||
FrameDetections(
|
||||
frame_index=0,
|
||||
timestamp=0,
|
||||
width=10,
|
||||
height=10,
|
||||
detections=(mismatch,),
|
||||
)
|
||||
|
||||
frame_one = FrameDetections(1, 1.0, 10, 10)
|
||||
frame_zero = FrameDetections(0, 0.0, 10, 10)
|
||||
with pytest.raises(ValueError, match="increasing"):
|
||||
DetectionSequence(10, 10, frames=(frame_one, frame_zero))
|
||||
|
||||
|
||||
def test_point_track_and_event_schemas_round_trip_without_tensors():
|
||||
sequence = sample_sequence()
|
||||
points = PointSequence(
|
||||
width=32,
|
||||
height=24,
|
||||
frame_count=2,
|
||||
fps=2,
|
||||
points=(
|
||||
VisionPoint(
|
||||
x=11,
|
||||
y=7.5,
|
||||
label="cat",
|
||||
score=0.8,
|
||||
frame_index=0,
|
||||
track_id=7,
|
||||
),
|
||||
),
|
||||
)
|
||||
assert PointSequence.from_json(points.to_json()).to_dict() == points.to_dict()
|
||||
|
||||
track = Track(
|
||||
track_id=7,
|
||||
label="cat",
|
||||
score=0.8,
|
||||
detections=sequence.all_detections(),
|
||||
)
|
||||
tracks = TrackSequence(
|
||||
width=32,
|
||||
height=24,
|
||||
frame_count=2,
|
||||
fps=2,
|
||||
tracks=(track,),
|
||||
)
|
||||
restored_tracks = TrackSequence.from_json(tracks.to_json())
|
||||
assert restored_tracks.to_dict() == tracks.to_dict()
|
||||
assert "mask" not in tracks.to_json()
|
||||
|
||||
events = EventSequence(
|
||||
duration=2.0,
|
||||
events=(
|
||||
TemporalEvent(
|
||||
start_time=0.25,
|
||||
end_time=1.5,
|
||||
label="movement",
|
||||
text="the cat moves",
|
||||
score=0.9,
|
||||
),
|
||||
),
|
||||
)
|
||||
assert EventSequence.from_json(events.to_json()).to_dict() == events.to_dict()
|
||||
with pytest.raises(ValueError, match="duration"):
|
||||
EventSequence(
|
||||
duration=1.0,
|
||||
events=(TemporalEvent(0.0, 2.0),),
|
||||
)
|
||||
@@ -0,0 +1,432 @@
|
||||
import importlib
|
||||
import inspect
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
PACKAGE = Path(__file__).resolve().parents[1].name
|
||||
vision_types = importlib.import_module(f"{PACKAGE}.nodes.vision_types")
|
||||
vision_utils = importlib.import_module(f"{PACKAGE}.nodes.vision_utils")
|
||||
|
||||
VLM_DETECTIONS = vision_types.VLM_DETECTIONS
|
||||
VLM_POINTS = vision_types.VLM_POINTS
|
||||
Detection = vision_types.Detection
|
||||
DetectionSequence = vision_types.DetectionSequence
|
||||
FrameDetections = vision_types.FrameDetections
|
||||
BOUNDING_BOXES = vision_utils.BOUNDING_BOXES
|
||||
NODE_CLASS_MAPPINGS = vision_utils.NODE_CLASS_MAPPINGS
|
||||
VLMCropDetections = vision_utils.VLMCropDetections
|
||||
VLMDetectionsFromJSON = vision_utils.VLMDetectionsFromJSON
|
||||
VLMDetectionsToBoundingBoxes = vision_utils.VLMDetectionsToBoundingBoxes
|
||||
VLMDetectionsToJSON = vision_utils.VLMDetectionsToJSON
|
||||
VLMDetectionsToMasks = vision_utils.VLMDetectionsToMasks
|
||||
VLMDetectionsToPoints = vision_utils.VLMDetectionsToPoints
|
||||
VLMFilterDetections = vision_utils.VLMFilterDetections
|
||||
VLMMaskComposite = vision_utils.VLMMaskComposite
|
||||
VLMMaskProcessor = vision_utils.VLMMaskProcessor
|
||||
VLMRenderDetections = vision_utils.VLMRenderDetections
|
||||
VLMSelectDetection = vision_utils.VLMSelectDetection
|
||||
bounding_boxes_payload = vision_utils.bounding_boxes_payload
|
||||
composite_with_mask = vision_utils.composite_with_mask
|
||||
crop_detections = vision_utils.crop_detections
|
||||
detection_centers = vision_utils.detection_centers
|
||||
filter_detection_sequence = vision_utils.filter_detection_sequence
|
||||
instance_map_images = vision_utils.instance_map_images
|
||||
masks_to_images = vision_utils.masks_to_images
|
||||
process_masks = vision_utils.process_masks
|
||||
render_detections = vision_utils.render_detections
|
||||
select_detection_sequence = vision_utils.select_detection_sequence
|
||||
sequence_masks = vision_utils.sequence_masks
|
||||
|
||||
|
||||
def sample_sequence() -> DetectionSequence:
|
||||
cat_mask = torch.zeros((16, 20))
|
||||
cat_mask[2:8, 3:10] = 1
|
||||
return DetectionSequence(
|
||||
width=20,
|
||||
height=16,
|
||||
frame_count=2,
|
||||
fps=2,
|
||||
frames=(
|
||||
FrameDetections(
|
||||
frame_index=0,
|
||||
timestamp=0,
|
||||
width=20,
|
||||
height=16,
|
||||
detections=(
|
||||
Detection(
|
||||
bbox_xyxy=(3.2, 2.1, 10.0, 8.0),
|
||||
label="cat",
|
||||
score=0.9,
|
||||
frame_index=0,
|
||||
timestamp=0,
|
||||
track_id=5,
|
||||
mask=cat_mask,
|
||||
),
|
||||
Detection(
|
||||
bbox_xyxy=(12, 4, 18, 13),
|
||||
label="dog",
|
||||
score=0.55,
|
||||
frame_index=0,
|
||||
timestamp=0,
|
||||
),
|
||||
),
|
||||
),
|
||||
FrameDetections(
|
||||
frame_index=1,
|
||||
timestamp=0.5,
|
||||
width=20,
|
||||
height=16,
|
||||
detections=(
|
||||
Detection(
|
||||
bbox_xyxy=(5, 3, 12, 10),
|
||||
label="cat",
|
||||
score=0.8,
|
||||
frame_index=1,
|
||||
timestamp=0.5,
|
||||
track_id=5,
|
||||
),
|
||||
),
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def empty_sequence() -> DetectionSequence:
|
||||
return DetectionSequence(
|
||||
width=20,
|
||||
height=16,
|
||||
frame_count=2,
|
||||
frames=(
|
||||
FrameDetections(0, 0, 20, 16),
|
||||
FrameDetections(1, 1, 20, 16),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_filter_and_select_by_all_supported_fields():
|
||||
sequence = sample_sequence()
|
||||
filtered = filter_detection_sequence(
|
||||
sequence,
|
||||
label="CAT",
|
||||
label_mode="exact",
|
||||
minimum_score=0.85,
|
||||
minimum_area=30,
|
||||
maximum_area=50,
|
||||
track_id=5,
|
||||
frame_index=0,
|
||||
)
|
||||
assert len(filtered.frames) == 1
|
||||
assert [item.label for item in filtered.all_detections()] == ["cat"]
|
||||
|
||||
contains = filter_detection_sequence(
|
||||
sequence,
|
||||
label="o",
|
||||
minimum_score=0.5,
|
||||
)
|
||||
assert [item.label for item in contains.all_detections()] == ["dog"]
|
||||
|
||||
selected = select_detection_sequence(sequence, 1)
|
||||
assert [item.label for item in selected.all_detections()] == ["dog"]
|
||||
missing = select_detection_sequence(sequence, 100)
|
||||
assert missing.all_detections() == ()
|
||||
assert len(missing.frames) == 2
|
||||
|
||||
with pytest.raises(ValueError, match="maximum_area"):
|
||||
filter_detection_sequence(
|
||||
sequence,
|
||||
minimum_area=10,
|
||||
maximum_area=5,
|
||||
)
|
||||
|
||||
|
||||
def test_core_bounding_boxes_are_integer_xywh_with_metadata():
|
||||
payload = bounding_boxes_payload(sample_sequence())
|
||||
assert payload[0] == {
|
||||
"x": 3,
|
||||
"y": 2,
|
||||
"width": 7,
|
||||
"height": 6,
|
||||
"metadata": {
|
||||
"bbox_xyxy": [3.2, 2.1, 10.0, 8.0],
|
||||
"frame_index": 0,
|
||||
"timestamp": 0.0,
|
||||
"label": "cat",
|
||||
"score": 0.9,
|
||||
"track_id": 5,
|
||||
},
|
||||
}
|
||||
assert payload[2]["x"] == 5
|
||||
assert payload[2]["metadata"]["frame_index"] == 1
|
||||
assert bounding_boxes_payload(empty_sequence()) == []
|
||||
|
||||
|
||||
def test_centers_and_masks_preserve_frame_mapping_and_empty_shapes():
|
||||
sequence = sample_sequence()
|
||||
points = detection_centers(sequence)
|
||||
assert points.points[0].x == pytest.approx(6.6)
|
||||
assert points.points[0].y == pytest.approx(5.05)
|
||||
assert points.points[0].track_id == 5
|
||||
assert json.loads(points.to_json())["schema"] == "comfyui-vlm/points"
|
||||
|
||||
unions, individuals, mapping = sequence_masks(sequence)
|
||||
assert unions.shape == (2, 16, 20)
|
||||
assert individuals.shape == (3, 16, 20)
|
||||
assert unions[0].sum() > 0
|
||||
assert mapping[0] == {
|
||||
"mask_index": 0,
|
||||
"frame_index": 0,
|
||||
"detection_index": 0,
|
||||
"label": "cat",
|
||||
"track_id": 5,
|
||||
}
|
||||
|
||||
empty_unions, empty_individuals, empty_mapping = sequence_masks(empty_sequence())
|
||||
assert empty_unions.shape == (2, 16, 20)
|
||||
assert empty_unions.sum() == 0
|
||||
assert empty_individuals.shape == (0, 16, 20)
|
||||
assert empty_mapping == []
|
||||
|
||||
|
||||
def test_creator_mask_outputs_are_binary_previewable_and_instance_colored():
|
||||
sequence = sample_sequence()
|
||||
unions, individuals, _mapping = sequence_masks(sequence)
|
||||
union_images = masks_to_images(unions)
|
||||
individual_images = masks_to_images(individuals)
|
||||
instance_maps = instance_map_images(sequence)
|
||||
|
||||
assert set(torch.unique(unions).tolist()) <= {0.0, 1.0}
|
||||
assert set(torch.unique(individuals).tolist()) <= {0.0, 1.0}
|
||||
assert union_images.shape == (2, 16, 20, 3)
|
||||
assert individual_images.shape == (3, 16, 20, 3)
|
||||
assert torch.equal(union_images[..., 0], unions)
|
||||
assert torch.equal(union_images[..., 0], union_images[..., 2])
|
||||
assert instance_maps.shape == (2, 16, 20, 3)
|
||||
assert instance_maps.sum() > 0
|
||||
assert torch.equal(instance_maps, instance_map_images(sequence))
|
||||
|
||||
|
||||
def test_mask_processing_grow_shrink_feather_and_inverse_are_batch_safe():
|
||||
mask = torch.zeros((2, 9, 9))
|
||||
mask[:, 4, 4] = 1
|
||||
grown, grown_binary, grown_inverse = process_masks(
|
||||
mask,
|
||||
threshold=0.5,
|
||||
grow_shrink=1,
|
||||
feather_radius=0,
|
||||
)
|
||||
assert grown.shape == mask.shape
|
||||
assert grown[0].sum() == 9
|
||||
assert torch.equal(grown, grown_binary)
|
||||
assert torch.allclose(grown + grown_inverse, torch.ones_like(grown))
|
||||
|
||||
soft, binary, inverse = process_masks(
|
||||
mask,
|
||||
threshold=0.5,
|
||||
grow_shrink=2,
|
||||
feather_radius=2,
|
||||
)
|
||||
assert binary[0].sum() == 25
|
||||
assert torch.any((soft > 0) & (soft < 1))
|
||||
assert torch.allclose(soft + inverse, torch.ones_like(soft), atol=1e-6)
|
||||
|
||||
full = torch.ones((1, 9, 9))
|
||||
shrunk, _, _ = process_masks(full, grow_shrink=-1)
|
||||
assert shrunk[0, 0].sum() == 0
|
||||
assert shrunk[0, -1].sum() == 0
|
||||
with pytest.raises(ValueError, match="threshold"):
|
||||
process_masks(mask, threshold=2)
|
||||
|
||||
|
||||
def test_mask_composite_splits_foreground_and_broadcasts_video_batches():
|
||||
image = torch.zeros((2, 4, 5, 3))
|
||||
image[..., 0] = 1
|
||||
mask = torch.zeros((1, 4, 5))
|
||||
mask[:, :, :2] = 1
|
||||
replacement = torch.zeros((1, 4, 5, 3))
|
||||
replacement[..., 2] = 1
|
||||
|
||||
composite, foreground, background_only, mask_image = composite_with_mask(
|
||||
image,
|
||||
mask,
|
||||
background=replacement,
|
||||
)
|
||||
assert composite.shape == image.shape
|
||||
assert foreground.shape == image.shape
|
||||
assert background_only.shape == image.shape
|
||||
assert mask_image.shape == image.shape
|
||||
assert torch.all(composite[:, :, :2, 0] == 1)
|
||||
assert torch.all(composite[:, :, 2:, 2] == 1)
|
||||
assert foreground[:, :, 2:].sum() == 0
|
||||
assert background_only[:, :, :2].sum() == 0
|
||||
|
||||
solid, *_ = composite_with_mask(
|
||||
image[:1],
|
||||
mask,
|
||||
background_color="#0f0",
|
||||
)
|
||||
assert torch.all(solid[:, :, 2:, 1] == 1)
|
||||
with pytest.raises(ValueError, match="dimensions"):
|
||||
composite_with_mask(image, torch.zeros((1, 3, 3)))
|
||||
|
||||
|
||||
def test_rendering_is_deterministic_batch_safe_and_empty_safe():
|
||||
image = torch.zeros((2, 16, 20, 3), dtype=torch.float32)
|
||||
sequence = sample_sequence()
|
||||
first = render_detections(
|
||||
image,
|
||||
sequence,
|
||||
draw_masks=True,
|
||||
draw_labels=True,
|
||||
)
|
||||
second = render_detections(
|
||||
image,
|
||||
sequence,
|
||||
draw_masks=True,
|
||||
draw_labels=True,
|
||||
)
|
||||
assert first.shape == image.shape
|
||||
assert torch.equal(first, second)
|
||||
assert first.sum() > 0
|
||||
|
||||
unchanged = render_detections(image, empty_sequence())
|
||||
assert torch.equal(unchanged, image)
|
||||
with pytest.raises(ValueError, match="dimensions"):
|
||||
render_detections(
|
||||
torch.zeros((2, 10, 10, 3)),
|
||||
sequence,
|
||||
)
|
||||
|
||||
|
||||
def test_padded_square_crops_form_a_non_distorted_image_batch():
|
||||
image = torch.zeros((2, 16, 20, 3), dtype=torch.float32)
|
||||
image[0, 2:8, 3:10, 0] = 1
|
||||
image[0, 4:13, 12:18, 1] = 1
|
||||
image[1, 3:10, 5:12, 2] = 1
|
||||
crops, metadata = crop_detections(
|
||||
image,
|
||||
sample_sequence(),
|
||||
padding=1,
|
||||
square=True,
|
||||
)
|
||||
assert crops.shape[0] == 3
|
||||
assert crops.ndim == 4
|
||||
assert crops.shape[1] == max(record["valid_height"] for record in metadata)
|
||||
assert crops.shape[2] == max(record["valid_width"] for record in metadata)
|
||||
assert all(record["batch_width"] == crops.shape[2] for record in metadata)
|
||||
assert metadata[0]["track_id"] == 5
|
||||
valid_crop = crops[
|
||||
0,
|
||||
: metadata[0]["valid_height"],
|
||||
: metadata[0]["valid_width"],
|
||||
]
|
||||
assert valid_crop.sum() > 0
|
||||
|
||||
empty_crops, empty_metadata = crop_detections(image, empty_sequence())
|
||||
assert empty_crops.shape == (0, 1, 1, 3)
|
||||
assert empty_metadata == []
|
||||
|
||||
|
||||
def test_utility_node_contracts_and_json_round_trip():
|
||||
sequence = sample_sequence()
|
||||
encoded = VLMDetectionsToJSON().serialize(sequence, pretty=False)[0]
|
||||
restored = VLMDetectionsFromJSON().parse(encoded)[0]
|
||||
assert restored.to_dict() == sequence.to_dict()
|
||||
|
||||
filtered = VLMFilterDetections().filter(
|
||||
sequence,
|
||||
label="cat",
|
||||
label_mode="exact",
|
||||
minimum_score=0,
|
||||
minimum_area=0,
|
||||
maximum_area=0,
|
||||
track_id=-1,
|
||||
frame_index=-1,
|
||||
)[0]
|
||||
assert len(filtered.all_detections()) == 2
|
||||
assert len(VLMSelectDetection().select(sequence, 0)[0].all_detections()) == 1
|
||||
|
||||
boxes, boxes_json = VLMDetectionsToBoundingBoxes().convert(sequence)
|
||||
assert json.loads(boxes_json) == boxes
|
||||
points, points_json = VLMDetectionsToPoints().convert(sequence)
|
||||
assert json.loads(points_json) == points.to_dict()
|
||||
(
|
||||
union,
|
||||
individual,
|
||||
mask_json,
|
||||
inverse,
|
||||
union_images,
|
||||
individual_images,
|
||||
instance_maps,
|
||||
) = VLMDetectionsToMasks().convert(sequence)
|
||||
assert union.shape[0] == 2
|
||||
assert individual.shape[0] == 3
|
||||
assert len(json.loads(mask_json)) == 3
|
||||
assert torch.allclose(union + inverse, torch.ones_like(union))
|
||||
assert union_images.shape == (2, 16, 20, 3)
|
||||
assert individual_images.shape == (3, 16, 20, 3)
|
||||
assert instance_maps.shape == (2, 16, 20, 3)
|
||||
|
||||
processed = VLMMaskProcessor().process(union, 0.5, 1, 1)
|
||||
assert [value.shape[0] for value in processed] == [2, 2, 2, 2]
|
||||
composited = VLMMaskComposite().composite(
|
||||
torch.ones((2, 16, 20, 3)),
|
||||
union,
|
||||
"#000000",
|
||||
)
|
||||
assert all(value.shape == (2, 16, 20, 3) for value in composited)
|
||||
|
||||
image = torch.zeros((2, 16, 20, 3))
|
||||
assert (
|
||||
VLMRenderDetections()
|
||||
.render(
|
||||
image,
|
||||
sequence,
|
||||
True,
|
||||
False,
|
||||
0.25,
|
||||
2,
|
||||
)[0]
|
||||
.shape
|
||||
== image.shape
|
||||
)
|
||||
crops, crop_json = VLMCropDetections().crop(image, sequence, 0, False)
|
||||
assert crops.shape[0] == 3
|
||||
assert len(json.loads(crop_json)) == 3
|
||||
|
||||
|
||||
def test_every_utility_node_accepts_its_declared_inputs():
|
||||
expected = {
|
||||
"VLMDetectionsFromJSON",
|
||||
"VLMDetectionsToJSON",
|
||||
"VLMFilterDetections",
|
||||
"VLMSelectDetection",
|
||||
"VLMDetectionsToBoundingBoxes",
|
||||
"VLMDetectionsToPoints",
|
||||
"VLMDetectionsToMasks",
|
||||
"VLMMaskProcessor",
|
||||
"VLMMaskComposite",
|
||||
"VLMRenderDetections",
|
||||
"VLMCropDetections",
|
||||
}
|
||||
assert set(NODE_CLASS_MAPPINGS) == expected
|
||||
assert VLMDetectionsFromJSON.RETURN_TYPES == (VLM_DETECTIONS,)
|
||||
assert VLMDetectionsToBoundingBoxes.RETURN_TYPES[0] == BOUNDING_BOXES
|
||||
assert VLMDetectionsToPoints.RETURN_TYPES[0] == VLM_POINTS
|
||||
|
||||
for node_class in NODE_CLASS_MAPPINGS.values():
|
||||
schema = node_class.INPUT_TYPES()
|
||||
declared = {
|
||||
name
|
||||
for group in schema.values()
|
||||
if isinstance(group, dict)
|
||||
for name in group
|
||||
}
|
||||
function = getattr(node_class, node_class.FUNCTION)
|
||||
accepted = set(inspect.signature(function).parameters)
|
||||
assert declared <= accepted, (
|
||||
node_class.__name__,
|
||||
declared - accepted,
|
||||
)
|
||||
@@ -0,0 +1,62 @@
|
||||
import { app } from "../../../scripts/app.js";
|
||||
|
||||
const LLM_NODE = "PromptGenerateAPI";
|
||||
const SAFE_SOURCE = "Provider environment variable";
|
||||
const NO_KEY_SOURCE = "No key (loopback custom endpoint only)";
|
||||
const SAFE_SOURCES = new Set([SAFE_SOURCE, NO_KEY_SOURCE]);
|
||||
const CREDENTIAL_WIDGET_INDEX = 2;
|
||||
|
||||
function visitGraphNodes(graphData, callback) {
|
||||
for (const node of graphData?.nodes ?? []) {
|
||||
callback(node);
|
||||
}
|
||||
for (const subgraph of graphData?.definitions?.subgraphs ?? []) {
|
||||
visitGraphNodes(subgraph, callback);
|
||||
}
|
||||
}
|
||||
|
||||
function scrubSerializedNode(node) {
|
||||
if (node?.type !== LLM_NODE) {
|
||||
return;
|
||||
}
|
||||
const values = node.widgets_values;
|
||||
if (Array.isArray(values)) {
|
||||
const saved = values[CREDENTIAL_WIDGET_INDEX];
|
||||
if (!SAFE_SOURCES.has(saved)) {
|
||||
values[CREDENTIAL_WIDGET_INDEX] = SAFE_SOURCE;
|
||||
}
|
||||
return;
|
||||
}
|
||||
if (values && typeof values === "object") {
|
||||
// Some frontend versions serialize widgets by name.
|
||||
delete values.api_key;
|
||||
if (!SAFE_SOURCES.has(values.credential_source)) {
|
||||
values.credential_source = SAFE_SOURCE;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
function enforceLiveWidget(node) {
|
||||
if (node?.type !== LLM_NODE) {
|
||||
return;
|
||||
}
|
||||
const widget = node.widgets?.find(
|
||||
(item) => item.name === "credential_source",
|
||||
);
|
||||
if (widget && !SAFE_SOURCES.has(widget.value)) {
|
||||
widget.value = SAFE_SOURCE;
|
||||
widget.callback?.(SAFE_SOURCE);
|
||||
}
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name: "gokayfem.vlm.api-credential-security",
|
||||
async beforeConfigureGraph(graphData) {
|
||||
// Runs on the cloned workflow before LiteGraph creates any widgets, so a
|
||||
// legacy key never reaches a DOM input or the active graph.
|
||||
visitGraphNodes(graphData, scrubSerializedNode);
|
||||
},
|
||||
loadedGraphNode(node) {
|
||||
enforceLiveWidget(node);
|
||||
},
|
||||
});
|
||||
+179
-83
@@ -3,83 +3,151 @@ import { api } from "../../../scripts/api.js";
|
||||
|
||||
const OUTPUT_NAME = "output_text";
|
||||
const VIEW_TEXT_NODE = "ViewText";
|
||||
const MODERN_VLM_NODE = "ModernVLM";
|
||||
const STREAMING_SOURCE_NODES = new Set([
|
||||
"ModernVLM",
|
||||
"Moondream31Query",
|
||||
"Moondream31Caption",
|
||||
"PromptGenerateAPI",
|
||||
"HostedVLMAPI",
|
||||
"VLMVideoTemporalReasoner",
|
||||
]);
|
||||
|
||||
function textMetrics(value) {
|
||||
const text = String(value ?? "");
|
||||
const words = text.trim() ? text.trim().split(/\s+/u).length : 0;
|
||||
const lines = text ? text.split("\n").length : 0;
|
||||
return `${text.length.toLocaleString()} chars · ${words.toLocaleString()} words · ${lines.toLocaleString()} lines`;
|
||||
}
|
||||
|
||||
function makeButton(label, title, handler) {
|
||||
const button = document.createElement("button");
|
||||
button.textContent = label;
|
||||
button.type = "button";
|
||||
button.title = title;
|
||||
button.addEventListener("click", handler);
|
||||
Object.assign(button.style, {
|
||||
color: "var(--input-text, #ddd)",
|
||||
background: "var(--comfy-input-bg, #202020)",
|
||||
border: "1px solid var(--border-color, #555)",
|
||||
borderRadius: "5px",
|
||||
padding: "3px 8px",
|
||||
cursor: "pointer",
|
||||
whiteSpace: "nowrap",
|
||||
});
|
||||
return button;
|
||||
}
|
||||
|
||||
function ensureOutputWidget(node) {
|
||||
let widget = node.widgets?.find((item) => item.name === OUTPUT_NAME);
|
||||
if (!widget) {
|
||||
const container = document.createElement("div");
|
||||
const header = document.createElement("div");
|
||||
const status = document.createElement("span");
|
||||
const copy = document.createElement("button");
|
||||
const output = document.createElement("textarea");
|
||||
|
||||
status.textContent = "Ready";
|
||||
copy.textContent = "Copy";
|
||||
copy.type = "button";
|
||||
copy.title = "Copy the complete VLM response";
|
||||
copy.addEventListener("click", async () => {
|
||||
const previous = copy.textContent;
|
||||
try {
|
||||
await navigator.clipboard.writeText(output.value);
|
||||
copy.textContent = "Copied";
|
||||
} catch {
|
||||
copy.textContent = "Copy failed";
|
||||
}
|
||||
window.setTimeout(() => {
|
||||
copy.textContent = previous;
|
||||
}, 1200);
|
||||
});
|
||||
|
||||
output.readOnly = true;
|
||||
output.setAttribute("aria-label", "VLM text output");
|
||||
header.append(status, copy);
|
||||
container.append(header, output);
|
||||
Object.assign(container.style, {
|
||||
display: "flex",
|
||||
flexDirection: "column",
|
||||
width: "100%",
|
||||
height: "100%",
|
||||
minHeight: "150px",
|
||||
gap: "6px",
|
||||
});
|
||||
Object.assign(header.style, {
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
justifyContent: "space-between",
|
||||
color: "var(--descrip-text, #aaa)",
|
||||
fontSize: "12px",
|
||||
});
|
||||
Object.assign(copy.style, {
|
||||
color: "var(--input-text, #ddd)",
|
||||
background: "var(--comfy-input-bg, #202020)",
|
||||
border: "1px solid var(--border-color, #555)",
|
||||
borderRadius: "5px",
|
||||
padding: "3px 9px",
|
||||
cursor: "pointer",
|
||||
});
|
||||
Object.assign(output.style, {
|
||||
width: "100%",
|
||||
flex: "1",
|
||||
minHeight: "120px",
|
||||
resize: "vertical",
|
||||
boxSizing: "border-box",
|
||||
color: "var(--input-text, #ddd)",
|
||||
background: "var(--comfy-input-bg, #202020)",
|
||||
border: "1px solid var(--border-color, #555)",
|
||||
borderRadius: "6px",
|
||||
padding: "8px",
|
||||
lineHeight: "1.45",
|
||||
whiteSpace: "pre-wrap",
|
||||
});
|
||||
widget = node.addDOMWidget(OUTPUT_NAME, "STRING", container, {
|
||||
serialize: false,
|
||||
hideOnZoom: false,
|
||||
});
|
||||
widget.serialize = false;
|
||||
widget.inputEl = output;
|
||||
widget.statusEl = status;
|
||||
if (widget) {
|
||||
return widget;
|
||||
}
|
||||
|
||||
const container = document.createElement("div");
|
||||
const header = document.createElement("div");
|
||||
const status = document.createElement("span");
|
||||
const meta = document.createElement("span");
|
||||
const actions = document.createElement("div");
|
||||
const output = document.createElement("textarea");
|
||||
|
||||
status.textContent = "Ready";
|
||||
meta.textContent = textMetrics("");
|
||||
output.readOnly = true;
|
||||
output.wrap = "soft";
|
||||
output.spellcheck = false;
|
||||
output.setAttribute("aria-label", "VLM text output");
|
||||
|
||||
const copy = makeButton("Copy", "Copy complete text", async () => {
|
||||
const previous = copy.textContent;
|
||||
try {
|
||||
await navigator.clipboard.writeText(output.value);
|
||||
copy.textContent = "Copied";
|
||||
} catch {
|
||||
copy.textContent = "Copy failed";
|
||||
}
|
||||
window.setTimeout(() => {
|
||||
copy.textContent = previous;
|
||||
}, 1200);
|
||||
});
|
||||
const download = makeButton("Save", "Download output as a UTF-8 text file", () => {
|
||||
const blob = new Blob([output.value], {
|
||||
type: "text/plain;charset=utf-8",
|
||||
});
|
||||
const url = URL.createObjectURL(blob);
|
||||
const anchor = document.createElement("a");
|
||||
anchor.href = url;
|
||||
anchor.download = `vlm-output-${new Date().toISOString().replaceAll(":", "-")}.txt`;
|
||||
anchor.click();
|
||||
URL.revokeObjectURL(url);
|
||||
});
|
||||
const wrap = makeButton("Wrap: on", "Toggle long-line wrapping", () => {
|
||||
const enabled = output.wrap !== "off";
|
||||
output.wrap = enabled ? "off" : "soft";
|
||||
output.style.whiteSpace = enabled ? "pre" : "pre-wrap";
|
||||
output.style.overflowX = enabled ? "auto" : "hidden";
|
||||
wrap.textContent = enabled ? "Wrap: off" : "Wrap: on";
|
||||
});
|
||||
const follow = makeButton("Follow: on", "Follow streaming output", () => {
|
||||
widget.followOutput = !widget.followOutput;
|
||||
follow.textContent = widget.followOutput ? "Follow: on" : "Follow: off";
|
||||
});
|
||||
|
||||
actions.append(wrap, follow, copy, download);
|
||||
header.append(status, meta, actions);
|
||||
container.append(header, output);
|
||||
|
||||
Object.assign(container.style, {
|
||||
display: "flex",
|
||||
flexDirection: "column",
|
||||
width: "100%",
|
||||
height: "100%",
|
||||
minHeight: "190px",
|
||||
gap: "6px",
|
||||
});
|
||||
Object.assign(header.style, {
|
||||
display: "grid",
|
||||
gridTemplateColumns: "auto minmax(0, 1fr) auto",
|
||||
alignItems: "center",
|
||||
gap: "9px",
|
||||
color: "var(--descrip-text, #aaa)",
|
||||
fontSize: "11px",
|
||||
});
|
||||
Object.assign(meta.style, {
|
||||
overflow: "hidden",
|
||||
textOverflow: "ellipsis",
|
||||
whiteSpace: "nowrap",
|
||||
});
|
||||
Object.assign(actions.style, {
|
||||
display: "flex",
|
||||
gap: "4px",
|
||||
justifyContent: "flex-end",
|
||||
});
|
||||
Object.assign(output.style, {
|
||||
width: "100%",
|
||||
flex: "1",
|
||||
minHeight: "160px",
|
||||
resize: "vertical",
|
||||
boxSizing: "border-box",
|
||||
color: "var(--input-text, #ddd)",
|
||||
background: "var(--comfy-input-bg, #202020)",
|
||||
border: "1px solid var(--border-color, #555)",
|
||||
borderRadius: "6px",
|
||||
padding: "9px",
|
||||
lineHeight: "1.45",
|
||||
whiteSpace: "pre-wrap",
|
||||
overflowWrap: "anywhere",
|
||||
tabSize: "4",
|
||||
});
|
||||
|
||||
widget = node.addDOMWidget(OUTPUT_NAME, "STRING", container, {
|
||||
serialize: false,
|
||||
hideOnZoom: false,
|
||||
});
|
||||
widget.serialize = false;
|
||||
widget.inputEl = output;
|
||||
widget.statusEl = status;
|
||||
widget.metaEl = meta;
|
||||
widget.followOutput = true;
|
||||
return widget;
|
||||
}
|
||||
|
||||
@@ -88,8 +156,10 @@ function setOutput(node, text, state = "Complete") {
|
||||
const value = Array.isArray(text) ? text.join("\n\n") : String(text ?? "");
|
||||
widget.value = value;
|
||||
widget.inputEl.value = value;
|
||||
if (widget.statusEl) {
|
||||
widget.statusEl.textContent = state;
|
||||
widget.statusEl.textContent = state;
|
||||
widget.metaEl.textContent = textMetrics(value);
|
||||
if (widget.followOutput && state === "Streaming…") {
|
||||
widget.inputEl.scrollTop = widget.inputEl.scrollHeight;
|
||||
}
|
||||
node.setDirtyCanvas?.(true, true);
|
||||
}
|
||||
@@ -104,18 +174,38 @@ function findNode(graph, id) {
|
||||
?? null;
|
||||
}
|
||||
|
||||
function linkFor(graph, linkId) {
|
||||
return graph?.links?.get?.(linkId)
|
||||
?? graph?._links?.get?.(linkId)
|
||||
?? null;
|
||||
}
|
||||
|
||||
function isReroute(node) {
|
||||
return String(node?.type ?? "").toLowerCase().includes("reroute");
|
||||
}
|
||||
|
||||
function connectedViewTextNodes(source) {
|
||||
if (!source?.graph) {
|
||||
return [];
|
||||
}
|
||||
const found = new Set();
|
||||
for (const output of source.outputs ?? []) {
|
||||
for (const linkId of output.links ?? []) {
|
||||
const link = source.graph.links?.get?.(linkId)
|
||||
?? source.graph._links?.get?.(linkId);
|
||||
const target = findNode(source.graph, link?.target_id);
|
||||
if (target?.type === VIEW_TEXT_NODE) {
|
||||
found.add(target);
|
||||
const visited = new Set([source.id]);
|
||||
const queue = [source];
|
||||
while (queue.length) {
|
||||
const current = queue.shift();
|
||||
for (const output of current.outputs ?? []) {
|
||||
for (const linkId of output.links ?? []) {
|
||||
const link = linkFor(source.graph, linkId);
|
||||
const target = findNode(source.graph, link?.target_id);
|
||||
if (!target || visited.has(target.id)) {
|
||||
continue;
|
||||
}
|
||||
visited.add(target.id);
|
||||
if (target.type === VIEW_TEXT_NODE) {
|
||||
found.add(target);
|
||||
} else if (isReroute(target)) {
|
||||
queue.push(target);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -131,7 +221,7 @@ function updateFromProgress({ nodeId, text }) {
|
||||
setOutput(source, text, "Streaming…");
|
||||
return;
|
||||
}
|
||||
if (source.type !== MODERN_VLM_NODE) {
|
||||
if (!STREAMING_SOURCE_NODES.has(source.type)) {
|
||||
return;
|
||||
}
|
||||
for (const target of connectedViewTextNodes(source)) {
|
||||
@@ -154,6 +244,12 @@ app.registerExtension({
|
||||
nodeType.prototype.onNodeCreated = function (...args) {
|
||||
const result = onCreated?.apply(this, args);
|
||||
ensureOutputWidget(this);
|
||||
if (Array.isArray(this.size)) {
|
||||
this.setSize?.([
|
||||
Math.max(this.size[0], 430),
|
||||
Math.max(this.size[1], 290),
|
||||
]);
|
||||
}
|
||||
return result;
|
||||
};
|
||||
const onExecuted = nodeType.prototype.onExecuted;
|
||||
|
||||
Reference in New Issue
Block a user