Compare commits
4
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
bb558b1495 | ||
|
|
c8882b5226 | ||
|
|
239c9045ab | ||
|
|
0da5070039 |
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 = {}
|
||||
|
||||
@@ -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
@@ -2,7 +2,11 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from numbers import Real
|
||||
|
||||
import torch
|
||||
from PIL import Image, ImageDraw
|
||||
@@ -24,24 +28,173 @@ from .runtime import (
|
||||
|
||||
MODELS = {
|
||||
"Florence-2 base FT (fast)": "florence-community/Florence-2-base-ft",
|
||||
"Florence-2 large FT (recommended)": (
|
||||
"florence-community/Florence-2-large-ft"
|
||||
),
|
||||
"Florence-2 large FT (recommended)": ("florence-community/Florence-2-large-ft"),
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class FlorenceTaskSpec:
|
||||
"""Declarative contract for one official Florence-2 task."""
|
||||
|
||||
token: str
|
||||
input_kind: str
|
||||
output_kind: str
|
||||
|
||||
|
||||
TASKS = {
|
||||
"Caption": "<CAPTION>",
|
||||
"Detailed caption": "<DETAILED_CAPTION>",
|
||||
"More detailed caption": "<MORE_DETAILED_CAPTION>",
|
||||
"OCR": "<OCR>",
|
||||
"OCR with regions": "<OCR_WITH_REGION>",
|
||||
"Object detection": "<OD>",
|
||||
"Dense region caption": "<DENSE_REGION_CAPTION>",
|
||||
"Region proposals": "<REGION_PROPOSAL>",
|
||||
"Referring expression segmentation": "<REFERRING_EXPRESSION_SEGMENTATION>",
|
||||
"Open vocabulary detection": "<OPEN_VOCABULARY_DETECTION>",
|
||||
"Caption": FlorenceTaskSpec("<CAPTION>", "none", "text"),
|
||||
"Detailed caption": FlorenceTaskSpec("<DETAILED_CAPTION>", "none", "text"),
|
||||
"More detailed caption": FlorenceTaskSpec(
|
||||
"<MORE_DETAILED_CAPTION>", "none", "text"
|
||||
),
|
||||
"OCR": FlorenceTaskSpec("<OCR>", "none", "text"),
|
||||
"OCR with regions": FlorenceTaskSpec("<OCR_WITH_REGION>", "none", "quad_boxes"),
|
||||
"Object detection": FlorenceTaskSpec("<OD>", "none", "boxes"),
|
||||
"Dense region caption": FlorenceTaskSpec("<DENSE_REGION_CAPTION>", "none", "boxes"),
|
||||
"Caption to phrase grounding": FlorenceTaskSpec(
|
||||
"<CAPTION_TO_PHRASE_GROUNDING>", "text", "boxes"
|
||||
),
|
||||
"Referring expression segmentation": FlorenceTaskSpec(
|
||||
"<REFERRING_EXPRESSION_SEGMENTATION>", "text", "polygons"
|
||||
),
|
||||
"Region to segmentation": FlorenceTaskSpec(
|
||||
"<REGION_TO_SEGMENTATION>", "region", "polygons"
|
||||
),
|
||||
"Open vocabulary detection": FlorenceTaskSpec(
|
||||
"<OPEN_VOCABULARY_DETECTION>", "text", "mixed"
|
||||
),
|
||||
"Region to category": FlorenceTaskSpec("<REGION_TO_CATEGORY>", "region", "text"),
|
||||
"Region to description": FlorenceTaskSpec(
|
||||
"<REGION_TO_DESCRIPTION>", "region", "text"
|
||||
),
|
||||
"Region to OCR": FlorenceTaskSpec("<REGION_TO_OCR>", "region", "text"),
|
||||
"Region proposals": FlorenceTaskSpec("<REGION_PROPOSAL>", "none", "boxes"),
|
||||
}
|
||||
|
||||
|
||||
def _clean_decoded_text(value):
|
||||
"""Remove generation wrappers without discarding Florence location tokens."""
|
||||
|
||||
text = str(value)
|
||||
for token in ("<s>", "</s>", "<pad>"):
|
||||
text = text.replace(token, "")
|
||||
return text.strip()
|
||||
|
||||
|
||||
def _select_region(region, image_index, batch_size):
|
||||
"""Select one core BOUNDING_BOX for the current image.
|
||||
|
||||
Core primitive boxes are dictionaries. Detection nodes may emit either a
|
||||
flat per-image list or a nested batch list, so both common shapes are
|
||||
accepted while ambiguous multi-region inputs fail explicitly.
|
||||
"""
|
||||
|
||||
if region is None or isinstance(region, dict):
|
||||
return region
|
||||
if not isinstance(region, (list, tuple)):
|
||||
raise TypeError("region must be a core BOUNDING_BOX dictionary.")
|
||||
if not region:
|
||||
return None
|
||||
|
||||
if all(isinstance(item, dict) for item in region):
|
||||
if len(region) == 1:
|
||||
return region[0]
|
||||
if len(region) == batch_size:
|
||||
return region[image_index]
|
||||
raise ValueError("Region tasks require exactly one BOUNDING_BOX per image.")
|
||||
|
||||
if len(region) != batch_size:
|
||||
raise ValueError("Batched BOUNDING_BOX input must contain one entry per image.")
|
||||
frame_regions = region[image_index]
|
||||
if isinstance(frame_regions, dict):
|
||||
return frame_regions
|
||||
if not isinstance(frame_regions, (list, tuple)) or len(frame_regions) != 1:
|
||||
raise ValueError(
|
||||
"Region tasks require exactly one BOUNDING_BOX per image; "
|
||||
"select a detection before connecting it."
|
||||
)
|
||||
if not isinstance(frame_regions[0], dict):
|
||||
raise TypeError("Each BOUNDING_BOX entry must be a dictionary.")
|
||||
return frame_regions[0]
|
||||
|
||||
|
||||
def _encode_region(region, image_size):
|
||||
"""Encode an absolute-pixel core BOUNDING_BOX as Florence location tokens."""
|
||||
|
||||
if not isinstance(region, dict):
|
||||
raise TypeError("region must be a core BOUNDING_BOX dictionary.")
|
||||
|
||||
try:
|
||||
x = float(region["x"])
|
||||
y = float(region["y"])
|
||||
box_width = float(region["width"])
|
||||
box_height = float(region["height"])
|
||||
except KeyError as exc:
|
||||
raise ValueError("region must contain x, y, width, and height.") from exc
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise ValueError("region coordinates must be numeric.") from exc
|
||||
|
||||
values = (x, y, box_width, box_height)
|
||||
if not all(math.isfinite(value) for value in values):
|
||||
raise ValueError("region coordinates must be finite.")
|
||||
if box_width <= 0 or box_height <= 0:
|
||||
raise ValueError("region width and height must be greater than zero.")
|
||||
|
||||
image_width, image_height = image_size
|
||||
if image_width <= 0 or image_height <= 0:
|
||||
raise ValueError("image dimensions must be greater than zero.")
|
||||
|
||||
x0 = max(0.0, min(float(image_width), x))
|
||||
y0 = max(0.0, min(float(image_height), y))
|
||||
x1 = max(0.0, min(float(image_width), x + box_width))
|
||||
y1 = max(0.0, min(float(image_height), y + box_height))
|
||||
if x1 <= x0 or y1 <= y0:
|
||||
raise ValueError("region does not overlap the input image.")
|
||||
|
||||
coordinates = (
|
||||
x0 / image_width,
|
||||
y0 / image_height,
|
||||
x1 / image_width,
|
||||
y1 / image_height,
|
||||
)
|
||||
bins = [
|
||||
max(0, min(999, math.floor(coordinate * 1000))) for coordinate in coordinates
|
||||
]
|
||||
return "".join(f"<loc_{value}>" for value in bins)
|
||||
|
||||
|
||||
def _task_extra_input(task_name, text_input, region, image_size):
|
||||
"""Validate and prepare the optional suffix for a Florence task prompt."""
|
||||
|
||||
try:
|
||||
spec = TASKS[task_name]
|
||||
except KeyError as exc:
|
||||
raise ValueError(f"Unsupported Florence-2 task: {task_name}") from exc
|
||||
|
||||
text = (text_input or "").strip()
|
||||
if spec.input_kind == "none":
|
||||
if text:
|
||||
raise ValueError(f"{task_name} does not accept text input.")
|
||||
if region is not None:
|
||||
raise ValueError(f"{task_name} does not accept a region input.")
|
||||
return ""
|
||||
if spec.input_kind == "text":
|
||||
if not text:
|
||||
raise ValueError(f"{task_name} requires text input.")
|
||||
if region is not None:
|
||||
raise ValueError(f"{task_name} does not accept a region input.")
|
||||
return text
|
||||
if spec.input_kind == "region":
|
||||
if text:
|
||||
raise ValueError(
|
||||
f"{task_name} uses the region input and does not accept text."
|
||||
)
|
||||
if region is None:
|
||||
raise ValueError(f"{task_name} requires a connected BOUNDING_BOX region.")
|
||||
return _encode_region(region, image_size)
|
||||
raise RuntimeError(f"Unknown Florence task input kind: {spec.input_kind}")
|
||||
|
||||
|
||||
class FlorencePredictor:
|
||||
def __init__(self, model_label):
|
||||
transformers = require_module("transformers")
|
||||
@@ -66,9 +219,7 @@ class FlorencePredictor:
|
||||
|
||||
def run(self, image, task_token, text, max_new_tokens, beams):
|
||||
prompt = task_token + (text.strip() if text.strip() else "")
|
||||
inputs = self.processor(
|
||||
text=prompt, images=image, return_tensors="pt"
|
||||
)
|
||||
inputs = self.processor(text=prompt, images=image, return_tensors="pt")
|
||||
model = self.handle.ensure_loaded()
|
||||
device = model_device(model)
|
||||
inputs = move_inputs(inputs, device, floating_dtype=self.dtype)
|
||||
@@ -80,9 +231,7 @@ class FlorencePredictor:
|
||||
do_sample=False,
|
||||
early_stopping=int(beams) > 1,
|
||||
)
|
||||
raw = self.processor.batch_decode(
|
||||
generated, skip_special_tokens=False
|
||||
)[0]
|
||||
raw = self.processor.batch_decode(generated, skip_special_tokens=False)[0]
|
||||
parsed = self.processor.post_process_generation(
|
||||
raw, task=task_token, image_size=image.size
|
||||
)
|
||||
@@ -95,39 +244,156 @@ def _json_default(value):
|
||||
return str(value)
|
||||
|
||||
|
||||
_SPATIAL_KEYS = frozenset(
|
||||
{
|
||||
"bboxes",
|
||||
"quad_boxes",
|
||||
"polygons",
|
||||
"labels",
|
||||
"bboxes_labels",
|
||||
"polygons_labels",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _spatial_result(parsed):
|
||||
if not isinstance(parsed, dict):
|
||||
return {}
|
||||
if _SPATIAL_KEYS.intersection(parsed):
|
||||
return parsed
|
||||
result = next(iter(parsed.values()), {})
|
||||
return result if isinstance(result, dict) else {}
|
||||
|
||||
|
||||
def _stable_color(kind, index, label):
|
||||
key = f"{kind}:{index}:{label}".encode("utf-8", errors="replace")
|
||||
digest = hashlib.blake2b(key, digest_size=3).digest()
|
||||
return tuple(64 + channel % 192 for channel in digest)
|
||||
|
||||
|
||||
def _points(values, image_size):
|
||||
if not isinstance(values, (list, tuple)) or len(values) < 6:
|
||||
return []
|
||||
width, height = image_size
|
||||
points = []
|
||||
for index in range(0, len(values) - 1, 2):
|
||||
x, y = values[index], values[index + 1]
|
||||
if not isinstance(x, Real) or not isinstance(y, Real):
|
||||
return []
|
||||
if not math.isfinite(float(x)) or not math.isfinite(float(y)):
|
||||
return []
|
||||
points.append(
|
||||
(
|
||||
max(0, min(width - 1, round(float(x)))),
|
||||
max(0, min(height - 1, round(float(y)))),
|
||||
)
|
||||
)
|
||||
return points
|
||||
|
||||
|
||||
def _box(values, image_size):
|
||||
if not isinstance(values, (list, tuple)) or len(values) < 4:
|
||||
return None
|
||||
if not all(isinstance(value, Real) for value in values[:4]):
|
||||
return None
|
||||
coordinates = [float(value) for value in values[:4]]
|
||||
if not all(math.isfinite(value) for value in coordinates):
|
||||
return None
|
||||
x0, y0, x1, y1 = coordinates
|
||||
x0, x1 = sorted((x0, x1))
|
||||
y0, y1 = sorted((y0, y1))
|
||||
width, height = image_size
|
||||
x0 = max(0, min(width - 1, round(x0)))
|
||||
x1 = max(0, min(width - 1, round(x1)))
|
||||
y0 = max(0, min(height - 1, round(y0)))
|
||||
y1 = max(0, min(height - 1, round(y1)))
|
||||
if x1 <= x0 or y1 <= y0:
|
||||
return None
|
||||
return x0, y0, x1, y1
|
||||
|
||||
|
||||
def _polygon_list(group):
|
||||
if not isinstance(group, (list, tuple)) or not group:
|
||||
return []
|
||||
if isinstance(group[0], Real):
|
||||
return [group]
|
||||
return [item for item in group if isinstance(item, (list, tuple))]
|
||||
|
||||
|
||||
def _label_with_score(labels, scores, index):
|
||||
label = str(labels[index]) if index < len(labels) else ""
|
||||
if index < len(scores) and isinstance(scores[index], Real):
|
||||
score = f"{float(scores[index]):.3f}"
|
||||
return f"{label} {score}".strip()
|
||||
return label
|
||||
|
||||
|
||||
def _draw_label(draw, position, text, color, image_size):
|
||||
if not text:
|
||||
return
|
||||
x, y = position
|
||||
try:
|
||||
left, top, right, bottom = draw.textbbox((0, 0), text)
|
||||
text_width, text_height = right - left, bottom - top
|
||||
except AttributeError:
|
||||
text_width, text_height = draw.textlength(text), 11
|
||||
width, height = image_size
|
||||
x = max(0, min(width - text_width - 4, x))
|
||||
y = max(0, min(height - text_height - 4, y))
|
||||
background = (0, 0, 0) if sum(color) > 360 else (255, 255, 255)
|
||||
foreground = (255, 255, 255) if background == (0, 0, 0) else (0, 0, 0)
|
||||
draw.rectangle(
|
||||
(x, y, x + text_width + 4, y + text_height + 4),
|
||||
fill=background,
|
||||
)
|
||||
draw.text((x + 2, y + 2), text, fill=foreground)
|
||||
|
||||
|
||||
def _visualize(image, parsed):
|
||||
result = next(iter(parsed.values()), parsed) if isinstance(parsed, dict) else {}
|
||||
result = _spatial_result(parsed)
|
||||
mask = Image.new("L", image.size, 0)
|
||||
visual = image.copy().convert("RGB")
|
||||
mask_draw = ImageDraw.Draw(mask)
|
||||
draw = ImageDraw.Draw(visual)
|
||||
labels = result.get("labels", []) if isinstance(result, dict) else []
|
||||
width = max(2, min(8, round(min(image.size) / 256 * 3)))
|
||||
labels = result.get("labels", [])
|
||||
scores = result.get("scores", [])
|
||||
|
||||
for index, box in enumerate(result.get("bboxes", [])):
|
||||
box = [float(value) for value in box]
|
||||
draw.rectangle(box, outline="#00ff88", width=3)
|
||||
if index < len(labels):
|
||||
draw.text((box[0] + 3, box[1] + 3), str(labels[index]), fill="#00ff88")
|
||||
box_labels = result.get("bboxes_labels", labels)
|
||||
for index, values in enumerate(result.get("bboxes", [])):
|
||||
box = _box(values, image.size)
|
||||
if box is None:
|
||||
continue
|
||||
label = _label_with_score(box_labels, scores, index)
|
||||
color = _stable_color("box", index, label)
|
||||
mask_draw.rectangle(box, fill=255)
|
||||
draw.rectangle(box, outline=color, width=width)
|
||||
_draw_label(draw, (box[0], box[1]), label, color, image.size)
|
||||
|
||||
for quad in result.get("quad_boxes", []):
|
||||
points = [
|
||||
(float(quad[index]), float(quad[index + 1]))
|
||||
for index in range(0, len(quad), 2)
|
||||
]
|
||||
draw.line(points + [points[0]], fill="#00c8ff", width=3)
|
||||
for index, values in enumerate(result.get("quad_boxes", [])):
|
||||
points = _points(values, image.size)
|
||||
if len(points) < 3:
|
||||
continue
|
||||
label = _label_with_score(labels, scores, index)
|
||||
color = _stable_color("quad", index, label)
|
||||
mask_draw.polygon(points, fill=255)
|
||||
draw.line(points + [points[0]], fill=color, width=width)
|
||||
_draw_label(draw, points[0], label, color, image.size)
|
||||
|
||||
polygons = result.get("polygons", [])
|
||||
for group in polygons:
|
||||
# Florence may return either one flat polygon or a list of polygons.
|
||||
groups = [group] if group and isinstance(group[0], (int, float)) else group
|
||||
for polygon in groups:
|
||||
points = [
|
||||
(float(polygon[index]), float(polygon[index + 1]))
|
||||
for index in range(0, len(polygon), 2)
|
||||
]
|
||||
if len(points) >= 3:
|
||||
mask_draw.polygon(points, fill=255)
|
||||
draw.line(points + [points[0]], fill="#ff4da6", width=3)
|
||||
polygon_labels = result.get("polygons_labels", labels)
|
||||
for index, group in enumerate(result.get("polygons", [])):
|
||||
label = _label_with_score(polygon_labels, scores, index)
|
||||
color = _stable_color("polygon", index, label)
|
||||
label_drawn = False
|
||||
for polygon in _polygon_list(group):
|
||||
points = _points(polygon, image.size)
|
||||
if len(points) < 3:
|
||||
continue
|
||||
mask_draw.polygon(points, fill=255)
|
||||
draw.line(points + [points[0]], fill=color, width=width)
|
||||
if not label_drawn:
|
||||
_draw_label(draw, points[0], label, color, image.size)
|
||||
label_drawn = True
|
||||
return mask, visual
|
||||
|
||||
|
||||
@@ -143,7 +409,10 @@ class Florence2(CachedModelNode):
|
||||
{
|
||||
"default": "",
|
||||
"multiline": True,
|
||||
"tooltip": "Required for referring-expression and open-vocabulary tasks.",
|
||||
"tooltip": (
|
||||
"Required only for phrase grounding, referring-expression "
|
||||
"segmentation, and open-vocabulary detection."
|
||||
),
|
||||
},
|
||||
),
|
||||
"model": (
|
||||
@@ -158,6 +427,15 @@ class Florence2(CachedModelNode):
|
||||
},
|
||||
"optional": {
|
||||
"unload_after": ("BOOLEAN", {"default": False}),
|
||||
"region": (
|
||||
"BOUNDING_BOX",
|
||||
{
|
||||
"tooltip": (
|
||||
"Core bounding box input required by Region to "
|
||||
"Segmentation/Category/Description/OCR."
|
||||
)
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -175,28 +453,52 @@ class Florence2(CachedModelNode):
|
||||
max_new_tokens,
|
||||
beams,
|
||||
unload_after=False,
|
||||
region=None,
|
||||
):
|
||||
predictor = self.get_or_create_model(
|
||||
model, lambda: FlorencePredictor(model)
|
||||
)
|
||||
images = tensor_batch_to_pil(image)
|
||||
if not images:
|
||||
raise ValueError("Florence-2 requires at least one input image.")
|
||||
try:
|
||||
spec = TASKS[task]
|
||||
except KeyError as exc:
|
||||
raise ValueError(f"Unsupported Florence-2 task: {task}") from exc
|
||||
|
||||
extra_inputs = []
|
||||
for index, pil_image in enumerate(images):
|
||||
selected_region = _select_region(region, index, len(images))
|
||||
extra_inputs.append(
|
||||
_task_extra_input(
|
||||
task,
|
||||
text_input,
|
||||
selected_region,
|
||||
pil_image.size,
|
||||
)
|
||||
)
|
||||
|
||||
predictor = self.get_or_create_model(model, lambda: FlorencePredictor(model))
|
||||
texts, records, masks, visuals = [], [], [], []
|
||||
try:
|
||||
for pil_image in tensor_batch_to_pil(image):
|
||||
for pil_image, extra_input in zip(images, extra_inputs):
|
||||
raw, parsed = predictor.run(
|
||||
pil_image,
|
||||
TASKS[task],
|
||||
text_input,
|
||||
spec.token,
|
||||
extra_input,
|
||||
max_new_tokens,
|
||||
beams,
|
||||
)
|
||||
texts.append(raw)
|
||||
texts.append(_clean_decoded_text(raw))
|
||||
records.append(parsed)
|
||||
mask, visual = _visualize(pil_image, parsed)
|
||||
masks.append(pil_mask_to_tensor(mask))
|
||||
visuals.append(pil_to_tensor(visual))
|
||||
return (
|
||||
batch_text(texts),
|
||||
json.dumps(records, ensure_ascii=False, default=_json_default),
|
||||
json.dumps(
|
||||
records,
|
||||
ensure_ascii=False,
|
||||
default=_json_default,
|
||||
sort_keys=True,
|
||||
),
|
||||
torch.cat(masks),
|
||||
torch.cat(visuals),
|
||||
)
|
||||
|
||||
@@ -0,0 +1,413 @@
|
||||
"""Dependency-light geometry, mask, color, and association primitives."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import colorsys
|
||||
import hashlib
|
||||
import math
|
||||
from 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",
|
||||
]
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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",
|
||||
}
|
||||
@@ -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
@@ -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 (
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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
@@ -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,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
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
@@ -0,0 +1,250 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from ComfyUI_VLM_nodes.nodes import florence2
|
||||
from PIL import Image
|
||||
|
||||
EXPECTED_TASKS = {
|
||||
"Caption": ("<CAPTION>", "none"),
|
||||
"Detailed caption": ("<DETAILED_CAPTION>", "none"),
|
||||
"More detailed caption": ("<MORE_DETAILED_CAPTION>", "none"),
|
||||
"OCR": ("<OCR>", "none"),
|
||||
"OCR with regions": ("<OCR_WITH_REGION>", "none"),
|
||||
"Object detection": ("<OD>", "none"),
|
||||
"Dense region caption": ("<DENSE_REGION_CAPTION>", "none"),
|
||||
"Caption to phrase grounding": ("<CAPTION_TO_PHRASE_GROUNDING>", "text"),
|
||||
"Referring expression segmentation": (
|
||||
"<REFERRING_EXPRESSION_SEGMENTATION>",
|
||||
"text",
|
||||
),
|
||||
"Region to segmentation": ("<REGION_TO_SEGMENTATION>", "region"),
|
||||
"Open vocabulary detection": ("<OPEN_VOCABULARY_DETECTION>", "text"),
|
||||
"Region to category": ("<REGION_TO_CATEGORY>", "region"),
|
||||
"Region to description": ("<REGION_TO_DESCRIPTION>", "region"),
|
||||
"Region to OCR": ("<REGION_TO_OCR>", "region"),
|
||||
"Region proposals": ("<REGION_PROPOSAL>", "none"),
|
||||
}
|
||||
|
||||
|
||||
def test_registry_covers_all_official_transformers_tasks():
|
||||
assert len(florence2.TASKS) == 15
|
||||
assert {
|
||||
name: (spec.token, spec.input_kind) for name, spec in florence2.TASKS.items()
|
||||
} == EXPECTED_TASKS
|
||||
assert {spec.output_kind for spec in florence2.TASKS.values()} == {
|
||||
"text",
|
||||
"boxes",
|
||||
"quad_boxes",
|
||||
"polygons",
|
||||
"mixed",
|
||||
}
|
||||
|
||||
|
||||
def test_node_contract_preserves_outputs_and_adds_core_region_input():
|
||||
schema = florence2.Florence2.INPUT_TYPES()
|
||||
assert florence2.NODE_CLASS_MAPPINGS["Florence2"] is florence2.Florence2
|
||||
assert florence2.Florence2.RETURN_TYPES[:4] == (
|
||||
"STRING",
|
||||
"STRING",
|
||||
"MASK",
|
||||
"IMAGE",
|
||||
)
|
||||
assert florence2.Florence2.RETURN_NAMES[:4] == (
|
||||
"text",
|
||||
"structured_json",
|
||||
"mask",
|
||||
"visualization",
|
||||
)
|
||||
assert schema["optional"]["region"][0] == "BOUNDING_BOX"
|
||||
assert "forceInput" not in repr(schema["optional"]["region"])
|
||||
|
||||
|
||||
def test_region_encoding_uses_core_xywh_and_florence_location_bins():
|
||||
region = {"x": 10, "y": 20, "width": 40, "height": 100}
|
||||
assert florence2._encode_region(region, (100, 200)) == (
|
||||
"<loc_100><loc_100><loc_500><loc_600>"
|
||||
)
|
||||
|
||||
clamped = {"x": -10, "y": -20, "width": 200, "height": 300}
|
||||
assert florence2._encode_region(clamped, (100, 200)) == (
|
||||
"<loc_0><loc_0><loc_999><loc_999>"
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="greater than zero"):
|
||||
florence2._encode_region(
|
||||
{"x": 0, "y": 0, "width": 0, "height": 10},
|
||||
(100, 100),
|
||||
)
|
||||
with pytest.raises(ValueError, match="does not overlap"):
|
||||
florence2._encode_region(
|
||||
{"x": 200, "y": 200, "width": 10, "height": 10},
|
||||
(100, 100),
|
||||
)
|
||||
|
||||
|
||||
def test_task_inputs_are_validated_before_inference():
|
||||
image_size = (100, 100)
|
||||
region = {"x": 10, "y": 10, "width": 20, "height": 20}
|
||||
|
||||
assert florence2._task_extra_input("Caption", "", None, image_size) == ""
|
||||
with pytest.raises(ValueError, match="does not accept text"):
|
||||
florence2._task_extra_input("Caption", "unexpected", None, image_size)
|
||||
with pytest.raises(ValueError, match="requires text"):
|
||||
florence2._task_extra_input("Open vocabulary detection", "", None, image_size)
|
||||
assert (
|
||||
florence2._task_extra_input(
|
||||
"Open vocabulary detection", "red car", None, image_size
|
||||
)
|
||||
== "red car"
|
||||
)
|
||||
with pytest.raises(ValueError, match="requires a connected BOUNDING_BOX"):
|
||||
florence2._task_extra_input("Region to OCR", "", None, image_size)
|
||||
assert (
|
||||
florence2._task_extra_input("Region to OCR", "", region, image_size)
|
||||
== "<loc_100><loc_100><loc_300><loc_300>"
|
||||
)
|
||||
with pytest.raises(ValueError, match="does not accept text"):
|
||||
florence2._task_extra_input("Region to OCR", "also text", region, image_size)
|
||||
|
||||
|
||||
def test_region_selection_accepts_core_and_batched_detector_shapes():
|
||||
first = {"x": 1, "y": 2, "width": 3, "height": 4}
|
||||
second = {"x": 5, "y": 6, "width": 7, "height": 8}
|
||||
assert florence2._select_region(first, 0, 2) is first
|
||||
assert florence2._select_region([first, second], 1, 2) is second
|
||||
assert florence2._select_region([[first], [second]], 0, 2) is first
|
||||
|
||||
with pytest.raises(ValueError, match="exactly one"):
|
||||
florence2._select_region([[first, second]], 0, 1)
|
||||
|
||||
|
||||
def test_visualization_is_deterministic_and_masks_every_spatial_shape():
|
||||
image = Image.new("RGB", (48, 36), "black")
|
||||
parsed = {
|
||||
"<OPEN_VOCABULARY_DETECTION>": {
|
||||
"bboxes": [[1, 1, 10, 10]],
|
||||
"bboxes_labels": ["box"],
|
||||
"quad_boxes": [[14, 1, 22, 1, 22, 10, 14, 10]],
|
||||
"labels": ["ocr"],
|
||||
"polygons": [[[26, 1, 40, 1, 40, 12, 26, 12]]],
|
||||
"polygons_labels": ["polygon"],
|
||||
}
|
||||
}
|
||||
|
||||
mask_a, visual_a = florence2._visualize(image, parsed)
|
||||
mask_b, visual_b = florence2._visualize(image, parsed)
|
||||
mask = np.asarray(mask_a)
|
||||
|
||||
assert mask[5, 5] == 255
|
||||
assert mask[5, 18] == 255
|
||||
assert mask[5, 30] == 255
|
||||
assert mask_a.tobytes() == mask_b.tobytes()
|
||||
assert visual_a.tobytes() == visual_b.tobytes()
|
||||
|
||||
|
||||
def test_predictor_generation_is_deterministic_without_downloads():
|
||||
calls = {}
|
||||
|
||||
class FakeProcessor:
|
||||
def __call__(self, text, images, return_tensors):
|
||||
calls["prompt"] = text
|
||||
assert images.size == (8, 8)
|
||||
assert return_tensors == "pt"
|
||||
return {
|
||||
"input_ids": torch.tensor([[1]], dtype=torch.long),
|
||||
"pixel_values": torch.zeros((1, 3, 8, 8)),
|
||||
}
|
||||
|
||||
def batch_decode(self, generated, skip_special_tokens):
|
||||
assert skip_special_tokens is False
|
||||
return ["<s>answer</s>"]
|
||||
|
||||
def post_process_generation(self, raw, task, image_size):
|
||||
return {task: "answer"}
|
||||
|
||||
class FakeModel(torch.nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.anchor = torch.nn.Parameter(torch.zeros(()))
|
||||
|
||||
def generate(self, **kwargs):
|
||||
calls["generation"] = kwargs
|
||||
return torch.tensor([[2]], dtype=torch.long)
|
||||
|
||||
class FakeHandle:
|
||||
def __init__(self):
|
||||
self.model = FakeModel()
|
||||
|
||||
def ensure_loaded(self):
|
||||
return self.model
|
||||
|
||||
predictor = object.__new__(florence2.FlorencePredictor)
|
||||
predictor.dtype = torch.float32
|
||||
predictor.processor = FakeProcessor()
|
||||
predictor.handle = FakeHandle()
|
||||
|
||||
raw, parsed = predictor.run(
|
||||
Image.new("RGB", (8, 8)),
|
||||
"<CAPTION>",
|
||||
"",
|
||||
32,
|
||||
3,
|
||||
)
|
||||
assert raw == "<s>answer</s>"
|
||||
assert parsed == {"<CAPTION>": "answer"}
|
||||
assert calls["prompt"] == "<CAPTION>"
|
||||
assert calls["generation"]["do_sample"] is False
|
||||
assert calls["generation"]["num_beams"] == 3
|
||||
assert calls["generation"]["early_stopping"] is True
|
||||
|
||||
|
||||
def test_node_cleans_text_preserves_structured_data_and_unloads_target():
|
||||
parsed = {
|
||||
"<OD>": {
|
||||
"bboxes": [[2, 2, 12, 12]],
|
||||
"labels": ["person"],
|
||||
"quad_boxes": [[14, 2, 22, 2, 22, 12, 14, 12]],
|
||||
"polygons": [[[24, 2, 30, 2, 30, 12, 24, 12]]],
|
||||
}
|
||||
}
|
||||
calls = []
|
||||
|
||||
class FakePredictor:
|
||||
def run(
|
||||
self,
|
||||
image,
|
||||
task_token,
|
||||
extra_input,
|
||||
max_new_tokens,
|
||||
beams,
|
||||
):
|
||||
calls.append((image.size, task_token, extra_input, max_new_tokens, beams))
|
||||
return "<s>person<loc_1><loc_2></s><pad>", parsed
|
||||
|
||||
node = florence2.Florence2()
|
||||
node.get_or_create_model = lambda key, factory: FakePredictor()
|
||||
unloads = []
|
||||
node.maybe_clear_model = unloads.append
|
||||
|
||||
output = node.run(
|
||||
torch.zeros((1, 32, 32, 3)),
|
||||
"Object detection",
|
||||
"",
|
||||
"Florence-2 base FT (fast)",
|
||||
64,
|
||||
1,
|
||||
unload_after=True,
|
||||
)
|
||||
|
||||
assert len(output) == 4
|
||||
assert output[0] == "person<loc_1><loc_2>"
|
||||
assert json.loads(output[1]) == [parsed]
|
||||
assert output[2].shape == (1, 32, 32)
|
||||
assert output[2].max().item() == 1.0
|
||||
assert output[3].shape == (1, 32, 32, 3)
|
||||
assert calls == [((32, 32), "<OD>", "", 64, 1)]
|
||||
assert unloads == [True]
|
||||
@@ -0,0 +1,166 @@
|
||||
import importlib
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
PACKAGE = Path(__file__).resolve().parents[1].name
|
||||
geometry = importlib.import_module(f"{PACKAGE}.nodes.geometry")
|
||||
vision_types = importlib.import_module(f"{PACKAGE}.nodes.vision_types")
|
||||
|
||||
associate_detections = geometry.associate_detections
|
||||
bbox_from_mask = geometry.bbox_from_mask
|
||||
bbox_iou = geometry.bbox_iou
|
||||
box_area = geometry.box_area
|
||||
box_center = geometry.box_center
|
||||
box_to_mask = geometry.box_to_mask
|
||||
clip_box = geometry.clip_box
|
||||
clip_polygon = geometry.clip_polygon
|
||||
denormalize_box = geometry.denormalize_box
|
||||
detection_to_mask = geometry.detection_to_mask
|
||||
deterministic_color = geometry.deterministic_color
|
||||
expand_box = geometry.expand_box
|
||||
individual_detection_masks = geometry.individual_detection_masks
|
||||
mask_iou = geometry.mask_iou
|
||||
normalize_box = geometry.normalize_box
|
||||
polygon_area = geometry.polygon_area
|
||||
polygon_to_mask = geometry.polygon_to_mask
|
||||
quad_to_mask = geometry.quad_to_mask
|
||||
translate_box = geometry.translate_box
|
||||
union_detection_mask = geometry.union_detection_mask
|
||||
Detection = vision_types.Detection
|
||||
|
||||
|
||||
def test_box_clipping_normalization_area_and_center():
|
||||
assert clip_box((0, 2, 25, 22), 20, 10) == (0, 2, 20, 10)
|
||||
normalized = normalize_box((5, 2, 15, 8), 20, 10)
|
||||
assert normalized == pytest.approx((0.25, 0.2, 0.75, 0.8))
|
||||
assert denormalize_box(normalized, 20, 10) == pytest.approx((5, 2, 15, 8))
|
||||
assert box_area((5, 2, 15, 8)) == 60
|
||||
assert box_center((5, 2, 15, 8)) == (10, 5)
|
||||
with pytest.raises(ValueError, match="Normalized"):
|
||||
denormalize_box((0, 0, 2, 1), 20, 10)
|
||||
with pytest.raises(ValueError, match="x2"):
|
||||
clip_box((2, 0, 1, 1), 20, 10)
|
||||
|
||||
|
||||
def test_polygon_clipping_and_area():
|
||||
polygon = ((-2, -3), (8, 0), (8, 5), (0, 5))
|
||||
assert clip_polygon(polygon, 6, 4) == (
|
||||
(0, 0),
|
||||
(6, 0),
|
||||
(6, 4),
|
||||
(0, 4),
|
||||
)
|
||||
assert polygon_area(((0, 0), (5, 0), (5, 4), (0, 4))) == 20
|
||||
assert polygon_area(((0, 0), (0, 4), (5, 4), (5, 0))) == 20
|
||||
with pytest.raises(ValueError, match="at least three"):
|
||||
polygon_area(((0, 0), (1, 1)))
|
||||
|
||||
|
||||
def test_bbox_and_mask_iou():
|
||||
assert bbox_iou((0, 0, 10, 10), (5, 0, 15, 10)) == pytest.approx(1 / 3)
|
||||
assert bbox_iou((0, 0, 1, 1), (2, 2, 3, 3)) == 0
|
||||
first = torch.zeros((4, 4))
|
||||
second = torch.zeros((4, 4))
|
||||
first[:2, :2] = 1
|
||||
second[1:3, :2] = 1
|
||||
assert mask_iou(first, second) == pytest.approx(1 / 3)
|
||||
assert mask_iou(torch.zeros((2, 2)), torch.zeros((2, 2))) == 0
|
||||
with pytest.raises(ValueError, match="same shape"):
|
||||
mask_iou(torch.zeros((2, 2)), torch.zeros((3, 2)))
|
||||
|
||||
|
||||
def test_box_polygon_and_quad_rasterization():
|
||||
box = box_to_mask((1.2, 2.1, 4.1, 5.2), 8, 7)
|
||||
assert box.shape == (7, 8)
|
||||
assert box.sum().item() == 16
|
||||
polygon = polygon_to_mask(((1, 1), (5, 1), (5, 5), (1, 5)), 8, 8)
|
||||
quad = quad_to_mask(((1, 1), (5, 1), (5, 5), (1, 5)), 8, 8)
|
||||
assert torch.equal(polygon, quad)
|
||||
assert polygon.sum() > 0
|
||||
with pytest.raises(ValueError, match="exactly four"):
|
||||
quad_to_mask(((0, 0), (1, 0), (1, 1)), 4, 4)
|
||||
|
||||
|
||||
def test_detection_mask_priority_union_individual_and_bbox():
|
||||
explicit = torch.zeros((8, 8))
|
||||
explicit[3:6, 2:5] = 1
|
||||
with_mask = Detection(
|
||||
bbox_xyxy=(0, 0, 8, 8),
|
||||
polygon=((0, 0), (8, 0), (8, 8), (0, 8)),
|
||||
mask=explicit,
|
||||
)
|
||||
polygon_only = Detection(
|
||||
bbox_xyxy=(1, 1, 5, 5),
|
||||
polygon=((1, 1), (5, 1), (5, 5), (1, 5)),
|
||||
)
|
||||
assert torch.equal(detection_to_mask(with_mask, 8, 8), explicit)
|
||||
masks = individual_detection_masks((with_mask, polygon_only), 8, 8)
|
||||
assert masks.shape == (2, 8, 8)
|
||||
union = union_detection_mask((with_mask, polygon_only), 8, 8)
|
||||
assert union.shape == (8, 8)
|
||||
assert torch.all(union >= masks[0])
|
||||
assert individual_detection_masks((), 8, 8).shape == (0, 8, 8)
|
||||
assert union_detection_mask((), 8, 8).sum() == 0
|
||||
assert bbox_from_mask(explicit) == (2, 3, 5, 6)
|
||||
assert bbox_from_mask(torch.zeros((2, 2))) is None
|
||||
|
||||
|
||||
def test_deterministic_color_and_box_expansion():
|
||||
assert deterministic_color("track-1") == deterministic_color("track-1")
|
||||
assert deterministic_color("track-1") != deterministic_color("track-2")
|
||||
assert all(0 <= channel <= 255 for channel in deterministic_color("object"))
|
||||
assert expand_box((4, 4, 8, 6), 12, 12, padding=1) == (3, 3, 9, 7)
|
||||
squared = expand_box((4, 4, 8, 6), 12, 12, square=True)
|
||||
assert squared == (4, 3, 8, 7)
|
||||
assert translate_box((1, 2, 3, 4), 2, 1) == (3, 3, 5, 5)
|
||||
|
||||
|
||||
def test_label_aware_stable_association_with_motion():
|
||||
previous = (
|
||||
Detection(
|
||||
bbox_xyxy=(0, 0, 10, 10),
|
||||
label="cat",
|
||||
track_id=4,
|
||||
),
|
||||
Detection(
|
||||
bbox_xyxy=(20, 0, 30, 10),
|
||||
label="dog",
|
||||
track_id=9,
|
||||
),
|
||||
)
|
||||
current = (
|
||||
Detection(bbox_xyxy=(5, 0, 15, 10), label="cat"),
|
||||
Detection(bbox_xyxy=(20, 0, 30, 10), label="bird"),
|
||||
Detection(bbox_xyxy=(40, 0, 50, 10), label="dog"),
|
||||
)
|
||||
without_motion = associate_detections(
|
||||
previous,
|
||||
current,
|
||||
minimum_iou=0.3,
|
||||
)
|
||||
assert without_motion.matches == ((0, 0, pytest.approx(1 / 3)),)
|
||||
assert without_motion.unmatched_previous == (1,)
|
||||
assert without_motion.unmatched_current == (1, 2)
|
||||
|
||||
with_motion = associate_detections(
|
||||
previous,
|
||||
current,
|
||||
minimum_iou=0.9,
|
||||
motion_by_track={4: (5, 0), 9: (20, 0)},
|
||||
)
|
||||
assert with_motion.matches == (
|
||||
(0, 0, 1.0),
|
||||
(1, 2, 1.0),
|
||||
)
|
||||
assert with_motion.unmatched_previous == ()
|
||||
assert with_motion.unmatched_current == (1,)
|
||||
|
||||
label_agnostic = associate_detections(
|
||||
previous,
|
||||
current,
|
||||
minimum_iou=0.9,
|
||||
label_aware=False,
|
||||
)
|
||||
assert label_agnostic.matches == ((1, 1, 1.0),)
|
||||
@@ -0,0 +1,141 @@
|
||||
from types import SimpleNamespace
|
||||
|
||||
import torch
|
||||
from ComfyUI_VLM_nodes.nodes.grounding import (
|
||||
MODEL_SPECS,
|
||||
VLMOpenVocabularyDetection,
|
||||
core_bounding_box_frames,
|
||||
core_bounding_boxes,
|
||||
detection_box_masks,
|
||||
parse_labels,
|
||||
result_to_detections,
|
||||
)
|
||||
from ComfyUI_VLM_nodes.nodes.vision_types import (
|
||||
DetectionSequence,
|
||||
FrameDetections,
|
||||
)
|
||||
|
||||
|
||||
def test_detector_catalog_is_small_fast_and_portable():
|
||||
assert "Grounding DINO Tiny (fast)" in MODEL_SPECS
|
||||
assert "OmDet Turbo Swin Tiny (fast)" in MODEL_SPECS
|
||||
assert all("/" in spec.model_id for spec in MODEL_SPECS.values())
|
||||
schema = VLMOpenVocabularyDetection.INPUT_TYPES()
|
||||
assert tuple(MODEL_SPECS) == schema["required"]["model"][0]
|
||||
|
||||
|
||||
def test_label_parser_preserves_phrases_and_removes_duplicates():
|
||||
assert parse_labels("red car, person\nsmall dog;person") == [
|
||||
"red car",
|
||||
"person",
|
||||
"small dog",
|
||||
]
|
||||
|
||||
|
||||
def test_transformers_results_are_clipped_sorted_and_normalized():
|
||||
result = {
|
||||
"boxes": torch.tensor([[-5.0, 2.0, 20.0, 12.0], [5.0, 5.0, 9.0, 9.0]]),
|
||||
"scores": torch.tensor([0.25, 0.9]),
|
||||
"text_labels": ["cat", "dog"],
|
||||
}
|
||||
detections = result_to_detections(
|
||||
result,
|
||||
labels=["cat", "dog"],
|
||||
width=16,
|
||||
height=10,
|
||||
frame_index=0,
|
||||
timestamp=0.0,
|
||||
source="test/model",
|
||||
max_detections=20,
|
||||
)
|
||||
assert [item.label for item in detections] == ["dog", "cat"]
|
||||
assert detections[1].bbox_xyxy == (0.0, 2.0, 16.0, 10.0)
|
||||
assert detections[0].metadata["model_id"] == "test/model"
|
||||
|
||||
|
||||
def test_max_detections_is_applied_after_confidence_sorting():
|
||||
detections = result_to_detections(
|
||||
{
|
||||
"boxes": [[0, 0, 1, 1], [1, 1, 2, 2], [2, 2, 3, 3]],
|
||||
"scores": [0.1, 0.9, 0.8],
|
||||
"text_labels": ["low", "best", "second"],
|
||||
},
|
||||
labels=["low", "best", "second"],
|
||||
width=4,
|
||||
height=4,
|
||||
frame_index=0,
|
||||
timestamp=0.0,
|
||||
source="test/model",
|
||||
max_detections=2,
|
||||
)
|
||||
assert [item.label for item in detections] == ["best", "second"]
|
||||
|
||||
|
||||
def test_box_masks_and_core_boxes_keep_geometry_and_metadata():
|
||||
detections = result_to_detections(
|
||||
{
|
||||
"boxes": torch.tensor([[1.0, 2.0, 4.0, 5.0]]),
|
||||
"scores": torch.tensor([0.8]),
|
||||
"labels": torch.tensor([0]),
|
||||
},
|
||||
labels=["cat"],
|
||||
width=8,
|
||||
height=6,
|
||||
frame_index=0,
|
||||
timestamp=0.0,
|
||||
source="test/model",
|
||||
max_detections=5,
|
||||
)
|
||||
sequence = DetectionSequence(
|
||||
width=8,
|
||||
height=6,
|
||||
frames=(
|
||||
FrameDetections(
|
||||
frame_index=0,
|
||||
timestamp=0.0,
|
||||
width=8,
|
||||
height=6,
|
||||
detections=detections,
|
||||
),
|
||||
),
|
||||
frame_count=1,
|
||||
)
|
||||
masks = detection_box_masks(sequence)
|
||||
assert masks.shape == (1, 6, 8)
|
||||
assert masks.sum().item() == 9
|
||||
boxes = core_bounding_boxes(sequence)
|
||||
assert boxes == [
|
||||
{
|
||||
"x": 1,
|
||||
"y": 2,
|
||||
"width": 3,
|
||||
"height": 3,
|
||||
"label": "cat",
|
||||
"score": detections[0].score,
|
||||
"metadata": {
|
||||
"frame_index": 0,
|
||||
"label": "cat",
|
||||
"score": detections[0].score,
|
||||
"source": "test/model",
|
||||
},
|
||||
}
|
||||
]
|
||||
assert core_bounding_box_frames(sequence) == [boxes]
|
||||
|
||||
|
||||
def test_result_label_indices_are_resolved():
|
||||
detections = result_to_detections(
|
||||
{
|
||||
"boxes": [[0, 0, 4, 4]],
|
||||
"scores": [SimpleNamespace(item=lambda: 0.5)],
|
||||
"classes": [SimpleNamespace(item=lambda: 1)],
|
||||
},
|
||||
labels=["cat", "dog"],
|
||||
width=4,
|
||||
height=4,
|
||||
frame_index=0,
|
||||
timestamp=0.0,
|
||||
source="omdet",
|
||||
max_detections=1,
|
||||
)
|
||||
assert detections[0].label == "dog"
|
||||
+288
-27
@@ -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)
|
||||
|
||||
@@ -0,0 +1,251 @@
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from ComfyUI_VLM_nodes.nodes.sam2 import (
|
||||
SAM2_MODELS,
|
||||
Sam2Spec,
|
||||
Sam2VideoPredictor,
|
||||
VLMSAM2VideoSegmentation,
|
||||
_core_box,
|
||||
_normalize_processed_masks,
|
||||
seed_boxes,
|
||||
)
|
||||
from ComfyUI_VLM_nodes.nodes.vision_types import (
|
||||
Detection,
|
||||
DetectionSequence,
|
||||
FrameDetections,
|
||||
)
|
||||
|
||||
|
||||
def test_sam2_catalog_leads_with_tiny_and_has_no_30b_models():
|
||||
assert next(iter(SAM2_MODELS)) == "SAM2.1 Hiera Tiny (fast)"
|
||||
assert all("30b" not in spec.model_id.lower() for spec in SAM2_MODELS.values())
|
||||
assert VLMSAM2VideoSegmentation.RETURN_NAMES[0:2] == ("tracks", "json")
|
||||
|
||||
|
||||
def test_core_box_is_converted_from_xywh():
|
||||
assert _core_box({"x": 2, "y": 3, "width": 5, "height": 7}) == (
|
||||
2.0,
|
||||
3.0,
|
||||
7.0,
|
||||
10.0,
|
||||
)
|
||||
with pytest.raises(ValueError, match="positive"):
|
||||
_core_box({"x": 2, "y": 3, "width": 0, "height": 7})
|
||||
|
||||
|
||||
def test_detection_seeds_keep_labels_and_ids():
|
||||
detection = Detection(
|
||||
bbox_xyxy=(1, 2, 9, 10),
|
||||
label="cat",
|
||||
frame_index=2,
|
||||
timestamp=0.2,
|
||||
track_id=7,
|
||||
)
|
||||
sequence = DetectionSequence(
|
||||
width=10,
|
||||
height=10,
|
||||
frames=(
|
||||
FrameDetections(
|
||||
frame_index=2,
|
||||
timestamp=0.2,
|
||||
width=10,
|
||||
height=10,
|
||||
detections=(detection,),
|
||||
),
|
||||
),
|
||||
frame_count=3,
|
||||
fps=10,
|
||||
)
|
||||
boxes, ids, labels = seed_boxes(
|
||||
width=10,
|
||||
height=10,
|
||||
frame_index=2,
|
||||
detections=sequence,
|
||||
bounding_box=None,
|
||||
)
|
||||
assert boxes == [[1.0, 2.0, 9.0, 10.0]]
|
||||
assert ids == [7]
|
||||
assert labels == {7: "cat"}
|
||||
|
||||
|
||||
def test_processed_mask_shapes_are_normalized():
|
||||
assert _normalize_processed_masks(torch.zeros(2, 1, 4, 5)).shape == (
|
||||
2,
|
||||
4,
|
||||
5,
|
||||
)
|
||||
assert _normalize_processed_masks(torch.zeros(4, 5)).shape == (1, 4, 5)
|
||||
with pytest.raises(RuntimeError, match="unsupported"):
|
||||
_normalize_processed_masks(torch.zeros(1, 2, 3, 4, 5))
|
||||
|
||||
|
||||
def test_video_session_runs_seed_frame_before_both_propagation_directions():
|
||||
class FakeProcessor:
|
||||
def init_video_session(self, **_kwargs):
|
||||
return SimpleNamespace(obj_ids=[1], seed_inferred=False)
|
||||
|
||||
def add_inputs_to_inference_session(self, **kwargs):
|
||||
kwargs["inference_session"].seed_frame = kwargs["frame_idx"]
|
||||
|
||||
def post_process_masks(self, masks, **_kwargs):
|
||||
return masks
|
||||
|
||||
class FakeModel(torch.nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.anchor = torch.nn.Parameter(torch.zeros(()))
|
||||
self.seed_calls = []
|
||||
self.propagation_directions = []
|
||||
|
||||
def forward(self, inference_session, frame_idx):
|
||||
inference_session.seed_inferred = True
|
||||
self.seed_calls.append(frame_idx)
|
||||
return SimpleNamespace(
|
||||
frame_idx=frame_idx,
|
||||
pred_masks=torch.ones(1, 1, 4, 4),
|
||||
)
|
||||
|
||||
def propagate_in_video_iterator(
|
||||
self,
|
||||
inference_session,
|
||||
start_frame_idx,
|
||||
reverse=False,
|
||||
**_kwargs,
|
||||
):
|
||||
assert inference_session.seed_inferred
|
||||
self.propagation_directions.append(reverse)
|
||||
indices = (
|
||||
range(start_frame_idx, 3)
|
||||
if not reverse
|
||||
else range(start_frame_idx, -1, -1)
|
||||
)
|
||||
for frame_index in indices:
|
||||
yield SimpleNamespace(
|
||||
frame_idx=frame_index,
|
||||
pred_masks=torch.ones(1, 1, 4, 4),
|
||||
)
|
||||
|
||||
model = FakeModel()
|
||||
predictor = Sam2VideoPredictor.__new__(Sam2VideoPredictor)
|
||||
predictor.processor = FakeProcessor()
|
||||
predictor.dtype = torch.float32
|
||||
predictor.spec = Sam2Spec("test/sam2", "test-sam2")
|
||||
predictor.handle = SimpleNamespace(ensure_loaded=lambda: model)
|
||||
detections = DetectionSequence(
|
||||
width=4,
|
||||
height=4,
|
||||
frames=(
|
||||
FrameDetections(
|
||||
frame_index=0,
|
||||
timestamp=0.0,
|
||||
width=4,
|
||||
height=4,
|
||||
detections=(
|
||||
Detection(
|
||||
bbox_xyxy=(0, 0, 4, 4),
|
||||
label="object",
|
||||
frame_index=0,
|
||||
timestamp=0.0,
|
||||
),
|
||||
),
|
||||
),
|
||||
),
|
||||
frame_count=1,
|
||||
)
|
||||
tracks, union, individual, preview = predictor.propagate(
|
||||
torch.zeros(3, 4, 4, 3),
|
||||
seed_frame=1,
|
||||
fps=10.0,
|
||||
detections=detections,
|
||||
bounding_box=None,
|
||||
seed_mask=None,
|
||||
mask_threshold=0.0,
|
||||
keep_video_on_cpu=True,
|
||||
mask_output="union_and_objects",
|
||||
render_preview=True,
|
||||
)
|
||||
assert model.seed_calls == [1]
|
||||
assert model.propagation_directions == [False, True]
|
||||
assert [item.frame_index for item in tracks.tracks[0].detections] == [0, 1, 2]
|
||||
assert union.shape == (3, 4, 4)
|
||||
assert individual.shape == (3, 4, 4)
|
||||
assert preview.shape == (3, 4, 4, 3)
|
||||
assert torch.equal(preview[0], preview[1])
|
||||
assert torch.equal(preview[1], preview[2])
|
||||
|
||||
|
||||
def test_multi_object_mask_seeds_are_passed_as_one_mask_per_object():
|
||||
class FakeProcessor:
|
||||
def __init__(self):
|
||||
self.received_masks = None
|
||||
|
||||
def init_video_session(self, **kwargs):
|
||||
assert str(kwargs["processing_device"]) == "cpu"
|
||||
return SimpleNamespace(obj_ids=[1, 2], seed_inferred=False)
|
||||
|
||||
def add_inputs_to_inference_session(self, **kwargs):
|
||||
self.received_masks = kwargs["input_masks"]
|
||||
|
||||
def post_process_masks(self, masks, **_kwargs):
|
||||
return masks
|
||||
|
||||
class FakeModel(torch.nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.anchor = torch.nn.Parameter(torch.zeros(()))
|
||||
|
||||
def forward(self, inference_session, frame_idx):
|
||||
inference_session.seed_inferred = True
|
||||
return SimpleNamespace(
|
||||
frame_idx=frame_idx,
|
||||
pred_masks=torch.ones(2, 1, 4, 4),
|
||||
)
|
||||
|
||||
def propagate_in_video_iterator(self, **_kwargs):
|
||||
return iter(())
|
||||
|
||||
processor = FakeProcessor()
|
||||
predictor = Sam2VideoPredictor.__new__(Sam2VideoPredictor)
|
||||
predictor.processor = processor
|
||||
predictor.dtype = torch.float32
|
||||
predictor.spec = Sam2Spec("test/sam2", "test-sam2")
|
||||
predictor.handle = SimpleNamespace(ensure_loaded=FakeModel)
|
||||
images = torch.zeros(1, 4, 4, 3)
|
||||
tracks, _union, individual, preview = predictor.propagate(
|
||||
images,
|
||||
seed_frame=0,
|
||||
fps=24.0,
|
||||
detections=None,
|
||||
bounding_box=None,
|
||||
seed_mask=torch.ones(2, 4, 4),
|
||||
mask_threshold=0.0,
|
||||
keep_video_on_cpu=True,
|
||||
mask_output="union_only",
|
||||
render_preview=False,
|
||||
)
|
||||
assert isinstance(processor.received_masks, list)
|
||||
assert len(processor.received_masks) == 2
|
||||
assert len(tracks.tracks) == 2
|
||||
assert individual.shape == (0, 4, 4)
|
||||
assert preview.data_ptr() == images.data_ptr()
|
||||
|
||||
|
||||
def test_nested_core_boxes_select_seed_frame_and_keep_top_level_labels():
|
||||
boxes, ids, labels = seed_boxes(
|
||||
width=20,
|
||||
height=20,
|
||||
frame_index=1,
|
||||
detections=None,
|
||||
bounding_box=[
|
||||
[{"x": 0, "y": 0, "width": 2, "height": 2, "label": "old"}],
|
||||
[
|
||||
{"x": 3, "y": 4, "width": 5, "height": 6, "label": "person"},
|
||||
{"x": 10, "y": 11, "width": 4, "height": 3},
|
||||
],
|
||||
],
|
||||
)
|
||||
assert boxes == [[3.0, 4.0, 8.0, 10.0], [10.0, 11.0, 14.0, 14.0]]
|
||||
assert ids == [1, 2]
|
||||
assert labels == {1: "person", 2: None}
|
||||
@@ -0,0 +1,213 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from ComfyUI_VLM_nodes.nodes.sam3_adapter import (
|
||||
VLMTrackReport,
|
||||
iter_sam3_masks,
|
||||
sam3_track_data_to_tracks,
|
||||
track_report_json,
|
||||
track_report_payload,
|
||||
track_report_text,
|
||||
unpack_sam3_mask,
|
||||
validate_sam3_track_data,
|
||||
)
|
||||
from ComfyUI_VLM_nodes.nodes.vision_types import (
|
||||
Detection,
|
||||
DetectionSequence,
|
||||
FrameDetections,
|
||||
)
|
||||
|
||||
|
||||
def _pack_masks(masks):
|
||||
masks = masks.to(torch.uint8)
|
||||
width = masks.shape[-1]
|
||||
assert width % 8 == 0
|
||||
bits = 1 << torch.arange(8, dtype=torch.int64)
|
||||
grouped = masks.reshape(*masks.shape[:-1], width // 8, 8)
|
||||
return (grouped * bits).sum(dim=-1).to(torch.uint8)
|
||||
|
||||
|
||||
def _sample_track_data():
|
||||
masks = torch.zeros(2, 2, 4, 8, dtype=torch.bool)
|
||||
masks[0, 0, 1:3, 2:5] = True
|
||||
masks[1, 0, 1:4, 3:6] = True
|
||||
masks[1, 1, 0:2, 0:2] = True
|
||||
return {
|
||||
"packed_masks": _pack_masks(masks),
|
||||
"n_frames": 2,
|
||||
"scores": [0.9, 0.75],
|
||||
"orig_size": (40, 80),
|
||||
}
|
||||
|
||||
|
||||
def _seed_detections():
|
||||
detections = (
|
||||
Detection(
|
||||
bbox_xyxy=(20, 10, 50, 30),
|
||||
label="cat",
|
||||
text="the cat",
|
||||
score=0.95,
|
||||
frame_index=0,
|
||||
timestamp=0.0,
|
||||
track_id=7,
|
||||
source="seed",
|
||||
),
|
||||
Detection(
|
||||
bbox_xyxy=(0, 0, 20, 20),
|
||||
label="fish",
|
||||
score=0.8,
|
||||
frame_index=0,
|
||||
timestamp=0.0,
|
||||
track_id=9,
|
||||
source="seed",
|
||||
),
|
||||
)
|
||||
return DetectionSequence(
|
||||
width=80,
|
||||
height=40,
|
||||
frames=(
|
||||
FrameDetections(
|
||||
frame_index=0,
|
||||
timestamp=0.0,
|
||||
width=80,
|
||||
height=40,
|
||||
detections=detections,
|
||||
),
|
||||
),
|
||||
frame_count=2,
|
||||
fps=10.0,
|
||||
)
|
||||
|
||||
|
||||
def test_validate_and_stream_unpack_without_expanding_the_video():
|
||||
track_data = _sample_track_data()
|
||||
layout = validate_sam3_track_data(track_data)
|
||||
assert (
|
||||
layout.n_frames,
|
||||
layout.n_objects,
|
||||
layout.mask_height,
|
||||
layout.mask_width,
|
||||
) == (2, 2, 4, 8)
|
||||
yielded = list(iter_sam3_masks(track_data, present_only=True))
|
||||
assert [(frame, obj) for frame, obj, _mask in yielded] == [
|
||||
(0, 0),
|
||||
(1, 0),
|
||||
(1, 1),
|
||||
]
|
||||
assert all(mask.shape == (4, 8) for _frame, _obj, mask in yielded)
|
||||
assert yielded[0][2].sum().item() == 6
|
||||
one = unpack_sam3_mask(track_data["packed_masks"][1, 1])
|
||||
assert one.dtype == torch.bool
|
||||
assert one.sum().item() == 4
|
||||
|
||||
|
||||
def test_adapter_derives_scaled_boxes_and_preserves_seed_identity():
|
||||
tracks = sam3_track_data_to_tracks(
|
||||
_sample_track_data(),
|
||||
seed_detections=_seed_detections(),
|
||||
fps=10.0,
|
||||
)
|
||||
assert (tracks.width, tracks.height, tracks.frame_count, tracks.fps) == (
|
||||
80,
|
||||
40,
|
||||
2,
|
||||
10.0,
|
||||
)
|
||||
assert [track.track_id for track in tracks.tracks] == [7, 9]
|
||||
assert [track.label for track in tracks.tracks] == ["cat", "fish"]
|
||||
cat, fish = tracks.tracks
|
||||
assert cat.detections[0].bbox_xyxy == (20.0, 10.0, 50.0, 30.0)
|
||||
assert cat.detections[1].bbox_xyxy == (30.0, 10.0, 60.0, 40.0)
|
||||
assert fish.detections[0].frame_index == 1
|
||||
assert fish.detections[0].bbox_xyxy == (0.0, 0.0, 20.0, 20.0)
|
||||
assert cat.detections[0].metadata["mask_ref"]["object_index"] == 0
|
||||
assert fish.detections[0].metadata["mask_ref"]["object_index"] == 1
|
||||
assert cat.metadata["seeded"] is True
|
||||
assert fish.metadata["seeded"] is True
|
||||
|
||||
serialized = tracks.to_json()
|
||||
assert "mask_ref" in serialized
|
||||
assert "packed_masks" not in serialized
|
||||
assert "tensor" not in serialized.casefold()
|
||||
|
||||
|
||||
def test_adapter_handles_an_empty_core_result():
|
||||
tracks = sam3_track_data_to_tracks(
|
||||
{
|
||||
"packed_masks": None,
|
||||
"n_frames": 3,
|
||||
"scores": [],
|
||||
"orig_size": (48, 64),
|
||||
},
|
||||
fps=24.0,
|
||||
)
|
||||
assert tracks.tracks == ()
|
||||
assert tracks.frame_count == 3
|
||||
assert tracks.metadata["object_slots"] == 0
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("track_data", "message"),
|
||||
(
|
||||
(
|
||||
{"packed_masks": None, "n_frames": 1, "scores": []},
|
||||
"orig_size",
|
||||
),
|
||||
(
|
||||
{
|
||||
"packed_masks": torch.zeros(1, 1, 2, 1),
|
||||
"n_frames": 1,
|
||||
"scores": [0.5],
|
||||
"orig_size": (2, 8),
|
||||
},
|
||||
"uint8",
|
||||
),
|
||||
(
|
||||
{
|
||||
"packed_masks": torch.zeros(2, 1, 2, 1, dtype=torch.uint8),
|
||||
"n_frames": 1,
|
||||
"scores": [0.5],
|
||||
"orig_size": (2, 8),
|
||||
},
|
||||
"n_frames",
|
||||
),
|
||||
(
|
||||
{
|
||||
"packed_masks": torch.zeros(1, 1, 2, 1, dtype=torch.uint8),
|
||||
"n_frames": 1,
|
||||
"scores": [1.5],
|
||||
"orig_size": (2, 8),
|
||||
},
|
||||
"0 to 1",
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_adapter_rejects_incompatible_private_payloads(track_data, message):
|
||||
with pytest.raises((TypeError, ValueError), match=message):
|
||||
validate_sam3_track_data(track_data)
|
||||
|
||||
|
||||
def test_track_report_is_small_deterministic_and_history_safe():
|
||||
tracks = sam3_track_data_to_tracks(
|
||||
_sample_track_data(),
|
||||
seed_detections=_seed_detections(),
|
||||
fps=10.0,
|
||||
)
|
||||
payload = track_report_payload(tracks)
|
||||
assert payload["track_count"] == 2
|
||||
assert payload["observation_count"] == 3
|
||||
assert payload["state_counts"] == {"active": 2}
|
||||
encoded = track_report_json(tracks)
|
||||
assert json.loads(encoded) == payload
|
||||
assert "packed_masks" not in encoded
|
||||
text = track_report_text(tracks)
|
||||
assert "Tracks: 2" in text
|
||||
assert "#7 cat" in text
|
||||
assert "#9 fish" in text
|
||||
|
||||
node_result = VLMTrackReport().report(tracks)
|
||||
assert node_result["result"] == (encoded, text)
|
||||
assert node_result["ui"]["text"] == [text]
|
||||
@@ -0,0 +1,376 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
import pytest
|
||||
from ComfyUI_VLM_nodes.nodes.spatial_parser import (
|
||||
COORDINATE_MODES,
|
||||
VLMSpatialPromptBuilder,
|
||||
VLMStructuredSpatialParser,
|
||||
build_spatial_prompt,
|
||||
load_json_document,
|
||||
parse_spatial_response,
|
||||
)
|
||||
from ComfyUI_VLM_nodes.nodes.vision_types import (
|
||||
VLM_DETECTIONS,
|
||||
VLM_POINTS,
|
||||
DetectionSequence,
|
||||
PointSequence,
|
||||
)
|
||||
|
||||
|
||||
def test_prompt_builder_is_explicit_and_provider_neutral():
|
||||
prompt = build_spatial_prompt(
|
||||
"Find every vehicle.",
|
||||
coordinate_mode="normalized_0_1000",
|
||||
width=1920,
|
||||
height=1080,
|
||||
frame_count=12,
|
||||
fps=24.0,
|
||||
)
|
||||
|
||||
assert prompt.startswith("Perform this visual analysis task:")
|
||||
assert "Find every vehicle." in prompt
|
||||
assert "Return only one valid JSON object" in prompt
|
||||
assert "normalized_0_1000" in prompt
|
||||
assert '"frame_count":12' in prompt
|
||||
assert '"fps":24.0' in prompt
|
||||
assert "zero-based frame_index" in prompt
|
||||
assert "bbox_xyxy" in prompt
|
||||
assert "polygon" in prompt
|
||||
assert '"point"' in prompt
|
||||
assert "score" in prompt
|
||||
example = json.loads(prompt.split("Required JSON shape:\n", 1)[1])
|
||||
assert example["coordinate_mode"] == "normalized_0_1000"
|
||||
|
||||
|
||||
def test_node_contracts_use_canonical_spatial_types():
|
||||
assert tuple(COORDINATE_MODES) == (
|
||||
"pixel",
|
||||
"normalized_0_1",
|
||||
"normalized_0_1000",
|
||||
)
|
||||
assert VLMSpatialPromptBuilder.RETURN_TYPES == ("STRING",)
|
||||
assert VLMStructuredSpatialParser.RETURN_TYPES == (
|
||||
VLM_DETECTIONS,
|
||||
VLM_POINTS,
|
||||
"STRING",
|
||||
)
|
||||
schema = VLMStructuredSpatialParser.INPUT_TYPES()
|
||||
assert tuple(schema["required"]["coordinate_mode"][0]) == COORDINATE_MODES
|
||||
|
||||
|
||||
def test_parser_accepts_only_complete_plain_or_fenced_json():
|
||||
assert load_json_document(" ") == {}
|
||||
assert load_json_document('```json\n{"frames":[]}\n```') == {"frames": []}
|
||||
|
||||
for invalid in (
|
||||
'Here is the result: {"frames":[]}',
|
||||
'Result:\n```json\n{"frames":[]}\n```',
|
||||
'{"frames":[]} trailing',
|
||||
"```python\n{}\n```",
|
||||
'{"x": 1, "x": 2}',
|
||||
'{"score": NaN}',
|
||||
):
|
||||
with pytest.raises(ValueError):
|
||||
load_json_document(invalid)
|
||||
|
||||
|
||||
def test_normalized_video_parse_clips_and_preserves_metadata():
|
||||
response = json.dumps(
|
||||
{
|
||||
"coordinate_mode": "normalized_0_1",
|
||||
"media": {
|
||||
"width": 200,
|
||||
"height": 100,
|
||||
"frame_count": 3,
|
||||
"fps": 2,
|
||||
"codec": "test-codec",
|
||||
},
|
||||
"source": "unit-vlm",
|
||||
"metadata": {"request_id": "abc"},
|
||||
"vendor": {"latency_ms": 12},
|
||||
"frames": [
|
||||
{
|
||||
"frame_index": 0,
|
||||
"metadata": {"scene": "start"},
|
||||
"detections": [
|
||||
{
|
||||
"class": "cat",
|
||||
"confidence": 0.75,
|
||||
"bbox": [-0.1, 0.2, 1.2, 0.8],
|
||||
"polygon": [
|
||||
[-0.5, 0.2],
|
||||
[0.5, 0.2],
|
||||
[0.5, 1.5],
|
||||
],
|
||||
"instance_id": "cat-1",
|
||||
"metadata": {"occluded": False},
|
||||
}
|
||||
],
|
||||
"points": [
|
||||
{
|
||||
"name": "nose",
|
||||
"point": [0.25, 0.5],
|
||||
"confidence": 0.9,
|
||||
"landmark_id": 7,
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"frame_index": 2,
|
||||
"detections": [
|
||||
{
|
||||
"label": "sign",
|
||||
"quad": [
|
||||
[0.1, 0.1],
|
||||
[0.9, 0.1],
|
||||
[0.9, 0.9],
|
||||
[0.1, 0.9],
|
||||
],
|
||||
"text": "STOP",
|
||||
}
|
||||
],
|
||||
},
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
detections, points, normalized_json = parse_spatial_response(
|
||||
f"```json\n{response}\n```",
|
||||
width=200,
|
||||
height=100,
|
||||
coordinate_mode="normalized_0_1",
|
||||
)
|
||||
|
||||
assert isinstance(detections, DetectionSequence)
|
||||
assert isinstance(points, PointSequence)
|
||||
assert detections.frame_count == points.frame_count == 3
|
||||
assert detections.fps == points.fps == 2.0
|
||||
assert [frame.frame_index for frame in detections.frames] == [0, 2]
|
||||
cat, sign = detections.all_detections()
|
||||
assert cat.bbox_xyxy == (0.0, 20.0, 200.0, 80.0)
|
||||
assert cat.polygon == ((0.0, 20.0), (100.0, 20.0), (100.0, 100.0))
|
||||
assert cat.label == "cat"
|
||||
assert cat.score == 0.75
|
||||
assert cat.metadata.to_dict() == {
|
||||
"instance_id": "cat-1",
|
||||
"occluded": False,
|
||||
}
|
||||
assert sign.bbox_xyxy == (20.0, 10.0, 180.0, 90.0)
|
||||
assert sign.quad is not None and len(sign.quad) == 4
|
||||
assert sign.text == "STOP"
|
||||
assert points.points[0].x == 50.0
|
||||
assert points.points[0].y == 50.0
|
||||
assert points.points[0].label == "nose"
|
||||
assert points.points[0].metadata["landmark_id"] == 7
|
||||
assert detections.frames[0].metadata["scene"] == "start"
|
||||
assert detections.metadata.to_dict() == {
|
||||
"coordinate_mode": "normalized_0_1",
|
||||
"media_metadata": {"codec": "test-codec"},
|
||||
"request_id": "abc",
|
||||
"vendor": {"latency_ms": 12},
|
||||
}
|
||||
|
||||
normalized = json.loads(normalized_json)
|
||||
assert normalized["schema"] == "comfyui-vlm/spatial"
|
||||
assert normalized["detections"] == detections.to_dict()
|
||||
assert normalized["points"] == points.to_dict()
|
||||
|
||||
|
||||
def test_pixel_aliases_xywh_flat_segmentation_and_multiple_points():
|
||||
response = json.dumps(
|
||||
{
|
||||
"objects": [
|
||||
{
|
||||
"name": "panel",
|
||||
"box": {"x": -2, "y": 5, "width": 15, "height": 30},
|
||||
"segmentation": [0, 5, 13, 5, 13, 35, 0, 35],
|
||||
"points": [[2, 7], [12, 30]],
|
||||
"score": 1,
|
||||
}
|
||||
]
|
||||
}
|
||||
)
|
||||
detections, points, _json = parse_spatial_response(
|
||||
response,
|
||||
width=10,
|
||||
height=20,
|
||||
coordinate_mode="pixel",
|
||||
frame_count=1,
|
||||
)
|
||||
|
||||
detection = detections.all_detections()[0]
|
||||
assert detection.bbox_xyxy == (0.0, 5.0, 10.0, 20.0)
|
||||
assert detection.polygon == (
|
||||
(0.0, 5.0),
|
||||
(10.0, 5.0),
|
||||
(10.0, 20.0),
|
||||
(0.0, 20.0),
|
||||
)
|
||||
assert [(point.x, point.y) for point in points.points] == [
|
||||
(2.0, 7.0),
|
||||
(10.0, 20.0),
|
||||
]
|
||||
|
||||
|
||||
def test_direct_multi_point_record_keeps_shared_fields_without_duplicates():
|
||||
detections, points, _json = parse_spatial_response(
|
||||
'{"label":"hand","confidence":0.8,"points":[[1,2],[3,4]],'
|
||||
'"metadata":{"side":"left"}}',
|
||||
width=10,
|
||||
height=10,
|
||||
)
|
||||
|
||||
assert detections.all_detections() == ()
|
||||
assert [(point.x, point.y) for point in points.points] == [
|
||||
(1.0, 2.0),
|
||||
(3.0, 4.0),
|
||||
]
|
||||
assert {point.label for point in points.points} == {"hand"}
|
||||
assert {point.score for point in points.points} == {0.8}
|
||||
assert points.points[0].metadata["side"] == "left"
|
||||
assert detections.frames[0].metadata.to_dict() == {}
|
||||
|
||||
|
||||
def test_top_level_record_batch_groups_video_frame_indices_and_timestamps():
|
||||
response = json.dumps(
|
||||
[
|
||||
{"frame_index": 2, "bbox_xyxy": [1, 2, 3, 4], "label": "late"},
|
||||
{"frame_index": 0, "point": [5, 6], "label": "early"},
|
||||
{"frame_index": 0, "box": [0, 0, 4, 5], "label": "first"},
|
||||
]
|
||||
)
|
||||
detections, points, _json = parse_spatial_response(
|
||||
response,
|
||||
width=20,
|
||||
height=10,
|
||||
fps=4,
|
||||
coordinate_mode="pixel",
|
||||
)
|
||||
|
||||
assert [frame.frame_index for frame in detections.frames] == [0, 2]
|
||||
assert detections.frames[1].timestamp == 0.5
|
||||
assert detections.frame_count == 3
|
||||
assert [item.label for item in detections.all_detections()] == [
|
||||
"first",
|
||||
"late",
|
||||
]
|
||||
assert points.points[0].frame_index == 0
|
||||
|
||||
|
||||
def test_normalized_1000_polygon_without_box_derives_clipped_bbox():
|
||||
detections, points, _json = parse_spatial_response(
|
||||
'{"polygon":[[-100,100],[500,100],[1200,900]],"label":"shape"}',
|
||||
width=300,
|
||||
height=200,
|
||||
coordinate_mode="normalized_0_1000",
|
||||
)
|
||||
|
||||
detection = detections.all_detections()[0]
|
||||
assert detection.bbox_xyxy == (0.0, 20.0, 300.0, 180.0)
|
||||
assert not points.points
|
||||
|
||||
|
||||
def test_empty_payloads_are_valid_and_predictable():
|
||||
for response in ("", "{}", "[]", '{"frames":[]}'):
|
||||
detections, points, normalized_json = parse_spatial_response(
|
||||
response,
|
||||
width=640,
|
||||
height=480,
|
||||
coordinate_mode="pixel",
|
||||
frame_count=5,
|
||||
fps=25,
|
||||
source="empty-test",
|
||||
)
|
||||
assert detections.frames == ()
|
||||
assert detections.frame_count == 5
|
||||
assert detections.fps == 25
|
||||
assert points.points == ()
|
||||
assert points.frame_count == 5
|
||||
normalized = json.loads(normalized_json)
|
||||
assert normalized["detections"]["frames"] == []
|
||||
assert normalized["points"]["points"] == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("response", "message"),
|
||||
[
|
||||
('{"bbox":[4,4,2,5]}', "x2 >= x1"),
|
||||
('{"polygon":[[1,1],[2,2]]}', "at least three"),
|
||||
('{"quad":[[0,0],[1,0],[1,1]]}', "exactly 4"),
|
||||
('{"point":[1]}', "two coordinates"),
|
||||
('{"score":2,"point":[1,1]}', "between 0 and 1"),
|
||||
(
|
||||
'{"bbox":[0,0,1,1],"polygon":[[0,0],[1,0],[0,1]],'
|
||||
'"quad":[[0,0],[1,0],[1,1],[0,1]]}',
|
||||
"either polygon",
|
||||
),
|
||||
(
|
||||
'{"bbox":[0,0,1,1],"coordinate_mode":"normalized_0_1"}',
|
||||
"parser is set",
|
||||
),
|
||||
('{"label":"no geometry"}', "requires bbox"),
|
||||
],
|
||||
)
|
||||
def test_strict_validation(response, message):
|
||||
with pytest.raises((TypeError, ValueError), match=message):
|
||||
parse_spatial_response(
|
||||
response,
|
||||
width=10,
|
||||
height=10,
|
||||
coordinate_mode="pixel",
|
||||
)
|
||||
|
||||
|
||||
def test_dimensions_timing_and_alias_conflicts_are_rejected():
|
||||
cases = [
|
||||
('{"media":{"width":20}}', {"width": 10}, "does not match"),
|
||||
('{"media":{"fps":30}}', {"fps": 24}, "does not match"),
|
||||
(
|
||||
'{"label":"a","class":"b","point":[1,1]}',
|
||||
{},
|
||||
"Conflicting aliases",
|
||||
),
|
||||
(
|
||||
'{"frames":[{"frame_index":0},{"frame_index":0}]}',
|
||||
{},
|
||||
"Duplicate frame_index",
|
||||
),
|
||||
(
|
||||
'{"detections":[],"objects":[]}',
|
||||
{},
|
||||
"multiple detection collection",
|
||||
),
|
||||
]
|
||||
for response, overrides, message in cases:
|
||||
arguments = {
|
||||
"width": 10,
|
||||
"height": 10,
|
||||
"coordinate_mode": "pixel",
|
||||
**overrides,
|
||||
}
|
||||
with pytest.raises(ValueError, match=message):
|
||||
parse_spatial_response(response, **arguments)
|
||||
|
||||
|
||||
def test_node_methods_return_direct_canonical_payloads():
|
||||
prompt = VLMSpatialPromptBuilder().build(
|
||||
"Locate the subject.",
|
||||
"pixel",
|
||||
100,
|
||||
80,
|
||||
1,
|
||||
0,
|
||||
)[0]
|
||||
detections, points, normalized = VLMStructuredSpatialParser().parse(
|
||||
'{"bbox_xyxy":[1,2,30,40],"point":[5,6]}',
|
||||
"pixel",
|
||||
100,
|
||||
80,
|
||||
)
|
||||
|
||||
assert isinstance(prompt, str)
|
||||
assert isinstance(detections, DetectionSequence)
|
||||
assert isinstance(points, PointSequence)
|
||||
assert json.loads(normalized)["version"] == 1
|
||||
@@ -0,0 +1,234 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
from ComfyUI_VLM_nodes.nodes.tracking import (
|
||||
VLMByteTracker,
|
||||
associate_detection_sequence,
|
||||
)
|
||||
from ComfyUI_VLM_nodes.nodes.vision_types import (
|
||||
Detection,
|
||||
DetectionSequence,
|
||||
FrameDetections,
|
||||
)
|
||||
|
||||
|
||||
def _frame(
|
||||
frame_index,
|
||||
detections,
|
||||
*,
|
||||
width=100,
|
||||
height=100,
|
||||
fps=10.0,
|
||||
):
|
||||
timestamp = frame_index / fps
|
||||
records = tuple(
|
||||
Detection(
|
||||
bbox_xyxy=record["box"],
|
||||
label=record.get("label"),
|
||||
score=record.get("score"),
|
||||
frame_index=frame_index,
|
||||
timestamp=timestamp,
|
||||
mask=record.get("mask"),
|
||||
)
|
||||
for record in detections
|
||||
)
|
||||
return FrameDetections(
|
||||
frame_index=frame_index,
|
||||
timestamp=timestamp,
|
||||
width=width,
|
||||
height=height,
|
||||
detections=records,
|
||||
)
|
||||
|
||||
|
||||
def _sequence(frames, *, frame_count=None, fps=10.0, width=100, height=100):
|
||||
return DetectionSequence(
|
||||
width=width,
|
||||
height=height,
|
||||
frames=tuple(frames),
|
||||
frame_count=frame_count or 0,
|
||||
fps=fps,
|
||||
source="test-detector",
|
||||
)
|
||||
|
||||
|
||||
def test_low_confidence_second_stage_keeps_the_track_id():
|
||||
sequence = _sequence(
|
||||
(
|
||||
_frame(
|
||||
0,
|
||||
[{"box": (10, 10, 30, 30), "score": 0.9, "label": "cat"}],
|
||||
),
|
||||
_frame(
|
||||
1,
|
||||
[{"box": (11, 10, 31, 30), "score": 0.25, "label": "cat"}],
|
||||
),
|
||||
)
|
||||
)
|
||||
tracks = associate_detection_sequence(
|
||||
sequence,
|
||||
high_threshold=0.6,
|
||||
low_threshold=0.1,
|
||||
min_hits=1,
|
||||
)
|
||||
assert len(tracks.tracks) == 1
|
||||
track = tracks.tracks[0]
|
||||
assert track.track_id == 0
|
||||
assert [item.track_id for item in track.detections] == [0, 0]
|
||||
assert track.detections[1].metadata["association_stage"] == "low"
|
||||
assert track.metadata["state"] == "active"
|
||||
|
||||
|
||||
def test_low_confidence_detection_cannot_start_a_track():
|
||||
sequence = _sequence(
|
||||
(
|
||||
_frame(
|
||||
0,
|
||||
[{"box": (10, 10, 30, 30), "score": 0.2, "label": "cat"}],
|
||||
),
|
||||
)
|
||||
)
|
||||
tracks = associate_detection_sequence(
|
||||
sequence,
|
||||
high_threshold=0.6,
|
||||
low_threshold=0.1,
|
||||
min_hits=1,
|
||||
)
|
||||
assert tracks.tracks == ()
|
||||
|
||||
|
||||
def test_label_aware_matching_prevents_cross_class_identity_reuse():
|
||||
frames = (
|
||||
_frame(
|
||||
0,
|
||||
[{"box": (10, 10, 30, 30), "score": 0.9, "label": "cat"}],
|
||||
),
|
||||
_frame(
|
||||
1,
|
||||
[{"box": (10, 10, 30, 30), "score": 0.9, "label": "dog"}],
|
||||
),
|
||||
)
|
||||
label_aware = associate_detection_sequence(
|
||||
_sequence(frames),
|
||||
min_hits=1,
|
||||
label_aware=True,
|
||||
emit_predictions=False,
|
||||
)
|
||||
class_agnostic = associate_detection_sequence(
|
||||
_sequence(frames),
|
||||
min_hits=1,
|
||||
label_aware=False,
|
||||
emit_predictions=False,
|
||||
)
|
||||
assert len(label_aware.tracks) == 2
|
||||
assert [track.label for track in label_aware.tracks] == ["cat", "dog"]
|
||||
assert len(class_agnostic.tracks) == 1
|
||||
assert [
|
||||
detection.frame_index for detection in class_agnostic.tracks[0].detections
|
||||
] == [0, 1]
|
||||
|
||||
|
||||
def test_max_age_seconds_uses_fps_and_marks_removed_deterministically():
|
||||
sequence = _sequence(
|
||||
(
|
||||
_frame(
|
||||
0,
|
||||
[{"box": (10, 10, 30, 30), "score": 0.9}],
|
||||
fps=2.0,
|
||||
),
|
||||
),
|
||||
frame_count=4,
|
||||
fps=2.0,
|
||||
)
|
||||
tracks = associate_detection_sequence(
|
||||
sequence,
|
||||
min_hits=1,
|
||||
max_age_seconds=1.0,
|
||||
emit_predictions=True,
|
||||
)
|
||||
track = tracks.tracks[0]
|
||||
assert [item.frame_index for item in track.detections] == [0, 1, 2]
|
||||
assert [item.metadata["observation"] for item in track.detections] == [
|
||||
"detected",
|
||||
"predicted",
|
||||
"predicted",
|
||||
]
|
||||
assert track.metadata["state"] == "removed"
|
||||
assert track.metadata["removed_frame"] == 3
|
||||
assert track.metadata["last_observed_frame"] == 0
|
||||
|
||||
|
||||
def test_mask_iou_can_rescue_a_zero_bbox_iou_match():
|
||||
mask = torch.zeros(100, 100)
|
||||
mask[40:60, 40:60] = 1
|
||||
sequence = _sequence(
|
||||
(
|
||||
_frame(
|
||||
0,
|
||||
[
|
||||
{
|
||||
"box": (0, 0, 10, 10),
|
||||
"score": 0.9,
|
||||
"mask": mask,
|
||||
}
|
||||
],
|
||||
),
|
||||
_frame(
|
||||
1,
|
||||
[
|
||||
{
|
||||
"box": (20, 0, 30, 10),
|
||||
"score": 0.9,
|
||||
"mask": mask,
|
||||
}
|
||||
],
|
||||
),
|
||||
)
|
||||
)
|
||||
tracks = VLMByteTracker(
|
||||
min_hits=1,
|
||||
motion_gate=1.0e12,
|
||||
).track(sequence)
|
||||
assert len(tracks.tracks) == 1
|
||||
assert len(tracks.tracks[0].detections) == 2
|
||||
|
||||
|
||||
def test_hungarian_results_and_serialization_are_repeatable():
|
||||
sequence = _sequence(
|
||||
(
|
||||
_frame(
|
||||
0,
|
||||
[
|
||||
{"box": (5, 5, 20, 20), "score": 0.9},
|
||||
{"box": (40, 5, 55, 20), "score": 0.9},
|
||||
],
|
||||
),
|
||||
_frame(
|
||||
1,
|
||||
[
|
||||
{"box": (41, 5, 56, 20), "score": 0.9},
|
||||
{"box": (6, 5, 21, 20), "score": 0.9},
|
||||
],
|
||||
),
|
||||
)
|
||||
)
|
||||
options = {"min_hits": 1, "emit_predictions": False}
|
||||
first = associate_detection_sequence(sequence, **options)
|
||||
second = associate_detection_sequence(sequence, **options)
|
||||
assert first.to_json() == second.to_json()
|
||||
assert [
|
||||
[detection.bbox_xyxy for detection in track.detections]
|
||||
for track in first.tracks
|
||||
] == [
|
||||
[(5.0, 5.0, 20.0, 20.0), (6.0, 5.0, 21.0, 20.0)],
|
||||
[(40.0, 5.0, 55.0, 20.0), (41.0, 5.0, 56.0, 20.0)],
|
||||
]
|
||||
|
||||
|
||||
def test_tracker_rejects_invalid_threshold_order():
|
||||
try:
|
||||
VLMByteTracker(high_threshold=0.2, low_threshold=0.3)
|
||||
except ValueError as error:
|
||||
assert "low_threshold" in str(error)
|
||||
else:
|
||||
raise AssertionError("Expected invalid threshold order to fail.")
|
||||
@@ -0,0 +1,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),),
|
||||
)
|
||||
@@ -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
@@ -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");
|
||||
}
|
||||
}
|
||||
},
|
||||
});
|
||||
|
||||
Reference in New Issue
Block a user