Instructions to use PRVSM01/sam2.1-memory-attention-dynamic-ONNX with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- sam2
How to use PRVSM01/sam2.1-memory-attention-dynamic-ONNX with sam2:
# Use SAM2 with images import torch from sam2.sam2_image_predictor import SAM2ImagePredictor predictor = SAM2ImagePredictor.from_pretrained(PRVSM01/sam2.1-memory-attention-dynamic-ONNX) with torch.inference_mode(), torch.autocast("cuda", dtype=torch.bfloat16): predictor.set_image(<your_image>) masks, _, _ = predictor.predict(<input_prompts>)# Use SAM2 with videos import torch from sam2.sam2_video_predictor import SAM2VideoPredictor predictor = SAM2VideoPredictor.from_pretrained(PRVSM01/sam2.1-memory-attention-dynamic-ONNX) with torch.inference_mode(), torch.autocast("cuda", dtype=torch.bfloat16): state = predictor.init_state(<your_video>) # add new prompts and instantly get the output on the same frame frame_idx, object_ids, masks = predictor.add_new_points(state, <your_prompts>): # propagate the prompts to get masklets throughout the video for frame_idx, object_ids, masks in predictor.propagate_in_video(state): ... - Notebooks
- Google Colab
- Kaggle
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 candidateiou_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
- -
Model tree for PRVSM01/sam2.1-memory-attention-dynamic-ONNX
Base model
facebook/sam2.1-hiera-base-plus