SAM2.1 custom ONNX exports (PRVSM Smart Mask pipeline)

Custom ONNX re-exports of submodules from facebook/sam2.1-hiera-base-plus (Meta, Apache-2.0), used by PRVSM's Smart Mask (SAM2.1-based video object segmentation, runs fully client-side in-browser via WebGPU/onnxruntime-web).

memory_attention_dynamic.onnx

Re-export of the memory_attention submodule via torch.onnx.export(..., dynamo=True, dynamic_shapes=...) so the memory sequence length (spatial memory + object pointer tokens) is a true dynamic axis.

Why this exists

The existing community ONNX export of SAM2's video pipeline (s-tornqvist/sam2-hiera-base-plus-video-ONNX) traces memory_attention with a fixed memory length (exactly 7 frames + 64 object-pointer tokens). That works once the memory bank is full, but produces incorrect (and eventually degenerate) masks for the first few frames after a click, since the real model was never trained on a zero-padded/repeated memory bank (verified by direct comparison against the real PyTorch model).

This export accepts any real memory length instead, matching the model's actual training/inference behavior for partially-filled memory banks.

Inputs / outputs

Same tensor names and shapes as the underlying Sam2VideoMemoryAttention module, split into separate spatial_memory / object_pointer_memory inputs (concatenated internally) so the true count of real object-pointer tokens is derivable from tensor shape rather than a traced-in constant:

  • current_vision_features: (4096, 1, 256)
  • current_vision_position_embeddings: (4096, 1, 256)
  • spatial_memory: (N*4096, 1, 64), N = number of memory frames (dynamic)
  • spatial_memory_pos: (N*4096, 1, 64)
  • object_pointer_memory: (M, 1, 64), M = number of object-pointer sub-tokens (dynamic)
  • object_pointer_memory_pos: (M, 1, 64)
  • output refined_vision_features: (1, 1, 4096, 256)

prompt_encoder_mask_decoder_video_raw.onnx

Re-export of the same prompt_encoder + mask_decoder pair as s-tornqvist/sam2-hiera-base-plus-video-ONNX's prompt_encoder_mask_decoder_video_fp32.onnx, but with the internal "select best candidate mask" step (_dynamic_multimask_via_stability in transformers' modeling_sam2_video.py) removed from the traced graph.

Why this exists

That internal selection step uses torch.gather on a 5D tensor ([batch, point_batch, num_candidates, H, W]), which onnxruntime-web's WebGPU (JSEP) backend cannot compile for the GatherElements op on 5D inputs (works fine on SAM3-tracker's 4D output, breaks on SAM2.1's 5D one). This is a kernel bug, not an unsupported-op case, so there is no automatic CPU fallback for just that node, forcing the entire decoder graph onto WASM in the community export.

This export returns the 4 raw mask candidates + their IoU scores instead of doing the selection inside the graph, so the graph has zero GatherElements nodes and runs entirely on WebGPU. The selection (candidate 0 if its "stability score" clears a threshold, otherwise the best of candidates 1-3 by IoU) is trivial to redo outside the graph on the small (256x256) already- computed arrays. Verified bit-exact (both branches) against the real PyTorch Sam2VideoMaskDecoder._dynamic_multimask_via_stability reference. Measured ~30ms on WebGPU vs. ~370ms forced onto WASM for the original graph (same environment, same input sizes).

Inputs / outputs

Same 5 inputs as the original prompt_encoder_mask_decoder_video_fp32.onnx:

  • image_embeddings.0: (1, 32, 256, 256)
  • image_embeddings.1: (1, 64, 128, 128)
  • image_embeddings.2: (1, 256, 64, 64)
  • image_positional_embeddings: (1, 256, 64, 64)
  • input_masks: (1, 1, 256, 256)

Outputs (all 4 mask candidates, no internal selection):

  • pred_masks: (1, 1, 4, 256, 256), low-res mask logits per candidate
  • iou_scores: (1, 1, 4)
  • mask_token: (1, 1, 256), always candidate 0's token (matches original graph's behavior, which never re-selects this regardless of which mask candidate wins)
  • object_score_logits: (1, 1, 1)

Consumer must reproduce the selection (delta=0.05, thresh=0.98):

area_i = count(pred_masks[..., 0, :, :] > delta)
area_u = count(pred_masks[..., 0, :, :] > -delta)
stability = area_u > 0 ? area_i / area_u : 1
if stability >= thresh: use candidate 0
else: use argmax(iou_scores[..., 1:]) + 1

License

Apache-2.0, inherited from facebook/sam2.1-hiera-base-plus. Derivative work (re-export only, same weights, no retraining), not an original model.

Downloads last month
-
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Model tree for PRVSM01/sam2.1-memory-attention-dynamic-ONNX

Quantized
(3)
this model