39 changed files with 10924 additions and 321 deletions
+3 -1
View File
@@ -52,9 +52,11 @@ jobs:
- name: Install ComfyUI and node dependencies
run: |
git clone --depth 1 https://github.com/Comfy-Org/ComfyUI.git ../ComfyUI
python -m pip install pytest packaging
python -m pip install pytest packaging build
python -m pip install -r ../ComfyUI/requirements.txt -r requirements.txt
- name: Test
run: python -m pytest -q
- name: Compile
run: python -m compileall -q .
- name: Build distribution
run: python -m build
+95 -1
View File
@@ -26,6 +26,69 @@ 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.
### 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.
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)
## Dependency behavior
- Python 3.10 through 3.13 is covered by CI.
@@ -58,7 +121,7 @@ official project currently publishes backend indexes and documents source
build flags:
```bash
# NVIDIA; replace cu124 with the CUDA index matching the environment.
# 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
@@ -87,11 +150,42 @@ Source builds use `GGML_CUDA=on`, `GGML_METAL=on`, `GGML_HIP=on`,
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
+73 -1
View File
@@ -1,6 +1,6 @@
# Model validation
Validated on 2026-07-28 with ComfyUI 0.28.0, Python 3.12, Transformers 5.14.1,
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.
@@ -15,6 +15,7 @@ on the D drive and executed through WSL.
| 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
@@ -29,6 +30,18 @@ workflow (`EmptyImage` -> `ModernVLM` -> `ViewText`) ran the cached LFM2.5-VL
`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
@@ -37,6 +50,65 @@ SmolVLM2 256M/500M/2.2B, LFM2.5 VL 450M/1.6B, InternVL 3.5 1B/2B, and Granite
Vision 3.3 2B/4.1 4B. Gemma 3 4B is the sixteenth entry and correctly requires
license acceptance plus `HF_TOKEN`.
## Structured vision validation
The versioned detection/track/point/event payloads, geometry and mask
conversion, strict spatial parser, Grounding-family adapters, SAM2.1 session
plumbing, SAM3 bit-packed payload adapter, and ByteTrack-style association pass
the local WSL contract suite. Those tests validate schemas, shapes, output
ordering, bounds, timestamps, deterministic IDs, and error handling.
Representative real-weight checks were then submitted through ComfyUI's local
`POST /prompt` API and verified from `/history/{prompt_id}`. The test machine
used ComfyUI 0.28.0, Python 3.12.12, PyTorch 2.13.0+cu126, Transformers 5.14.1,
and an NVIDIA RTX 3090. Input media, checkpoints, model caches, ComfyUI, and
this checkout all remained on the D drive under WSL.
| Family | Representative checkpoint policy | Real-weight status |
| --- | --- | --- |
| Grounding DINO | Tiny; Base uses the same loader/processor contract | **Passed**: FP16, four real 640x360 video frames in two-frame micro-batches; person and bird boxes/labels were visually checked, serialized, timestamped, and in bounds |
| OWLv2 | Base Ensemble | Pending |
| OmDet Turbo | Swin Tiny | Pending |
| SAM2.1 video | Hiera Tiny; sibling sizes use the same session adapter | **Passed**: FP16, real 12-frame 640x360 clip at 24 FPS with CPU preprocessing/state. Grounding's core `BOUNDING_BOX` output connected directly: the forward union-only run kept one person ID on frames 0-11; a last-frame reverse run kept two IDs for 24 observations and emitted 24 frame-major object masks. All geometry was in bounds and first/last overlays and masks were visually checked |
| Comfy core SAM3.1 | `sam3.1_multiplex_fp16.safetensors`, only after license/access is available | Pending |
| SAM3 adapter/report | Synthetic core payload contract | Passed without weights; real core handoff pending |
| ByteTrack-style tracker | Deterministic synthetic crossing, missed-frame, and expiry cases | Passed; no model weights exist |
| Florence-2 multitask | Base FT; Large uses the same native Transformers contract | **Passed**: real object-detection API run produced bounded woman, face, and clothing boxes plus a visually checked overlay |
The SAM2 API check initially exposed a real session-lifecycle defect that unit
fixtures did not: prompt insertion must be followed by inference on the seeded
frame before propagation. The implementation now performs that seed pass and
also propagates in reverse when `seed_frame` is greater than zero. Later live
checks exercised nested multi-object core boxes, CPU preprocessing/state,
union-only low-memory output, optional object-mask output, disabled preview
rendering, reverse propagation, and `unload_after=true` for both models. The
final unload run returned total reported GPU memory use to within 4 MiB of the
pre-run `nvidia-smi` baseline.
Grounding DINO and SAM2 sibling sizes are catalog-available but were not
downloaded or executed. OWLv2, OmDet Turbo, and gated SAM3 remain explicitly
unverified; the UI never presents them as locally tested simply because their
schemas import.
The acceptance run for each model family must record:
1. Exact checkpoint revision, ComfyUI/Python/PyTorch/Transformers versions,
device, dtype, peak accelerator allocation, and wall time.
2. A real image or short bounded video with manually verified boxes, labels,
masks, timestamps, and stable IDs.
3. The canonical JSON schema/version and every advertised output socket,
including preview/report output through ComfyUI's local `/prompt` API.
4. A second queue using the cached model, followed by an `unload_after=true`
run where that option exists.
5. Failure behavior for an absent checkpoint or gated access without exposing
a token.
One checkpoint per distinct implementation family is enough for sibling model
sizes that share the same code path. Validation prioritizes the smallest useful
checkpoint and will not download or execute a 30B model. A larger variant is
tested only when it has a different loader, processor, postprocessor, or
quantization path.
## Not marked passed
- Qwen 3 VL 30B-A3B: weights are available locally, but inference validation
+181 -2
View File
@@ -1,10 +1,11 @@
# ComfyUI VLM Nodes
Production-oriented vision-language, structured prompting, audio, and utility
nodes for ComfyUI. Version 2.1 supports ComfyUI's selected NVIDIA CUDA, AMD
nodes for ComfyUI. Version 2.3 supports ComfyUI's selected NVIDIA CUDA, AMD
ROCm, Apple Metal, Intel XPU, and CPU device without replacing its PyTorch
build. It removes startup installers and global accelerator cache flushes,
adds real image/video batches, and uses ComfyUI model residency and offloading.
adds real image/video batches and live token streaming, and uses ComfyUI model
residency and offloading.
## Modern model coverage
@@ -31,6 +32,15 @@ is enabled only when the explicit custom-model option requires it. Florence-2
uses the Transformers-native converted checkpoints instead of Microsoft’s
legacy repository code.
## Live text output
`Modern VLM` streams decoded text through ComfyUI's native `progress_text`
WebSocket channel by default. A connected `ViewText` node updates while tokens
arrive, shows the final response after execution, and restores the last result
when ComfyUI rehydrates workflow output history. Disable `stream_output` for
API-only or headless runs that do not need incremental UI updates. Streaming is
best-effort and never changes the final `STRING` output or makes inference fail.
Specialized nodes remain available where a generic chat node would discard
useful model capabilities:
@@ -46,6 +56,157 @@ useful model capabilities:
- **llama.cpp LLaVA/GGUF**, structured prompt suggestions, OpenAI-compatible
prompting, and AudioLDM2.
## Structured detection, segmentation, and tracking
The vision nodes use stable, typed sockets instead of passing model-specific
lists between nodes:
| Socket | JSON schema | Purpose |
| --- | --- | --- |
| `VLM_DETECTIONS` | `comfyui-vlm/detections`, version 1 | Per-frame boxes, labels, scores, optional polygons/quads, and in-process masks |
| `VLM_TRACKS` | `comfyui-vlm/tracks`, version 1 | Durable object IDs with ordered observations over time |
| `VLM_POINTS` | `comfyui-vlm/points`, version 1 | Pixel-coordinate points, including detection centers |
| `VLM_EVENTS` | `comfyui-vlm/events`, version 1 | Ordered temporal events for downstream video analysis |
All spatial coordinates are source-image pixels. Bounding boxes are
`[x1, y1, x2, y2]` with an exclusive right/bottom edge; polygons contain at
least three points and quads exactly four. JSON roots contain `schema`,
`version`, media dimensions/frame count/FPS, and their ordered records. Dense
mask tensors remain in-process and are deliberately omitted from JSON so API
results do not unexpectedly grow by hundreds of megabytes.
The utility layer converts without model-specific glue:
- `VLMStructuredSpatialParser` strictly parses pixel, normalized 0–1, or
normalized 0–1000 JSON from any VLM into `VLM_DETECTIONS` and `VLM_POINTS`.
`VLMSpatialPromptBuilder` creates the matching constrained prompt.
- `VLMDetectionsToBoundingBoxes`, `VLMDetectionsToPoints`, and
`VLMDetectionsToMasks` emit Comfy core boxes, center points, combined and
individual binary masks, inverse masks, ready-to-preview black-and-white
images, and stable-color instance maps. Polygon/quad masks are rasterized
when present, otherwise the bounding box is used. Existing output indexes
remain stable; the creator-facing mask images and instance map are appended.
- `VLMFilterDetections`, `VLMSelectDetection`, `VLMCropDetections`, and
`VLMRenderDetections` provide label/score/area/frame selection, padded crops,
and deterministic overlays.
- `VLMMaskProcessor` accepts any Comfy `MASK`, including SAM2/SAM3 masks, and
returns a feathered matte, strict binary mask, inverse mask, and
black-and-white image. Its grow/shrink and Gaussian feathering run in Torch
without OpenCV or SciPy.
- `VLMMaskComposite` applies still-image or video mask batches to a source and
returns the replacement composite, isolated foreground, original
background-only plate, and black-and-white mask image. A single mask or
background broadcasts safely across a video batch.
- `VLMDetectionsFromJSON` and `VLMDetectionsToJSON` are the explicit API and
persistence boundary for the versioned detection schema.
### Open-vocabulary image and video detection
`VLMOpenVocabularyDetection` exposes one interface for:
- Grounding DINO Tiny and Base
- OWLv2 Base Ensemble
- OmDet Turbo Swin Tiny
It accepts a still image or an `IMAGE` batch of video frames and processes the
batch frame by frame. Outputs, in socket order, are `detections`, `json`,
`preview`, `box_mask`, and Comfy core `bounding_boxes`. Connect the FPS output
of `GetVideoComponents` when the input is video so every timestamp is correct.
For tracking-by-detection, run detection over the complete bounded batch and
connect it to `VLMTrackDetections`.
`VLMTrackDetections` uses a ByteTrack-style two-stage high/low-confidence
association, motion prediction, label-aware matching, and time-based expiry.
IDs are durable within the supplied sequence and survive short missed
detections when `emit_predictions` is enabled. Independent Comfy queue runs or
independently sliced chunks are separate tracking sessions; they do not
silently reuse IDs.
### SAM2.1 and Comfy core SAM3.1
`VLMSAM2VideoSegmentation` propagates first-frame detections, one core
`BOUNDING_BOX`, or seed masks through an `IMAGE` batch using SAM2.1 Hiera Tiny,
Small, Base+, or Large. It returns `VLM_TRACKS`, report JSON, per-frame union
masks, frame-major individual object masks, and an overlay batch. The object
IDs assigned at the seed frame remain stable for that video session.
`VLMSAM3TrackAdapter` is intentionally an adapter, not a second SAM3 loader. It
validates ComfyUI core `SAM3_TRACK_DATA`, preserves the core bit-packed mask
payload unchanged, and exposes lightweight `VLM_TRACKS` metadata with mask
references. Connect its passthrough output to core `SAM3_TrackPreview` or
`SAM3_TrackToMask`, and connect `tracks` to `VLMTrackReport`. This avoids
duplicating dense masks in memory or JSON.
SAM3 weights use Meta's SAM License. The upstream `facebook/sam3` repository
requires accepting access terms and sharing the requested account information;
the ComfyUI checkpoint is also marked `sam-license`. Review and accept the
license before downloading. The example names ComfyUI's
`sam3.1_multiplex_fp16.safetensors`; if it is unavailable, use the SAM2.1
workflow rather than substituting an unrelated checkpoint.
### Florence-2 task coverage
`Florence2` exposes all 15 supported task contracts:
| Task | Extra input | Structured result |
| --- | --- | --- |
| Caption | none | text |
| Detailed caption | none | text |
| More detailed caption | none | text |
| OCR | none | text |
| OCR with regions | none | text plus quadrilateral regions |
| Object detection | none | labeled boxes |
| Dense region caption | none | captions with boxes |
| Caption to phrase grounding | `text_input` | phrase boxes |
| Referring expression segmentation | `text_input` | polygons and mask |
| Region to segmentation | one `BOUNDING_BOX` per image | polygons and mask |
| Open vocabulary detection | `text_input` | model-provided spatial records |
| Region to category | one `BOUNDING_BOX` per image | text |
| Region to description | one `BOUNDING_BOX` per image | text |
| Region to OCR | one `BOUNDING_BOX` per image | text |
| Region proposals | none | boxes |
Every task returns `text`, `structured_json`, `mask`, and `visualization`.
Tasks that do not produce a spatial result return an empty mask and the source
image visualization. Region tasks reject ambiguous multi-box input; use
`VLMSelectDetection` to isolate the record, then supply exactly one core
`BOUNDING_BOX` with the same pixel coordinates.
### Video memory strategy
- Trim long media with core `Video Slice`, then use `GetVideoComponents`.
Downscale the complete frame batch before detection or segmentation and keep
every frame at identical dimensions.
- Grounding detection supports configurable micro-batches; keep `batch_size=1`
for minimum VRAM or increase it when memory allows. It returns both nested
per-frame core `BOUNDING_BOX` values and flat metadata-rich
`BOUNDING_BOXES`.
- SAM2.1 stores source video frames on CPU, keeps its inference state on CPU by
default, and limits the vision-feature cache to one frame. Union masks and
previews return on CPU. Full per-object mask volumes are opt-in with
`mask_output=union_and_objects`; disable `render_preview` to avoid another
full-resolution overlay copy on long clips.
- Start with Grounding DINO Tiny plus SAM2.1 Hiera Tiny. Increase detector or
segmenter size only after the pipeline is correct. `unload_after=false`
caches one model per node instance; use `true` when another large model must
run immediately afterward.
- A `Video Slice` is an independent propagation session. For very long media,
use bounded slices, reseed each slice, and keep the overlap/output mapping in
the caller. The pack does not pretend IDs are globally stable across separate
queues.
- The SAM3 adapter never unpacks the complete mask volume for its report. Use
core `SAM3_TrackToMask` only when a dense selected mask is actually needed.
API-format examples are in [`examples/vision`](examples/vision):
- [`grounding_dino_image_api.json`](examples/vision/grounding_dino_image_api.json)
- [`sam2_video_tracking_api.json`](examples/vision/sam2_video_tracking_api.json)
- [`sam3_core_adapter_blueprint_api.json`](examples/vision/sam3_core_adapter_blueprint_api.json)
Upload the named media to ComfyUI's input directory, adjust the filenames and
labels, then submit the JSON object as the `prompt` value to `/prompt`. These
are API graphs, not frontend workflow-export JSON.
## Install
Install through ComfyUI Manager, or clone into `ComfyUI/custom_nodes` and run:
@@ -70,6 +231,20 @@ python -m pip install -r ComfyUI/custom_nodes/ComfyUI_VLM_nodes/requirements-lla
See [COMPATIBILITY.md](COMPATIBILITY.md) for the tested matrix and official
backend-specific GGUF commands.
The GGUF loaders now query the installed llama.cpp build instead of inferring
its capabilities from PyTorch. Accelerator offload automatically falls back to
CPU when a CPU-only wheel is installed. Advanced optional inputs expose logical
and physical prompt batching (`n_batch`/`n_ubatch`), flash-attention policy,
mmap, and CUDA/ROCm multi-GPU layer/row splitting without changing legacy
workflow sockets. `Auto` flash attention retries the portable path if a
backend/model pair rejects it.
The **LLaVA Vision Projector Loader** supports metadata-driven MTMD plus
explicit handlers for LLaVA 1.5/1.6, MiniCPM-V 2.6, Moondream2, NanoLLaVA,
Qwen2.5-VL, Gemma 4, Llama 3 Vision Alpha, and Obsidian. Use the default
metadata-driven handler for current GGUF + mmproj pairs; select the named
legacy handler when a model card requires it.
Models are downloaded only when their node first executes and are stored below
`ComfyUI/models/LLavacheckpoints`. Hugging Face downloads respect `HF_TOKEN`.
Gemma 3 and PaLI-Gemma require accepting their model licenses on Hugging Face.
@@ -85,6 +260,9 @@ Gemma 3 and PaLI-Gemma require accepting their model licenses on Hugging Face.
stay on ComfyUI's active device instead of assuming GPU zero. Large-model
Accelerate placement is enabled on CUDA/ROCm/XPU; any disk offload remains
inside the model's ComfyUI directory.
- llama.cpp model and projector bytes are included in the pre-load reservation.
The runtime reports llama.cpp's own compiled backend, GPU-offload, mmap, and
mlock capabilities in **VLM Runtime Diagnostics**.
- `unload_after=false` caches one model per node instance for fast repeated
queues. Turn it on for maximum reclamation between prompts.
- A connected `video_frames` batch becomes the primary visual input. The
@@ -136,6 +314,7 @@ Real-weight checks are opt-in because they download multi-gigabyte checkpoints:
```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
```
See [MODEL_VALIDATION.md](MODEL_VALIDATION.md) for the exact real-weight and
+6
View File
@@ -10,6 +10,7 @@ node_list = [
"audioldm2",
"diagnostics",
"florence2",
"grounding",
"joytag",
"kosmos2",
"llavaloader",
@@ -22,9 +23,14 @@ node_list = [
"paligemma",
"playmusic",
"qwen2vl",
"sam2",
"sam3_adapter",
"simpletext",
"spatial_parser",
"suggest",
"tracking",
"uform",
"vision_utils",
]
NODE_CLASS_MAPPINGS = {}
+92
View File
@@ -0,0 +1,92 @@
# 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.
### `sam2_video_tracking_api.json`
Runs this bounded pipeline:
`LoadVideo` → `Video Slice` → `GetVideoComponents` → `ImageScale` →
`ImageFromBatch` → Grounding DINO first-frame detection → SAM2.1 propagation.
The example limits the source to two seconds, scales its largest dimension to
768 pixels while preserving aspect ratio, unloads Grounding DINO after
seeding, and keeps SAM2.1 video state on CPU. The example requests only the
union mask volume; change `mask_output` to `union_and_objects` only when every
per-object mask is required. `VLMTrackReport` is an output node and the final
`PreviewImage` displays SAM2.1 output index 4.
For a longer source, change `start_time` and keep a bounded `duration`.
Independent slices create independent object-ID sessions.
### `sam3_core_adapter_blueprint_api.json`
Uses ComfyUI core nodes to load and run SAM3.1, then passes core
`SAM3_TRACK_DATA` through `VLMSAM3TrackAdapter`. The adapter's output 1 is the
unchanged core payload consumed by `SAM3_TrackPreview`; output 0 is canonical
`VLM_TRACKS` consumed by `VLMTrackReport`.
The graph intentionally names:
`ComfyUI/models/checkpoints/sam3.1_multiplex_fp16.safetensors`
The checkpoint is not bundled. Review the SAM License before downloading
[Comfy-Org/sam3.1](https://huggingface.co/Comfy-Org/sam3.1). ComfyUI rejects
the graph at prompt validation when the named checkpoint is absent. Use the
SAM2.1 example when SAM3.1 access or compatible core support is unavailable.
## Output history
ComfyUI returns image/video previews in the execution history and text reports
in the output-node UI payload. Canonical JSON is also available on the linked
string outputs. Dense masks intentionally stay as tensors rather than being
embedded in the JSON report.
## Creator mask outputs
`VLM Detections to Masks` preserves its original first three outputs and
appends creator-ready derivatives:
| Index | Output |
| ---: | --- |
| 0 | Per-frame combined/union `MASK` |
| 1 | Flattened per-object `MASK` batch |
| 2 | JSON mapping each object mask to its frame/detection/track |
| 3 | Per-frame inverse/background `MASK` |
| 4 | Combined masks as black-and-white `IMAGE` batches |
| 5 | Individual masks as black-and-white `IMAGE` batches |
| 6 | Stable-color per-frame instance maps |
All binary mask values are exactly zero or one. `VLM Mask Processor` can grow,
shrink, and feather any of these masks and returns processed, binary, inverse,
and black-and-white image outputs. `VLM Mask Composite` accepts the resulting
mask plus still-image or video frames and returns a composite, isolated
foreground, background-only plate, and mask image. Connect an optional
background image/video batch to replace the solid background color.
@@ -0,0 +1,45 @@
{
"1": {
"class_type": "LoadImage",
"inputs": {
"image": "grounding_input.png"
}
},
"2": {
"class_type": "VLMOpenVocabularyDetection",
"inputs": {
"image": [
"1",
0
],
"model": "Grounding DINO Tiny (fast)",
"labels": "person, dog, bicycle",
"box_threshold": 0.3,
"text_threshold": 0.25,
"max_detections": 100,
"fps": 1.0,
"nms_threshold": 0.5,
"precision": "auto",
"batch_size": 1,
"unload_after": false
}
},
"3": {
"class_type": "PreviewImage",
"inputs": {
"images": [
"2",
2
]
}
},
"4": {
"class_type": "ViewText",
"inputs": {
"text": [
"2",
1
]
}
}
}
@@ -0,0 +1,116 @@
{
"1": {
"class_type": "LoadVideo",
"inputs": {
"file": "tracking_input.mp4"
}
},
"2": {
"class_type": "Video Slice",
"inputs": {
"video": [
"1",
0
],
"start_time": 0.0,
"duration": 2.0,
"strict_duration": false
}
},
"3": {
"class_type": "GetVideoComponents",
"inputs": {
"video": [
"2",
0
]
}
},
"4": {
"class_type": "ImageScaleToMaxDimension",
"inputs": {
"image": [
"3",
0
],
"upscale_method": "area",
"largest_size": 768
}
},
"5": {
"class_type": "ImageFromBatch",
"inputs": {
"image": [
"4",
0
],
"batch_index": 0,
"length": 1
}
},
"6": {
"class_type": "VLMOpenVocabularyDetection",
"inputs": {
"image": [
"5",
0
],
"model": "Grounding DINO Tiny (fast)",
"labels": "person, dog, vehicle",
"box_threshold": 0.3,
"text_threshold": 0.25,
"max_detections": 16,
"fps": [
"3",
2
],
"nms_threshold": 0.5,
"precision": "auto",
"batch_size": 1,
"unload_after": true
}
},
"7": {
"class_type": "VLMSAM2VideoSegmentation",
"inputs": {
"images": [
"4",
0
],
"model": "SAM2.1 Hiera Tiny (fast)",
"seed_frame": 0,
"fps": [
"3",
2
],
"detections": [
"6",
0
],
"mask_threshold": 0.0,
"precision": "auto",
"keep_video_on_cpu": true,
"mask_output": "union_only",
"render_preview": true,
"unload_after": false
}
},
"8": {
"class_type": "VLMTrackReport",
"inputs": {
"tracks": [
"7",
0
]
}
},
"9": {
"class_type": "PreviewImage",
"inputs": {
"images": [
"7",
4
]
}
}
}
@@ -0,0 +1,116 @@
{
"1": {
"class_type": "LoadVideo",
"inputs": {
"file": "tracking_input.mp4"
}
},
"2": {
"class_type": "Video Slice",
"inputs": {
"video": [
"1",
0
],
"start_time": 0.0,
"duration": 2.0,
"strict_duration": false
}
},
"3": {
"class_type": "GetVideoComponents",
"inputs": {
"video": [
"2",
0
]
}
},
"4": {
"class_type": "ImageScaleToMaxDimension",
"inputs": {
"image": [
"3",
0
],
"upscale_method": "area",
"largest_size": 768
}
},
"5": {
"class_type": "CheckpointLoaderSimple",
"inputs": {
"ckpt_name": "sam3.1_multiplex_fp16.safetensors"
}
},
"6": {
"class_type": "CLIPTextEncode",
"inputs": {
"text": "person, dog, vehicle",
"clip": [
"5",
1
]
}
},
"7": {
"class_type": "SAM3_VideoTrack",
"inputs": {
"images": [
"4",
0
],
"model": [
"5",
0
],
"conditioning": [
"6",
0
],
"detection_threshold": 0.5,
"max_objects": 8,
"detect_interval": 1
}
},
"8": {
"class_type": "VLMSAM3TrackAdapter",
"inputs": {
"track_data": [
"7",
0
],
"fps": [
"3",
2
]
}
},
"9": {
"class_type": "VLMTrackReport",
"inputs": {
"tracks": [
"8",
0
]
}
},
"10": {
"class_type": "SAM3_TrackPreview",
"inputs": {
"track_data": [
"8",
1
],
"images": [
"4",
0
],
"opacity": 0.5,
"fps": [
"3",
2
]
}
}
}
+355 -53
View File
@@ -2,7 +2,11 @@
from __future__ import annotations
import hashlib
import json
import math
from dataclasses import dataclass
from numbers import Real
import torch
from PIL import Image, ImageDraw
@@ -24,24 +28,173 @@ from .runtime import (
MODELS = {
"Florence-2 base FT (fast)": "florence-community/Florence-2-base-ft",
"Florence-2 large FT (recommended)": (
"florence-community/Florence-2-large-ft"
),
"Florence-2 large FT (recommended)": ("florence-community/Florence-2-large-ft"),
}
@dataclass(frozen=True)
class FlorenceTaskSpec:
"""Declarative contract for one official Florence-2 task."""
token: str
input_kind: str
output_kind: str
TASKS = {
"Caption": "<CAPTION>",
"Detailed caption": "<DETAILED_CAPTION>",
"More detailed caption": "<MORE_DETAILED_CAPTION>",
"OCR": "<OCR>",
"OCR with regions": "<OCR_WITH_REGION>",
"Object detection": "<OD>",
"Dense region caption": "<DENSE_REGION_CAPTION>",
"Region proposals": "<REGION_PROPOSAL>",
"Referring expression segmentation": "<REFERRING_EXPRESSION_SEGMENTATION>",
"Open vocabulary detection": "<OPEN_VOCABULARY_DETECTION>",
"Caption": FlorenceTaskSpec("<CAPTION>", "none", "text"),
"Detailed caption": FlorenceTaskSpec("<DETAILED_CAPTION>", "none", "text"),
"More detailed caption": FlorenceTaskSpec(
"<MORE_DETAILED_CAPTION>", "none", "text"
),
"OCR": FlorenceTaskSpec("<OCR>", "none", "text"),
"OCR with regions": FlorenceTaskSpec("<OCR_WITH_REGION>", "none", "quad_boxes"),
"Object detection": FlorenceTaskSpec("<OD>", "none", "boxes"),
"Dense region caption": FlorenceTaskSpec("<DENSE_REGION_CAPTION>", "none", "boxes"),
"Caption to phrase grounding": FlorenceTaskSpec(
"<CAPTION_TO_PHRASE_GROUNDING>", "text", "boxes"
),
"Referring expression segmentation": FlorenceTaskSpec(
"<REFERRING_EXPRESSION_SEGMENTATION>", "text", "polygons"
),
"Region to segmentation": FlorenceTaskSpec(
"<REGION_TO_SEGMENTATION>", "region", "polygons"
),
"Open vocabulary detection": FlorenceTaskSpec(
"<OPEN_VOCABULARY_DETECTION>", "text", "mixed"
),
"Region to category": FlorenceTaskSpec("<REGION_TO_CATEGORY>", "region", "text"),
"Region to description": FlorenceTaskSpec(
"<REGION_TO_DESCRIPTION>", "region", "text"
),
"Region to OCR": FlorenceTaskSpec("<REGION_TO_OCR>", "region", "text"),
"Region proposals": FlorenceTaskSpec("<REGION_PROPOSAL>", "none", "boxes"),
}
def _clean_decoded_text(value):
"""Remove generation wrappers without discarding Florence location tokens."""
text = str(value)
for token in ("<s>", "</s>", "<pad>"):
text = text.replace(token, "")
return text.strip()
def _select_region(region, image_index, batch_size):
"""Select one core BOUNDING_BOX for the current image.
Core primitive boxes are dictionaries. Detection nodes may emit either a
flat per-image list or a nested batch list, so both common shapes are
accepted while ambiguous multi-region inputs fail explicitly.
"""
if region is None or isinstance(region, dict):
return region
if not isinstance(region, (list, tuple)):
raise TypeError("region must be a core BOUNDING_BOX dictionary.")
if not region:
return None
if all(isinstance(item, dict) for item in region):
if len(region) == 1:
return region[0]
if len(region) == batch_size:
return region[image_index]
raise ValueError("Region tasks require exactly one BOUNDING_BOX per image.")
if len(region) != batch_size:
raise ValueError("Batched BOUNDING_BOX input must contain one entry per image.")
frame_regions = region[image_index]
if isinstance(frame_regions, dict):
return frame_regions
if not isinstance(frame_regions, (list, tuple)) or len(frame_regions) != 1:
raise ValueError(
"Region tasks require exactly one BOUNDING_BOX per image; "
"select a detection before connecting it."
)
if not isinstance(frame_regions[0], dict):
raise TypeError("Each BOUNDING_BOX entry must be a dictionary.")
return frame_regions[0]
def _encode_region(region, image_size):
"""Encode an absolute-pixel core BOUNDING_BOX as Florence location tokens."""
if not isinstance(region, dict):
raise TypeError("region must be a core BOUNDING_BOX dictionary.")
try:
x = float(region["x"])
y = float(region["y"])
box_width = float(region["width"])
box_height = float(region["height"])
except KeyError as exc:
raise ValueError("region must contain x, y, width, and height.") from exc
except (TypeError, ValueError) as exc:
raise ValueError("region coordinates must be numeric.") from exc
values = (x, y, box_width, box_height)
if not all(math.isfinite(value) for value in values):
raise ValueError("region coordinates must be finite.")
if box_width <= 0 or box_height <= 0:
raise ValueError("region width and height must be greater than zero.")
image_width, image_height = image_size
if image_width <= 0 or image_height <= 0:
raise ValueError("image dimensions must be greater than zero.")
x0 = max(0.0, min(float(image_width), x))
y0 = max(0.0, min(float(image_height), y))
x1 = max(0.0, min(float(image_width), x + box_width))
y1 = max(0.0, min(float(image_height), y + box_height))
if x1 <= x0 or y1 <= y0:
raise ValueError("region does not overlap the input image.")
coordinates = (
x0 / image_width,
y0 / image_height,
x1 / image_width,
y1 / image_height,
)
bins = [
max(0, min(999, math.floor(coordinate * 1000))) for coordinate in coordinates
]
return "".join(f"<loc_{value}>" for value in bins)
def _task_extra_input(task_name, text_input, region, image_size):
"""Validate and prepare the optional suffix for a Florence task prompt."""
try:
spec = TASKS[task_name]
except KeyError as exc:
raise ValueError(f"Unsupported Florence-2 task: {task_name}") from exc
text = (text_input or "").strip()
if spec.input_kind == "none":
if text:
raise ValueError(f"{task_name} does not accept text input.")
if region is not None:
raise ValueError(f"{task_name} does not accept a region input.")
return ""
if spec.input_kind == "text":
if not text:
raise ValueError(f"{task_name} requires text input.")
if region is not None:
raise ValueError(f"{task_name} does not accept a region input.")
return text
if spec.input_kind == "region":
if text:
raise ValueError(
f"{task_name} uses the region input and does not accept text."
)
if region is None:
raise ValueError(f"{task_name} requires a connected BOUNDING_BOX region.")
return _encode_region(region, image_size)
raise RuntimeError(f"Unknown Florence task input kind: {spec.input_kind}")
class FlorencePredictor:
def __init__(self, model_label):
transformers = require_module("transformers")
@@ -66,9 +219,7 @@ class FlorencePredictor:
def run(self, image, task_token, text, max_new_tokens, beams):
prompt = task_token + (text.strip() if text.strip() else "")
inputs = self.processor(
text=prompt, images=image, return_tensors="pt"
)
inputs = self.processor(text=prompt, images=image, return_tensors="pt")
model = self.handle.ensure_loaded()
device = model_device(model)
inputs = move_inputs(inputs, device, floating_dtype=self.dtype)
@@ -80,9 +231,7 @@ class FlorencePredictor:
do_sample=False,
early_stopping=int(beams) > 1,
)
raw = self.processor.batch_decode(
generated, skip_special_tokens=False
)[0]
raw = self.processor.batch_decode(generated, skip_special_tokens=False)[0]
parsed = self.processor.post_process_generation(
raw, task=task_token, image_size=image.size
)
@@ -95,39 +244,156 @@ def _json_default(value):
return str(value)
_SPATIAL_KEYS = frozenset(
{
"bboxes",
"quad_boxes",
"polygons",
"labels",
"bboxes_labels",
"polygons_labels",
}
)
def _spatial_result(parsed):
if not isinstance(parsed, dict):
return {}
if _SPATIAL_KEYS.intersection(parsed):
return parsed
result = next(iter(parsed.values()), {})
return result if isinstance(result, dict) else {}
def _stable_color(kind, index, label):
key = f"{kind}:{index}:{label}".encode("utf-8", errors="replace")
digest = hashlib.blake2b(key, digest_size=3).digest()
return tuple(64 + channel % 192 for channel in digest)
def _points(values, image_size):
if not isinstance(values, (list, tuple)) or len(values) < 6:
return []
width, height = image_size
points = []
for index in range(0, len(values) - 1, 2):
x, y = values[index], values[index + 1]
if not isinstance(x, Real) or not isinstance(y, Real):
return []
if not math.isfinite(float(x)) or not math.isfinite(float(y)):
return []
points.append(
(
max(0, min(width - 1, round(float(x)))),
max(0, min(height - 1, round(float(y)))),
)
)
return points
def _box(values, image_size):
if not isinstance(values, (list, tuple)) or len(values) < 4:
return None
if not all(isinstance(value, Real) for value in values[:4]):
return None
coordinates = [float(value) for value in values[:4]]
if not all(math.isfinite(value) for value in coordinates):
return None
x0, y0, x1, y1 = coordinates
x0, x1 = sorted((x0, x1))
y0, y1 = sorted((y0, y1))
width, height = image_size
x0 = max(0, min(width - 1, round(x0)))
x1 = max(0, min(width - 1, round(x1)))
y0 = max(0, min(height - 1, round(y0)))
y1 = max(0, min(height - 1, round(y1)))
if x1 <= x0 or y1 <= y0:
return None
return x0, y0, x1, y1
def _polygon_list(group):
if not isinstance(group, (list, tuple)) or not group:
return []
if isinstance(group[0], Real):
return [group]
return [item for item in group if isinstance(item, (list, tuple))]
def _label_with_score(labels, scores, index):
label = str(labels[index]) if index < len(labels) else ""
if index < len(scores) and isinstance(scores[index], Real):
score = f"{float(scores[index]):.3f}"
return f"{label} {score}".strip()
return label
def _draw_label(draw, position, text, color, image_size):
if not text:
return
x, y = position
try:
left, top, right, bottom = draw.textbbox((0, 0), text)
text_width, text_height = right - left, bottom - top
except AttributeError:
text_width, text_height = draw.textlength(text), 11
width, height = image_size
x = max(0, min(width - text_width - 4, x))
y = max(0, min(height - text_height - 4, y))
background = (0, 0, 0) if sum(color) > 360 else (255, 255, 255)
foreground = (255, 255, 255) if background == (0, 0, 0) else (0, 0, 0)
draw.rectangle(
(x, y, x + text_width + 4, y + text_height + 4),
fill=background,
)
draw.text((x + 2, y + 2), text, fill=foreground)
def _visualize(image, parsed):
result = next(iter(parsed.values()), parsed) if isinstance(parsed, dict) else {}
result = _spatial_result(parsed)
mask = Image.new("L", image.size, 0)
visual = image.copy().convert("RGB")
mask_draw = ImageDraw.Draw(mask)
draw = ImageDraw.Draw(visual)
labels = result.get("labels", []) if isinstance(result, dict) else []
width = max(2, min(8, round(min(image.size) / 256 * 3)))
labels = result.get("labels", [])
scores = result.get("scores", [])
for index, box in enumerate(result.get("bboxes", [])):
box = [float(value) for value in box]
draw.rectangle(box, outline="#00ff88", width=3)
if index < len(labels):
draw.text((box[0] + 3, box[1] + 3), str(labels[index]), fill="#00ff88")
box_labels = result.get("bboxes_labels", labels)
for index, values in enumerate(result.get("bboxes", [])):
box = _box(values, image.size)
if box is None:
continue
label = _label_with_score(box_labels, scores, index)
color = _stable_color("box", index, label)
mask_draw.rectangle(box, fill=255)
draw.rectangle(box, outline=color, width=width)
_draw_label(draw, (box[0], box[1]), label, color, image.size)
for quad in result.get("quad_boxes", []):
points = [
(float(quad[index]), float(quad[index + 1]))
for index in range(0, len(quad), 2)
]
draw.line(points + [points[0]], fill="#00c8ff", width=3)
for index, values in enumerate(result.get("quad_boxes", [])):
points = _points(values, image.size)
if len(points) < 3:
continue
label = _label_with_score(labels, scores, index)
color = _stable_color("quad", index, label)
mask_draw.polygon(points, fill=255)
draw.line(points + [points[0]], fill=color, width=width)
_draw_label(draw, points[0], label, color, image.size)
polygons = result.get("polygons", [])
for group in polygons:
# Florence may return either one flat polygon or a list of polygons.
groups = [group] if group and isinstance(group[0], (int, float)) else group
for polygon in groups:
points = [
(float(polygon[index]), float(polygon[index + 1]))
for index in range(0, len(polygon), 2)
]
if len(points) >= 3:
mask_draw.polygon(points, fill=255)
draw.line(points + [points[0]], fill="#ff4da6", width=3)
polygon_labels = result.get("polygons_labels", labels)
for index, group in enumerate(result.get("polygons", [])):
label = _label_with_score(polygon_labels, scores, index)
color = _stable_color("polygon", index, label)
label_drawn = False
for polygon in _polygon_list(group):
points = _points(polygon, image.size)
if len(points) < 3:
continue
mask_draw.polygon(points, fill=255)
draw.line(points + [points[0]], fill=color, width=width)
if not label_drawn:
_draw_label(draw, points[0], label, color, image.size)
label_drawn = True
return mask, visual
@@ -143,7 +409,10 @@ class Florence2(CachedModelNode):
{
"default": "",
"multiline": True,
"tooltip": "Required for referring-expression and open-vocabulary tasks.",
"tooltip": (
"Required only for phrase grounding, referring-expression "
"segmentation, and open-vocabulary detection."
),
},
),
"model": (
@@ -158,6 +427,15 @@ class Florence2(CachedModelNode):
},
"optional": {
"unload_after": ("BOOLEAN", {"default": False}),
"region": (
"BOUNDING_BOX",
{
"tooltip": (
"Core bounding box input required by Region to "
"Segmentation/Category/Description/OCR."
)
},
),
},
}
@@ -175,28 +453,52 @@ class Florence2(CachedModelNode):
max_new_tokens,
beams,
unload_after=False,
region=None,
):
predictor = self.get_or_create_model(
model, lambda: FlorencePredictor(model)
)
images = tensor_batch_to_pil(image)
if not images:
raise ValueError("Florence-2 requires at least one input image.")
try:
spec = TASKS[task]
except KeyError as exc:
raise ValueError(f"Unsupported Florence-2 task: {task}") from exc
extra_inputs = []
for index, pil_image in enumerate(images):
selected_region = _select_region(region, index, len(images))
extra_inputs.append(
_task_extra_input(
task,
text_input,
selected_region,
pil_image.size,
)
)
predictor = self.get_or_create_model(model, lambda: FlorencePredictor(model))
texts, records, masks, visuals = [], [], [], []
try:
for pil_image in tensor_batch_to_pil(image):
for pil_image, extra_input in zip(images, extra_inputs):
raw, parsed = predictor.run(
pil_image,
TASKS[task],
text_input,
spec.token,
extra_input,
max_new_tokens,
beams,
)
texts.append(raw)
texts.append(_clean_decoded_text(raw))
records.append(parsed)
mask, visual = _visualize(pil_image, parsed)
masks.append(pil_mask_to_tensor(mask))
visuals.append(pil_to_tensor(visual))
return (
batch_text(texts),
json.dumps(records, ensure_ascii=False, default=_json_default),
json.dumps(
records,
ensure_ascii=False,
default=_json_default,
sort_keys=True,
),
torch.cat(masks),
torch.cat(visuals),
)
+413
View File
@@ -0,0 +1,413 @@
"""Dependency-light geometry, mask, color, and association primitives."""
from __future__ import annotations
import colorsys
import hashlib
import math
from dataclasses import dataclass
from typing import Iterable, Mapping
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",
]
+532
View File
@@ -0,0 +1,532 @@
"""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)
model = model_class.from_pretrained(model_path, dtype=dtype)
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",
}
+145 -50
View File
@@ -7,11 +7,16 @@ 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,
@@ -35,7 +40,11 @@ def _make_handle(
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,
@@ -43,6 +52,7 @@ def _make_handle(
n_threads=n_threads,
chat_handler_factory=_clip_factory(clip),
seed=seed,
**options,
)
@@ -59,13 +69,6 @@ def _vision_messages(system_msg: str, prompt: str, data_uri: str):
]
def _content(response: dict[str, Any]) -> str:
try:
return str(response["choices"][0]["message"]["content"])
except (KeyError, IndexError, TypeError) as exc:
raise RuntimeError(f"llama.cpp returned an unexpected response: {response!r}") from exc
def _run_batch(
image,
model,
@@ -78,12 +81,10 @@ def _run_batch(
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)
),
messages=_vision_messages(system_msg, prompt, image_data_uri(pil_image)),
**generation,
)
responses.append(_content(response))
responses.append(llama_chat_content(response))
return batch_text(responses)
@@ -92,23 +93,27 @@ class LLavaLoader:
def INPUT_TYPES(cls):
return {
"required": {
"ckpt_name": (
folder_paths.get_filename_list("LLavacheckpoints"),
),
"ckpt_name": (folder_paths.get_filename_list("LLavacheckpoints"),),
"max_ctx": (
"INT",
{"default": 4096, "min": 128, "max": 131072, "step": 64},
),
"gpu_layers": (
"INT",
{"default": 27, "min": -1, "max": 1000, "step": 1},
{"default": -1, "min": -1, "max": 1000, "step": 1},
),
"n_threads": (
"INT",
{"default": 8, "min": 1, "max": 256, "step": 1},
{
"default": default_llama_threads(),
"min": 1,
"max": 256,
"step": 1,
},
),
"clip": ("CUSTOM", {"default": ""}),
}
},
"optional": llama_runtime_input_types(),
}
RETURN_TYPES = ("CUSTOM",)
@@ -117,12 +122,37 @@ class LLavaLoader:
CATEGORY = "VLM Nodes/LLava"
def load_llava_checkpoint(
self, ckpt_name, max_ctx, gpu_layers, n_threads, clip
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
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,
),
),
)
@@ -132,10 +162,14 @@ class LlavaClipLoader:
def INPUT_TYPES(cls):
return {
"required": {
"clip_name": (
folder_paths.get_filename_list("LLavacheckpoints"),
)
}
"clip_name": (folder_paths.get_filename_list("LLavacheckpoints"),),
},
"optional": {
"handler": (
list(LLAMA_VISION_HANDLER_CHOICES),
{"default": "Auto (GGUF chat template)"},
),
},
}
RETURN_TYPES = ("CUSTOM",)
@@ -143,8 +177,8 @@ class LlavaClipLoader:
FUNCTION = "load_clip_checkpoint"
CATEGORY = "VLM Nodes/LLava"
def load_clip_checkpoint(self, clip_name):
return (LlavaClipConfig(resolve_model_path(clip_name)),)
def load_clip_checkpoint(self, clip_name, handler="LLaVA 1.5"):
return (LlavaClipConfig(resolve_model_path(clip_name), handler),)
class LLavaSamplerSimple:
@@ -276,6 +310,8 @@ class _CachedLlavaBase:
gpu_layers,
n_threads,
seed=42,
handler="LLaVA 1.5",
**runtime_options,
):
key = (
ckpt_name,
@@ -284,10 +320,12 @@ class _CachedLlavaBase:
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))
clip = LlavaClipConfig(resolve_model_path(clip_name), handler)
self._handle = _make_handle(
ckpt_name,
max_ctx,
@@ -295,6 +333,7 @@ class _CachedLlavaBase:
n_threads,
clip,
seed=seed,
runtime_options=runtime_options,
)
self._key = key
return self._handle
@@ -311,23 +350,24 @@ class LLavaOptionalMemoryFreeSimple(_CachedLlavaBase):
def INPUT_TYPES(cls):
return {
"required": {
"ckpt_name": (
folder_paths.get_filename_list("LLavacheckpoints"),
),
"clip_name": (
folder_paths.get_filename_list("LLavacheckpoints"),
),
"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": 27, "min": -1, "max": 1000, "step": 1},
{"default": -1, "min": -1, "max": 1000, "step": 1},
),
"n_threads": (
"INT",
{"default": 8, "min": 1, "max": 256, "step": 1},
{
"default": default_llama_threads(),
"min": 1,
"max": 256,
"step": 1,
},
),
"image": ("IMAGE",),
"prompt": ("STRING", {"default": "", "multiline": True}),
@@ -336,7 +376,14 @@ class LLavaOptionalMemoryFreeSimple(_CachedLlavaBase):
{"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(),
},
}
RETURN_TYPES = ("STRING",)
@@ -354,9 +401,32 @@ class LLavaOptionalMemoryFreeSimple(_CachedLlavaBase):
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
ckpt_name,
clip_name,
max_ctx,
gpu_layers,
n_threads,
handler=handler,
**options,
)
try:
result = _run_batch(
@@ -375,32 +445,29 @@ 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"),
),
"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": 27, "min": -1, "max": 1000, "step": 1},
{"default": -1, "min": -1, "max": 1000, "step": 1},
),
"n_threads": (
"INT",
{"default": 8, "min": 1, "max": 256, "step": 1},
{
"default": default_llama_threads(),
"min": 1,
"max": 256,
"step": 1,
},
),
"image": ("IMAGE",),
"system_msg": (
"STRING",
{
"default": (
"You are an assistant who accurately describes images."
)
},
{"default": ("You are an assistant who accurately describes images.")},
),
"prompt": ("STRING", {"default": "", "multiline": True}),
"max_tokens": (
@@ -431,7 +498,16 @@ class LLavaOptionalMemoryFreeAdvanced(_CachedLlavaBase):
"seed": ("INT", {"default": 42, "step": 1}),
"unload": ("BOOLEAN", {"default": False}),
}
return {"required": required}
return {
"required": required,
"optional": {
"handler": (
list(LLAMA_VISION_HANDLER_CHOICES),
{"default": "Auto (GGUF chat template)"},
),
**llama_runtime_input_types(),
},
}
RETURN_TYPES = ("STRING",)
FUNCTION = "generate_text_advanced"
@@ -456,7 +532,24 @@ class LLavaOptionalMemoryFreeAdvanced(_CachedLlavaBase):
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,
@@ -464,6 +557,8 @@ class LLavaOptionalMemoryFreeAdvanced(_CachedLlavaBase):
gpu_layers,
n_threads,
seed,
handler,
**options,
)
try:
result = _run_batch(
+40 -22
View File
@@ -5,10 +5,14 @@ from __future__ import annotations
from .runtime import (
CachedModelNode,
LlamaHandle,
LlavaClipConfig,
batch_text,
default_llama_threads,
hf_download,
image_data_uri,
require_module,
llama_chat_content,
llama_runtime_input_types,
llama_runtime_options,
tensor_batch_to_pil,
)
@@ -30,6 +34,7 @@ class MiniCPMPredictor:
context_length,
gpu_layers,
n_threads,
runtime_options=None,
):
model_path = hf_download(
MODEL_REPO,
@@ -42,19 +47,10 @@ class MiniCPMPredictor:
"minicpm-v-2_6-gguf",
)
def create_handler():
chat = require_module(
"llama_cpp.llama_chat_format", "llama-cpp-python"
)
handler_class = getattr(chat, "MiniCPMv26ChatHandler", None)
if handler_class is None:
raise RuntimeError(
"Your llama-cpp-python build is too old for MiniCPM-V 2.6. "
"Install a current CUDA or CPU wheel."
)
return handler_class(
clip_model_path=str(projector_path), verbose=False
)
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,
@@ -62,6 +58,8 @@ class MiniCPMPredictor:
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):
@@ -87,9 +85,7 @@ class MiniCPMPredictor:
"content": [
{
"type": "image_url",
"image_url": {
"url": image_data_uri(image)
},
"image_url": {"url": image_data_uri(image)},
},
{"type": "text", "text": prompt},
],
@@ -101,9 +97,7 @@ class MiniCPMPredictor:
top_k=int(top_k),
repeat_penalty=float(repeat_penalty),
)
results.append(
str(response["choices"][0]["message"]["content"]).strip()
)
results.append(llama_chat_content(response))
return batch_text(results)
@@ -149,13 +143,18 @@ class MiniCPMNode(CachedModelNode):
),
"n_threads": (
"INT",
{"default": 8, "min": 1, "max": 256},
{
"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(),
},
}
@@ -174,15 +173,33 @@ class MiniCPMNode(CachedModelNode):
top_k=100,
repeat_penalty=1.05,
gpu_layers=-1,
n_threads=8,
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,
@@ -191,6 +208,7 @@ class MiniCPMNode(CachedModelNode):
context_length,
gpu_layers,
n_threads,
options,
),
)
try:
+119 -9
View File
@@ -7,8 +7,9 @@ small and large VLM families while keeping downloads and VRAM allocation lazy.
from __future__ import annotations
import threading
from dataclasses import dataclass
from typing import Any
from typing import Any, Callable
import torch
@@ -198,6 +199,32 @@ MEMORY_MODES = (
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)
@@ -218,6 +245,7 @@ class ModernVLMPredictor:
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)
@@ -416,6 +444,7 @@ class ModernVLMPredictor:
video_frames=None,
fps: float = 1.0,
enable_thinking: bool = False,
stream_callback: Callable[[str], None] | None = None,
) -> str:
primary_images = (
tensor_batch_to_pil(images) if images is not None else []
@@ -490,16 +519,78 @@ class ModernVLMPredictor:
generation.update(
temperature=float(temperature), top_p=float(top_p)
)
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,
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,
)[0].strip()
)
)
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)
@@ -557,7 +648,18 @@ class ModernVLM(CachedModelNode):
),
"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",)
@@ -580,7 +682,14 @@ class ModernVLM(CachedModelNode):
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"
@@ -605,6 +714,7 @@ class ModernVLM(CachedModelNode):
video_frames,
fps,
enable_thinking,
stream_callback,
),
)
finally:
+429 -70
View File
@@ -16,6 +16,7 @@ import io
import logging
import os
import platform
import re
import threading
from contextlib import nullcontext
from dataclasses import dataclass
@@ -23,14 +24,27 @@ from importlib import metadata
from pathlib import Path
from typing import Any, Callable, Iterable, Mapping
import folder_paths
import numpy as np
import torch
from PIL import Image
import folder_paths
LOGGER = logging.getLogger("ComfyUI_VLM_nodes")
GGUF_EXTENSIONS = {".gguf"}
LLAMA_FLASH_ATTENTION_CHOICES = ("Auto", "Enabled", "Disabled")
LLAMA_SPLIT_MODE_CHOICES = ("Layer", "Row", "Single GPU")
LLAMA_VISION_HANDLER_CHOICES = (
"Auto (GGUF chat template)",
"LLaVA 1.5",
"LLaVA 1.6",
"MiniCPM-V 2.6",
"Moondream2",
"NanoLLaVA",
"Qwen2.5-VL",
"Gemma 4",
"Llama 3 Vision Alpha",
"Obsidian",
)
class OptionalDependencyError(RuntimeError):
@@ -96,9 +110,7 @@ def resolve_model_path(filename: str) -> Path:
return Path(getter("LLavacheckpoints", filename))
path = folder_paths.get_full_path("LLavacheckpoints", filename)
if path is None:
raise FileNotFoundError(
f"Model '{filename}' was not found in {model_root()}."
)
raise FileNotFoundError(f"Model '{filename}' was not found in {model_root()}.")
return Path(path)
@@ -126,16 +138,12 @@ def snapshot_download(repo_id: str, subdirectory: str, **kwargs: Any) -> Path:
}
download_kwargs.update(kwargs)
# local_dir_use_symlinks was removed from newer huggingface-hub versions.
if "local_dir_use_symlinks" in inspect.signature(
hub.snapshot_download
).parameters:
if "local_dir_use_symlinks" in inspect.signature(hub.snapshot_download).parameters:
download_kwargs.setdefault("local_dir_use_symlinks", False)
return Path(hub.snapshot_download(**download_kwargs))
def hf_download(
repo_id: str, filename: str, subdirectory: str, **kwargs: Any
) -> Path:
def hf_download(repo_id: str, filename: str, subdirectory: str, **kwargs: Any) -> Path:
hub = require_module("huggingface_hub", "huggingface-hub")
destination = model_cache_dir(subdirectory)
download_kwargs = {
@@ -144,9 +152,7 @@ def hf_download(
"local_dir": str(destination),
}
download_kwargs.update(kwargs)
if "local_dir_use_symlinks" in inspect.signature(
hub.hf_hub_download
).parameters:
if "local_dir_use_symlinks" in inspect.signature(hub.hf_hub_download).parameters:
download_kwargs.setdefault("local_dir_use_symlinks", False)
return Path(hub.hf_hub_download(**download_kwargs))
@@ -179,9 +185,7 @@ def tensor_to_pil(image: torch.Tensor, index: int = 0) -> Image.Image:
)
if value.numel() and (value.max() > 1.0 or value.min() < 0.0):
value = value / 255.0
array = (
value.clamp(0.0, 1.0).mul(255.0).round().to(torch.uint8).numpy()
)
array = value.clamp(0.0, 1.0).mul(255.0).round().to(torch.uint8).numpy()
if array.shape[-1] == 1:
array = np.repeat(array, 3, axis=-1)
elif array.shape[-1] == 4:
@@ -223,6 +227,23 @@ def batch_text(responses: Iterable[str]) -> str:
)
def llama_chat_content(response: Any) -> str:
"""Extract non-empty text from llama.cpp's OpenAI-compatible response."""
try:
content = response["choices"][0]["message"]["content"]
except (KeyError, IndexError, TypeError) as exc:
raise RuntimeError(
f"llama.cpp returned an unexpected response: {response!r}"
) from exc
if content is None or not str(content).strip():
raise RuntimeError(
"llama.cpp returned an empty response. Check that the GGUF chat "
"template/vision handler matches the model and that max_tokens is positive."
)
return str(content).strip()
def execution_device() -> torch.device:
"""Return ComfyUI's selected device, with portable standalone fallbacks."""
@@ -314,26 +335,14 @@ def torch_dtype(
if requested in {"float32", "fp32"}:
return torch.float32
if requested in {"float16", "fp16"}:
return (
torch.float16
if device.type in {"cuda", "mps", "xpu"}
else torch.float32
)
return torch.float16 if device.type in {"cuda", "mps", "xpu"} else torch.float32
if requested in {"bfloat16", "bf16"}:
if supports_bfloat16(device):
return torch.bfloat16
return (
torch.float16
if device.type in {"cuda", "mps", "xpu"}
else torch.float32
)
return torch.float16 if device.type in {"cuda", "mps", "xpu"} else torch.float32
if supports_bfloat16(device):
return torch.bfloat16
return (
torch.float16
if device.type in {"cuda", "mps", "xpu"}
else torch.float32
)
return torch.float16 if device.type in {"cuda", "mps", "xpu"} else torch.float32
def _release_tuple(distribution: str) -> tuple[int, ...]:
@@ -361,10 +370,10 @@ def require_quantization_backend(feature: str) -> torch.device:
f"{feature} is not supported on the active {backend} backend. "
"Use ComfyUI managed precision or CPU mode."
)
if (
platform.system() == "Darwin"
and platform.machine().lower() not in {"arm64", "aarch64"}
):
if platform.system() == "Darwin" and platform.machine().lower() not in {
"arm64",
"aarch64",
}:
raise RuntimeError(
f"{feature} needs bitsandbytes, which has no official Intel-macOS "
"wheel. Use ComfyUI managed precision, or use an Apple Silicon Mac."
@@ -417,6 +426,7 @@ def runtime_diagnostics() -> dict[str, Any]:
"torch_cuda": getattr(getattr(torch, "version", None), "cuda", None),
"torch_hip": getattr(getattr(torch, "version", None), "hip", None),
"packages": packages,
"llama_cpp": llama_cpp_diagnostics(),
}
@@ -484,13 +494,9 @@ class ManagedTorchModel:
self.load_device = load_device or model_management.get_torch_device()
self.offload_device = offload_device or (
torch.device("cpu")
if self.load_device.type != "cpu"
else self.load_device
)
self.model = _ManagedModelAdapter(
model.eval(), self.offload_device
torch.device("cpu") if self.load_device.type != "cpu" else self.load_device
)
self.model = _ManagedModelAdapter(model.eval(), self.offload_device)
self.processor = processor
self.patcher = ModelPatcher(
self.model,
@@ -608,19 +614,257 @@ def reserve_external_vram(memory_required: int) -> None:
LOGGER.debug("Could not reserve VRAM through ComfyUI: %s", exc)
def default_llama_threads() -> int:
"""A portable generation-thread default that avoids oversubscribing ComfyUI."""
return max(1, min(16, (os.cpu_count() or 4) // 2))
def llama_runtime_input_types() -> dict[str, tuple[Any, ...]]:
"""Advanced llama.cpp inputs shared by every GGUF loader."""
return {
"n_batch": (
"INT",
{
"default": 512,
"min": 1,
"max": 8192,
"step": 1,
"tooltip": "Logical prompt batch. Lower this if context loading runs out of memory.",
},
),
"n_ubatch": (
"INT",
{
"default": 512,
"min": 1,
"max": 8192,
"step": 1,
"tooltip": "Physical prompt micro-batch. Never exceeds n_batch.",
},
),
"flash_attention": (
list(LLAMA_FLASH_ATTENTION_CHOICES),
{
"default": "Auto",
"tooltip": "Auto enables llama.cpp flash attention only with accelerator offload and safely retries without it when unsupported.",
},
),
"use_mmap": (
"BOOLEAN",
{
"default": True,
"tooltip": "Memory-map GGUF weights when the installed backend supports it.",
},
),
"split_mode": (
list(LLAMA_SPLIT_MODE_CHOICES),
{
"default": "Layer",
"tooltip": "How llama.cpp distributes tensors across multiple accelerators.",
},
),
"main_gpu": (
"INT",
{
"default": 0,
"min": 0,
"max": 31,
"step": 1,
},
),
"tensor_split": (
"STRING",
{
"default": "",
"tooltip": "Optional comma-separated accelerator proportions, for example 0.6,0.4.",
},
),
}
def llama_runtime_options(
*,
n_batch: int = 512,
n_ubatch: int = 512,
flash_attention: str = "Auto",
use_mmap: bool = True,
split_mode: str = "Layer",
main_gpu: int = 0,
tensor_split: str = "",
) -> dict[str, Any]:
"""Normalize node inputs into keyword arguments accepted by LlamaHandle."""
return {
"n_batch": int(n_batch),
"n_ubatch": int(n_ubatch),
"flash_attention": str(flash_attention),
"use_mmap": bool(use_mmap),
"split_mode": str(split_mode),
"main_gpu": int(main_gpu),
"tensor_split": str(tensor_split),
}
def _call_llama_probe(module: Any, name: str) -> bool | None:
for source in (module, getattr(module, "llama_cpp", None)):
probe = getattr(source, name, None)
if callable(probe):
try:
return bool(probe())
except Exception:
return None
return None
def _llama_system_info(module: Any) -> str | None:
for source in (module, getattr(module, "llama_cpp", None)):
probe = getattr(source, "llama_print_system_info", None)
if callable(probe):
try:
value = probe()
if isinstance(value, bytes):
value = value.decode("utf-8", errors="replace")
return str(value).strip() or None
except Exception:
return None
return None
def _llama_backend_labels(system_info: str | None) -> list[str]:
if not system_info:
return []
value = system_info.upper()
candidates = (
("CUDA", "cuda"),
("ROCM", "rocm"),
("HIP", "hip"),
("METAL", "metal"),
("VULKAN", "vulkan"),
("SYCL", "sycl"),
("CANN", "cann"),
("OPENCL", "opencl"),
("BLAS", "blas"),
)
return [label for marker, label in candidates if marker in value]
def llama_cpp_diagnostics(module: Any | None = None) -> dict[str, Any]:
"""Describe the installed llama.cpp build without allocating model memory."""
try:
version = metadata.version("llama-cpp-python")
except metadata.PackageNotFoundError:
version = None
if module is None:
try:
module = importlib.import_module("llama_cpp")
except Exception as exc:
return {
"installed": version is not None,
"version": version,
"import_error": f"{type(exc).__name__}: {exc}",
"gpu_offload": None,
"mmap": None,
"mlock": None,
"backends": [],
"system_info": None,
}
system_info = _llama_system_info(module)
return {
"installed": True,
"version": getattr(module, "__version__", None) or version,
"import_error": None,
"gpu_offload": _call_llama_probe(module, "llama_supports_gpu_offload"),
"mmap": _call_llama_probe(module, "llama_supports_mmap"),
"mlock": _call_llama_probe(module, "llama_supports_mlock"),
"backends": _llama_backend_labels(system_info),
"system_info": system_info,
}
def _parse_tensor_split(value: str | Iterable[float] | None) -> list[float] | None:
if value is None or value == "":
return None
if isinstance(value, str):
parts = [part.strip() for part in value.split(",") if part.strip()]
try:
values = [float(part) for part in parts]
except ValueError as exc:
raise ValueError(
"tensor_split must be comma-separated numbers, for example 0.6,0.4."
) from exc
else:
values = [float(part) for part in value]
if not values or any(not np.isfinite(part) or part < 0 for part in values):
raise ValueError("tensor_split values must be finite and non-negative.")
if not any(values):
raise ValueError("tensor_split must give at least one accelerator a weight.")
return values
def _resolve_split_mode(module: Any, value: str) -> Any:
constants = {
"Layer": "LLAMA_SPLIT_MODE_LAYER",
"Row": "LLAMA_SPLIT_MODE_ROW",
"Single GPU": "LLAMA_SPLIT_MODE_NONE",
}
if value not in constants:
raise ValueError(
f"Unknown llama.cpp split mode {value!r}; choose one of {LLAMA_SPLIT_MODE_CHOICES}."
)
return getattr(module, constants[value], None)
@dataclass(frozen=True)
class LlavaClipConfig:
model_path: Path
# Preserve the prompt format used by workflows saved before handler
# selection existed. New loader widgets explicitly pass the Auto choice.
handler: str = "LLaVA 1.5"
def create(self):
def create(self, *, use_gpu: bool = True):
module = require_module("llama_cpp.llama_chat_format", "llama-cpp-python")
return module.Llava15ChatHandler(
clip_model_path=str(self.model_path), verbose=False
)
handlers = {
"LLaVA 1.5": "Llava15ChatHandler",
"LLaVA 1.6": "Llava16ChatHandler",
"MiniCPM-V 2.6": "MiniCPMv26ChatHandler",
"Moondream2": "MoondreamChatHandler",
"NanoLLaVA": "NanoLlavaChatHandler",
"Qwen2.5-VL": "Qwen25VLChatHandler",
"Gemma 4": "Gemma4ChatHandler",
"Llama 3 Vision Alpha": "Llama3VisionAlphaChatHandler",
"Obsidian": "ObsidianChatHandler",
}
if self.handler == "Auto (GGUF chat template)":
handler_class = getattr(module, "MTMDChatHandler", None)
if handler_class is not None:
return handler_class(
clip_model_path=str(self.model_path),
verbose=False,
use_gpu=use_gpu,
)
handler_name = "Llava15ChatHandler"
else:
try:
handler_name = handlers[self.handler]
except KeyError as exc:
raise ValueError(
f"Unknown vision handler {self.handler!r}; choose one of "
f"{LLAMA_VISION_HANDLER_CHOICES}."
) from exc
handler_class = getattr(module, handler_name, None)
if handler_class is None:
raise RuntimeError(
f"The installed llama-cpp-python does not provide {handler_name}. "
"Install a current wheel for CUDA, ROCm/HIP, Metal, Vulkan, SYCL, or CPU."
)
return handler_class(clip_model_path=str(self.model_path), verbose=False)
class LlamaHandle:
"""Lazy llama.cpp handle that owns and closes its exact GPU allocations."""
"""Lazy llama.cpp handle that owns and closes its accelerator allocations."""
def __init__(
self,
@@ -630,8 +874,17 @@ class LlamaHandle:
n_gpu_layers: int,
n_threads: int,
chat_format: str | None = None,
chat_handler_factory: Callable[[], Any] | None = None,
chat_handler_factory: Callable[..., Any] | None = None,
seed: int = 42,
n_batch: int = 512,
n_ubatch: int = 512,
n_threads_batch: int | None = None,
flash_attention: str = "Auto",
use_mmap: bool = True,
split_mode: str = "Layer",
main_gpu: int = 0,
tensor_split: str | Iterable[float] | None = None,
projector_path: Path | None = None,
) -> None:
self.model_path = Path(model_path)
self.n_ctx = int(n_ctx)
@@ -640,6 +893,21 @@ class LlamaHandle:
self.chat_format = chat_format
self.chat_handler_factory = chat_handler_factory
self.seed = int(seed)
self.n_batch = max(1, int(n_batch))
self.n_ubatch = max(1, int(n_ubatch))
self.n_threads_batch = (
None if n_threads_batch is None else max(1, int(n_threads_batch))
)
if flash_attention not in LLAMA_FLASH_ATTENTION_CHOICES:
raise ValueError(
f"flash_attention must be one of {LLAMA_FLASH_ATTENTION_CHOICES}."
)
self.flash_attention = flash_attention
self.use_mmap = bool(use_mmap)
self.split_mode = split_mode
self.main_gpu = max(0, int(main_gpu))
self.tensor_split = _parse_tensor_split(tensor_split)
self.projector_path = None if projector_path is None else Path(projector_path)
self._llm = None
self._chat_handler = None
self._lock = threading.RLock()
@@ -652,8 +920,56 @@ class LlamaHandle:
self.n_gpu_layers,
self.n_threads,
self.chat_format,
self.seed,
self.n_batch,
self.n_ubatch,
self.n_threads_batch,
self.flash_attention,
self.use_mmap,
self.split_mode,
self.main_gpu,
tuple(self.tensor_split or ()),
str(self.projector_path) if self.projector_path else None,
)
@staticmethod
def _filter_constructor_kwargs(llama_class: Any, requested: dict[str, Any]):
signature = inspect.signature(llama_class.__init__)
named_parameters = {
name
for name, parameter in signature.parameters.items()
if name != "self"
and parameter.kind
in {
inspect.Parameter.POSITIONAL_OR_KEYWORD,
inspect.Parameter.KEYWORD_ONLY,
}
}
accepts_kwargs = any(
parameter.kind == inspect.Parameter.VAR_KEYWORD
for parameter in signature.parameters.values()
)
return {
key: value
for key, value in requested.items()
if value is not None
and (
key in named_parameters
or (accepts_kwargs and not named_parameters)
)
}
def _create_chat_handler(self, *, use_gpu: bool):
if self.chat_handler_factory is None:
return None
try:
signature = inspect.signature(self.chat_handler_factory)
if "use_gpu" in signature.parameters:
return self.chat_handler_factory(use_gpu=use_gpu)
except (TypeError, ValueError):
pass
return self.chat_handler_factory()
def ensure_loaded(self):
if self._llm is not None:
return self._llm
@@ -663,37 +979,83 @@ class LlamaHandle:
if not self.model_path.is_file():
raise FileNotFoundError(f"GGUF model not found: {self.model_path}")
# llama.cpp owns its accelerator allocator, so reserve enough room
# through ComfyUI instead of emptying a global backend cache.
if self.n_gpu_layers != 0:
reserve_external_vram(self.model_path.stat().st_size)
llama_cpp = require_module("llama_cpp", "llama-cpp-python")
if self.chat_handler_factory is not None:
self._chat_handler = self.chat_handler_factory()
capabilities = llama_cpp_diagnostics(llama_cpp)
supports_gpu = capabilities["gpu_offload"]
effective_gpu_layers = self.n_gpu_layers
if effective_gpu_layers != 0 and supports_gpu is False:
LOGGER.warning(
"The installed llama-cpp-python build has no accelerator "
"offload; falling back from n_gpu_layers=%s to CPU.",
effective_gpu_layers,
)
effective_gpu_layers = 0
# llama.cpp owns its allocator, so ask ComfyUI to make room without
# globally clearing unrelated accelerator caches.
if effective_gpu_layers != 0:
memory_required = self.model_path.stat().st_size
if self.projector_path is not None and self.projector_path.is_file():
memory_required += self.projector_path.stat().st_size
reserve_external_vram(memory_required)
use_gpu = effective_gpu_layers != 0 and supports_gpu is not False
self._chat_handler = self._create_chat_handler(use_gpu=use_gpu)
effective_batch = min(
self.n_batch, self.n_ctx if self.n_ctx > 0 else self.n_batch
)
effective_ubatch = min(self.n_ubatch, effective_batch)
flash_attention = self.flash_attention == "Enabled" or (
self.flash_attention == "Auto" and use_gpu
)
requested = {
"model_path": str(self.model_path),
"chat_handler": self._chat_handler,
"chat_format": self.chat_format,
"n_ctx": self.n_ctx,
"n_gpu_layers": self.n_gpu_layers,
"n_gpu_layers": effective_gpu_layers,
"n_threads": self.n_threads,
"n_batch": min(1024, self.n_ctx),
"offload_kqv": self.n_gpu_layers != 0,
"flash_attn": self.n_gpu_layers != 0,
# None lets llama.cpp choose its optimized prompt-processing
# thread count independently from token-generation threads.
"n_threads_batch": self.n_threads_batch,
"n_batch": effective_batch,
"n_ubatch": effective_ubatch,
"split_mode": _resolve_split_mode(llama_cpp, self.split_mode),
"main_gpu": self.main_gpu,
"tensor_split": self.tensor_split,
"offload_kqv": use_gpu,
"op_offload": use_gpu,
"flash_attn": flash_attention,
"use_mmap": self.use_mmap and capabilities["mmap"] is not False,
"use_mlock": False,
"embedding": False,
"verbose": False,
"seed": self.seed,
}
signature = inspect.signature(llama_cpp.Llama.__init__)
kwargs = {
key: value
for key, value in requested.items()
if key in signature.parameters and value is not None
}
self._llm = llama_cpp.Llama(**kwargs)
kwargs = self._filter_constructor_kwargs(llama_cpp.Llama, requested)
try:
self._llm = llama_cpp.Llama(**kwargs)
except Exception as exc:
# Some backend/model pairs reject flash attention at context
# creation. Auto is an optimization, never a compatibility
# requirement, so retry once with the official safe default.
if (
self.flash_attention == "Auto"
and kwargs.get("flash_attn")
and re.search(r"flash|attn|attention", str(exc), re.IGNORECASE)
):
LOGGER.warning(
"llama.cpp flash attention was unavailable; retrying "
"with the portable attention path: %s",
exc,
)
kwargs["flash_attn"] = False
gc.collect()
self._llm = llama_cpp.Llama(**kwargs)
else:
raise
return self._llm
def close(self) -> None:
@@ -743,9 +1105,6 @@ def close_handle(handle: Any) -> None:
def inference_context(device: torch.device, dtype: torch.dtype):
if (
device.type in {"cuda", "xpu"}
and dtype in {torch.float16, torch.bfloat16}
):
if device.type in {"cuda", "xpu"} and dtype in {torch.float16, torch.bfloat16}:
return torch.autocast(device.type, dtype=dtype)
return nullcontext()
+539
View File
@@ -0,0 +1,539 @@
"""Prompt-seeded SAM2.1 image-batch/video segmentation and tracking."""
from __future__ import annotations
import math
from dataclasses import dataclass
from typing import Any
import torch
from .geometry import bbox_from_mask, deterministic_color
from .runtime import (
CachedModelNode,
ManagedTorchModel,
inference_context,
model_device,
require_module,
snapshot_download,
tensor_batch_to_pil,
torch_dtype,
)
from .vision_types import (
VLM_DETECTIONS,
VLM_TRACKS,
Detection,
DetectionSequence,
Track,
TrackSequence,
)
@dataclass(frozen=True)
class Sam2Spec:
model_id: str
cache_name: str
SAM2_MODELS = {
"SAM2.1 Hiera Tiny (fast)": Sam2Spec(
"facebook/sam2.1-hiera-tiny", "sam2.1-hiera-tiny"
),
"SAM2.1 Hiera Small": Sam2Spec("facebook/sam2.1-hiera-small", "sam2.1-hiera-small"),
"SAM2.1 Hiera Base+": Sam2Spec(
"facebook/sam2.1-hiera-base-plus", "sam2.1-hiera-base-plus"
),
"SAM2.1 Hiera Large": Sam2Spec("facebook/sam2.1-hiera-large", "sam2.1-hiera-large"),
}
def _core_box(value: dict[str, Any]) -> tuple[float, float, float, float]:
try:
x, y = float(value["x"]), float(value["y"])
width, height = float(value["width"]), float(value["height"])
except (KeyError, TypeError, ValueError) as exc:
raise ValueError(
"BOUNDING_BOX must contain numeric x, y, width, and height."
) from exc
if width <= 0 or height <= 0:
raise ValueError("BOUNDING_BOX width and height must be positive.")
return x, y, x + width, y + height
def _core_box_frames(value: Any) -> list[list[dict[str, Any]]]:
"""Normalize core dict/flat/nested BOUNDING_BOX values to frame lists."""
if value is None:
return []
if isinstance(value, dict):
return [[value]]
if not isinstance(value, list):
raise TypeError("BOUNDING_BOX must be a dict, list of dicts, or frame list.")
if not value:
return []
if all(isinstance(item, dict) for item in value):
return [value]
if all(
isinstance(frame, list) and all(isinstance(item, dict) for item in frame)
for frame in value
):
return value
raise TypeError("BOUNDING_BOX contains an unsupported nested value.")
def _box_label(value: dict[str, Any]) -> str | None:
label = value.get("label")
if isinstance(label, str) and label.strip():
return label.strip()
metadata = value.get("metadata")
if isinstance(metadata, dict):
label = metadata.get("label")
if isinstance(label, str) and label.strip():
return label.strip()
return None
def seed_boxes(
*,
width: int,
height: int,
frame_index: int,
detections: DetectionSequence | None,
bounding_box: Any,
) -> tuple[list[list[float]], list[int], dict[int, str | None]]:
boxes: list[list[float]] = []
object_ids: list[int] = []
labels: dict[int, str | None] = {}
if detections is not None:
if not isinstance(detections, DetectionSequence):
raise TypeError("detections must be a VLM Detection Sequence.")
if detections.width != width or detections.height != height:
raise ValueError(
"Detection dimensions must exactly match the SAM2 video frames."
)
frame = detections.frame(frame_index)
if frame is None and len(detections.frames) == 1:
# A detector run over a selected single image is an explicit seed
# annotation and may be applied to any chosen video frame.
frame = detections.frames[0]
elif frame is None and detections.frames:
raise ValueError(
f"Detections do not contain the requested seed frame {frame_index}."
)
for index, detection in enumerate(frame.detections if frame else (), 1):
track_id = detection.track_id
object_id = int(track_id if track_id is not None else index)
while object_id in object_ids:
object_id += 1
x1, y1, x2, y2 = detection.bbox_xyxy
boxes.append(
[
min(width, max(0.0, x1)),
min(height, max(0.0, y1)),
min(width, max(0.0, x2)),
min(height, max(0.0, y2)),
]
)
object_ids.append(object_id)
labels[object_id] = detection.label
box_frames = _core_box_frames(bounding_box)
if len(box_frames) == 1:
selected_boxes = box_frames[0]
elif box_frames and frame_index < len(box_frames):
selected_boxes = box_frames[frame_index]
elif box_frames:
raise ValueError(
f"BOUNDING_BOX has {len(box_frames)} frames but seed_frame is "
f"{frame_index}."
)
else:
selected_boxes = []
for value in selected_boxes:
x1, y1, x2, y2 = _core_box(value)
object_id = max(object_ids, default=0) + 1
boxes.append(
[
min(width, max(0.0, x1)),
min(height, max(0.0, y1)),
min(width, max(0.0, x2)),
min(height, max(0.0, y2)),
]
)
object_ids.append(object_id)
labels[object_id] = _box_label(value)
valid = []
for box, object_id in zip(boxes, object_ids):
if box[2] > box[0] and box[3] > box[1]:
valid.append((box, object_id))
return (
[box for box, _object_id in valid],
[object_id for _box, object_id in valid],
labels,
)
def _normalize_processed_masks(value: torch.Tensor) -> torch.Tensor:
masks = value.detach().to(device="cpu")
if masks.ndim == 4 and masks.shape[1] == 1:
masks = masks[:, 0]
elif masks.ndim == 4 and masks.shape[0] == 1:
masks = masks[0]
if masks.ndim == 2:
masks = masks.unsqueeze(0)
if masks.ndim != 3:
raise RuntimeError(
f"SAM2 returned an unsupported mask shape {tuple(masks.shape)}."
)
return masks if masks.dtype == torch.bool else masks > 0.5
class Sam2VideoPredictor:
def __init__(self, spec: Sam2Spec, precision: str):
transformers = require_module("transformers")
if not hasattr(transformers, "Sam2VideoModel"):
raise RuntimeError(
"SAM2 video requires Transformers with Sam2VideoModel support."
)
model_path = snapshot_download(
spec.model_id,
spec.cache_name,
ignore_patterns=["*.pt", "*.bin", "*.onnx", "*.tflite"],
)
self.processor = transformers.Sam2VideoProcessor.from_pretrained(model_path)
self.dtype = torch_dtype(precision)
model = transformers.Sam2VideoModel.from_pretrained(
model_path, dtype=self.dtype
)
model.eval()
self.handle = ManagedTorchModel(model, processor=self.processor)
self.spec = spec
def close(self):
self.handle.close()
def propagate(
self,
images: torch.Tensor,
*,
seed_frame: int,
fps: float,
detections: DetectionSequence | None,
bounding_box: Any,
seed_mask: torch.Tensor | None,
mask_threshold: float,
keep_video_on_cpu: bool,
mask_output: str,
render_preview: bool,
) -> tuple[TrackSequence, torch.Tensor, torch.Tensor, torch.Tensor]:
if not math.isfinite(fps) or fps <= 0:
raise ValueError("fps must be finite and positive.")
if mask_output not in {"union_only", "union_and_objects"}:
raise ValueError(f"Unsupported mask_output mode {mask_output!r}.")
pil_images = tensor_batch_to_pil(images)
if not pil_images:
raise ValueError("SAM2 requires at least one image.")
if not 0 <= seed_frame < len(pil_images):
raise ValueError(
f"seed_frame {seed_frame} is outside the {len(pil_images)}-frame batch."
)
width, height = pil_images[0].size
if any(image.size != (width, height) for image in pil_images):
raise ValueError("Every video frame must have identical dimensions.")
boxes, object_ids, labels = seed_boxes(
width=width,
height=height,
frame_index=seed_frame,
detections=detections,
bounding_box=bounding_box,
)
masks_for_seed = None
if not boxes and seed_mask is not None:
masks_for_seed = seed_mask.detach().to(device="cpu", dtype=torch.float32)
if masks_for_seed.ndim == 2:
masks_for_seed = masks_for_seed.unsqueeze(0)
if masks_for_seed.ndim != 3 or tuple(masks_for_seed.shape[-2:]) != (
height,
width,
):
raise ValueError("seed_mask must have shape [objects, height, width].")
object_ids = list(range(1, masks_for_seed.shape[0] + 1))
labels = {object_id: None for object_id in 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
@@ -0,0 +1,560 @@
"""Guarded adapters for ComfyUI core ``SAM3_TRACK_DATA`` payloads.
Core SAM3 keeps masks bit-packed for memory efficiency. This module preserves
that payload untouched and emits small canonical ``VLM_TRACKS`` metadata with
mask references instead of embedding dense masks in JSON.
"""
from __future__ import annotations
import json
import math
from collections.abc import Iterator, Mapping, Sequence
from dataclasses import dataclass
from typing import Any
import torch
from .geometry import bbox_from_mask, clip_box
from .vision_types import (
VLM_DETECTIONS,
VLM_TRACKS,
Detection,
DetectionSequence,
Track,
TrackSequence,
)
SAM3_TRACK_DATA = "SAM3_TRACK_DATA"
SAM3_ADAPTER_SOURCE = "comfyui-core-sam3"
_REQUIRED_KEYS = frozenset({"packed_masks", "n_frames", "scores", "orig_size"})
@dataclass(frozen=True, slots=True)
class SAM3TrackLayout:
n_frames: int
n_objects: int
mask_height: int
mask_width: int
orig_height: int
orig_width: int
scores: tuple[float | None, ...]
@dataclass(frozen=True, slots=True)
class _SeedIdentity:
track_id: int | None
label: str | None
text: str | None
score: float | None
source: str | None
def _integer(value: Any, name: str, *, minimum: int = 0) -> int:
if isinstance(value, bool) or not isinstance(value, int):
raise TypeError(f"{name} must be an integer.")
if value < minimum:
raise ValueError(f"{name} must be at least {minimum}.")
return value
def _scores(
values: Any,
*,
n_objects: int,
) -> tuple[float | None, ...]:
if isinstance(values, torch.Tensor):
if values.ndim != 1:
raise ValueError("SAM3 scores tensor must have shape [objects].")
items = values.detach().cpu().tolist()
elif isinstance(values, Sequence) and not isinstance(values, (str, bytes)):
items = list(values)
else:
raise TypeError("SAM3 scores must be a one-dimensional sequence.")
if len(items) > n_objects:
raise ValueError("SAM3 scores contain more entries than mask objects.")
parsed: list[float | None] = []
for value in items:
if value is None:
parsed.append(None)
continue
score = float(value)
if not math.isfinite(score) or not 0.0 <= score <= 1.0:
raise ValueError("SAM3 scores must be finite values from 0 to 1.")
parsed.append(score)
parsed.extend([None] * (n_objects - len(parsed)))
return tuple(parsed)
def validate_sam3_track_data(track_data: Any) -> SAM3TrackLayout:
"""Validate the private core payload before interpreting its bit layout."""
if not isinstance(track_data, Mapping):
raise TypeError("SAM3_TRACK_DATA must be a mapping.")
missing = sorted(_REQUIRED_KEYS - set(track_data))
if missing:
raise ValueError(
"SAM3_TRACK_DATA is missing required keys: " + ", ".join(missing)
)
n_frames = _integer(track_data["n_frames"], "n_frames")
orig_size = track_data["orig_size"]
if not isinstance(orig_size, (tuple, list)) or len(orig_size) != 2:
raise TypeError("SAM3 orig_size must be (height, width).")
orig_height = _integer(orig_size[0], "orig_height", minimum=1)
orig_width = _integer(orig_size[1], "orig_width", minimum=1)
packed = track_data["packed_masks"]
if packed is None:
scores = _scores(track_data["scores"], n_objects=0)
return SAM3TrackLayout(
n_frames=n_frames,
n_objects=0,
mask_height=0,
mask_width=0,
orig_height=orig_height,
orig_width=orig_width,
scores=scores,
)
if not isinstance(packed, torch.Tensor):
raise TypeError("SAM3 packed_masks must be a torch.Tensor or None.")
if packed.dtype != torch.uint8:
raise TypeError("SAM3 packed_masks must use torch.uint8.")
if packed.ndim != 4:
raise ValueError(
"SAM3 packed_masks must have shape [frames, objects, height, packed_width]."
)
if packed.shape[0] != n_frames:
raise ValueError("SAM3 n_frames does not match packed_masks.")
n_objects = int(packed.shape[1])
mask_height = int(packed.shape[2])
packed_width = int(packed.shape[3])
if n_objects < 1 or mask_height < 1 or packed_width < 1:
raise ValueError("SAM3 packed_masks dimensions must be positive.")
scores = _scores(track_data["scores"], n_objects=n_objects)
return SAM3TrackLayout(
n_frames=n_frames,
n_objects=n_objects,
mask_height=mask_height,
mask_width=packed_width * 8,
orig_height=orig_height,
orig_width=orig_width,
scores=scores,
)
def unpack_sam3_mask(packed_mask: torch.Tensor) -> torch.Tensor:
"""Unpack exactly one object/frame mask, avoiding full-video expansion."""
if not isinstance(packed_mask, torch.Tensor):
raise TypeError("packed_mask must be a torch.Tensor.")
if packed_mask.dtype != torch.uint8 or packed_mask.ndim != 2:
raise ValueError("packed_mask must be uint8 with shape [height, packed_width].")
bits = torch.tensor(
(1, 2, 4, 8, 16, 32, 64, 128),
dtype=torch.uint8,
device=packed_mask.device,
)
return (
torch.bitwise_and(packed_mask.unsqueeze(-1), bits)
.ne(0)
.reshape(packed_mask.shape[0], packed_mask.shape[1] * 8)
)
def iter_sam3_masks(
track_data: Mapping[str, Any],
*,
present_only: bool = False,
) -> Iterator[tuple[int, int, torch.Tensor]]:
"""Yield one unpacked mask at a time as ``(frame, object, mask)``."""
layout = validate_sam3_track_data(track_data)
packed = track_data["packed_masks"]
if packed is None:
return
for frame_index in range(layout.n_frames):
for object_index in range(layout.n_objects):
mask = unpack_sam3_mask(packed[frame_index, object_index])
if present_only and not bool(mask.any().item()):
continue
yield frame_index, object_index, mask
def _seeds_from_detections(
sequence: DetectionSequence,
) -> list[_SeedIdentity]:
for frame in sequence.frames:
if frame.detections:
return [
_SeedIdentity(
track_id=detection.track_id,
label=detection.label,
text=detection.text,
score=detection.score,
source=detection.source,
)
for detection in frame.detections
]
return []
def _seeds_from_tracks(sequence: TrackSequence) -> list[_SeedIdentity]:
return [
_SeedIdentity(
track_id=track.track_id,
label=track.label or track.detections[0].label,
text=track.detections[0].text,
score=track.score,
source=track.source,
)
for track in sequence.tracks
]
def _seed_identities(
*,
seed_detections: DetectionSequence | None,
seed_tracks: TrackSequence | None,
n_objects: int,
) -> tuple[tuple[_SeedIdentity, ...], int]:
if seed_detections is not None and seed_tracks is not None:
raise ValueError("Connect seed_detections or seed_tracks, not both.")
if seed_detections is not None:
if not isinstance(seed_detections, DetectionSequence):
raise TypeError("seed_detections must be a DetectionSequence.")
seeds = _seeds_from_detections(seed_detections)
elif seed_tracks is not None:
if not isinstance(seed_tracks, TrackSequence):
raise TypeError("seed_tracks must be a TrackSequence.")
seeds = _seeds_from_tracks(seed_tracks)
else:
seeds = []
used_ids: set[int] = set()
next_id = 0
identities = []
for object_index in range(n_objects):
seed = seeds[object_index] if object_index < len(seeds) else None
preferred = None if seed is None else seed.track_id
if preferred is not None and preferred not in used_ids:
track_id = preferred
else:
while next_id in used_ids:
next_id += 1
track_id = next_id
next_id += 1
used_ids.add(track_id)
identities.append(
_SeedIdentity(
track_id=track_id,
label=None if seed is None else seed.label,
text=None if seed is None else seed.text,
score=None if seed is None else seed.score,
source=None if seed is None else seed.source,
)
)
return tuple(identities), min(len(seeds), n_objects)
def _scaled_bbox(
bbox: tuple[float, float, float, float],
layout: SAM3TrackLayout,
) -> tuple[float, float, float, float]:
scale_x = layout.orig_width / layout.mask_width
scale_y = layout.orig_height / layout.mask_height
x1, y1, x2, y2 = bbox
return clip_box(
(
x1 * scale_x,
y1 * scale_y,
x2 * scale_x,
y2 * scale_y,
),
layout.orig_width,
layout.orig_height,
)
def sam3_track_data_to_tracks(
track_data: Mapping[str, Any],
*,
seed_detections: DetectionSequence | None = None,
seed_tracks: TrackSequence | None = None,
fps: float | None = None,
source: str = SAM3_ADAPTER_SOURCE,
) -> TrackSequence:
"""Create canonical sparse metadata while retaining packed masks separately."""
layout = validate_sam3_track_data(track_data)
if fps is not None:
fps = float(fps)
if not math.isfinite(fps) or fps <= 0:
raise ValueError("fps must be finite and positive or None.")
identities, seed_count = _seed_identities(
seed_detections=seed_detections,
seed_tracks=seed_tracks,
n_objects=layout.n_objects,
)
detections_by_object: list[list[Detection]] = [
[] for _index in range(layout.n_objects)
]
for frame_index, object_index, mask in iter_sam3_masks(
track_data, present_only=True
):
bbox = bbox_from_mask(mask)
if bbox is None:
continue
identity = identities[object_index]
timestamp = frame_index / fps if fps is not None else 0.0
score = layout.scores[object_index]
if score is None:
score = identity.score
detections_by_object[object_index].append(
Detection(
bbox_xyxy=_scaled_bbox(bbox, layout),
label=identity.label,
text=identity.text,
score=score,
frame_index=frame_index,
timestamp=timestamp,
track_id=identity.track_id,
source=source,
metadata={
"observation": "propagated",
"visibility": "visible",
"sam3_object_index": object_index,
"mask_ref": {
"type": SAM3_TRACK_DATA,
"frame_index": frame_index,
"object_index": object_index,
},
},
)
)
tracks = []
for object_index, detections in enumerate(detections_by_object):
if not detections:
continue
identity = identities[object_index]
present_frames = len(detections)
final_state = (
"active" if detections[-1].frame_index == layout.n_frames - 1 else "lost"
)
score = layout.scores[object_index]
if score is None:
score = identity.score
tracks.append(
Track(
track_id=identity.track_id,
detections=tuple(detections),
label=identity.label,
score=score,
source=source,
metadata={
"state": final_state,
"sam3_object_index": object_index,
"first_frame": detections[0].frame_index,
"last_observed_frame": detections[-1].frame_index,
"present_frames": present_frames,
"presence_ratio": (
present_frames / layout.n_frames if layout.n_frames else 0.0
),
"seeded": object_index < seed_count,
},
)
)
return TrackSequence(
width=layout.orig_width,
height=layout.orig_height,
tracks=tuple(sorted(tracks, key=lambda item: item.track_id)),
frame_count=layout.n_frames,
fps=fps,
source=source,
metadata={
"adapter": "sam3-track-data/v1",
"mask_payload": {
"type": SAM3_TRACK_DATA,
"encoding": "little-endian-bitpack",
"mask_width": layout.mask_width,
"mask_height": layout.mask_height,
},
"object_slots": layout.n_objects,
},
)
def track_report_payload(tracks: TrackSequence) -> dict[str, Any]:
"""Return a compact, history-safe report with no tensor content."""
if not isinstance(tracks, TrackSequence):
raise TypeError("tracks must be a TrackSequence.")
records = []
state_counts: dict[str, int] = {}
total_observations = 0
for track in tracks.tracks:
observations: dict[str, int] = {}
for detection in track.detections:
kind = str(detection.metadata.to_dict().get("observation", "detected"))
observations[kind] = observations.get(kind, 0) + 1
total_observations += len(track.detections)
state = str(track.metadata.to_dict().get("state", "unknown"))
state_counts[state] = state_counts.get(state, 0) + 1
records.append(
{
"track_id": track.track_id,
"label": track.label,
"state": state,
"score": track.score,
"first_frame": track.detections[0].frame_index,
"last_frame": track.detections[-1].frame_index,
"observation_count": len(track.detections),
"observations": observations,
}
)
media: dict[str, Any] = {
"width": tracks.width,
"height": tracks.height,
"frame_count": tracks.frame_count,
}
if tracks.fps is not None:
media["fps"] = tracks.fps
return {
"schema": "comfyui-vlm/track-report",
"version": 1,
"media": media,
"track_count": len(tracks.tracks),
"observation_count": total_observations,
"state_counts": dict(sorted(state_counts.items())),
"tracks": records,
}
def track_report_json(
tracks: TrackSequence,
*,
indent: int | None = 2,
) -> str:
return json.dumps(
track_report_payload(tracks),
ensure_ascii=False,
allow_nan=False,
sort_keys=True,
indent=indent,
)
def track_report_text(tracks: TrackSequence) -> str:
report = track_report_payload(tracks)
media = report["media"]
lines = [
(
f"Tracks: {report['track_count']} | "
f"Observations: {report['observation_count']} | "
f"Frames: {media['frame_count']} | "
f"Size: {media['width']}x{media['height']}"
)
]
if "fps" in media:
lines[0] += f" | FPS: {media['fps']:g}"
for track in report["tracks"]:
label = track["label"] or "(unlabeled)"
lines.append(
f"#{track['track_id']} {label}: {track['state']}, "
f"frames {track['first_frame']}-{track['last_frame']}, "
f"{track['observation_count']} observations"
)
return "\n".join(lines)
class VLMSAM3TrackAdapter:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"track_data": (SAM3_TRACK_DATA,),
"fps": (
"FLOAT",
{
"default": 0.0,
"min": 0.0,
"max": 1000.0,
"step": 0.01,
"tooltip": "0 keeps timestamps unknown.",
},
),
},
"optional": {
"seed_detections": (VLM_DETECTIONS,),
"seed_tracks": (VLM_TRACKS,),
},
}
RETURN_TYPES = (VLM_TRACKS, SAM3_TRACK_DATA)
RETURN_NAMES = ("tracks", "track_data")
FUNCTION = "adapt"
CATEGORY = "VLM Nodes/Vision/Tracking"
def adapt(
self,
track_data,
fps,
seed_detections=None,
seed_tracks=None,
):
tracks = sam3_track_data_to_tracks(
track_data,
seed_detections=seed_detections,
seed_tracks=seed_tracks,
fps=None if fps <= 0 else fps,
)
return tracks, track_data
class VLMTrackReport:
@classmethod
def INPUT_TYPES(cls):
return {"required": {"tracks": (VLM_TRACKS,)}}
RETURN_TYPES = ("STRING", "STRING")
RETURN_NAMES = ("report_json", "report_text")
FUNCTION = "report"
CATEGORY = "VLM Nodes/Vision/Tracking"
OUTPUT_NODE = True
def report(self, tracks):
report_json = track_report_json(tracks)
report_text = track_report_text(tracks)
return {
"ui": {"text": [report_text]},
"result": (report_json, report_text),
}
NODE_CLASS_MAPPINGS = {
"VLMSAM3TrackAdapter": VLMSAM3TrackAdapter,
"VLMTrackReport": VLMTrackReport,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"VLMSAM3TrackAdapter": "VLM SAM3 Track Adapter",
"VLMTrackReport": "VLM Track Report",
}
__all__ = [
"NODE_CLASS_MAPPINGS",
"NODE_DISPLAY_NAME_MAPPINGS",
"SAM3TrackLayout",
"SAM3_ADAPTER_SOURCE",
"SAM3_TRACK_DATA",
"VLMSAM3TrackAdapter",
"VLMTrackReport",
"iter_sam3_masks",
"sam3_track_data_to_tracks",
"track_report_json",
"track_report_payload",
"track_report_text",
"unpack_sam3_mask",
"validate_sam3_track_data",
]
File diff suppressed because it is too large Load Diff
+148 -70
View File
@@ -12,15 +12,18 @@ import os
import re
from typing import Any, Literal, Optional
import folder_paths
import torch
from pydantic import BaseModel, Field
import folder_paths
from .prompts import system_msg_prompts, system_msg_simple
from .runtime import (
LlamaHandle,
close_handle,
default_llama_threads,
llama_chat_content,
llama_runtime_input_types,
llama_runtime_options,
require_module,
resolve_model_path,
unwrap_llm,
@@ -36,17 +39,13 @@ ANY = AnyType("*")
class Analysis(BaseModel):
main_character: list[str] = Field(
..., description="Main subjects and objects."
)
main_character: list[str] = Field(..., description="Main subjects and objects.")
artform: list[str] = Field(..., description="Art forms present.")
photo_type: list[str] = Field(..., description="Photographic genres.")
color_with_objects: list[str] = Field(
..., description="Objects paired with their colors."
)
digital_artform: list[str] = Field(
..., description="Digital art techniques."
)
digital_artform: list[str] = Field(..., description="Digital art techniques.")
background: list[str] = Field(..., description="Background details.")
lighting: list[str] = Field(..., description="Lighting details.")
@@ -86,9 +85,7 @@ class ArtPromptSpecification(BaseModel):
techniques: ArtisticTechniques
theme: ImageryTheme
style: VisualStyle
creative_descriptions: list[ArtInspirationNarrative] = Field(
default_factory=list
)
creative_descriptions: list[ArtInspirationNarrative] = Field(default_factory=list)
def _schema(model_class: type[BaseModel]) -> dict[str, Any]:
@@ -98,12 +95,7 @@ def _schema(model_class: type[BaseModel]) -> dict[str, Any]:
def _response_content(response: dict[str, Any]) -> str:
try:
return str(response["choices"][0]["message"]["content"])
except (KeyError, IndexError, TypeError) as exc:
raise RuntimeError(
f"The model returned an unexpected response: {response!r}"
) from exc
return llama_chat_content(response)
def _chat(
@@ -161,9 +153,7 @@ def _structured_chat(
try:
parsed = json.loads(raw)
except json.JSONDecodeError as exc:
raise RuntimeError(
f"The model did not return valid JSON: {raw[:500]}"
) from exc
raise RuntimeError(f"The model did not return valid JSON: {raw[:500]}") from exc
return raw, parsed
@@ -253,8 +243,7 @@ class PromptGenerateAPI:
{
"default": "",
"tooltip": (
"OpenAI-compatible base URL, e.g. "
"http://127.0.0.1:8000/v1."
"OpenAI-compatible base URL, e.g. http://127.0.0.1:8000/v1."
),
},
),
@@ -298,9 +287,7 @@ class PromptGenerateAPI:
model, route_url, route_mode = route
model = (model_override or model).strip()
if not model:
raise ValueError(
"A model ID is required for Custom / OpenAI-compatible."
)
raise ValueError("A model ID is required for Custom / OpenAI-compatible.")
effective_url = (base_url or route_url or "").strip() or None
mode = route_mode if api_mode == "Auto" else api_mode
return model, effective_url, mode
@@ -326,11 +313,7 @@ class PromptGenerateAPI:
)
key = (
api_key.strip()
or (
os.getenv("DEEPSEEK_API_KEY", "")
if model_name == "DeepSeek"
else ""
)
or (os.getenv("DEEPSEEK_API_KEY", "") if model_name == "DeepSeek" else "")
or os.getenv("VLM_API_KEY", "")
or os.getenv("OPENAI_API_KEY", "")
)
@@ -355,9 +338,7 @@ class PromptGenerateAPI:
f"Optional question:\n{question.strip()}"
).strip()
history_limit = max(0, int(context_size)) * 2
history = (
self.session_history[-history_limit:] if history_limit else []
)
history = self.session_history[-history_limit:] if history_limit else []
if mode == "Responses":
response = client.responses.create(
@@ -397,20 +378,23 @@ class LLMLoader:
def INPUT_TYPES(cls):
return {
"required": {
"ckpt_name": (
folder_paths.get_filename_list("LLavacheckpoints"),
),
"ckpt_name": (folder_paths.get_filename_list("LLavacheckpoints"),),
"max_ctx": (
"INT",
{"default": 2048, "min": 128, "max": 131072, "step": 64},
),
"gpu_layers": (
"INT",
{"default": 27, "min": -1, "max": 1000, "step": 1},
{"default": -1, "min": -1, "max": 1000, "step": 1},
),
"n_threads": (
"INT",
{"default": 8, "min": 1, "max": 256, "step": 1},
{
"default": default_llama_threads(),
"min": 1,
"max": 256,
"step": 1,
},
),
},
"optional": {
@@ -422,7 +406,8 @@ class LLMLoader:
"Leave blank to use the chat template embedded in GGUF."
),
},
)
),
**llama_runtime_input_types(),
},
}
@@ -432,7 +417,19 @@ class LLMLoader:
CATEGORY = "VLM Nodes/LLM"
def load_llm_checkpoint(
self, ckpt_name, max_ctx, gpu_layers, n_threads, chat_format=""
self,
ckpt_name,
max_ctx,
gpu_layers,
n_threads,
chat_format="",
n_batch=512,
n_ubatch=512,
flash_attention="Auto",
use_mmap=True,
split_mode="Layer",
main_gpu=0,
tensor_split="",
):
return (
LlamaHandle(
@@ -441,6 +438,15 @@ class LLMLoader:
n_gpu_layers=gpu_layers,
n_threads=n_threads,
chat_format=chat_format.strip() or None,
**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,
),
),
)
@@ -693,9 +699,9 @@ class ChatMusician:
symusic = require_module("symusic", "symusic")
score = symusic.Score.from_abc(abc)
rendered = symusic.Synthesizer(
sample_rate=int(sample_rate)
).render(score, stereo=True)
rendered = symusic.Synthesizer(sample_rate=int(sample_rate)).render(
score, stereo=True
)
waveform = torch.as_tensor(rendered, dtype=torch.float32)
if waveform.ndim == 1:
waveform = waveform.unsqueeze(0)
@@ -794,9 +800,7 @@ class CreativeArtPromptGenerator:
techniques = ", ".join(parsed["techniques"]["preferred"])
theme = parsed["theme"]["core_subject"]
styles = ", ".join(parsed["style"]["desired"])
return (
f"{theme}. Techniques: {techniques}. Visual style: {styles}.",
)
return (f"{theme}. Techniques: {techniques}. Visual style: {styles}.",)
class Suggester:
@@ -889,13 +893,9 @@ class StructuredOutput:
"float": "number",
"bool": "boolean",
}
property_schema: dict[str, Any] = {
"description": attribute_description.strip()
}
property_schema: dict[str, Any] = {"description": attribute_description.strip()}
if attribute_type == "Category":
values = [
value.strip() for value in categories.split(",") if value.strip()
]
values = [value.strip() for value in categories.split(",") if value.strip()]
if not values:
raise ValueError(
"Category requires at least one comma-separated value."
@@ -917,9 +917,7 @@ class StructuredOutput:
temperature=temperature,
)
value = parsed[name]
return (
value if isinstance(value, str) else json.dumps(value),
)
return (value if isinstance(value, str) else json.dumps(value),)
class _CachedLLMBase:
@@ -927,13 +925,24 @@ class _CachedLLMBase:
self._handle = None
self._key = None
def _model(self, ckpt_name, max_ctx, gpu_layers, n_threads, seed=42):
def _model(
self,
ckpt_name,
max_ctx,
gpu_layers,
n_threads,
seed=42,
chat_format="",
**runtime_options,
):
key = (
ckpt_name,
int(max_ctx),
int(gpu_layers),
int(n_threads),
int(seed),
chat_format.strip(),
tuple(sorted(runtime_options.items())),
)
if self._handle is None or self._key != key:
close_handle(self._handle)
@@ -943,6 +952,8 @@ class _CachedLLMBase:
n_gpu_layers=gpu_layers,
n_threads=n_threads,
seed=seed,
chat_format=chat_format.strip() or None,
**runtime_options,
)
self._key = key
return self._handle
@@ -959,20 +970,23 @@ class LLMOptionalMemoryFreeSimple(_CachedLLMBase):
def INPUT_TYPES(cls):
return {
"required": {
"ckpt_name": (
folder_paths.get_filename_list("LLavacheckpoints"),
),
"ckpt_name": (folder_paths.get_filename_list("LLavacheckpoints"),),
"max_ctx": (
"INT",
{"default": 4096, "min": 128, "max": 131072, "step": 64},
),
"gpu_layers": (
"INT",
{"default": 27, "min": -1, "max": 1000, "step": 1},
{"default": -1, "min": -1, "max": 1000, "step": 1},
),
"n_threads": (
"INT",
{"default": 8, "min": 1, "max": 256, "step": 1},
{
"default": default_llama_threads(),
"min": 1,
"max": 256,
"step": 1,
},
),
"prompt": (
"STRING",
@@ -983,7 +997,14 @@ class LLMOptionalMemoryFreeSimple(_CachedLLMBase):
{"default": 0.1, "min": 0.0, "max": 2.0, "step": 0.01},
),
"unload": ("BOOLEAN", {"default": False}),
}
},
"optional": {
"chat_format": (
"STRING",
{"default": ""},
),
**llama_runtime_input_types(),
},
}
RETURN_TYPES = ("STRING",)
@@ -999,9 +1020,31 @@ class LLMOptionalMemoryFreeSimple(_CachedLLMBase):
prompt,
temperature,
unload,
chat_format="",
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, max_ctx, gpu_layers, n_threads
ckpt_name,
max_ctx,
gpu_layers,
n_threads,
chat_format=chat_format,
**options,
)
try:
return (
@@ -1020,20 +1063,23 @@ class LLMOptionalMemoryFreeAdvanced(_CachedLLMBase):
@classmethod
def INPUT_TYPES(cls):
required = {
"ckpt_name": (
folder_paths.get_filename_list("LLavacheckpoints"),
),
"ckpt_name": (folder_paths.get_filename_list("LLavacheckpoints"),),
"max_ctx": (
"INT",
{"default": 4096, "min": 128, "max": 131072, "step": 64},
),
"gpu_layers": (
"INT",
{"default": 27, "min": -1, "max": 1000, "step": 1},
{"default": -1, "min": -1, "max": 1000, "step": 1},
),
"n_threads": (
"INT",
{"default": 8, "min": 1, "max": 256, "step": 1},
{
"default": default_llama_threads(),
"min": 1,
"max": 256,
"step": 1,
},
),
"system_msg": (
"STRING",
@@ -1074,7 +1120,16 @@ class LLMOptionalMemoryFreeAdvanced(_CachedLLMBase):
"seed": ("INT", {"default": 42, "step": 1}),
"unload": ("BOOLEAN", {"default": False}),
}
return {"required": required}
return {
"required": required,
"optional": {
"chat_format": (
"STRING",
{"default": ""},
),
**llama_runtime_input_types(),
},
}
RETURN_TYPES = ("STRING",)
FUNCTION = "generate_text_advanced"
@@ -1097,9 +1152,32 @@ class LLMOptionalMemoryFreeAdvanced(_CachedLLMBase):
repeat_penalty,
seed,
unload,
chat_format="",
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, max_ctx, gpu_layers, n_threads, seed
ckpt_name,
max_ctx,
gpu_layers,
n_threads,
seed,
chat_format,
**options,
)
try:
return (
+766
View File
@@ -0,0 +1,766 @@
"""Deterministic tracking-by-detection for canonical VLM vision payloads.
The tracker intentionally owns only temporal association and identity. Dense
mask propagation remains the responsibility of SAM-style video models. This
keeps the baseline portable across CUDA, ROCm, MPS, XPU, and CPU systems.
"""
from __future__ import annotations
import math
from dataclasses import dataclass, field
from typing import Iterable
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",
]
+998
View File
@@ -0,0 +1,998 @@
"""Canonical, immutable spatial payloads shared by VLM vision nodes.
Coordinates use source-image pixels. Bounding boxes are always ``xyxy`` with
an exclusive right/bottom edge. Masks are optional in-process tensors and are
deliberately omitted from every JSON representation.
"""
from __future__ import annotations
import json
import math
from collections.abc import Iterator, Mapping
from dataclasses import dataclass, field
from typing import Any
import torch
VLM_DETECTIONS = "VLM_DETECTIONS"
VLM_TRACKS = "VLM_TRACKS"
VLM_POINTS = "VLM_POINTS"
VLM_EVENTS = "VLM_EVENTS"
SCHEMA_VERSION = 1
DETECTIONS_SCHEMA = "comfyui-vlm/detections"
TRACKS_SCHEMA = "comfyui-vlm/tracks"
POINTS_SCHEMA = "comfyui-vlm/points"
EVENTS_SCHEMA = "comfyui-vlm/events"
PointXY = tuple[float, float]
BoxXYXY = tuple[float, float, float, float]
Polygon = tuple[PointXY, ...]
class FrozenDict(Mapping[str, Any]):
"""Small recursively immutable mapping used for metadata."""
__slots__ = ("_items", "_lookup")
def __init__(self, values: Mapping[str, Any] | None = None):
items = []
for key, value in (values or {}).items():
if not isinstance(key, str):
raise TypeError("Metadata keys must be strings.")
items.append((key, _freeze_json(value)))
self._items = tuple(sorted(items))
self._lookup = dict(self._items)
def __getitem__(self, key: str) -> Any:
return self._lookup[key]
def __iter__(self) -> Iterator[str]:
return (key for key, _value in self._items)
def __len__(self) -> int:
return len(self._items)
def __repr__(self) -> str:
return f"FrozenDict({dict(self._items)!r})"
def __hash__(self) -> int:
return hash(self._items)
def to_dict(self) -> dict[str, Any]:
return {key: _thaw_json(value) for key, value in self._items}
def _freeze_json(value: Any) -> Any:
if value is None or isinstance(value, (str, bool, int)):
return value
if isinstance(value, float):
if not math.isfinite(value):
raise ValueError("Metadata numbers must be finite.")
return value
if isinstance(value, Mapping):
return FrozenDict(value)
if isinstance(value, (list, tuple)):
return tuple(_freeze_json(item) for item in value)
raise TypeError(f"Metadata must contain JSON values, got {type(value).__name__}.")
def _thaw_json(value: Any) -> Any:
if isinstance(value, FrozenDict):
return value.to_dict()
if isinstance(value, tuple):
return [_thaw_json(item) for item in value]
return value
def _metadata(value: Mapping[str, Any] | FrozenDict | None) -> FrozenDict:
return value if isinstance(value, FrozenDict) else FrozenDict(value)
def _finite(value: Any, name: str) -> float:
number = float(value)
if not math.isfinite(number):
raise ValueError(f"{name} must be finite.")
return number
def _non_negative(value: Any, name: str) -> float:
number = _finite(value, name)
if number < 0:
raise ValueError(f"{name} must be non-negative.")
return number
def _optional_score(value: Any) -> float | None:
if value is None:
return None
score = _finite(value, "score")
if not 0.0 <= score <= 1.0:
raise ValueError("score must be between 0 and 1.")
return score
def _optional_text(value: Any, name: str) -> str | None:
if value is None:
return None
if not isinstance(value, str):
raise TypeError(f"{name} must be a string or None.")
return value
def _box(value: Any) -> BoxXYXY:
if not isinstance(value, (list, tuple)) or len(value) != 4:
raise TypeError("bbox_xyxy must contain exactly four numbers.")
x1, y1, x2, y2 = (
_non_negative(component, "bbox coordinate") for component in value
)
if x2 < x1 or y2 < y1:
raise ValueError("bbox_xyxy must satisfy x2 >= x1 and y2 >= y1.")
return x1, y1, x2, y2
def _point(value: Any) -> PointXY:
if not isinstance(value, (list, tuple)) or len(value) != 2:
raise TypeError("A point must contain exactly two numbers.")
return (
_non_negative(value[0], "point x"),
_non_negative(value[1], "point y"),
)
def _polygon(
value: Any,
*,
name: str,
exact_points: int | None = None,
) -> Polygon | None:
if value is None:
return None
if not isinstance(value, (list, tuple)):
raise TypeError(f"{name} must be a sequence of points.")
points = tuple(_point(item) for item in value)
if exact_points is not None and len(points) != exact_points:
raise ValueError(f"{name} must contain exactly {exact_points} points.")
if exact_points is None and len(points) < 3:
raise ValueError(f"{name} must contain at least three points.")
return points
def _mask(value: Any) -> torch.Tensor | None:
if value is None:
return None
if not isinstance(value, torch.Tensor):
raise TypeError("mask must be a torch.Tensor or None.")
if value.ndim != 2:
raise ValueError("mask must have shape [height, width].")
return value.detach().to(dtype=torch.float32).clamp(0, 1).clone()
def _base_record(
*,
label: str | None,
text: str | None,
score: float | None,
frame_index: int,
timestamp: float,
track_id: int | None,
source: str | None,
metadata: FrozenDict,
) -> dict[str, Any]:
record: dict[str, Any] = {
"frame_index": frame_index,
"timestamp": timestamp,
}
if label is not None:
record["label"] = label
if text is not None:
record["text"] = text
if score is not None:
record["score"] = score
if track_id is not None:
record["track_id"] = track_id
if source is not None:
record["source"] = source
if metadata:
record["metadata"] = metadata.to_dict()
return record
@dataclass(frozen=True, slots=True)
class Detection:
bbox_xyxy: BoxXYXY
label: str | None = None
text: str | None = None
score: float | None = None
polygon: Polygon | None = None
quad: Polygon | None = None
frame_index: int = 0
timestamp: float = 0.0
track_id: int | None = None
source: str | None = None
metadata: FrozenDict = field(default_factory=FrozenDict)
mask: torch.Tensor | None = field(default=None, repr=False, compare=False)
def __post_init__(self) -> None:
object.__setattr__(self, "bbox_xyxy", _box(self.bbox_xyxy))
object.__setattr__(self, "label", _optional_text(self.label, "label"))
object.__setattr__(self, "text", _optional_text(self.text, "text"))
object.__setattr__(self, "score", _optional_score(self.score))
object.__setattr__(
self,
"polygon",
_polygon(self.polygon, name="polygon"),
)
object.__setattr__(
self,
"quad",
_polygon(self.quad, name="quad", exact_points=4),
)
if not isinstance(self.frame_index, int) or self.frame_index < 0:
raise ValueError("frame_index must be a non-negative integer.")
object.__setattr__(
self,
"timestamp",
_non_negative(self.timestamp, "timestamp"),
)
if self.track_id is not None and (
not isinstance(self.track_id, int) or self.track_id < 0
):
raise ValueError("track_id must be a non-negative integer or None.")
object.__setattr__(self, "source", _optional_text(self.source, "source"))
object.__setattr__(self, "metadata", _metadata(self.metadata))
object.__setattr__(self, "mask", _mask(self.mask))
@property
def area(self) -> float:
x1, y1, x2, y2 = self.bbox_xyxy
return (x2 - x1) * (y2 - y1)
@property
def center(self) -> PointXY:
x1, y1, x2, y2 = self.bbox_xyxy
return (x1 + x2) * 0.5, (y1 + y2) * 0.5
def to_dict(self) -> dict[str, Any]:
record = _base_record(
label=self.label,
text=self.text,
score=self.score,
frame_index=self.frame_index,
timestamp=self.timestamp,
track_id=self.track_id,
source=self.source,
metadata=self.metadata,
)
record["bbox_xyxy"] = list(self.bbox_xyxy)
if self.polygon is not None:
record["polygon"] = [list(point) for point in self.polygon]
if self.quad is not None:
record["quad"] = [list(point) for point in self.quad]
return record
@classmethod
def from_dict(cls, value: Mapping[str, Any]) -> Detection:
if not isinstance(value, Mapping):
raise TypeError("A detection must be a JSON object.")
return cls(
bbox_xyxy=value["bbox_xyxy"],
label=value.get("label"),
text=value.get("text"),
score=value.get("score"),
polygon=value.get("polygon"),
quad=value.get("quad"),
frame_index=value.get("frame_index", 0),
timestamp=value.get("timestamp", 0.0),
track_id=value.get("track_id"),
source=value.get("source"),
metadata=value.get("metadata"),
)
@dataclass(frozen=True, slots=True)
class FrameDetections:
frame_index: int
timestamp: float
width: int
height: int
detections: tuple[Detection, ...] = ()
metadata: FrozenDict = field(default_factory=FrozenDict)
def __post_init__(self) -> None:
if not isinstance(self.frame_index, int) or self.frame_index < 0:
raise ValueError("frame_index must be a non-negative integer.")
object.__setattr__(
self,
"timestamp",
_non_negative(self.timestamp, "timestamp"),
)
if not isinstance(self.width, int) or self.width <= 0:
raise ValueError("width must be a positive integer.")
if not isinstance(self.height, int) or self.height <= 0:
raise ValueError("height must be a positive integer.")
detections = tuple(self.detections)
for detection in detections:
if not isinstance(detection, Detection):
raise TypeError("detections must contain Detection values.")
if detection.frame_index != self.frame_index:
raise ValueError("Detection frame_index does not match its frame.")
if not math.isclose(detection.timestamp, self.timestamp):
raise ValueError("Detection timestamp does not match its frame.")
x1, y1, x2, y2 = detection.bbox_xyxy
if x1 > self.width or x2 > self.width:
raise ValueError("Detection x coordinates exceed the frame width.")
if y1 > self.height or y2 > self.height:
raise ValueError("Detection y coordinates exceed the frame height.")
for shape in (detection.polygon, detection.quad):
if shape is not None and any(
x > self.width or y > self.height for x, y in shape
):
raise ValueError("Detection geometry exceeds the frame bounds.")
if detection.mask is not None and tuple(detection.mask.shape) != (
self.height,
self.width,
):
raise ValueError("Detection mask shape does not match its frame.")
object.__setattr__(self, "detections", detections)
object.__setattr__(self, "metadata", _metadata(self.metadata))
def to_dict(self) -> dict[str, Any]:
record: dict[str, Any] = {
"frame_index": self.frame_index,
"timestamp": self.timestamp,
"detections": [item.to_dict() for item in self.detections],
}
if self.metadata:
record["metadata"] = self.metadata.to_dict()
return record
@classmethod
def from_dict(
cls,
value: Mapping[str, Any],
*,
width: int,
height: int,
) -> FrameDetections:
if not isinstance(value, Mapping):
raise TypeError("A frame must be a JSON object.")
frame_index = value.get("frame_index", 0)
timestamp = value.get("timestamp", 0.0)
detections = []
for record in value.get("detections", []):
merged = dict(record)
merged.setdefault("frame_index", frame_index)
merged.setdefault("timestamp", timestamp)
detections.append(Detection.from_dict(merged))
return cls(
frame_index=frame_index,
timestamp=timestamp,
width=width,
height=height,
detections=tuple(detections),
metadata=value.get("metadata"),
)
@dataclass(frozen=True, slots=True)
class DetectionSequence:
width: int
height: int
frames: tuple[FrameDetections, ...] = ()
frame_count: int = 0
fps: float | None = None
source: str | None = None
metadata: FrozenDict = field(default_factory=FrozenDict)
version: int = SCHEMA_VERSION
def __post_init__(self) -> None:
if self.version != SCHEMA_VERSION:
raise ValueError(f"Unsupported detection schema version {self.version}.")
if not isinstance(self.width, int) or self.width <= 0:
raise ValueError("width must be a positive integer.")
if not isinstance(self.height, int) or self.height <= 0:
raise ValueError("height must be a positive integer.")
frames = tuple(self.frames)
if any(not isinstance(frame, FrameDetections) for frame in frames):
raise TypeError("frames must contain FrameDetections values.")
indices = [frame.frame_index for frame in frames]
if indices != sorted(indices) or len(indices) != len(set(indices)):
raise ValueError("Frame indices must be unique and increasing.")
timestamps = [frame.timestamp for frame in frames]
if timestamps != sorted(timestamps):
raise ValueError("Frame timestamps must be increasing.")
if any(
frame.width != self.width or frame.height != self.height for frame in frames
):
raise ValueError("Every frame must match the sequence dimensions.")
frame_count = self.frame_count
if not isinstance(frame_count, int) or frame_count < 0:
raise ValueError("frame_count must be a non-negative integer.")
minimum_count = indices[-1] + 1 if indices else 0
if frame_count == 0:
frame_count = minimum_count
elif frame_count < minimum_count:
raise ValueError("frame_count is smaller than the largest frame index.")
fps = None if self.fps is None else _finite(self.fps, "fps")
if fps is not None and fps <= 0:
raise ValueError("fps must be positive.")
object.__setattr__(self, "frames", frames)
object.__setattr__(self, "frame_count", frame_count)
object.__setattr__(self, "fps", fps)
object.__setattr__(self, "source", _optional_text(self.source, "source"))
object.__setattr__(self, "metadata", _metadata(self.metadata))
def all_detections(self) -> tuple[Detection, ...]:
return tuple(
detection for frame in self.frames for detection in frame.detections
)
def frame(self, frame_index: int) -> FrameDetections | None:
return next(
(frame for frame in self.frames if frame.frame_index == frame_index),
None,
)
def to_dict(self) -> dict[str, Any]:
media: dict[str, Any] = {
"width": self.width,
"height": self.height,
"frame_count": self.frame_count,
}
if self.fps is not None:
media["fps"] = self.fps
record: dict[str, Any] = {
"schema": DETECTIONS_SCHEMA,
"version": self.version,
"media": media,
"frames": [frame.to_dict() for frame in self.frames],
}
if self.source is not None:
record["source"] = self.source
if self.metadata:
record["metadata"] = self.metadata.to_dict()
return record
def to_json(self, *, indent: int | None = None) -> str:
return json.dumps(
self.to_dict(),
ensure_ascii=False,
allow_nan=False,
sort_keys=True,
indent=indent,
)
@classmethod
def from_dict(cls, value: Mapping[str, Any]) -> DetectionSequence:
if not isinstance(value, Mapping):
raise TypeError("Detection JSON must contain an object.")
if value.get("schema") != DETECTIONS_SCHEMA:
raise ValueError(f"Expected schema {DETECTIONS_SCHEMA!r}.")
if value.get("version") != SCHEMA_VERSION:
raise ValueError(
f"Unsupported detection schema version {value.get('version')!r}."
)
media = value.get("media")
if not isinstance(media, Mapping):
raise ValueError("Detection JSON requires a media object.")
width, height = media.get("width"), media.get("height")
frames = tuple(
FrameDetections.from_dict(frame, width=width, height=height)
for frame in value.get("frames", [])
)
return cls(
width=width,
height=height,
frames=frames,
frame_count=media.get("frame_count", 0),
fps=media.get("fps"),
source=value.get("source"),
metadata=value.get("metadata"),
version=value["version"],
)
@classmethod
def from_json(cls, value: str) -> DetectionSequence:
if not isinstance(value, str):
raise TypeError("Detection JSON must be a string.")
try:
parsed = json.loads(value)
except json.JSONDecodeError as exc:
raise ValueError(f"Invalid detection JSON: {exc.msg}.") from exc
return cls.from_dict(parsed)
@dataclass(frozen=True, slots=True)
class VisionPoint:
x: float
y: float
label: str | None = None
text: str | None = None
score: float | None = None
frame_index: int = 0
timestamp: float = 0.0
track_id: int | None = None
source: str | None = None
metadata: FrozenDict = field(default_factory=FrozenDict)
def __post_init__(self) -> None:
object.__setattr__(self, "x", _non_negative(self.x, "point x"))
object.__setattr__(self, "y", _non_negative(self.y, "point y"))
object.__setattr__(self, "label", _optional_text(self.label, "label"))
object.__setattr__(self, "text", _optional_text(self.text, "text"))
object.__setattr__(self, "score", _optional_score(self.score))
if not isinstance(self.frame_index, int) or self.frame_index < 0:
raise ValueError("frame_index must be a non-negative integer.")
object.__setattr__(
self,
"timestamp",
_non_negative(self.timestamp, "timestamp"),
)
if self.track_id is not None and (
not isinstance(self.track_id, int) or self.track_id < 0
):
raise ValueError("track_id must be a non-negative integer or None.")
object.__setattr__(self, "source", _optional_text(self.source, "source"))
object.__setattr__(self, "metadata", _metadata(self.metadata))
def to_dict(self) -> dict[str, Any]:
record = _base_record(
label=self.label,
text=self.text,
score=self.score,
frame_index=self.frame_index,
timestamp=self.timestamp,
track_id=self.track_id,
source=self.source,
metadata=self.metadata,
)
record.update(x=self.x, y=self.y)
return record
@classmethod
def from_dict(cls, value: Mapping[str, Any]) -> VisionPoint:
return cls(
x=value["x"],
y=value["y"],
label=value.get("label"),
text=value.get("text"),
score=value.get("score"),
frame_index=value.get("frame_index", 0),
timestamp=value.get("timestamp", 0.0),
track_id=value.get("track_id"),
source=value.get("source"),
metadata=value.get("metadata"),
)
@dataclass(frozen=True, slots=True)
class PointSequence:
width: int
height: int
points: tuple[VisionPoint, ...] = ()
frame_count: int = 0
fps: float | None = None
source: str | None = None
metadata: FrozenDict = field(default_factory=FrozenDict)
version: int = SCHEMA_VERSION
def __post_init__(self) -> None:
if self.version != SCHEMA_VERSION:
raise ValueError(f"Unsupported point schema version {self.version}.")
if not isinstance(self.width, int) or self.width <= 0:
raise ValueError("width must be a positive integer.")
if not isinstance(self.height, int) or self.height <= 0:
raise ValueError("height must be a positive integer.")
points = tuple(self.points)
if any(not isinstance(point, VisionPoint) for point in points):
raise TypeError("points must contain VisionPoint values.")
if any(point.x > self.width or point.y > self.height for point in points):
raise ValueError("Point coordinates exceed the sequence bounds.")
minimum_count = max((point.frame_index for point in points), default=-1) + 1
frame_count = self.frame_count or minimum_count
if not isinstance(frame_count, int) or frame_count < minimum_count:
raise ValueError("frame_count is inconsistent with point frame indices.")
fps = None if self.fps is None else _finite(self.fps, "fps")
if fps is not None and fps <= 0:
raise ValueError("fps must be positive.")
object.__setattr__(self, "points", points)
object.__setattr__(self, "frame_count", frame_count)
object.__setattr__(self, "fps", fps)
object.__setattr__(self, "source", _optional_text(self.source, "source"))
object.__setattr__(self, "metadata", _metadata(self.metadata))
def to_dict(self) -> dict[str, Any]:
media: dict[str, Any] = {
"width": self.width,
"height": self.height,
"frame_count": self.frame_count,
}
if self.fps is not None:
media["fps"] = self.fps
result: dict[str, Any] = {
"schema": POINTS_SCHEMA,
"version": self.version,
"media": media,
"points": [point.to_dict() for point in self.points],
}
if self.source is not None:
result["source"] = self.source
if self.metadata:
result["metadata"] = self.metadata.to_dict()
return result
def to_json(self, *, indent: int | None = None) -> str:
return json.dumps(
self.to_dict(),
ensure_ascii=False,
allow_nan=False,
sort_keys=True,
indent=indent,
)
@classmethod
def from_dict(cls, value: Mapping[str, Any]) -> PointSequence:
if not isinstance(value, Mapping):
raise TypeError("Point JSON must contain an object.")
if value.get("schema") != POINTS_SCHEMA:
raise ValueError(f"Expected schema {POINTS_SCHEMA!r}.")
if value.get("version") != SCHEMA_VERSION:
raise ValueError(
f"Unsupported point schema version {value.get('version')!r}."
)
media = value.get("media")
if not isinstance(media, Mapping):
raise ValueError("Point JSON requires a media object.")
return cls(
width=media.get("width"),
height=media.get("height"),
points=tuple(
VisionPoint.from_dict(point) for point in value.get("points", [])
),
frame_count=media.get("frame_count", 0),
fps=media.get("fps"),
source=value.get("source"),
metadata=value.get("metadata"),
version=value["version"],
)
@classmethod
def from_json(cls, value: str) -> PointSequence:
try:
return cls.from_dict(json.loads(value))
except json.JSONDecodeError as exc:
raise ValueError(f"Invalid point JSON: {exc.msg}.") from exc
@dataclass(frozen=True, slots=True)
class Track:
track_id: int
detections: tuple[Detection, ...]
label: str | None = None
score: float | None = None
source: str | None = None
metadata: FrozenDict = field(default_factory=FrozenDict)
def __post_init__(self) -> None:
if not isinstance(self.track_id, int) or self.track_id < 0:
raise ValueError("track_id must be a non-negative integer.")
detections = tuple(self.detections)
if not detections:
raise ValueError("A track requires at least one detection.")
if any(not isinstance(item, Detection) for item in detections):
raise TypeError("detections must contain Detection values.")
indices = [item.frame_index for item in detections]
if indices != sorted(indices) or len(indices) != len(set(indices)):
raise ValueError("Track detections must have increasing unique frames.")
if any(
item.track_id is not None and item.track_id != self.track_id
for item in detections
):
raise ValueError("Detection track_id does not match its Track.")
object.__setattr__(self, "detections", detections)
object.__setattr__(self, "label", _optional_text(self.label, "label"))
object.__setattr__(self, "score", _optional_score(self.score))
object.__setattr__(self, "source", _optional_text(self.source, "source"))
object.__setattr__(self, "metadata", _metadata(self.metadata))
def to_dict(self) -> dict[str, Any]:
record: dict[str, Any] = {
"track_id": self.track_id,
"detections": [item.to_dict() for item in self.detections],
}
if self.label is not None:
record["label"] = self.label
if self.score is not None:
record["score"] = self.score
if self.source is not None:
record["source"] = self.source
if self.metadata:
record["metadata"] = self.metadata.to_dict()
return record
@classmethod
def from_dict(cls, value: Mapping[str, Any]) -> Track:
if not isinstance(value, Mapping):
raise TypeError("A track must be a JSON object.")
return cls(
track_id=value["track_id"],
detections=tuple(
Detection.from_dict(item) for item in value.get("detections", [])
),
label=value.get("label"),
score=value.get("score"),
source=value.get("source"),
metadata=value.get("metadata"),
)
@dataclass(frozen=True, slots=True)
class TrackSequence:
width: int
height: int
tracks: tuple[Track, ...]
frame_count: int
fps: float | None = None
source: str | None = None
metadata: FrozenDict = field(default_factory=FrozenDict)
version: int = SCHEMA_VERSION
def __post_init__(self) -> None:
if self.version != SCHEMA_VERSION:
raise ValueError(f"Unsupported track schema version {self.version}.")
if not isinstance(self.width, int) or self.width <= 0:
raise ValueError("width must be a positive integer.")
if not isinstance(self.height, int) or self.height <= 0:
raise ValueError("height must be a positive integer.")
if not isinstance(self.frame_count, int) or self.frame_count < 0:
raise ValueError("frame_count must be a non-negative integer.")
tracks = tuple(self.tracks)
if any(not isinstance(track, Track) for track in tracks):
raise TypeError("tracks must contain Track values.")
ids = [track.track_id for track in tracks]
if len(ids) != len(set(ids)):
raise ValueError("Track IDs must be unique.")
for track in tracks:
for detection in track.detections:
if detection.frame_index >= self.frame_count:
raise ValueError("Track detection exceeds frame_count.")
x1, y1, x2, y2 = detection.bbox_xyxy
if x1 > self.width or x2 > self.width:
raise ValueError("Track detection exceeds the frame width.")
if y1 > self.height or y2 > self.height:
raise ValueError("Track detection exceeds the frame height.")
if detection.mask is not None and tuple(detection.mask.shape) != (
self.height,
self.width,
):
raise ValueError("Track mask shape does not match the sequence.")
fps = None if self.fps is None else _finite(self.fps, "fps")
if fps is not None and fps <= 0:
raise ValueError("fps must be positive.")
object.__setattr__(self, "tracks", tracks)
object.__setattr__(self, "fps", fps)
object.__setattr__(self, "source", _optional_text(self.source, "source"))
object.__setattr__(self, "metadata", _metadata(self.metadata))
def to_dict(self) -> dict[str, Any]:
media: dict[str, Any] = {
"width": self.width,
"height": self.height,
"frame_count": self.frame_count,
}
if self.fps is not None:
media["fps"] = self.fps
result: dict[str, Any] = {
"schema": TRACKS_SCHEMA,
"version": self.version,
"media": media,
"tracks": [track.to_dict() for track in self.tracks],
}
if self.source is not None:
result["source"] = self.source
if self.metadata:
result["metadata"] = self.metadata.to_dict()
return result
def to_json(self, *, indent: int | None = None) -> str:
return json.dumps(
self.to_dict(),
ensure_ascii=False,
allow_nan=False,
sort_keys=True,
indent=indent,
)
@classmethod
def from_dict(cls, value: Mapping[str, Any]) -> TrackSequence:
if not isinstance(value, Mapping):
raise TypeError("Track JSON must contain an object.")
if value.get("schema") != TRACKS_SCHEMA:
raise ValueError(f"Expected schema {TRACKS_SCHEMA!r}.")
if value.get("version") != SCHEMA_VERSION:
raise ValueError(
f"Unsupported track schema version {value.get('version')!r}."
)
media = value.get("media")
if not isinstance(media, Mapping):
raise ValueError("Track JSON requires a media object.")
return cls(
width=media.get("width"),
height=media.get("height"),
frame_count=media.get("frame_count", 0),
fps=media.get("fps"),
tracks=tuple(Track.from_dict(track) for track in value.get("tracks", [])),
source=value.get("source"),
metadata=value.get("metadata"),
version=value["version"],
)
@classmethod
def from_json(cls, value: str) -> TrackSequence:
try:
return cls.from_dict(json.loads(value))
except json.JSONDecodeError as exc:
raise ValueError(f"Invalid track JSON: {exc.msg}.") from exc
@dataclass(frozen=True, slots=True)
class TemporalEvent:
start_time: float
end_time: float
label: str | None = None
text: str | None = None
score: float | None = None
source: str | None = None
metadata: FrozenDict = field(default_factory=FrozenDict)
def __post_init__(self) -> None:
start = _non_negative(self.start_time, "start_time")
end = _non_negative(self.end_time, "end_time")
if end < start:
raise ValueError("end_time must be greater than or equal to start_time.")
object.__setattr__(self, "start_time", start)
object.__setattr__(self, "end_time", end)
object.__setattr__(self, "label", _optional_text(self.label, "label"))
object.__setattr__(self, "text", _optional_text(self.text, "text"))
object.__setattr__(self, "score", _optional_score(self.score))
object.__setattr__(self, "source", _optional_text(self.source, "source"))
object.__setattr__(self, "metadata", _metadata(self.metadata))
def to_dict(self) -> dict[str, Any]:
result: dict[str, Any] = {
"start_time": self.start_time,
"end_time": self.end_time,
}
if self.label is not None:
result["label"] = self.label
if self.text is not None:
result["text"] = self.text
if self.score is not None:
result["score"] = self.score
if self.source is not None:
result["source"] = self.source
if self.metadata:
result["metadata"] = self.metadata.to_dict()
return result
@classmethod
def from_dict(cls, value: Mapping[str, Any]) -> TemporalEvent:
if not isinstance(value, Mapping):
raise TypeError("An event must be a JSON object.")
return cls(
start_time=value["start_time"],
end_time=value["end_time"],
label=value.get("label"),
text=value.get("text"),
score=value.get("score"),
source=value.get("source"),
metadata=value.get("metadata"),
)
@dataclass(frozen=True, slots=True)
class EventSequence:
events: tuple[TemporalEvent, ...] = ()
duration: float | None = None
source: str | None = None
metadata: FrozenDict = field(default_factory=FrozenDict)
version: int = SCHEMA_VERSION
def __post_init__(self) -> None:
if self.version != SCHEMA_VERSION:
raise ValueError(f"Unsupported event schema version {self.version}.")
events = tuple(self.events)
if any(not isinstance(event, TemporalEvent) for event in events):
raise TypeError("events must contain TemporalEvent values.")
if list(events) != sorted(
events,
key=lambda event: (event.start_time, event.end_time),
):
raise ValueError("Events must be ordered by start_time.")
duration = (
None if self.duration is None else _non_negative(self.duration, "duration")
)
if duration is not None and any(event.end_time > duration for event in events):
raise ValueError("An event extends beyond the media duration.")
object.__setattr__(self, "events", events)
object.__setattr__(self, "duration", duration)
object.__setattr__(self, "source", _optional_text(self.source, "source"))
object.__setattr__(self, "metadata", _metadata(self.metadata))
def to_dict(self) -> dict[str, Any]:
result: dict[str, Any] = {
"schema": EVENTS_SCHEMA,
"version": self.version,
"events": [event.to_dict() for event in self.events],
}
if self.duration is not None:
result["duration"] = self.duration
if self.source is not None:
result["source"] = self.source
if self.metadata:
result["metadata"] = self.metadata.to_dict()
return result
def to_json(self, *, indent: int | None = None) -> str:
return json.dumps(
self.to_dict(),
ensure_ascii=False,
allow_nan=False,
sort_keys=True,
indent=indent,
)
@classmethod
def from_dict(cls, value: Mapping[str, Any]) -> EventSequence:
if not isinstance(value, Mapping):
raise TypeError("Event JSON must contain an object.")
if value.get("schema") != EVENTS_SCHEMA:
raise ValueError(f"Expected schema {EVENTS_SCHEMA!r}.")
if value.get("version") != SCHEMA_VERSION:
raise ValueError(
f"Unsupported event schema version {value.get('version')!r}."
)
return cls(
events=tuple(
TemporalEvent.from_dict(event) for event in value.get("events", [])
),
duration=value.get("duration"),
source=value.get("source"),
metadata=value.get("metadata"),
version=value["version"],
)
@classmethod
def from_json(cls, value: str) -> EventSequence:
try:
return cls.from_dict(json.loads(value))
except json.JSONDecodeError as exc:
raise ValueError(f"Invalid event JSON: {exc.msg}.") from exc
__all__ = [
"BoxXYXY",
"DETECTIONS_SCHEMA",
"Detection",
"DetectionSequence",
"EVENTS_SCHEMA",
"EventSequence",
"FrameDetections",
"FrozenDict",
"POINTS_SCHEMA",
"PointSequence",
"PointXY",
"Polygon",
"SCHEMA_VERSION",
"TRACKS_SCHEMA",
"TemporalEvent",
"Track",
"TrackSequence",
"VLM_DETECTIONS",
"VLM_EVENTS",
"VLM_POINTS",
"VLM_TRACKS",
"VisionPoint",
]
File diff suppressed because it is too large Load Diff
+25 -3
View File
@@ -1,10 +1,11 @@
[project]
name = "comfyui_vlm_nodes"
version = "2.1.0"
version = "3.0.0"
description = "Production-ready local and API vision-language nodes for ComfyUI"
readme = "README.md"
requires-python = ">=3.10"
license = { file = "LICENSE" }
license = "MIT"
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')",
@@ -15,6 +16,7 @@ dependencies = [
"pydantic>=2.7,<3",
"qwen-vl-utils>=0.0.14",
"safetensors>=0.4.3",
"scipy>=1.10,<2",
"soundfile>=0.12",
"symusic>=0.5",
"transformers>=5.4,<6",
@@ -39,7 +41,7 @@ quantization = [
"bitsandbytes>=0.50,<1",
]
gguf = [
"llama-cpp-python>=0.3.15",
"llama-cpp-python>=0.3.20,<1",
]
[project.urls]
@@ -50,3 +52,23 @@ Issues = "https://github.com/gokayfem/ComfyUI_VLM_nodes/issues"
PublisherId = "gokayfem"
DisplayName = "ComfyUI VLM Nodes"
Icon = ""
[tool.setuptools]
packages = [
"comfyui_vlm_nodes",
"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",
"requirements*.txt",
]
"comfyui_vlm_nodes.web.js" = ["*.js"]
+1 -1
View File
@@ -1,4 +1,4 @@
# 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.15
llama-cpp-python>=0.3.20,<1
+1
View File
@@ -11,6 +11,7 @@ openai>=1.30,<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
transformers>=5.4,<6
+16 -1
View File
@@ -2,10 +2,10 @@
from __future__ import annotations
import importlib.util
import sys
from pathlib import Path
REPOSITORY = Path(__file__).resolve().parents[1]
for candidate in (
REPOSITORY.parent,
@@ -14,3 +14,18 @@ for candidate in (
):
if candidate.exists():
sys.path.insert(0, str(candidate))
# Git worktrees are often intentionally named after a feature branch rather
# than the import package. Load this checkout explicitly so tests can never
# pass by silently importing a sibling clone with the canonical directory name.
if REPOSITORY.name != "ComfyUI_VLM_nodes":
specification = importlib.util.spec_from_file_location(
"ComfyUI_VLM_nodes",
REPOSITORY / "__init__.py",
submodule_search_locations=[str(REPOSITORY)],
)
if specification is None or specification.loader is None:
raise RuntimeError(f"Could not load package from {REPOSITORY}.")
package = importlib.util.module_from_spec(specification)
sys.modules["ComfyUI_VLM_nodes"] = package
specification.loader.exec_module(package)
+98
View File
@@ -0,0 +1,98 @@
"""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 ComfyUI_VLM_nodes.nodes.runtime import (
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()
+250
View File
@@ -0,0 +1,250 @@
from __future__ import annotations
import json
import numpy as np
import pytest
import torch
from ComfyUI_VLM_nodes.nodes import florence2
from PIL import Image
EXPECTED_TASKS = {
"Caption": ("<CAPTION>", "none"),
"Detailed caption": ("<DETAILED_CAPTION>", "none"),
"More detailed caption": ("<MORE_DETAILED_CAPTION>", "none"),
"OCR": ("<OCR>", "none"),
"OCR with regions": ("<OCR_WITH_REGION>", "none"),
"Object detection": ("<OD>", "none"),
"Dense region caption": ("<DENSE_REGION_CAPTION>", "none"),
"Caption to phrase grounding": ("<CAPTION_TO_PHRASE_GROUNDING>", "text"),
"Referring expression segmentation": (
"<REFERRING_EXPRESSION_SEGMENTATION>",
"text",
),
"Region to segmentation": ("<REGION_TO_SEGMENTATION>", "region"),
"Open vocabulary detection": ("<OPEN_VOCABULARY_DETECTION>", "text"),
"Region to category": ("<REGION_TO_CATEGORY>", "region"),
"Region to description": ("<REGION_TO_DESCRIPTION>", "region"),
"Region to OCR": ("<REGION_TO_OCR>", "region"),
"Region proposals": ("<REGION_PROPOSAL>", "none"),
}
def test_registry_covers_all_official_transformers_tasks():
assert len(florence2.TASKS) == 15
assert {
name: (spec.token, spec.input_kind) for name, spec in florence2.TASKS.items()
} == EXPECTED_TASKS
assert {spec.output_kind for spec in florence2.TASKS.values()} == {
"text",
"boxes",
"quad_boxes",
"polygons",
"mixed",
}
def test_node_contract_preserves_outputs_and_adds_core_region_input():
schema = florence2.Florence2.INPUT_TYPES()
assert florence2.NODE_CLASS_MAPPINGS["Florence2"] is florence2.Florence2
assert florence2.Florence2.RETURN_TYPES[:4] == (
"STRING",
"STRING",
"MASK",
"IMAGE",
)
assert florence2.Florence2.RETURN_NAMES[:4] == (
"text",
"structured_json",
"mask",
"visualization",
)
assert schema["optional"]["region"][0] == "BOUNDING_BOX"
assert "forceInput" not in repr(schema["optional"]["region"])
def test_region_encoding_uses_core_xywh_and_florence_location_bins():
region = {"x": 10, "y": 20, "width": 40, "height": 100}
assert florence2._encode_region(region, (100, 200)) == (
"<loc_100><loc_100><loc_500><loc_600>"
)
clamped = {"x": -10, "y": -20, "width": 200, "height": 300}
assert florence2._encode_region(clamped, (100, 200)) == (
"<loc_0><loc_0><loc_999><loc_999>"
)
with pytest.raises(ValueError, match="greater than zero"):
florence2._encode_region(
{"x": 0, "y": 0, "width": 0, "height": 10},
(100, 100),
)
with pytest.raises(ValueError, match="does not overlap"):
florence2._encode_region(
{"x": 200, "y": 200, "width": 10, "height": 10},
(100, 100),
)
def test_task_inputs_are_validated_before_inference():
image_size = (100, 100)
region = {"x": 10, "y": 10, "width": 20, "height": 20}
assert florence2._task_extra_input("Caption", "", None, image_size) == ""
with pytest.raises(ValueError, match="does not accept text"):
florence2._task_extra_input("Caption", "unexpected", None, image_size)
with pytest.raises(ValueError, match="requires text"):
florence2._task_extra_input("Open vocabulary detection", "", None, image_size)
assert (
florence2._task_extra_input(
"Open vocabulary detection", "red car", None, image_size
)
== "red car"
)
with pytest.raises(ValueError, match="requires a connected BOUNDING_BOX"):
florence2._task_extra_input("Region to OCR", "", None, image_size)
assert (
florence2._task_extra_input("Region to OCR", "", region, image_size)
== "<loc_100><loc_100><loc_300><loc_300>"
)
with pytest.raises(ValueError, match="does not accept text"):
florence2._task_extra_input("Region to OCR", "also text", region, image_size)
def test_region_selection_accepts_core_and_batched_detector_shapes():
first = {"x": 1, "y": 2, "width": 3, "height": 4}
second = {"x": 5, "y": 6, "width": 7, "height": 8}
assert florence2._select_region(first, 0, 2) is first
assert florence2._select_region([first, second], 1, 2) is second
assert florence2._select_region([[first], [second]], 0, 2) is first
with pytest.raises(ValueError, match="exactly one"):
florence2._select_region([[first, second]], 0, 1)
def test_visualization_is_deterministic_and_masks_every_spatial_shape():
image = Image.new("RGB", (48, 36), "black")
parsed = {
"<OPEN_VOCABULARY_DETECTION>": {
"bboxes": [[1, 1, 10, 10]],
"bboxes_labels": ["box"],
"quad_boxes": [[14, 1, 22, 1, 22, 10, 14, 10]],
"labels": ["ocr"],
"polygons": [[[26, 1, 40, 1, 40, 12, 26, 12]]],
"polygons_labels": ["polygon"],
}
}
mask_a, visual_a = florence2._visualize(image, parsed)
mask_b, visual_b = florence2._visualize(image, parsed)
mask = np.asarray(mask_a)
assert mask[5, 5] == 255
assert mask[5, 18] == 255
assert mask[5, 30] == 255
assert mask_a.tobytes() == mask_b.tobytes()
assert visual_a.tobytes() == visual_b.tobytes()
def test_predictor_generation_is_deterministic_without_downloads():
calls = {}
class FakeProcessor:
def __call__(self, text, images, return_tensors):
calls["prompt"] = text
assert images.size == (8, 8)
assert return_tensors == "pt"
return {
"input_ids": torch.tensor([[1]], dtype=torch.long),
"pixel_values": torch.zeros((1, 3, 8, 8)),
}
def batch_decode(self, generated, skip_special_tokens):
assert skip_special_tokens is False
return ["<s>answer</s>"]
def post_process_generation(self, raw, task, image_size):
return {task: "answer"}
class FakeModel(torch.nn.Module):
def __init__(self):
super().__init__()
self.anchor = torch.nn.Parameter(torch.zeros(()))
def generate(self, **kwargs):
calls["generation"] = kwargs
return torch.tensor([[2]], dtype=torch.long)
class FakeHandle:
def __init__(self):
self.model = FakeModel()
def ensure_loaded(self):
return self.model
predictor = object.__new__(florence2.FlorencePredictor)
predictor.dtype = torch.float32
predictor.processor = FakeProcessor()
predictor.handle = FakeHandle()
raw, parsed = predictor.run(
Image.new("RGB", (8, 8)),
"<CAPTION>",
"",
32,
3,
)
assert raw == "<s>answer</s>"
assert parsed == {"<CAPTION>": "answer"}
assert calls["prompt"] == "<CAPTION>"
assert calls["generation"]["do_sample"] is False
assert calls["generation"]["num_beams"] == 3
assert calls["generation"]["early_stopping"] is True
def test_node_cleans_text_preserves_structured_data_and_unloads_target():
parsed = {
"<OD>": {
"bboxes": [[2, 2, 12, 12]],
"labels": ["person"],
"quad_boxes": [[14, 2, 22, 2, 22, 12, 14, 12]],
"polygons": [[[24, 2, 30, 2, 30, 12, 24, 12]]],
}
}
calls = []
class FakePredictor:
def run(
self,
image,
task_token,
extra_input,
max_new_tokens,
beams,
):
calls.append((image.size, task_token, extra_input, max_new_tokens, beams))
return "<s>person<loc_1><loc_2></s><pad>", parsed
node = florence2.Florence2()
node.get_or_create_model = lambda key, factory: FakePredictor()
unloads = []
node.maybe_clear_model = unloads.append
output = node.run(
torch.zeros((1, 32, 32, 3)),
"Object detection",
"",
"Florence-2 base FT (fast)",
64,
1,
unload_after=True,
)
assert len(output) == 4
assert output[0] == "person<loc_1><loc_2>"
assert json.loads(output[1]) == [parsed]
assert output[2].shape == (1, 32, 32)
assert output[2].max().item() == 1.0
assert output[3].shape == (1, 32, 32, 3)
assert calls == [((32, 32), "<OD>", "", 64, 1)]
assert unloads == [True]
+166
View File
@@ -0,0 +1,166 @@
import importlib
from pathlib import Path
import pytest
import torch
PACKAGE = Path(__file__).resolve().parents[1].name
geometry = importlib.import_module(f"{PACKAGE}.nodes.geometry")
vision_types = importlib.import_module(f"{PACKAGE}.nodes.vision_types")
associate_detections = geometry.associate_detections
bbox_from_mask = geometry.bbox_from_mask
bbox_iou = geometry.bbox_iou
box_area = geometry.box_area
box_center = geometry.box_center
box_to_mask = geometry.box_to_mask
clip_box = geometry.clip_box
clip_polygon = geometry.clip_polygon
denormalize_box = geometry.denormalize_box
detection_to_mask = geometry.detection_to_mask
deterministic_color = geometry.deterministic_color
expand_box = geometry.expand_box
individual_detection_masks = geometry.individual_detection_masks
mask_iou = geometry.mask_iou
normalize_box = geometry.normalize_box
polygon_area = geometry.polygon_area
polygon_to_mask = geometry.polygon_to_mask
quad_to_mask = geometry.quad_to_mask
translate_box = geometry.translate_box
union_detection_mask = geometry.union_detection_mask
Detection = vision_types.Detection
def test_box_clipping_normalization_area_and_center():
assert clip_box((0, 2, 25, 22), 20, 10) == (0, 2, 20, 10)
normalized = normalize_box((5, 2, 15, 8), 20, 10)
assert normalized == pytest.approx((0.25, 0.2, 0.75, 0.8))
assert denormalize_box(normalized, 20, 10) == pytest.approx((5, 2, 15, 8))
assert box_area((5, 2, 15, 8)) == 60
assert box_center((5, 2, 15, 8)) == (10, 5)
with pytest.raises(ValueError, match="Normalized"):
denormalize_box((0, 0, 2, 1), 20, 10)
with pytest.raises(ValueError, match="x2"):
clip_box((2, 0, 1, 1), 20, 10)
def test_polygon_clipping_and_area():
polygon = ((-2, -3), (8, 0), (8, 5), (0, 5))
assert clip_polygon(polygon, 6, 4) == (
(0, 0),
(6, 0),
(6, 4),
(0, 4),
)
assert polygon_area(((0, 0), (5, 0), (5, 4), (0, 4))) == 20
assert polygon_area(((0, 0), (0, 4), (5, 4), (5, 0))) == 20
with pytest.raises(ValueError, match="at least three"):
polygon_area(((0, 0), (1, 1)))
def test_bbox_and_mask_iou():
assert bbox_iou((0, 0, 10, 10), (5, 0, 15, 10)) == pytest.approx(1 / 3)
assert bbox_iou((0, 0, 1, 1), (2, 2, 3, 3)) == 0
first = torch.zeros((4, 4))
second = torch.zeros((4, 4))
first[:2, :2] = 1
second[1:3, :2] = 1
assert mask_iou(first, second) == pytest.approx(1 / 3)
assert mask_iou(torch.zeros((2, 2)), torch.zeros((2, 2))) == 0
with pytest.raises(ValueError, match="same shape"):
mask_iou(torch.zeros((2, 2)), torch.zeros((3, 2)))
def test_box_polygon_and_quad_rasterization():
box = box_to_mask((1.2, 2.1, 4.1, 5.2), 8, 7)
assert box.shape == (7, 8)
assert box.sum().item() == 16
polygon = polygon_to_mask(((1, 1), (5, 1), (5, 5), (1, 5)), 8, 8)
quad = quad_to_mask(((1, 1), (5, 1), (5, 5), (1, 5)), 8, 8)
assert torch.equal(polygon, quad)
assert polygon.sum() > 0
with pytest.raises(ValueError, match="exactly four"):
quad_to_mask(((0, 0), (1, 0), (1, 1)), 4, 4)
def test_detection_mask_priority_union_individual_and_bbox():
explicit = torch.zeros((8, 8))
explicit[3:6, 2:5] = 1
with_mask = Detection(
bbox_xyxy=(0, 0, 8, 8),
polygon=((0, 0), (8, 0), (8, 8), (0, 8)),
mask=explicit,
)
polygon_only = Detection(
bbox_xyxy=(1, 1, 5, 5),
polygon=((1, 1), (5, 1), (5, 5), (1, 5)),
)
assert torch.equal(detection_to_mask(with_mask, 8, 8), explicit)
masks = individual_detection_masks((with_mask, polygon_only), 8, 8)
assert masks.shape == (2, 8, 8)
union = union_detection_mask((with_mask, polygon_only), 8, 8)
assert union.shape == (8, 8)
assert torch.all(union >= masks[0])
assert individual_detection_masks((), 8, 8).shape == (0, 8, 8)
assert union_detection_mask((), 8, 8).sum() == 0
assert bbox_from_mask(explicit) == (2, 3, 5, 6)
assert bbox_from_mask(torch.zeros((2, 2))) is None
def test_deterministic_color_and_box_expansion():
assert deterministic_color("track-1") == deterministic_color("track-1")
assert deterministic_color("track-1") != deterministic_color("track-2")
assert all(0 <= channel <= 255 for channel in deterministic_color("object"))
assert expand_box((4, 4, 8, 6), 12, 12, padding=1) == (3, 3, 9, 7)
squared = expand_box((4, 4, 8, 6), 12, 12, square=True)
assert squared == (4, 3, 8, 7)
assert translate_box((1, 2, 3, 4), 2, 1) == (3, 3, 5, 5)
def test_label_aware_stable_association_with_motion():
previous = (
Detection(
bbox_xyxy=(0, 0, 10, 10),
label="cat",
track_id=4,
),
Detection(
bbox_xyxy=(20, 0, 30, 10),
label="dog",
track_id=9,
),
)
current = (
Detection(bbox_xyxy=(5, 0, 15, 10), label="cat"),
Detection(bbox_xyxy=(20, 0, 30, 10), label="bird"),
Detection(bbox_xyxy=(40, 0, 50, 10), label="dog"),
)
without_motion = associate_detections(
previous,
current,
minimum_iou=0.3,
)
assert without_motion.matches == ((0, 0, pytest.approx(1 / 3)),)
assert without_motion.unmatched_previous == (1,)
assert without_motion.unmatched_current == (1, 2)
with_motion = associate_detections(
previous,
current,
minimum_iou=0.9,
motion_by_track={4: (5, 0), 9: (20, 0)},
)
assert with_motion.matches == (
(0, 0, 1.0),
(1, 2, 1.0),
)
assert with_motion.unmatched_previous == ()
assert with_motion.unmatched_current == (1,)
label_agnostic = associate_detections(
previous,
current,
minimum_iou=0.9,
label_aware=False,
)
assert label_agnostic.matches == ((1, 1, 1.0),)
+141
View File
@@ -0,0 +1,141 @@
from types import SimpleNamespace
import torch
from ComfyUI_VLM_nodes.nodes.grounding import (
MODEL_SPECS,
VLMOpenVocabularyDetection,
core_bounding_box_frames,
core_bounding_boxes,
detection_box_masks,
parse_labels,
result_to_detections,
)
from ComfyUI_VLM_nodes.nodes.vision_types import (
DetectionSequence,
FrameDetections,
)
def test_detector_catalog_is_small_fast_and_portable():
assert "Grounding DINO Tiny (fast)" in MODEL_SPECS
assert "OmDet Turbo Swin Tiny (fast)" in MODEL_SPECS
assert all("/" in spec.model_id for spec in MODEL_SPECS.values())
schema = VLMOpenVocabularyDetection.INPUT_TYPES()
assert tuple(MODEL_SPECS) == schema["required"]["model"][0]
def test_label_parser_preserves_phrases_and_removes_duplicates():
assert parse_labels("red car, person\nsmall dog;person") == [
"red car",
"person",
"small dog",
]
def test_transformers_results_are_clipped_sorted_and_normalized():
result = {
"boxes": torch.tensor([[-5.0, 2.0, 20.0, 12.0], [5.0, 5.0, 9.0, 9.0]]),
"scores": torch.tensor([0.25, 0.9]),
"text_labels": ["cat", "dog"],
}
detections = result_to_detections(
result,
labels=["cat", "dog"],
width=16,
height=10,
frame_index=0,
timestamp=0.0,
source="test/model",
max_detections=20,
)
assert [item.label for item in detections] == ["dog", "cat"]
assert detections[1].bbox_xyxy == (0.0, 2.0, 16.0, 10.0)
assert detections[0].metadata["model_id"] == "test/model"
def test_max_detections_is_applied_after_confidence_sorting():
detections = result_to_detections(
{
"boxes": [[0, 0, 1, 1], [1, 1, 2, 2], [2, 2, 3, 3]],
"scores": [0.1, 0.9, 0.8],
"text_labels": ["low", "best", "second"],
},
labels=["low", "best", "second"],
width=4,
height=4,
frame_index=0,
timestamp=0.0,
source="test/model",
max_detections=2,
)
assert [item.label for item in detections] == ["best", "second"]
def test_box_masks_and_core_boxes_keep_geometry_and_metadata():
detections = result_to_detections(
{
"boxes": torch.tensor([[1.0, 2.0, 4.0, 5.0]]),
"scores": torch.tensor([0.8]),
"labels": torch.tensor([0]),
},
labels=["cat"],
width=8,
height=6,
frame_index=0,
timestamp=0.0,
source="test/model",
max_detections=5,
)
sequence = DetectionSequence(
width=8,
height=6,
frames=(
FrameDetections(
frame_index=0,
timestamp=0.0,
width=8,
height=6,
detections=detections,
),
),
frame_count=1,
)
masks = detection_box_masks(sequence)
assert masks.shape == (1, 6, 8)
assert masks.sum().item() == 9
boxes = core_bounding_boxes(sequence)
assert boxes == [
{
"x": 1,
"y": 2,
"width": 3,
"height": 3,
"label": "cat",
"score": detections[0].score,
"metadata": {
"frame_index": 0,
"label": "cat",
"score": detections[0].score,
"source": "test/model",
},
}
]
assert core_bounding_box_frames(sequence) == [boxes]
def test_result_label_indices_are_resolved():
detections = result_to_detections(
{
"boxes": [[0, 0, 4, 4]],
"scores": [SimpleNamespace(item=lambda: 0.5)],
"classes": [SimpleNamespace(item=lambda: 1)],
},
labels=["cat", "dog"],
width=4,
height=4,
frame_index=0,
timestamp=0.0,
source="omdet",
max_detections=1,
)
assert detections[0].label == "dog"
+288 -27
View File
@@ -1,14 +1,14 @@
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 PIL import Image
import ComfyUI_VLM_nodes as package
from ComfyUI_VLM_nodes.nodes import (
audioldm2,
florence2,
@@ -16,16 +16,24 @@ from ComfyUI_VLM_nodes.nodes import (
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():
@@ -86,6 +94,7 @@ def test_runtime_report_and_device_map_are_supportable():
"torch_cuda",
"torch_hip",
"packages",
"llama_cpp",
} <= report.keys()
device_map = external_device_map()
assert set(device_map) == {""}
@@ -137,6 +146,202 @@ def test_dependency_metadata_matches_installer_requirements():
{"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(
@@ -181,18 +386,14 @@ def test_florence_rendering_supports_boxes_quads_and_nested_polygons():
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
]
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 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
@@ -212,12 +413,8 @@ def test_modern_catalog_has_current_quality_and_low_vram_tiers():
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
)
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):
@@ -250,13 +447,77 @@ def test_modern_video_is_primary_input_and_thinking_is_explicit():
assert captured["video_metadata"]["frames_indices"] == [0, 1, 2, 3]
def test_internvl_video_uses_an_even_vision_patch_grid():
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", "InternVL 3.5", 1.0, video=True
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
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:
@@ -285,12 +546,14 @@ def test_qwen2_legacy_quantized_labels_use_maintained_backends():
"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)")
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):
@@ -300,9 +563,7 @@ def test_audioldm_keeps_legacy_outputs_and_adds_standard_audio(monkeypatch):
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"
)
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)
+251
View File
@@ -0,0 +1,251 @@
from types import SimpleNamespace
import pytest
import torch
from ComfyUI_VLM_nodes.nodes.sam2 import (
SAM2_MODELS,
Sam2Spec,
Sam2VideoPredictor,
VLMSAM2VideoSegmentation,
_core_box,
_normalize_processed_masks,
seed_boxes,
)
from ComfyUI_VLM_nodes.nodes.vision_types import (
Detection,
DetectionSequence,
FrameDetections,
)
def test_sam2_catalog_leads_with_tiny_and_has_no_30b_models():
assert next(iter(SAM2_MODELS)) == "SAM2.1 Hiera Tiny (fast)"
assert all("30b" not in spec.model_id.lower() for spec in SAM2_MODELS.values())
assert VLMSAM2VideoSegmentation.RETURN_NAMES[0:2] == ("tracks", "json")
def test_core_box_is_converted_from_xywh():
assert _core_box({"x": 2, "y": 3, "width": 5, "height": 7}) == (
2.0,
3.0,
7.0,
10.0,
)
with pytest.raises(ValueError, match="positive"):
_core_box({"x": 2, "y": 3, "width": 0, "height": 7})
def test_detection_seeds_keep_labels_and_ids():
detection = Detection(
bbox_xyxy=(1, 2, 9, 10),
label="cat",
frame_index=2,
timestamp=0.2,
track_id=7,
)
sequence = DetectionSequence(
width=10,
height=10,
frames=(
FrameDetections(
frame_index=2,
timestamp=0.2,
width=10,
height=10,
detections=(detection,),
),
),
frame_count=3,
fps=10,
)
boxes, ids, labels = seed_boxes(
width=10,
height=10,
frame_index=2,
detections=sequence,
bounding_box=None,
)
assert boxes == [[1.0, 2.0, 9.0, 10.0]]
assert ids == [7]
assert labels == {7: "cat"}
def test_processed_mask_shapes_are_normalized():
assert _normalize_processed_masks(torch.zeros(2, 1, 4, 5)).shape == (
2,
4,
5,
)
assert _normalize_processed_masks(torch.zeros(4, 5)).shape == (1, 4, 5)
with pytest.raises(RuntimeError, match="unsupported"):
_normalize_processed_masks(torch.zeros(1, 2, 3, 4, 5))
def test_video_session_runs_seed_frame_before_both_propagation_directions():
class FakeProcessor:
def init_video_session(self, **_kwargs):
return SimpleNamespace(obj_ids=[1], seed_inferred=False)
def add_inputs_to_inference_session(self, **kwargs):
kwargs["inference_session"].seed_frame = kwargs["frame_idx"]
def post_process_masks(self, masks, **_kwargs):
return masks
class FakeModel(torch.nn.Module):
def __init__(self):
super().__init__()
self.anchor = torch.nn.Parameter(torch.zeros(()))
self.seed_calls = []
self.propagation_directions = []
def forward(self, inference_session, frame_idx):
inference_session.seed_inferred = True
self.seed_calls.append(frame_idx)
return SimpleNamespace(
frame_idx=frame_idx,
pred_masks=torch.ones(1, 1, 4, 4),
)
def propagate_in_video_iterator(
self,
inference_session,
start_frame_idx,
reverse=False,
**_kwargs,
):
assert inference_session.seed_inferred
self.propagation_directions.append(reverse)
indices = (
range(start_frame_idx, 3)
if not reverse
else range(start_frame_idx, -1, -1)
)
for frame_index in indices:
yield SimpleNamespace(
frame_idx=frame_index,
pred_masks=torch.ones(1, 1, 4, 4),
)
model = FakeModel()
predictor = Sam2VideoPredictor.__new__(Sam2VideoPredictor)
predictor.processor = FakeProcessor()
predictor.dtype = torch.float32
predictor.spec = Sam2Spec("test/sam2", "test-sam2")
predictor.handle = SimpleNamespace(ensure_loaded=lambda: model)
detections = DetectionSequence(
width=4,
height=4,
frames=(
FrameDetections(
frame_index=0,
timestamp=0.0,
width=4,
height=4,
detections=(
Detection(
bbox_xyxy=(0, 0, 4, 4),
label="object",
frame_index=0,
timestamp=0.0,
),
),
),
),
frame_count=1,
)
tracks, union, individual, preview = predictor.propagate(
torch.zeros(3, 4, 4, 3),
seed_frame=1,
fps=10.0,
detections=detections,
bounding_box=None,
seed_mask=None,
mask_threshold=0.0,
keep_video_on_cpu=True,
mask_output="union_and_objects",
render_preview=True,
)
assert model.seed_calls == [1]
assert model.propagation_directions == [False, True]
assert [item.frame_index for item in tracks.tracks[0].detections] == [0, 1, 2]
assert union.shape == (3, 4, 4)
assert individual.shape == (3, 4, 4)
assert preview.shape == (3, 4, 4, 3)
assert torch.equal(preview[0], preview[1])
assert torch.equal(preview[1], preview[2])
def test_multi_object_mask_seeds_are_passed_as_one_mask_per_object():
class FakeProcessor:
def __init__(self):
self.received_masks = None
def init_video_session(self, **kwargs):
assert str(kwargs["processing_device"]) == "cpu"
return SimpleNamespace(obj_ids=[1, 2], seed_inferred=False)
def add_inputs_to_inference_session(self, **kwargs):
self.received_masks = kwargs["input_masks"]
def post_process_masks(self, masks, **_kwargs):
return masks
class FakeModel(torch.nn.Module):
def __init__(self):
super().__init__()
self.anchor = torch.nn.Parameter(torch.zeros(()))
def forward(self, inference_session, frame_idx):
inference_session.seed_inferred = True
return SimpleNamespace(
frame_idx=frame_idx,
pred_masks=torch.ones(2, 1, 4, 4),
)
def propagate_in_video_iterator(self, **_kwargs):
return iter(())
processor = FakeProcessor()
predictor = Sam2VideoPredictor.__new__(Sam2VideoPredictor)
predictor.processor = processor
predictor.dtype = torch.float32
predictor.spec = Sam2Spec("test/sam2", "test-sam2")
predictor.handle = SimpleNamespace(ensure_loaded=FakeModel)
images = torch.zeros(1, 4, 4, 3)
tracks, _union, individual, preview = predictor.propagate(
images,
seed_frame=0,
fps=24.0,
detections=None,
bounding_box=None,
seed_mask=torch.ones(2, 4, 4),
mask_threshold=0.0,
keep_video_on_cpu=True,
mask_output="union_only",
render_preview=False,
)
assert isinstance(processor.received_masks, list)
assert len(processor.received_masks) == 2
assert len(tracks.tracks) == 2
assert individual.shape == (0, 4, 4)
assert preview.data_ptr() == images.data_ptr()
def test_nested_core_boxes_select_seed_frame_and_keep_top_level_labels():
boxes, ids, labels = seed_boxes(
width=20,
height=20,
frame_index=1,
detections=None,
bounding_box=[
[{"x": 0, "y": 0, "width": 2, "height": 2, "label": "old"}],
[
{"x": 3, "y": 4, "width": 5, "height": 6, "label": "person"},
{"x": 10, "y": 11, "width": 4, "height": 3},
],
],
)
assert boxes == [[3.0, 4.0, 8.0, 10.0], [10.0, 11.0, 14.0, 14.0]]
assert ids == [1, 2]
assert labels == {1: "person", 2: None}
+213
View File
@@ -0,0 +1,213 @@
from __future__ import annotations
import json
import pytest
import torch
from ComfyUI_VLM_nodes.nodes.sam3_adapter import (
VLMTrackReport,
iter_sam3_masks,
sam3_track_data_to_tracks,
track_report_json,
track_report_payload,
track_report_text,
unpack_sam3_mask,
validate_sam3_track_data,
)
from ComfyUI_VLM_nodes.nodes.vision_types import (
Detection,
DetectionSequence,
FrameDetections,
)
def _pack_masks(masks):
masks = masks.to(torch.uint8)
width = masks.shape[-1]
assert width % 8 == 0
bits = 1 << torch.arange(8, dtype=torch.int64)
grouped = masks.reshape(*masks.shape[:-1], width // 8, 8)
return (grouped * bits).sum(dim=-1).to(torch.uint8)
def _sample_track_data():
masks = torch.zeros(2, 2, 4, 8, dtype=torch.bool)
masks[0, 0, 1:3, 2:5] = True
masks[1, 0, 1:4, 3:6] = True
masks[1, 1, 0:2, 0:2] = True
return {
"packed_masks": _pack_masks(masks),
"n_frames": 2,
"scores": [0.9, 0.75],
"orig_size": (40, 80),
}
def _seed_detections():
detections = (
Detection(
bbox_xyxy=(20, 10, 50, 30),
label="cat",
text="the cat",
score=0.95,
frame_index=0,
timestamp=0.0,
track_id=7,
source="seed",
),
Detection(
bbox_xyxy=(0, 0, 20, 20),
label="fish",
score=0.8,
frame_index=0,
timestamp=0.0,
track_id=9,
source="seed",
),
)
return DetectionSequence(
width=80,
height=40,
frames=(
FrameDetections(
frame_index=0,
timestamp=0.0,
width=80,
height=40,
detections=detections,
),
),
frame_count=2,
fps=10.0,
)
def test_validate_and_stream_unpack_without_expanding_the_video():
track_data = _sample_track_data()
layout = validate_sam3_track_data(track_data)
assert (
layout.n_frames,
layout.n_objects,
layout.mask_height,
layout.mask_width,
) == (2, 2, 4, 8)
yielded = list(iter_sam3_masks(track_data, present_only=True))
assert [(frame, obj) for frame, obj, _mask in yielded] == [
(0, 0),
(1, 0),
(1, 1),
]
assert all(mask.shape == (4, 8) for _frame, _obj, mask in yielded)
assert yielded[0][2].sum().item() == 6
one = unpack_sam3_mask(track_data["packed_masks"][1, 1])
assert one.dtype == torch.bool
assert one.sum().item() == 4
def test_adapter_derives_scaled_boxes_and_preserves_seed_identity():
tracks = sam3_track_data_to_tracks(
_sample_track_data(),
seed_detections=_seed_detections(),
fps=10.0,
)
assert (tracks.width, tracks.height, tracks.frame_count, tracks.fps) == (
80,
40,
2,
10.0,
)
assert [track.track_id for track in tracks.tracks] == [7, 9]
assert [track.label for track in tracks.tracks] == ["cat", "fish"]
cat, fish = tracks.tracks
assert cat.detections[0].bbox_xyxy == (20.0, 10.0, 50.0, 30.0)
assert cat.detections[1].bbox_xyxy == (30.0, 10.0, 60.0, 40.0)
assert fish.detections[0].frame_index == 1
assert fish.detections[0].bbox_xyxy == (0.0, 0.0, 20.0, 20.0)
assert cat.detections[0].metadata["mask_ref"]["object_index"] == 0
assert fish.detections[0].metadata["mask_ref"]["object_index"] == 1
assert cat.metadata["seeded"] is True
assert fish.metadata["seeded"] is True
serialized = tracks.to_json()
assert "mask_ref" in serialized
assert "packed_masks" not in serialized
assert "tensor" not in serialized.casefold()
def test_adapter_handles_an_empty_core_result():
tracks = sam3_track_data_to_tracks(
{
"packed_masks": None,
"n_frames": 3,
"scores": [],
"orig_size": (48, 64),
},
fps=24.0,
)
assert tracks.tracks == ()
assert tracks.frame_count == 3
assert tracks.metadata["object_slots"] == 0
@pytest.mark.parametrize(
("track_data", "message"),
(
(
{"packed_masks": None, "n_frames": 1, "scores": []},
"orig_size",
),
(
{
"packed_masks": torch.zeros(1, 1, 2, 1),
"n_frames": 1,
"scores": [0.5],
"orig_size": (2, 8),
},
"uint8",
),
(
{
"packed_masks": torch.zeros(2, 1, 2, 1, dtype=torch.uint8),
"n_frames": 1,
"scores": [0.5],
"orig_size": (2, 8),
},
"n_frames",
),
(
{
"packed_masks": torch.zeros(1, 1, 2, 1, dtype=torch.uint8),
"n_frames": 1,
"scores": [1.5],
"orig_size": (2, 8),
},
"0 to 1",
),
),
)
def test_adapter_rejects_incompatible_private_payloads(track_data, message):
with pytest.raises((TypeError, ValueError), match=message):
validate_sam3_track_data(track_data)
def test_track_report_is_small_deterministic_and_history_safe():
tracks = sam3_track_data_to_tracks(
_sample_track_data(),
seed_detections=_seed_detections(),
fps=10.0,
)
payload = track_report_payload(tracks)
assert payload["track_count"] == 2
assert payload["observation_count"] == 3
assert payload["state_counts"] == {"active": 2}
encoded = track_report_json(tracks)
assert json.loads(encoded) == payload
assert "packed_masks" not in encoded
text = track_report_text(tracks)
assert "Tracks: 2" in text
assert "#7 cat" in text
assert "#9 fish" in text
node_result = VLMTrackReport().report(tracks)
assert node_result["result"] == (encoded, text)
assert node_result["ui"]["text"] == [text]
+376
View File
@@ -0,0 +1,376 @@
from __future__ import annotations
import json
import pytest
from ComfyUI_VLM_nodes.nodes.spatial_parser import (
COORDINATE_MODES,
VLMSpatialPromptBuilder,
VLMStructuredSpatialParser,
build_spatial_prompt,
load_json_document,
parse_spatial_response,
)
from ComfyUI_VLM_nodes.nodes.vision_types import (
VLM_DETECTIONS,
VLM_POINTS,
DetectionSequence,
PointSequence,
)
def test_prompt_builder_is_explicit_and_provider_neutral():
prompt = build_spatial_prompt(
"Find every vehicle.",
coordinate_mode="normalized_0_1000",
width=1920,
height=1080,
frame_count=12,
fps=24.0,
)
assert prompt.startswith("Perform this visual analysis task:")
assert "Find every vehicle." in prompt
assert "Return only one valid JSON object" in prompt
assert "normalized_0_1000" in prompt
assert '"frame_count":12' in prompt
assert '"fps":24.0' in prompt
assert "zero-based frame_index" in prompt
assert "bbox_xyxy" in prompt
assert "polygon" in prompt
assert '"point"' in prompt
assert "score" in prompt
example = json.loads(prompt.split("Required JSON shape:\n", 1)[1])
assert example["coordinate_mode"] == "normalized_0_1000"
def test_node_contracts_use_canonical_spatial_types():
assert tuple(COORDINATE_MODES) == (
"pixel",
"normalized_0_1",
"normalized_0_1000",
)
assert VLMSpatialPromptBuilder.RETURN_TYPES == ("STRING",)
assert VLMStructuredSpatialParser.RETURN_TYPES == (
VLM_DETECTIONS,
VLM_POINTS,
"STRING",
)
schema = VLMStructuredSpatialParser.INPUT_TYPES()
assert tuple(schema["required"]["coordinate_mode"][0]) == COORDINATE_MODES
def test_parser_accepts_only_complete_plain_or_fenced_json():
assert load_json_document(" ") == {}
assert load_json_document('```json\n{"frames":[]}\n```') == {"frames": []}
for invalid in (
'Here is the result: {"frames":[]}',
'Result:\n```json\n{"frames":[]}\n```',
'{"frames":[]} trailing',
"```python\n{}\n```",
'{"x": 1, "x": 2}',
'{"score": NaN}',
):
with pytest.raises(ValueError):
load_json_document(invalid)
def test_normalized_video_parse_clips_and_preserves_metadata():
response = json.dumps(
{
"coordinate_mode": "normalized_0_1",
"media": {
"width": 200,
"height": 100,
"frame_count": 3,
"fps": 2,
"codec": "test-codec",
},
"source": "unit-vlm",
"metadata": {"request_id": "abc"},
"vendor": {"latency_ms": 12},
"frames": [
{
"frame_index": 0,
"metadata": {"scene": "start"},
"detections": [
{
"class": "cat",
"confidence": 0.75,
"bbox": [-0.1, 0.2, 1.2, 0.8],
"polygon": [
[-0.5, 0.2],
[0.5, 0.2],
[0.5, 1.5],
],
"instance_id": "cat-1",
"metadata": {"occluded": False},
}
],
"points": [
{
"name": "nose",
"point": [0.25, 0.5],
"confidence": 0.9,
"landmark_id": 7,
}
],
},
{
"frame_index": 2,
"detections": [
{
"label": "sign",
"quad": [
[0.1, 0.1],
[0.9, 0.1],
[0.9, 0.9],
[0.1, 0.9],
],
"text": "STOP",
}
],
},
],
}
)
detections, points, normalized_json = parse_spatial_response(
f"```json\n{response}\n```",
width=200,
height=100,
coordinate_mode="normalized_0_1",
)
assert isinstance(detections, DetectionSequence)
assert isinstance(points, PointSequence)
assert detections.frame_count == points.frame_count == 3
assert detections.fps == points.fps == 2.0
assert [frame.frame_index for frame in detections.frames] == [0, 2]
cat, sign = detections.all_detections()
assert cat.bbox_xyxy == (0.0, 20.0, 200.0, 80.0)
assert cat.polygon == ((0.0, 20.0), (100.0, 20.0), (100.0, 100.0))
assert cat.label == "cat"
assert cat.score == 0.75
assert cat.metadata.to_dict() == {
"instance_id": "cat-1",
"occluded": False,
}
assert sign.bbox_xyxy == (20.0, 10.0, 180.0, 90.0)
assert sign.quad is not None and len(sign.quad) == 4
assert sign.text == "STOP"
assert points.points[0].x == 50.0
assert points.points[0].y == 50.0
assert points.points[0].label == "nose"
assert points.points[0].metadata["landmark_id"] == 7
assert detections.frames[0].metadata["scene"] == "start"
assert detections.metadata.to_dict() == {
"coordinate_mode": "normalized_0_1",
"media_metadata": {"codec": "test-codec"},
"request_id": "abc",
"vendor": {"latency_ms": 12},
}
normalized = json.loads(normalized_json)
assert normalized["schema"] == "comfyui-vlm/spatial"
assert normalized["detections"] == detections.to_dict()
assert normalized["points"] == points.to_dict()
def test_pixel_aliases_xywh_flat_segmentation_and_multiple_points():
response = json.dumps(
{
"objects": [
{
"name": "panel",
"box": {"x": -2, "y": 5, "width": 15, "height": 30},
"segmentation": [0, 5, 13, 5, 13, 35, 0, 35],
"points": [[2, 7], [12, 30]],
"score": 1,
}
]
}
)
detections, points, _json = parse_spatial_response(
response,
width=10,
height=20,
coordinate_mode="pixel",
frame_count=1,
)
detection = detections.all_detections()[0]
assert detection.bbox_xyxy == (0.0, 5.0, 10.0, 20.0)
assert detection.polygon == (
(0.0, 5.0),
(10.0, 5.0),
(10.0, 20.0),
(0.0, 20.0),
)
assert [(point.x, point.y) for point in points.points] == [
(2.0, 7.0),
(10.0, 20.0),
]
def test_direct_multi_point_record_keeps_shared_fields_without_duplicates():
detections, points, _json = parse_spatial_response(
'{"label":"hand","confidence":0.8,"points":[[1,2],[3,4]],'
'"metadata":{"side":"left"}}',
width=10,
height=10,
)
assert detections.all_detections() == ()
assert [(point.x, point.y) for point in points.points] == [
(1.0, 2.0),
(3.0, 4.0),
]
assert {point.label for point in points.points} == {"hand"}
assert {point.score for point in points.points} == {0.8}
assert points.points[0].metadata["side"] == "left"
assert detections.frames[0].metadata.to_dict() == {}
def test_top_level_record_batch_groups_video_frame_indices_and_timestamps():
response = json.dumps(
[
{"frame_index": 2, "bbox_xyxy": [1, 2, 3, 4], "label": "late"},
{"frame_index": 0, "point": [5, 6], "label": "early"},
{"frame_index": 0, "box": [0, 0, 4, 5], "label": "first"},
]
)
detections, points, _json = parse_spatial_response(
response,
width=20,
height=10,
fps=4,
coordinate_mode="pixel",
)
assert [frame.frame_index for frame in detections.frames] == [0, 2]
assert detections.frames[1].timestamp == 0.5
assert detections.frame_count == 3
assert [item.label for item in detections.all_detections()] == [
"first",
"late",
]
assert points.points[0].frame_index == 0
def test_normalized_1000_polygon_without_box_derives_clipped_bbox():
detections, points, _json = parse_spatial_response(
'{"polygon":[[-100,100],[500,100],[1200,900]],"label":"shape"}',
width=300,
height=200,
coordinate_mode="normalized_0_1000",
)
detection = detections.all_detections()[0]
assert detection.bbox_xyxy == (0.0, 20.0, 300.0, 180.0)
assert not points.points
def test_empty_payloads_are_valid_and_predictable():
for response in ("", "{}", "[]", '{"frames":[]}'):
detections, points, normalized_json = parse_spatial_response(
response,
width=640,
height=480,
coordinate_mode="pixel",
frame_count=5,
fps=25,
source="empty-test",
)
assert detections.frames == ()
assert detections.frame_count == 5
assert detections.fps == 25
assert points.points == ()
assert points.frame_count == 5
normalized = json.loads(normalized_json)
assert normalized["detections"]["frames"] == []
assert normalized["points"]["points"] == []
@pytest.mark.parametrize(
("response", "message"),
[
('{"bbox":[4,4,2,5]}', "x2 >= x1"),
('{"polygon":[[1,1],[2,2]]}', "at least three"),
('{"quad":[[0,0],[1,0],[1,1]]}', "exactly 4"),
('{"point":[1]}', "two coordinates"),
('{"score":2,"point":[1,1]}', "between 0 and 1"),
(
'{"bbox":[0,0,1,1],"polygon":[[0,0],[1,0],[0,1]],'
'"quad":[[0,0],[1,0],[1,1],[0,1]]}',
"either polygon",
),
(
'{"bbox":[0,0,1,1],"coordinate_mode":"normalized_0_1"}',
"parser is set",
),
('{"label":"no geometry"}', "requires bbox"),
],
)
def test_strict_validation(response, message):
with pytest.raises((TypeError, ValueError), match=message):
parse_spatial_response(
response,
width=10,
height=10,
coordinate_mode="pixel",
)
def test_dimensions_timing_and_alias_conflicts_are_rejected():
cases = [
('{"media":{"width":20}}', {"width": 10}, "does not match"),
('{"media":{"fps":30}}', {"fps": 24}, "does not match"),
(
'{"label":"a","class":"b","point":[1,1]}',
{},
"Conflicting aliases",
),
(
'{"frames":[{"frame_index":0},{"frame_index":0}]}',
{},
"Duplicate frame_index",
),
(
'{"detections":[],"objects":[]}',
{},
"multiple detection collection",
),
]
for response, overrides, message in cases:
arguments = {
"width": 10,
"height": 10,
"coordinate_mode": "pixel",
**overrides,
}
with pytest.raises(ValueError, match=message):
parse_spatial_response(response, **arguments)
def test_node_methods_return_direct_canonical_payloads():
prompt = VLMSpatialPromptBuilder().build(
"Locate the subject.",
"pixel",
100,
80,
1,
0,
)[0]
detections, points, normalized = VLMStructuredSpatialParser().parse(
'{"bbox_xyxy":[1,2,30,40],"point":[5,6]}',
"pixel",
100,
80,
)
assert isinstance(prompt, str)
assert isinstance(detections, DetectionSequence)
assert isinstance(points, PointSequence)
assert json.loads(normalized)["version"] == 1
+234
View File
@@ -0,0 +1,234 @@
from __future__ import annotations
import torch
from ComfyUI_VLM_nodes.nodes.tracking import (
VLMByteTracker,
associate_detection_sequence,
)
from ComfyUI_VLM_nodes.nodes.vision_types import (
Detection,
DetectionSequence,
FrameDetections,
)
def _frame(
frame_index,
detections,
*,
width=100,
height=100,
fps=10.0,
):
timestamp = frame_index / fps
records = tuple(
Detection(
bbox_xyxy=record["box"],
label=record.get("label"),
score=record.get("score"),
frame_index=frame_index,
timestamp=timestamp,
mask=record.get("mask"),
)
for record in detections
)
return FrameDetections(
frame_index=frame_index,
timestamp=timestamp,
width=width,
height=height,
detections=records,
)
def _sequence(frames, *, frame_count=None, fps=10.0, width=100, height=100):
return DetectionSequence(
width=width,
height=height,
frames=tuple(frames),
frame_count=frame_count or 0,
fps=fps,
source="test-detector",
)
def test_low_confidence_second_stage_keeps_the_track_id():
sequence = _sequence(
(
_frame(
0,
[{"box": (10, 10, 30, 30), "score": 0.9, "label": "cat"}],
),
_frame(
1,
[{"box": (11, 10, 31, 30), "score": 0.25, "label": "cat"}],
),
)
)
tracks = associate_detection_sequence(
sequence,
high_threshold=0.6,
low_threshold=0.1,
min_hits=1,
)
assert len(tracks.tracks) == 1
track = tracks.tracks[0]
assert track.track_id == 0
assert [item.track_id for item in track.detections] == [0, 0]
assert track.detections[1].metadata["association_stage"] == "low"
assert track.metadata["state"] == "active"
def test_low_confidence_detection_cannot_start_a_track():
sequence = _sequence(
(
_frame(
0,
[{"box": (10, 10, 30, 30), "score": 0.2, "label": "cat"}],
),
)
)
tracks = associate_detection_sequence(
sequence,
high_threshold=0.6,
low_threshold=0.1,
min_hits=1,
)
assert tracks.tracks == ()
def test_label_aware_matching_prevents_cross_class_identity_reuse():
frames = (
_frame(
0,
[{"box": (10, 10, 30, 30), "score": 0.9, "label": "cat"}],
),
_frame(
1,
[{"box": (10, 10, 30, 30), "score": 0.9, "label": "dog"}],
),
)
label_aware = associate_detection_sequence(
_sequence(frames),
min_hits=1,
label_aware=True,
emit_predictions=False,
)
class_agnostic = associate_detection_sequence(
_sequence(frames),
min_hits=1,
label_aware=False,
emit_predictions=False,
)
assert len(label_aware.tracks) == 2
assert [track.label for track in label_aware.tracks] == ["cat", "dog"]
assert len(class_agnostic.tracks) == 1
assert [
detection.frame_index for detection in class_agnostic.tracks[0].detections
] == [0, 1]
def test_max_age_seconds_uses_fps_and_marks_removed_deterministically():
sequence = _sequence(
(
_frame(
0,
[{"box": (10, 10, 30, 30), "score": 0.9}],
fps=2.0,
),
),
frame_count=4,
fps=2.0,
)
tracks = associate_detection_sequence(
sequence,
min_hits=1,
max_age_seconds=1.0,
emit_predictions=True,
)
track = tracks.tracks[0]
assert [item.frame_index for item in track.detections] == [0, 1, 2]
assert [item.metadata["observation"] for item in track.detections] == [
"detected",
"predicted",
"predicted",
]
assert track.metadata["state"] == "removed"
assert track.metadata["removed_frame"] == 3
assert track.metadata["last_observed_frame"] == 0
def test_mask_iou_can_rescue_a_zero_bbox_iou_match():
mask = torch.zeros(100, 100)
mask[40:60, 40:60] = 1
sequence = _sequence(
(
_frame(
0,
[
{
"box": (0, 0, 10, 10),
"score": 0.9,
"mask": mask,
}
],
),
_frame(
1,
[
{
"box": (20, 0, 30, 10),
"score": 0.9,
"mask": mask,
}
],
),
)
)
tracks = VLMByteTracker(
min_hits=1,
motion_gate=1.0e12,
).track(sequence)
assert len(tracks.tracks) == 1
assert len(tracks.tracks[0].detections) == 2
def test_hungarian_results_and_serialization_are_repeatable():
sequence = _sequence(
(
_frame(
0,
[
{"box": (5, 5, 20, 20), "score": 0.9},
{"box": (40, 5, 55, 20), "score": 0.9},
],
),
_frame(
1,
[
{"box": (41, 5, 56, 20), "score": 0.9},
{"box": (6, 5, 21, 20), "score": 0.9},
],
),
)
)
options = {"min_hits": 1, "emit_predictions": False}
first = associate_detection_sequence(sequence, **options)
second = associate_detection_sequence(sequence, **options)
assert first.to_json() == second.to_json()
assert [
[detection.bbox_xyxy for detection in track.detections]
for track in first.tracks
] == [
[(5.0, 5.0, 20.0, 20.0), (6.0, 5.0, 21.0, 20.0)],
[(40.0, 5.0, 55.0, 20.0), (41.0, 5.0, 56.0, 20.0)],
]
def test_tracker_rejects_invalid_threshold_order():
try:
VLMByteTracker(high_threshold=0.2, low_threshold=0.3)
except ValueError as error:
assert "low_threshold" in str(error)
else:
raise AssertionError("Expected invalid threshold order to fail.")
+265
View File
@@ -0,0 +1,265 @@
import importlib
import json
from dataclasses import FrozenInstanceError
from pathlib import Path
import pytest
import torch
PACKAGE = Path(__file__).resolve().parents[1].name
vision_types = importlib.import_module(f"{PACKAGE}.nodes.vision_types")
DETECTIONS_SCHEMA = vision_types.DETECTIONS_SCHEMA
EVENTS_SCHEMA = vision_types.EVENTS_SCHEMA
POINTS_SCHEMA = vision_types.POINTS_SCHEMA
SCHEMA_VERSION = vision_types.SCHEMA_VERSION
TRACKS_SCHEMA = vision_types.TRACKS_SCHEMA
VLM_DETECTIONS = vision_types.VLM_DETECTIONS
VLM_EVENTS = vision_types.VLM_EVENTS
VLM_POINTS = vision_types.VLM_POINTS
VLM_TRACKS = vision_types.VLM_TRACKS
Detection = vision_types.Detection
DetectionSequence = vision_types.DetectionSequence
EventSequence = vision_types.EventSequence
FrameDetections = vision_types.FrameDetections
FrozenDict = vision_types.FrozenDict
PointSequence = vision_types.PointSequence
TemporalEvent = vision_types.TemporalEvent
Track = vision_types.Track
TrackSequence = vision_types.TrackSequence
VisionPoint = vision_types.VisionPoint
def sample_sequence() -> DetectionSequence:
mask = torch.zeros((24, 32), dtype=torch.float32)
mask[3:12, 4:18] = 1
first = Detection(
bbox_xyxy=(4, 3, 18, 12),
label="cat",
text="sleeping cat",
score=0.875,
polygon=((4, 3), (18, 3), (18, 12), (4, 12)),
quad=((4, 3), (18, 3), (18, 12), (4, 12)),
frame_index=0,
timestamp=0.0,
track_id=7,
source="unit-test",
metadata={"attributes": ["small", "red"], "visible": True},
mask=mask,
)
second = Detection(
bbox_xyxy=(6.5, 5.0, 20.25, 15.0),
label="cat",
score=0.8,
frame_index=1,
timestamp=0.5,
track_id=7,
)
return DetectionSequence(
width=32,
height=24,
frame_count=2,
fps=2.0,
source="synthetic",
metadata={"nested": {"value": 3}},
frames=(
FrameDetections(
frame_index=0,
timestamp=0.0,
width=32,
height=24,
detections=(first,),
),
FrameDetections(
frame_index=1,
timestamp=0.5,
width=32,
height=24,
detections=(second,),
),
),
)
def test_public_socket_and_schema_names_are_stable():
assert VLM_DETECTIONS == "VLM_DETECTIONS"
assert VLM_TRACKS == "VLM_TRACKS"
assert VLM_POINTS == "VLM_POINTS"
assert VLM_EVENTS == "VLM_EVENTS"
assert SCHEMA_VERSION == 1
assert DETECTIONS_SCHEMA == "comfyui-vlm/detections"
assert TRACKS_SCHEMA == "comfyui-vlm/tracks"
assert POINTS_SCHEMA == "comfyui-vlm/points"
assert EVENTS_SCHEMA == "comfyui-vlm/events"
def test_detection_payload_is_validated_immutable_and_mask_safe():
original = torch.ones((4, 5))
detection = Detection(
bbox_xyxy=(0, 0, 5, 4),
label="object",
score=1.0,
metadata={"items": [1, {"ready": True}]},
mask=original,
)
original.zero_()
assert detection.mask.sum().item() == 20
assert isinstance(detection.metadata, FrozenDict)
assert detection.metadata["items"][1]["ready"] is True
with pytest.raises(FrozenInstanceError):
detection.label = "changed"
with pytest.raises(TypeError):
detection.metadata["new"] = "value"
with pytest.raises(TypeError, match="Metadata"):
Detection(
bbox_xyxy=(0, 0, 1, 1),
metadata={"tensor": torch.ones(1)},
)
record = detection.to_dict()
assert "mask" not in record
assert "mask" not in json.dumps(record)
@pytest.mark.parametrize(
("kwargs", "error"),
[
({"bbox_xyxy": (2, 0, 1, 2)}, "x2"),
({"bbox_xyxy": (-1, 0, 1, 2)}, "non-negative"),
({"bbox_xyxy": (0, 0, 1, 2), "score": 1.1}, "between 0 and 1"),
({"bbox_xyxy": (0, 0, 1, 2), "frame_index": -1}, "frame_index"),
({"bbox_xyxy": (0, 0, 1, 2), "timestamp": -0.1}, "timestamp"),
(
{"bbox_xyxy": (0, 0, 1, 2), "quad": ((0, 0), (1, 0), (1, 1))},
"exactly 4",
),
(
{"bbox_xyxy": (0, 0, 1, 2), "mask": torch.ones(1, 2, 3)},
"shape",
),
],
)
def test_detection_rejects_invalid_values(kwargs, error):
with pytest.raises((TypeError, ValueError), match=error):
Detection(**kwargs)
def test_detection_sequence_json_round_trip_is_versioned_and_tensor_free():
sequence = sample_sequence()
encoded = sequence.to_json(indent=2)
decoded_json = json.loads(encoded)
assert decoded_json["schema"] == DETECTIONS_SCHEMA
assert decoded_json["version"] == SCHEMA_VERSION
assert decoded_json["media"] == {
"fps": 2.0,
"frame_count": 2,
"height": 24,
"width": 32,
}
assert "mask" not in encoded
restored = DetectionSequence.from_json(encoded)
assert restored.to_dict() == sequence.to_dict()
assert restored.frames[0].detections[0].mask is None
assert restored.all_detections()[1].center == pytest.approx((13.375, 10.0))
assert restored.frame(99) is None
decoded_json["version"] = 99
with pytest.raises(ValueError, match="Unsupported"):
DetectionSequence.from_dict(decoded_json)
decoded_json["version"] = SCHEMA_VERSION
decoded_json["schema"] = "other"
with pytest.raises(ValueError, match="Expected schema"):
DetectionSequence.from_dict(decoded_json)
def test_frame_and_sequence_enforce_dimensions_order_and_timestamps():
detection = Detection(
bbox_xyxy=(0, 0, 11, 5),
frame_index=0,
timestamp=0,
)
with pytest.raises(ValueError, match="width"):
FrameDetections(
frame_index=0,
timestamp=0,
width=10,
height=10,
detections=(detection,),
)
mismatch = Detection(
bbox_xyxy=(0, 0, 1, 1),
frame_index=1,
timestamp=0,
)
with pytest.raises(ValueError, match="frame_index"):
FrameDetections(
frame_index=0,
timestamp=0,
width=10,
height=10,
detections=(mismatch,),
)
frame_one = FrameDetections(1, 1.0, 10, 10)
frame_zero = FrameDetections(0, 0.0, 10, 10)
with pytest.raises(ValueError, match="increasing"):
DetectionSequence(10, 10, frames=(frame_one, frame_zero))
def test_point_track_and_event_schemas_round_trip_without_tensors():
sequence = sample_sequence()
points = PointSequence(
width=32,
height=24,
frame_count=2,
fps=2,
points=(
VisionPoint(
x=11,
y=7.5,
label="cat",
score=0.8,
frame_index=0,
track_id=7,
),
),
)
assert PointSequence.from_json(points.to_json()).to_dict() == points.to_dict()
track = Track(
track_id=7,
label="cat",
score=0.8,
detections=sequence.all_detections(),
)
tracks = TrackSequence(
width=32,
height=24,
frame_count=2,
fps=2,
tracks=(track,),
)
restored_tracks = TrackSequence.from_json(tracks.to_json())
assert restored_tracks.to_dict() == tracks.to_dict()
assert "mask" not in tracks.to_json()
events = EventSequence(
duration=2.0,
events=(
TemporalEvent(
start_time=0.25,
end_time=1.5,
label="movement",
text="the cat moves",
score=0.9,
),
),
)
assert EventSequence.from_json(events.to_json()).to_dict() == events.to_dict()
with pytest.raises(ValueError, match="duration"):
EventSequence(
duration=1.0,
events=(TemporalEvent(0.0, 2.0),),
)
+432
View File
@@ -0,0 +1,432 @@
import importlib
import inspect
import json
from pathlib import Path
import pytest
import torch
PACKAGE = Path(__file__).resolve().parents[1].name
vision_types = importlib.import_module(f"{PACKAGE}.nodes.vision_types")
vision_utils = importlib.import_module(f"{PACKAGE}.nodes.vision_utils")
VLM_DETECTIONS = vision_types.VLM_DETECTIONS
VLM_POINTS = vision_types.VLM_POINTS
Detection = vision_types.Detection
DetectionSequence = vision_types.DetectionSequence
FrameDetections = vision_types.FrameDetections
BOUNDING_BOXES = vision_utils.BOUNDING_BOXES
NODE_CLASS_MAPPINGS = vision_utils.NODE_CLASS_MAPPINGS
VLMCropDetections = vision_utils.VLMCropDetections
VLMDetectionsFromJSON = vision_utils.VLMDetectionsFromJSON
VLMDetectionsToBoundingBoxes = vision_utils.VLMDetectionsToBoundingBoxes
VLMDetectionsToJSON = vision_utils.VLMDetectionsToJSON
VLMDetectionsToMasks = vision_utils.VLMDetectionsToMasks
VLMDetectionsToPoints = vision_utils.VLMDetectionsToPoints
VLMFilterDetections = vision_utils.VLMFilterDetections
VLMMaskComposite = vision_utils.VLMMaskComposite
VLMMaskProcessor = vision_utils.VLMMaskProcessor
VLMRenderDetections = vision_utils.VLMRenderDetections
VLMSelectDetection = vision_utils.VLMSelectDetection
bounding_boxes_payload = vision_utils.bounding_boxes_payload
composite_with_mask = vision_utils.composite_with_mask
crop_detections = vision_utils.crop_detections
detection_centers = vision_utils.detection_centers
filter_detection_sequence = vision_utils.filter_detection_sequence
instance_map_images = vision_utils.instance_map_images
masks_to_images = vision_utils.masks_to_images
process_masks = vision_utils.process_masks
render_detections = vision_utils.render_detections
select_detection_sequence = vision_utils.select_detection_sequence
sequence_masks = vision_utils.sequence_masks
def sample_sequence() -> DetectionSequence:
cat_mask = torch.zeros((16, 20))
cat_mask[2:8, 3:10] = 1
return DetectionSequence(
width=20,
height=16,
frame_count=2,
fps=2,
frames=(
FrameDetections(
frame_index=0,
timestamp=0,
width=20,
height=16,
detections=(
Detection(
bbox_xyxy=(3.2, 2.1, 10.0, 8.0),
label="cat",
score=0.9,
frame_index=0,
timestamp=0,
track_id=5,
mask=cat_mask,
),
Detection(
bbox_xyxy=(12, 4, 18, 13),
label="dog",
score=0.55,
frame_index=0,
timestamp=0,
),
),
),
FrameDetections(
frame_index=1,
timestamp=0.5,
width=20,
height=16,
detections=(
Detection(
bbox_xyxy=(5, 3, 12, 10),
label="cat",
score=0.8,
frame_index=1,
timestamp=0.5,
track_id=5,
),
),
),
),
)
def empty_sequence() -> DetectionSequence:
return DetectionSequence(
width=20,
height=16,
frame_count=2,
frames=(
FrameDetections(0, 0, 20, 16),
FrameDetections(1, 1, 20, 16),
),
)
def test_filter_and_select_by_all_supported_fields():
sequence = sample_sequence()
filtered = filter_detection_sequence(
sequence,
label="CAT",
label_mode="exact",
minimum_score=0.85,
minimum_area=30,
maximum_area=50,
track_id=5,
frame_index=0,
)
assert len(filtered.frames) == 1
assert [item.label for item in filtered.all_detections()] == ["cat"]
contains = filter_detection_sequence(
sequence,
label="o",
minimum_score=0.5,
)
assert [item.label for item in contains.all_detections()] == ["dog"]
selected = select_detection_sequence(sequence, 1)
assert [item.label for item in selected.all_detections()] == ["dog"]
missing = select_detection_sequence(sequence, 100)
assert missing.all_detections() == ()
assert len(missing.frames) == 2
with pytest.raises(ValueError, match="maximum_area"):
filter_detection_sequence(
sequence,
minimum_area=10,
maximum_area=5,
)
def test_core_bounding_boxes_are_integer_xywh_with_metadata():
payload = bounding_boxes_payload(sample_sequence())
assert payload[0] == {
"x": 3,
"y": 2,
"width": 7,
"height": 6,
"metadata": {
"bbox_xyxy": [3.2, 2.1, 10.0, 8.0],
"frame_index": 0,
"timestamp": 0.0,
"label": "cat",
"score": 0.9,
"track_id": 5,
},
}
assert payload[2]["x"] == 5
assert payload[2]["metadata"]["frame_index"] == 1
assert bounding_boxes_payload(empty_sequence()) == []
def test_centers_and_masks_preserve_frame_mapping_and_empty_shapes():
sequence = sample_sequence()
points = detection_centers(sequence)
assert points.points[0].x == pytest.approx(6.6)
assert points.points[0].y == pytest.approx(5.05)
assert points.points[0].track_id == 5
assert json.loads(points.to_json())["schema"] == "comfyui-vlm/points"
unions, individuals, mapping = sequence_masks(sequence)
assert unions.shape == (2, 16, 20)
assert individuals.shape == (3, 16, 20)
assert unions[0].sum() > 0
assert mapping[0] == {
"mask_index": 0,
"frame_index": 0,
"detection_index": 0,
"label": "cat",
"track_id": 5,
}
empty_unions, empty_individuals, empty_mapping = sequence_masks(empty_sequence())
assert empty_unions.shape == (2, 16, 20)
assert empty_unions.sum() == 0
assert empty_individuals.shape == (0, 16, 20)
assert empty_mapping == []
def test_creator_mask_outputs_are_binary_previewable_and_instance_colored():
sequence = sample_sequence()
unions, individuals, _mapping = sequence_masks(sequence)
union_images = masks_to_images(unions)
individual_images = masks_to_images(individuals)
instance_maps = instance_map_images(sequence)
assert set(torch.unique(unions).tolist()) <= {0.0, 1.0}
assert set(torch.unique(individuals).tolist()) <= {0.0, 1.0}
assert union_images.shape == (2, 16, 20, 3)
assert individual_images.shape == (3, 16, 20, 3)
assert torch.equal(union_images[..., 0], unions)
assert torch.equal(union_images[..., 0], union_images[..., 2])
assert instance_maps.shape == (2, 16, 20, 3)
assert instance_maps.sum() > 0
assert torch.equal(instance_maps, instance_map_images(sequence))
def test_mask_processing_grow_shrink_feather_and_inverse_are_batch_safe():
mask = torch.zeros((2, 9, 9))
mask[:, 4, 4] = 1
grown, grown_binary, grown_inverse = process_masks(
mask,
threshold=0.5,
grow_shrink=1,
feather_radius=0,
)
assert grown.shape == mask.shape
assert grown[0].sum() == 9
assert torch.equal(grown, grown_binary)
assert torch.allclose(grown + grown_inverse, torch.ones_like(grown))
soft, binary, inverse = process_masks(
mask,
threshold=0.5,
grow_shrink=2,
feather_radius=2,
)
assert binary[0].sum() == 25
assert torch.any((soft > 0) & (soft < 1))
assert torch.allclose(soft + inverse, torch.ones_like(soft), atol=1e-6)
full = torch.ones((1, 9, 9))
shrunk, _, _ = process_masks(full, grow_shrink=-1)
assert shrunk[0, 0].sum() == 0
assert shrunk[0, -1].sum() == 0
with pytest.raises(ValueError, match="threshold"):
process_masks(mask, threshold=2)
def test_mask_composite_splits_foreground_and_broadcasts_video_batches():
image = torch.zeros((2, 4, 5, 3))
image[..., 0] = 1
mask = torch.zeros((1, 4, 5))
mask[:, :, :2] = 1
replacement = torch.zeros((1, 4, 5, 3))
replacement[..., 2] = 1
composite, foreground, background_only, mask_image = composite_with_mask(
image,
mask,
background=replacement,
)
assert composite.shape == image.shape
assert foreground.shape == image.shape
assert background_only.shape == image.shape
assert mask_image.shape == image.shape
assert torch.all(composite[:, :, :2, 0] == 1)
assert torch.all(composite[:, :, 2:, 2] == 1)
assert foreground[:, :, 2:].sum() == 0
assert background_only[:, :, :2].sum() == 0
solid, *_ = composite_with_mask(
image[:1],
mask,
background_color="#0f0",
)
assert torch.all(solid[:, :, 2:, 1] == 1)
with pytest.raises(ValueError, match="dimensions"):
composite_with_mask(image, torch.zeros((1, 3, 3)))
def test_rendering_is_deterministic_batch_safe_and_empty_safe():
image = torch.zeros((2, 16, 20, 3), dtype=torch.float32)
sequence = sample_sequence()
first = render_detections(
image,
sequence,
draw_masks=True,
draw_labels=True,
)
second = render_detections(
image,
sequence,
draw_masks=True,
draw_labels=True,
)
assert first.shape == image.shape
assert torch.equal(first, second)
assert first.sum() > 0
unchanged = render_detections(image, empty_sequence())
assert torch.equal(unchanged, image)
with pytest.raises(ValueError, match="dimensions"):
render_detections(
torch.zeros((2, 10, 10, 3)),
sequence,
)
def test_padded_square_crops_form_a_non_distorted_image_batch():
image = torch.zeros((2, 16, 20, 3), dtype=torch.float32)
image[0, 2:8, 3:10, 0] = 1
image[0, 4:13, 12:18, 1] = 1
image[1, 3:10, 5:12, 2] = 1
crops, metadata = crop_detections(
image,
sample_sequence(),
padding=1,
square=True,
)
assert crops.shape[0] == 3
assert crops.ndim == 4
assert crops.shape[1] == max(record["valid_height"] for record in metadata)
assert crops.shape[2] == max(record["valid_width"] for record in metadata)
assert all(record["batch_width"] == crops.shape[2] for record in metadata)
assert metadata[0]["track_id"] == 5
valid_crop = crops[
0,
: metadata[0]["valid_height"],
: metadata[0]["valid_width"],
]
assert valid_crop.sum() > 0
empty_crops, empty_metadata = crop_detections(image, empty_sequence())
assert empty_crops.shape == (0, 1, 1, 3)
assert empty_metadata == []
def test_utility_node_contracts_and_json_round_trip():
sequence = sample_sequence()
encoded = VLMDetectionsToJSON().serialize(sequence, pretty=False)[0]
restored = VLMDetectionsFromJSON().parse(encoded)[0]
assert restored.to_dict() == sequence.to_dict()
filtered = VLMFilterDetections().filter(
sequence,
label="cat",
label_mode="exact",
minimum_score=0,
minimum_area=0,
maximum_area=0,
track_id=-1,
frame_index=-1,
)[0]
assert len(filtered.all_detections()) == 2
assert len(VLMSelectDetection().select(sequence, 0)[0].all_detections()) == 1
boxes, boxes_json = VLMDetectionsToBoundingBoxes().convert(sequence)
assert json.loads(boxes_json) == boxes
points, points_json = VLMDetectionsToPoints().convert(sequence)
assert json.loads(points_json) == points.to_dict()
(
union,
individual,
mask_json,
inverse,
union_images,
individual_images,
instance_maps,
) = VLMDetectionsToMasks().convert(sequence)
assert union.shape[0] == 2
assert individual.shape[0] == 3
assert len(json.loads(mask_json)) == 3
assert torch.allclose(union + inverse, torch.ones_like(union))
assert union_images.shape == (2, 16, 20, 3)
assert individual_images.shape == (3, 16, 20, 3)
assert instance_maps.shape == (2, 16, 20, 3)
processed = VLMMaskProcessor().process(union, 0.5, 1, 1)
assert [value.shape[0] for value in processed] == [2, 2, 2, 2]
composited = VLMMaskComposite().composite(
torch.ones((2, 16, 20, 3)),
union,
"#000000",
)
assert all(value.shape == (2, 16, 20, 3) for value in composited)
image = torch.zeros((2, 16, 20, 3))
assert (
VLMRenderDetections()
.render(
image,
sequence,
True,
False,
0.25,
2,
)[0]
.shape
== image.shape
)
crops, crop_json = VLMCropDetections().crop(image, sequence, 0, False)
assert crops.shape[0] == 3
assert len(json.loads(crop_json)) == 3
def test_every_utility_node_accepts_its_declared_inputs():
expected = {
"VLMDetectionsFromJSON",
"VLMDetectionsToJSON",
"VLMFilterDetections",
"VLMSelectDetection",
"VLMDetectionsToBoundingBoxes",
"VLMDetectionsToPoints",
"VLMDetectionsToMasks",
"VLMMaskProcessor",
"VLMMaskComposite",
"VLMRenderDetections",
"VLMCropDetections",
}
assert set(NODE_CLASS_MAPPINGS) == expected
assert VLMDetectionsFromJSON.RETURN_TYPES == (VLM_DETECTIONS,)
assert VLMDetectionsToBoundingBoxes.RETURN_TYPES[0] == BOUNDING_BOXES
assert VLMDetectionsToPoints.RETURN_TYPES[0] == VLM_POINTS
for node_class in NODE_CLASS_MAPPINGS.values():
schema = node_class.INPUT_TYPES()
declared = {
name
for group in schema.values()
if isinstance(group, dict)
for name in group
}
function = getattr(node_class, node_class.FUNCTION)
accepted = set(inspect.signature(function).parameters)
assert declared <= accepted, (
node_class.__name__,
declared - accepted,
)
+126 -10
View File
@@ -1,16 +1,66 @@
import { app } from "../../../scripts/app.js";
import { api } from "../../../scripts/api.js";
const OUTPUT_NAME = "output_text";
const VIEW_TEXT_NODE = "ViewText";
const MODERN_VLM_NODE = "ModernVLM";
function ensureOutputWidget(node) {
let widget = node.widgets?.find((item) => item.name === OUTPUT_NAME);
if (!widget) {
const container = document.createElement("div");
const header = document.createElement("div");
const status = document.createElement("span");
const copy = document.createElement("button");
const output = document.createElement("textarea");
status.textContent = "Ready";
copy.textContent = "Copy";
copy.type = "button";
copy.title = "Copy the complete VLM response";
copy.addEventListener("click", async () => {
const previous = copy.textContent;
try {
await navigator.clipboard.writeText(output.value);
copy.textContent = "Copied";
} catch {
copy.textContent = "Copy failed";
}
window.setTimeout(() => {
copy.textContent = previous;
}, 1200);
});
output.readOnly = true;
output.setAttribute("aria-label", "VLM text output");
Object.assign(output.style, {
header.append(status, copy);
container.append(header, output);
Object.assign(container.style, {
display: "flex",
flexDirection: "column",
width: "100%",
height: "100%",
minHeight: "150px",
gap: "6px",
});
Object.assign(header.style, {
display: "flex",
alignItems: "center",
justifyContent: "space-between",
color: "var(--descrip-text, #aaa)",
fontSize: "12px",
});
Object.assign(copy.style, {
color: "var(--input-text, #ddd)",
background: "var(--comfy-input-bg, #202020)",
border: "1px solid var(--border-color, #555)",
borderRadius: "5px",
padding: "3px 9px",
cursor: "pointer",
});
Object.assign(output.style, {
width: "100%",
flex: "1",
minHeight: "120px",
resize: "vertical",
boxSizing: "border-box",
@@ -19,21 +69,85 @@ function ensureOutputWidget(node) {
border: "1px solid var(--border-color, #555)",
borderRadius: "6px",
padding: "8px",
lineHeight: "1.45",
whiteSpace: "pre-wrap",
});
widget = node.addDOMWidget(OUTPUT_NAME, "STRING", output, {
widget = node.addDOMWidget(OUTPUT_NAME, "STRING", container, {
serialize: false,
hideOnZoom: false,
});
widget.serialize = false;
widget.inputEl = output;
widget.statusEl = status;
}
return widget;
}
function setOutput(node, text, state = "Complete") {
const widget = ensureOutputWidget(node);
const value = Array.isArray(text) ? text.join("\n\n") : String(text ?? "");
widget.value = value;
widget.inputEl.value = value;
if (widget.statusEl) {
widget.statusEl.textContent = state;
}
node.setDirtyCanvas?.(true, true);
}
function findNode(graph, id) {
if (!graph || id == null) {
return null;
}
return graph.getNodeById?.(id)
?? graph.getNodeById?.(String(id))
?? graph.getNodeById?.(Number(id))
?? null;
}
function connectedViewTextNodes(source) {
if (!source?.graph) {
return [];
}
const found = new Set();
for (const output of source.outputs ?? []) {
for (const linkId of output.links ?? []) {
const link = source.graph.links?.get?.(linkId)
?? source.graph._links?.get?.(linkId);
const target = findNode(source.graph, link?.target_id);
if (target?.type === VIEW_TEXT_NODE) {
found.add(target);
}
}
}
return [...found];
}
function updateFromProgress({ nodeId, text }) {
const source = findNode(app.rootGraph ?? app.graph, nodeId);
if (!source) {
return;
}
if (source.type === VIEW_TEXT_NODE) {
setOutput(source, text, "Streaming…");
return;
}
if (source.type !== MODERN_VLM_NODE) {
return;
}
for (const target of connectedViewTextNodes(source)) {
setOutput(target, text, "Streaming…");
}
}
app.registerExtension({
name: "gokayfem.vlm.view-text",
async setup() {
api.addEventListener("progress_text", ({ detail }) => {
updateFromProgress(detail ?? {});
});
},
async beforeRegisterNodeDef(nodeType, nodeData) {
if (nodeData.name !== "ViewText") {
if (nodeData.name !== VIEW_TEXT_NODE) {
return;
}
const onCreated = nodeType.prototype.onNodeCreated;
@@ -45,14 +159,16 @@ app.registerExtension({
const onExecuted = nodeType.prototype.onExecuted;
nodeType.prototype.onExecuted = function (message) {
const result = onExecuted?.apply(this, arguments);
const values = Array.isArray(message?.text)
? message.text
: [message?.text ?? ""];
const widget = ensureOutputWidget(this);
widget.value = values.join("");
widget.inputEl.value = widget.value;
this.setDirtyCanvas?.(true, true);
setOutput(this, message?.text, "Complete");
return result;
};
},
onNodeOutputsUpdated(nodeOutputs) {
for (const [nodeId, output] of Object.entries(nodeOutputs ?? {})) {
const node = findNode(app.rootGraph ?? app.graph, nodeId);
if (node?.type === VIEW_TEXT_NODE && output?.text != null) {
setOutput(node, output.text, "Complete");
}
}
},
});