2 Commits
Author SHA1 Message Date
Gökay Aydoğan 2a22af273f Update requirements.txt 2024-11-06 18:53:28 +03:00
gokayfem 23bd87c30a req 2024-11-06 18:43:05 +03:00
113 changed files with 5233 additions and 34984 deletions
-104
View File
@@ -1,104 +0,0 @@
name: Bug report
description: A node fails, errors, or produces wrong output.
labels: ["bug"]
body:
- type: markdown
attributes:
value: |
Most unresolvable reports are missing the environment details below.
Please run the **VLM Runtime Diagnostics** node and paste its output —
it captures your OS, Python, PyTorch, accelerator backend, and which
optional backends are installed.
- type: input
id: version
attributes:
label: Node pack version
description: From ComfyUI Manager, or the `version` in `pyproject.toml`.
placeholder: "3.3.1"
validations:
required: true
- type: dropdown
id: install
attributes:
label: How did you install it?
options:
- ComfyUI Manager
- Comfy Registry
- git clone into custom_nodes
- Other (describe below)
validations:
required: true
- type: dropdown
id: comfy
attributes:
label: ComfyUI flavour
options:
- ComfyUI Desktop
- ComfyUI Portable (python_embeded)
- Manual install (venv)
- Manual install (conda)
- Cloud / RunPod / other host
validations:
required: true
- type: textarea
id: diagnostics
attributes:
label: VLM Runtime Diagnostics output
description: Add the node to any workflow, run it, and paste the result.
render: text
validations:
required: true
- type: input
id: node
attributes:
label: Which node fails?
placeholder: "LLMSampler, LLavaSamplerSimple, ModernVLM, ..."
validations:
required: true
- type: input
id: model
attributes:
label: Which model / GGUF file?
description: Include the exact filename or Hugging Face repo id.
placeholder: "Qwen 3 VL 4B Instruct, or llava-1.6-mistral-7b.Q4_K_M.gguf"
validations:
required: true
- type: textarea
id: expected
attributes:
label: What did you expect, and what happened instead?
validations:
required: true
- type: textarea
id: traceback
attributes:
label: Full console output
description: |
The complete traceback from the ComfyUI terminal, not just the last
line. Include the startup log if the pack failed to import.
render: shell
validations:
required: true
- type: checkboxes
id: checks
attributes:
label: Before submitting
options:
- label: I updated to the latest version of this node pack and ComfyUI.
required: true
- label: I searched existing open and closed issues.
required: true
- label: >-
If this involves GGUF or `llama-cpp-python`, I installed it with
the arguments for my accelerator from the
[llama-cpp-python install guide](https://github.com/abetlen/llama-cpp-python#installation).
required: false
-11
View File
@@ -1,11 +0,0 @@
blank_issues_enabled: false
contact_links:
- name: llama-cpp-python installation help
url: https://github.com/abetlen/llama-cpp-python#installation
about: >-
Build or GPU-offload failures for GGUF nodes are almost always
llama-cpp-python installation issues. Install the wheel matching your
accelerator first.
- name: ComfyUI Manager and installation problems
url: https://github.com/Comfy-Org/ComfyUI-Manager/issues
about: For problems installing or updating custom nodes in general.
-46
View File
@@ -1,46 +0,0 @@
name: Model or feature request
description: Ask for support for a new VLM/LLM, or a new node.
labels: ["enhancement"]
body:
- type: textarea
id: what
attributes:
label: What would you like added?
validations:
required: true
- type: input
id: model
attributes:
label: Model repository (if requesting a model)
description: A Hugging Face repo id, so the architecture can be checked.
placeholder: "Qwen/Qwen3-VL-8B-Instruct"
- type: dropdown
id: backend
attributes:
label: Which backend would it use?
options:
- transformers (safetensors)
- llama.cpp (GGUF)
- Hosted API
- Not sure
validations:
required: true
- type: textarea
id: why
attributes:
label: What does it let you do that current nodes cannot?
validations:
required: true
- type: checkboxes
id: checks
attributes:
label: Before submitting
options:
- label: >-
I checked the README node reference to confirm this is not already
supported.
required: true
-33
View File
@@ -1,33 +0,0 @@
## What does this change?
<!-- One or two sentences. Link any issue it closes: "Closes #123". -->
## Type of change
- [ ] Bug fix
- [ ] New model support
- [ ] New node
- [ ] Refactor / maintenance
- [ ] Documentation
## Checklist
- [ ] `python -m pytest -q` passes.
- [ ] `python -m ruff check .` passes.
- [ ] Importing the pack still performs no network access, compilation, or
package install.
- [ ] If a node schema changed, existing widget order is preserved (Comfy
serializes widget values by position, so reordering breaks saved
workflows).
- [ ] New optional dependencies fail only the node that needs them, with an
actionable error.
- [ ] `pyproject.toml` `version` is bumped if this is user-visible, and
`CHANGELOG.md` has an entry. Releases only publish on a version change.
## Testing
<!--
Which nodes did you run, on which backend (CUDA / ROCm / Metal / XPU / CPU),
and with which model? Real-weight checks are opt-in:
python tests/manual_model_smoke.py --model "Qwen 3 VL 4B Instruct"
-->
-14
View File
@@ -1,14 +0,0 @@
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: ["*"]
-96
View File
@@ -1,96 +0,0 @@
name: CI
on:
push:
pull_request:
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 }}
timeout-minutes: 35
strategy:
fail-fast: false
matrix:
include:
- label: Linux / Python 3.10
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
with:
python-version: ${{ matrix.python }}
cache: pip
cache-dependency-path: requirements.txt
- name: Install CPU PyTorch
if: matrix.cpu_index == true
run: |
python -m pip install --upgrade pip
python -m pip install torch --index-url https://download.pytorch.org/whl/cpu
- name: Install macOS PyTorch
if: matrix.cpu_index == false
run: |
python -m pip install --upgrade pip
python -m pip install torch
- name: Install ComfyUI and node dependencies
run: |
git clone --depth 1 https://github.com/Comfy-Org/ComfyUI.git ../ComfyUI
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
-171
View File
@@ -1,171 +0,0 @@
name: Publish Comfy node fleet
on:
workflow_dispatch:
inputs:
target:
description: Node repository to check
required: true
default: all
type: choice
options:
- all
- vlm
- depth
- dream
- texture
schedule:
- cron: "17 * * * *"
push:
branches:
- main
paths:
- ".github/workflows/publish-fleet.yml"
permissions:
contents: read
concurrency:
group: comfy-registry-fleet
cancel-in-progress: false
jobs:
publish:
name: Check ${{ matrix.target }}
runs-on: ubuntu-latest
timeout-minutes: 10
strategy:
fail-fast: false
matrix:
include:
- target: vlm
repository: gokayfem/ComfyUI_VLM_nodes
node_id: comfyui_vlm_nodes
- target: depth
repository: gokayfem/ComfyUI-Depth-Visualization
node_id: comfyui-depth-visualization
- target: dream
repository: gokayfem/ComfyUI-Dream-Interpreter
node_id: comfyui-dream-interpreter
- target: texture
repository: gokayfem/ComfyUI-Texture-Simple
node_id: comfyui-texture-simple
steps:
- name: Select target
id: select
env:
REQUESTED_TARGET: ${{ inputs.target || 'all' }}
MATRIX_TARGET: ${{ matrix.target }}
run: |
if [[ "$REQUESTED_TARGET" == "all" || "$REQUESTED_TARGET" == "$MATRIX_TARGET" ]]; then
echo "selected=true" >> "$GITHUB_OUTPUT"
else
echo "selected=false" >> "$GITHUB_OUTPUT"
fi
- name: Check out node
if: steps.select.outputs.selected == 'true'
uses: actions/checkout@v7
with:
repository: ${{ matrix.repository }}
ref: main
path: node
persist-credentials: false
- name: Set up Python
if: steps.select.outputs.selected == 'true'
uses: actions/setup-python@v7
with:
python-version: "3.12"
- name: Read and verify release metadata
if: steps.select.outputs.selected == 'true'
id: metadata
working-directory: node
env:
EXPECTED_NODE_ID: ${{ matrix.node_id }}
run: |
python - <<'PY'
import os
import tomllib
from pathlib import Path
metadata = tomllib.loads(Path("pyproject.toml").read_text(encoding="utf-8"))
node_id = metadata["project"]["name"]
version = metadata["project"]["version"]
publisher = metadata["tool"]["comfy"]["PublisherId"]
expected = os.environ["EXPECTED_NODE_ID"]
if node_id != expected:
raise SystemExit(f"Expected node id {expected!r}, found {node_id!r}")
if publisher != "gokayfem":
raise SystemExit(f"Expected publisher 'gokayfem', found {publisher!r}")
with Path(os.environ["GITHUB_OUTPUT"]).open("a", encoding="utf-8") as output:
print(f"node_id={node_id}", file=output)
print(f"version={version}", file=output)
PY
- name: Check Registry version
if: steps.select.outputs.selected == 'true'
id: registry
env:
NODE_ID: ${{ steps.metadata.outputs.node_id }}
VERSION: ${{ steps.metadata.outputs.version }}
run: |
python - <<'PY'
import json
import os
import urllib.parse
import urllib.request
from pathlib import Path
node_id = urllib.parse.quote(os.environ["NODE_ID"], safe="")
url = f"https://api.comfy.org/nodes/{node_id}/versions"
request = urllib.request.Request(
url,
headers={"Accept": "application/json", "User-Agent": "comfy-node-fleet-publisher"},
)
with urllib.request.urlopen(request, timeout=30) as response:
versions = json.load(response)
wanted = os.environ["VERSION"]
exists = any(item.get("version") == wanted for item in versions)
with Path(os.environ["GITHUB_OUTPUT"]).open("a", encoding="utf-8") as output:
print(f"exists={'true' if exists else 'false'}", file=output)
PY
- name: Require publisher credential
if: steps.select.outputs.selected == 'true' && steps.registry.outputs.exists != 'true'
env:
REGISTRY_ACCESS_TOKEN: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
run: |
if [[ -z "$REGISTRY_ACCESS_TOKEN" ]]; then
echo "::error title=Missing registry token::Add the publisher API key as the REGISTRY_ACCESS_TOKEN repository secret."
exit 1
fi
- name: Install pinned publisher
if: steps.select.outputs.selected == 'true' && steps.registry.outputs.exists != 'true'
run: python -m pip install --disable-pip-version-check --no-input "comfy-cli==1.13.0"
- name: Publish missing version
if: steps.select.outputs.selected == 'true' && steps.registry.outputs.exists != 'true'
working-directory: node
env:
REGISTRY_ACCESS_TOKEN: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
run: comfy --skip-prompt --no-enable-telemetry node publish --token "$REGISTRY_ACCESS_TOKEN"
- name: Record result
if: steps.select.outputs.selected == 'true'
env:
NODE_ID: ${{ steps.metadata.outputs.node_id }}
VERSION: ${{ steps.metadata.outputs.version }}
ALREADY_PUBLISHED: ${{ steps.registry.outputs.exists }}
run: |
if [[ "$ALREADY_PUBLISHED" == "true" ]]; then
echo "### $NODE_ID $VERSION already published" >> "$GITHUB_STEP_SUMMARY"
else
echo "### Published $NODE_ID $VERSION" >> "$GITHUB_STEP_SUMMARY"
fi
+5 -98
View File
@@ -6,109 +6,16 @@ on:
- main
paths:
- "pyproject.toml"
- ".github/workflows/publish.yml"
concurrency:
group: comfy-registry-${{ github.repository }}
cancel-in-progress: false
env:
COMFY_CLI_VERSION: "1.13.0"
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
timeout-minutes: 10
permissions:
contents: read
steps:
- name: Check out code
uses: actions/checkout@v7
- name: Set up Python
uses: actions/setup-python@v7
with:
python-version: "3.12"
- name: Read release metadata
id: metadata
run: |
python - <<'PY'
import os
import tomllib
from pathlib import Path
metadata = tomllib.loads(Path("pyproject.toml").read_text(encoding="utf-8"))
node_id = metadata["project"]["name"]
version = metadata["project"]["version"]
publisher = metadata["tool"]["comfy"]["PublisherId"]
if publisher != "gokayfem":
raise SystemExit(f"Expected publisher 'gokayfem', found {publisher!r}")
with Path(os.environ["GITHUB_OUTPUT"]).open("a", encoding="utf-8") as output:
print(f"node_id={node_id}", file=output)
print(f"version={version}", file=output)
PY
- name: Check Registry version
id: registry
env:
NODE_ID: ${{ steps.metadata.outputs.node_id }}
VERSION: ${{ steps.metadata.outputs.version }}
run: |
python - <<'PY'
import json
import os
import urllib.parse
import urllib.request
from pathlib import Path
node_id = urllib.parse.quote(os.environ["NODE_ID"], safe="")
request = urllib.request.Request(
f"https://api.comfy.org/nodes/{node_id}/versions",
headers={"Accept": "application/json", "User-Agent": "comfy-node-publisher"},
)
with urllib.request.urlopen(request, timeout=30) as response:
versions = json.load(response)
exists = any(item.get("version") == os.environ["VERSION"] for item in versions)
with Path(os.environ["GITHUB_OUTPUT"]).open("a", encoding="utf-8") as output:
print(f"exists={'true' if exists else 'false'}", file=output)
PY
- name: Check publisher credential
if: steps.registry.outputs.exists != 'true'
id: credentials
env:
REGISTRY_ACCESS_TOKEN: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
run: |
if [[ -n "$REGISTRY_ACCESS_TOKEN" ]]; then
echo "available=true" >> "$GITHUB_OUTPUT"
else
echo "available=false" >> "$GITHUB_OUTPUT"
echo "::notice title=Central publisher enabled::The secure fleet publisher will publish this release within one hour."
fi
- name: Install pinned Comfy CLI
if: steps.registry.outputs.exists != 'true' && steps.credentials.outputs.available == 'true'
shell: bash
run: python -m pip install --disable-pip-version-check "comfy-cli==${COMFY_CLI_VERSION}"
uses: actions/checkout@v4
- name: Publish Custom Node
if: steps.registry.outputs.exists != 'true' && steps.credentials.outputs.available == 'true'
id: publish
continue-on-error: true
shell: bash
env:
REGISTRY_ACCESS_TOKEN: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
run: comfy --skip-prompt --no-enable-telemetry node publish --token "$REGISTRY_ACCESS_TOKEN"
- name: Record publication result
env:
NODE_ID: ${{ steps.metadata.outputs.node_id }}
VERSION: ${{ steps.metadata.outputs.version }}
ALREADY_PUBLISHED: ${{ steps.registry.outputs.exists }}
PUBLISH_OUTCOME: ${{ steps.publish.outcome }}
run: |
if [[ "$ALREADY_PUBLISHED" == "true" ]]; then
echo "### $NODE_ID $VERSION already published" >> "$GITHUB_STEP_SUMMARY"
elif [[ "$PUBLISH_OUTCOME" == "success" ]]; then
echo "### Published $NODE_ID $VERSION" >> "$GITHUB_STEP_SUMMARY"
else
echo "::notice title=Central publishing handoff::The secure fleet publisher will retry this release within one hour."
echo "### $NODE_ID $VERSION queued for the fleet publisher" >> "$GITHUB_STEP_SUMMARY"
fi
uses: Comfy-Org/publish-node-action@main
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
-199
View File
@@ -1,199 +0,0 @@
# Changelog
All notable changes to this project are documented here.
The format follows [Keep a Changelog](https://keepachangelog.com/en/1.1.0/),
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
Versions are published to the [Comfy Registry](https://registry.comfy.org/)
from `pyproject.toml`. A release is only published when `version` changes, so
every user-visible fix needs a version bump.
## [3.5.0] - 2026-07-31
### Added
- A MiniMax music node with fixed global and China endpoints, generation and
cover model selection, regional request fields, URL and hexadecimal response
decoding, and MP3, WAV, and PCM output through the existing audio contract.
### Security
- MiniMax credentials are read only from `MINIMAX_API_KEY`; workflows cannot
supply a key or redirect it to a custom endpoint, and request errors redact
the resolved value before reaching ComfyUI.
## [3.4.0] - 2026-07-31
### Added
- A robotics-safe VLA layer with typed embodiment, observation, and action
contracts; bounded multi-camera history; trajectory inspection and preview;
action-chunk replanning; and explicit bounds, rate, dimension, horizon, and
non-finite-value checks before handoff.
- Native policy clients for OpenPI's WebSocket protocol and NVIDIA Isaac
GR00T's ZeroMQ protocol, plus a portable authenticated HTTP/JPEG protocol for
isolated policy runtimes.
- An isolated current-LeRobot policy server with pre/postprocessor support,
serialized inference, optional idle CPU offload, checkpoint feature metadata,
and environment-only bearer authentication.
- A curated 15-model VLA catalog covering SmolVLA, X-VLA, the OpenPI family,
GR00T N1.7, WALL-OSS, MolmoAct2, VLA-JEPA, LingBot-VA, FastWAM, EO-1, EVO-1,
OpenVLA-OFT, and Octo with explicit readiness and fine-tuning requirements.
- A complete API workflow, setup guide, compatibility matrix, security
guidance, and real-weight SmolVLA validation on an RTX 3090.
### Security
- Workflow JSON never stores robotics API keys. The clients read only
`VLA_POLICY_TOKEN`, `OPENPI_API_KEY`, or `GROOT_API_TOKEN` from the
environment, redact them from errors/reports, reject embedded URL
credentials, and require encrypted transports plus explicit opt-in for
remote endpoints where the upstream protocol supports encryption.
- The included policy server bounds request, camera, history, and response
sizes and never uses pickle across the network.
## [3.3.1] - 2026-07-30
### Fixed
- Package metadata declared `license = "MIT"` while the bundled `LICENSE` has
been Apache-2.0 since the initial commit. Built wheels therefore contained
contradictory MIT metadata and Apache-2.0 license text. The Registry already
referenced the license file and was unaffected. Metadata now says
`Apache-2.0`.
- Moondream 2 and Moondream 3.1 local inference (`b8ae298`).
- SmolVLM setup dependencies (`c13ee23`).
The two fixes above landed on `main` after 3.3.0 without a version bump, so
the Registry publish workflow saw 3.3.0 already published and skipped them.
They reach Registry users for the first time in 3.3.1.
### Added
- Test coverage for the GGUF text and multimodal node families, which
previously had none: `nodes/suggest.py` (0% to 98%) and
`nodes/llavaloader.py` (0% to 99%). The new cases pin the behaviours behind
the pack's longest-running bug reports: widget ordering (#156), sampling
kwarg plumbing (#144), and handle teardown on both success and failure
(#137).
- `ruff` lint gate and a coverage floor in CI, plus `requirements-dev.txt`
for the tooling.
- `CHANGELOG.md`, `CONTRIBUTING.md`, issue and pull request templates, and a
Dependabot configuration.
- A complete node reference in the README covering all 78 registered nodes.
## [3.3.0] - 2026-07-29
### Added
- Moondream Photon support and universal VLM acceleration utilities, including
the image pixel-budget and performance-profile nodes (`102f166`).
## [3.2.0] - 2026-07-29
### Added
- Adaptive video intelligence with temporal reasoning, plus the text workflow
toolkit (join, template, clean, replace, split, JSON extract, inspect)
(`44fefcb`).
## [3.1.0] - 2026-07-29
### Changed
- Hosted LLM and VLM API nodes modernized and hardened, with provider profiles
for OpenAI, Google Gemini, Anthropic, xAI, DeepSeek, and others (`505b324`).
## [3.0.0] - 2026-07-29
### Added
- Unified vision stack: open-vocabulary detection (Grounding DINO, OWLv2,
OmDet), SAM2.1 and SAM3.1 segmentation, tracking, and creator mask tools,
with structured detection/segmentation schemas (`39fc116`).
### Changed
- **Breaking:** detection and segmentation nodes now emit structured data
types rather than loose strings. Workflows wiring these outputs into text
nodes need the new converter utilities.
## [2.3.0] - 2026-07-29
### Added
- Reliable streaming VLM text output (`239c904`).
## [2.2.0] - 2026-07-29
### Changed
- llama.cpp GGUF runtime modernized. `llama-cpp-agent` was removed in favour
of llama-cpp-python's native JSON Schema support, which resolves the
unstable wrapper API behind the `unexpected keyword argument 'temperature'`
crashes (#144).
## [2.1.0] - 2026-07-28
### Added
- Cross-platform runtime support across NVIDIA CUDA, AMD ROCm, Apple Metal,
Intel XPU, and CPU, without replacing ComfyUI's PyTorch (`4c200c4`).
## [2.0.1] - 2026-07-28
### Added
- Small VLM catalog and real-weight model validation evidence
(see `MODEL_VALIDATION.md`) (`460b27a`).
## [2.0.0] - 2026-07-28
### Changed
- **Breaking:** node pack modernized with an explicit GPU lifecycle. Models
now load lazily on first execution and register with ComfyUI's model manager
so they participate in smart VRAM offloading, which addresses models
remaining resident after generation (#137) (`b89f628`).
- **Breaking:** `forceInput` string hacks removed from node schemas. They
corrupted the widget index during serialization and shifted inputs on saved
workflows (#156). Use the native right-click "Convert to Input" instead.
- Import is now failure-isolated: a broken optional model cannot prevent
unrelated nodes from loading (#94, #145).
- `numpy` is no longer pinned. The old `numpy<2.0.0` pin crashed startup on
NumPy 2.x environments (#157).
- Model coverage moved to current releases, including Qwen 3 / 3.5 VL,
SmolVLM2, InternVL, Granite Vision, and Gemma 3 (#148, #151). The
unmaintained InternLM-XComposer2 nodes were dropped (#139).
### Removed
- **Breaking:** `llama-cpp-agent` dependency (see 2.2.0).
- **Breaking:** InternLM-XComposer2 nodes, which depended on an AutoGPTQ stack
that pinned incompatible PyTorch versions (#139).
## 1.0.0 - 1.0.6 (2024-05-20 to 2024-11-03)
Initial packaged releases, predating changelog tracking. This line covered
LLaVA GGUF loaders and samplers, Moondream, Kosmos-2, JoyTag, UForm,
MiniCPM-V, PaLI-Gemma, Florence-2, Molmo, Qwen2-VL, the LLM prompt and
suggestion generators, AudioLDM2, and ChatMusician. See the
[commit history](https://github.com/gokayfem/ComfyUI_VLM_nodes/commits/main)
for detail.
Tagging began at 3.3.0. Earlier versions link to the commit that declared
them, because retroactively tagging them would run current CI against code
that predates it.
[3.4.0]: https://github.com/gokayfem/ComfyUI_VLM_nodes/compare/v3.3.1...v3.4.0
[3.3.1]: https://github.com/gokayfem/ComfyUI_VLM_nodes/compare/v3.3.0...v3.3.1
[3.3.0]: https://github.com/gokayfem/ComfyUI_VLM_nodes/releases/tag/v3.3.0
[3.2.0]: https://github.com/gokayfem/ComfyUI_VLM_nodes/commit/44fefcb
[3.1.0]: https://github.com/gokayfem/ComfyUI_VLM_nodes/commit/505b324
[3.0.0]: https://github.com/gokayfem/ComfyUI_VLM_nodes/commit/39fc116
[2.3.0]: https://github.com/gokayfem/ComfyUI_VLM_nodes/commit/239c904
[2.2.0]: https://github.com/gokayfem/ComfyUI_VLM_nodes/commit/0da5070
[2.1.0]: https://github.com/gokayfem/ComfyUI_VLM_nodes/commit/4c200c4
[2.0.1]: https://github.com/gokayfem/ComfyUI_VLM_nodes/commit/460b27a
[2.0.0]: https://github.com/gokayfem/ComfyUI_VLM_nodes/commit/b89f628
-21
View File
@@ -1,21 +0,0 @@
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"
-293
View File
@@ -1,293 +0,0 @@
# Platform and accelerator compatibility
ComfyUI owns PyTorch. This node pack deliberately does not depend on `torch`,
`torchvision`, or a vendor wheel, because installing a generic PyPI build can
silently replace a working CUDA, ROCm, XPU, or Metal environment.
Install `requirements.txt` with the same Python executable that starts ComfyUI.
The **VLM Runtime Diagnostics** node reports the environment seen by the pack
without downloading a model.
## Support matrix
| Platform | Managed Transformers | bitsandbytes 4/8-bit | GGUF acceleration |
| --- | --- | --- | --- |
| Linux + NVIDIA | CUDA, BF16/FP16 capability detected | Official wheel | CUDA or Vulkan |
| Windows + NVIDIA | CUDA, BF16/FP16 capability detected | Official wheel | CUDA or Vulkan |
| Linux + AMD | ROCm through PyTorch's `cuda` API | Official ROCm wheel for listed GPU architectures | ROCm/HIP or Vulkan |
| Windows + AMD | Current ComfyUI/AMD ROCm PyTorch builds | Official ROCm Windows wheel for listed GPU architectures | HIP Radeon or Vulkan |
| Apple Silicon macOS | MPS, BF16 on supported macOS/PyTorch; FP16 fallback | Official arm64 wheel | Metal |
| Intel GPU | XPU with BF16 capability detection | Official XPU/CPU wheel | SYCL or Vulkan |
| CPU | FP32 | Official wheels on supported architectures | OpenBLAS or default CPU |
| Intel macOS | CPU/legacy MPS environment as provided by ComfyUI | No official bitsandbytes wheel; dependency is skipped | CPU build |
The default **ComfyUI managed** mode is the portable path. Quantization is an
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.
- `transformers>=5.4,<6` and `huggingface-hub>=1.5,<2` are paired intentionally;
Transformers 5.4 requires Hub 1.5 or newer.
- `bitsandbytes>=0.50` is the first dependency floor used here for the current
multi-backend releases. Environment markers prevent an unsupported wheel
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
python -m pip install -r ComfyUI/custom_nodes/ComfyUI_VLM_nodes/requirements.txt
```
If quantization was skipped but the machine has a supported custom build:
```bash
python -m pip install -r ComfyUI/custom_nodes/ComfyUI_VLM_nodes/requirements-quantization.txt
```
## llama.cpp / GGUF
`llama-cpp-python` must be compiled or selected for the actual backend. Its
official project currently publishes backend indexes and documents source
build flags:
```bash
# NVIDIA; choose a wheel supported by the installed driver.
python -m pip install llama-cpp-python \
--extra-index-url https://abetlen.github.io/llama-cpp-python/whl/cu124
# Apple Metal
python -m pip install llama-cpp-python \
--extra-index-url https://abetlen.github.io/llama-cpp-python/whl/metal
# Linux ROCm
python -m pip install llama-cpp-python \
--extra-index-url https://abetlen.github.io/llama-cpp-python/whl/rocm72
# Linux or Windows Vulkan
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
python -m pip install llama-cpp-python `
--extra-index-url https://abetlen.github.io/llama-cpp-python/whl/hip-radeon
```
Source builds use `GGML_CUDA=on`, `GGML_METAL=on`, `GGML_HIP=on`,
`GGML_VULKAN=on`, or `GGML_SYCL=on` through `CMAKE_ARGS`. Use an arm64 Python
on Apple Silicon; an x86 Python builds the wrong architecture and is
dramatically slower.
The llama.cpp wheel is an independent native runtime; it does not have to use
the same accelerator API as ComfyUI's PyTorch wheel. For example, a Vulkan
llama.cpp wheel can coexist with a CUDA or CPU PyTorch build. The nodes query
`llama_supports_gpu_offload`, `llama_supports_mmap`, and llama.cpp's system
information at runtime. They never label a wheel CUDA/ROCm/Metal based only on
`torch`.
### GGUF runtime controls
- `gpu_layers=-1` requests full accelerator offload. A build that reports no
offload support is automatically clamped to `0` and continues on CPU.
- `n_batch` is the logical prompt batch and `n_ubatch` is the physical
micro-batch. The runtime clamps both to the selected context and guarantees
`n_ubatch <= n_batch`.
- **Auto** flash attention enables the optimized path for accelerator offload
and retries once without it only when llama.cpp reports an attention-related
initialization failure. **Enabled** remains strict; **Disabled** is the
maximum-compatibility setting.
- `use_mmap` is honored only when the compiled backend reports mmap support.
- Layer, row, and single-device split modes plus `main_gpu` and
comma-separated `tensor_split` weights are passed through when supported by
the installed binding. Parallel multi-GPU is primarily a CUDA/ROCm feature;
Vulkan and SYCL support is more limited.
- Current multimodal GGUFs should use **Auto (GGUF chat template)**, which maps
to llama.cpp's MTMD handler. Named legacy handlers remain selectable for
model cards that require an exact prompt format.
- Every model handle is lazy, mutex-protected, cache-keyed by all performance
settings, and closes its exact model and projector handler on unload.
Authoritative installation references:
- [ComfyUI installation and hardware backends](https://github.com/Comfy-Org/ComfyUI)
- [bitsandbytes installation and supported hardware](https://huggingface.co/docs/bitsandbytes/installation)
- [llama-cpp-python supported backends](https://github.com/abetlen/llama-cpp-python#supported-backends)
- [llama-cpp-python API reference](https://llama-cpp-python.readthedocs.io/en/latest/api-reference/)
- [llama.cpp backend feature matrix](https://github.com/ggml-org/llama.cpp/wiki/Feature-matrix)
## Attention and offloading
- **Auto (SDPA)** lets PyTorch choose its maintained kernel and is the default
on every backend.
- **Flash Attention 2** is preflighted for CUDA/ROCm only. A compatible
`flash-attn` build is still required.
- ComfyUI-managed models participate in its normal model patcher lifecycle.
- External bitsandbytes and llama.cpp allocations ask ComfyUI to free space
first, then release only their owned model on unload.
- Automatic CPU/disk device mapping is used for large CUDA/ROCm/XPU models.
MPS unified memory and CPU use an explicit active-device map.
- AudioLDM2 uses FP16 on capable accelerators, FP32 on CPU, CUDA-API CPU
offload for NVIDIA/ROCm, and a portable CPU random generator on MPS.
## What CI proves
Every push installs current ComfyUI plus this complete `requirements.txt` and
runs imports, schemas, runtime contracts, tests, and byte-compilation on:
- Ubuntu, Python 3.10
- Ubuntu, Python 3.13
- Windows, Python 3.12
- macOS, Python 3.12
Hosted runners do not contain production NVIDIA, AMD, or Intel GPUs. CI
therefore tests backend selection and dtype/device-map contracts, while real
GPU model smoke tests remain explicit hardware validation. It does not claim
that a CPU simulation executed a vendor kernel.
-106
View File
@@ -1,106 +0,0 @@
# 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.
-186
View File
@@ -1,186 +0,0 @@
# Model validation
Validated on 2026-07-29 with ComfyUI 0.28.0, Python 3.12, Transformers 5.14.1,
PyTorch 2.13.0+cu126, and an RTX 3090 24 GB. All models and caches were stored
on the D drive and executed through WSL.
## Real-weight passes
| Family | Representative result | Peak CUDA |
| --- | --- | ---: |
| Qwen 3.5 | 0.8B BF16 image/video; 0.8B NF4; 2B, 4B, and 9B images | 0.82–17.62 GiB |
| Qwen 3 VL | 2B, 4B, and 8B images returned the correct red object | 3.99–16.37 GiB |
| SmolVLM2 | 500M image/video and 2.2B video returned the correct object | 2.29–5.41 GiB |
| LFM2.5 VL | 450M returned “red … rectangle” | 0.88 GiB |
| InternVL 3.5 | 1B video returned “green rectangle” after the 448px patch-grid fix | 2.14 GiB |
| Granite Vision 4.1 | 4B returned “solid red square” through native Transformers code | 7.61 GiB |
| Florence-2 | Native converted base-FT returned and parsed a bright-red-square caption | 0.59 GiB |
| llama.cpp GGUF | Official Qwen3.5-0.8B Q4_0 with llama-cpp-python 0.3.34 CUDA loaded in 20.015s and generated the exact requested response in 0.654s | < 1 GiB model weights |
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
node repositories enabled and no custom-node import failures. A real local API
workflow (`EmptyImage` -> `ModernVLM` -> `ViewText`) ran the cached LFM2.5-VL
450M checkpoint on a solid red input, returned `Red.`, and completed with
`unload_after=true`. Prompt ID:
`919f92cd-ecb2-487b-abf0-19f5e4d88229`.
A second real local API workflow (`LLMLoader` -> `LLMSampler` -> `ViewText`)
used the official 563 MB `ggml-org/Qwen3.5-0.8B-GGUF` Q4_0 checkpoint with
full GPU offload, `n_batch=256`, `n_ubatch=128`, mmap, and Auto flash
attention. It returned exactly `ComfyUI llama API ready` and completed
successfully. Prompt ID: `eed8458d-de7f-47ac-8ebf-e48e4dacc2d6`.
The installed llama.cpp CUDA 12.4 wheel reported GPU offload, mmap, and mlock
support directly. CPU-only fallback, Metal/Vulkan/SYCL/ROCm-independent
capability detection, multi-GPU options, and flash-attention retry are covered
by simulated backend contract tests; those vendor kernels were not claimed as
real hardware passes on the NVIDIA test machine.
## Catalog validation
Configuration and processor resolution passed for all 15 ungated entries in
the small/fast catalog: Qwen 3.5 0.8B/2B/4B, Qwen 3 VL 2B/4B, Qwen 2.5 VL 3B,
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
was stopped at the user's request and will not be repeated.
- Moondream2 2025-06-21: its pinned remote wrapper needed Transformers 5 loading
metadata, but this Torch/CUDA stack produced NaN probabilities when sampling
and immediate EOS with greedy decoding. The node defaults to the
non-destructive greedy path and raises an actionable error on an empty result.
- PaLI-Gemma and Gemma 3: gated checkpoints were not accessible without an
accepted license and token.
+127 -835
View File
File diff suppressed because it is too large Load Diff
-120
View File
@@ -1,120 +0,0 @@
# 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.
+50 -38
View File
@@ -1,67 +1,79 @@
import importlib.util
import os
import importlib
import logging
import pkg_resources
import sys
import subprocess
import folder_paths
from .nodes.runtime import register_model_folder
supported_LLava_extensions = set(['.gguf'])
LOGGER = logging.getLogger("ComfyUI_VLM_nodes")
register_model_folder()
try:
folder_paths.folder_names_and_paths["LLavacheckpoints"] = (folder_paths.folder_names_and_paths["LLavacheckpoints"][0], supported_LLava_extensions)
except:
# check if LLavacheckpoints exists otherwise create
if not os.path.isdir(os.path.join(folder_paths.models_dir, "LLavacheckpoints")):
os.mkdir(os.path.join(folder_paths.models_dir, "LLavacheckpoints"))
folder_paths.folder_names_and_paths["LLavacheckpoints"] = ([os.path.join(folder_paths.models_dir, "LLavacheckpoints")], supported_LLava_extensions)
# Define the check_requirements_installed function here or import it
def check_requirements_installed(requirements_path):
with open(requirements_path, 'r') as f:
requirements = [pkg_resources.Requirement.parse(line.strip()) for line in f if line.strip()]
installed_packages = {pkg.key: pkg for pkg in pkg_resources.working_set}
installed_packages_set = set(installed_packages.keys())
missing_packages = []
for requirement in requirements:
if requirement.key not in installed_packages_set or not installed_packages[requirement.key] in requirement:
missing_packages.append(str(requirement))
if missing_packages:
print(f"Missing or outdated packages: {', '.join(missing_packages)}")
print("Installing/Updating missing packages...")
subprocess.check_call([sys.executable, '-s', '-m', 'pip', 'install', *missing_packages])
else:
print("All packages from requirements.txt are installed and up to date.")
requirements_path = os.path.join(os.path.dirname(os.path.realpath(__file__)), "requirements.txt")
check_requirements_installed(requirements_path)
from .install_init import init, get_system_info, install_llama
system_info = get_system_info()
install_llama(system_info)
llama_cpp_agent_path = os.path.join(os.path.dirname(os.path.realpath(__file__)), "cpp_agent_req.txt")
check_requirements_installed(llama_cpp_agent_path)
init()
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 = {}
NODE_DISPLAY_NAME_MAPPINGS = {}
IMPORT_ERRORS = {}
for module_name in node_list:
try:
imported_module = importlib.import_module(f".nodes.{module_name}", __name__)
except Exception as exc:
# A broken optional model must never prevent unrelated nodes from loading.
IMPORT_ERRORS[module_name] = f"{type(exc).__name__}: {exc}"
LOGGER.exception("Could not load optional node module %s", module_name)
continue
NODE_CLASS_MAPPINGS.update(
getattr(imported_module, "NODE_CLASS_MAPPINGS", {})
)
NODE_DISPLAY_NAME_MAPPINGS.update(
getattr(imported_module, "NODE_DISPLAY_NAME_MAPPINGS", {})
)
imported_module = importlib.import_module(f".nodes.{module_name}", __name__)
NODE_CLASS_MAPPINGS = {**NODE_CLASS_MAPPINGS, **imported_module.NODE_CLASS_MAPPINGS}
NODE_DISPLAY_NAME_MAPPINGS = {**NODE_DISPLAY_NAME_MAPPINGS, **imported_module.NODE_DISPLAY_NAME_MAPPINGS}
WEB_DIRECTORY = "./web"
__all__ = [
"NODE_CLASS_MAPPINGS",
"NODE_DISPLAY_NAME_MAPPINGS",
"WEB_DIRECTORY",
]
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS", "WEB_DIRECTORY"]
+5
View File
@@ -0,0 +1,5 @@
llama-cpp-agent
mkdocs
mkdocs-material
mkdocstrings[python]
docstring-parser
-266
View File
@@ -1,266 +0,0 @@
# Robotics and VLA workflows
The robotics nodes make ComfyUI a policy-development, inspection, and
simulation surface. They do **not** send commands to motors, ROS, CAN, serial,
or a robot SDK.
The boundary is intentional:
```text
camera/state/task
|
v
VLA Observation Builder
|
v
isolated policy server ---> raw action chunk
|
v
VLA Action Safety Gate
|
+-------------+-------------+
| |
v v
inspect / plot / record simulator or your own
supervised controller bridge
```
A real controller bridge must independently enforce a deadman, watchdog,
emergency stop, collision/workspace limits, timestamps, command freshness, and
the manufacturer's limits. A `safe_for_handoff=true` workflow result only
means that the declared ComfyUI profile checks passed.
## Why policy runtimes are isolated
LeRobot, openpi, Isaac-GR00T, OpenVLA/OFT, and Octo use different PyTorch/JAX,
CUDA, Transformers, compiler, and operating-system combinations. Installing
all of those into ComfyUI would replace or constrain the working accelerator
stack and make Windows, macOS, ROCm, and XPU support worse.
The ComfyUI package therefore contains only:
- typed state/action/camera contracts;
- bounded image serialization;
- a dependency-light universal HTTPS/loopback HTTP client;
- exact clients for the official openpi MessagePack WebSocket and GR00T
MessagePack/ZeroMQ protocols;
- action validation, horizon control, inspection, and plotting.
The heavyweight policy stays in its own process, container, WSL distribution,
Linux machine, Mac, or GPU server. This also allows ComfyUI to use AMD ROCm,
Apple Metal, Intel XPU, or CPU while a policy runs on an NVIDIA Linux server.
## Fast path: SmolVLA through the universal sidecar
Use a separate LeRobot environment. Current LeRobot documentation recommends
Python 3.12 and exposes policy-specific extras. On this computer, keep it on
the D drive:
```bash
# WSL
python3.12 -m venv /mnt/d/vla-runtime/lerobot-smolvla
source /mnt/d/vla-runtime/lerobot-smolvla/bin/activate
python -m pip install --upgrade pip
python -m pip install "lerobot[smolvla]"
export VLA_POLICY_TOKEN="$(python -c 'import secrets; print(secrets.token_urlsafe(32))')"
python /mnt/d/ComfyUI_windows_portable/ComfyUI/custom_nodes/ComfyUI_VLM_nodes/examples/robotics/lerobot_policy_server.py \
--policy-type smolvla \
--policy-path YOUR_FINE_TUNED_SMOLVLA_CHECKPOINT \
--device auto \
--actions-per-chunk 16 \
--idle-offload-seconds 300
```
Set the same `VLA_POLICY_TOKEN` in the environment that launches ComfyUI.
Never put it in a workflow. In `VLA Policy — Universal HTTP`, use
`http://127.0.0.1:8787`.
`lerobot/smolvla_base` is a base model. It is a useful fine-tuning starting
point, not a universal zero-shot controller. Use an embodiment-specific
checkpoint whose feature names, action dimensions, state dimensions,
normalization statistics, and camera keys match the workflow.
The sidecar:
- loads only the chosen policy and its serialized pre/post-processors;
- uses `predict_action_chunk` when provided and falls back to `select_action`;
- keeps the model resident by default for low latency;
- can move it to CPU after an idle interval and move it back on demand;
- accepts one request at a time per policy, preventing stateful policy races;
- uses bounded JSON/JPEG rather than pickle;
- never returns tracebacks, environment variables, request data, or
authorization headers.
For a real local API acceptance run, start ComfyUI and the policy sidecar, put
an image in ComfyUI's `input` directory, then run:
```bash
python tests/manual_robotics_smoke.py \
--comfy-url http://127.0.0.1:8188 \
--policy-url http://127.0.0.1:8787 \
--image robot_front.png
```
The script queues the graph through `POST /prompt`, waits on its history entry,
and prints the policy report, safety report, first action, and preview filename.
Install the relevant official LeRobot extra for another policy. Examples are
`lerobot[pi]` for π0/π0.5/π0-FAST and `lerobot[smolvla]` for SmolVLA. Some
newer policy integrations may require installing current LeRobot from source
with their documented extra.
## Native openpi server
Install the small ComfyUI client dependencies:
```bash
python -m pip install -r requirements-robotics-client.txt
```
Run the official openpi policy WebSocket server in its own supported
environment. The upstream runtime is currently tested on Ubuntu 22.04 and an
NVIDIA GPU with more than 8 GB for inference; use WSL/Docker or a remote Linux
server rather than forcing it into a macOS/Windows ComfyUI environment.
Use:
- `Flat keys (DROID / LIBERO)` for observations such as
`observation/image`, `observation/wrist_image`, and `observation/state`;
- `Nested images (ALOHA)` for `state`, an `images` mapping such as
`cam_high`/wrist cameras, and `prompt`.
The workflow supplies key names, but the policy's own transform still defines
the exact shapes and normalization. `OPENPI_API_KEY` is read only from the
ComfyUI server environment. Remote endpoints require WSS and explicit
`allow_remote=true`.
## Native Isaac-GR00T N1.7 server
Install the same lightweight robotics client requirements in ComfyUI. Run the
official GR00T `PolicyServer` beside an embodiment-compatible `Gr00tPolicy`.
The node sends the documented nested contract:
```text
video.<camera> uint8 [batch=1, history, height, width, RGB=3]
state.state float32[batch=1, history, state_dim]
language.task string [batch=1, 1]
```
The official server returns one or more physical-unit action streams with
shape `[batch, horizon, dimension]`. The node flattens those streams while
preserving their named slices. `GROOT_API_TOKEN` remains in the ComfyUI
environment.
GR00T N1.7 currently targets NVIDIA CUDA/Jetson Linux and needs an
embodiment-compatible base or post-trained checkpoint. A ComfyUI client on
Windows, macOS, ROCm, or another machine may call that server over a trusted
network, but remote access must be explicitly enabled. Native GR00T ZMQ does
not encrypt traffic; use a private authenticated network/tunnel. Prefer the
universal HTTPS bridge when transport-layer encryption is required.
## Model catalog: what “available” means
`VLA Model Catalog` distinguishes these states:
| Family | Example checkpoint | Route | Important qualification |
| --- | --- | --- | --- |
| SmolVLA | `lerobot/smolvla_base` | LeRobot HTTP sidecar | 450M and the best small starting point; fine-tune for the robot |
| X-VLA | `lerobot/xvla-base` | LeRobot HTTP sidecar | 0.9B cross-embodiment base; use a matching domain checkpoint |
| π0 | `lerobot/pi0_base` | LeRobot or openpi | Base/fine-tuning model, not a universal drop-in controller |
| π0-FAST | `lerobot/pi0fast-base` | LeRobot or openpi | Faster tokenized action generation |
| π0.5 | `lerobot/pi05_base` | LeRobot or openpi | Open-world generalization; still embodiment-specific |
| GR00T N1.7 | `nvidia/GR00T-N1.7-3B` | GR00T ZMQ or LeRobot | Base has specific zero-shot tags; other robots need post-training |
| X-Square WALL-OSS | `x-square-robot/wall-oss-flow` | LeRobot HTTP sidecar | MoE research model; validate checkpoint terms and embodiment |
| MolmoAct2 | `lerobot/MolmoAct2-SO100_101-LeRobot` | LeRobot HTTP sidecar | Converted SO-100/SO-101 checkpoint |
| VLA-JEPA | `lerobot/VLA-JEPA-Pretrain` | LeRobot HTTP sidecar | DROID pretrain plus LIBERO/SimplerEnv checkpoints |
| LingBot-VA | `lerobot/lingbot_va_base` | LeRobot HTTP sidecar | Prefer its LIBERO-Long/RoboTwin post-train where applicable |
| FastWAM | released LIBERO checkpoint | LeRobot HTTP sidecar | Heavy world-action research runtime |
| EO-1 / EVO-1 | your trained checkpoint | LeRobot HTTP sidecar | Architecture support, not a universal ready-made controller |
| OpenVLA-OFT | compatible OFT fine-tune | dedicated sidecar | OFT is the faster multi-image/high-frequency OpenVLA route |
| Octo small | Octo small 27M | dedicated JAX sidecar | Lightweight legacy research baseline |
The catalog is a verified runtime/checkpoint map, not a promise that a base
checkpoint understands an arbitrary robot. Exact data transforms and
fine-tuning are part of the policy.
## Observation history and real-time use
Connect an `IMAGE` batch to a camera input to represent temporal history. All
camera batches must have the same length, although a one-frame camera may
broadcast. Use:
`Video Slice` → `VLM Adaptive Frame Sampler` or a live capture source →
`VLM Image Pixel Budget` → `VLA Observation Builder`
For closed-loop robotics, do not run an unbounded ComfyUI queue for each motor
tick. Use ComfyUI to prototype and inspect the observation/policy/safety
contract, and use the policy runtime's asynchronous control support for the
actual high-frequency loop. LeRobot supports asynchronous action chunks and
GR00T supports TensorRT deployment; both are better places for timing-critical
execution.
## Action safety semantics
The `VLA Action Safety Gate` checks:
- policy action dimension against the embodiment;
- NaN and infinity;
- minimum and maximum values;
- maximum change per action dimension and control step;
- the requested execution horizon.
Modes:
- `Block unsafe`: raise and stop the workflow on any violation.
- `Clamp safely`: replace non-finite values conservatively, then clamp bounds
and sequential per-step deltas.
- `Hold position on unsafe`: replace the whole chunk with the explicitly
supplied previous/current command.
- `Report only`: preserve the raw trajectory and set
`safe_for_handoff=false`.
For delta-action policies, `previous_action_json` means the previous delta
command, not an absolute joint pose. Define the profile in the same units and
semantics as the policy output.
`VLA Actions From JSON` imports recorded/simulator trajectories without a
network policy. `VLA Action Chunk Replan` blends the unexecuted edge of an old
chunk into a new chunk to reduce discontinuities, then the result should pass
through the safety gate again. This deterministic blend is useful for workflow
experiments but does not replace LeRobot's asynchronous controller or a
policy-specific real-time chunking implementation.
## API example
`vla_http_policy_safety_api.json` is a ComfyUI API prompt graph. Put
`robot_front.png` in `ComfyUI/input`, start a compatible sidecar, then POST:
```json
{"prompt": {"...": "contents of vla_http_policy_safety_api.json"}}
```
It builds an observation, calls the policy, clamps it against the explicit
profile, renders the trajectory, and outputs both inference and safety JSON.
## Security checklist
- Keep all policy tokens in environment variables.
- Leave `allow_remote=false` for local servers.
- Remote universal endpoints must use HTTPS; remote openpi endpoints must use
WSS.
- Never expose GR00T ZMQ directly to an untrusted network.
- Pin checkpoint revisions when reproducibility matters.
- Treat camera images, task language, and robot state as sensitive data.
- Do not connect action JSON directly to hardware without a separate
supervised controller bridge and independent safety system.
Authoritative upstream documentation:
- [LeRobot installation](https://huggingface.co/docs/lerobot/main/en/installation)
- [LeRobot SmolVLA](https://huggingface.co/docs/lerobot/smolvla)
- [LeRobot asynchronous inference](https://huggingface.co/docs/lerobot/async)
- [Physical Intelligence openpi](https://github.com/Physical-Intelligence/openpi)
- [NVIDIA Isaac-GR00T](https://github.com/NVIDIA/Isaac-GR00T)
- [OpenVLA and OFT](https://github.com/openvla/openvla)
- [Octo](https://github.com/octo-models/octo)
-1
View File
@@ -1 +0,0 @@
"""Runnable, dependency-isolated robotics policy bridge examples."""
-434
View File
@@ -1,434 +0,0 @@
#!/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()
@@ -1,130 +0,0 @@
{
"1": {
"class_type": "LoadImage",
"inputs": {
"image": "robot_front.png"
}
},
"2": {
"class_type": "VLAEmbodimentProfile",
"inputs": {
"preset": "Generic 7-DoF joint + gripper",
"control_hz": 20.0,
"state_names_json": "",
"action_names_json": "",
"action_min_json": "",
"action_max_json": "",
"max_delta_json": "",
"camera_names_json": "",
"action_mode_override": ""
}
},
"3": {
"class_type": "VLAObservationBuilder",
"inputs": {
"task": "Pick up the blue cube and place it in the tray.",
"state_json": "[0, 0, 0, 0, 0, 0, 0, 0]",
"primary_image": [
"1",
0
],
"primary_camera": "observation.images.front",
"history_fps": 10.0,
"timestamp": 0.0,
"embodiment": [
"2",
0
]
}
},
"4": {
"class_type": "VLAHTTPPolicy",
"inputs": {
"observation": [
"3",
0
],
"endpoint": "http://127.0.0.1:8787",
"timeout_seconds": 120.0,
"include_history": true,
"allow_remote": false
}
},
"5": {
"class_type": "VLAActionSafety",
"inputs": {
"actions": [
"4",
0
],
"embodiment": [
"2",
0
],
"mode": "Clamp safely",
"execution_horizon": 8,
"previous_action_json": "[0, 0, 0, 0, 0, 0, 0, 0]"
}
},
"6": {
"class_type": "VLATrajectoryPreview",
"inputs": {
"actions": [
"5",
0
],
"width": 960,
"height": 480,
"embodiment": [
"2",
0
]
}
},
"7": {
"class_type": "VLAActionInspect",
"inputs": {
"actions": [
"5",
0
],
"step_index": 0
}
},
"8": {
"class_type": "PreviewImage",
"inputs": {
"images": [
"6",
0
]
}
},
"9": {
"class_type": "ViewText",
"inputs": {
"text": [
"4",
1
]
}
},
"10": {
"class_type": "ViewText",
"inputs": {
"text": [
"5",
1
]
}
},
"11": {
"class_type": "ViewText",
"inputs": {
"text": [
"7",
0
]
}
}
}
-57
View File
@@ -1,57 +0,0 @@
{
"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
]
}
}
}
-124
View File
@@ -1,124 +0,0 @@
# 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.
@@ -1,45 +0,0 @@
{
"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
]
}
}
}
@@ -1,76 +0,0 @@
{
"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
]
}
}
}
@@ -1,74 +0,0 @@
{
"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
]
}
}
}
@@ -1,116 +0,0 @@
{
"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
]
}
}
}
@@ -1,116 +0,0 @@
{
"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
]
}
}
}
@@ -1,91 +0,0 @@
{
"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
]
}
}
}
@@ -1,98 +0,0 @@
{
"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
]
}
}
}
+486
View File
@@ -0,0 +1,486 @@
import os
import json
import shutil
from os.path import join, dirname, abspath, exists
from os import makedirs, symlink, readlink
import platform
import subprocess
import sys
import importlib.util
import re
import torch
import cpuinfo
import packaging.tags
from requests import get
import asyncio
import inspect
import aiohttp
from server import PromptServer
from tqdm import tqdm
import pkg_resources
def verify_python_support():
"""Verify Python version meets minimum requirements."""
version = tuple(map(int, platform.python_version_tuple()[:2]))
if version < (3, 8):
print("Warning: Python 3.8 or higher is required")
return False
return True
def verify_pypy_support(system_info):
"""Verify if the current PyPy version/platform combination is supported."""
if 'pp' in system_info['python_version']:
pp_ver = system_info['python_version'][2:4]
if pp_ver not in ['38', '39', '310']:
print("Warning: Current PyPy version may not be supported")
return False
if system_info['platform_tag'] not in ['linux_i686', 'linux_x86_64', 'win_amd64',
'macosx_10_15_x86_64', 'macosx_10_9_x86_64']:
print("Warning: Current platform may not be supported for PyPy")
return False
return True
def get_python_version():
"""Return the Python version in a format matching wheel tags, e.g., 'cp39' for Python 3.9."""
version = platform.python_version_tuple()[:2]
impl = 'pp' if platform.python_implementation().lower() == 'pypy' else 'cp'
return f"{impl}{version[0]}{version[1]}"
def get_system_info():
"""Gather system information related to platform architecture, Python version, and OS."""
system_info = {
'gpu': False,
'cuda_version': None,
'rocm_version': None,
'python_version': get_python_version(),
'os': platform.system().lower(),
'arch': platform.machine().lower(),
'platform_tag': None
}
# Determine platform-specific tags
if system_info['os'] == 'linux':
if system_info['arch'] == 'x86_64':
system_info['platform_tag'] = 'linux_x86_64'
elif system_info['arch'] == 'i686':
system_info['platform_tag'] = 'linux_i686'
elif system_info['arch'] == 'aarch64':
system_info['platform_tag'] = 'linux_aarch64'
elif system_info['os'] == 'windows':
if system_info['arch'] == 'amd64':
system_info['platform_tag'] = 'win_amd64'
elif system_info['arch'] == 'x86':
system_info['platform_tag'] = 'win32'
elif system_info['os'] == 'darwin':
if system_info['arch'] == 'x86_64':
# Intel Mac
if 'pp' in system_info['python_version']:
system_info['platform_tag'] = 'macosx_10_15_x86_64'
else:
py_ver = int(system_info['python_version'][3:])
if py_ver >= 12:
system_info['platform_tag'] = 'macosx_10_13_x86_64'
else:
system_info['platform_tag'] = 'macosx_10_9_x86_64'
elif system_info['arch'] == 'arm64':
# Apple Silicon (M1/M2/M3)
print("Apple Silicon detected. llama-cpp-python will be built with Metal support")
system_info['platform_tag'] = None # Force source build for optimal Metal support
system_info['metal'] = True
# Check for GPU support
if importlib.util.find_spec('torch'):
try:
import torch
if hasattr(torch.version, 'hip') and torch.version.hip is not None:
system_info['gpu'] = True
system_info['rocm_version'] = f"rocm{torch.version.hip}"
elif torch.cuda.is_available():
system_info['gpu'] = True
system_info['cuda_version'] = "cu" + torch.version.cuda.replace(".", "").strip()
except:
pass
return system_info
def latest_lamacpp():
"""Fetch the latest version of llama-cpp-python, with fallback."""
try:
response = get("https://api.github.com/repos/abetlen/llama-cpp-python/releases/latest", timeout=10)
response.raise_for_status()
return response.json()["tag_name"].replace("v", "")
except Exception as e:
print(f"Failed to fetch latest version: {e}")
return "0.3.1" # Fallback to known working version
def package_is_installed(package_name):
"""Check if a Python package is installed."""
return importlib.util.find_spec(package_name) is not None
def install_package(package_name, extra_args=None):
"""Install a Python package with pip."""
command = [sys.executable, "-m", "pip", "install", package_name, "--no-cache-dir"]
if extra_args:
command.extend(extra_args.split())
subprocess.check_call(command)
def install_llama(system_info):
"""Install llama-cpp-python using the appropriate method based on system capabilities."""
if not verify_python_support():
print("ERROR: Unsupported Python version")
return False
if not verify_pypy_support(system_info):
print("WARNING: Unsupported PyPy configuration")
imported = package_is_installed("llama-cpp-python") or package_is_installed("llama_cpp")
if imported:
print("llama-cpp installed")
return True
# Simple pip install for Linux
if system_info['os'] == 'linux':
try:
print("Installing llama-cpp-python via pip")
install_package("llama-cpp-python")
return True
except Exception as e:
print(f"Installation failed: {e}")
return False
# If pre-built wheels fail, try GitHub release wheels
try:
version = latest_lamacpp()
platform_tag = system_info['platform_tag']
if platform_tag:
python_version = system_info['python_version']
wheel_name = f"llama_cpp_python-{version}-{python_version}-{python_version}-{platform_tag}.whl"
wheel_url = f"https://github.com/abetlen/llama-cpp-python/releases/download/v{version}/{wheel_name}"
print(f"Attempting to install from {wheel_url}")
install_package(wheel_url)
print(f"Successfully installed llama-cpp-python v{version}")
return True
except Exception as e:
print(f"GitHub wheel installation failed: {e}")
print("Attempting source build with acceleration...")
# Build from source with appropriate acceleration
try:
if system_info.get('metal', False):
print("Building llama-cpp-python from source with Metal support")
os.environ['CMAKE_ARGS'] = "-DGGML_METAL=on"
install_package("llama-cpp-python")
return True
elif system_info['gpu']:
if system_info.get('cuda_version'):
print("Building llama-cpp-python from source with CUDA support")
# Add ZLUDA support check
if os.environ.get('ZLUDA_PATH'):
print("ZLUDA detected, building with ZLUDA support")
os.environ['CMAKE_ARGS'] = "-DGGML_CUDA=on -DGGML_CUDA_ZLUDA=on"
else:
os.environ['CMAKE_ARGS'] = "-DGGML_CUDA=on"
install_package("llama-cpp-python")
return True
elif system_info.get('rocm_version'):
print("Building llama-cpp-python from source with ROCm support")
os.environ['CMAKE_ARGS'] = "-DGGML_HIPBLAS=on"
install_package("llama-cpp-python")
return True
except Exception as e:
print(f"Accelerated build failed: {e}")
print("Falling back to CPU-only version")
# Final fallback - basic CPU version
try:
print("Installing CPU-only version")
install_package("llama-cpp-python")
return True
except Exception as e:
print(f"CPU installation failed: {e}")
return False
config = None
def is_logging_enabled():
config = get_extension_config()
if "logging" not in config:
return False
return config["logging"]
def log(message, type=None, always=False, name=None):
if not always and not is_logging_enabled():
return
if type is not None:
message = f"[{type}] {message}"
if name is None:
name = get_extension_config()["name"]
print(f"(vlmnodes:{name}) {message}")
def get_ext_dir(subpath=None, mkdir=False):
dir = os.path.dirname(__file__)
if subpath is not None:
dir = os.path.join(dir, subpath)
dir = os.path.abspath(dir)
if mkdir and not os.path.exists(dir):
os.makedirs(dir)
return dir
def get_comfy_dir(subpath=None, mkdir=False):
dir = os.path.dirname(inspect.getfile(PromptServer))
if subpath is not None:
dir = os.path.join(dir, subpath)
dir = os.path.abspath(dir)
if mkdir and not os.path.exists(dir):
os.makedirs(dir)
return dir
def get_web_ext_dir():
config = get_extension_config()
name = config["name"]
dir = get_comfy_dir("web/extensions/vlmnodes")
if not os.path.exists(dir):
os.makedirs(dir)
dir = os.path.join(dir, name)
return dir
def get_extension_config(reload=False):
global config
if reload == False and config is not None:
return config
config_path = get_ext_dir("vlmnodes.json")
default_config_path = get_ext_dir("vlmnodes.default.json")
if not os.path.exists(config_path):
if os.path.exists(default_config_path):
shutil.copy(default_config_path, config_path)
if not os.path.exists(config_path):
log(f"Failed to create config at {config_path}", type="ERROR", always=True, name="???")
print(f"Extension path: {get_ext_dir()}")
return {"name": "Unknown", "version": -1}
else:
log("Missing pysssss.default.json, this extension may not work correctly. Please reinstall the extension.",
type="ERROR", always=True, name="???")
print(f"Extension path: {get_ext_dir()}")
return {"name": "Unknown", "version": -1}
with open(config_path, "r") as f:
config = json.loads(f.read())
return config
def link_js(src, dst):
src = os.path.abspath(src)
dst = os.path.abspath(dst)
if os.name == "nt":
try:
import _winapi
_winapi.CreateJunction(src, dst)
return True
except:
pass
try:
os.symlink(src, dst)
return True
except:
import logging
logging.exception('')
return False
def is_junction(path):
if os.name != "nt":
return False
try:
return bool(os.readlink(path))
except OSError:
return False
def install_js():
src_dir = get_ext_dir("web/js")
if not os.path.exists(src_dir):
log("No JS")
return
should_install = should_install_js()
if should_install:
log("it looks like you're running an old version of ComfyUI that requires manual setup of web files, it is recommended you update your installation.", "warning", True)
dst_dir = get_web_ext_dir()
linked = os.path.islink(dst_dir) or is_junction(dst_dir)
if linked or os.path.exists(dst_dir):
if linked:
if should_install:
log("JS already linked")
else:
os.unlink(dst_dir)
log("JS unlinked, PromptServer will serve extension")
elif not should_install:
shutil.rmtree(dst_dir)
log("JS deleted, PromptServer will serve extension")
return
if not should_install:
log("JS skipped, PromptServer will serve extension")
return
if link_js(src_dir, dst_dir):
log("JS linked")
return
log("Copying JS files")
shutil.copytree(src_dir, dst_dir, dirs_exist_ok=True)
def should_install_js():
return not hasattr(PromptServer.instance, "supports") or "custom_nodes_from_web" not in PromptServer.instance.supports
def init(check_imports=None):
log("Init")
if check_imports is not None:
import importlib.util
for imp in check_imports:
spec = importlib.util.find_spec(imp)
if spec is None:
log(f"{imp} is required, please check requirements are installed.",
type="ERROR", always=True)
return False
install_js()
return True
def get_async_loop():
loop = None
try:
loop = asyncio.get_event_loop()
except:
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
return loop
def get_http_session():
loop = get_async_loop()
return aiohttp.ClientSession(loop=loop)
async def download(url, stream, update_callback=None, session=None):
close_session = False
if session is None:
close_session = True
session = get_http_session()
try:
async with session.get(url) as response:
size = int(response.headers.get('content-length', 0)) or None
with tqdm(
unit='B', unit_scale=True, miniters=1, desc=url.split('/')[-1], total=size,
) as progressbar:
perc = 0
async for chunk in response.content.iter_chunked(2048):
stream.write(chunk)
progressbar.update(len(chunk))
if update_callback is not None and progressbar.total is not None and progressbar.total != 0:
last = perc
perc = round(progressbar.n / progressbar.total, 2)
if perc != last:
last = perc
await update_callback(perc)
finally:
if close_session and session is not None:
await session.close()
async def download_to_file(url, destination, update_callback=None, is_ext_subpath=True, session=None):
if is_ext_subpath:
destination = get_ext_dir(destination)
with open(destination, mode='wb') as f:
download(url, f, update_callback, session)
def wait_for_async(async_fn, loop=None):
res = []
async def run_async():
r = await async_fn()
res.append(r)
if loop is None:
try:
loop = asyncio.get_event_loop()
except:
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
loop.run_until_complete(run_async())
return res[0]
def update_node_status(client_id, node, text, progress=None):
if client_id is None:
client_id = PromptServer.instance.client_id
if client_id is None:
return
PromptServer.instance.send_sync("vlmnodes/update_status", {
"node": node,
"progress": progress,
"text": text
}, client_id)
async def update_node_status_async(client_id, node, text, progress=None):
if client_id is None:
client_id = PromptServer.instance.client_id
if client_id is None:
return
await PromptServer.instance.send("vlmnodes/update_status", {
"node": node,
"progress": progress,
"text": text
}, client_id)
def get_config_value(key, default=None, throw=False):
split = key.split(".")
obj = get_extension_config()
for s in split:
if s in obj:
obj = obj[s]
else:
if throw:
raise KeyError("Configuration key missing: " + key)
else:
return default
return obj
def is_inside_dir(root_dir, check_path):
root_dir = os.path.abspath(root_dir)
if not os.path.isabs(check_path):
check_path = os.path.abspath(os.path.join(root_dir, check_path))
return os.path.commonpath([check_path, root_dir]) == root_dir
def get_child_dir(root_dir, child_path, throw_if_outside=True):
child_path = os.path.abspath(os.path.join(root_dir, child_path))
if is_inside_dir(root_dir, child_path):
return child_path
if throw_if_outside:
raise NotADirectoryError(
"Saving outside the target folder is not allowed.")
return None
-279
View File
@@ -1,279 +0,0 @@
"""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",
}
+98 -185
View File
@@ -1,202 +1,105 @@
"""Lazy AudioLDM2 generation with legacy and standard ComfyUI AUDIO outputs."""
from __future__ import annotations
from huggingface_hub import snapshot_download
from pathlib import Path
import torch
import os
import soundfile as sf
from folder_paths import output_directory
import folder_paths
import datetime
from pathlib import Path
import folder_paths
import numpy as np
import torch
from .runtime import (
CachedModelNode,
execution_device,
require_module,
reserve_external_vram,
snapshot_download,
torch_dtype,
)
# Define the directory for saving files related to the audio model
files_for_audio_model = Path(folder_paths.folder_names_and_paths["LLavacheckpoints"][0][0]) / "files_for_audioldm2"
files_for_audio_model.mkdir(parents=True, exist_ok=True) # Ensure the directory exists
class AnyType(str):
def __ne__(self, other):
def __ne__(self, __value: object) -> bool:
return False
base_path = os.path.dirname(os.path.realpath(__file__))
# Our any instance wants to be a wildcard string
any = AnyType("*")
class AudioLDM2ModelPredictor:
def __init__(self):
from diffusers import AudioLDM2Pipeline
self.device = "cuda" if torch.cuda.is_available() else "cpu"
torch_dtype = torch.float16 if self.device == "cuda" else torch.float32
ANY = AnyType("*")
# Use snapshot_download to manage the model download/cache
self.model_path = snapshot_download("cvssp/audioldm2",
local_dir=files_for_audio_model,
force_download=False, # Set to True to always download
local_files_only=False, # Download if not available locally
use_auth_token=False, # Set to True if using a private model
local_dir_use_symlinks="auto", # Auto-manage symlinks
ignore_patterns=["*.bin", "*.jpg", "*.png"]) # Ignore unrelated files
self.pipeline = AudioLDM2Pipeline.from_pretrained(self.model_path,
torch_dtype=torch_dtype).to(self.device)
self.generator = torch.Generator(self.device)
class AudioLDM2Predictor:
def __init__(self, cpu_offload=True):
diffusers = require_module("diffusers")
path = snapshot_download(
"cvssp/audioldm2",
"audioldm2",
ignore_patterns=["*.bin", "*.jpg", "*.png"],
)
self.device = execution_device()
dtype = torch_dtype("float16", self.device)
if self.device.type != "cpu":
reserve_external_vram(8 * 1024**3)
self.pipeline = diffusers.AudioLDM2Pipeline.from_pretrained(
path, torch_dtype=dtype
)
# Accelerate's model CPU offload is currently reliable on the CUDA API,
# which covers both NVIDIA CUDA and AMD ROCm PyTorch builds.
if self.device.type == "cuda" and cpu_offload:
require_module("accelerate")
self.pipeline.enable_model_cpu_offload()
else:
self.pipeline.to(self.device)
def generate_audio(self, text, negative_prompt, duration, guidance_scale, random_seed, sample_rate, n_candidates=1, extension="wav"):
if text is None:
raise ValueError("Please provide a text input.")
# Manual seed for reproducibility
self.generator.manual_seed(int(random_seed))
def close(self):
self.pipeline = None
import gc
gc.collect()
try:
import comfy.model_management as model_management
model_management.soft_empty_cache()
except Exception:
pass
def generate(self, text, negative, duration, guidance, seed, count, steps):
# MPS generators are not supported by every PyTorch/Diffusers pairing.
# A CPU generator remains deterministic and works with every pipeline.
generator_device = (
self.device if self.device.type in {"cuda", "xpu"} else "cpu"
)
generator = torch.Generator(device=generator_device).manual_seed(
int(seed)
)
audios = self.pipeline(
# Generate audio
waveforms = self.pipeline(
text,
negative_prompt=negative or None,
audio_length_in_s=float(duration),
guidance_scale=float(guidance),
num_inference_steps=int(steps),
num_waveforms_per_prompt=int(count),
generator=generator,
).audios
array = np.asarray(audios, dtype=np.float32)
if array.ndim == 1:
array = array[None, :]
native_rate = int(
getattr(
getattr(getattr(self.pipeline, "vae", None), "config", None),
"sampling_rate",
16000,
)
)
return array, native_rate
audio_length_in_s=duration,
guidance_scale=guidance_scale,
num_inference_steps=200,
negative_prompt=negative_prompt,
num_waveforms_per_prompt=n_candidates,
generator=self.generator,
)["audios"]
final_waveforms = waveforms[0].tolist()
return (final_waveforms, sample_rate) # Return the path of the generated audio file
class AudioLDM2Node(CachedModelNode):
class AudioLDM2Node:
def __init__(self):
self.predictor = AudioLDM2ModelPredictor()
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"text": ("STRING", {"default": "", "multiline": True}),
"negative_prompt": (
"STRING",
{"default": "", "multiline": True},
),
"duration": (
"INT",
{"default": 10, "min": 1, "max": 60},
),
"guidance_scale": (
"FLOAT",
{"default": 3.5, "min": 0.1, "max": 20.0, "step": 0.1},
),
"seed": ("INT", {"default": 42, "min": 0}),
"n_candidates": (
"INT",
{"default": 1, "min": 1, "max": 10},
),
"sample_rate": (
"INT",
{"default": 16000, "min": 8000, "max": 48000},
),
"extension": (["wav", "flac"],),
},
"optional": {
"steps": ("INT", {"default": 100, "min": 10, "max": 500}),
"cpu_offload": ("BOOLEAN", {"default": True}),
"unload_after": ("BOOLEAN", {"default": False}),
},
"text": ("STRING",{"default": "", "forceInput": True}),
"negative_prompt": ("STRING",{"default": "", "forceInput": True}),
"duration": ("INT",{"default": 10, "min": 1, "max": 60, "step": 1}),
"guidance_scale": ("FLOAT", {"default": 3.5, "min": 0.1, "max": 20.0, "step": 0.1}),
"seed": ("INT", {"default": 42, "step": 1}),
"n_candidates": ("INT", {"default": 3, "min": 1, "max": 10, "step": 1}),
"sample_rate": ("INT", {"default": 16000, "min": 8000, "max": 48000, "step": 1}),
"extension": (["wav", "mp3", "flac"], {"default": "wav"}),
}
}
RETURN_NAMES = ("wave_form", "sample_rate", "audio")
RETURN_TYPES = (ANY, "INT", "AUDIO")
RETURN_NAMES = ("wave_form", "sample_rate", )
RETURN_TYPES = (any, "INT", )
OUTPUT_NODE = True
FUNCTION = "generate_audio_final"
CATEGORY = "VLM Nodes/Audio"
def generate_audio_final(
self,
text,
negative_prompt,
duration,
guidance_scale,
sample_rate,
seed,
n_candidates,
extension,
steps=100,
cpu_offload=True,
unload_after=False,
):
del extension
predictor = self.get_or_create_model(
("audioldm2", bool(cpu_offload)),
lambda: AudioLDM2Predictor(cpu_offload),
)
try:
waveforms, native_rate = predictor.generate(
text,
negative_prompt,
duration,
guidance_scale,
seed,
n_candidates,
steps,
)
if int(sample_rate) != native_rate:
samples = torch.from_numpy(waveforms).unsqueeze(1)
target_length = round(
samples.shape[-1] * int(sample_rate) / native_rate
)
waveforms = (
torch.nn.functional.interpolate(
samples,
size=target_length,
mode="linear",
align_corners=False,
)
.squeeze(1)
.numpy()
)
# Standard Comfy AUDIO is [batch, channels, samples].
audio = {
"waveform": torch.from_numpy(waveforms).unsqueeze(1),
"sample_rate": int(sample_rate),
}
return (waveforms[0].tolist(), int(sample_rate), audio)
finally:
self.maybe_clear_model(unload_after)
def generate_audio_final(self, text, negative_prompt, duration, guidance_scale, sample_rate, seed, n_candidates, extension):
wave_form, sample_rate_final = self.predictor.generate_audio(text, negative_prompt, duration, guidance_scale, seed, sample_rate, n_candidates, extension)
return (wave_form, sample_rate_final, )
class SaveAudioNode:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"waveforms": (ANY,),
"sample_rate": ("INT",),
"extension": (["wav", "flac"],),
"filename": ("STRING", {"default": "audio"}),
"waveforms": (any, {}),
"sample_rate": ("INT", {"forceInput": True}),
"extension": (["wav", "mp3", "flac"], {"default": "wav"}),
"filename": ("STRING", {"default": "audio", "forceInput": True}) # Input for filename
}
}
@@ -206,25 +109,35 @@ class SaveAudioNode:
OUTPUT_NODE = True
def save_audio(self, waveforms, sample_rate, extension, filename):
soundfile = require_module("soundfile")
safe_name = Path(filename).name.strip() or "audio"
output = Path(folder_paths.output_directory)
output.mkdir(parents=True, exist_ok=True)
base = output / safe_name
path = base.with_suffix(f".{extension}")
counter = 2
while path.exists():
path = output / f"{safe_name}_{counter:05d}.{extension}"
counter += 1
soundfile.write(path, np.asarray(waveforms), int(sample_rate))
return ()
# Build the base audio path
base_path = Path(output_directory) / filename
# Initialize a counter
counter = 1
# Check if the file exists and append a number if it does
while True:
# Format the filename with leading zeros for numbering
if counter == 1:
audio_path = base_path.with_suffix(f".{extension}") # First instance
else:
audio_path = base_path.with_name(f"{filename}_{counter:05d}").with_suffix(f".{extension}")
if not audio_path.exists():
break # Found a unique filename
counter += 1 # Increment the counter
# Save the audio file
sf.write(audio_path.as_posix(), waveforms, sample_rate)
return ()
NODE_CLASS_MAPPINGS = {
"AudioLDM2Node": AudioLDM2Node,
"SaveAudioNode": SaveAudioNode,
"SaveAudioNode": SaveAudioNode
}
NODE_DISPLAY_NAME_MAPPINGS = {
"AudioLDM2Node": "AudioLDM2",
"SaveAudioNode": "Save Audio",
"AudioLDM2Node": "AudioLDM-2 Node",
"SaveAudioNode": "Save Audio Node"
}
-35
View File
@@ -1,35 +0,0 @@
"""A zero-download runtime report for portable support requests."""
from __future__ import annotations
import json
from .runtime import runtime_diagnostics
class VLMRuntimeDiagnostics:
@classmethod
def INPUT_TYPES(cls):
return {"required": {}}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("runtime_report",)
FUNCTION = "report"
CATEGORY = "VLM Nodes/Diagnostics"
OUTPUT_NODE = True
def report(self):
return (
json.dumps(
runtime_diagnostics(),
ensure_ascii=False,
indent=2,
sort_keys=True,
),
)
NODE_CLASS_MAPPINGS = {"VLMRuntimeDiagnostics": VLMRuntimeDiagnostics}
NODE_DISPLAY_NAME_MAPPINGS = {
"VLMRuntimeDiagnostics": "VLM Runtime Diagnostics"
}
-510
View File
@@ -1,510 +0,0 @@
"""Florence-2 multitask caption, OCR, detection and segmentation node."""
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
from .runtime import (
CachedModelNode,
ManagedTorchModel,
batch_text,
inference_context,
model_device,
move_inputs,
pil_mask_to_tensor,
pil_to_tensor,
require_module,
snapshot_download,
tensor_batch_to_pil,
torch_dtype,
)
MODELS = {
"Florence-2 base FT (fast)": "florence-community/Florence-2-base-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": 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")
repo_id = MODELS[model_label]
path = snapshot_download(
repo_id,
f"florence2/{repo_id.replace('/', '--')}",
ignore_patterns=["*.bin"],
)
self.dtype = torch_dtype("float16")
self.processor = transformers.Florence2Processor.from_pretrained(path)
model = transformers.Florence2ForConditionalGeneration.from_pretrained(
path,
dtype=self.dtype,
)
model.eval()
self.handle = ManagedTorchModel(model, processor=self.processor)
def close(self):
self.handle.close()
self.processor = None
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")
model = self.handle.ensure_loaded()
device = model_device(model)
inputs = move_inputs(inputs, device, floating_dtype=self.dtype)
with torch.inference_mode(), inference_context(device, self.dtype):
generated = model.generate(
**inputs,
max_new_tokens=int(max_new_tokens),
num_beams=int(beams),
do_sample=False,
early_stopping=int(beams) > 1,
)
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
)
return raw, parsed
def _json_default(value):
if hasattr(value, "tolist"):
return value.tolist()
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 = _spatial_result(parsed)
mask = Image.new("L", image.size, 0)
visual = image.copy().convert("RGB")
mask_draw = ImageDraw.Draw(mask)
draw = ImageDraw.Draw(visual)
width = max(2, min(8, round(min(image.size) / 256 * 3)))
labels = result.get("labels", [])
scores = result.get("scores", [])
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 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)
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
class Florence2(CachedModelNode):
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"task": (list(TASKS),),
"text_input": (
"STRING",
{
"default": "",
"multiline": True,
"tooltip": (
"Required only for phrase grounding, referring-expression "
"segmentation, and open-vocabulary detection."
),
},
),
"model": (
list(MODELS),
{"default": "Florence-2 large FT (recommended)"},
),
"max_new_tokens": (
"INT",
{"default": 1024, "min": 1, "max": 4096},
),
"beams": ("INT", {"default": 3, "min": 1, "max": 8}),
},
"optional": {
"unload_after": ("BOOLEAN", {"default": False}),
"region": (
"BOUNDING_BOX",
{
"tooltip": (
"Core bounding box input required by Region to "
"Segmentation/Category/Description/OCR."
)
},
),
},
}
RETURN_TYPES = ("STRING", "STRING", "MASK", "IMAGE")
RETURN_NAMES = ("text", "structured_json", "mask", "visualization")
FUNCTION = "run"
CATEGORY = "VLM Nodes/Florence-2"
def run(
self,
image,
task,
text_input,
model,
max_new_tokens,
beams,
unload_after=False,
region=None,
):
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, extra_input in zip(images, extra_inputs):
raw, parsed = predictor.run(
pil_image,
spec.token,
extra_input,
max_new_tokens,
beams,
)
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,
sort_keys=True,
),
torch.cat(masks),
torch.cat(visuals),
)
finally:
self.maybe_clear_model(unload_after)
NODE_CLASS_MAPPINGS = {"Florence2": Florence2}
NODE_DISPLAY_NAME_MAPPINGS = {"Florence2": "Florence-2 Multitask Vision"}
-413
View File
@@ -1,413 +0,0 @@
"""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",
]
-537
View File
@@ -1,537 +0,0 @@
"""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
View File
File diff suppressed because it is too large Load Diff
+121 -120
View File
@@ -1,140 +1,141 @@
"""JoyTag image tagging with cached, ComfyUI-managed model weights."""
from __future__ import annotations
import numpy as np
import torch
from .joytagger import Models
from PIL import Image
import torch.amp.autocast_mode
from pathlib import Path
import torch
import torchvision.transforms.functional as TVF
from huggingface_hub import snapshot_download
from torchvision import transforms
import folder_paths
from .runtime import (
CachedModelNode,
ManagedTorchModel,
batch_text,
inference_context,
model_device,
snapshot_download,
tensor_batch_to_pil,
torch_dtype,
)
THRESHOLD = 0.4
MODEL_ID = "fancyfeast/joytag"
# Define your local directory where you want to save the files
files_for_joytagger = Path(folder_paths.folder_names_and_paths["LLavacheckpoints"][0][0]) / "files_for_joytagger"
# Check if the directory exists, create if it doesn't (optional)
files_for_joytagger.mkdir(parents=True, exist_ok=True)
def download_joytag():
# Ensure the correct behavior based on the existence of the local directory
print(f"Target directory for download: {files_for_joytagger}")
# Call snapshot_download with specified parameters
path = snapshot_download(
"fancyfeast/joytag", # Example repo_id
local_dir=files_for_joytagger,
force_download=False, # Set to True if you always want to download, regardless of local copy
local_files_only=False, # Set to False to allow downloading if not available locally
local_dir_use_symlinks="auto" # or set to True/False based on your symlink preference
)
print(f"Model path: {path}")
return path
def prepare_image(image: Image.Image, target_size: int) -> torch.Tensor:
width, height = image.size
side = max(width, height)
canvas = Image.new("RGB", (side, side), (255, 255, 255))
canvas.paste(image.convert("RGB"), ((side - width) // 2, (side - height) // 2))
if side != target_size:
canvas = canvas.resize(
(target_size, target_size), Image.Resampling.BICUBIC
)
array = np.asarray(canvas, dtype=np.float32) / 255.0
tensor = torch.from_numpy(array.copy()).permute(2, 0, 1)
mean = torch.tensor([0.48145466, 0.4578275, 0.40821073])[:, None, None]
std = torch.tensor([0.26862954, 0.26130258, 0.27577711])[:, None, None]
return (tensor - mean) / std
# Pad image to square
image_shape = image.size
max_dim = max(image_shape)
pad_left = (max_dim - image_shape[0]) // 2
pad_top = (max_dim - image_shape[1]) // 2
padded_image = Image.new('RGB', (max_dim, max_dim), (255, 255, 255))
padded_image.paste(image, (pad_left, pad_top))
# Resize image
if max_dim != target_size:
padded_image = padded_image.resize((target_size, target_size), Image.BICUBIC)
# Convert to tensor
image_tensor = TVF.pil_to_tensor(padded_image) / 255.0
# Normalize
image_tensor = TVF.normalize(image_tensor, mean=[0.48145466, 0.4578275, 0.40821073], std=[0.26862954, 0.26130258, 0.27577711])
return image_tensor
def clean_tag(tag: str) -> str:
return (
tag.replace("(medium)", "")
.replace("\\", "")
.replace("m/", "")
.replace("_", " ")
.strip(" -")
)
class JoyTagPredictor:
def __init__(self):
from .joytagger import Models
# Extract and process the tags
def process_tag(tag):
tag = tag.replace("(medium)", "") # Remove (medium)
tag = tag.replace("\\", "") # Remove \
tag = tag.replace("m/", "") # Remove m/
tag = tag.replace("-", "") # Remove -
tag = tag.replace("_", " ") # Replace underscores with spaces
tag = tag.strip() # Remove leading and trailing spaces
return tag
path = snapshot_download(MODEL_ID, "joytag")
model = Models.VisionModel.load_model(path, device=None).eval()
self.tags = [
line.strip()
for line in (path / "top_tags.txt").read_text(
encoding="utf-8"
).splitlines()
if line.strip()
]
self.dtype = torch_dtype("float16")
self.handle = ManagedTorchModel(model)
class Joytag:
def __init__(self):
pass
def close(self):
self.handle.close()
self.tags = []
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"tag_number": ("INT", {
"default": 1,
"min": 1, #Minimum value
"max": 100, #Maximum value
"step": 1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
},
}
def predict(self, images, count: int, threshold: float):
results = []
for image in tensor_batch_to_pil(images):
model = self.handle.ensure_loaded()
device = model_device(model)
tensor = prepare_image(image, model.image_size).unsqueeze(0).to(device)
with torch.inference_mode(), inference_context(device, self.dtype):
predictions = model({"image": tensor})["tags"].sigmoid()[0]
scores = predictions.float().cpu()
ranked = torch.argsort(scores, descending=True).tolist()
selected = [
index
for index in ranked
if scores[index].item() >= float(threshold)
][: int(count)]
# Always return up to tag_number useful results, even when the
# threshold is deliberately high.
if not selected:
selected = ranked[: int(count)]
tags = [clean_tag(self.tags[index]) for index in selected]
results.append(", ".join(tag for tag in tags if tag))
return batch_text(results)
RETURN_TYPES = ("STRING",)
FUNCTION = "tags"
class Joytag(CachedModelNode):
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"tag_number": (
"INT",
{
"default": 20,
"min": 1,
"max": 100,
"step": 1,
"display": "number",
},
),
},
"optional": {
"threshold": (
"FLOAT",
{"default": 0.4, "min": 0.0, "max": 1.0, "step": 0.01},
),
"unload_after": ("BOOLEAN", {"default": False}),
},
}
CATEGORY = "VLM Nodes/JoyTag"
RETURN_TYPES = ("STRING",)
FUNCTION = "tags"
CATEGORY = "VLM Nodes/Vision/Tagging"
def tags(self, image, tag_number):
path = download_joytag()
print(f"Model path: {path}")
model = Models.VisionModel.load_model(Path(path), device='cuda')
model.eval()
with open(Path(path) / 'top_tags.txt', 'r') as f:
top_tags = [line.strip() for line in f.readlines() if line.strip()]
def tags(
self,
image,
tag_number,
threshold=0.4,
unload_after=False,
):
predictor = self.get_or_create_model(MODEL_ID, JoyTagPredictor)
try:
return (
predictor.predict(image, tag_number, threshold),
)
finally:
self.maybe_clear_model(unload_after)
@torch.no_grad()
def predict(image: Image.Image):
image_tensor = prepare_image(image, model.image_size)
batch = {
'image': image_tensor.unsqueeze(0).to('cuda'),
}
with torch.amp.autocast_mode.autocast('cuda', enabled=True):
preds = model(batch)
tag_preds = preds['tags'].sigmoid().cpu()
scores = {top_tags[i]: tag_preds[0][i] for i in range(len(top_tags))}
predicted_tags = [tag for tag, score in scores.items() if score > THRESHOLD]
tag_string = ', '.join(predicted_tags)
return tag_string, scores
image = transforms.ToPILImage()(image[0].permute(2, 0, 1))
_, scores = predict(image)
# Get the top 50 tag and score pairs
top_tags_scores = sorted(scores.items(), key=lambda x: x[1], reverse=True)[:tag_number]
# Extract the tags from the pairs
top_tags_processed = [process_tag(tag) for tag, _ in top_tags_scores]
top_tags_full = [tag for tag in top_tags_processed if tag]
# Concatenate the tags with a comma separator
top_50_tags_string = ', '.join(top_tags_full)
return (top_50_tags_string, )
# A dictionary that contains all nodes you want to export with their names
NODE_CLASS_MAPPINGS = {"Joytag": Joytag}
NODE_DISPLAY_NAME_MAPPINGS = {"Joytag": "JoyTag"}
# A dictionary that contains the friendly/humanly readable titles for the nodes
NODE_DISPLAY_NAME_MAPPINGS = {"Joytag": "Joytag Node"}
+7 -4
View File
@@ -2,6 +2,7 @@ import json
from pathlib import Path
from typing import Optional
import torch
import torch.backends.cuda
import torch.nn as nn
import torch.nn.functional as F
import torchvision
@@ -210,8 +211,9 @@ class FastCLIPAttention2(nn.Module):
v_states = v_states.view(bsz, src_len, self.num_heads, self.head_dim).transpose(1, 2) # (bsz, num_heads, src_len, head_dim)
# Performs scale of query_states, attention, and softmax
x = F.scaled_dot_product_attention(q_states, k_states, v_states) # (bsz, num_heads, tgt_len, head_dim)
x = x.transpose(1, 2).contiguous().view(bsz, tgt_len, embed_dim) # (bsz, tgt_len, embed_dim)
with torch.backends.cuda.sdp_kernel(enable_math=False):
x = F.scaled_dot_product_attention(q_states, k_states, v_states) # (bsz, num_heads, tgt_len, head_dim)
x = x.transpose(1, 2).contiguous().view(bsz, tgt_len, embed_dim) # (bsz, tgt_len, embed_dim)
# Projection
x = self.out_proj(x) # (bsz, tgt_len, out_dim)
@@ -863,8 +865,9 @@ class ViTBlock(nn.Module):
k_states = qkv_states[1].view(bsz, src_len, self.num_heads, embed_dim // self.num_heads).transpose(1, 2) # (bsz, num_heads, src_len, embed_dim // num_heads)
v_states = qkv_states[2].view(bsz, src_len, self.num_heads, embed_dim // self.num_heads).transpose(1, 2) # (bsz, num_heads, src_len, embed_dim // num_heads)
out = F.scaled_dot_product_attention(q_states, k_states, v_states) # (bsz, num_heads, tgt_len, head_dim)
out = out.transpose(1, 2).contiguous().view(bsz, src_len, embed_dim) # (bsz, tgt_len, embed_dim)
with torch.backends.cuda.sdp_kernel(enable_math=False):
out = F.scaled_dot_product_attention(q_states, k_states, v_states) # (bsz, num_heads, tgt_len, head_dim)
out = out.transpose(1, 2).contiguous().view(bsz, src_len, embed_dim) # (bsz, tgt_len, embed_dim)
out = self.out_proj(out)
+61 -98
View File
@@ -1,84 +1,59 @@
"""Kosmos-2 grounding/caption node with lazy, Comfy-managed loading."""
from __future__ import annotations
from transformers import AutoModelForVision2Seq, AutoProcessor
from PIL import Image
from pathlib import Path
import torch
from .runtime import (
CachedModelNode,
ManagedTorchModel,
batch_text,
inference_context,
model_device,
move_inputs,
require_module,
snapshot_download,
tensor_batch_to_pil,
torch_dtype,
)
MODEL_ID = "microsoft/kosmos-2-patch14-224"
from torchvision.transforms import ToPILImage
from huggingface_hub import snapshot_download
import folder_paths
# Define the directory for saving files related to your new model
files_for_new_model = Path(folder_paths.folder_names_and_paths["LLavacheckpoints"][0][0]) / "files_for_kosmos2"
files_for_new_model.mkdir(parents=True, exist_ok=True) # Ensure the directory exists
class KosmosModelPredictor:
def __init__(self):
transformers = require_module("transformers")
model_path = snapshot_download(
MODEL_ID, "kosmos2", ignore_patterns=["*.bin"]
self.model_path = snapshot_download("microsoft/kosmos-2-patch14-224",
local_dir=files_for_new_model,
force_download=False, # Set to True if you always want to download, regardless of local copy
local_files_only=False, # Set to False to allow downloading if not available locally
local_dir_use_symlinks="auto",
ignore_patterns=["*.bin", "*.jpg", "*.png"]) # or set to True/False based on your symlink preference
self.device = "cuda:0" if torch.cuda.is_available() else "cpu"
self.model = AutoModelForVision2Seq.from_pretrained(self.model_path).to(self.device)
self.processor = AutoProcessor.from_pretrained(self.model_path)
def generate_predictions(self, image_path, main_text):
# Load the image
image_input = Image.open(image_path).convert("RGB")
text_input = f"<grounding>{main_text}: "
# Process the inputs
inputs = self.processor(text=text_input, images=image_input, return_tensors="pt").to(self.device)
# Generate predictions
generated_ids = self.model.generate(
pixel_values=inputs["pixel_values"],
input_ids=inputs["input_ids"],
attention_mask=inputs["attention_mask"],
image_embeds=None,
image_embeds_position_mask=inputs["image_embeds_position_mask"],
use_cache=True,
max_new_tokens=128,
)
self.dtype = torch_dtype("bfloat16")
model_class = getattr(
transformers,
"Kosmos2ForConditionalGeneration",
getattr(transformers, "AutoModelForImageTextToText", None),
)
if model_class is None:
raise RuntimeError(
"This Transformers version does not include Kosmos-2 support."
)
model = model_class.from_pretrained(
model_path, torch_dtype=self.dtype
).eval()
self.processor = transformers.AutoProcessor.from_pretrained(model_path)
self.handle = ManagedTorchModel(model, processor=self.processor)
def close(self):
self.handle.close()
self.processor = None
# Decode the generated IDs
generated_text = self.processor.batch_decode(generated_ids, skip_special_tokens=True)[0]
def generate(self, images, text, max_new_tokens):
results = []
for image in tensor_batch_to_pil(images):
prompt = f"<grounding>{text.strip()}"
inputs = self.processor(
text=prompt, images=image, return_tensors="pt"
)
model = self.handle.ensure_loaded()
device = model_device(model)
inputs = move_inputs(inputs, device)
with torch.inference_mode(), inference_context(device, self.dtype):
output = model.generate(
**inputs,
use_cache=True,
max_new_tokens=int(max_new_tokens),
)
decoded = self.processor.batch_decode(
output, skip_special_tokens=True
)[0]
post_process = getattr(
self.processor, "post_process_generation", None
)
if callable(post_process):
processed, _entities = post_process(decoded)
else:
processed = decoded
if processed.startswith(text):
processed = processed[len(text) :].lstrip(": \n")
results.append(processed.strip())
return batch_text(results)
# By default, the generated text is cleanup and the entities are extracted.
processed_text, entities = self.processor.post_process_generation(generated_text)
return processed_text[len(main_text)+2:]
# Example of integrating NewModelPredictor into a node-like structure
class Kosmos2model:
def __init__(self):
self.predictor = KosmosModelPredictor()
class Kosmos2model(CachedModelNode):
@classmethod
def INPUT_TYPES(cls):
return {
@@ -86,39 +61,27 @@ class Kosmos2model(CachedModelNode):
"image": ("IMAGE",),
"text_input": (
"STRING",
{"multiline": True, "default": "Describe the image."},
{
"multiline": True,
"default": "",
},
),
},
"optional": {
"max_new_tokens": (
"INT",
{"default": 128, "min": 1, "max": 2048},
),
"unload_after": ("BOOLEAN", {"default": False}),
},
}
RETURN_TYPES = ("STRING",)
FUNCTION = "new_model_generate_predictions"
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
def new_model_generate_predictions(
self,
image,
text_input,
max_new_tokens=128,
unload_after=False,
):
predictor = self.get_or_create_model(
MODEL_ID, KosmosModelPredictor
)
try:
return (
predictor.generate(image, text_input, max_new_tokens),
)
finally:
self.maybe_clear_model(unload_after)
CATEGORY = "VLM Nodes/Kosmos-2"
def new_model_generate_predictions(self, image, text_input):
pil_image = ToPILImage()(image[0].permute(2, 0, 1))
temp_path = files_for_new_model / "temp_image.png"
pil_image.save(temp_path)
response = self.predictor.generate_predictions(temp_path, text_input)
return (response, )
NODE_CLASS_MAPPINGS = {"Kosmos2model": Kosmos2model}
NODE_DISPLAY_NAME_MAPPINGS = {"Kosmos2model": "Kosmos-2"}
NODE_DISPLAY_NAME_MAPPINGS = {"Kosmos2model": "Kosmos-2 Node"}
+298 -516
View File
@@ -1,198 +1,76 @@
"""llama.cpp multimodal nodes with lazy loading and owned GPU cleanup."""
from __future__ import annotations
from typing import Any
import folder_paths
from .runtime import (
LLAMA_VISION_HANDLER_CHOICES,
LlamaHandle,
LlavaClipConfig,
batch_text,
close_handle,
default_llama_threads,
image_data_uri,
llama_chat_content,
llama_runtime_input_types,
llama_runtime_options,
resolve_model_path,
tensor_batch_to_pil,
unwrap_llm,
)
import os
from io import BytesIO
from llama_cpp import Llama
from llama_cpp.llama_chat_format import Llava15ChatHandler
import base64
from torchvision.transforms import ToPILImage
import gc
import torch
def _clip_factory(clip: Any):
if isinstance(clip, LlavaClipConfig):
return clip.create
if callable(getattr(clip, "create", None)):
return clip.create
# Compatibility with workflows that pass a pre-created llama.cpp handler.
return lambda: clip
def _make_handle(
ckpt_name: str,
max_ctx: int,
gpu_layers: int,
n_threads: int,
clip: Any,
*,
seed: int = 42,
runtime_options: dict[str, Any] | None = None,
) -> LlamaHandle:
options = dict(runtime_options or {})
if isinstance(clip, LlavaClipConfig):
options.setdefault("projector_path", clip.model_path)
return LlamaHandle(
resolve_model_path(ckpt_name),
n_ctx=max_ctx,
n_gpu_layers=gpu_layers,
n_threads=n_threads,
chat_handler_factory=_clip_factory(clip),
seed=seed,
**options,
)
def _vision_messages(system_msg: str, prompt: str, data_uri: str):
return [
{"role": "system", "content": system_msg},
{
"role": "user",
"content": [
{"type": "image_url", "image_url": {"url": data_uri}},
{"type": "text", "text": prompt},
],
},
]
def _run_batch(
image,
model,
*,
system_msg: str,
prompt: str,
**generation: Any,
) -> str:
llm = unwrap_llm(model)
responses = []
for pil_image in tensor_batch_to_pil(image):
response = llm.create_chat_completion(
messages=_vision_messages(system_msg, prompt, image_data_uri(pil_image)),
**generation,
)
responses.append(llama_chat_content(response))
return batch_text(responses)
supported_LLava_extensions = set(['.gguf'])
try:
folder_paths.folder_names_and_paths["LLavacheckpoints"] = (folder_paths.folder_names_and_paths["LLavacheckpoints"][0], supported_LLava_extensions)
except:
# check if LLavacheckpoints exists otherwise create
if not os.path.isdir(os.path.join(folder_paths.models_dir, "LLavacheckpoints")):
os.mkdir(os.path.join(folder_paths.models_dir, "LLavacheckpoints"))
folder_paths.folder_names_and_paths["LLavacheckpoints"] = ([os.path.join(folder_paths.models_dir, "LLavacheckpoints")], supported_LLava_extensions)
class LLavaLoader:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"ckpt_name": (folder_paths.get_filename_list("LLavacheckpoints"),),
"max_ctx": (
"INT",
{"default": 4096, "min": 128, "max": 131072, "step": 64},
),
"gpu_layers": (
"INT",
{"default": -1, "min": -1, "max": 1000, "step": 1},
),
"n_threads": (
"INT",
{
"default": default_llama_threads(),
"min": 1,
"max": 256,
"step": 1,
},
),
"clip": ("CUSTOM", {"default": ""}),
},
"optional": llama_runtime_input_types(),
}
def INPUT_TYPES(s):
return {"required": {
"ckpt_name": (folder_paths.get_filename_list("LLavacheckpoints"), ),
"max_ctx": ("INT", {"default": 4096, "min": 128, "max": 8192, "step": 64}),
"gpu_layers": ("INT", {"default": 27, "min": 0, "max": 100, "step": 1}),
"n_threads": ("INT", {"default": 8, "min": 1, "max": 100, "step": 1}),
"clip": ("CUSTOM", {"default": ""}),
}}
RETURN_TYPES = ("CUSTOM",)
RETURN_NAMES = ("model",)
FUNCTION = "load_llava_checkpoint"
CATEGORY = "VLM Nodes/LLava"
def load_llava_checkpoint(
self,
ckpt_name,
max_ctx,
gpu_layers,
n_threads,
clip,
n_batch=512,
n_ubatch=512,
flash_attention="Auto",
use_mmap=True,
split_mode="Layer",
main_gpu=0,
tensor_split="",
):
# The GGUF and mmproj are loaded only when a sampler actually executes.
return (
_make_handle(
ckpt_name,
max_ctx,
gpu_layers,
n_threads,
clip,
runtime_options=llama_runtime_options(
n_batch=n_batch,
n_ubatch=n_ubatch,
flash_attention=flash_attention,
use_mmap=use_mmap,
split_mode=split_mode,
main_gpu=main_gpu,
tensor_split=tensor_split,
),
),
)
def load_llava_checkpoint(self, ckpt_name, max_ctx, gpu_layers, n_threads, clip ):
ckpt_path = folder_paths.get_full_path("LLavacheckpoints", ckpt_name)
llm = Llama(model_path = ckpt_path, chat_handler=clip,offload_kqv=True, f16_kv=True, use_mlock=False, embedding=False, n_batch=1024, last_n_tokens_size=1024, verbose=True, seed=42, n_ctx = max_ctx, n_gpu_layers=gpu_layers, n_threads=n_threads, logits_all=True, echo=False)
return (llm, )
class LlavaClipLoader:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"clip_name": (folder_paths.get_filename_list("LLavacheckpoints"),),
},
"optional": {
"handler": (
list(LLAMA_VISION_HANDLER_CHOICES),
{"default": "Auto (GGUF chat template)"},
),
},
}
RETURN_TYPES = ("CUSTOM",)
RETURN_NAMES = ("clip",)
def INPUT_TYPES(s):
return {"required": {
"clip_name": (folder_paths.get_filename_list("LLavacheckpoints"), ),
}}
RETURN_TYPES = ("CUSTOM", )
RETURN_NAMES = ("clip", )
FUNCTION = "load_clip_checkpoint"
CATEGORY = "VLM Nodes/LLava"
def load_clip_checkpoint(self, clip_name):
clip_path = folder_paths.get_full_path("LLavacheckpoints", clip_name)
clip = Llava15ChatHandler(clip_model_path = clip_path, verbose=False)
return (clip, )
def load_clip_checkpoint(self, clip_name, handler="LLaVA 1.5"):
return (LlavaClipConfig(resolve_model_path(clip_name), handler),)
class LLavaSamplerSimple:
class LLavaSamplerSimple:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"prompt": ("STRING", {"default": "", "multiline": True}),
"prompt": ("STRING",{"forceInput": True} ),
"model": ("CUSTOM", {"default": ""}),
"temperature": (
"FLOAT",
{"default": 0.1, "min": 0.0, "max": 2.0, "step": 0.01},
),
"temperature": ("FLOAT", {"default": 0.1, "min": 0.01, "max": 1.0, "step": 0.01}),
}
}
@@ -201,62 +79,62 @@ class LLavaSamplerSimple:
CATEGORY = "VLM Nodes/LLava"
def generate_text(self, image, prompt, model, temperature):
return (
_run_batch(
image,
model,
system_msg="You are an assistant who accurately describes images.",
prompt=prompt,
temperature=temperature,
),
# Assuming 'image' is a PyTorch tensor of shape [C, H, W]
# Convert the PyTorch tensor to a PIL image
pil_image = ToPILImage()(image[0].permute(2, 0, 1))
# Convert the PIL image to a bytes buffer
buffer = BytesIO()
pil_image.save(buffer, format="PNG") # You can change the format if needed
# Get the bytes from the buffer
image_bytes = buffer.getvalue()
# Encode the bytes to base64
base64_string = f"data:image/jpeg;base64,{base64.b64encode(image_bytes).decode('utf-8')}"
# Now, `base64_string` contains the base64-encoded string of the image
llm = model
response = llm.create_chat_completion(
messages = [
{"role": "system", "content": "You are an assistant who perfectly describes images."},
{
"role": "user",
"content": [
{"type": "image_url", "image_url": {"url" : base64_string}},
{"type" : "text", "text": f"{prompt}"}
]
}
],
temperature = temperature,
)
class LLavaSamplerAdvanced:
return (f"{response['choices'][0]['message']['content']}", )
class LLavaSamplerAdvanced:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"system_msg": (
"STRING",
{
"default": (
"You are an assistant who accurately describes images."
)
},
),
"prompt": (
"STRING",
{"default": "", "multiline": True},
),
"system_msg": ("STRING",{"default" : "You are an assistant who perfectly describes images."}),
"prompt": ("STRING",{"forceInput": True, "default": ""}),
"model": ("CUSTOM", {"default": ""}),
"max_tokens": (
"INT",
{"default": 512, "min": 1, "max": 8192, "step": 1},
),
"temperature": (
"FLOAT",
{"default": 0.1, "min": 0.0, "max": 2.0, "step": 0.01},
),
"top_p": (
"FLOAT",
{"default": 0.95, "min": 0.0, "max": 1.0, "step": 0.01},
),
"top_k": ("INT", {"default": 40, "min": 0, "step": 1}),
"frequency_penalty": (
"FLOAT",
{"default": 0.0, "min": -2.0, "max": 2.0, "step": 0.01},
),
"presence_penalty": (
"FLOAT",
{"default": 0.0, "min": -2.0, "max": 2.0, "step": 0.01},
),
"repeat_penalty": (
"FLOAT",
{"default": 1.1, "min": 0.0, "max": 2.0, "step": 0.01},
),
"seed": ("INT", {"default": 42, "step": 1}),
"max_tokens": ("INT", {"default": 512, "min": 1, "max": 2048, "step": 1}),
"temperature": ("FLOAT", {"default": 0.1, "min": 0.01, "max": 1.0, "step": 0.01}),
"top_p": ("FLOAT", {"default": 0.95, "min": 0.1, "max": 1.0, "step": 0.01}),
"top_k": ("INT", {"default": 40, "step": 1}),
"frequency_penalty": ("FLOAT", {"default": 0.0, "step": 0.01}),
"presence_penalty": ("FLOAT", {"default": 0.0, "step": 0.01}),
"repeat_penalty": ("FLOAT", {"default": 1.1, "step": 0.01}),
"seed": ("INT", {"default": 42, "step":1})
}
}
@@ -264,321 +142,225 @@ class LLavaSamplerAdvanced:
FUNCTION = "generate_text_advanced"
CATEGORY = "VLM Nodes/LLava"
def generate_text_advanced(
self,
image,
system_msg,
prompt,
model,
max_tokens,
temperature,
top_p,
top_k,
frequency_penalty,
presence_penalty,
repeat_penalty,
seed,
):
return (
_run_batch(
image,
model,
system_msg=system_msg,
prompt=prompt,
max_tokens=max_tokens,
temperature=temperature,
top_p=top_p,
top_k=top_k,
frequency_penalty=frequency_penalty,
presence_penalty=presence_penalty,
repeat_penalty=repeat_penalty,
seed=seed,
),
def generate_text_advanced(self, image, system_msg, prompt, model, max_tokens, temperature, top_p, frequency_penalty, presence_penalty, repeat_penalty, top_k,seed):
# Assuming 'image' is a PyTorch tensor of shape [C, H, W]
# Convert the PyTorch tensor to a PIL image
pil_image = ToPILImage()(image[0].permute(2, 0, 1))
# Convert the PIL image to a bytes buffer
buffer = BytesIO()
pil_image.save(buffer, format="PNG") # You can change the format if needed
# Get the bytes from the buffer
image_bytes = buffer.getvalue()
# Encode the bytes to base64
base64_string = f"data:image/jpeg;base64,{base64.b64encode(image_bytes).decode('utf-8')}"
# Now, `base64_string` contains the base64-encoded string of the image
llm = model
response = llm.create_chat_completion(
messages = [
{"role": "system", "content": system_msg},
{
"role": "user",
"content": [
{"type": "image_url", "image_url": {"url" : base64_string}},
{"type" : "text", "text": f"{prompt}"}
]
}
],
max_tokens = max_tokens,
temperature = temperature,
top_p = top_p,
top_k = top_k,
frequency_penalty = frequency_penalty,
presence_penalty = presence_penalty,
repeat_penalty = repeat_penalty,
seed=seed
)
class _CachedLlavaBase:
return (f"{response['choices'][0]['message']['content']}", )
class LLavaOptionalMemoryFreeSimple:
def __init__(self):
self._handle = None
self._key = None
self.llm = None # Store the model instance
self.clip = None # Store the clip instance
def _model(
self,
ckpt_name,
clip_name,
max_ctx,
gpu_layers,
n_threads,
seed=42,
handler="LLaVA 1.5",
**runtime_options,
):
key = (
ckpt_name,
clip_name,
int(max_ctx),
int(gpu_layers),
int(n_threads),
int(seed),
handler,
tuple(sorted(runtime_options.items())),
)
if self._handle is None or self._key != key:
close_handle(self._handle)
clip = LlavaClipConfig(resolve_model_path(clip_name), handler)
self._handle = _make_handle(
ckpt_name,
max_ctx,
gpu_layers,
n_threads,
clip,
seed=seed,
runtime_options=runtime_options,
)
self._key = key
return self._handle
def _maybe_unload(self, unload):
if unload:
close_handle(self._handle)
self._handle = None
self._key = None
class LLavaOptionalMemoryFreeSimple(_CachedLlavaBase):
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"ckpt_name": (folder_paths.get_filename_list("LLavacheckpoints"),),
"clip_name": (folder_paths.get_filename_list("LLavacheckpoints"),),
"max_ctx": (
"INT",
{"default": 4096, "min": 128, "max": 131072, "step": 64},
),
"gpu_layers": (
"INT",
{"default": -1, "min": -1, "max": 1000, "step": 1},
),
"n_threads": (
"INT",
{
"default": default_llama_threads(),
"min": 1,
"max": 256,
"step": 1,
},
),
"ckpt_name": (folder_paths.get_filename_list("LLavacheckpoints"), ),
"clip_name": (folder_paths.get_filename_list("LLavacheckpoints"), ),
"max_ctx": ("INT", {"default": 4096, "min": 128, "max": 128000, "step": 64}),
"gpu_layers": ("INT", {"default": 27, "min": 0, "max": 100, "step": 1}),
"n_threads": ("INT", {"default": 8, "min": 1, "max": 100, "step": 1}),
"image": ("IMAGE",),
"prompt": ("STRING", {"default": "", "multiline": True}),
"temperature": (
"FLOAT",
{"default": 0.1, "min": 0.0, "max": 2.0, "step": 0.01},
),
"unload": ("BOOLEAN", {"default": False}),
},
"optional": {
"handler": (
list(LLAMA_VISION_HANDLER_CHOICES),
{"default": "Auto (GGUF chat template)"},
),
**llama_runtime_input_types(),
},
"prompt": ("STRING", {"forceInput": True}),
"temperature": ("FLOAT", {"default": 0.1, "min": 0.01, "max": 1.0, "step": 0.01}),
"unload": ("BOOLEAN", {"default": False}), # Add unload parameter
}
}
RETURN_TYPES = ("STRING",)
FUNCTION = "generate_text"
CATEGORY = "VLM Nodes/LLava"
def generate_text(
self,
ckpt_name,
clip_name,
max_ctx,
gpu_layers,
n_threads,
image,
prompt,
temperature,
unload,
handler="LLaVA 1.5",
n_batch=512,
n_ubatch=512,
flash_attention="Auto",
use_mmap=True,
split_mode="Layer",
main_gpu=0,
tensor_split="",
):
options = llama_runtime_options(
n_batch=n_batch,
n_ubatch=n_ubatch,
flash_attention=flash_attention,
use_mmap=use_mmap,
split_mode=split_mode,
main_gpu=main_gpu,
tensor_split=tensor_split,
)
model = self._model(
ckpt_name,
clip_name,
max_ctx,
gpu_layers,
n_threads,
handler=handler,
**options,
)
try:
result = _run_batch(
image,
model,
system_msg="You are an assistant who accurately describes images.",
prompt=prompt,
temperature=temperature,
)
return (result,)
finally:
self._maybe_unload(unload)
def generate_text(self, ckpt_name, clip_name, max_ctx, gpu_layers, n_threads, image, prompt, temperature, unload):
# Load the model
# Load the clip
clip_path = folder_paths.get_full_path("LLavacheckpoints", clip_name)
self.clip = Llava15ChatHandler(clip_model_path=clip_path, verbose=False)
# Load model
ckpt_path = folder_paths.get_full_path("LLavacheckpoints", ckpt_name)
self.llm = Llama(model_path = ckpt_path, chat_handler=self.clip, offload_kqv=True, f16_kv=True, use_mlock=False, embedding=False, n_batch=1024, last_n_tokens_size=1024, verbose=True, seed=42, n_ctx = max_ctx, n_gpu_layers=gpu_layers, n_threads=n_threads, logits_all=True, echo=False)
# Assuming 'image' is a PyTorch tensor of shape [C, H, W]
# Convert the PyTorch tensor to a PIL image
pil_image = ToPILImage()(image[0].permute(2, 0, 1))
# Convert the PIL image to a bytes buffer
buffer = BytesIO()
pil_image.save(buffer, format="PNG") # You can change the format if needed
# Get the bytes from the buffer
image_bytes = buffer.getvalue()
# Encode the bytes to base64
base64_string = f"data:image/jpeg;base64,{base64.b64encode(image_bytes).decode('utf-8')}"
# Now, `base64_string` contains the base64-encoded string of the image
response = self.llm.create_chat_completion(
messages=[
{"role": "system", "content": "You are an assistant who perfectly describes images."},
{
"role": "user",
"content": [
{"type": "image_url", "image_url": {"url": base64_string}},
{"type": "text", "text": f"{prompt}"}
]
}
],
temperature=temperature,
)
if unload and self.llm is not None:
del self.llm # Unload the model
self.llm = None # Remove reference to the model
gc.collect()
torch.cuda.empty_cache()
if unload and self.clip is not None:
del self.clip # Unload the clip
self.clip = None # Remove reference to the clip
gc.collect()
torch.cuda.empty_cache()
return (f"{response['choices'][0]['message']['content']}", )
class LLavaOptionalMemoryFreeAdvanced:
def __init__(self):
self.llm = None # Store the model instance
self.clip = None # Store the clip instance
class LLavaOptionalMemoryFreeAdvanced(_CachedLlavaBase):
@classmethod
def INPUT_TYPES(cls):
required = {
"ckpt_name": (folder_paths.get_filename_list("LLavacheckpoints"),),
"clip_name": (folder_paths.get_filename_list("LLavacheckpoints"),),
"max_ctx": (
"INT",
{"default": 4096, "min": 128, "max": 131072, "step": 64},
),
"gpu_layers": (
"INT",
{"default": -1, "min": -1, "max": 1000, "step": 1},
),
"n_threads": (
"INT",
{
"default": default_llama_threads(),
"min": 1,
"max": 256,
"step": 1,
},
),
"image": ("IMAGE",),
"system_msg": (
"STRING",
{"default": ("You are an assistant who accurately describes images.")},
),
"prompt": ("STRING", {"default": "", "multiline": True}),
"max_tokens": (
"INT",
{"default": 512, "min": 1, "max": 8192, "step": 1},
),
"temperature": (
"FLOAT",
{"default": 0.1, "min": 0.0, "max": 2.0, "step": 0.01},
),
"top_p": (
"FLOAT",
{"default": 0.95, "min": 0.0, "max": 1.0, "step": 0.01},
),
"top_k": ("INT", {"default": 40, "min": 0, "step": 1}),
"frequency_penalty": (
"FLOAT",
{"default": 0.0, "min": -2.0, "max": 2.0, "step": 0.01},
),
"presence_penalty": (
"FLOAT",
{"default": 0.0, "min": -2.0, "max": 2.0, "step": 0.01},
),
"repeat_penalty": (
"FLOAT",
{"default": 1.1, "min": 0.0, "max": 2.0, "step": 0.01},
),
"seed": ("INT", {"default": 42, "step": 1}),
"unload": ("BOOLEAN", {"default": False}),
}
return {
"required": required,
"optional": {
"handler": (
list(LLAMA_VISION_HANDLER_CHOICES),
{"default": "Auto (GGUF chat template)"},
),
**llama_runtime_input_types(),
},
"required": {
"ckpt_name": (folder_paths.get_filename_list("LLavacheckpoints"), ),
"clip_name": (folder_paths.get_filename_list("LLavacheckpoints"), ),
"max_ctx": ("INT", {"default": 4096, "min": 128, "max": 128000, "step": 64}),
"gpu_layers": ("INT", {"default": 27, "min": 0, "max": 100, "step": 1}),
"n_threads": ("INT", {"default": 8, "min": 1, "max": 100, "step": 1}),
"image": ("IMAGE",),
"system_msg": ("STRING", {"default": "You are an assistant who perfectly describes images."}),
"prompt": ("STRING", {"forceInput": True, "default": ""}),
"max_tokens": ("INT", {"default": 512, "min": 1, "max": 2048, "step": 1}),
"temperature": ("FLOAT", {"default": 0.1, "min": 0.01, "max": 1.0, "step": 0.01}),
"top_p": ("FLOAT", {"default": 0.95, "min": 0.1, "max": 1.0, "step": 0.01}),
"top_k": ("INT", {"default": 40, "step": 1}),
"frequency_penalty": ("FLOAT", {"default": 0.0, "step": 0.01}),
"presence_penalty": ("FLOAT", {"default": 0.0, "step": 0.01}),
"repeat_penalty": ("FLOAT", {"default": 1.1, "step": 0.01}),
"seed": ("INT", {"default": 42, "step": 1}),
"unload": ("BOOLEAN", {"default": False}), # Add unload parameter
}
}
RETURN_TYPES = ("STRING",)
FUNCTION = "generate_text_advanced"
CATEGORY = "VLM Nodes/LLava"
def generate_text_advanced(
self,
ckpt_name,
clip_name,
max_ctx,
gpu_layers,
n_threads,
image,
system_msg,
prompt,
max_tokens,
temperature,
top_p,
top_k,
frequency_penalty,
presence_penalty,
repeat_penalty,
seed,
unload,
handler="LLaVA 1.5",
n_batch=512,
n_ubatch=512,
flash_attention="Auto",
use_mmap=True,
split_mode="Layer",
main_gpu=0,
tensor_split="",
):
options = llama_runtime_options(
n_batch=n_batch,
n_ubatch=n_ubatch,
flash_attention=flash_attention,
use_mmap=use_mmap,
split_mode=split_mode,
main_gpu=main_gpu,
tensor_split=tensor_split,
)
model = self._model(
ckpt_name,
clip_name,
max_ctx,
gpu_layers,
n_threads,
seed,
handler,
**options,
)
try:
result = _run_batch(
image,
model,
system_msg=system_msg,
prompt=prompt,
max_tokens=max_tokens,
temperature=temperature,
top_p=top_p,
top_k=top_k,
frequency_penalty=frequency_penalty,
presence_penalty=presence_penalty,
repeat_penalty=repeat_penalty,
seed=seed,
)
return (result,)
finally:
self._maybe_unload(unload)
def generate_text_advanced(self, ckpt_name, clip_name, max_ctx, gpu_layers, n_threads, image, system_msg, prompt, max_tokens, temperature, top_p, top_k, frequency_penalty, presence_penalty, repeat_penalty, seed, unload):
# Load the clip
clip_path = folder_paths.get_full_path("LLavacheckpoints", clip_name)
self.clip = Llava15ChatHandler(clip_model_path=clip_path, verbose=False)
# Load model
ckpt_path = folder_paths.get_full_path("LLavacheckpoints", ckpt_name)
self.llm = Llama(model_path = ckpt_path, chat_handler=self.clip, offload_kqv=True, f16_kv=True, use_mlock=False, embedding=False, n_batch=1024, last_n_tokens_size=1024, verbose=True, seed=42, n_ctx = max_ctx, n_gpu_layers=gpu_layers, n_threads=n_threads, logits_all=True, echo=False)
# Assuming 'image' is a PyTorch tensor of shape [C, H, W]
# Convert the PyTorch tensor to a PIL image
pil_image = ToPILImage()(image[0].permute(2, 0, 1))
# Convert the PIL image to a bytes buffer
buffer = BytesIO()
pil_image.save(buffer, format="PNG") # You can change the format if needed
# Get the bytes from the buffer
image_bytes = buffer.getvalue()
# Encode the bytes to base64
base64_string = f"data:image/jpeg;base64,{base64.b64encode(image_bytes).decode('utf-8')}"
# Now, `base64_string` contains the base64-encoded string of the image
response = self.llm.create_chat_completion(
messages=[
{"role": "system", "content": system_msg},
{
"role": "user",
"content": [
{"type": "image_url", "image_url": {"url": base64_string}},
{"type": "text", "text": f"{prompt}"}
]
}
],
max_tokens=max_tokens,
temperature=temperature,
top_p=top_p,
top_k=top_k,
frequency_penalty=frequency_penalty,
presence_penalty=presence_penalty,
repeat_penalty=repeat_penalty,
seed=seed,
)
if unload and self.llm is not None:
del self.llm # Unload the model
self.llm = None # Remove reference to the model
gc.collect()
torch.cuda.empty_cache()
if unload and self.clip is not None:
del self.clip # Unload the clip
self.clip = None # Remove reference to the clip
gc.collect()
torch.cuda.empty_cache()
return (f"{response['choices'][0]['message']['content']}", )
NODE_CLASS_MAPPINGS = {
"LLava Loader Simple": LLavaLoader,
@@ -588,12 +370,12 @@ NODE_CLASS_MAPPINGS = {
"LLavaOptionalMemoryFreeSimple": LLavaOptionalMemoryFreeSimple,
"LLavaOptionalMemoryFreeAdvanced": LLavaOptionalMemoryFreeAdvanced,
}
# A dictionary that contains the friendly/humanly readable titles for the nodes
NODE_DISPLAY_NAME_MAPPINGS = {
"LLava Loader Simple": "LLaVA Loader",
"LLavaSamplerSimple": "LLaVA Sampler",
"LlavaClipLoader": "LLaVA Vision Projector Loader",
"LLavaSamplerAdvanced": "LLaVA Sampler (Advanced)",
"LLavaOptionalMemoryFreeSimple": "LLaVA (Managed Cache)",
"LLavaOptionalMemoryFreeAdvanced": "LLaVA (Managed Cache, Advanced)",
"LLava Loader Simple": "LLava Loader Simple",
"LLavaSamplerSimple": "LLava Sampler Simple",
"LlavaClipLoader": "Llava Clip Loader",
"LLavaSamplerAdvanced": "LLava Sampler Advanced",
"LLavaOptionalMemoryFreeSimple": "LLava Optional Memory Free Simple",
"LLavaOptionalMemoryFreeAdvanced": "LLava Optional Memory Free Advanced",
}
+66 -140
View File
@@ -1,162 +1,88 @@
"""MC-LLaVA node with in-memory images and ComfyUI-managed weights."""
from __future__ import annotations
from transformers import AutoModelForCausalLM, AutoProcessor
from PIL import Image
from pathlib import Path
import torch
from .runtime import (
CachedModelNode,
ManagedTorchModel,
batch_text,
inference_context,
model_device,
move_inputs,
require_module,
snapshot_download,
tensor_batch_to_pil,
torch_dtype,
)
MODEL_ID = "visheratin/MC-LLaVA-3b"
from huggingface_hub import snapshot_download
from torchvision.transforms import ToPILImage
import io
from PIL import Image
import folder_paths
# Define the directory for saving files related to the MCLLaVA model
files_for_mcllava_model = Path(folder_paths.folder_names_and_paths["LLavacheckpoints"][0][0]) / "files_for_mcllava"
files_for_mcllava_model.mkdir(parents=True, exist_ok=True) # Ensure the directory exists
class MCLLaVAModelPredictor:
def __init__(self):
transformers = require_module("transformers")
model_path = snapshot_download(
MODEL_ID, "mcllava", ignore_patterns=["*.bin"]
)
self.dtype = torch_dtype("float16")
model = transformers.AutoModelForCausalLM.from_pretrained(
model_path,
torch_dtype=self.dtype,
trust_remote_code=True,
).eval()
self.processor = transformers.AutoProcessor.from_pretrained(
model_path, trust_remote_code=True
)
self.handle = ManagedTorchModel(model, processor=self.processor)
self.model_path = snapshot_download("visheratin/MC-LLaVA-3b",
local_dir=files_for_mcllava_model,
force_download=False, # Set to True if you always want to download, regardless of local copy
local_files_only=False, # Set to False to allow downloading if not available locally
local_dir_use_symlinks="auto", # or set to True/False based on your symlink preference
ignore_patterns=["*.bin", "*.jpg", "*.png"]) # Exclude certain file types
self.device = "cuda" if torch.cuda.is_available() else "cpu"
self.model = AutoModelForCausalLM.from_pretrained(self.model_path, torch_dtype=torch.float16, trust_remote_code=True).to(self.device)
self.processor = AutoProcessor.from_pretrained(self.model_path, trust_remote_code=True)
def close(self):
self.handle.close()
self.processor = None
def generate_predictions(self, pil_image, prompt, temperature, top_p, max_crops, num_tokens):
# Load the image
# Save the PIL image to a bytes buffer instead of a file on disk.
buffer = io.BytesIO()
pil_image.save(buffer, format='PNG')
def generate(
self,
images,
prompt,
temperature,
top_p,
max_crops,
num_tokens,
max_new_tokens,
):
results = []
formatted = (
"<|im_start|>user\n<image>\n"
f"{prompt}<|im_end|>\n<|im_start|>assistant\n"
)
for image in tensor_batch_to_pil(images):
model = self.handle.ensure_loaded()
device = model_device(model)
inputs = self.processor(
formatted,
[image],
model,
max_crops=int(max_crops),
num_tokens=int(num_tokens),
)
inputs = move_inputs(inputs, device)
do_sample = float(temperature) > 0.0
generation = {
"max_new_tokens": int(max_new_tokens),
"do_sample": do_sample,
"use_cache": True,
"eos_token_id": self.processor.tokenizer.eos_token_id,
}
if do_sample:
generation.update(
temperature=float(temperature), top_p=float(top_p)
)
with torch.inference_mode(), inference_context(device, self.dtype):
output = model.generate(**inputs, **generation)
input_length = inputs["input_ids"].shape[-1]
text = self.processor.tokenizer.decode(
output[0, input_length:], skip_special_tokens=True
)
results.append(text.strip())
return batch_text(results)
# Move to the beginning of the buffer so Image.open can read from it.
buffer.seek(0)
# Open the image as if it was a 'raw' image from an HTTP response.
image_input = Image.open(buffer)
final_prompt = f"""<|im_start|>user
<image>
{prompt}<|im_end|>
<|im_start|>assistant
"""
with torch.inference_mode():
inputs = self.processor(final_prompt, [image_input], self.model, max_crops=max_crops, num_tokens=num_tokens)
with torch.inference_mode():
output = self.model.generate(**inputs, max_new_tokens=200, do_sample=False, use_cache=False, top_p=top_p, temperature=temperature, eos_token_id=self.processor.tokenizer.eos_token_id)
generated_text = self.processor.tokenizer.decode(output[0]).replace(final_prompt, "").replace("<|im_end|>", "")
return generated_text
# Example of integrating MCLLaVAModelPredictor into a node-like structure
class MCLLaVAModel:
def __init__(self):
self.predictor = MCLLaVAModelPredictor()
class MCLLaVAModel(CachedModelNode):
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"prompt": (
"STRING",
{"multiline": True, "default": "Describe the image."},
),
"temperature": (
"FLOAT",
{"default": 0.1, "min": 0.0, "max": 2.0, "step": 0.01},
),
"top_p": (
"FLOAT",
{"default": 0.9, "min": 0.0, "max": 1.0, "step": 0.01},
),
"max_crops": (
"INT",
{"default": 100, "min": 1, "max": 300, "step": 1},
),
"num_tokens": (
"INT",
{"default": 728, "min": 1, "max": 4096, "step": 1},
),
},
"optional": {
"max_new_tokens": (
"INT",
{"default": 200, "min": 1, "max": 4096},
),
"unload_after": ("BOOLEAN", {"default": False}),
"prompt": ( "STRING",{"multiline": True, "default": "", },),
"temperature": ( "FLOAT", {"default": 0.1, "min": 0.0, "max": 1.0, "step": 0.01},),
"top_p": ( "FLOAT", {"default": 0.9, "min": 0.0, "max": 1.0, "step": 0.01},),
"max_crops": ( "INT", {"default": 100, "min": 1, "max": 300, "step": 1},),
"num_tokens": ( "INT", {"default": 728, "min": 1, "max": 2048, "step": 1},),
},
}
RETURN_TYPES = ("STRING",)
FUNCTION = "generate_image_description"
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
def generate_image_description(
self,
image,
prompt,
temperature,
top_p,
max_crops,
num_tokens,
max_new_tokens=200,
unload_after=False,
):
predictor = self.get_or_create_model(
MODEL_ID, MCLLaVAModelPredictor
)
try:
return (
predictor.generate(
image,
prompt,
temperature,
top_p,
max_crops,
num_tokens,
max_new_tokens,
),
)
finally:
self.maybe_clear_model(unload_after)
CATEGORY = "VLM Nodes/MC-LLaVA"
def generate_image_description(self, image, prompt, temperature, top_p, max_crops, num_tokens):
pil_image = ToPILImage()(image[0].permute(2, 0, 1))
response = self.predictor.generate_predictions(pil_image, prompt, temperature, top_p, max_crops, num_tokens)
return (response, )
NODE_CLASS_MAPPINGS = {"MCLLaVAModel": MCLLaVAModel}
NODE_DISPLAY_NAME_MAPPINGS = {"MCLLaVAModel": "MC-LLaVA"}
NODE_DISPLAY_NAME_MAPPINGS = {"MCLLaVAModel": "MC-LLaVA Node"}
+161 -204
View File
@@ -1,231 +1,188 @@
"""MiniCPM-V 2.6 GGUF node using llama.cpp's native vision handler."""
import os
import subprocess
import torch
import numpy as np
from PIL import Image
from pathlib import Path
from huggingface_hub import hf_hub_download
import folder_paths
from transformers import AutoModel, AutoTokenizer
from __future__ import annotations
# Define the directory for saving MiniCPM files
MINICPM_PATH = Path(folder_paths.folder_names_and_paths["LLavacheckpoints"][0][0]) / "minicpm_files"
MINICPM_PATH.mkdir(parents=True, exist_ok=True)
from .runtime import (
CachedModelNode,
LlamaHandle,
LlavaClipConfig,
batch_text,
default_llama_threads,
hf_download,
image_data_uri,
llama_chat_content,
llama_runtime_input_types,
llama_runtime_options,
tensor_batch_to_pil,
)
MODEL_REPO = "openbmb/MiniCPM-V-2_6-gguf"
# Available GGUF model variants and their file sizes (in GB)
GGUF_MODELS = {
"Q2_K (3GB)": "ggml-model-Q2_K.gguf",
"Q3_K (3.8GB)": "ggml-model-Q3_K.gguf",
"Q4_K_M (4.7GB)": "ggml-model-Q4_K_M.gguf",
"Q5_K_M (5.4GB)": "ggml-model-Q5_K_M.gguf",
"Q8_0 (8.1GB)": "ggml-model-Q8_0.gguf",
"F16 (15.2GB)": "ggml-model-f16.gguf",
"F16 (15.2GB)": "ggml-model-f16.gguf"
}
class MiniCPMPredictor:
def __init__(
self,
model_variant,
context_length,
gpu_layers,
n_threads,
runtime_options=None,
):
model_path = hf_download(
MODEL_REPO,
GGUF_MODELS[model_variant],
"minicpm-v-2_6-gguf",
)
projector_path = hf_download(
MODEL_REPO,
"mmproj-model-f16.gguf",
"minicpm-v-2_6-gguf",
)
clip = LlavaClipConfig(projector_path, "MiniCPM-V 2.6")
def create_handler(*, use_gpu=True):
return clip.create(use_gpu=use_gpu)
self.handle = LlamaHandle(
model_path,
n_ctx=int(context_length),
n_gpu_layers=int(gpu_layers),
n_threads=int(n_threads),
chat_handler_factory=create_handler,
projector_path=projector_path,
**dict(runtime_options or {}),
)
def close(self):
self.handle.close()
def generate(
self,
images,
prompt,
temperature,
top_p,
top_k,
repeat_penalty,
max_tokens,
):
llm = self.handle.ensure_loaded()
results = []
for image in tensor_batch_to_pil(images):
response = llm.create_chat_completion(
messages=[
{
"role": "user",
"content": [
{
"type": "image_url",
"image_url": {"url": image_data_uri(image)},
},
{"type": "text", "text": prompt},
],
}
],
max_tokens=int(max_tokens),
temperature=float(temperature),
top_p=float(top_p),
top_k=int(top_k),
repeat_penalty=float(repeat_penalty),
def __init__(self, model_name='openbmb/MiniCPM-V-2_6', context_length=4096, temp=0.7,
top_p=0.8, top_k=100, repeat_penalty=1.05):
self.context_length = context_length
self.temp = temp
self.top_p = top_p
self.top_k = top_k
self.repeat_penalty = repeat_penalty
# Load model and tokenizer
print(f"Loading model: {model_name}...")
self.model = AutoModel.from_pretrained(model_name, trust_remote_code=True,
attn_implementation='sdpa', torch_dtype=torch.bfloat16).eval().cuda()
self.tokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True)
print("Model loaded successfully.")
def generate(self, image_path, prompt):
"""Generate response using the model"""
image = Image.open(image_path).convert('RGB')
msgs = [{'role': 'user', 'content': [image, prompt]}]
try:
response = self.model.chat(
image=None,
msgs=msgs,
tokenizer=self.tokenizer
)
results.append(llama_chat_content(response))
return batch_text(results)
return response
except Exception as e:
return f"Error generating response: {str(e)}"
class MiniCPMNode:
def __init__(self):
self.predictor = None
self.current_model = None
class MiniCPMNode(CachedModelNode):
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"prompt": (
"STRING",
{
"multiline": True,
"default": "Describe this image in detail.",
},
),
"model_variant": (list(GGUF_MODELS),),
"context_length": (
"INT",
{"default": 4096, "min": 512, "max": 131072},
),
"temperature": (
"FLOAT",
{"default": 0.2, "min": 0.0, "max": 2.0, "step": 0.1},
),
"top_p": (
"FLOAT",
{"default": 0.8, "min": 0.0, "max": 1.0, "step": 0.05},
),
"top_k": (
"INT",
{"default": 100, "min": 0, "max": 1000},
),
"repeat_penalty": (
"FLOAT",
{"default": 1.05, "min": 0.0, "max": 2.0, "step": 0.05},
),
},
"optional": {
"gpu_layers": (
"INT",
{"default": -1, "min": -1, "max": 1000},
),
"n_threads": (
"INT",
{
"default": default_llama_threads(),
"min": 1,
"max": 256,
},
),
"max_tokens": (
"INT",
{"default": 512, "min": 1, "max": 8192},
),
"unload_after": ("BOOLEAN", {"default": False}),
**llama_runtime_input_types(),
},
"image": ("IMAGE", {"tooltip": "Input image to be analyzed by MiniCPM-V"}),
"prompt": ("STRING", {
"multiline": True,
"default": "Describe this image in detail.",
"tooltip": "Instructions for the model. Be specific about what aspects of the image you want analyzed."
}),
"model_variant": (list(GGUF_MODELS.keys()), {
"tooltip": "Model size/quality tradeoff. Smaller models (Q2-Q4) are faster but less accurate. Larger models (Q8, F16) provide better quality but require more VRAM."
}),
"context_length": ("INT", {
"default": 4096,
"min": 512,
"max": 8192,
"tooltip": "Maximum length of text context. Larger values allow longer conversations but use more memory. Default 4096 works well for most cases."
}),
"temperature": ("FLOAT", {
"default": 0.7,
"min": 0.1,
"max": 2.0,
"step": 0.1,
"tooltip": "Controls randomness in generation. Lower values (0.1-0.5) are more focused and deterministic. Higher values (0.8-2.0) increase creativity and variance."
}),
"top_p": ("FLOAT", {
"default": 0.8,
"min": 0.1,
"max": 1.0,
"step": 0.1,
"tooltip": "Nucleus sampling threshold. Lower values make responses more focused. Higher values allow more diverse word choices."
}),
"top_k": ("INT", {
"default": 100,
"min": 1,
"max": 1000,
"tooltip": "Limits the number of tokens considered for each generation step. Lower values increase focus, higher values allow more variety."
}),
"repeat_penalty": ("FLOAT", {
"default": 1.05,
"min": 1.0,
"max": 2.0,
"step": 0.05,
"tooltip": "Penalizes word repetition. Values above 1.0 discourage repeated phrases. Higher values (>1.3) may affect fluency."
})
}
}
RETURN_TYPES = ("STRING",)
FUNCTION = "generate"
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
CATEGORY = "VLM Nodes/MiniCPM-V"
def generate(
self,
image,
prompt,
model_variant,
context_length=4096,
temperature=0.2,
top_p=0.8,
top_k=100,
repeat_penalty=1.05,
gpu_layers=-1,
n_threads=None,
max_tokens=512,
unload_after=False,
n_batch=512,
n_ubatch=512,
flash_attention="Auto",
use_mmap=True,
split_mode="Layer",
main_gpu=0,
tensor_split="",
):
n_threads = default_llama_threads() if n_threads is None else int(n_threads)
options = llama_runtime_options(
n_batch=n_batch,
n_ubatch=n_ubatch,
flash_attention=flash_attention,
use_mmap=use_mmap,
split_mode=split_mode,
main_gpu=main_gpu,
tensor_split=tensor_split,
)
key = (
model_variant,
int(context_length),
int(gpu_layers),
int(n_threads),
tuple(sorted(options.items())),
)
predictor = self.get_or_create_model(
key,
lambda: MiniCPMPredictor(
model_variant,
context_length,
gpu_layers,
n_threads,
options,
),
)
def download_model(self, model_filename):
"""Download model files from Huggingface"""
try:
return (
predictor.generate(
image,
prompt,
temperature,
top_p,
top_k,
repeat_penalty,
max_tokens,
),
print(f"Downloading model: {model_filename}...")
model_path = hf_hub_download(
repo_id="openbmb/MiniCPM-V-2_6-gguf",
filename=model_filename,
local_dir=MINICPM_PATH,
local_dir_use_symlinks=False
)
finally:
self.maybe_clear_model(unload_after)
print("Downloading mmproj model if not exists...")
mmproj_path = hf_hub_download(
repo_id="openbmb/MiniCPM-V-2_6-gguf",
filename="mmproj-model-f16.gguf",
local_dir=MINICPM_PATH,
local_dir_use_symlinks=False
)
print("Download complete.")
return Path(model_path), Path(mmproj_path)
except Exception as e:
raise RuntimeError(f"Error downloading model: {str(e)}")
def generate(self, image, prompt, model_variant, context_length=4096,
temperature=0.7, top_p=0.8, top_k=100, repeat_penalty=1.05):
# Get model filename from variant name
model_filename = GGUF_MODELS[model_variant]
# Initialize or update predictor if needed
if (self.predictor is None or
self.current_model != model_filename):
# Download model if needed
model_path, mmproj_path = self.download_model(model_filename)
# Initialize predictor
try:
self.predictor = MiniCPMPredictor(
model_name='openbmb/MiniCPM-V-2_6',
context_length=context_length,
temp=temperature,
top_p=top_p,
top_k=top_k,
repeat_penalty=repeat_penalty
)
self.current_model = model_filename
except Exception as e:
return (f"Error initializing model: {str(e)}",)
# Save input image temporarily
temp_image = MINICPM_PATH / "temp_input.png"
Image.fromarray(np.uint8(image[0] * 255)).save(temp_image)
try:
# Generate response
response = self.predictor.generate(temp_image, prompt)
# Clean up
temp_image.unlink(missing_ok=True)
return (response,)
except Exception as e:
return (f"Error during generation: {str(e)}",)
NODE_CLASS_MAPPINGS = {"MiniCPMNode": MiniCPMNode}
NODE_DISPLAY_NAME_MAPPINGS = {"MiniCPMNode": "MiniCPM-V 2.6 (GGUF)"}
# Register the node
NODE_CLASS_MAPPINGS = {
"MiniCPMNode": MiniCPMNode
}
NODE_DISPLAY_NAME_MAPPINGS = {
"MiniCPMNode": "MiniCPM-V Model"
}
-478
View File
@@ -1,478 +0,0 @@
"""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",
]
-826
View File
@@ -1,826 +0,0 @@
"""Modern, chat-template based vision-language models.
This node intentionally uses the Transformers multimodal auto classes instead
of model-specific glue. It provides one stable ComfyUI surface for current
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
import torch
from .runtime import (
CachedModelNode,
ExternalTorchModel,
ManagedTorchModel,
accelerator_backend,
batch_text,
execution_device,
external_device_map,
inference_context,
model_device,
move_inputs,
normalize_hf_model_id,
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)
class ModelSpec:
repo_id: str
family: str
estimated_gib: float
gated: bool = False
video: bool = False
small_fast: bool = False
trust_remote_code: bool = False
# Deliberately curated: these are useful tiers, not every redundant checkpoint.
MODEL_CATALOG = {
"Qwen 3.5 0.8B (fastest current)": ModelSpec(
"Qwen/Qwen3.5-0.8B",
"Qwen 3.5",
2.0,
video=True,
small_fast=True,
),
"Qwen 3.5 2B": ModelSpec(
"Qwen/Qwen3.5-2B",
"Qwen 3.5",
4.5,
video=True,
small_fast=True,
),
"Qwen 3.5 4B (recommended)": ModelSpec(
"Qwen/Qwen3.5-4B",
"Qwen 3.5",
8.5,
video=True,
small_fast=True,
),
"Qwen 3.5 9B": ModelSpec(
"Qwen/Qwen3.5-9B", "Qwen 3.5", 19.0, video=True
),
"Qwen 3.5 27B (4-bit recommended)": ModelSpec(
"Qwen/Qwen3.5-27B", "Qwen 3.5", 55.0, video=True
),
"Qwen 3.5 35B-A3B (4-bit recommended)": ModelSpec(
"Qwen/Qwen3.5-35B-A3B", "Qwen 3.5", 72.0, video=True
),
"Qwen 3.6 27B (4-bit recommended)": ModelSpec(
"Qwen/Qwen3.6-27B", "Qwen 3.6", 55.0, video=True
),
"Qwen 3 VL 2B Instruct": ModelSpec(
"Qwen/Qwen3-VL-2B-Instruct",
"Qwen 3 VL",
5.0,
video=True,
small_fast=True,
),
"Qwen 3 VL 4B Instruct": ModelSpec(
"Qwen/Qwen3-VL-4B-Instruct",
"Qwen 3 VL",
9.0,
video=True,
small_fast=True,
),
"Qwen 3 VL 8B Instruct": ModelSpec(
"Qwen/Qwen3-VL-8B-Instruct", "Qwen 3 VL", 18.0, video=True
),
"Qwen 3 VL 30B-A3B Instruct (4-bit recommended)": ModelSpec(
"Qwen/Qwen3-VL-30B-A3B-Instruct", "Qwen 3 VL", 61.0, video=True
),
"Qwen 2.5 VL 3B Instruct (legacy workflows)": ModelSpec(
"Qwen/Qwen2.5-VL-3B-Instruct",
"Qwen 2.5 VL",
7.0,
video=True,
small_fast=True,
),
"Qwen 2.5 VL 7B Instruct (legacy workflows)": ModelSpec(
"Qwen/Qwen2.5-VL-7B-Instruct", "Qwen 2.5 VL", 16.0, video=True
),
"Gemma 3 4B IT (license acceptance required)": ModelSpec(
"google/gemma-3-4b-it",
"Gemma 3",
9.0,
gated=True,
small_fast=True,
),
"Gemma 3 12B IT (license acceptance required)": ModelSpec(
"google/gemma-3-12b-it", "Gemma 3", 25.0, gated=True
),
"Gemma 3 27B IT (4-bit recommended, gated)": ModelSpec(
"google/gemma-3-27b-it", "Gemma 3", 55.0, gated=True
),
"SmolVLM2 256M Video (smallest)": ModelSpec(
"HuggingFaceTB/SmolVLM2-256M-Video-Instruct",
"SmolVLM2",
1.4,
video=True,
small_fast=True,
),
"SmolVLM2 500M Video (low VRAM)": ModelSpec(
"HuggingFaceTB/SmolVLM2-500M-Video-Instruct",
"SmolVLM2",
1.8,
video=True,
small_fast=True,
),
"SmolVLM2 2.2B Video": ModelSpec(
"HuggingFaceTB/SmolVLM2-2.2B-Instruct",
"SmolVLM2",
5.2,
video=True,
small_fast=True,
),
"LFM2.5 VL 450M (edge)": ModelSpec(
"LiquidAI/LFM2.5-VL-450M",
"LFM2.5 VL",
1.5,
small_fast=True,
),
"LFM2.5 VL 1.6B": ModelSpec(
"LiquidAI/LFM2.5-VL-1.6B",
"LFM2.5 VL",
4.0,
small_fast=True,
),
"InternVL 3.5 1B HF": ModelSpec(
"OpenGVLab/InternVL3_5-1B-HF",
"InternVL 3.5",
2.5,
video=True,
small_fast=True,
),
"InternVL 3.5 2B HF": ModelSpec(
"OpenGVLab/InternVL3_5-2B-HF",
"InternVL 3.5",
5.0,
video=True,
small_fast=True,
),
"Granite Vision 3.3 2B (documents/OCR)": ModelSpec(
"ibm-granite/granite-vision-3.3-2b",
"Granite Vision 3.3",
6.5,
small_fast=True,
),
"Granite Vision 4.1 4B (structured documents)": ModelSpec(
"ibm-granite/granite-vision-4.1-4b",
"Granite Vision 4.1",
9.0,
small_fast=True,
),
"Custom Hugging Face model": ModelSpec(
"",
"Custom",
8.0,
trust_remote_code=True,
),
}
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)",
"8-bit (bitsandbytes)",
"CPU",
)
ATTENTION_MODES = ("Auto (SDPA)", "Flash Attention 2", "Eager")
def _progress_text_sender(node_id: str | None) -> Callable[[str], None] | None:
"""Return a best-effort sender for ComfyUI's native progress-text channel."""
if node_id is None:
return None
try:
from server import PromptServer
server = PromptServer.instance
except (ImportError, AttributeError):
return None
def send(text: str) -> None:
try:
server.send_progress_text(
text,
str(node_id),
server.client_id,
)
except Exception:
# Streaming is a UI enhancement and must never fail inference.
return
return send
def _model_class(transformers):
for name in ("AutoModelForImageTextToText", "AutoModelForMultimodalLM"):
model_class = getattr(transformers, name, None)
if model_class is not None:
return model_class
raise RuntimeError(
"Modern VLMs require a current Transformers release with "
"AutoModelForImageTextToText support."
)
class ModernVLMPredictor:
def __init__(
self,
model_label: str,
custom_model_id: str,
memory_mode: str,
attention_mode: str,
) -> None:
transformers = require_module("transformers")
self.streamer_class = getattr(transformers, "TextIteratorStreamer", None)
spec = MODEL_CATALOG[model_label]
repo_id = (
normalize_hf_model_id(custom_model_id)
if spec.family == "Custom"
else spec.repo_id
)
self.spec = spec
self.dtype = torch_dtype("bfloat16")
if (
attention_mode == "Flash Attention 2"
and accelerator_backend(execution_device())
not in {"nvidia-cuda", "amd-rocm"}
):
raise RuntimeError(
"Flash Attention 2 requires a supported CUDA or ROCm build. "
"Select Auto (SDPA) on Apple Metal, Intel XPU, or CPU."
)
quantization_device = None
if memory_mode in {
"4-bit NF4 (bitsandbytes)",
"8-bit (bitsandbytes)",
}:
# Validate before downloading a multi-gigabyte checkpoint.
quantization_device = require_quantization_backend(memory_mode)
try:
model_path = snapshot_download(
repo_id,
f"modern-vlm/{repo_id.replace('/', '--')}",
ignore_patterns=["*.bin", "*.msgpack", "*.h5", "*.onnx"],
)
except Exception as exc:
if spec.gated:
raise RuntimeError(
f"{repo_id} is gated. Accept its Hugging Face license and "
"set HF_TOKEN before running this node."
) from exc
raise
self.processor = transformers.AutoProcessor.from_pretrained(
model_path,
trust_remote_code=spec.trust_remote_code,
)
attention = {
# Let each architecture choose its maintained native kernel. Most
# current PyTorch models select SDPA here, while hybrid edge models
# can retain their own attention implementation.
"Auto (SDPA)": None,
"Flash Attention 2": "flash_attention_2",
"Eager": "eager",
}[attention_mode]
kwargs: dict[str, Any] = {
"dtype": self.dtype,
"trust_remote_code": spec.trust_remote_code,
}
if attention is not None:
kwargs["attn_implementation"] = attention
external = memory_mode != "ComfyUI managed (BF16)"
if memory_mode in {"4-bit NF4 (bitsandbytes)", "8-bit (bitsandbytes)"}:
assert quantization_device is not None
needs_offload = spec.estimated_gib >= 40.0
kwargs["quantization_config"] = transformers.BitsAndBytesConfig(
load_in_4bit=memory_mode.startswith("4-bit"),
load_in_8bit=memory_mode.startswith("8-bit"),
bnb_4bit_compute_dtype=self.dtype,
bnb_4bit_quant_type="nf4",
bnb_4bit_use_double_quant=True,
llm_int8_enable_fp32_cpu_offload=needs_offload,
)
if needs_offload:
# Automatic CPU/disk placement is maintained for CUDA, ROCm,
# and XPU. MPS uses unified memory and CPU already runs in RAM,
# so both stay on their explicit active device.
kwargs["device_map"] = external_device_map(
allow_auto_offload=True
)
kwargs["offload_folder"] = str(model_path / ".offload")
else:
# Avoid accidental dispatch to device zero when ComfyUI chose
# another GPU, Apple Metal, Intel XPU, or CPU.
kwargs["device_map"] = external_device_map()
divisor = 4 if memory_mode.startswith("4-bit") else 2
if quantization_device.type != "cpu":
reserve_external_vram(
int(spec.estimated_gib * 1024**3 / divisor)
)
elif memory_mode == "CPU":
kwargs["dtype"] = torch.float32
try:
model = _model_class(transformers).from_pretrained(
model_path, **kwargs
).eval()
except OSError as exc:
if spec.gated:
raise RuntimeError(
f"{repo_id} is gated. Accept its Hugging Face license and "
"set HF_TOKEN before running this node."
) from exc
raise
except ImportError as exc:
if attention_mode == "Flash Attention 2":
raise RuntimeError(
"Flash Attention 2 is unavailable for this Python/PyTorch "
"build. Select Auto (SDPA), or install a matching wheel."
) from exc
raise
self.handle = (
ExternalTorchModel(model, processor=self.processor)
if external
else ManagedTorchModel(model, processor=self.processor)
)
def close(self) -> None:
self.handle.close()
self.processor = None
def _inputs(
self,
messages,
enable_thinking: bool = False,
*,
video_metadata: dict[str, Any] | None = None,
):
"""Use the standard multimodal template, with an older-template fallback."""
template_kwargs = (
{"enable_thinking": bool(enable_thinking)}
if self.spec.family in {"Qwen 3.5", "Qwen 3.6"}
else {}
)
processor_kwargs = (
{
"video_metadata": [[video_metadata]],
# ComfyUI already supplied the selected frames as a batch.
"do_sample_frames": False,
}
if video_metadata is not None
else None
)
if processor_kwargs is not None and self.spec.family == "InternVL 3.5":
# The published InternVL 3.5 video preprocessor uses 384px, which
# makes a 27x27 patch grid with its 14px vision patches. The
# model's 0.5 pixel shuffle requires even spatial dimensions.
image_size = getattr(self.processor.image_processor, "size", None)
processor_kwargs["size"] = (
dict(image_size)
if image_size is not None
else {"height": 448, "width": 448}
)
try:
return self.processor.apply_chat_template(
messages,
add_generation_prompt=True,
tokenize=True,
return_dict=True,
return_tensors="pt",
processor_kwargs=processor_kwargs,
**template_kwargs,
)
except (TypeError, ValueError, KeyError):
media = []
portable_messages = []
for message in messages:
content = []
for part in message["content"]:
if part["type"] == "image":
media.append(part["image"])
content.append({"type": "image"})
elif part["type"] == "video":
media.extend(part["video"])
content.extend({"type": "image"} for _ in part["video"])
else:
content.append(part)
portable_messages.append(
{"role": message["role"], "content": content}
)
prompt = self.processor.apply_chat_template(
portable_messages,
add_generation_prompt=True,
tokenize=False,
**template_kwargs,
)
return self.processor(
text=[prompt], images=media, return_tensors="pt"
)
def generate(
self,
images,
prompt: str,
system_prompt: str,
max_new_tokens: int,
temperature: float,
top_p: float,
video_frames=None,
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 []
)
video = (
tensor_batch_to_pil(video_frames)
if video_frames is not None
else None
)
if video is None and not primary_images:
raise ValueError("Connect either image or video_frames.")
if video is not None and not self.spec.video:
raise ValueError(
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
# required still image as well makes small video models attend to the
# still and silently ignore the frames.
runs = [None] if video is not None else primary_images
for image in runs:
messages = []
if system_prompt.strip():
messages.append(
{
"role": "system",
"content": [
{"type": "text", "text": system_prompt.strip()}
],
}
)
content = (
[{"type": "video", "video": video}]
if video is not None
else [{"type": "image", "image": image}]
)
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:
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,
video_metadata=metadata,
)
model = self.handle.ensure_loaded()
device = model_device(model)
inputs = move_inputs(inputs, device)
input_length = inputs["input_ids"].shape[-1]
generation: dict[str, Any] = {
"max_new_tokens": int(max_new_tokens),
"do_sample": float(temperature) > 0,
}
if generation["do_sample"]:
generation.update(
temperature=float(temperature), top_p=float(top_p)
)
streamer_class = self.streamer_class
tokenizer = getattr(self.processor, "tokenizer", self.processor)
if stream_callback is not None and streamer_class is not None:
streamer = streamer_class(
tokenizer,
skip_prompt=True,
skip_special_tokens=True,
clean_up_tokenization_spaces=False,
)
generated = []
errors: list[BaseException] = []
def generate_in_background() -> None:
try:
with (
torch.inference_mode(),
inference_context(device, self.dtype),
):
generated.append(
model.generate(
**inputs,
**generation,
streamer=streamer,
)
)
except BaseException as exc:
errors.append(exc)
# Unblock TextIteratorStreamer if generation exits
# before it can publish its normal stop signal.
streamer.end()
worker = threading.Thread(
target=generate_in_background,
name="ComfyUI-VLM-token-stream",
daemon=True,
)
worker.start()
chunks = []
for chunk in streamer:
chunks.append(chunk)
current = batch_text(
[*results, "".join(chunks).strip()]
)
if current:
stream_callback(current)
worker.join()
if errors:
raise errors[0]
decoded = "".join(chunks).strip()
if not decoded and generated:
new_tokens = generated[0][:, input_length:]
decoded = self.processor.batch_decode(
new_tokens,
skip_special_tokens=True,
clean_up_tokenization_spaces=False,
)[0].strip()
results.append(decoded)
else:
with (
torch.inference_mode(),
inference_context(device, self.dtype),
):
output = model.generate(**inputs, **generation)
new_tokens = output[:, input_length:]
results.append(
self.processor.batch_decode(
new_tokens,
skip_special_tokens=True,
clean_up_tokenization_spaces=False,
)[0].strip()
)
return batch_text(results)
class ModernVLM(CachedModelNode):
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"prompt": (
"STRING",
{
"multiline": True,
"default": "Describe this image precisely and in detail.",
},
),
"model": (
list(RECOMMENDED_MODEL_LABELS),
{"default": "Qwen 3 VL 2B Instruct"},
),
"custom_model_id": ("STRING", {"default": ""}),
"memory_mode": (
MEMORY_MODES,
{"default": "ComfyUI managed (BF16)"},
),
"max_new_tokens": (
"INT",
{"default": 512, "min": 1, "max": 16384},
),
"temperature": (
"FLOAT",
{"default": 0.1, "min": 0.0, "max": 2.0, "step": 0.05},
),
"top_p": (
"FLOAT",
{"default": 0.9, "min": 0.01, "max": 1.0, "step": 0.01},
),
},
"optional": {
"image": ("IMAGE",),
"system_prompt": (
"STRING",
{
"multiline": True,
"default": "You are an expert visual analyst.",
},
),
"video_frames": ("IMAGE",),
"video_selection": (VLM_VIDEO_SELECTION,),
"fps": (
"FLOAT",
{"default": 1.0, "min": 0.1, "max": 60.0, "step": 0.1},
),
"attention_mode": (
ATTENTION_MODES,
{"default": "Auto (SDPA)"},
),
"enable_thinking": ("BOOLEAN", {"default": False}),
"unload_after": ("BOOLEAN", {"default": False}),
"stream_output": (
"BOOLEAN",
{
"default": True,
"tooltip": (
"Stream generated text through ComfyUI's native "
"progress-text WebSocket while inference runs."
),
},
),
},
"hidden": {"unique_id": "UNIQUE_ID"},
}
RETURN_TYPES = ("STRING",)
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,
model,
custom_model_id,
memory_mode,
max_new_tokens,
temperature,
top_p,
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,
unload_after=False,
stream_output=True,
unique_id=None,
):
stream_callback = (
_progress_text_sender(unique_id) if stream_output else None
)
if stream_callback is not None:
stream_callback("Preparing model…")
effective_custom_id = (
normalize_hf_model_id(custom_model_id)
if model == "Custom Hugging Face model"
else ""
)
key = (model, effective_custom_id, memory_mode, attention_mode)
predictor = self.get_or_create_model(
key,
lambda: ModernVLMPredictor(
model, effective_custom_id, memory_mode, attention_mode
),
)
try:
return (
predictor.generate(
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)
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",
}
+317 -168
View File
@@ -1,196 +1,345 @@
"""AllenAI Molmo nodes with batch support and deterministic model ownership."""
from __future__ import annotations
from typing import Any
import torch
import os
from PIL import Image
from pathlib import Path
import folder_paths
import logging
import warnings
from transformers import AutoModelForCausalLM, AutoProcessor, GenerationConfig, BitsAndBytesConfig
from huggingface_hub import snapshot_download
import torch.amp.autocast_mode
import psutil
from .runtime import (
CachedModelNode,
ExternalTorchModel,
ManagedTorchModel,
batch_text,
external_device_map,
inference_context,
model_device,
require_module,
require_quantization_backend,
reserve_external_vram,
snapshot_download,
tensor_batch_to_pil,
torch_dtype,
)
# Configure logging
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger('MolmoNode')
# Filter specific warnings
warnings.filterwarnings('ignore', message='.*The model weights are not tied.*')
warnings.filterwarnings('ignore', message='.*You should use.*max_memory.*')
# Define the directory for saving Molmo files
MOLMO_PATH = Path(folder_paths.folder_names_and_paths["LLavacheckpoints"][0][0]) / "files_for_molmo"
MOLMO_PATH.mkdir(parents=True, exist_ok=True)
# Memory configurations with detailed descriptions
MEMORY_MODES = {
"Full Precision (45GB+ Required)": "managed",
"8-bit Quantized (25GB+ Required)": "8bit",
"4-bit Quantized (15GB+ Required)": "4bit",
"4-bit + CPU Offload (12GB+ Required)": "4bit-offload",
}
MOLMO_MODELS = {
"MolmoE-1B (Efficient)": "allenai/MolmoE-1B-0924",
"Molmo-7B-D (Best 7B)": "allenai/Molmo-7B-D-0924",
"Molmo-7B-O (Alternative 7B)": "allenai/Molmo-7B-O-0924",
"Full Precision (45GB+ Required)": {
"description": "Uses full FP16 precision. Requires ~45GB total system RAM, including 24GB+ VRAM.",
"load_in_8bit": False,
"load_in_4bit": False,
"double_quant": False,
"cpu_offload": False
},
"8-bit Quantized (25GB+ Required)": {
"description": "Uses 8-bit quantization. Requires ~25GB total system RAM. Good balance of quality and memory usage.",
"load_in_8bit": True,
"load_in_4bit": False,
"double_quant": False,
"cpu_offload": False
},
"4-bit Quantized (15GB+ Required)": {
"description": "Uses 4-bit quantization. Requires ~15GB total system RAM. Lowest memory usage, slight quality impact.",
"load_in_8bit": False,
"load_in_4bit": True,
"double_quant": True,
"cpu_offload": False
},
"4-bit + CPU Offload (12GB+ Required)": {
"description": "Uses 4-bit quantization with CPU offloading. Slowest but lowest VRAM usage (~12GB).",
"load_in_8bit": False,
"load_in_4bit": True,
"double_quant": True,
"cpu_offload": True
}
}
# Available Molmo models
MOLMO_MODELS = {
"MolmoE-1B (Efficient)": {
"repo": "allenai/MolmoE-1B-0924",
"description": "Mixture-of-Experts model, smallest option (still requires significant RAM)"
},
"Molmo-7B-D (Best 7B)": {
"repo": "allenai/Molmo-7B-D-0924",
"description": "⚠️ Very large model, requires more RAM than MolmoE-1B"
},
"Molmo-7B-O (Alternative 7B)": {
"repo": "allenai/Molmo-7B-O-0924",
"description": "⚠️ Very large model, requires more RAM than MolmoE-1B"
}
}
class SystemResources:
@staticmethod
def get_system_memory():
return psutil.virtual_memory().total / (1024 ** 3) # GB
@staticmethod
def get_available_vram():
if not torch.cuda.is_available():
return 0
return torch.cuda.get_device_properties(0).total_memory / (1024 ** 3) # GB
@staticmethod
def check_memory_requirements(memory_mode):
config = MEMORY_MODES[memory_mode]
required_ram = 15 if config["load_in_4bit"] else (25 if config["load_in_8bit"] else 45)
available_ram = SystemResources.get_system_memory()
available_vram = SystemResources.get_available_vram()
warnings = []
if available_ram < required_ram:
warnings.append(f"WARNING: This memory mode requires {required_ram}GB total RAM, but only {available_ram:.1f}GB available")
min_vram = 12 if config["cpu_offload"] else 24
if available_vram < min_vram:
warnings.append(f"WARNING: Recommended minimum {min_vram}GB VRAM, but only {available_vram:.1f}GB available")
return warnings
class MolmoPredictor:
def __init__(self, model_name, memory_mode, use_autocast):
transformers = require_module("transformers")
repo_id = MOLMO_MODELS[model_name]
mode = MEMORY_MODES[memory_mode]
external = mode != "managed"
if external:
# Validate before downloading a multi-gigabyte checkpoint.
require_quantization_backend(memory_mode)
path = snapshot_download(
repo_id,
f"molmo/{repo_id.replace('/', '--')}",
ignore_patterns=["*.bin"],
def __init__(self, model_name, memory_mode="4-bit Quantized (15GB+ Required)", use_autocast=True):
self.model_name = MOLMO_MODELS[model_name]["repo"]
self.memory_config = MEMORY_MODES[memory_mode]
self.use_autocast = use_autocast and torch.cuda.is_available()
# Check system resources
warnings = SystemResources.check_memory_requirements(memory_mode)
for warning in warnings:
logger.warning(warning)
# Download model if needed
logger.info(f"Downloading/loading {model_name} in {memory_mode} mode...")
self.model_path = snapshot_download(
self.model_name,
local_dir=MOLMO_PATH / model_name,
local_dir_use_symlinks="auto"
)
self.dtype = torch_dtype("bfloat16")
self.use_autocast = bool(use_autocast)
self.processor = transformers.AutoProcessor.from_pretrained(
path, trust_remote_code=True
)
kwargs: dict[str, Any] = {
"trust_remote_code": True,
"dtype": self.dtype,
}
if external:
kwargs["quantization_config"] = transformers.BitsAndBytesConfig(
load_in_8bit=mode == "8bit",
load_in_4bit=mode.startswith("4bit"),
bnb_4bit_compute_dtype=self.dtype,
bnb_4bit_quant_type="nf4",
bnb_4bit_use_double_quant=True,
try:
# Configure quantization
compute_dtype = torch.float16 if torch.cuda.is_available() else torch.float32
quant_config = None
if self.memory_config["load_in_4bit"] or self.memory_config["load_in_8bit"]:
quant_config = BitsAndBytesConfig(
load_in_8bit=self.memory_config["load_in_8bit"],
load_in_4bit=self.memory_config["load_in_4bit"],
bnb_4bit_compute_dtype=compute_dtype,
bnb_4bit_use_double_quant=self.memory_config["double_quant"],
bnb_4bit_quant_type="nf4" # More accurate than fp4
)
# Load processor
self.processor = AutoProcessor.from_pretrained(
self.model_path,
trust_remote_code=True
)
kwargs["device_map"] = external_device_map(
allow_auto_offload=mode == "4bit-offload"
# Load model with optimizations
device_map = "auto" if self.memory_config["cpu_offload"] else None
self.model = AutoModelForCausalLM.from_pretrained(
self.model_path,
trust_remote_code=True,
quantization_config=quant_config,
device_map=device_map,
torch_dtype=compute_dtype
)
reserve_external_vram(
(5 if "1B" in model_name else 12) * 1024**3
logger.info(f"Successfully loaded {model_name}")
except RuntimeError as e:
if "out of memory" in str(e):
torch.cuda.empty_cache()
raise RuntimeError(
f"Out of memory while loading model. Current mode: {memory_mode}\n"
"Try:\n"
"1. Using a more aggressive memory saving mode\n"
"2. Closing other applications\n"
"3. Restarting ComfyUI"
) from e
raise
def generate(self, image, prompt, max_new_tokens=200, temperature=0.7, top_p=0.9, top_k=50):
try:
# Process inputs
inputs = self.processor.process(
images=[image],
text=prompt
)
model = transformers.AutoModelForCausalLM.from_pretrained(
path, **kwargs
).eval()
self.handle = (
ExternalTorchModel(model, processor=self.processor)
if external
else ManagedTorchModel(model, processor=self.processor)
)
def close(self):
self.handle.close()
self.processor = None
def generate(self, image, prompt, max_new_tokens, temperature, top_p, top_k):
model = self.handle.ensure_loaded()
device = model_device(model)
inputs = self.processor.process(images=[image], text=prompt)
inputs = {
key: value.to(device).unsqueeze(0)
for key, value in inputs.items()
}
config = require_module("transformers").GenerationConfig(
max_new_tokens=int(max_new_tokens),
do_sample=float(temperature) > 0,
temperature=max(float(temperature), 1e-5),
top_p=float(top_p),
top_k=int(top_k),
stop_strings="<|endoftext|>",
pad_token_id=self.processor.tokenizer.pad_token_id,
eos_token_id=self.processor.tokenizer.eos_token_id,
)
context = (
inference_context(device, self.dtype)
if self.use_autocast
else torch.no_grad()
)
with torch.inference_mode(), context:
output = model.generate_from_batch(
inputs, config, tokenizer=self.processor.tokenizer
# Move inputs to device and create batch
device = next(self.model.parameters()).device
inputs = {k: v.to(device).unsqueeze(0) for k, v in inputs.items()}
# Configure generation
generation_config = GenerationConfig(
max_new_tokens=max_new_tokens,
do_sample=True,
temperature=temperature,
top_p=top_p,
top_k=top_k,
stop_strings="<|endoftext|>",
pad_token_id=self.processor.tokenizer.pad_token_id,
eos_token_id=self.processor.tokenizer.eos_token_id
)
return self.processor.tokenizer.decode(
output[0, inputs["input_ids"].shape[1]:],
skip_special_tokens=True,
).strip()
# Generate with autocast if enabled
if self.use_autocast:
with torch.amp.autocast(device_type="cuda", dtype=torch.bfloat16):
output = self.model.generate_from_batch(
inputs,
generation_config,
tokenizer=self.processor.tokenizer
)
else:
output = self.model.generate_from_batch(
inputs,
generation_config,
tokenizer=self.processor.tokenizer
)
# Get input size before cleanup
input_size = inputs['input_ids'].size(1)
# Clean up
del inputs
torch.cuda.empty_cache()
# Extract and decode generated tokens using saved size
generated_tokens = output[0, input_size:]
return self.processor.tokenizer.decode(generated_tokens, skip_special_tokens=True)
except RuntimeError as e:
if "out of memory" in str(e):
torch.cuda.empty_cache()
raise RuntimeError(
"Out of memory during generation. Try:\n"
"1. Using a more aggressive memory saving mode\n"
"2. Reducing max_new_tokens\n"
"3. Clearing ComfyUI cache\n"
"4. Restarting ComfyUI"
) from e
raise
class MolmoNode:
def __init__(self):
self.predictor = None
self.current_model = None
self.current_memory_mode = None
self.current_autocast = None
class MolmoNode(CachedModelNode):
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"prompt": (
"STRING",
{"multiline": True, "default": "Describe this image in detail."},
),
"model_name": (list(MOLMO_MODELS),),
"memory_mode": (
list(MEMORY_MODES),
{"default": "4-bit Quantized (15GB+ Required)"},
),
"max_new_tokens": (
"INT",
{"default": 200, "min": 1, "max": 2048},
),
"temperature": (
"FLOAT",
{"default": 0.2, "min": 0.0, "max": 2.0, "step": 0.1},
),
"top_p": (
"FLOAT",
{"default": 0.9, "min": 0.01, "max": 1.0, "step": 0.01},
),
"top_k": ("INT", {"default": 50, "min": 1, "max": 100}),
"use_autocast": ("BOOLEAN", {"default": True}),
},
"optional": {
"unload_after": ("BOOLEAN", {"default": False}),
},
"image": ("IMAGE", {
"tooltip": "Input image to be analyzed by Molmo"
}),
"prompt": ("STRING", {
"multiline": True,
"default": "Describe this image in detail.",
"tooltip": "Instructions for the model. Be specific about what aspects of the image you want analyzed."
}),
"model_name": (list(MOLMO_MODELS.keys()), {
"tooltip": "⚠️ WARNING: These are very large models requiring significant RAM/VRAM. Start with MolmoE-1B."
}),
"memory_mode": (list(MEMORY_MODES.keys()), {
"default": "4-bit Quantized (15GB+ Required)",
"tooltip": "Controls RAM/VRAM usage. Use most aggressive option that works on your system."
}),
"max_new_tokens": ("INT", {
"default": 200,
"min": 1,
"max": 2048,
"tooltip": "Maximum tokens to generate. Higher values need more VRAM. Start small (200) and increase if needed."
}),
"temperature": ("FLOAT", {
"default": 0.7,
"min": 0.1,
"max": 2.0,
"step": 0.1,
"tooltip": "Controls randomness. Lower (0.1-0.5) = more focused, higher (0.8-2.0) = more creative."
}),
"top_p": ("FLOAT", {
"default": 0.9,
"min": 0.1,
"max": 1.0,
"step": 0.1,
"tooltip": "Nucleus sampling. Lower = more focused on likely tokens, higher = more diverse vocabulary."
}),
"top_k": ("INT", {
"default": 50,
"min": 1,
"max": 100,
"tooltip": "Limits token choices to top K most likely. Lower = more focused, higher = more variety."
}),
"use_autocast": ("BOOLEAN", {
"default": True,
"tooltip": "Enables mixed precision. Keeps quality while reducing VRAM usage. Recommended ON."
})
}
}
RETURN_TYPES = ("STRING",)
FUNCTION = "generate"
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
CATEGORY = "VLM Nodes/Molmo"
def generate(
self,
image,
prompt,
model_name,
memory_mode="4-bit Quantized (15GB+ Required)",
max_new_tokens=200,
temperature=0.2,
top_p=0.9,
top_k=50,
use_autocast=True,
unload_after=False,
):
predictor = self.get_or_create_model(
(model_name, memory_mode, bool(use_autocast)),
lambda: MolmoPredictor(model_name, memory_mode, use_autocast),
)
def generate(self, image, prompt, model_name, memory_mode="4-bit Quantized (15GB+ Required)",
max_new_tokens=200, temperature=0.7, top_p=0.9, top_k=50, use_autocast=True):
try:
return (
batch_text(
predictor.generate(
pil,
prompt,
max_new_tokens,
temperature,
top_p,
top_k,
)
for pil in tensor_batch_to_pil(image)
),
# Initialize or update predictor if needed
if (self.predictor is None or
self.current_model != model_name or
self.current_memory_mode != memory_mode or
self.current_autocast != use_autocast):
# Clean up old model if it exists
if self.predictor is not None:
del self.predictor.model
del self.predictor.processor
torch.cuda.empty_cache()
self.predictor = MolmoPredictor(
model_name,
memory_mode=memory_mode,
use_autocast=use_autocast
)
self.current_model = model_name
self.current_memory_mode = memory_mode
self.current_autocast = use_autocast
# Convert tensor to PIL Image
pil_image = Image.fromarray((image[0] * 255).numpy().astype('uint8'))
# Generate response
response = self.predictor.generate(
pil_image,
prompt,
max_new_tokens=max_new_tokens,
temperature=temperature,
top_p=top_p,
top_k=top_k
)
finally:
self.maybe_clear_model(unload_after)
return (response,)
except Exception as e:
# Clean up on error
if hasattr(self, 'predictor') and self.predictor is not None:
del self.predictor.model
del self.predictor.processor
self.predictor = None
torch.cuda.empty_cache()
return (f"Error: {str(e)}",)
# Register the node
NODE_CLASS_MAPPINGS = {
"MolmoNode": MolmoNode
}
NODE_CLASS_MAPPINGS = {"MolmoNode": MolmoNode}
NODE_DISPLAY_NAME_MAPPINGS = {"MolmoNode": "Molmo Vision-Language Model"}
NODE_DISPLAY_NAME_MAPPINGS = {
"MolmoNode": "Molmo Vision-Language Model"
}
+2
View File
@@ -0,0 +1,2 @@
from .vision_encoder import VisionEncoder
from .text_model import TextModel
+66
View File
@@ -0,0 +1,66 @@
# Copyright (c) Microsoft Corporation.
# Licensed under the MIT license.
import math
from typing import Optional
from transformers import PretrainedConfig
class PhiConfig(PretrainedConfig):
"""Phi configuration."""
model_type = "phi-msft"
attribute_map = {
"max_position_embeddings": "n_positions",
"hidden_size": "n_embd",
"num_attention_heads": "n_head",
"num_hidden_layers": "n_layer",
}
def __init__(
self,
vocab_size: int = 50304,
n_positions: int = 2048,
n_embd: int = 1024,
n_layer: int = 20,
n_inner: Optional[int] = None,
n_head: int = 16,
n_head_kv: Optional[int] = None,
rotary_dim: Optional[int] = 32,
activation_function: Optional[str] = "gelu_new",
flash_attn: bool = False,
flash_rotary: bool = False,
fused_dense: bool = False,
attn_pdrop: float = 0.0,
embd_pdrop: float = 0.0,
resid_pdrop: float = 0.0,
layer_norm_epsilon: float = 1e-5,
initializer_range: float = 0.02,
tie_word_embeddings: bool = False,
pad_vocab_size_multiple: int = 64,
gradient_checkpointing: bool = False,
**kwargs
) -> None:
self.vocab_size = int(
math.ceil(vocab_size / pad_vocab_size_multiple) * pad_vocab_size_multiple
)
self.n_positions = n_positions
self.n_embd = n_embd
self.n_layer = n_layer
self.n_inner = n_inner
self.n_head = n_head
self.n_head_kv = n_head_kv
self.rotary_dim = min(rotary_dim, n_embd // n_head)
self.activation_function = activation_function
self.flash_attn = flash_attn
self.flash_rotary = flash_rotary
self.fused_dense = fused_dense
self.attn_pdrop = attn_pdrop
self.embd_pdrop = embd_pdrop
self.resid_pdrop = resid_pdrop
self.layer_norm_epsilon = layer_norm_epsilon
self.initializer_range = initializer_range
self.gradient_checkpointing = gradient_checkpointing
super().__init__(tie_word_embeddings=tie_word_embeddings, **kwargs)
File diff suppressed because it is too large Load Diff
+86
View File
@@ -0,0 +1,86 @@
import torch
import transformers
from transformers import CodeGenTokenizerFast as Tokenizer
from accelerate import init_empty_weights, load_checkpoint_and_dispatch
from .phi.configuration_phi import PhiConfig
from .phi.modeling_phi import PhiForCausalLM
import re
transformers.logging.set_verbosity_error()
class TextModel:
def __init__(self, model_path: str = "model") -> None:
super().__init__()
self.tokenizer = Tokenizer.from_pretrained(f"{model_path}/tokenizer")
phi_config = PhiConfig.from_pretrained(f"{model_path}/text_model_cfg.json")
with init_empty_weights():
self.model = PhiForCausalLM(phi_config)
self.model = load_checkpoint_and_dispatch(
self.model,
f"{model_path}/text_model.pt",
device_map="auto",
)
self.text_emb = self.model.get_input_embeddings()
def input_embeds(self, prompt, image_embeds):
embeds = []
def _add_toks(toks):
embeds.append(self.text_emb(toks))
def _tokenize(txt):
return self.tokenizer(
txt, return_tensors="pt", add_special_tokens=False
).input_ids.to(self.model.device)
# Add BOS token
_add_toks(
torch.tensor([[self.tokenizer.bos_token_id]], device=self.model.device)
)
if "<image>" not in prompt:
embeds.append(self.text_emb(_tokenize(prompt)))
else:
assert prompt.count("<image>") == 1
before, after = prompt.split("<image>")
embeds.append(self.text_emb(_tokenize(f"{before}<image>")))
embeds.append(image_embeds.to(self.model.device))
embeds.append(self.text_emb(_tokenize(f"</image>{after}")))
return torch.cat(embeds, dim=1)
def generate(
self, image_embeds, prompt, eos_text="Human:", max_new_tokens=128, **kwargs
):
eos_tokens = self.tokenizer(eos_text, add_special_tokens=False)[0].ids
generate_config = {
"eos_token_id": eos_tokens,
"bos_token_id": self.tokenizer.bos_token_id,
"pad_token_id": self.tokenizer.eos_token_id,
"max_new_tokens": max_new_tokens,
**kwargs,
}
with torch.no_grad():
inputs_embeds = self.input_embeds(prompt, image_embeds)
output_ids = self.model.generate(
inputs_embeds=inputs_embeds, **generate_config
)
return self.tokenizer.batch_decode(output_ids, skip_special_tokens=True)
def answer_question(self, image_embeds, question):
prompt = f"<image>\n\nQuestion: {question}\n\nAnswer:"
answer = self.generate(
image_embeds,
prompt,
eos_text="<END>",
max_new_tokens=128,
)[0]
return re.sub("<$", "", re.sub("END$", "", answer)).strip()
+35
View File
@@ -0,0 +1,35 @@
import torch
from PIL import Image
from einops import rearrange
from torchvision.transforms.v2 import (
Compose,
Resize,
InterpolationMode,
ToImage,
ToDtype,
Normalize,
)
class VisionEncoder:
def __init__(self, model_path: str = "model") -> None:
self.model = torch.jit.load(f"{model_path}/vision.pt").to(dtype=torch.float32)
self.preprocess = Compose(
[
Resize(size=(384, 384), interpolation=InterpolationMode.BICUBIC),
ToImage(),
ToDtype(torch.float32, scale=True),
Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]),
]
)
def __call__(self, image: Image) -> torch.Tensor:
with torch.no_grad():
image_vec = self.preprocess(image.convert("RGB")).unsqueeze(0)
image_vec = image_vec[:, :, :-6, :-6]
image_vec = rearrange(
image_vec, "b c (h p1) (w p2) -> b (h w) (c p1 p2)", p1=14, p2=14
)
return self.model(image_vec)
+41 -174
View File
@@ -1,143 +1,43 @@
"""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 transformers import AutoModelForCausalLM, AutoTokenizer
from PIL import Image
from pathlib import Path
from types import ModuleType
import torch
from .runtime import (
CachedModelNode,
ManagedTorchModel,
batch_text,
inference_context,
model_device,
require_module,
snapshot_download,
tensor_batch_to_pil,
torch_dtype,
)
MODEL_ID = "vikhyatk/moondream2"
MODEL_REVISION = "2025-06-21"
_CHECKPOINT_PACKAGE = "_comfyui_vlm_moondream2_checkpoint"
from torchvision.transforms import ToPILImage
from huggingface_hub import snapshot_download
import folder_paths
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()
# Define the directory for saving files related to your new model
files_for_moondream2 = Path(folder_paths.folder_names_and_paths["LLavacheckpoints"][0][0]) / "files_for_moondream2"
files_for_moondream2.mkdir(parents=True, exist_ok=True) # Ensure the directory exists
class Moondream2Predictor:
def __init__(self):
model_path = snapshot_download(
MODEL_ID,
"moondream2",
revision=MODEL_REVISION,
ignore_patterns=["*.bin", "*.gguf"],
)
self.dtype = torch_dtype("bfloat16")
model = _load_native_checkpoint(model_path)
self.handle = ManagedTorchModel(model)
self.model_path = snapshot_download("vikhyatk/moondream2",
local_dir=files_for_moondream2,
force_download=False, # Set to True if you always want to download, regardless of local copy
local_files_only=False, # Set to False to allow downloading if not available locally
revision="2024-04-02", # Specify the revision date for version control
local_dir_use_symlinks="auto", # or set to True/False based on your symlink preference
ignore_patterns=["*.bin", "*.jpg", "*.png", "*.gguf"]) # Customize based on need
self.device = "cuda:0" if torch.cuda.is_available() else "cpu"
self.model = AutoModelForCausalLM.from_pretrained(self.model_path, trust_remote_code=True).to(self.device).eval()
self.tokenizer = AutoTokenizer.from_pretrained(self.model_path)
def close(self):
self.handle.close()
def generate_predictions(self, image_path, question):
# Load and process the image
image_input = Image.open(image_path).convert("RGB")
enc_image = self.model.encode_image(image_input)
def generate(
self,
images,
question,
max_tokens=256,
temperature=0.0,
top_p=0.3,
reasoning=False,
):
results = []
for image in tensor_batch_to_pil(images):
model = self.handle.ensure_loaded()
device = model_device(model)
with torch.inference_mode(), inference_context(device, self.dtype):
response = model.query(
image,
question,
reasoning=bool(reasoning),
settings={
"max_tokens": int(max_tokens),
"temperature": float(temperature),
"top_p": float(top_p),
# Moondream's encoder indexes this optional key
# directly; None selects the base checkpoint.
"variant": None,
},
)
if isinstance(response, dict):
response = response.get("answer", response)
if not str(response).strip():
raise RuntimeError(
"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)
# Generate predictions
generated_text = self.model.answer_question(enc_image, question, self.tokenizer)
return generated_text
class Moondream2model:
def __init__(self):
self.predictor = Moondream2Predictor()
class Moondream2model(CachedModelNode):
@classmethod
def INPUT_TYPES(cls):
return {
@@ -147,59 +47,26 @@ class Moondream2model(CachedModelNode):
"STRING",
{
"multiline": True,
"default": "Describe this image in detail.",
"default": "",
},
),
},
"optional": {
"max_tokens": (
"INT",
{"default": 256, "min": 1, "max": 2048},
),
"temperature": (
"FLOAT",
{"default": 0.0, "min": 0.0, "max": 2.0, "step": 0.05},
),
"top_p": (
"FLOAT",
{"default": 0.3, "min": 0.01, "max": 1.0, "step": 0.01},
),
"reasoning": ("BOOLEAN", {"default": False}),
"unload_after": ("BOOLEAN", {"default": False}),
},
}
RETURN_TYPES = ("STRING",)
FUNCTION = "moondream2_generate_predictions"
CATEGORY = "VLM Nodes/Modern/Edge"
def moondream2_generate_predictions(
self,
image,
text_input,
max_tokens=256,
temperature=0.0,
top_p=0.3,
reasoning=False,
unload_after=False,
):
predictor = self.get_or_create_model(
(MODEL_ID, MODEL_REVISION), Moondream2Predictor
)
try:
return (
predictor.generate(
image,
text_input,
max_tokens,
temperature,
top_p,
reasoning,
),
)
finally:
self.maybe_clear_model(unload_after)
CATEGORY = "VLM Nodes/Moondream2"
def moondream2_generate_predictions(self, image, text_input):
# Convert tensor image to PIL Image
pil_image = ToPILImage()(image[0].permute(2, 0, 1))
temp_path = files_for_moondream2 / "temp_image.png"
pil_image.save(temp_path)
response = self.predictor.generate_predictions(temp_path, text_input)
return (response, )
NODE_CLASS_MAPPINGS = {"Moondream2model": Moondream2model}
NODE_DISPLAY_NAME_MAPPINGS = {"Moondream2model": "Moondream 2"}
NODE_DISPLAY_NAME_MAPPINGS = {"Moondream2model": "Moondream-2 Node"}
-1645
View File
File diff suppressed because it is too large Load Diff
-388
View File
@@ -1,388 +0,0 @@
"""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())
+63 -18
View File
@@ -1,10 +1,38 @@
"""Backward-compatible MoonDream node powered by the current Moondream 2."""
from .moondream import VisionEncoder, TextModel
from huggingface_hub import snapshot_download
import torch
import os
import hashlib
from torchvision import transforms
from pathlib import Path
import folder_paths
from .moondream2 import MODEL_ID, MODEL_REVISION, Moondream2Predictor
from .runtime import CachedModelNode
if torch.cuda.is_available():
DEVICE = "cuda"
DTYPE = torch.float16
else:
DEVICE = "cpu"
DTYPE = torch.float32
class MoonDream(CachedModelNode):
files_for_moondream = Path(folder_paths.folder_names_and_paths["LLavacheckpoints"][0][0]) / "files_for__moondream"
files_for_moondream.mkdir(parents=True, exist_ok=True)
output_directory = os.path.join(files_for_moondream , "output")
# Define your local directory where you want to save the files
image_encoder_cache_path = os.path.join(output_directory, "image_encoder_cache")
class MoonDream:
def __init__(self):
self.model_path = snapshot_download("vikhyatk/moondream1",
revision="5cd8d1ecd7e0d8d95222543e1960d340ddffbfef",
local_dir=files_for_moondream,
force_download=False, # Set to True if you always want to download, regardless of local copy
local_files_only=False, # Set to False to allow downloading if not available locally
local_dir_use_symlinks="auto" # or set to True/False based on your symlink preference
)
self.vision_encoder = VisionEncoder(self.model_path)
self.text_model = TextModel(self.model_path)
@classmethod
def INPUT_TYPES(cls):
return {
@@ -14,28 +42,45 @@ class MoonDream(CachedModelNode):
"STRING",
{
"multiline": True,
"default": "Describe this image in detail.",
"default": "",
},
),
},
"optional": {
"unload_after": ("BOOLEAN", {"default": False})
},
}
RETURN_TYPES = ("STRING",)
FUNCTION = "answer_questions"
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
def answer_questions(self, image, question, unload_after=False):
predictor = self.get_or_create_model(
(MODEL_ID, MODEL_REVISION), Moondream2Predictor
)
try:
return (predictor.generate(image, question),)
finally:
self.maybe_clear_model(unload_after)
CATEGORY = "VLM Nodes/MoonDream"
def process_image(self, image):
# Calculate checksum of the image
image_array = image.numpy() # Convert Tensor to NumPy array
image_hash = hashlib.sha256(image_array.tobytes()).hexdigest()
image = transforms.ToPILImage()(image[0].permute(2, 0, 1))
# Check if `image_encoder_cache/{image_hash}.pt` exists, if so load and return it.
# Otherwise, save the encoded image to `image_encoder_cache/{image_hash}.pt` and return it.
cache_path = f"{image_encoder_cache_path}/{image_hash}.pt"
if os.path.exists(cache_path):
return torch.load(cache_path).to(DEVICE, dtype=DTYPE)
else:
image_vec = self.vision_encoder(image)
os.makedirs(image_encoder_cache_path, exist_ok=True)
torch.save(image_vec, cache_path)
return image_vec.to(DEVICE, dtype=DTYPE)
def answer_questions(self, image, question):
image_embeds = self.process_image(image)
full_sentence = self.text_model.answer_question(image_embeds, question)
return (full_sentence,)
# A dictionary that contains all nodes you want to export with their names
NODE_CLASS_MAPPINGS = {"MoonDream": MoonDream}
NODE_DISPLAY_NAME_MAPPINGS = {"MoonDream": "MoonDream (Moondream 2)"}
# A dictionary that contains the friendly/humanly readable titles for the nodes
NODE_DISPLAY_NAME_MAPPINGS = {"MoonDream": "MoonDream Node"}
+792 -349
View File
File diff suppressed because it is too large Load Diff
+3 -3
View File
@@ -14,8 +14,8 @@ class PlayMusic:
return {"required": {
"mode": (["always", "on empty queue"], {}),
"volume": ("FLOAT", {"min": 0, "max": 1, "step": 0.1, "default": 0.5}),
"wave_form": (any,),
"sample_rate": ("INT",),
"wave_form": ([], {"forceInput": True}),
"sample_rate": ("INT", {"forceInput": True}),
}}
FUNCTION = "nop"
@@ -30,7 +30,7 @@ class PlayMusic:
return float("NaN")
def nop(self, mode, volume, wave_form, sample_rate):
return {"ui": {"a": wave_form, "b": sample_rate}, "result": (wave_form,)}
return {"ui": {"a": wave_form, "b": sample_rate}, "result": (any,)}
NODE_CLASS_MAPPINGS = {
+1 -1
View File
@@ -51,4 +51,4 @@ Optional: If asked to create a random prompt create one.
# Define the system message
system_msg_simple = """
You are an helpful asistant. Answer optional questions or help the user for their optional queries.
"""
"""
+367 -371
View File
@@ -1,413 +1,409 @@
"""Qwen2-VL with real image/video batches and ComfyUI-aware VRAM handling."""
from __future__ import annotations
from typing import Any
import torch
import psutil
import os
from PIL import Image
from pathlib import Path
from torchvision.transforms import ToPILImage
from huggingface_hub import snapshot_download
import folder_paths
from transformers import AutoModelForVision2Seq, AutoTokenizer, AutoProcessor, BitsAndBytesConfig
from qwen_vl_utils import process_vision_info
from .runtime import (
CachedModelNode,
ExternalTorchModel,
ManagedTorchModel,
accelerator_backend,
batch_text,
execution_device,
external_device_map,
inference_context,
model_device,
move_inputs,
require_module,
require_quantization_backend,
reserve_external_vram,
snapshot_download,
tensor_batch_to_pil,
torch_dtype,
)
def check_flash_attention():
"""Check if flash attention 2 is available"""
try:
from flash_attn import flash_attn_func
return True
except ImportError:
return False
FLASH_ATTENTION_AVAILABLE = check_flash_attention()
# Define the directory for saving Qwen2-VL files
files_for_qwen2vl = Path(folder_paths.folder_names_and_paths["LLavacheckpoints"][0][0]) / "files_for_qwen2vl"
files_for_qwen2vl.mkdir(parents=True, exist_ok=True)
# Model VRAM requirements (approximate, in GB)
MODEL_VRAM_REQUIREMENTS = {
"Qwen2-VL-2B": 4,
"Qwen2-VL-7B": 14,
"Qwen2-VL-72B": 40,
"Qwen2-VL-2B-AWQ": 2,
"Qwen2-VL-2B-GPTQ-Int4": 2,
"Qwen2-VL-2B-GPTQ-Int8": 3,
"Qwen2-VL-7B-AWQ": 5,
"Qwen2-VL-7B-GPTQ-Int4": 5,
"Qwen2-VL-7B-GPTQ-Int8": 8,
"Qwen2-VL-72B-AWQ": 20,
"Qwen2-VL-72B-GPTQ-Int4": 20,
"Qwen2-VL-72B-GPTQ-Int8": 25,
}
QWEN2_VL_MODELS = {
"Qwen2-VL-2B": "Qwen/Qwen2-VL-2B-Instruct",
"Qwen2-VL-7B": "Qwen/Qwen2-VL-7B-Instruct",
"Qwen2-VL-72B": "Qwen/Qwen2-VL-72B-Instruct",
}
# Old workflows used separate AWQ/GPTQ repositories whose integration breaks
# across Transformers/AutoGPTQ releases. Resolve those labels to the same base
# weights and the maintained bitsandbytes path instead.
LEGACY_QUANTIZED_ALIASES = {
f"Qwen2-VL-{size}-{quant}": (
f"Qwen2-VL-{size}",
(
"Balanced (8-bit)"
if quant.endswith("Int8")
else "Maximum Savings (4-bit)"
),
)
for size in ("2B", "7B", "72B")
for quant in ("AWQ", "GPTQ-Int4", "GPTQ-Int8")
}
QWEN2_VL_CHOICES = ("Qwen2-VL-2B", "Qwen2-VL-7B")
MEMORY_MODES = [
"ComfyUI managed (BF16)",
"Balanced (8-bit)",
"Maximum Savings (4-bit)",
"CPU Offload",
"Default",
]
ESTIMATED_MODEL_BYTES = {
"Qwen2-VL-2B": 5 * 1024**3,
"Qwen2-VL-7B": 16 * 1024**3,
"Qwen2-VL-72B": 145 * 1024**3,
"Qwen2-VL-2B-AWQ": "Qwen/Qwen2-VL-2B-Instruct-AWQ",
"Qwen2-VL-2B-GPTQ-Int4": "Qwen/Qwen2-VL-2B-Instruct-GPTQ-Int4",
"Qwen2-VL-2B-GPTQ-Int8": "Qwen/Qwen2-VL-2B-Instruct-GPTQ-Int8",
"Qwen2-VL-7B-AWQ": "Qwen/Qwen2-VL-7B-Instruct-AWQ",
"Qwen2-VL-7B-GPTQ-Int4": "Qwen/Qwen2-VL-7B-Instruct-GPTQ-Int4",
"Qwen2-VL-7B-GPTQ-Int8": "Qwen/Qwen2-VL-7B-Instruct-GPTQ-Int8",
"Qwen2-VL-72B-AWQ": "Qwen/Qwen2-VL-72B-Instruct-AWQ",
"Qwen2-VL-72B-GPTQ-Int4": "Qwen/Qwen2-VL-72B-Instruct-GPTQ-Int4",
"Qwen2-VL-72B-GPTQ-Int8": "Qwen/Qwen2-VL-72B-Instruct-GPTQ-Int8",
}
MEMORY_EFFICIENT_CONFIGS = {
"Balanced (8-bit)": {
"load_in_8bit": True,
"load_in_4bit": False,
"cpu_offload": False,
"attention_mode": "flash_attention_2" if FLASH_ATTENTION_AVAILABLE else None,
},
"Maximum Savings (4-bit)": {
"load_in_8bit": False,
"load_in_4bit": True,
"cpu_offload": True,
"attention_mode": "flash_attention_2" if FLASH_ATTENTION_AVAILABLE else None,
},
"CPU Offload": {
"load_in_8bit": False,
"load_in_4bit": False,
"cpu_offload": True,
"attention_mode": None,
},
"Default": {
"load_in_8bit": False,
"load_in_4bit": False,
"cpu_offload": False,
"attention_mode": None,
}
}
def _model_class(transformers):
for name in (
"Qwen2VLForConditionalGeneration",
"AutoModelForImageTextToText",
"AutoModelForMultimodalLM",
):
cls = getattr(transformers, name, None)
if cls is not None:
return cls
raise RuntimeError(
"This Transformers version does not include a Qwen2-VL model class."
)
class SystemResources:
@staticmethod
def get_available_memory():
"""Get available system memory in GB"""
return psutil.virtual_memory().available / (1024 * 1024 * 1024)
@staticmethod
def get_available_vram():
"""Get available VRAM in GB"""
if not torch.cuda.is_available():
return 0
try:
torch.cuda.empty_cache() # Clear unused cached memory
return torch.cuda.get_device_properties(0).total_memory / (1024 * 1024 * 1024)
except:
return 0
def _attention_value(mode: str) -> str:
return {
"Auto (SDPA)": "sdpa",
"Flash Attention 2": "flash_attention_2",
"Eager": "eager",
}[mode]
@staticmethod
def check_resources(model_name, memory_mode):
"""Check if system has enough resources for the model"""
required_vram = MODEL_VRAM_REQUIREMENTS.get(model_name, 0)
config = MEMORY_EFFICIENT_CONFIGS[memory_mode]
# Adjust VRAM requirements based on memory mode
if config["load_in_8bit"]:
required_vram = required_vram * 0.5 # Approximately half VRAM usage
elif config["load_in_4bit"]:
required_vram = required_vram * 0.25 # Approximately quarter VRAM usage
elif config["cpu_offload"]:
required_vram = required_vram * 0.7 # Rough estimate for CPU offloading
available_vram = SystemResources.get_available_vram()
available_memory = SystemResources.get_available_memory()
# Need at least 2GB system memory buffer
required_system_memory = required_vram + 2
error_messages = []
if available_vram < required_vram:
error_messages.append(
f"Insufficient VRAM: Model {model_name} requires {required_vram:.1f}GB VRAM, "
f"but only {available_vram:.1f}GB available. "
"Try using a more aggressive memory saving mode."
)
if available_memory < required_system_memory:
error_messages.append(
f"Insufficient system memory: Need at least {required_system_memory:.1f}GB, "
f"but only {available_memory:.1f}GB available"
)
return error_messages
class Qwen2VLPredictor:
def __init__(
self,
model_name: str,
memory_mode: str,
attention_mode: str,
min_pixels: int,
max_pixels: int,
):
transformers = require_module("transformers")
if model_name in LEGACY_QUANTIZED_ALIASES:
model_name, memory_mode = LEGACY_QUANTIZED_ALIASES[model_name]
if (
attention_mode == "Flash Attention 2"
and accelerator_backend(execution_device())
not in {"nvidia-cuda", "amd-rocm"}
):
raise RuntimeError(
"Flash Attention 2 requires a supported CUDA or ROCm build. "
"Select Auto (SDPA) on Apple Metal, Intel XPU, or CPU."
)
if memory_mode in {"Balanced (8-bit)", "Maximum Savings (4-bit)"}:
# Validate before downloading a multi-gigabyte checkpoint.
require_quantization_backend(memory_mode)
repo_id = QWEN2_VL_MODELS[model_name]
model_path = snapshot_download(
repo_id,
f"qwen2vl/{model_name}",
ignore_patterns=["*.bin"],
def __init__(self, model_name, memory_mode="Balanced (8-bit)"):
# Check system resources
error_messages = SystemResources.check_resources(model_name, memory_mode)
if error_messages:
raise RuntimeError("\n".join(error_messages))
self.model_path = snapshot_download(
QWEN2_VL_MODELS[model_name],
local_dir=files_for_qwen2vl / model_name,
force_download=False,
local_files_only=False,
revision="main"
)
self.dtype = torch_dtype("bfloat16")
self.processor = transformers.AutoProcessor.from_pretrained(
model_path,
min_pixels=int(min_pixels),
max_pixels=int(max_pixels),
)
kwargs: dict[str, Any] = {
"torch_dtype": self.dtype,
"attn_implementation": _attention_value(attention_mode),
}
external = memory_mode in {
"Balanced (8-bit)",
"Maximum Savings (4-bit)",
"CPU Offload",
}
if memory_mode in {"Balanced (8-bit)", "Maximum Savings (4-bit)"}:
kwargs["quantization_config"] = transformers.BitsAndBytesConfig(
load_in_8bit=memory_mode == "Balanced (8-bit)",
load_in_4bit=memory_mode == "Maximum Savings (4-bit)",
bnb_4bit_compute_dtype=self.dtype,
bnb_4bit_quant_type="nf4",
bnb_4bit_use_double_quant=True,
)
if external:
require_module("accelerate")
estimate = ESTIMATED_MODEL_BYTES.get(
model_name.split("-AWQ", 1)[0].split("-GPTQ", 1)[0],
8 * 1024**3,
)
reserve_external_vram(
estimate // (4 if memory_mode == "Maximum Savings (4-bit)" else 2)
)
kwargs["device_map"] = external_device_map(
allow_auto_offload=memory_mode == "CPU Offload"
)
self.device = "cuda" if torch.cuda.is_available() else "cpu"
try:
model = _model_class(transformers).from_pretrained(
model_path, **kwargs
).eval()
except ImportError as exc:
if attention_mode == "Flash Attention 2":
# Get memory configuration
config = MEMORY_EFFICIENT_CONFIGS[memory_mode]
# Base model kwargs
model_kwargs = {
"trust_remote_code": True,
"device_map": "auto" if config["cpu_offload"] else None,
}
# Setup quantization config if needed
if config["load_in_8bit"] or config["load_in_4bit"]:
model_kwargs.update({
"load_in_8bit": config["load_in_8bit"],
"load_in_4bit": config["load_in_4bit"],
"bnb_4bit_compute_dtype": torch.float16,
"bnb_4bit_use_double_quant": True,
})
# Add attention optimization if specified and available
if config["attention_mode"]:
try:
model_kwargs["attn_implementation"] = config["attention_mode"]
except Exception as e:
print(f"Warning: Flash Attention 2 requested but not available: {str(e)}")
# Set appropriate dtype based on model type
if "GPTQ" in model_name or "AWQ" in model_name:
model_kwargs["torch_dtype"] = "auto"
else:
model_kwargs["torch_dtype"] = torch.float16 if torch.cuda.is_available() else torch.float32
self.model = AutoModelForVision2Seq.from_pretrained(
self.model_path,
**model_kwargs
)
self.processor = AutoProcessor.from_pretrained(self.model_path)
self.tokenizer = AutoTokenizer.from_pretrained(self.model_path, trust_remote_code=True)
except RuntimeError as e:
if "out of memory" in str(e):
process = psutil.Process()
mem_info = process.memory_info()
torch.cuda.empty_cache()
raise RuntimeError(
"Flash Attention 2 was selected but flash-attn is not "
"installed for this PyTorch accelerator build. Use Auto "
"(SDPA), or install a matching flash-attn wheel."
) from exc
f"Out of VRAM while loading {model_name}. Try:\n"
"1. Using a more aggressive memory saving mode\n"
"2. Using a smaller model (e.g., 2B instead of 7B)\n"
"3. Using a quantized version (AWQ/GPTQ)\n"
"4. Clearing other models from memory\n"
"5. Restarting ComfyUI\n"
f"Process memory: {mem_info.rss / 1024**3:.1f}GB"
) from e
raise
self.handle = (
ExternalTorchModel(model, processor=self.processor)
if external
else ManagedTorchModel(model, processor=self.processor)
)
def close(self):
self.handle.close()
self.processor = None
def _generate_messages(
self,
messages,
*,
max_new_tokens,
temperature,
top_p,
) -> str:
process_vision_info = require_module(
"qwen_vl_utils", "qwen-vl-utils"
).process_vision_info
text = self.processor.apply_chat_template(
messages, tokenize=False, add_generation_prompt=True
)
image_inputs, video_inputs = process_vision_info(messages)
inputs = self.processor(
text=[text],
images=image_inputs,
videos=video_inputs,
padding=True,
return_tensors="pt",
)
model = self.handle.ensure_loaded()
device = model_device(model)
inputs = move_inputs(inputs, device)
generation: dict[str, Any] = {
"max_new_tokens": int(max_new_tokens),
"do_sample": float(temperature) > 0.0,
}
if generation["do_sample"]:
generation.update(
temperature=float(temperature), top_p=float(top_p)
)
tokenizer = getattr(self.processor, "tokenizer", None)
if tokenizer is not None:
generation["pad_token_id"] = tokenizer.pad_token_id
generation["eos_token_id"] = tokenizer.eos_token_id
with torch.inference_mode(), inference_context(device, self.dtype):
output_ids = model.generate(**inputs, **generation)
trimmed = [
output[len(input_ids) :]
for input_ids, output in zip(inputs["input_ids"], output_ids)
]
return self.processor.batch_decode(
trimmed,
skip_special_tokens=True,
clean_up_tokenization_spaces=False,
)[0].strip()
def generate_images(
self, images, prompt, max_new_tokens, temperature, top_p
) -> str:
results = []
for image in tensor_batch_to_pil(images):
messages = [
{
"role": "user",
"content": [
{"type": "image", "image": image},
{"type": "text", "text": prompt},
],
}
]
results.append(
self._generate_messages(
messages,
max_new_tokens=max_new_tokens,
temperature=temperature,
top_p=top_p,
)
)
return batch_text(results)
def generate_video(
self,
primary_image,
frames,
prompt,
max_new_tokens,
temperature,
top_p,
fps,
) -> str:
# The still IMAGE socket is required by ComfyUI for backwards
# compatibility, but a connected frame batch is the visual source for
# video inference. Mixing both causes small VLMs to answer from the
# still and ignore temporal content.
del primary_image
frame_list = tensor_batch_to_pil(frames)
def process_video(self, video_frames, fps=1.0):
"""Process video frames for video understanding"""
messages = [
{
"role": "user",
"content": [
{
"type": "video",
"video": frame_list,
"fps": float(fps),
},
{
"type": "text",
"text": (
f"The video frames are sampled at {float(fps):g} "
f"FPS.\n\n{prompt}"
),
},
],
"video": video_frames,
"fps": fps
}
]
}
]
return self._generate_messages(
messages,
max_new_tokens=max_new_tokens,
temperature=temperature,
top_p=top_p,
)
return messages
def generate_predictions(self, image_path, prompt, max_new_tokens=512, temperature=0.7, top_p=0.9, video_frames=None, fps=1.0):
try:
# Handle video input if provided
if video_frames:
messages = self.process_video(video_frames, fps)
messages[0]["content"].append({"type": "text", "text": prompt})
else:
# Standard image processing
messages = [
{
"role": "user",
"content": [
{"type": "image", "image": str(image_path)},
{"type": "text", "text": prompt}
]
}
]
# Process the inputs
text = self.processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
image_inputs, video_inputs = process_vision_info(messages)
inputs = self.processor(
text=[text],
images=image_inputs,
videos=video_inputs,
return_tensors="pt",
padding=True
)
try:
inputs = {k: v.to(self.device) for k, v in inputs.items()}
# Generate response
output_ids = self.model.generate(
**inputs,
max_new_tokens=max_new_tokens,
temperature=temperature,
top_p=top_p,
pad_token_id=self.tokenizer.pad_token_id,
eos_token_id=self.tokenizer.eos_token_id
)
# Decode and return the response
generated_text = self.tokenizer.decode(output_ids[0][inputs["input_ids"].shape[1]:], skip_special_tokens=True)
return generated_text.strip()
except RuntimeError as e:
if "out of memory" in str(e):
torch.cuda.empty_cache()
raise RuntimeError(
"Out of VRAM during generation. Try:\n"
"1. Using a more aggressive memory saving mode\n"
"2. Reducing max_new_tokens\n"
"3. Using a smaller model\n"
"4. Using a quantized version (AWQ/GPTQ)"
) from e
raise
except Exception as e:
return f"Error during generation: {str(e)}"
class Qwen2VLNode(CachedModelNode):
class Qwen2VLNode:
def __init__(self):
self.predictor = None
self.current_model = None
self.current_memory_mode = None
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"text_input": (
"STRING",
{
"multiline": True,
"default": "Describe this image in detail.",
},
),
"model_name": (list(QWEN2_VL_CHOICES),),
"memory_mode": (
MEMORY_MODES,
{"default": "ComfyUI managed (BF16)"},
),
"max_new_tokens": (
"INT",
{"default": 512, "min": 1, "max": 8192},
),
"temperature": (
"FLOAT",
{"default": 0.2, "min": 0.0, "max": 2.0, "step": 0.1},
),
"top_p": (
"FLOAT",
{"default": 0.9, "min": 0.0, "max": 1.0, "step": 0.05},
),
"image": ("IMAGE",),
"text_input": ("STRING", {
"multiline": True,
"default": "Describe this image in detail."
}),
"model_name": (list(QWEN2_VL_MODELS.keys()),),
"memory_mode": (list(MEMORY_EFFICIENT_CONFIGS.keys()),),
"max_new_tokens": ("INT", {
"default": 512,
"min": 1,
"max": 2048
}),
"temperature": ("FLOAT", {
"default": 0.7,
"min": 0.1,
"max": 2.0,
"step": 0.1
}),
"top_p": ("FLOAT", {
"default": 0.9,
"min": 0.1,
"max": 1.0,
"step": 0.1
})
},
"optional": {
"image": ("IMAGE",),
"video_frames": ("IMAGE",),
"fps": (
"FLOAT",
{"default": 1.0, "min": 0.1, "max": 60.0, "step": 0.1},
),
"attention_mode": (
["Auto (SDPA)", "Flash Attention 2", "Eager"],
{"default": "Auto (SDPA)"},
),
"min_pixels": (
"INT",
{"default": 256 * 28 * 28, "min": 28 * 28},
),
"max_pixels": (
"INT",
{"default": 1280 * 28 * 28, "min": 28 * 28},
),
"unload_after": ("BOOLEAN", {"default": False}),
},
"fps": ("FLOAT", {
"default": 1.0,
"min": 0.1,
"max": 30.0,
"step": 0.1
})
}
}
RETURN_TYPES = ("STRING",)
FUNCTION = "generate"
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
CATEGORY = "VLM Nodes/Qwen2-VL"
def generate(
self,
text_input,
model_name,
memory_mode="ComfyUI managed (BF16)",
max_new_tokens=512,
temperature=0.2,
top_p=0.9,
image=None,
video_frames=None,
fps=1.0,
attention_mode="Auto (SDPA)",
min_pixels=256 * 28 * 28,
max_pixels=1280 * 28 * 28,
unload_after=False,
):
if min_pixels > max_pixels:
raise ValueError("min_pixels cannot be greater than max_pixels.")
if image is None and video_frames is None:
raise ValueError("Connect either image or video_frames.")
key = (
model_name,
memory_mode,
attention_mode,
int(min_pixels),
int(max_pixels),
)
predictor = self.get_or_create_model(
key,
lambda: Qwen2VLPredictor(
model_name,
memory_mode,
attention_mode,
min_pixels,
max_pixels,
),
)
def generate(self, image, text_input, model_name, memory_mode="Balanced (8-bit)",
max_new_tokens=512, temperature=0.7, top_p=0.9, video_frames=None, fps=1.0):
# Initialize or update predictor if model or memory mode changed
if (self.predictor is None or self.current_model != model_name or
self.current_memory_mode != memory_mode):
# Clean up old model
if self.predictor is not None:
del self.predictor.model
del self.predictor.processor
del self.predictor.tokenizer
torch.cuda.empty_cache()
try:
self.predictor = Qwen2VLPredictor(model_name, memory_mode)
self.current_model = model_name
self.current_memory_mode = memory_mode
except Exception as e:
return (f"Error initializing model: {str(e)}",)
# Convert tensor image to PIL Image and save temporarily
pil_image = ToPILImage()(image[0].permute(2, 0, 1))
temp_path = files_for_qwen2vl / "temp_image.png"
pil_image.save(temp_path)
video_frame_list = None
if video_frames is not None:
video_frame_list = [str(temp_path)] # Use current image as first frame
# Add additional video frames if provided
for frame in video_frames[1:]:
frame_path = files_for_qwen2vl / f"temp_frame_{len(video_frame_list)}.png"
ToPILImage()(frame.permute(2, 0, 1)).save(frame_path)
video_frame_list.append(str(frame_path))
try:
if video_frames is None:
result = predictor.generate_images(
image,
text_input,
max_new_tokens,
temperature,
top_p,
)
else:
result = predictor.generate_video(
image,
video_frames,
text_input,
max_new_tokens,
temperature,
top_p,
fps,
)
return (result,)
finally:
self.maybe_clear_model(unload_after)
# Generate response
response = self.predictor.generate_predictions(
temp_path,
text_input,
max_new_tokens=max_new_tokens,
temperature=temperature,
top_p=top_p,
video_frames=video_frame_list,
fps=fps
)
# Clean up all temporary files
try:
os.remove(temp_path)
if video_frame_list:
for frame_path in video_frame_list[1:]:
try:
os.remove(frame_path)
except:
pass
except:
pass
return (response,)
except Exception as e:
return (f"Error during generation: {str(e)}",)
# Register the node
NODE_CLASS_MAPPINGS = {
"Qwen2VLNode": Qwen2VLNode
}
NODE_CLASS_MAPPINGS = {"Qwen2VLNode": Qwen2VLNode}
NODE_DISPLAY_NAME_MAPPINGS = {"Qwen2VLNode": "Qwen2-VL"}
NODE_DISPLAY_NAME_MAPPINGS = {
"Qwen2VLNode": "Qwen2-VL Model"
}
-2457
View File
File diff suppressed because it is too large Load Diff
-1167
View File
File diff suppressed because it is too large Load Diff
-539
View File
@@ -1,539 +0,0 @@
"""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",
}
-560
View File
@@ -1,560 +0,0 @@
"""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",
]
+101 -1074
View File
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+627 -892
View File
File diff suppressed because it is too large Load Diff
-766
View File
@@ -1,766 +0,0 @@
"""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",
]
+82 -92
View File
@@ -1,86 +1,84 @@
"""UForm Gen2 Qwen node with safe lazy loading."""
from __future__ import annotations
from pathlib import Path
from transformers import AutoModel, AutoProcessor, StoppingCriteria, StoppingCriteriaList
import torch
from PIL import Image
from torchvision.transforms import ToPILImage
from huggingface_hub import snapshot_download
import folder_paths
# Define the directory for saving files related to uform-gen2-qwen
files_for_uform_gen2_qwen = Path(folder_paths.folder_names_and_paths["LLavacheckpoints"][0][0]) / "files_for_uform_gen2_qwen"
files_for_uform_gen2_qwen.mkdir(parents=True, exist_ok=True) # Ensure the directory exists
from .runtime import (
CachedModelNode,
ManagedTorchModel,
batch_text,
inference_context,
model_device,
require_module,
snapshot_download,
tensor_batch_to_pil,
torch_dtype,
)
MODEL_ID = "unum-cloud/uform-gen2-qwen-500m"
class StopOnTokens(StoppingCriteria):
def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor, **kwargs) -> bool:
stop_ids = [151645] # Define stop tokens as per your model's specifics
for stop_id in stop_ids:
if input_ids[0][-1] == stop_id:
return True
return False
class UformGen2QwenChat:
def __init__(self):
transformers = require_module("transformers")
model_path = snapshot_download(
MODEL_ID, "uform-gen2-qwen", ignore_patterns=["*.bin"]
self.model_path = snapshot_download("unum-cloud/uform-gen2-qwen-500m",
local_dir=files_for_uform_gen2_qwen,
force_download=False, # Set to True if you always want to download, regardless of local copy
local_files_only=False, # Set to False to allow downloading if not available locally
local_dir_use_symlinks="auto") # or set to True/False based on your symlink preference
self.device = "cuda:0" if torch.cuda.is_available() else "cpu"
self.model = AutoModel.from_pretrained(self.model_path, trust_remote_code=True).to(self.device)
self.processor = AutoProcessor.from_pretrained(self.model_path, trust_remote_code=True)
def chat_response(self, message, history, image_path):
stop = StopOnTokens()
messages = [{"role": "system", "content": "You are a helpful Assistant."}]
for user_msg, assistant_msg in history:
messages.append({"role": "user", "content": user_msg})
messages.append({"role": "assistant", "content": assistant_msg})
if len(messages) == 1:
message = f" <image>{message}"
messages.append({"role": "user", "content": message})
model_inputs = self.processor.tokenizer.apply_chat_template(
messages,
add_generation_prompt=True,
return_tensors="pt"
)
self.dtype = torch_dtype("float16")
model = transformers.AutoModel.from_pretrained(
model_path,
trust_remote_code=True,
torch_dtype=self.dtype,
).eval()
self.processor = transformers.AutoProcessor.from_pretrained(
model_path, trust_remote_code=True
image = Image.open(image_path) # Load image using PIL
image_tensor = (
self.processor.feature_extractor(image)
.unsqueeze(0)
)
self.handle = ManagedTorchModel(model, processor=self.processor)
def close(self):
self.handle.close()
self.processor = None
attention_mask = torch.ones(
1, model_inputs.shape[1] + self.processor.num_image_latents - 1
)
def chat(self, images, question, max_new_tokens):
results = []
for image in tensor_batch_to_pil(images):
messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": f"<image>{question}"},
]
input_ids = self.processor.tokenizer.apply_chat_template(
messages,
add_generation_prompt=True,
return_tensors="pt",
)
image_tensor = self.processor.feature_extractor(image).unsqueeze(0)
attention_mask = torch.ones(
1,
input_ids.shape[1] + self.processor.num_image_latents - 1,
dtype=torch.long,
)
model = self.handle.ensure_loaded()
device = model_device(model)
model_inputs = {
"input_ids": input_ids.to(device),
"images": image_tensor.to(device),
"attention_mask": attention_mask.to(device),
}
with torch.inference_mode(), inference_context(device, self.dtype):
output = model.generate(
**model_inputs,
max_new_tokens=int(max_new_tokens),
eos_token_id=self.processor.tokenizer.eos_token_id,
)
generated = output[0, input_ids.shape[-1] :]
results.append(
self.processor.tokenizer.decode(
generated, skip_special_tokens=True
).strip()
)
return batch_text(results)
model_inputs = {
"input_ids": model_inputs,
"images": image_tensor,
"attention_mask": attention_mask
}
model_inputs = {k: v.to(self.device) for k, v in model_inputs.items()}
output = self.model.generate(
**model_inputs,
max_new_tokens=1024,
stopping_criteria=StoppingCriteriaList([stop])
)
response_text = self.processor.tokenizer.decode(output[0], skip_special_tokens=True)
return response_text
# Example of integrating UformGen2QwenChat into a node-like structure
class UformGen2QwenNode:
def __init__(self):
self.chat_model = UformGen2QwenChat()
class UformGen2QwenNode(CachedModelNode):
@classmethod
def INPUT_TYPES(cls):
return {
@@ -90,34 +88,26 @@ class UformGen2QwenNode(CachedModelNode):
"STRING",
{
"multiline": True,
"default": "Describe this image in detail.",
"default": "",
},
),
},
"optional": {
"max_new_tokens": (
"INT",
{"default": 512, "min": 1, "max": 4096},
),
"unload_after": ("BOOLEAN", {"default": False}),
},
}
RETURN_TYPES = ("STRING",)
FUNCTION = "uform_gen2_qwen_chat"
CATEGORY = "VLM Nodes/Legacy/Model Loaders"
def uform_gen2_qwen_chat(
self, image, question, max_new_tokens=512, unload_after=False
):
predictor = self.get_or_create_model(
MODEL_ID, UformGen2QwenChat
)
try:
return (predictor.chat(image, question, max_new_tokens),)
finally:
self.maybe_clear_model(unload_after)
CATEGORY = "VLM Nodes/UformGen2Qwen"
def uform_gen2_qwen_chat(self, image, question):
history = [] # Example empty history
pil_image = ToPILImage()(image[0].permute(2, 0, 1))
temp_path = files_for_uform_gen2_qwen / "temp.png"
pil_image.save(temp_path)
response = self.chat_model.chat_response(question, history, temp_path)
return (response.split("assistant\n", 1)[1], )
NODE_CLASS_MAPPINGS = {"UformGen2QwenNode": UformGen2QwenNode}
NODE_DISPLAY_NAME_MAPPINGS = {"UformGen2QwenNode": "UForm Gen2 Qwen"}
NODE_DISPLAY_NAME_MAPPINGS = {"UformGen2QwenNode": "UformGen2 Qwen Node"}
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
+6 -117
View File
@@ -1,126 +1,15 @@
[project]
name = "comfyui_vlm_nodes"
version = "3.5.0"
description = "Production-ready local and API vision-language nodes for ComfyUI"
readme = "README.md"
requires-python = ">=3.10"
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",
"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 = [
"Operating System :: Microsoft :: Windows",
"Operating System :: POSIX :: Linux",
"Operating System :: MacOS",
"Environment :: GPU :: NVIDIA CUDA",
"Environment :: GPU :: AMD ROCm",
"Environment :: GPU :: Intel Arc",
"Environment :: GPU :: Apple Metal",
"Programming Language :: Python :: 3.10",
"Programming Language :: Python :: 3.11",
"Programming Language :: Python :: 3.12",
"Programming Language :: Python :: 3.13",
]
[project.optional-dependencies]
quantization = [
"accelerate>=1.1,<2",
"bitsandbytes>=0.50,<1",
]
gguf = [
"llama-cpp-python>=0.3.20,<1",
]
robotics-client = [
"msgpack>=1.0.8,<2",
"pyzmq>=26,<28",
"websockets>=14,<17",
]
description = "Custom Nodes for Vision Language Models (VLM) , Large Language Models (LLM), Image Captioning, Automatic Prompt Generation, Creative and Consistent Prompt Suggestion, Keyword Extraction"
version = "1.0.6"
license = { file = "LICENSE" }
dependencies = ["accelerate>=0.27.0", "bitsandbytes", "cffi", "decord" , "diffusers" , "diskcache" , "einops>=0.7.0" , "gitpython", "huggingface-hub>=0.20.3", "moviepy", "openai>=0.27.8", "opencv-python", "optimum>=1.17.0", "pillow>=9.4.0", "py-cpuinfo>=3.3.0", "python-dateutil>=2.7.0", "pytz", "qwen-vl-utils", "safetensors>=0.4.1", "scikit-build", "six", "soundfile", "symusic", "torch>=2.0.1,<3.0.0", "torchvision>=0.15.2", "transformers>=4.38.2", "typing"]
[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"
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"]
Models = [{location = "/checkpoints/model.safetensor", model_url = "https://example.com/model.zip"}]
-8
View File
@@ -1,8 +0,0 @@
# 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
-4
View File
@@ -1,4 +0,0 @@
# Optional GGUF backend. This default may build the CPU backend from source.
# Prefer the official CUDA, Metal, ROCm/HIP, Vulkan, or SYCL wheel/build from:
# https://github.com/abetlen/llama-cpp-python
llama-cpp-python>=0.3.20,<1
-11
View File
@@ -1,11 +0,0 @@
# Install this file only into the isolated Moondream sidecar environment.
# Do not install it into ComfyUI's main environment: moondream 1.3 pins
# Pillow <11 while current ComfyUI uses a newer Pillow release.
moondream==1.3.0
# moondream 1.3.0 expects this exact runtime API. 0.4.7+ renamed the
# prefix-mask kernel and is not source-compatible with kestrel 0.4.2.
kestrel-kernels==0.4.6
# Kestrel's CUDA 12 AOT kernels call cudaLibraryLoadData. PyTorch's cu126
# runtime (12.6.77) does not export it; 12.9.79 does and remains within the
# CUDA 12 ABI. Keep this inside the isolated Photon environment only.
nvidia-cuda-runtime-cu12==12.9.79; (sys_platform == "linux" and platform_machine == "x86_64") or (sys_platform == "win32" and platform_machine == "AMD64")
-5
View File
@@ -1,5 +0,0 @@
# Optional maintained 4-bit/8-bit backend.
# Official 0.50+ wheels cover NVIDIA CUDA, AMD ROCm, Intel XPU/CPU,
# Apple Silicon, and supported Windows/Linux CPU architectures.
accelerate>=1.1,<2
bitsandbytes>=0.50,<1
-5
View File
@@ -1,5 +0,0 @@
# 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
+29 -21
View File
@@ -1,21 +1,29 @@
# ComfyUI provides torch, torchvision, numpy and Pillow.
# Keep this list resolver-friendly; no package is installed during node import.
accelerate>=1.1,<2
# Official wheels: Linux x86_64/aarch64, Windows AMD64/ARM64, macOS arm64.
# Unsupported machines keep every non-quantized node instead of failing install.
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
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
accelerate>=0.32.1
bitsandbytes
cffi
decord
diffusers >=0.31.0
diskcache
einops>=0.7.0
gitpython
huggingface-hub>=0.26.2
matplotlib
moviepy
numpy>=1.22.4,<2.0.0
openai>=0.27.8
opencv-python
optimum>=1.17.0
pillow>=9.4.0
py-cpuinfo>=3.3.0
python-dateutil>=2.7.0
pytz
qwen-vl-utils
safetensors>=0.4.1
scikit-build
six
soundfile
symusic
torch>=2.0.1
torchvision>=0.15.2
transformers>=4.46
typing
-50
View File
@@ -1,50 +0,0 @@
"""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)
-31
View File
@@ -1,31 +0,0 @@
"""Make the source checkout importable on every supported test runner."""
from __future__ import annotations
import importlib.util
import sys
from pathlib import Path
REPOSITORY = Path(__file__).resolve().parents[1]
for candidate in (
REPOSITORY.parent,
REPOSITORY.parent / "ComfyUI",
REPOSITORY.parents[1],
):
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)
-49
View File
@@ -1,49 +0,0 @@
"""Validate curated Hugging Face IDs without downloading model weights.
This opt-in network check resolves each repository's configuration and
processor through the installed Transformers version. It complements, but does
not replace, the real-weight smoke tests.
python tests/manual_catalog_probe.py
"""
from __future__ import annotations
import json
from _bootstrap import bootstrap
bootstrap()
from ComfyUI_VLM_nodes.nodes.modern_vlm import MODEL_CATALOG # noqa: E402
from transformers import AutoConfig, AutoProcessor # noqa: E402
def main() -> int:
records = []
for label, spec in MODEL_CATALOG.items():
if not spec.small_fast or spec.gated:
continue
config = AutoConfig.from_pretrained(
spec.repo_id,
trust_remote_code=spec.trust_remote_code,
)
processor = AutoProcessor.from_pretrained(
spec.repo_id,
trust_remote_code=spec.trust_remote_code,
)
records.append(
{
"label": label,
"repo_id": spec.repo_id,
"model_type": config.model_type,
"config_class": type(config).__name__,
"processor_class": type(processor).__name__,
}
)
print("CATALOG_PROBE_JSON=" + json.dumps(records, ensure_ascii=False))
return 0
if __name__ == "__main__":
raise SystemExit(main())
-102
View File
@@ -1,102 +0,0 @@
"""Opt-in real-weight smoke test for the shared llama.cpp runtime.
The default model is a small official ggml-org Qwen checkpoint. Nothing is
downloaded unless --download is supplied.
"""
from __future__ import annotations
import argparse
import json
import time
from pathlib import Path
from _bootstrap import bootstrap
bootstrap()
from ComfyUI_VLM_nodes.nodes.runtime import ( # noqa: E402
LlamaHandle,
default_llama_threads,
hf_download,
llama_cpp_diagnostics,
)
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(
"--prompt",
default="Reply with exactly: llama.cpp runtime ready",
)
parser.add_argument("--max-tokens", type=int, default=32)
parser.add_argument("--n-gpu-layers", type=int, default=-1)
return parser.parse_args()
def main() -> None:
args = parse_args()
model_path = args.model
if model_path is None:
if not args.download:
raise SystemExit(
"Pass --model /path/to/model.gguf, or explicitly allow the "
"small default download with --download."
)
model_path = hf_download(
args.repo,
args.filename,
"llama-cpp-smoke",
)
started = time.perf_counter()
handle = LlamaHandle(
model_path,
n_ctx=2048,
n_gpu_layers=args.n_gpu_layers,
n_threads=default_llama_threads(),
n_batch=512,
n_ubatch=512,
flash_attention="Auto",
)
try:
llm = handle.ensure_loaded()
loaded = time.perf_counter()
response = llm.create_chat_completion(
messages=[{"role": "user", "content": args.prompt}],
max_tokens=args.max_tokens,
temperature=0.0,
seed=42,
)
finished = time.perf_counter()
content = response["choices"][0]["message"]["content"]
print(
json.dumps(
{
"model": str(model_path),
"model_bytes": model_path.stat().st_size,
"llama_cpp": llama_cpp_diagnostics(),
"load_seconds": round(loaded - started, 3),
"generation_seconds": round(finished - loaded, 3),
"response": content,
},
ensure_ascii=False,
indent=2,
)
)
finally:
handle.close()
if __name__ == "__main__":
main()
-178
View File
@@ -1,178 +0,0 @@
"""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()
-126
View File
@@ -1,126 +0,0 @@
"""Opt-in real-weight smoke test for the Modern VLM node.
This is intentionally excluded from pytest because it downloads multi-gigabyte
models. Run one checkpoint per process so CUDA and file-handle cleanup are also
exercised:
python tests/manual_model_smoke.py --model "Qwen 3.5 2B"
"""
from __future__ import annotations
import argparse
import json
import time
import torch
from _bootstrap import bootstrap
bootstrap()
from ComfyUI_VLM_nodes.nodes.modern_vlm import ( # noqa: E402
MODEL_CATALOG,
ModernVLMPredictor,
)
def test_image() -> torch.Tensor:
image = torch.zeros((1, 96, 128, 3), dtype=torch.float32)
image[:, 20:76, 28:104, 0] = 1.0
return image
def test_video() -> torch.Tensor:
frames = torch.zeros((4, 96, 128, 3), dtype=torch.float32)
for index in range(4):
left = 12 + index * 18
frames[index, 30:66, left : left + 24, 1] = 1.0
return frames
def main() -> int:
parser = argparse.ArgumentParser()
parser.add_argument("--model", required=True, choices=MODEL_CATALOG)
parser.add_argument(
"--memory-mode",
default="ComfyUI managed (BF16)",
choices=[
"ComfyUI managed (BF16)",
"4-bit NF4 (bitsandbytes)",
"8-bit (bitsandbytes)",
"CPU",
],
)
parser.add_argument("--video", action="store_true")
parser.add_argument("--max-new-tokens", type=int, default=48)
args = parser.parse_args()
started = time.perf_counter()
if torch.cuda.is_available():
torch.cuda.reset_peak_memory_stats()
free_before, total = torch.cuda.mem_get_info()
else:
free_before = total = 0
predictor = ModernVLMPredictor(
args.model,
"",
args.memory_mode,
"Auto (SDPA)",
)
try:
prompt = (
"In this four-frame video, what color object moves horizontally? "
"Answer with the color and shape."
if args.video
else (
"Describe the dominant colors, shapes, and motion in one "
"short factual sentence."
)
)
response = predictor.generate(
None if args.video else test_image(),
prompt,
"",
args.max_new_tokens,
0.0,
0.9,
test_video() if args.video else None,
2.0,
)
if not response.strip():
raise RuntimeError("The model returned an empty response.")
if args.video and "green" not in response.lower():
raise RuntimeError(
f"The video frames were not understood; response was: {response}"
)
if not args.video and "red" not in response.lower():
raise RuntimeError(
f"The image was not understood; response was: {response}"
)
finally:
predictor.close()
if torch.cuda.is_available():
peak = torch.cuda.max_memory_allocated()
free_after, _ = torch.cuda.mem_get_info()
else:
peak = free_after = 0
record = {
"model": args.model,
"repo_id": MODEL_CATALOG[args.model].repo_id,
"memory_mode": args.memory_mode,
"video": args.video,
"response": response,
"seconds": round(time.perf_counter() - started, 2),
"cuda_total_gib": round(total / 1024**3, 2),
"cuda_free_before_gib": round(free_before / 1024**3, 2),
"cuda_free_after_gib": round(free_after / 1024**3, 2),
"cuda_peak_allocated_gib": round(peak / 1024**3, 2),
}
print("MODEL_SMOKE_JSON=" + json.dumps(record, ensure_ascii=False))
return 0
if __name__ == "__main__":
raise SystemExit(main())
-158
View File
@@ -1,158 +0,0 @@
"""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()
-296
View File
@@ -1,296 +0,0 @@
"""Opt-in real-weight smoke tests for specialized model backends.
Each invocation downloads and runs one real checkpoint. Keeping one model per
process verifies teardown and prevents one backend's CUDA state from masking
another backend's behavior.
"""
from __future__ import annotations
import argparse
import json
import time
import torch
from _bootstrap import bootstrap
bootstrap()
BACKENDS = (
"florence-base",
"florence-large",
"moondream2",
"qwen2vl-2b",
"qwen2vl-2b-video",
"qwen2vl-7b-4bit",
"molmo-1b",
"molmo-7b-d-4bit",
"molmo-7b-o-4bit",
"kosmos2",
"uform",
"mcllava",
"joytag",
"paligemma-caption",
"minicpm-gguf-q4",
"audioldm2",
)
def test_image() -> torch.Tensor:
image = torch.zeros((1, 192, 256, 3), dtype=torch.float32)
image[:, 48:144, 56:200, 0] = 1.0
return image
def test_video() -> torch.Tensor:
frames = torch.zeros((4, 96, 128, 3), dtype=torch.float32)
for index in range(4):
left = 12 + index * 18
frames[index, 30:66, left : left + 24, 1] = 1.0
return frames
def _run(backend: str):
image = test_image()
prompt = "What color is the large rectangle? Answer briefly."
if backend.startswith("florence-"):
from ComfyUI_VLM_nodes.nodes.florence2 import FlorencePredictor
from ComfyUI_VLM_nodes.nodes.runtime import tensor_batch_to_pil
label = {
"florence-base": "Florence-2 base FT (fast)",
"florence-large": "Florence-2 large FT (recommended)",
}[backend]
predictor = FlorencePredictor(label)
try:
raw, parsed = predictor.run(
tensor_batch_to_pil(image)[0],
"<MORE_DETAILED_CAPTION>",
"",
96,
3,
)
return {"response": raw, "parsed": parsed}
finally:
predictor.close()
if backend == "moondream2":
from ComfyUI_VLM_nodes.nodes.moondream2 import Moondream2Predictor
predictor = Moondream2Predictor()
try:
return {"response": predictor.generate(image, prompt)}
finally:
predictor.close()
if backend.startswith("qwen2vl-"):
from ComfyUI_VLM_nodes.nodes.qwen2vl import Qwen2VLPredictor
model_name, memory_mode = {
"qwen2vl-2b": ("Qwen2-VL-2B", "ComfyUI managed (BF16)"),
"qwen2vl-2b-video": (
"Qwen2-VL-2B",
"ComfyUI managed (BF16)",
),
"qwen2vl-7b-4bit": ("Qwen2-VL-7B", "Maximum Savings (4-bit)"),
}[backend]
predictor = Qwen2VLPredictor(
model_name,
memory_mode,
"Auto (SDPA)",
256 * 28 * 28,
1280 * 28 * 28,
)
try:
if backend.endswith("-video"):
return {
"response": predictor.generate_video(
None,
test_video(),
(
"What color object moves horizontally? Answer with "
"the color and shape."
),
48,
0.0,
0.9,
2.0,
)
}
return {
"response": predictor.generate_images(
image, prompt, 48, 0.0, 0.9
)
}
finally:
predictor.close()
if backend.startswith("molmo-"):
from ComfyUI_VLM_nodes.nodes.molmo import MolmoPredictor
from ComfyUI_VLM_nodes.nodes.runtime import tensor_batch_to_pil
model_name, memory_mode = {
"molmo-1b": (
"MolmoE-1B (Efficient)",
"Full Precision (45GB+ Required)",
),
"molmo-7b-d-4bit": (
"Molmo-7B-D (Best 7B)",
"4-bit Quantized (15GB+ Required)",
),
"molmo-7b-o-4bit": (
"Molmo-7B-O (Alternative 7B)",
"4-bit Quantized (15GB+ Required)",
),
}[backend]
predictor = MolmoPredictor(model_name, memory_mode, True)
try:
response = predictor.generate(
tensor_batch_to_pil(image)[0], prompt, 48, 0.0, 0.9, 20
)
return {"response": response}
finally:
predictor.close()
if backend == "kosmos2":
from ComfyUI_VLM_nodes.nodes.kosmos2 import KosmosModelPredictor
predictor = KosmosModelPredictor()
try:
return {"response": predictor.generate(image, prompt, 48)}
finally:
predictor.close()
if backend == "uform":
from ComfyUI_VLM_nodes.nodes.uform import UformGen2QwenChat
predictor = UformGen2QwenChat()
try:
return {"response": predictor.chat(image, prompt, 48)}
finally:
predictor.close()
if backend == "mcllava":
from ComfyUI_VLM_nodes.nodes.mcllava import MCLLaVAModelPredictor
predictor = MCLLaVAModelPredictor()
try:
return {
"response": predictor.generate(
image, prompt, 0.0, 0.9, 4, 728, 48
)
}
finally:
predictor.close()
if backend == "joytag":
from ComfyUI_VLM_nodes.nodes.joytag import JoyTagPredictor
predictor = JoyTagPredictor()
try:
return {"response": predictor.predict(image, 10, 0.1)}
finally:
predictor.close()
if backend == "paligemma-caption":
from ComfyUI_VLM_nodes.nodes.paligemma import (
PALIGEMMA_MODELS,
PaliPredictor,
)
from ComfyUI_VLM_nodes.nodes.runtime import tensor_batch_to_pil
predictor = PaliPredictor(PALIGEMMA_MODELS[0], "bfloat16", "None")
try:
return {
"response": predictor.generate(
tensor_batch_to_pil(image)[0],
"caption en",
max_new_tokens=64,
do_sample=False,
)
}
finally:
predictor.close()
if backend == "minicpm-gguf-q4":
from ComfyUI_VLM_nodes.nodes.minicpm import MiniCPMPredictor
predictor = MiniCPMPredictor("Q4_K_M (4.7GB)", 4096, -1, 8)
try:
return {
"response": predictor.generate(
image, prompt, 0.0, 0.9, 40, 1.05, 48
)
}
finally:
predictor.close()
if backend == "audioldm2":
from ComfyUI_VLM_nodes.nodes.audioldm2 import AudioLDM2Predictor
predictor = AudioLDM2Predictor(cpu_offload=True)
try:
audio, sample_rate = predictor.generate(
"a short clean bell chime",
"",
1.0,
2.5,
123,
1,
2,
)
return {
"response": f"audio {audio.shape}",
"sample_rate": sample_rate,
"finite": bool(torch.isfinite(torch.from_numpy(audio)).all()),
}
finally:
predictor.close()
raise AssertionError(f"Unhandled backend: {backend}")
def main() -> int:
parser = argparse.ArgumentParser()
parser.add_argument("--backend", required=True, choices=BACKENDS)
args = parser.parse_args()
started = time.perf_counter()
if torch.cuda.is_available():
torch.cuda.reset_peak_memory_stats()
free_before, total = torch.cuda.mem_get_info()
else:
free_before = total = 0
result = _run(args.backend)
response = str(result.get("response", ""))
if not response.strip():
raise RuntimeError("The model returned an empty response.")
expected = "green" if args.backend.endswith("-video") else "red"
if args.backend != "audioldm2" and expected not in response.lower():
raise RuntimeError(
f"The model did not identify the {expected} test object: {response}"
)
if args.backend == "audioldm2" and not result["finite"]:
raise RuntimeError("AudioLDM2 returned non-finite samples.")
if torch.cuda.is_available():
peak = torch.cuda.max_memory_allocated()
free_after, _ = torch.cuda.mem_get_info()
else:
peak = free_after = 0
result.update(
backend=args.backend,
seconds=round(time.perf_counter() - started, 2),
cuda_total_gib=round(total / 1024**3, 2),
cuda_free_before_gib=round(free_before / 1024**3, 2),
cuda_free_after_gib=round(free_after / 1024**3, 2),
cuda_peak_allocated_gib=round(peak / 1024**3, 2),
)
print("SPECIALIZED_SMOKE_JSON=" + json.dumps(result, ensure_ascii=False, default=str))
return 0
if __name__ == "__main__":
raise SystemExit(main())
-176
View File
@@ -1,176 +0,0 @@
"""Run adaptive temporal reasoning on a real local video and real VLM.
Example:
python tests/manual_video_intelligence_smoke.py \
/mnt/d/002.mp4 \
--model "Qwen 3 VL 2B Instruct" \
--output /mnt/d/comfyui-repair/video-intelligence-audit/result.json
"""
from __future__ import annotations
import argparse
import importlib.util
import json
import sys
import time
from pathlib import Path
import av
import torch
REPOSITORY = Path(__file__).resolve().parents[1]
if REPOSITORY.name != "ComfyUI_VLM_nodes":
specification = importlib.util.spec_from_file_location(
"ComfyUI_VLM_nodes",
REPOSITORY / "__init__.py",
submodule_search_locations=[str(REPOSITORY)],
)
if specification is None or specification.loader is None:
raise RuntimeError(f"Could not load package from {REPOSITORY}.")
package = importlib.util.module_from_spec(specification)
sys.modules["ComfyUI_VLM_nodes"] = package
specification.loader.exec_module(package)
from ComfyUI_VLM_nodes.nodes.modern_vlm import ModernVLMPredictor
from ComfyUI_VLM_nodes.nodes.video_intelligence import (
build_video_reasoning_prompt,
parse_video_reasoning_output,
resize_video_for_analysis,
sample_video_frames,
)
def load_video(path: Path) -> tuple[torch.Tensor, float]:
container = av.open(str(path))
try:
stream = container.streams.video[0]
rate = stream.average_rate or stream.guessed_rate
if rate is None:
raise RuntimeError("The video does not report a frame rate.")
frames = [
torch.from_numpy(frame.to_ndarray(format="rgb24")).to(torch.float32)
/ 255.0
for frame in container.decode(stream)
]
finally:
container.close()
if not frames:
raise RuntimeError("The video contains no decodable frames.")
return torch.stack(frames), float(rate)
def main() -> int:
parser = argparse.ArgumentParser()
parser.add_argument("video", type=Path)
parser.add_argument(
"--model",
default="Qwen 3 VL 2B Instruct",
)
parser.add_argument("--max-frames", type=int, default=12)
parser.add_argument("--analysis-max-side", type=int, default=448)
parser.add_argument("--max-new-tokens", type=int, default=512)
parser.add_argument("--output", type=Path)
args = parser.parse_args()
frames, fps = load_video(args.video)
sampled, selection, diagnostics = sample_video_frames(
frames,
fps=fps,
max_frames=args.max_frames,
strategy="Hybrid: scene + motion + tracks",
minimum_gap_seconds=0.2,
)
prompt = build_video_reasoning_prompt(
selection,
task="Detailed temporal summary",
question="What happens, and how do the people behave over time?",
max_events=12,
)
analysis_frames = resize_video_for_analysis(
sampled,
max_side=args.analysis_max_side,
)
predictor = ModernVLMPredictor(
args.model,
"",
"ComfyUI managed (BF16)",
"Auto (SDPA)",
)
started = time.perf_counter()
try:
raw = predictor.generate(
images=None,
prompt=prompt,
system_prompt=(
"You are a precise temporal video analyst. Return one JSON "
"object that obeys the supplied schema."
),
max_new_tokens=args.max_new_tokens,
temperature=0.0,
top_p=1.0,
video_frames=analysis_frames,
fps=fps,
video_selection=selection,
)
finally:
predictor.close()
reasoning_seconds = time.perf_counter() - started
result = {
"video": str(args.video),
"model": args.model,
"source_shape": list(frames.shape),
"fps": fps,
"selection": selection.to_dict(),
"sampling": diagnostics,
"analysis_shape": list(analysis_frames.shape),
"reasoning_seconds": reasoning_seconds,
"raw_response": raw,
"cuda_peak_gib": (
torch.cuda.max_memory_allocated() / 2**30
if torch.cuda.is_available()
else 0.0
),
}
try:
summary, events, normalized = parse_video_reasoning_output(raw, selection)
except (TypeError, ValueError) as exc:
result["structured_output_valid"] = False
result["structured_output_error"] = str(exc)
if args.output:
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(
json.dumps(
result,
ensure_ascii=False,
allow_nan=False,
indent=2,
sort_keys=True,
),
encoding="utf-8",
)
raise
result.update(
{
"structured_output_valid": True,
"summary": summary,
"events": events.to_dict(),
"normalized_response": json.loads(normalized),
}
)
encoded = json.dumps(
result,
ensure_ascii=False,
allow_nan=False,
indent=2,
sort_keys=True,
)
if args.output:
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(encoded, encoding="utf-8")
print(encoded)
return 0
if __name__ == "__main__":
raise SystemExit(main())
-111
View File
@@ -1,111 +0,0 @@
import json
import threading
import time
import pytest
import torch
from ComfyUI_VLM_nodes.nodes.acceleration import (
VLMImagePixelBudget,
VLMPerformanceProfile,
optimize_image_pixels,
)
from ComfyUI_VLM_nodes.nodes.runtime import (
CachedModelNode,
tensor_batch_to_pil,
tensor_to_pil,
)
def test_batch_conversion_matches_single_frame_contract():
images = torch.tensor(
[
[
[[float("nan"), 0.5, 2.0], [-1.0, 0.25, 1.0]],
[[0.0, 0.0, 0.0], [1.0, 1.0, 1.0]],
],
[
[[255.0, 128.0, 0.0], [0.0, 64.0, 255.0]],
[[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]],
],
]
)
batch = tensor_batch_to_pil(images)
assert len(batch) == 2
for index, converted in enumerate(batch):
assert converted.mode == "RGB"
assert converted.size == (2, 2)
assert converted.tobytes() == tensor_to_pil(images, index).tobytes()
with pytest.raises(IndexError, match="only has batch index 0"):
tensor_to_pil(images[0], 1)
def test_pixel_budget_preserves_aspect_and_patch_multiple():
images = torch.rand((3, 1080, 1920, 3), dtype=torch.float32)
output, report = optimize_image_pixels(
images,
max_megapixels=0.5,
max_edge=1024,
multiple=14,
resize_quality="Fast (area)",
)
assert output.ndim == 4
assert output.shape[0] == 3
assert output.shape[1] % 14 == 0
assert output.shape[2] % 14 == 0
assert output.shape[1] * output.shape[2] <= 500_000
assert output.shape[2] <= 1024
assert report["visual_work_reduction"] > 4
assert output.shape[2] / output.shape[1] == pytest.approx(16 / 9, rel=0.03)
def test_pixel_budget_never_upscales():
image = torch.rand((240, 320, 3), dtype=torch.float32)
output, report = optimize_image_pixels(
image,
max_megapixels=2.0,
max_edge=2048,
multiple=1,
resize_quality="Quality (bicubic)",
)
assert output is image
assert report["resized"] is False
def test_performance_nodes_return_standard_comfy_values():
profile = VLMPerformanceProfile().profile("Live / robotics")
assert profile[:5] == (24, 0.5, 896, 8, False)
assert json.loads(profile[5])["profile"] == "Live / robotics"
optimized = VLMImagePixelBudget().optimize(
torch.rand((1, 1000, 1600, 3)),
0.5,
1024,
"14",
"Fast (area)",
)
assert optimized[1] % 14 == 0
assert optimized[2] % 14 == 0
def test_cached_model_node_prevents_duplicate_concurrent_loads():
node = CachedModelNode()
factory_calls = []
handles = []
def factory():
factory_calls.append(1)
time.sleep(0.02)
return object()
def load():
handles.append(node.get_or_create_model("same-model", factory))
threads = [threading.Thread(target=load) for _ in range(8)]
for thread in threads:
thread.start()
for thread in threads:
thread.join()
assert len(factory_calls) == 1
assert len({id(handle) for handle in handles}) == 1
-110
View File
@@ -1,110 +0,0 @@
"""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
-250
View File
@@ -1,250 +0,0 @@
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]
-166
View File
@@ -1,166 +0,0 @@
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),)
-141
View File
@@ -1,141 +0,0 @@
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"
-904
View File
@@ -1,904 +0,0 @@
from __future__ import annotations
import json
from pathlib import Path
from types import SimpleNamespace
import pytest
import torch
from ComfyUI_VLM_nodes.nodes import hosted_api
class FakeHttpClient:
def __init__(self, **kwargs):
self.kwargs = kwargs
class FakeResponses:
def __init__(self, calls, failure=None, response_text="secure response"):
self.calls = calls
self.failure = failure
self.response_text = response_text
def create(self, **kwargs):
self.calls.append(("responses", kwargs))
if self.failure is not None:
raise self.failure
return SimpleNamespace(output_text=self.response_text)
class FakeChatCompletions:
def __init__(self, calls, failure=None, response_text="secure response"):
self.calls = calls
self.failure = failure
self.response_text = response_text
def create(self, **kwargs):
self.calls.append(("chat", kwargs))
if self.failure is not None:
raise self.failure
message = SimpleNamespace(content=self.response_text)
return SimpleNamespace(choices=[SimpleNamespace(message=message)])
def fake_openai_module(calls, failure=None, response_text="secure response"):
class FakeOpenAI:
def __init__(self, **kwargs):
calls.append(("client", kwargs))
self.responses = FakeResponses(
calls,
failure=failure,
response_text=response_text,
)
self.chat = SimpleNamespace(
completions=FakeChatCompletions(
calls,
failure=failure,
response_text=response_text,
)
)
def close(self):
calls.append(("close", {}))
return SimpleNamespace(
OpenAI=FakeOpenAI,
DefaultHttpxClient=FakeHttpClient,
)
def test_api_schemas_never_accept_plaintext_keys():
for node_class in (hosted_api.PromptGenerateAPI, hosted_api.HostedVLMAPI):
schema = node_class.INPUT_TYPES()
all_inputs = {
**schema.get("required", {}),
**schema.get("optional", {}),
**schema.get("hidden", {}),
}
assert "api_key" not in all_inputs
assert "credential_source" in all_inputs
assert "STRING" not in repr(all_inputs["credential_source"][0])
assert "web_search" in all_inputs
assert "output_format" in all_inputs
assert "json_schema" in all_inputs
assert "schema_api_style" in all_inputs
def test_json_schema_parser_blocks_remote_refs_and_bounds_input():
for keyword in ("$ref", "$dynamicRef", "$recursiveRef"):
with pytest.raises(ValueError, match="only local fragment"):
hosted_api.parse_json_schema(
"JSON Schema",
json.dumps(
{
"type": "object",
"properties": {
"payload": {
keyword: "https://attacker.example/schema.json"
}
},
}
),
)
with pytest.raises(ValueError, match="64,000"):
hosted_api.parse_json_schema("JSON Schema", "x" * 64_001)
def test_local_structured_output_validation_is_strict_and_normalized():
schema_text = json.dumps(
{
"type": "object",
"properties": {"count": {"type": "integer"}},
"required": ["count"],
"additionalProperties": False,
}
)
schema = hosted_api.parse_json_schema("JSON Schema", schema_text)
assert hosted_api.validate_structured_output(
'```json\n{"count": 2}\n```',
"JSON Schema",
schema,
) == '{\n "count": 2\n}'
with pytest.raises(RuntimeError, match=r"\$\.count \(type constraint\)"):
hosted_api.validate_structured_output(
'{"count": "two"}',
"JSON Schema",
schema,
)
with pytest.raises(RuntimeError, match="valid JSON"):
hosted_api.validate_structured_output(
'{"count":',
"JSON Schema",
schema,
)
def test_provider_catalog_uses_current_bound_credentials_and_endpoints():
assert len(hosted_api.PROVIDER_PROFILES) >= 18
expected = {
"OpenAI": "OPENAI_API_KEY",
"Google Gemini": "GEMINI_API_KEY",
"Anthropic": "ANTHROPIC_API_KEY",
"xAI": "XAI_API_KEY",
"DeepSeek": "DEEPSEEK_API_KEY",
"Groq": "GROQ_API_KEY",
"Mistral": "MISTRAL_API_KEY",
"Together AI": "TOGETHER_API_KEY",
"OpenRouter": "OPENROUTER_API_KEY",
"Custom / Local": "CUSTOM_API_KEY",
}
providers = {
profile.provider: profile.api_key_env
for profile in hosted_api.PROVIDER_PROFILES.values()
}
assert expected.items() <= providers.items()
for profile in hosted_api.PROVIDER_PROFILES.values():
if profile.base_url is not None:
assert profile.base_url.startswith("https://")
@pytest.mark.parametrize(
"url",
[
"http://example.com/v1",
"ftp://127.0.0.1/v1",
"https://user:secret@example.com/v1",
"https://example.com/v1?api_key=secret",
"not-a-url",
],
)
def test_custom_endpoint_rejects_unsafe_urls(url):
with pytest.raises(ValueError):
hosted_api.validate_custom_base_url(url)
@pytest.mark.parametrize(
"url",
[
"http://127.0.0.1:8000/v1",
"http://[::1]:11434/v1",
"http://localhost:1234/v1",
"https://example.com/v1/",
],
)
def test_custom_endpoint_accepts_https_or_loopback(url):
normalized, loopback = hosted_api.validate_custom_base_url(url)
assert normalized.startswith(("http://", "https://"))
assert loopback is (url.startswith("http://"))
def test_built_in_key_cannot_be_redirected(monkeypatch):
profile = hosted_api.provider_profile("OpenAI — GPT-5.6 Terra")
monkeypatch.setenv("OPENAI_API_KEY", "sk-real-secret-value")
with pytest.raises(ValueError, match="pinned to official hosts"):
hosted_api.resolve_endpoint(profile, "https://attacker.example/v1")
def test_legacy_plaintext_value_is_rejected_without_echo(monkeypatch):
profile = hosted_api.provider_profile("OpenAI — GPT-5.6 Terra")
secret = "sk-legacy-plaintext-that-must-not-appear"
with pytest.raises(ValueError) as captured:
hosted_api.resolve_api_key(profile, secret, loopback=False)
assert secret not in str(captured.value)
assert "legacy plaintext API key was removed" in str(captured.value)
def test_redaction_removes_exact_encoded_and_header_credentials():
secret = "sk-ant-example-SECRET_123456789"
message = (
f"Authorization: Bearer {secret}; api_key={secret}; "
f"url=https://user:{secret}@example.com; encoded={secret}"
)
redacted = hosted_api.redact_sensitive(message, (secret,))
assert secret not in redacted
assert "Bearer" not in redacted
assert "[REDACTED]" in redacted
def test_responses_call_is_stateless_private_and_provider_bound(monkeypatch):
calls = []
secret = "sk-openai-provider-bound-secret"
monkeypatch.setenv("OPENAI_API_KEY", secret)
monkeypatch.setattr(
hosted_api,
"require_module",
lambda *_args: fake_openai_module(calls),
)
node = hosted_api.PromptGenerateAPI()
assert not hasattr(node, "session_history")
result = node.generate_prompt(
"OpenAI — GPT-5.6 Terra",
False,
hosted_api.PROVIDER_CREDENTIAL,
"A scene",
"Improve it",
0,
0,
stream_output=False,
)
assert result == ("secure response",)
client_kwargs = next(payload for kind, payload in calls if kind == "client")
assert client_kwargs["api_key"] == secret
assert "base_url" not in client_kwargs
assert client_kwargs["http_client"].kwargs["follow_redirects"] is False
assert client_kwargs["http_client"].kwargs["trust_env"] is False
request = next(payload for kind, payload in calls if kind == "responses")
assert request["model"] == "gpt-5.6-terra"
assert request["store"] is False
assert "previous_response_id" not in request
assert "metadata" not in request
def test_openai_combines_web_search_structured_output_and_stream_contract(
monkeypatch,
):
calls = []
monkeypatch.setenv("OPENAI_API_KEY", "sk-test-structured-search")
monkeypatch.setattr(
hosted_api,
"require_module",
lambda import_name, *_args: (
fake_openai_module(calls, response_text='{"answer":"grounded"}')
if import_name == "openai"
else __import__(import_name)
),
)
schema = json.dumps(
{
"type": "object",
"properties": {"answer": {"type": "string"}},
"required": ["answer"],
"additionalProperties": False,
}
)
result = hosted_api.PromptGenerateAPI().generate_prompt(
"OpenAI — GPT-5.6 Sol",
False,
hosted_api.PROVIDER_CREDENTIAL,
"Find a current fact",
"",
0,
0,
web_search=True,
output_format="JSON Schema",
json_schema=schema,
stream_output=False,
)
assert result == ('{\n "answer": "grounded"\n}',)
request = next(payload for kind, payload in calls if kind == "responses")
assert request["tools"] == [{"type": "web_search"}]
assert request["text"]["format"]["type"] == "json_schema"
assert request["text"]["format"]["strict"] is True
assert request["text"]["format"]["schema"]["required"] == ["answer"]
assert "JSON Schema:" in request["instructions"]
def test_unsupported_web_search_fails_before_network(monkeypatch):
monkeypatch.setenv("DEEPSEEK_API_KEY", "deepseek-test-secret")
with pytest.raises(ValueError, match="does not expose native web search"):
hosted_api.PromptGenerateAPI().generate_prompt(
"DeepSeek — V4 Flash",
False,
hosted_api.PROVIDER_CREDENTIAL,
"Search now",
"",
0,
0,
web_search=True,
stream_output=False,
)
def test_responses_stream_collects_deltas_and_closes():
class Stream(list):
closed = False
def close(self):
self.closed = True
stream = Stream(
[
SimpleNamespace(type="response.created"),
SimpleNamespace(type="response.output_text.delta", delta="hello "),
SimpleNamespace(type="response.output_text.delta", delta="world"),
]
)
client = SimpleNamespace(
responses=SimpleNamespace(create=lambda **kwargs: stream)
)
assert hosted_api._stream_responses(client, {"model": "test"}, None) == (
"hello world"
)
assert stream.closed is True
def test_chat_stream_collects_deltas_and_closes():
class Stream(list):
closed = False
def close(self):
self.closed = True
stream = Stream(
[
SimpleNamespace(
choices=[
SimpleNamespace(delta=SimpleNamespace(content="frame "))
]
),
SimpleNamespace(
choices=[
SimpleNamespace(delta=SimpleNamespace(content="ready"))
]
),
]
)
client = SimpleNamespace(
chat=SimpleNamespace(
completions=SimpleNamespace(create=lambda **kwargs: stream)
)
)
assert hosted_api._stream_chat(client, {"model": "test"}, None) == (
"frame ready"
)
assert stream.closed is True
def test_provider_failure_never_echoes_api_key(monkeypatch):
calls = []
secret = "sk-secret-reflected-by-provider-123456"
failure = RuntimeError(f"Authorization: Bearer {secret} api_key={secret}")
monkeypatch.setenv("OPENAI_API_KEY", secret)
monkeypatch.setattr(
hosted_api,
"require_module",
lambda *_args: fake_openai_module(calls, failure=failure),
)
with pytest.raises(RuntimeError) as captured:
hosted_api.PromptGenerateAPI().generate_prompt(
"OpenAI — GPT-5.6 Terra",
False,
hosted_api.PROVIDER_CREDENTIAL,
"hello",
"",
0,
0,
stream_output=False,
)
assert secret not in str(captured.value)
assert "[REDACTED]" in str(captured.value)
def test_anthropic_uses_native_messages_and_keeps_key_out_of_body(monkeypatch):
calls = []
secret = "sk-ant-native-secret"
class FakeResponse:
def raise_for_status(self):
return None
def json(self):
return {
"content": [
{"type": "text", "text": "native Anthropic response"}
]
}
class FakeClient:
def __init__(self, **kwargs):
calls.append(("client", kwargs))
def post(self, url, **kwargs):
calls.append(("post", {"url": url, **kwargs}))
return FakeResponse()
def close(self):
calls.append(("close", {}))
monkeypatch.setenv("ANTHROPIC_API_KEY", secret)
monkeypatch.setattr(
hosted_api,
"require_module",
lambda import_name, *_args: (
SimpleNamespace(Client=FakeClient)
if import_name == "httpx"
else pytest.fail(f"Unexpected module request: {import_name}")
),
)
result = hosted_api.HostedVLMAPI().analyze(
"Anthropic — Claude Sonnet 5",
hosted_api.PROVIDER_CREDENTIAL,
"Read this image.",
"Be concise.",
1,
512,
80,
"auto",
images=torch.rand((1, 48, 64, 3)),
stream_output=False,
)
assert result == (
"native Anthropic response",
"claude-sonnet-5",
1,
)
request = next(payload for kind, payload in calls if kind == "post")
assert request["url"] == "https://api.anthropic.com/v1/messages"
assert request["headers"]["x-api-key"] == secret
assert secret not in repr(request["json"])
content = request["json"]["messages"][0]["content"]
assert content[1]["type"] == "image"
assert content[1]["source"]["type"] == "base64"
assert request["json"]["stream"] is False
def test_anthropic_native_stream_collects_text_deltas(monkeypatch):
class FakeStreamResponse:
def __enter__(self):
return self
def __exit__(self, *_args):
return False
def raise_for_status(self):
return None
def iter_lines(self):
yield 'event: content_block_delta'
yield (
'data: {"type":"content_block_delta","delta":'
'{"type":"text_delta","text":"hello "}}'
)
yield (
'data: {"type":"content_block_delta","delta":'
'{"type":"text_delta","text":"world"}}'
)
class FakeClient:
def __init__(self, **_kwargs):
pass
def stream(self, *_args, **_kwargs):
return FakeStreamResponse()
def close(self):
return None
monkeypatch.setattr(
hosted_api,
"require_module",
lambda import_name, *_args: (
SimpleNamespace(Client=FakeClient)
if import_name == "httpx"
else __import__(import_name)
),
)
profile = hosted_api.provider_profile("Anthropic — Claude Sonnet 5")
result = hosted_api._call_anthropic_api(
profile=profile,
model=profile.model,
endpoint=profile.base_url,
api_key="sk-ant-stream",
system_prompt="Be concise.",
prompt="Hello",
image_data=[],
timeout_seconds=30,
max_output_tokens=100,
stream_output=True,
use_system_proxy=False,
unique_id=None,
web_search=False,
output_format="Text",
output_schema=None,
)
assert result == "hello world"
def test_anthropic_native_search_and_structured_contracts(monkeypatch):
calls = []
class FakeResponse:
def raise_for_status(self):
return None
def json(self):
return {"content": [{"type": "text", "text": '{"answer":"yes"}'}]}
class FakeClient:
def __init__(self, **kwargs):
calls.append(("client", kwargs))
def post(self, url, **kwargs):
calls.append(("post", {"url": url, **kwargs}))
return FakeResponse()
def close(self):
return None
monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-ant-contract")
monkeypatch.setattr(
hosted_api,
"require_module",
lambda import_name, *_args: (
SimpleNamespace(Client=FakeClient)
if import_name == "httpx"
else __import__(import_name)
),
)
schema = (
'{"type":"object","properties":{"answer":{"type":"string"}},'
'"required":["answer"],"additionalProperties":false}'
)
result = hosted_api.PromptGenerateAPI().generate_prompt(
"Anthropic — Claude Sonnet 5",
False,
hosted_api.PROVIDER_CREDENTIAL,
"Return a value",
"",
0,
0,
output_format="JSON Schema",
json_schema=schema,
stream_output=False,
)
assert json.loads(result[0]) == {"answer": "yes"}
structured = next(payload for kind, payload in calls if kind == "post")
assert structured["json"]["output_config"]["format"]["type"] == "json_schema"
calls.clear()
hosted_api.PromptGenerateAPI().generate_prompt(
"Anthropic — Claude Sonnet 5",
False,
hosted_api.PROVIDER_CREDENTIAL,
"Search the web",
"",
0,
0,
web_search=True,
output_format="JSON Schema",
json_schema=schema,
stream_output=False,
)
searched = next(payload for kind, payload in calls if kind == "post")
assert searched["json"]["tools"][0]["type"] == "web_search_20260318"
assert searched["json"]["tools"][0]["allowed_callers"] == ["direct"]
assert "output_config" not in searched["json"]
def test_gemini_native_search_vision_and_schema_contract(monkeypatch):
calls = []
secret = "gemini-provider-secret"
class FakeResponse:
def raise_for_status(self):
return None
def json(self):
return {
"candidates": [
{
"content": {
"parts": [{"text": '{"objects":["tree"]}'}]
}
}
]
}
class FakeClient:
def __init__(self, **kwargs):
calls.append(("client", kwargs))
def post(self, url, **kwargs):
calls.append(("post", {"url": url, **kwargs}))
return FakeResponse()
def close(self):
return None
monkeypatch.setenv("GEMINI_API_KEY", secret)
monkeypatch.setattr(
hosted_api,
"require_module",
lambda import_name, *_args: (
SimpleNamespace(Client=FakeClient)
if import_name == "httpx"
else __import__(import_name)
),
)
schema = (
'{"type":"object","properties":{"objects":{"type":"array",'
'"items":{"type":"string"}}},"required":["objects"]}'
)
result = hosted_api.HostedVLMAPI().analyze(
"Google — Gemini 3.6 Flash",
hosted_api.PROVIDER_CREDENTIAL,
"Identify objects using current context.",
"Be precise.",
1,
512,
80,
"auto",
images=torch.rand((1, 32, 48, 3)),
web_search=True,
output_format="JSON Schema",
json_schema=schema,
stream_output=False,
)
assert json.loads(result[0]) == {"objects": ["tree"]}
request = next(payload for kind, payload in calls if kind == "post")
assert request["url"].endswith(
"/models/gemini-3.6-flash:generateContent"
)
assert request["headers"]["x-goog-api-key"] == secret
assert secret not in repr(request["json"])
assert request["json"]["tools"] == [{"google_search": {}}]
assert (
request["json"]["generationConfig"]["responseFormat"]["text"]["schema"][
"required"
]
== ["objects"]
)
inline = request["json"]["contents"][0]["parts"][1]["inlineData"]
assert inline["mimeType"] == "image/jpeg"
assert inline["data"]
def test_gemini_native_stream_collects_sse_deltas(monkeypatch):
class FakeStreamResponse:
def __enter__(self):
return self
def __exit__(self, *_args):
return False
def raise_for_status(self):
return None
def iter_lines(self):
yield (
'data: {"candidates":[{"content":{"parts":'
'[{"text":"frame "}]}}]}'
)
yield (
'data: {"candidates":[{"content":{"parts":'
'[{"text":"ready"}]}}]}'
)
class FakeClient:
def __init__(self, **_kwargs):
pass
def stream(self, *_args, **_kwargs):
return FakeStreamResponse()
def close(self):
return None
monkeypatch.setattr(
hosted_api,
"require_module",
lambda import_name, *_args: (
SimpleNamespace(Client=FakeClient)
if import_name == "httpx"
else __import__(import_name)
),
)
profile = hosted_api.provider_profile("Google — Gemini 3.6 Flash")
result = hosted_api._call_gemini_api(
profile=profile,
model=profile.model,
api_key="gemini-stream",
system_prompt="Be concise.",
prompt="Describe.",
image_data=[],
timeout_seconds=30,
max_output_tokens=100,
stream_output=True,
use_system_proxy=False,
unique_id=None,
web_search=True,
output_format="Text",
output_schema=None,
)
assert result == "frame ready"
def test_vlm_uniformly_samples_and_bounds_image_batch(monkeypatch):
calls = []
monkeypatch.setenv("OPENAI_API_KEY", "sk-test-only")
monkeypatch.setattr(
hosted_api,
"require_module",
lambda *_args: fake_openai_module(calls),
)
images = torch.rand((10, 96, 128, 3), dtype=torch.float32)
result = hosted_api.HostedVLMAPI().analyze(
"OpenAI — GPT-5.6 Terra",
hosted_api.PROVIDER_CREDENTIAL,
"Compare the sampled frames.",
"Be precise.",
4,
768,
82,
"low",
images=images,
stream_output=False,
)
assert result == ("secure response", "gpt-5.6-terra", 4)
request = next(payload for kind, payload in calls if kind == "responses")
content = request["input"][0]["content"]
image_parts = [part for part in content if part["type"] == "input_image"]
assert len(image_parts) == 4
assert all(part["image_url"].startswith("data:image/jpeg;base64,") for part in image_parts)
assert all(part["detail"] == "low" for part in image_parts)
def test_open_source_vlm_llama_cpp_schema_dialect_and_local_validation(
monkeypatch,
):
calls = []
monkeypatch.setattr(
hosted_api,
"require_module",
lambda import_name, *_args: (
fake_openai_module(calls, response_text='{"objects":["cat"]}')
if import_name == "openai"
else __import__(import_name)
),
)
schema = json.dumps(
{
"type": "object",
"properties": {
"objects": {
"type": "array",
"items": {"type": "string"},
}
},
"required": ["objects"],
"additionalProperties": False,
}
)
result = hosted_api.HostedVLMAPI().analyze(
"Custom / Local — OpenAI compatible",
hosted_api.LOCAL_NO_KEY,
"List visible objects.",
"Be precise.",
1,
512,
80,
"auto",
images=torch.rand((1, 32, 48, 3)),
base_url="http://127.0.0.1:8080/v1",
model_override="local-vlm",
output_format="JSON Schema",
json_schema=schema,
schema_api_style="llama.cpp JSON Schema",
stream_output=False,
)
assert result == (
'{\n "objects": [\n "cat"\n ]\n}',
"local-vlm",
1,
)
request = next(payload for kind, payload in calls if kind == "chat")
assert request["response_format"] == {
"type": "json_schema",
"schema": json.loads(schema),
}
image = request["messages"][1]["content"][1]
assert image["type"] == "image_url"
assert image["image_url"]["url"].startswith("data:image/jpeg;base64,")
def test_custom_openai_schema_style_uses_standard_wrapper(monkeypatch):
calls = []
monkeypatch.setattr(
hosted_api,
"require_module",
lambda import_name, *_args: (
fake_openai_module(calls, response_text='{"ok":true}')
if import_name == "openai"
else __import__(import_name)
),
)
schema = '{"type":"object","properties":{"ok":{"type":"boolean"}},"required":["ok"]}'
result, _, _ = hosted_api.execute_hosted(
model_name="Custom / Local — OpenAI compatible",
credential_source=hosted_api.LOCAL_NO_KEY,
prompt="Return status.",
system_prompt="Be exact.",
base_url="http://localhost:8000/v1",
model_override="local",
api_mode="Chat Completions",
timeout_seconds=30,
max_output_tokens=100,
reasoning_effort="none",
seed=0,
stream_output=False,
use_system_proxy=False,
unique_id=None,
output_format="JSON Schema",
json_schema=schema,
)
assert json.loads(result) == {"ok": True}
request = next(payload for kind, payload in calls if kind == "chat")
assert request["response_format"]["json_schema"]["strict"] is True
assert request["response_format"]["json_schema"]["schema"]["required"] == [
"ok"
]
def test_groq_auto_uses_documented_chat_route_for_structured_output(monkeypatch):
calls = []
monkeypatch.setenv("GROQ_API_KEY", "gsk-test-structured")
monkeypatch.setattr(
hosted_api,
"require_module",
lambda import_name, *_args: (
fake_openai_module(calls, response_text='{"ok":true}')
if import_name == "openai"
else __import__(import_name)
),
)
hosted_api.execute_hosted(
model_name="Groq — GPT-OSS 20B",
credential_source=hosted_api.PROVIDER_CREDENTIAL,
prompt="Return status.",
system_prompt="Be exact.",
base_url="",
model_override="",
api_mode="Auto",
timeout_seconds=30,
max_output_tokens=100,
reasoning_effort="none",
seed=0,
stream_output=False,
use_system_proxy=False,
unique_id=None,
output_format="JSON Schema",
json_schema=(
'{"type":"object","properties":{"ok":{"type":"boolean"}},'
'"required":["ok"],"additionalProperties":false}'
),
)
assert any(kind == "chat" for kind, _payload in calls)
assert not any(kind == "responses" for kind, _payload in calls)
def test_frontend_scrubs_legacy_key_before_graph_configuration():
web_root = Path(__file__).resolve().parents[1] / "web" / "js"
source = (
web_root / "apiSecurity.js"
).read_text("utf-8")
assert "beforeConfigureGraph" in source
assert "delete values.api_key" in source
assert "CREDENTIAL_WIDGET_INDEX = 2" in source
view_text = (web_root / "viewText.js").read_text("utf-8")
assert '"PromptGenerateAPI"' in view_text
assert '"HostedVLMAPI"' in view_text
-496
View File
@@ -1,496 +0,0 @@
"""Contract tests for the llama.cpp multimodal nodes in ``nodes/llavaloader.py``.
Covers batch handling, the vision message envelope, projector wiring, and the
cached-handle lifecycle. No llama.cpp wheel, mmproj, or GGUF weights required.
"""
from __future__ import annotations
import base64
from pathlib import Path
import pytest
import torch
from ComfyUI_VLM_nodes.nodes import llavaloader
from ComfyUI_VLM_nodes.nodes.runtime import LlamaHandle, LlavaClipConfig
MODEL_FILE = "llava.gguf"
CLIP_FILE = "mmproj.gguf"
class FakeLlama:
def __init__(self, contents: list[str] | None = None):
self.contents = contents or ["a description"]
self.calls: list[dict] = []
def create_chat_completion(self, **kwargs):
index = min(len(self.calls), len(self.contents) - 1)
self.calls.append(kwargs)
return {"choices": [{"message": {"content": self.contents[index]}}]}
class FakeHandle:
instances: list[FakeHandle] = []
def __init__(self, model_path, **kwargs):
self.model_path = model_path
self.kwargs = kwargs
self.closed = False
self.llama = FakeLlama()
FakeHandle.instances.append(self)
def ensure_loaded(self):
return self.llama
def close(self):
self.closed = True
@pytest.fixture
def resolved_paths(monkeypatch):
root = Path("/models/LLavacheckpoints")
monkeypatch.setattr(llavaloader, "resolve_model_path", lambda name: root / name)
return root
@pytest.fixture
def fake_handles(monkeypatch):
FakeHandle.instances = []
monkeypatch.setattr(llavaloader, "LlamaHandle", FakeHandle)
return FakeHandle
def image_batch(count: int = 1, size: int = 4) -> torch.Tensor:
"""A ComfyUI BHWC float image batch."""
return torch.rand(count, size, size, 3)
# --------------------------------------------------------------------------
# Widget ordering (see issue #156).
# --------------------------------------------------------------------------
def test_llava_sampler_simple_widget_order_is_frozen():
assert list(llavaloader.LLavaSamplerSimple.INPUT_TYPES()["required"]) == [
"image",
"prompt",
"model",
"temperature",
]
def test_llava_sampler_advanced_widget_order_is_frozen():
assert list(llavaloader.LLavaSamplerAdvanced.INPUT_TYPES()["required"]) == [
"image",
"system_msg",
"prompt",
"model",
"max_tokens",
"temperature",
"top_p",
"top_k",
"frequency_penalty",
"presence_penalty",
"repeat_penalty",
"seed",
]
def test_llava_loader_widget_order_is_frozen():
schema = llavaloader.LLavaLoader.INPUT_TYPES()
assert list(schema["required"]) == [
"ckpt_name",
"max_ctx",
"gpu_layers",
"n_threads",
"clip",
]
def test_optional_memory_free_simple_widget_order_is_frozen():
schema = llavaloader.LLavaOptionalMemoryFreeSimple.INPUT_TYPES()
assert list(schema["required"]) == [
"ckpt_name",
"clip_name",
"max_ctx",
"gpu_layers",
"n_threads",
"image",
"prompt",
"temperature",
"unload",
]
assert list(schema["optional"])[0] == "handler"
def test_every_llava_node_declares_a_callable_function_and_return_types():
for name, node_class in llavaloader.NODE_CLASS_MAPPINGS.items():
assert isinstance(node_class.RETURN_TYPES, tuple), name
assert node_class.RETURN_TYPES, name
assert callable(getattr(node_class, node_class.FUNCTION, None)), name
assert node_class.CATEGORY.startswith("VLM Nodes"), name
def test_display_names_cover_every_registered_node():
assert set(llavaloader.NODE_CLASS_MAPPINGS) == set(
llavaloader.NODE_DISPLAY_NAME_MAPPINGS
)
# --------------------------------------------------------------------------
# Vision message envelope.
# --------------------------------------------------------------------------
def test_vision_messages_place_the_image_before_the_text():
messages = llavaloader._vision_messages("sys", "what is this?", "data:image/png;b")
assert messages[0] == {"role": "system", "content": "sys"}
content = messages[1]["content"]
assert messages[1]["role"] == "user"
# llama.cpp vision handlers require the image part first.
assert content[0]["type"] == "image_url"
assert content[0]["image_url"]["url"] == "data:image/png;b"
assert content[1] == {"type": "text", "text": "what is this?"}
def test_run_batch_sends_a_png_data_uri_per_image():
llama = FakeLlama()
llavaloader._run_batch(
image_batch(1), llama, system_msg="sys", prompt="p", temperature=0.1
)
(call,) = llama.calls
url = call["messages"][1]["content"][0]["image_url"]["url"]
assert url.startswith("data:image/png;base64,")
# The payload must be real decodable PNG bytes.
decoded = base64.b64decode(url.split(",", 1)[1])
assert decoded.startswith(b"\x89PNG\r\n\x1a\n")
def test_run_batch_calls_the_model_once_per_batch_item():
llama = FakeLlama(["first", "second", "third"])
text = llavaloader._run_batch(
image_batch(3), llama, system_msg="sys", prompt="p", temperature=0.1
)
assert len(llama.calls) == 3
# Every batch item must survive into the response.
assert "first" in text
assert "second" in text
assert "third" in text
assert "--- Image 1 ---" in text
assert "--- Image 3 ---" in text
def test_run_batch_returns_bare_text_for_a_single_image():
llama = FakeLlama(["only one"])
text = llavaloader._run_batch(
image_batch(1), llama, system_msg="sys", prompt="p", temperature=0.1
)
assert text == "only one"
def test_run_batch_forwards_generation_kwargs_unchanged():
llama = FakeLlama()
llavaloader._run_batch(
image_batch(1),
llama,
system_msg="sys",
prompt="p",
max_tokens=32,
temperature=0.3,
top_p=0.7,
top_k=10,
seed=99,
)
(call,) = llama.calls
assert call["max_tokens"] == 32
assert call["temperature"] == 0.3
assert call["top_p"] == 0.7
assert call["top_k"] == 10
assert call["seed"] == 99
def test_sampler_simple_returns_a_single_string_output():
llama = FakeLlama(["a cat on a mat"])
result = llavaloader.LLavaSamplerSimple().generate_text(
image=image_batch(1), prompt="describe", model=llama, temperature=0.1
)
assert result == ("a cat on a mat",)
def test_sampler_advanced_uses_the_supplied_system_message():
llama = FakeLlama()
llavaloader.LLavaSamplerAdvanced().generate_text_advanced(
image=image_batch(1),
system_msg="answer in French",
prompt="describe",
model=llama,
max_tokens=16,
temperature=0.1,
top_p=0.9,
top_k=5,
frequency_penalty=0.0,
presence_penalty=0.0,
repeat_penalty=1.0,
seed=7,
)
(call,) = llama.calls
assert call["messages"][0] == {"role": "system", "content": "answer in French"}
# --------------------------------------------------------------------------
# Projector / clip wiring.
# --------------------------------------------------------------------------
def test_clip_factory_uses_the_config_create_hook():
config = LlavaClipConfig(Path("/models/mmproj.gguf"), "LLaVA 1.6")
assert llavaloader._clip_factory(config) == config.create
def test_clip_factory_accepts_a_precreated_handler():
sentinel = object()
factory = llavaloader._clip_factory(sentinel)
# Workflows saved before handler selection passed the handler itself.
assert factory() is sentinel
def test_make_handle_derives_the_projector_from_the_clip_config(resolved_paths):
config = LlavaClipConfig(resolved_paths / CLIP_FILE, "LLaVA 1.5")
handle = llavaloader._make_handle(MODEL_FILE, 4096, -1, 4, config)
assert isinstance(handle, LlamaHandle)
assert handle.projector_path == resolved_paths / CLIP_FILE
assert handle.chat_handler_factory == config.create
assert handle.n_ctx == 4096
# Still lazy: no llama.cpp object was constructed.
assert handle._llm is None
def test_make_handle_keeps_an_explicit_projector_override(resolved_paths):
config = LlavaClipConfig(resolved_paths / CLIP_FILE, "LLaVA 1.5")
override = Path("/models/other-mmproj.gguf")
handle = llavaloader._make_handle(
MODEL_FILE,
4096,
-1,
4,
config,
runtime_options={"projector_path": override},
)
assert handle.projector_path == override
def test_llava_loader_does_not_load_weights(resolved_paths):
config = LlavaClipConfig(resolved_paths / CLIP_FILE, "LLaVA 1.5")
(handle,) = llavaloader.LLavaLoader().load_llava_checkpoint(
ckpt_name=MODEL_FILE,
max_ctx=2048,
gpu_layers=10,
n_threads=8,
clip=config,
)
assert isinstance(handle, LlamaHandle)
assert handle._llm is None
assert handle.n_gpu_layers == 10
assert handle.n_threads == 8
def test_clip_loader_returns_a_frozen_config_with_the_chosen_handler(resolved_paths):
(config,) = llavaloader.LlavaClipLoader().load_clip_checkpoint(
CLIP_FILE, handler="MiniCPM-V 2.6"
)
assert isinstance(config, LlavaClipConfig)
assert config.model_path == resolved_paths / CLIP_FILE
assert config.handler == "MiniCPM-V 2.6"
def test_clip_config_rejects_an_unknown_handler(resolved_paths):
config = LlavaClipConfig(resolved_paths / CLIP_FILE, "Not A Handler")
with pytest.raises((ValueError, RuntimeError)) as error:
config.create()
# Either an unknown-handler rejection or a missing-wheel report is correct;
# a silent fallback to the wrong prompt format is not.
assert "handler" in str(error.value).lower() or "llama" in str(error.value).lower()
def test_clip_loader_defaults_to_the_embedded_gguf_chat_template():
handler = llavaloader.LlavaClipLoader.INPUT_TYPES()["optional"]["handler"]
choices, options = handler[0], handler[1]
assert options["default"] == "Auto (GGUF chat template)"
assert options["default"] in choices
assert "LLaVA 1.5" in choices
# --------------------------------------------------------------------------
# Cached-handle lifecycle (issue #137: "model never unloads").
# --------------------------------------------------------------------------
def test_cached_llava_reuses_one_handle_for_identical_settings(
fake_handles, resolved_paths
):
node = llavaloader.LLavaOptionalMemoryFreeSimple()
first = node._model(MODEL_FILE, CLIP_FILE, 4096, -1, 4)
second = node._model(MODEL_FILE, CLIP_FILE, 4096, -1, 4)
assert first is second
assert len(fake_handles.instances) == 1
def test_cached_llava_rebuilds_when_the_projector_changes(fake_handles, resolved_paths):
node = llavaloader.LLavaOptionalMemoryFreeSimple()
first = node._model(MODEL_FILE, CLIP_FILE, 4096, -1, 4)
second = node._model(MODEL_FILE, "other-mmproj.gguf", 4096, -1, 4)
assert first is not second
assert first.closed is True
assert len(fake_handles.instances) == 2
def test_cached_llava_rebuilds_when_the_handler_changes(fake_handles, resolved_paths):
node = llavaloader.LLavaOptionalMemoryFreeSimple()
node._model(MODEL_FILE, CLIP_FILE, 4096, -1, 4, handler="LLaVA 1.5")
node._model(MODEL_FILE, CLIP_FILE, 4096, -1, 4, handler="LLaVA 1.6")
assert len(fake_handles.instances) == 2
def test_cached_llava_unload_releases_the_handle(fake_handles, resolved_paths):
node = llavaloader.LLavaOptionalMemoryFreeSimple()
handle = node._model(MODEL_FILE, CLIP_FILE, 4096, -1, 4)
node._maybe_unload(False)
assert handle.closed is False
node._maybe_unload(True)
assert handle.closed is True
assert node._handle is None
assert node._key is None
def test_cached_llava_unload_is_safe_before_any_load(fake_handles, resolved_paths):
node = llavaloader.LLavaOptionalMemoryFreeSimple()
# Must not raise when nothing was ever loaded.
node._maybe_unload(True)
assert node._handle is None
def _memory_free_kwargs(**overrides):
kwargs = {
"ckpt_name": MODEL_FILE,
"clip_name": CLIP_FILE,
"max_ctx": 4096,
"gpu_layers": -1,
"n_threads": 4,
"image": image_batch(1),
"prompt": "describe this",
"temperature": 0.1,
"unload": False,
}
kwargs.update(overrides)
return kwargs
def test_memory_free_simple_generates_through_the_cached_handle(
fake_handles, resolved_paths
):
node = llavaloader.LLavaOptionalMemoryFreeSimple()
(text,) = node.generate_text(**_memory_free_kwargs())
assert text == "a description"
(handle,) = fake_handles.instances
assert handle.closed is False
(call,) = handle.llama.calls
assert call["temperature"] == 0.1
assert call["messages"][1]["content"][1]["text"] == "describe this"
def test_memory_free_simple_processes_every_image_in_the_batch(
fake_handles, resolved_paths
):
node = llavaloader.LLavaOptionalMemoryFreeSimple()
node.generate_text(**_memory_free_kwargs(image=image_batch(2)))
(handle,) = fake_handles.instances
assert len(handle.llama.calls) == 2
def test_memory_free_simple_unloads_after_generating_when_asked(
fake_handles, resolved_paths
):
node = llavaloader.LLavaOptionalMemoryFreeSimple()
node.generate_text(**_memory_free_kwargs(unload=True))
(handle,) = fake_handles.instances
assert handle.closed is True
assert node._handle is None
def test_memory_free_simple_unloads_even_when_generation_fails(
fake_handles, resolved_paths, monkeypatch
):
"""Issue #137: a failed generation must not strand the model in VRAM."""
def explode(*args, **kwargs):
raise RuntimeError("llama.cpp exploded")
monkeypatch.setattr(llavaloader, "_run_batch", explode)
node = llavaloader.LLavaOptionalMemoryFreeSimple()
with pytest.raises(RuntimeError, match="exploded"):
node.generate_text(**_memory_free_kwargs(unload=True))
(handle,) = fake_handles.instances
assert handle.closed is True
assert node._handle is None
def test_memory_free_advanced_forwards_the_system_message_and_sampling(
fake_handles, resolved_paths
):
node = llavaloader.LLavaOptionalMemoryFreeAdvanced()
(text,) = node.generate_text_advanced(
ckpt_name=MODEL_FILE,
clip_name=CLIP_FILE,
max_ctx=4096,
gpu_layers=-1,
n_threads=4,
image=image_batch(1),
system_msg="answer in German",
prompt="describe",
max_tokens=64,
temperature=0.4,
top_p=0.85,
top_k=25,
frequency_penalty=0.0,
presence_penalty=0.0,
repeat_penalty=1.05,
seed=5,
unload=False,
)
assert text == "a description"
(handle,) = fake_handles.instances
(call,) = handle.llama.calls
assert call["messages"][0] == {"role": "system", "content": "answer in German"}
assert call["max_tokens"] == 64
assert call["temperature"] == 0.4
assert call["top_p"] == 0.85
assert call["top_k"] == 25
assert call["seed"] == 5
def test_cached_llava_key_is_insensitive_to_runtime_option_ordering(
fake_handles, resolved_paths
):
node = llavaloader.LLavaOptionalMemoryFreeSimple()
node._model(MODEL_FILE, CLIP_FILE, 4096, -1, 4, n_batch=256, main_gpu=1)
node._model(MODEL_FILE, CLIP_FILE, 4096, -1, 4, main_gpu=1, n_batch=256)
assert len(fake_handles.instances) == 1
-291
View File
@@ -1,291 +0,0 @@
from types import SimpleNamespace
import numpy as np
import pytest
from ComfyUI_VLM_nodes.nodes import minimax_music
def generation_request(**overrides):
values = {
"region": "global_en",
"model": "music-3.0",
"prompt": "Reflective acoustic pop",
"lyrics": "[Verse]\nA quiet road under evening light",
"stream": False,
"output_format": "hex",
"audio_format": "mp3",
"sample_rate": 44100,
"bitrate": 256000,
"lyrics_optimizer": False,
"is_instrumental": False,
"aigc_watermark": False,
"audio_url": "",
"audio_base64": "",
"cover_feature_id": "",
}
values.update(overrides)
return minimax_music.build_music_request(**values)
def test_music_contract_matches_current_models_regions_and_formats():
assert minimax_music.REGION_ENDPOINTS == {
"global_en": "https://api.minimax.io/v1/music_generation",
"cn_zh": "https://api.minimaxi.com/v1/music_generation",
}
assert minimax_music.DEFAULT_MODEL == "music-3.0"
assert minimax_music.GENERATION_MODELS == (
"music-3.0",
"music-2.6",
"music-3.0-free",
"music-2.6-free",
)
assert minimax_music.COVER_MODELS == ("music-cover", "music-cover-free")
assert minimax_music.OUTPUT_FORMATS == ("url", "hex")
assert minimax_music.AUDIO_FORMATS == ("mp3", "wav", "pcm")
assert {
"model",
"prompt",
"lyrics",
"stream",
"output_format",
"audio_setting",
"lyrics_optimizer",
"is_instrumental",
"audio_url",
"audio_base64",
"cover_feature_id",
} == minimax_music.REQUEST_FIELDS
assert minimax_music.REGIONAL_FIELDS == {
"global_en": (),
"cn_zh": ("aigc_watermark",),
}
def test_generation_request_covers_generation_and_cn_fields():
request = generation_request(
region="cn_zh",
stream=True,
lyrics="",
lyrics_optimizer=True,
aigc_watermark=True,
audio_format="wav",
sample_rate=32000,
bitrate=128000,
)
assert request == {
"model": "music-3.0",
"prompt": "Reflective acoustic pop",
"stream": True,
"output_format": "hex",
"audio_setting": {
"sample_rate": 32000,
"bitrate": 128000,
"format": "wav",
},
"lyrics_optimizer": True,
"is_instrumental": False,
"aigc_watermark": True,
}
@pytest.mark.parametrize(
("source", "value"),
[
("audio_url", "https://media.example/reference.wav"),
("audio_base64", "dGVzdA=="),
("cover_feature_id", "feature-123"),
],
)
def test_cover_request_supports_each_documented_source(source, value):
overrides = {
"model": "music-cover",
"prompt": "Warm orchestral cover",
"lyrics": "Updated words for the cover",
source: value,
}
request = generation_request(**overrides)
assert request[source] == value
assert "lyrics_optimizer" not in request
assert "is_instrumental" not in request
def test_streaming_requires_hex_and_cover_sources_are_exclusive():
with pytest.raises(ValueError, match="output_format='hex'"):
generation_request(stream=True, output_format="url")
with pytest.raises(ValueError, match="exactly one"):
generation_request(
model="music-cover",
prompt="Warm orchestral cover",
audio_url="https://media.example/reference.wav",
audio_base64="dGVzdA==",
)
def test_stream_response_joins_hex_chunks_and_requires_completion():
class Response:
def iter_lines(self):
return iter(
[
'data: {"data":{"status":1,"audio":"0001"},'
'"base_resp":{"status_code":0}}',
'data: {"data":{"status":2,"audio":"0203"},'
'"extra_info":{"music_sample_rate":32000,"music_channel":2},'
'"base_resp":{"status_code":0}}',
"data: [DONE]",
]
)
audio, metadata = minimax_music._stream_audio(Response())
assert audio == "00010203"
assert metadata == {"music_sample_rate": 32000, "music_channel": 2}
def test_url_and_hex_response_decoding():
class DownloadResponse:
def __enter__(self):
return self
def __exit__(self, *_args):
return False
def raise_for_status(self):
return None
def iter_bytes(self):
return iter((b"ab", b"cd"))
class Client:
def stream(self, method, url):
assert method == "GET"
assert url == "https://media.example/music.wav"
return DownloadResponse()
client = Client()
assert minimax_music._audio_bytes(client, "61626364", "hex") == b"abcd"
assert (
minimax_music._audio_bytes(
client,
"https://media.example/music.wav",
"url",
)
== b"abcd"
)
def test_pcm_decoding_uses_response_sample_rate_and_channel_count():
content = np.array([0, 32767, -32768, 0], dtype="<i2").tobytes()
samples, sample_rate = minimax_music._decode_audio(
content,
"pcm",
44100,
{"music_sample_rate": 32000, "music_channel": 2},
)
assert samples.shape == (2, 2)
assert sample_rate == 32000
assert samples[0, 1] == pytest.approx(32767 / 32768)
def test_node_posts_to_fixed_region_and_returns_comfy_audio(monkeypatch):
captured = {}
class Response:
def raise_for_status(self):
return None
def json(self):
return {
"data": {"status": 2, "audio": "0102"},
"extra_info": {"music_sample_rate": 44100, "music_channel": 2},
"base_resp": {"status_code": 0},
}
class Client:
def __init__(self, **kwargs):
captured["client"] = kwargs
def post(self, endpoint, *, headers, json):
captured["endpoint"] = endpoint
captured["headers"] = headers
captured["request"] = json
return Response()
def close(self):
captured["closed"] = True
def fake_soundfile_read(buffer, **kwargs):
assert buffer.read() == b"\x01\x02"
assert kwargs == {"dtype": "float32", "always_2d": True}
return np.zeros((8, 2), dtype=np.float32), 44100
def fake_require_module(name, *_args):
if name == "httpx":
return SimpleNamespace(Client=Client)
if name == "soundfile":
return SimpleNamespace(read=fake_soundfile_read)
raise AssertionError(name)
monkeypatch.setenv(minimax_music.API_KEY_ENV, "test-key-not-for-production")
monkeypatch.setattr(minimax_music, "require_module", fake_require_module)
result = minimax_music.MiniMaxMusicNode().generate_music(
region="global_en",
model="music-3.0",
prompt="Reflective acoustic pop",
lyrics="[Verse]\nA quiet road under evening light",
stream=False,
output_format="hex",
audio_format="wav",
sample_rate=44100,
bitrate=256000,
lyrics_optimizer=False,
is_instrumental=False,
aigc_watermark=False,
)
assert captured["endpoint"] == minimax_music.REGION_ENDPOINTS["global_en"]
assert captured["headers"]["Authorization"].startswith("Bearer ")
assert captured["client"] == {
"timeout": 600.0,
"follow_redirects": False,
"trust_env": False,
}
assert captured["closed"] is True
assert len(result) == 3
assert result[1] == 44100
assert result[2]["waveform"].shape == (1, 2, 8)
def test_request_failures_redact_the_resolved_key(monkeypatch):
resolved_value = "unit-key"
class Client:
def __init__(self, **_kwargs):
pass
def post(self, *_args, **_kwargs):
raise RuntimeError(f"Authorization: Bearer {resolved_value}")
def close(self):
pass
monkeypatch.setenv(minimax_music.API_KEY_ENV, resolved_value)
monkeypatch.setattr(
minimax_music,
"require_module",
lambda *_args: SimpleNamespace(Client=Client),
)
with pytest.raises(RuntimeError) as captured:
minimax_music.MiniMaxMusicNode().generate_music(
region="global_en",
model="music-3.0",
prompt="Reflective acoustic pop",
lyrics="[Verse]\nA quiet road under evening light",
stream=False,
output_format="hex",
audio_format="mp3",
sample_rate=44100,
bitrate=256000,
lyrics_optimizer=False,
is_instrumental=False,
aigc_watermark=False,
)
assert resolved_value not in str(captured.value)
assert "[REDACTED]" in str(captured.value)
-77
View File
@@ -1,77 +0,0 @@
from __future__ import annotations
import importlib
import sys
from pathlib import Path
from types import ModuleType, SimpleNamespace
import torch
PACKAGE = __package__.split(".")[0] if __package__ else "ComfyUI_VLM_nodes"
module = importlib.import_module(f"{PACKAGE}.nodes.moondream2")
def test_native_checkpoint_loader_bypasses_transformers_from_pretrained(
tmp_path: Path,
monkeypatch,
):
package = ModuleType(module._CHECKPOINT_PACKAGE)
package.__path__ = [str(tmp_path.resolve())]
package.__package__ = module._CHECKPOINT_PACKAGE
checkpoint = ModuleType(f"{module._CHECKPOINT_PACKAGE}.hf_moondream")
calls = {}
class FakeConfig:
@classmethod
def from_pretrained(cls, model_path, **kwargs):
calls["config"] = (Path(model_path), kwargs)
return cls()
class FakeModel(torch.nn.Module):
def __init__(self, config):
super().__init__()
self.weight = torch.nn.Parameter(torch.zeros(1))
calls["model_config"] = config
checkpoint.HfConfig = FakeConfig
checkpoint.HfMoondream = FakeModel
monkeypatch.setitem(sys.modules, module._CHECKPOINT_PACKAGE, package)
monkeypatch.setitem(
sys.modules,
f"{module._CHECKPOINT_PACKAGE}.hf_moondream",
checkpoint,
)
weights = tmp_path / "model.safetensors"
weights.write_bytes(b"test")
def load_model(model, filename, *, strict):
calls["weights"] = (model, Path(filename), strict)
model.weight.data.fill_(1)
return set(), []
monkeypatch.setattr(
module,
"require_module",
lambda name: (
SimpleNamespace(load_model=load_model)
if name == "safetensors.torch"
else None
),
)
model = module._load_native_checkpoint(tmp_path)
assert isinstance(model, FakeModel)
assert not model.training
assert model.weight.item() == 1
assert calls["config"] == (tmp_path, {"local_files_only": True})
assert calls["weights"] == (model, weights, True)
def test_photon_requirements_pin_cuda_runtime_with_required_symbol():
requirements = (
Path(module.__file__).resolve().parents[1] / "requirements-moondream31.txt"
).read_text(encoding="utf-8")
assert "kestrel-kernels==0.4.6" in requirements
assert "nvidia-cuda-runtime-cu12==12.9.79" in requirements
-361
View File
@@ -1,361 +0,0 @@
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()
-632
View File
@@ -1,632 +0,0 @@
import base64
import inspect
import io
from contextlib import nullcontext
from pathlib import Path
from types import SimpleNamespace
import ComfyUI_VLM_nodes as package
import numpy as np
import pytest
import torch
from ComfyUI_VLM_nodes.nodes import (
audioldm2,
florence2,
modern_vlm,
paligemma,
qwen2vl,
)
from ComfyUI_VLM_nodes.nodes import (
runtime as vlm_runtime,
)
from ComfyUI_VLM_nodes.nodes.runtime import (
LlamaHandle,
LlavaClipConfig,
accelerator_backend,
external_device_map,
image_data_uri,
llama_chat_content,
llama_cpp_diagnostics,
pil_mask_to_tensor,
pil_to_tensor,
runtime_diagnostics,
tensor_batch_to_pil,
torch_dtype,
)
from PIL import Image
def test_every_module_imports_and_expected_nodes_exist():
assert package.IMPORT_ERRORS == {}
expected = {
"ModernVLM",
"LegacyModernVLM",
"VLMRuntimeDiagnostics",
"Florence2",
"Paligemma",
"MolmoNode",
"Qwen2VLNode",
"Moondream2model",
"MiniCPMNode",
}
assert expected <= package.NODE_CLASS_MAPPINGS.keys()
def test_node_schemas_do_not_use_force_input():
for node_class in package.NODE_CLASS_MAPPINGS.values():
schema = node_class.INPUT_TYPES()
assert "forceInput" not in repr(schema)
def test_source_has_no_runtime_installer_or_direct_cuda_cache():
root = Path(package.__file__).parent
source = "\n".join(
path.read_text(encoding="utf-8", errors="replace")
for path in (root / "nodes").rglob("*.py")
)
assert "torch.cuda.empty_cache" not in source
assert "subprocess.run" not in source
assert "pip install" not in source
def test_portable_device_dtype_and_backend_contracts(monkeypatch):
assert torch_dtype("float16", torch.device("cpu")) == torch.float32
assert torch_dtype("float16", torch.device("mps")) == torch.float16
assert torch_dtype("float16", torch.device("xpu")) == torch.float16
assert accelerator_backend(torch.device("mps")) == "apple-metal"
assert accelerator_backend(torch.device("xpu")) == "intel-xpu"
monkeypatch.setattr(torch.version, "hip", None, raising=False)
assert accelerator_backend(torch.device("cuda")) == "nvidia-cuda"
monkeypatch.setattr(torch.version, "hip", "7.2", raising=False)
assert accelerator_backend(torch.device("cuda")) == "amd-rocm"
def test_runtime_report_and_device_map_are_supportable():
report = runtime_diagnostics()
assert {
"platform",
"machine",
"python",
"torch",
"device",
"backend",
"bf16",
"torch_cuda",
"torch_hip",
"packages",
"llama_cpp",
} <= report.keys()
device_map = external_device_map()
assert set(device_map) == {""}
assert device_map[""] == report["device"]
def test_dependency_metadata_matches_installer_requirements():
try:
import tomllib
except ModuleNotFoundError:
pytest.skip("tomllib is built into Python 3.11+")
from packaging.requirements import Requirement
root = Path(package.__file__).parent
metadata = tomllib.loads((root / "pyproject.toml").read_text("utf-8"))
project_requirements = {
str(Requirement(value)) for value in metadata["project"]["dependencies"]
}
installer_requirements = {
str(Requirement(line))
for line in (root / "requirements.txt").read_text("utf-8").splitlines()
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)
for value in metadata["project"]["dependencies"]
if Requirement(value).name == "bitsandbytes"
)
assert bitsandbytes.marker is not None
supported = (
("linux", "x86_64"),
("linux", "aarch64"),
("win32", "AMD64"),
("win32", "ARM64"),
("darwin", "arm64"),
)
unsupported = (
("darwin", "x86_64"),
("linux", "ppc64le"),
)
for system, machine in supported:
assert bitsandbytes.marker.evaluate(
{"sys_platform": system, "platform_machine": machine}
)
for system, machine in unsupported:
assert not bitsandbytes.marker.evaluate(
{"sys_platform": system, "platform_machine": machine}
)
gguf_extra = {
str(Requirement(value))
for value in metadata["project"]["optional-dependencies"]["gguf"]
}
gguf_requirements = {
str(Requirement(line))
for line in (root / "requirements-llama-cpp.txt")
.read_text("utf-8")
.splitlines()
if line.strip() and not line.lstrip().startswith("#")
}
assert gguf_extra == gguf_requirements
def _fake_llama_module(llama_class, *, gpu=True, mmap=True):
return SimpleNamespace(
__version__="0.3.34",
Llama=llama_class,
LLAMA_SPLIT_MODE_LAYER=1,
LLAMA_SPLIT_MODE_ROW=2,
LLAMA_SPLIT_MODE_NONE=0,
llama_supports_gpu_offload=lambda: gpu,
llama_supports_mmap=lambda: mmap,
llama_supports_mlock=lambda: False,
llama_print_system_info=lambda: (
b"GGML_CUDA = 1 | BLAS = 1" if gpu else b"BLAS = 1"
),
)
def test_llama_cpp_diagnostics_reports_its_own_backend():
class FakeLlama:
pass
report = llama_cpp_diagnostics(_fake_llama_module(FakeLlama))
assert report["version"] == "0.3.34"
assert report["gpu_offload"] is True
assert report["mmap"] is True
assert report["backends"] == ["cuda", "blas"]
def test_llama_chat_content_rejects_empty_or_malformed_responses():
assert (
llama_chat_content({"choices": [{"message": {"content": " ready "}}]})
== "ready"
)
with pytest.raises(RuntimeError, match="empty response"):
llama_chat_content({"choices": [{"message": {"content": None}}]})
with pytest.raises(RuntimeError, match="unexpected response"):
llama_chat_content({"choices": []})
def test_llama_handle_falls_back_to_cpu_for_cpu_only_build(monkeypatch, tmp_path):
calls = []
class FakeLlama:
def __init__(self, **kwargs):
calls.append(kwargs)
def close(self):
calls.append("closed")
module = _fake_llama_module(FakeLlama, gpu=False, mmap=False)
monkeypatch.setattr(vlm_runtime, "require_module", lambda *_args: module)
reserved = []
monkeypatch.setattr(vlm_runtime, "reserve_external_vram", reserved.append)
model_path = tmp_path / "model.gguf"
model_path.write_bytes(b"gguf")
handler_gpu = []
class Handler:
def close(self):
handler_gpu.append("closed")
def handler_factory(*, use_gpu):
handler_gpu.append(use_gpu)
return Handler()
handle = LlamaHandle(
model_path,
n_ctx=0,
n_gpu_layers=-1,
n_threads=4,
n_batch=1024,
n_ubatch=768,
flash_attention="Auto",
use_mmap=True,
chat_handler_factory=handler_factory,
)
handle.ensure_loaded()
assert reserved == []
assert handler_gpu == [False]
assert calls[0]["n_gpu_layers"] == 0
assert calls[0]["n_batch"] == 1024
assert calls[0]["n_ubatch"] == 768
assert calls[0]["offload_kqv"] is False
assert calls[0]["op_offload"] is False
assert calls[0]["flash_attn"] is False
assert calls[0]["use_mmap"] is False
handle.close()
assert calls[-1] == "closed"
assert handler_gpu[-1] == "closed"
def test_llama_handle_uses_accelerator_batching_and_multi_gpu(monkeypatch, tmp_path):
calls = []
class FakeLlama:
def __init__(self, **kwargs):
calls.append(kwargs)
module = _fake_llama_module(FakeLlama)
monkeypatch.setattr(vlm_runtime, "require_module", lambda *_args: module)
reserved = []
monkeypatch.setattr(vlm_runtime, "reserve_external_vram", reserved.append)
model_path = tmp_path / "model.gguf"
projector_path = tmp_path / "mmproj.gguf"
model_path.write_bytes(b"1234")
projector_path.write_bytes(b"123")
handle = LlamaHandle(
model_path,
n_ctx=256,
n_gpu_layers=-1,
n_threads=6,
n_batch=512,
n_ubatch=1024,
split_mode="Row",
main_gpu=1,
tensor_split="0.25, 0.75",
projector_path=projector_path,
)
handle.ensure_loaded()
assert reserved == [7]
assert calls[0]["n_batch"] == 256
assert calls[0]["n_ubatch"] == 256
assert "n_threads_batch" not in calls[0]
assert calls[0]["split_mode"] == 2
assert calls[0]["main_gpu"] == 1
assert calls[0]["tensor_split"] == [0.25, 0.75]
assert calls[0]["flash_attn"] is True
assert calls[0]["offload_kqv"] is True
def test_llama_handle_auto_flash_attention_retries_portably(monkeypatch, tmp_path):
calls = []
class FakeLlama:
def __init__(self, **kwargs):
calls.append(kwargs)
if kwargs["flash_attn"]:
raise RuntimeError("flash attention is not supported")
monkeypatch.setattr(
vlm_runtime,
"require_module",
lambda *_args: _fake_llama_module(FakeLlama),
)
monkeypatch.setattr(vlm_runtime, "reserve_external_vram", lambda _size: None)
model_path = tmp_path / "model.gguf"
model_path.write_bytes(b"gguf")
LlamaHandle(
model_path,
n_ctx=128,
n_gpu_layers=-1,
n_threads=2,
).ensure_loaded()
assert [call["flash_attn"] for call in calls] == [True, False]
def test_llava_handler_auto_and_explicit_selection(monkeypatch, tmp_path):
calls = []
class MTMD:
def __init__(self, **kwargs):
calls.append(("auto", kwargs))
class MiniCPM:
def __init__(self, **kwargs):
calls.append(("minicpm", kwargs))
monkeypatch.setattr(
vlm_runtime,
"require_module",
lambda *_args: SimpleNamespace(
MTMDChatHandler=MTMD,
MiniCPMv26ChatHandler=MiniCPM,
),
)
projector = tmp_path / "mmproj.gguf"
projector.write_bytes(b"gguf")
LlavaClipConfig(projector, "Auto (GGUF chat template)").create(use_gpu=False)
LlavaClipConfig(projector, "MiniCPM-V 2.6").create(use_gpu=True)
assert calls[0][0] == "auto"
assert calls[0][1]["use_gpu"] is False
assert calls[1][0] == "minicpm"
def test_image_roundtrip_and_png_data_uri():
tensor = torch.tensor(
[[[[0.0, 0.5, 1.0], [1.0, float("nan"), 0.0]]]],
dtype=torch.float32,
)
images = tensor_batch_to_pil(tensor)
assert images[0].size == (2, 1)
uri = image_data_uri(images[0])
payload = base64.b64decode(uri.split(",", 1)[1])
assert Image.open(io.BytesIO(payload)).format == "PNG"
assert pil_to_tensor(images[0]).shape == (1, 1, 2, 3)
assert pil_mask_to_tensor(Image.new("L", (2, 3))).shape == (1, 3, 2)
def test_paligemma_parser_uses_normalized_boxes_and_16_codes():
codes = "".join(f"<seg{index:03d}>" for index in range(16))
parsed = paligemma.parse_segments(
f"<loc0100><loc0200><loc0900><loc0800>{codes} cat"
)
assert len(parsed) == 1
box, values, label = parsed[0]
assert box == pytest.approx((100 / 1024, 200 / 1024, 900 / 1024, 800 / 1024))
assert values == list(range(16))
assert label == "cat"
def test_florence_rendering_supports_boxes_quads_and_nested_polygons():
image = Image.new("RGB", (32, 24), "black")
parsed = {
"<TASK>": {
"bboxes": [[1, 1, 10, 10]],
"labels": ["box"],
"quad_boxes": [[2, 2, 8, 2, 8, 8, 2, 8]],
"polygons": [[[4, 4, 20, 4, 20, 20, 4, 20]]],
}
}
mask, visual = florence2._visualize(image, parsed)
assert np.asarray(mask).max() == 255
assert visual.size == image.size
def test_modern_catalog_has_current_quality_and_low_vram_tiers():
repositories = {spec.repo_id for spec in modern_vlm.MODEL_CATALOG.values()}
small_fast = [spec for spec in modern_vlm.MODEL_CATALOG.values() if spec.small_fast]
assert 10 <= len(small_fast) <= 20
assert all(
not spec.trust_remote_code
for spec in modern_vlm.MODEL_CATALOG.values()
if spec.family != "Custom"
)
assert modern_vlm.MODEL_CATALOG["Custom Hugging Face model"].trust_remote_code
assert "Qwen/Qwen3.5-4B" in repositories
assert "Qwen/Qwen3.5-35B-A3B" in repositories
assert "Qwen/Qwen3.6-27B" in repositories
assert "Qwen/Qwen3-VL-8B-Instruct" in repositories
assert "Qwen/Qwen2.5-VL-3B-Instruct" in repositories
assert "google/gemma-3-4b-it" in repositories
assert "HuggingFaceTB/SmolVLM2-256M-Video-Instruct" in repositories
assert "HuggingFaceTB/SmolVLM2-500M-Video-Instruct" in repositories
assert "LiquidAI/LFM2.5-VL-450M" in repositories
assert "LiquidAI/LFM2.5-VL-1.6B" in repositories
assert "OpenGVLab/InternVL3_5-1B-HF" in repositories
assert "OpenGVLab/InternVL3_5-2B-HF" in repositories
assert "ibm-granite/granite-vision-3.3-2b" in repositories
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"]
predictor = modern_vlm.ModernVLMPredictor.__new__(modern_vlm.ModernVLMPredictor)
predictor.spec = modern_vlm.ModelSpec("test/model", "Qwen 3.5", 1.0, video=True)
captured = {}
def capture(messages, enable_thinking=False, **kwargs):
captured["messages"] = messages
captured["enable_thinking"] = enable_thinking
captured.update(kwargs)
raise RuntimeError("captured before inference")
predictor._inputs = capture
frames = torch.zeros((4, 8, 8, 3), dtype=torch.float32)
with pytest.raises(RuntimeError, match="captured before inference"):
predictor.generate(
None,
"What moves?",
"",
8,
0.0,
0.9,
frames,
2.0,
True,
)
content = captured["messages"][-1]["content"]
assert [part["type"] for part in content] == ["video", "text"]
assert len(content[0]["video"]) == 4
assert "2 FPS" in content[1]["text"]
assert captured["enable_thinking"] is True
assert captured["video_metadata"]["fps"] == 2.0
assert captured["video_metadata"]["frames_indices"] == [0, 1, 2, 3]
def test_modern_vlm_streams_cumulative_text_without_changing_final_output(
monkeypatch,
):
class FakeStreamer:
def __init__(self, _tokenizer, **kwargs):
assert kwargs["skip_prompt"] is True
self.chunks = ["Hello ", "from ", "the VLM."]
def __iter__(self):
return iter(self.chunks)
def end(self):
pass
class FakeModel:
def generate(self, **kwargs):
assert isinstance(kwargs["streamer"], FakeStreamer)
return torch.tensor([[10, 11, 12]], dtype=torch.long)
class FakeProcessor:
tokenizer = object()
def batch_decode(self, *_args, **_kwargs):
return ["fallback"]
predictor = modern_vlm.ModernVLMPredictor.__new__(
modern_vlm.ModernVLMPredictor
)
predictor.spec = modern_vlm.ModelSpec("test/model", "Test", 1.0)
predictor.dtype = torch.float32
predictor.processor = FakeProcessor()
predictor.streamer_class = FakeStreamer
predictor.handle = SimpleNamespace(ensure_loaded=lambda: FakeModel())
predictor._inputs = lambda *_args, **_kwargs: {
"input_ids": torch.tensor([[1, 2]], dtype=torch.long)
}
monkeypatch.setattr(modern_vlm, "model_device", lambda _model: torch.device("cpu"))
monkeypatch.setattr(modern_vlm, "move_inputs", lambda inputs, _device: inputs)
monkeypatch.setattr(
modern_vlm,
"inference_context",
lambda *_args: nullcontext(),
)
partials = []
result = predictor.generate(
torch.zeros((1, 8, 8, 3), dtype=torch.float32),
"Describe it.",
"",
16,
0.0,
0.9,
stream_callback=partials.append,
)
assert result == "Hello from the VLM."
assert partials == ["Hello", "Hello from", "Hello from the VLM."]
def test_view_text_frontend_rehydrates_and_uses_native_progress_channel():
source = (
Path(package.__file__).parent / "web" / "js" / "viewText.js"
).read_text(encoding="utf-8")
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():
predictor = modern_vlm.ModernVLMPredictor.__new__(modern_vlm.ModernVLMPredictor)
predictor.spec = modern_vlm.ModelSpec("test/model", "InternVL 3.5", 1.0, video=True)
captured = {}
class ImageProcessor:
size = {"height": 448, "width": 448}
class Processor:
image_processor = ImageProcessor()
def apply_chat_template(self, _messages, **kwargs):
captured.update(kwargs)
return {"input_ids": torch.ones((1, 1), dtype=torch.long)}
predictor.processor = Processor()
predictor._inputs(
[{"role": "user", "content": [{"type": "text", "text": "test"}]}],
video_metadata={"fps": 2.0},
)
assert captured["processor_kwargs"]["size"] == {
"height": 448,
"width": 448,
}
def test_qwen2_legacy_quantized_labels_use_maintained_backends():
assert list(qwen2vl.QWEN2_VL_CHOICES) == [
"Qwen2-VL-2B",
"Qwen2-VL-7B",
]
assert qwen2vl.LEGACY_QUANTIZED_ALIASES["Qwen2-VL-7B-GPTQ-Int8"] == (
"Qwen2-VL-7B",
"Balanced (8-bit)",
)
assert qwen2vl.LEGACY_QUANTIZED_ALIASES["Qwen2-VL-7B-AWQ"] == (
"Qwen2-VL-7B",
"Maximum Savings (4-bit)",
)
def test_audioldm_keeps_legacy_outputs_and_adds_standard_audio(monkeypatch):
class FakePredictor:
def generate(self, *_args):
return np.zeros((2, 16), dtype=np.float32), 16000
node = audioldm2.AudioLDM2Node()
monkeypatch.setattr(node, "get_or_create_model", lambda *_args: FakePredictor())
result = node.generate_audio_final("rain", "", 1, 3.5, 16000, 42, 2, "wav")
assert len(result) == 3
assert result[1] == 16000
assert result[2]["waveform"].shape == (2, 1, 16)
def test_node_functions_accept_every_declared_input_name():
for node_class in package.NODE_CLASS_MAPPINGS.values():
function = getattr(node_class, node_class.FUNCTION)
signature = inspect.signature(function)
if any(
parameter.kind == inspect.Parameter.VAR_KEYWORD
for parameter in signature.parameters.values()
):
continue
declared = {
name
for group in node_class.INPUT_TYPES().values()
if isinstance(group, dict)
for name in group
}
accepted = set(signature.parameters)
assert declared <= accepted, (
node_class.__name__,
declared - accepted,
)
-753
View File
@@ -1,753 +0,0 @@
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

Some files were not shown because too many files have changed in this diff Show More