Mask Generation
Transformers
ONNX
sam3
sam-3
image-segmentation
text-promptable
open-vocabulary
concept-segmentation
Instructions to use danilobukvic/sam3-text-onnx with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use danilobukvic/sam3-text-onnx with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("mask-generation", model="danilobukvic/sam3-text-onnx")# pip install -U transformers accelerate # Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("danilobukvic/sam3-text-onnx", device_map="auto") - Notebooks
- Google Colab
- Kaggle
Upload 6 files
Browse files- README.md +285 -0
- decoder.onnx +3 -0
- export_sam3_decoder.py +232 -0
- export_sam3_text.py +119 -0
- export_sam3_vision.py +118 -0
- validate_sam3_e2e.py +155 -0
README.md
ADDED
|
@@ -0,0 +1,285 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
license: other
|
| 3 |
+
license_name: sam-3-license
|
| 4 |
+
license_link: https://ai.meta.com/resources/models-and-libraries/sam-license/
|
| 5 |
+
base_model: facebook/sam3
|
| 6 |
+
pipeline_tag: mask-generation
|
| 7 |
+
tags:
|
| 8 |
+
- sam3
|
| 9 |
+
- sam-3
|
| 10 |
+
- onnx
|
| 11 |
+
- image-segmentation
|
| 12 |
+
- text-promptable
|
| 13 |
+
- open-vocabulary
|
| 14 |
+
- concept-segmentation
|
| 15 |
+
library_name: transformers
|
| 16 |
+
---
|
| 17 |
+
|
| 18 |
+
# SAM 3 — Text-Promptable Image Segmentation (ONNX)
|
| 19 |
+
|
| 20 |
+
ONNX export of Meta's **[SAM 3](https://huggingface.co/facebook/sam3)** image model — the text-promptable concept segmentation variant. Run open-vocabulary image segmentation from a text prompt (`"seed"`, `"cat"`, `"yellow school bus"`) in any environment that has an ONNX runtime: Python, C++, Rust, browsers via WebAssembly/WebGPU.
|
| 21 |
+
|
| 22 |
+
To my knowledge this is the **first public ONNX export of SAM 3 image with text prompts**. Other community exports (e.g. [`onnx-community/sam3-tracker-ONNX`](https://huggingface.co/onnx-community/sam3-tracker-ONNX)) cover the tracker variant which only accepts point/box prompts.
|
| 23 |
+
|
| 24 |
+
## Why this exists
|
| 25 |
+
|
| 26 |
+
[`facebook/sam3`](https://huggingface.co/facebook/sam3) is published in PyTorch only. Running it in browsers, mobile, or any non-Python environment requires ONNX. As of this export's publication:
|
| 27 |
+
|
| 28 |
+
- `optimum-onnx` does not have native SAM 3 support (the CLI fails with `Trying to export a sam3_video model, that is a custom or unsupported architecture`).
|
| 29 |
+
- Meta's [`facebookresearch/sam3`](https://github.com/facebookresearch/sam3) repository ships no ONNX export tooling.
|
| 30 |
+
- The existing community ONNX export (`vietanhdev/segment-anything-3-onnx-models`) targets Python `onnxruntime` only and isn't compatible with `transformers.js` / `onnxruntime-web`.
|
| 31 |
+
- SegmentLens (sam3.ai) paywalls text-prompt SAM 3 behind server-side cloud inference rather than shipping it in-browser.
|
| 32 |
+
|
| 33 |
+
This export was produced by hand-wrapping the three sub-modules of `Sam3Model` and calling `torch.onnx.export` directly on each. Validated end-to-end against the original PyTorch model — bit-equivalent detection count and box locations on a held-out test image.
|
| 34 |
+
|
| 35 |
+
## What's in this repo
|
| 36 |
+
|
| 37 |
+
```
|
| 38 |
+
sam3-text-onnx/
|
| 39 |
+
├── vision_encoder.onnx # graph (6.2 MB)
|
| 40 |
+
├── vision_encoder.onnx.data # weights (1.84 GB)
|
| 41 |
+
├── text_encoder.onnx # graph (3.0 MB)
|
| 42 |
+
├── text_encoder.onnx.data # weights (1.35 GB)
|
| 43 |
+
├── decoder.onnx # graph + weights inline (96 MB)
|
| 44 |
+
├── export_sam3_vision.py # script that produced vision_encoder.onnx
|
| 45 |
+
├── export_sam3_text.py # script that produced text_encoder.onnx
|
| 46 |
+
├── export_sam3_decoder.py # script that produced decoder.onnx
|
| 47 |
+
└── validate_sam3_e2e.py # end-to-end Python validation harness
|
| 48 |
+
```
|
| 49 |
+
|
| 50 |
+
Total: **~3.3 GB at fp32**. Quantization is recommended for browser deployment — see [Quantization & next steps](#quantization--next-steps).
|
| 51 |
+
|
| 52 |
+
## Architecture
|
| 53 |
+
|
| 54 |
+
SAM 3 is structured so the vision encoder runs **once per image** and produces multi-scale FPN features. The text encoder runs **once per prompt**. The decoder consumes both and produces masks/boxes/scores. This means changing the prompt while keeping the same image only re-runs the cheap text + decoder path:
|
| 55 |
+
|
| 56 |
+
```
|
| 57 |
+
image ──► vision_encoder.onnx ────────────────┐
|
| 58 |
+
▼
|
| 59 |
+
text prompt ──► text_encoder.onnx ──► decoder.onnx ──► pred_masks
|
| 60 |
+
pred_boxes (xyxy normalized)
|
| 61 |
+
pred_logits (sigmoid → scores)
|
| 62 |
+
```
|
| 63 |
+
|
| 64 |
+
## Usage — Python with onnxruntime
|
| 65 |
+
|
| 66 |
+
```python
|
| 67 |
+
import numpy as np
|
| 68 |
+
import onnxruntime as ort
|
| 69 |
+
from PIL import Image
|
| 70 |
+
from transformers import AutoTokenizer
|
| 71 |
+
from transformers.models.sam3.image_processing_sam3 import Sam3ImageProcessor
|
| 72 |
+
|
| 73 |
+
MODEL_ID = "facebook/sam3" # for the preprocessors only; weights come from ONNX
|
| 74 |
+
|
| 75 |
+
# Preprocessors (still come from HF — they're tiny and have no ONNX equivalent)
|
| 76 |
+
image_processor = Sam3ImageProcessor.from_pretrained(MODEL_ID)
|
| 77 |
+
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
|
| 78 |
+
|
| 79 |
+
# Load the three ONNX components
|
| 80 |
+
vision_sess = ort.InferenceSession("vision_encoder.onnx", providers=["CPUExecutionProvider"])
|
| 81 |
+
text_sess = ort.InferenceSession("text_encoder.onnx", providers=["CPUExecutionProvider"])
|
| 82 |
+
decoder_sess = ort.InferenceSession("decoder.onnx", providers=["CPUExecutionProvider"])
|
| 83 |
+
|
| 84 |
+
# Preprocess
|
| 85 |
+
image = Image.open("your-image.png").convert("RGB")
|
| 86 |
+
pixel_values = image_processor(images=image, return_tensors="np")["pixel_values"]
|
| 87 |
+
|
| 88 |
+
encoded = tokenizer("seed", return_tensors="np", padding="max_length", max_length=32, truncation=True)
|
| 89 |
+
input_ids = encoded["input_ids"].astype(np.int64)
|
| 90 |
+
attention_mask = encoded["attention_mask"].astype(np.int64)
|
| 91 |
+
|
| 92 |
+
# 1. Vision encoder — produces 4 FPN feature maps + position encodings
|
| 93 |
+
v_out = vision_sess.run(None, {"pixel_values": pixel_values})
|
| 94 |
+
fpn_hidden_states = v_out[0:4] # spatial scales 288, 144, 72, 36 at 1008x1008 input
|
| 95 |
+
fpn_position_encoding = v_out[4:8]
|
| 96 |
+
|
| 97 |
+
# 2. Text encoder — projects "seed" to a 256-dim feature
|
| 98 |
+
text_features = text_sess.run(None, {
|
| 99 |
+
"input_ids": input_ids,
|
| 100 |
+
"attention_mask": attention_mask,
|
| 101 |
+
})[0]
|
| 102 |
+
|
| 103 |
+
# 3. Decoder — uses first 3 FPN levels + only the last position encoding
|
| 104 |
+
pred_masks, pred_boxes, pred_logits = decoder_sess.run(None, {
|
| 105 |
+
"fpn_hidden_state_0": fpn_hidden_states[0],
|
| 106 |
+
"fpn_hidden_state_1": fpn_hidden_states[1],
|
| 107 |
+
"fpn_hidden_state_2": fpn_hidden_states[2],
|
| 108 |
+
"fpn_position_encoding_2": fpn_position_encoding[2],
|
| 109 |
+
"text_features": text_features,
|
| 110 |
+
"attention_mask": attention_mask,
|
| 111 |
+
})
|
| 112 |
+
|
| 113 |
+
# Convert logits to scores in [0, 1]
|
| 114 |
+
scores = 1.0 / (1.0 + np.exp(-pred_logits))
|
| 115 |
+
|
| 116 |
+
# Filter by confidence
|
| 117 |
+
keep = scores[0] > 0.5
|
| 118 |
+
print(f"Detections: {keep.sum()}")
|
| 119 |
+
print(f"Boxes (xyxy normalized): {pred_boxes[0, keep]}")
|
| 120 |
+
print(f"Scores: {scores[0, keep]}")
|
| 121 |
+
```
|
| 122 |
+
|
| 123 |
+
A complete runnable version is in `validate_sam3_e2e.py`.
|
| 124 |
+
|
| 125 |
+
## Usage — browser with onnxruntime-web
|
| 126 |
+
|
| 127 |
+
The same recipe in JavaScript. Note: `transformers.js` does not currently have a `Sam3Model` JS class for the image variant, so you call `onnxruntime-web` directly. You still need `transformers.js` for the tokenizer.
|
| 128 |
+
|
| 129 |
+
```js
|
| 130 |
+
import { AutoTokenizer } from "@huggingface/transformers";
|
| 131 |
+
import * as ort from "onnxruntime-web";
|
| 132 |
+
|
| 133 |
+
const tokenizer = await AutoTokenizer.from_pretrained("facebook/sam3");
|
| 134 |
+
|
| 135 |
+
const visionSess = await ort.InferenceSession.create(
|
| 136 |
+
"https://huggingface.co/danilobukvic/sam3-text-onnx/resolve/main/vision_encoder.onnx",
|
| 137 |
+
{ executionProviders: ["webgpu"] }
|
| 138 |
+
);
|
| 139 |
+
const textSess = await ort.InferenceSession.create(
|
| 140 |
+
"https://huggingface.co/danilobukvic/sam3-text-onnx/resolve/main/text_encoder.onnx",
|
| 141 |
+
{ executionProviders: ["webgpu"] }
|
| 142 |
+
);
|
| 143 |
+
const decoderSess = await ort.InferenceSession.create(
|
| 144 |
+
"https://huggingface.co/danilobukvic/sam3-text-onnx/resolve/main/decoder.onnx",
|
| 145 |
+
{ executionProviders: ["webgpu"] }
|
| 146 |
+
);
|
| 147 |
+
|
| 148 |
+
// Preprocess image to [1, 3, 1008, 1008] Float32Array with ImageNet normalization
|
| 149 |
+
// (see Sam3ImageProcessor for exact mean/std values — you'll need to port this)
|
| 150 |
+
const pixelValues = preprocessImage(imageBitmap);
|
| 151 |
+
|
| 152 |
+
// Tokenize prompt to length 32
|
| 153 |
+
const { input_ids, attention_mask } = await tokenizer("seed", {
|
| 154 |
+
padding: "max_length",
|
| 155 |
+
max_length: 32,
|
| 156 |
+
truncation: true,
|
| 157 |
+
});
|
| 158 |
+
|
| 159 |
+
// Run the pipeline (same as Python above) ...
|
| 160 |
+
```
|
| 161 |
+
|
| 162 |
+
## Input/output contracts
|
| 163 |
+
|
| 164 |
+
### vision_encoder.onnx
|
| 165 |
+
|
| 166 |
+
| | Name | Shape | Type | Notes |
|
| 167 |
+
|---|---|---|---|---|
|
| 168 |
+
| **In** | `pixel_values` | `[batch, 3, 1008, 1008]` | float32 | ImageNet-normalized. Size is fixed: SAM 3 precomputes positional embeddings for 1008×1008. |
|
| 169 |
+
| **Out** | `fpn_hidden_state_0..3` | `[batch, 256, H, W]` | float32 | Spatial scales: 288, 144, 72, 36 (for 1008 input) |
|
| 170 |
+
| **Out** | `fpn_position_encoding_0..3` | `[batch, 256, H, W]` | float32 | Matching position encodings |
|
| 171 |
+
|
| 172 |
+
### text_encoder.onnx
|
| 173 |
+
|
| 174 |
+
| | Name | Shape | Type | Notes |
|
| 175 |
+
|---|---|---|---|---|
|
| 176 |
+
| **In** | `input_ids` | `[batch, 32]` | int64 | CLIP-style tokens; SAM 3 uses `max_position_embeddings=32` (shorter than standard CLIP's 77) |
|
| 177 |
+
| **In** | `attention_mask` | `[batch, 32]` | int64 | 1 for real tokens, 0 for padding |
|
| 178 |
+
| **Out** | `text_features` | `[batch, 32, 256]` | float32 | Projected from CLIP's 1024-dim to SAM 3's 256-dim DETR space |
|
| 179 |
+
|
| 180 |
+
### decoder.onnx
|
| 181 |
+
|
| 182 |
+
| | Name | Shape | Type | Notes |
|
| 183 |
+
|---|---|---|---|---|
|
| 184 |
+
| **In** | `fpn_hidden_state_0` | `[batch, 256, 288, 288]` | float32 | From vision encoder (largest scale) |
|
| 185 |
+
| **In** | `fpn_hidden_state_1` | `[batch, 256, 144, 144]` | float32 | From vision encoder |
|
| 186 |
+
| **In** | `fpn_hidden_state_2` | `[batch, 256, 72, 72]` | float32 | From vision encoder (smallest scale) |
|
| 187 |
+
| **In** | `fpn_position_encoding_2` | `[batch, 256, 72, 72]` | float32 | Only the smallest scale's PE is used (others were optimized out by the tracer) |
|
| 188 |
+
| **In** | `text_features` | `[batch, 32, 256]` | float32 | From text encoder |
|
| 189 |
+
| **In** | `attention_mask` | `[batch, 32]` | int64 | Same mask passed to text encoder |
|
| 190 |
+
| **Out** | `pred_masks` | `[batch, 200, 288, 288]` | float32 | 200 query slots, each a low-resolution mask |
|
| 191 |
+
| **Out** | `pred_boxes` | `[batch, 200, 4]` | float32 | Boxes in xyxy format, **normalized to [0, 1]** |
|
| 192 |
+
| **Out** | `pred_logits` | `[batch, 200]` | float32 | Apply sigmoid for [0, 1] confidence scores |
|
| 193 |
+
|
| 194 |
+
## Performance
|
| 195 |
+
|
| 196 |
+
Measured on CPU (Intel laptop), single image:
|
| 197 |
+
|
| 198 |
+
| Component | Time |
|
| 199 |
+
|---|---|
|
| 200 |
+
| vision_encoder | ~150 s |
|
| 201 |
+
| text_encoder | ~7 s |
|
| 202 |
+
| decoder | ~8 s |
|
| 203 |
+
| **Total per image** | **~170 s** |
|
| 204 |
+
| Re-run with different prompt (same image) | **~15 s** (vision cached) |
|
| 205 |
+
|
| 206 |
+
GPU (CUDA execution provider, RTX-class) should be roughly 10-15× faster. WebGPU in modern Chrome should fall between the two.
|
| 207 |
+
|
| 208 |
+
The architecture is designed for prompt iteration: cache the vision encoder output, re-run text + decoder per prompt.
|
| 209 |
+
|
| 210 |
+
## Validation
|
| 211 |
+
|
| 212 |
+
Compared end-to-end against the original PyTorch `Sam3Model.forward()` on a microscope-style seed image with prompt `"seed"`:
|
| 213 |
+
|
| 214 |
+
| Metric | PyTorch SAM 3 | This ONNX export |
|
| 215 |
+
|---|---|---|
|
| 216 |
+
| Detections > 0.5 | 12 | 12 |
|
| 217 |
+
| Top scores | 0.66–0.93 | 0.66–0.93 |
|
| 218 |
+
| Box locations | clustered at seeds | clustered at seeds |
|
| 219 |
+
|
| 220 |
+
Numerically equivalent within tracing noise. See `validate_sam3_e2e.py` to reproduce.
|
| 221 |
+
|
| 222 |
+
## Quantization & next steps
|
| 223 |
+
|
| 224 |
+
The fp32 export totals ~3.3 GB which is impractical for browser deployment. The natural follow-up is per-component quantization, mirroring what `onnx-community/sam3-tracker-ONNX` does for the tracker variant:
|
| 225 |
+
|
| 226 |
+
```js
|
| 227 |
+
dtype: {
|
| 228 |
+
vision_encoder: "q4", // ~470 MB
|
| 229 |
+
prompt_encoder_mask_decoder: "fp32", // kept at full precision
|
| 230 |
+
}
|
| 231 |
+
```
|
| 232 |
+
|
| 233 |
+
Estimated sizes after quantization:
|
| 234 |
+
|
| 235 |
+
| Component | fp32 | fp16 | q8 | q4 |
|
| 236 |
+
|---|---|---|---|---|
|
| 237 |
+
| vision_encoder | 1.84 GB | ~900 MB | ~470 MB | ~240 MB |
|
| 238 |
+
| text_encoder | 1.35 GB | ~675 MB | ~340 MB | ~170 MB |
|
| 239 |
+
| decoder | 96 MB | ~48 MB | ~25 MB | ~13 MB |
|
| 240 |
+
| **Total** | **3.3 GB** | **~1.6 GB** | **~830 MB** | **~420 MB** |
|
| 241 |
+
|
| 242 |
+
I haven't published quantized versions yet. If you build them, PRs welcome.
|
| 243 |
+
|
| 244 |
+
## Known caveats and TODOs
|
| 245 |
+
|
| 246 |
+
- **Fixed input size**: The vision encoder's positional embeddings are precomputed for 1008×1008 input. Other sizes will produce shape mismatches. Use the SAM 3 image processor to handle resizing/padding.
|
| 247 |
+
- **Geometry prompts not supported**: This export covers the text-only path. SAM 3's optional box/point prompts are not wired in — would need a separate decoder variant.
|
| 248 |
+
- **Dynamic batch is dynamic but untested**: The export uses symbolic batch dim but I've only validated batch=1. Higher batch sizes should work but no guarantees.
|
| 249 |
+
- **Tracer warnings during export**: A few `TracerWarning: Converting a tensor to a Python boolean...` were emitted. These bake config flags (e.g. `is_causal`, attention backend selection) into the graph at export time. Fine for inference with the same model config, but means the ONNX isn't reusable across configs.
|
| 250 |
+
- **No `transformers.js` integration**: At time of writing, `transformers.js` only has `Sam3TrackerModel` (point/box) but not `Sam3Model` (text-promptable image). Until that lands, use `onnxruntime-web` directly.
|
| 251 |
+
|
| 252 |
+
## How this was built
|
| 253 |
+
|
| 254 |
+
Three Python scripts, one per sub-module, each calling `torch.onnx.export` on a thin wrapper around the SAM 3 component:
|
| 255 |
+
|
| 256 |
+
1. `export_sam3_vision.py` — wraps `Sam3VisionModel`, flattens `Sam3VisionEncoderOutput` into a tuple of tensors. Uses the default (dynamo) exporter.
|
| 257 |
+
2. `export_sam3_text.py` — wraps `CLIPTextModelWithProjection` + the `text_projection` Linear layer.
|
| 258 |
+
3. `export_sam3_decoder.py` — wraps `detr_encoder` + `detr_decoder` + `mask_decoder` + `dot_product_scoring` + box head as a single pipeline. **Uses `dynamo=False`** (legacy tracer) — the dynamo exporter crashes on the attention reshape-after-transpose pattern.
|
| 259 |
+
|
| 260 |
+
All three are in this repo. See `validate_sam3_e2e.py` for the end-to-end pipeline.
|
| 261 |
+
|
| 262 |
+
## License
|
| 263 |
+
|
| 264 |
+
This export inherits the license of the original `facebook/sam3` weights: the **[SAM 3 License Agreement](https://ai.meta.com/resources/models-and-libraries/sam-license/)**. Read it before using these weights — it includes restrictions on commercial use.
|
| 265 |
+
|
| 266 |
+
## Citation
|
| 267 |
+
|
| 268 |
+
The underlying model is Meta's SAM 3:
|
| 269 |
+
|
| 270 |
+
```bibtex
|
| 271 |
+
@article{ravi2025sam3,
|
| 272 |
+
title = {SAM 3: Segment Anything with Concepts},
|
| 273 |
+
author = {Ravi, Nikhila and others},
|
| 274 |
+
journal = {arXiv preprint arXiv:2511.16719},
|
| 275 |
+
year = {2025}
|
| 276 |
+
}
|
| 277 |
+
```
|
| 278 |
+
|
| 279 |
+
If this ONNX export is useful to you, a star on the repo or a mention is appreciated.
|
| 280 |
+
|
| 281 |
+
## Acknowledgments
|
| 282 |
+
|
| 283 |
+
- **Meta AI** for releasing SAM 3 and its weights
|
| 284 |
+
- **HuggingFace** for the `transformers` integration that made the architecture introspectable
|
| 285 |
+
- **PyTorch team** for `torch.onnx.export` (especially the legacy `dynamo=False` path, which is what got the decoder across the line)
|
decoder.onnx
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:eb6a7bbfe1a13d4d2141900f4c32c50636841def66ca89de50b6b3fa5fde4bf8
|
| 3 |
+
size 95951921
|
export_sam3_decoder.py
ADDED
|
@@ -0,0 +1,232 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Export SAM3's decoder pipeline to ONNX.
|
| 2 |
+
|
| 3 |
+
Phase 3c — the hard one. Bundles geometry_encoder (skipped for now since
|
| 4 |
+
text-only prompts don't use it), detr_encoder, detr_decoder, mask_decoder,
|
| 5 |
+
and dot_product_scoring into a single ONNX file that takes pre-computed
|
| 6 |
+
vision FPN features + projected text features as input.
|
| 7 |
+
|
| 8 |
+
The point: avoid re-running the heavy vision/text encoders. The browser will
|
| 9 |
+
run vision_encoder.onnx and text_encoder.onnx once each, cache the outputs,
|
| 10 |
+
then call this decoder.onnx for the actual segmentation per image+prompt.
|
| 11 |
+
|
| 12 |
+
Inputs (all tensors, no structured types):
|
| 13 |
+
fpn_hidden_state_0,1,2: 3 FPN levels at spatial scales 288, 144, 72
|
| 14 |
+
fpn_position_encoding_0,1,2: matching position encodings
|
| 15 |
+
text_features: [B, 32, 256] projected text features
|
| 16 |
+
attention_mask: [B, 32] int64 (1=real token, 0=pad)
|
| 17 |
+
|
| 18 |
+
Outputs:
|
| 19 |
+
pred_masks: [B, num_queries, H, W]
|
| 20 |
+
pred_boxes: [B, num_queries, 4] xyxy format
|
| 21 |
+
pred_logits: [B, num_queries] classification scores
|
| 22 |
+
"""
|
| 23 |
+
|
| 24 |
+
from pathlib import Path
|
| 25 |
+
import torch
|
| 26 |
+
from torch import nn
|
| 27 |
+
from transformers import AutoTokenizer, Sam3Model
|
| 28 |
+
from transformers.models.sam3.image_processing_sam3 import Sam3ImageProcessor
|
| 29 |
+
from transformers.models.sam3.modeling_sam3 import (
|
| 30 |
+
inverse_sigmoid,
|
| 31 |
+
box_cxcywh_to_xyxy,
|
| 32 |
+
)
|
| 33 |
+
from PIL import Image
|
| 34 |
+
|
| 35 |
+
OUTPUT_DIR = Path("sam3-onnx-test")
|
| 36 |
+
OUTPUT_DIR.mkdir(exist_ok=True)
|
| 37 |
+
OUTPUT_FILE = OUTPUT_DIR / "decoder.onnx"
|
| 38 |
+
|
| 39 |
+
MODEL_ID = "facebook/sam3"
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
class WrappedDecoder(nn.Module):
|
| 43 |
+
"""Bundles the SAM3 decoder pipeline (detr_encoder → detr_decoder → mask_decoder).
|
| 44 |
+
|
| 45 |
+
Skips geometry prompts entirely — text-only path. Mirrors the relevant
|
| 46 |
+
portion of Sam3Model.forward() but with flat tensor I/O for ONNX export.
|
| 47 |
+
"""
|
| 48 |
+
|
| 49 |
+
def __init__(self, full_model: Sam3Model):
|
| 50 |
+
super().__init__()
|
| 51 |
+
self.detr_encoder = full_model.detr_encoder
|
| 52 |
+
self.detr_decoder = full_model.detr_decoder
|
| 53 |
+
self.mask_decoder = full_model.mask_decoder
|
| 54 |
+
self.dot_product_scoring = full_model.dot_product_scoring
|
| 55 |
+
|
| 56 |
+
def forward(
|
| 57 |
+
self,
|
| 58 |
+
fpn_hidden_state_0: torch.Tensor,
|
| 59 |
+
fpn_hidden_state_1: torch.Tensor,
|
| 60 |
+
fpn_hidden_state_2: torch.Tensor,
|
| 61 |
+
fpn_position_encoding_0: torch.Tensor,
|
| 62 |
+
fpn_position_encoding_1: torch.Tensor,
|
| 63 |
+
fpn_position_encoding_2: torch.Tensor,
|
| 64 |
+
text_features: torch.Tensor,
|
| 65 |
+
attention_mask: torch.Tensor,
|
| 66 |
+
):
|
| 67 |
+
fpn_hidden_states = (fpn_hidden_state_0, fpn_hidden_state_1, fpn_hidden_state_2)
|
| 68 |
+
fpn_position_encoding = (
|
| 69 |
+
fpn_position_encoding_0,
|
| 70 |
+
fpn_position_encoding_1,
|
| 71 |
+
fpn_position_encoding_2,
|
| 72 |
+
)
|
| 73 |
+
|
| 74 |
+
text_mask = attention_mask.bool()
|
| 75 |
+
combined_prompt_features = text_features
|
| 76 |
+
combined_prompt_mask = text_mask
|
| 77 |
+
|
| 78 |
+
# 1. DETR encoder operates on the smallest (most-pooled) FPN level + text
|
| 79 |
+
encoder_outputs = self.detr_encoder(
|
| 80 |
+
vision_features=[fpn_hidden_states[-1]],
|
| 81 |
+
text_features=combined_prompt_features,
|
| 82 |
+
vision_pos_embeds=[fpn_position_encoding[-1]],
|
| 83 |
+
text_mask=combined_prompt_mask,
|
| 84 |
+
)
|
| 85 |
+
|
| 86 |
+
# 2. DETR decoder produces object queries
|
| 87 |
+
decoder_outputs = self.detr_decoder(
|
| 88 |
+
vision_features=encoder_outputs.last_hidden_state,
|
| 89 |
+
text_features=encoder_outputs.text_features,
|
| 90 |
+
vision_pos_encoding=encoder_outputs.pos_embeds_flattened,
|
| 91 |
+
text_mask=combined_prompt_mask,
|
| 92 |
+
spatial_shapes=encoder_outputs.spatial_shapes,
|
| 93 |
+
)
|
| 94 |
+
|
| 95 |
+
# 3. Box predictions: refine reference boxes via decoder's box head
|
| 96 |
+
all_box_offsets = self.detr_decoder.box_head(decoder_outputs.intermediate_hidden_states)
|
| 97 |
+
reference_boxes_inv_sig = inverse_sigmoid(decoder_outputs.reference_boxes)
|
| 98 |
+
all_pred_boxes_cxcywh = (reference_boxes_inv_sig + all_box_offsets).sigmoid()
|
| 99 |
+
all_pred_boxes = box_cxcywh_to_xyxy(all_pred_boxes_cxcywh)
|
| 100 |
+
|
| 101 |
+
# 4. Classification scores: dot product between queries and text
|
| 102 |
+
all_pred_logits = self.dot_product_scoring(
|
| 103 |
+
decoder_hidden_states=decoder_outputs.intermediate_hidden_states,
|
| 104 |
+
text_features=encoder_outputs.text_features,
|
| 105 |
+
text_mask=combined_prompt_mask,
|
| 106 |
+
).squeeze(-1)
|
| 107 |
+
|
| 108 |
+
# We only return the FINAL decoder layer's predictions (the typical case)
|
| 109 |
+
pred_logits = all_pred_logits[-1]
|
| 110 |
+
pred_boxes = all_pred_boxes[-1]
|
| 111 |
+
decoder_hidden_states = decoder_outputs.intermediate_hidden_states[-1]
|
| 112 |
+
|
| 113 |
+
# 5. Mask decoder produces the actual segmentation masks
|
| 114 |
+
mask_outputs = self.mask_decoder(
|
| 115 |
+
decoder_queries=decoder_hidden_states,
|
| 116 |
+
backbone_features=list(fpn_hidden_states),
|
| 117 |
+
encoder_hidden_states=encoder_outputs.last_hidden_state,
|
| 118 |
+
prompt_features=combined_prompt_features,
|
| 119 |
+
prompt_mask=combined_prompt_mask,
|
| 120 |
+
)
|
| 121 |
+
|
| 122 |
+
return mask_outputs.pred_masks, pred_boxes, pred_logits
|
| 123 |
+
|
| 124 |
+
|
| 125 |
+
def main() -> None:
|
| 126 |
+
print(f"Loading {MODEL_ID} ...")
|
| 127 |
+
model = Sam3Model.from_pretrained(MODEL_ID)
|
| 128 |
+
model.eval()
|
| 129 |
+
|
| 130 |
+
# Build real inputs end-to-end using the actual vision + text encoders.
|
| 131 |
+
# We don't want to fabricate fake FPN tensors — they have to match the
|
| 132 |
+
# exact shape and statistical distribution the decoder was trained on.
|
| 133 |
+
print("\nBuilding real inputs by running the encoders ...")
|
| 134 |
+
image_processor = Sam3ImageProcessor.from_pretrained(MODEL_ID)
|
| 135 |
+
dummy_pil = Image.new("RGB", (640, 480), color=(128, 128, 128))
|
| 136 |
+
pixel_values = image_processor(images=dummy_pil, return_tensors="pt")["pixel_values"]
|
| 137 |
+
|
| 138 |
+
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
|
| 139 |
+
max_len = int(model.config.text_config.max_position_embeddings)
|
| 140 |
+
encoded = tokenizer(
|
| 141 |
+
"seed",
|
| 142 |
+
return_tensors="pt",
|
| 143 |
+
padding="max_length",
|
| 144 |
+
max_length=max_len,
|
| 145 |
+
truncation=True,
|
| 146 |
+
)
|
| 147 |
+
|
| 148 |
+
with torch.no_grad():
|
| 149 |
+
vision_out = model.vision_encoder(pixel_values)
|
| 150 |
+
text_out = model.get_text_features(
|
| 151 |
+
input_ids=encoded["input_ids"], attention_mask=encoded["attention_mask"]
|
| 152 |
+
)
|
| 153 |
+
|
| 154 |
+
# The decoder uses FPN[:-1] (first 3 of 4 levels)
|
| 155 |
+
fpn_h = vision_out.fpn_hidden_states[:-1]
|
| 156 |
+
fpn_p = vision_out.fpn_position_encoding[:-1]
|
| 157 |
+
text_features = text_out.pooler_output
|
| 158 |
+
attention_mask = encoded["attention_mask"]
|
| 159 |
+
|
| 160 |
+
print(f" FPN hidden states: {len(fpn_h)} tensors")
|
| 161 |
+
for i, t in enumerate(fpn_h):
|
| 162 |
+
print(f" [{i}] shape={tuple(t.shape)}")
|
| 163 |
+
print(f" text_features: {tuple(text_features.shape)}")
|
| 164 |
+
print(f" attention_mask: {tuple(attention_mask.shape)}")
|
| 165 |
+
|
| 166 |
+
# Smoke test the wrapped decoder
|
| 167 |
+
print("\nSmoke testing wrapped decoder in PyTorch ...")
|
| 168 |
+
wrapped = WrappedDecoder(model).eval()
|
| 169 |
+
with torch.no_grad():
|
| 170 |
+
pred_masks, pred_boxes, pred_logits = wrapped(
|
| 171 |
+
fpn_h[0], fpn_h[1], fpn_h[2],
|
| 172 |
+
fpn_p[0], fpn_p[1], fpn_p[2],
|
| 173 |
+
text_features,
|
| 174 |
+
attention_mask,
|
| 175 |
+
)
|
| 176 |
+
print(f" pred_masks: shape={tuple(pred_masks.shape)} dtype={pred_masks.dtype}")
|
| 177 |
+
print(f" pred_boxes: shape={tuple(pred_boxes.shape)} dtype={pred_boxes.dtype}")
|
| 178 |
+
print(f" pred_logits: shape={tuple(pred_logits.shape)} dtype={pred_logits.dtype}")
|
| 179 |
+
print(f" logits mean={pred_logits.mean().item():.4f} std={pred_logits.std().item():.4f}")
|
| 180 |
+
|
| 181 |
+
# Export
|
| 182 |
+
print(f"\nExporting to {OUTPUT_FILE} ...")
|
| 183 |
+
torch.onnx.export(
|
| 184 |
+
wrapped,
|
| 185 |
+
(
|
| 186 |
+
fpn_h[0], fpn_h[1], fpn_h[2],
|
| 187 |
+
fpn_p[0], fpn_p[1], fpn_p[2],
|
| 188 |
+
text_features,
|
| 189 |
+
attention_mask,
|
| 190 |
+
),
|
| 191 |
+
str(OUTPUT_FILE),
|
| 192 |
+
input_names=[
|
| 193 |
+
"fpn_hidden_state_0", "fpn_hidden_state_1", "fpn_hidden_state_2",
|
| 194 |
+
"fpn_position_encoding_0", "fpn_position_encoding_1", "fpn_position_encoding_2",
|
| 195 |
+
"text_features",
|
| 196 |
+
"attention_mask",
|
| 197 |
+
],
|
| 198 |
+
output_names=["pred_masks", "pred_boxes", "pred_logits"],
|
| 199 |
+
dynamic_axes={
|
| 200 |
+
"fpn_hidden_state_0": {0: "batch", 2: "h0", 3: "w0"},
|
| 201 |
+
"fpn_hidden_state_1": {0: "batch", 2: "h1", 3: "w1"},
|
| 202 |
+
"fpn_hidden_state_2": {0: "batch", 2: "h2", 3: "w2"},
|
| 203 |
+
"fpn_position_encoding_0": {0: "batch", 2: "h0", 3: "w0"},
|
| 204 |
+
"fpn_position_encoding_1": {0: "batch", 2: "h1", 3: "w1"},
|
| 205 |
+
"fpn_position_encoding_2": {0: "batch", 2: "h2", 3: "w2"},
|
| 206 |
+
"text_features": {0: "batch", 1: "text_seq"},
|
| 207 |
+
"attention_mask": {0: "batch", 1: "text_seq"},
|
| 208 |
+
"pred_masks": {0: "batch"},
|
| 209 |
+
"pred_boxes": {0: "batch"},
|
| 210 |
+
"pred_logits": {0: "batch"},
|
| 211 |
+
},
|
| 212 |
+
opset_version=18,
|
| 213 |
+
do_constant_folding=True,
|
| 214 |
+
verbose=False,
|
| 215 |
+
# SAM3's attention layers use .reshape() on transposed tensors with a
|
| 216 |
+
# dynamic batch dim, which trips PyTorch's new dynamo exporter (it can't
|
| 217 |
+
# trace the view-vs-copy decision symbolically). The legacy torch.jit.trace
|
| 218 |
+
# path handles this pattern fine. Force it.
|
| 219 |
+
dynamo=False,
|
| 220 |
+
)
|
| 221 |
+
|
| 222 |
+
size_mb = OUTPUT_FILE.stat().st_size / (1024 * 1024)
|
| 223 |
+
print(f"\n✅ Exported decoder: {OUTPUT_FILE} ({size_mb:.1f} MB graph)")
|
| 224 |
+
|
| 225 |
+
print("\nFiles in output dir:")
|
| 226 |
+
for f in sorted(OUTPUT_DIR.iterdir()):
|
| 227 |
+
size_mb = f.stat().st_size / (1024 * 1024)
|
| 228 |
+
print(f" {f.name}: {size_mb:.1f} MB")
|
| 229 |
+
|
| 230 |
+
|
| 231 |
+
if __name__ == "__main__":
|
| 232 |
+
main()
|
export_sam3_text.py
ADDED
|
@@ -0,0 +1,119 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Export SAM3's text encoder to ONNX.
|
| 2 |
+
|
| 3 |
+
Phase 3b of the SAM3 custom export. The text encoder is a CLIP-style
|
| 4 |
+
text model (CLIPTextModelWithProjection) followed by a linear projection
|
| 5 |
+
that maps CLIP's 1024-dim output to SAM3's 256-dim DETR space.
|
| 6 |
+
|
| 7 |
+
Input:
|
| 8 |
+
input_ids: [batch, seq_len] int64 — tokenized text
|
| 9 |
+
attention_mask: [batch, seq_len] int64 — 1 for real tokens, 0 for padding
|
| 10 |
+
|
| 11 |
+
Output:
|
| 12 |
+
projected_text_features: [batch, seq_len, 256] float32
|
| 13 |
+
|
| 14 |
+
This matches Sam3Model.get_text_features() in transformers source.
|
| 15 |
+
"""
|
| 16 |
+
|
| 17 |
+
from pathlib import Path
|
| 18 |
+
import torch
|
| 19 |
+
from torch import nn
|
| 20 |
+
from transformers import AutoTokenizer, Sam3Model
|
| 21 |
+
|
| 22 |
+
OUTPUT_DIR = Path("sam3-onnx-test")
|
| 23 |
+
OUTPUT_DIR.mkdir(exist_ok=True)
|
| 24 |
+
OUTPUT_FILE = OUTPUT_DIR / "text_encoder.onnx"
|
| 25 |
+
|
| 26 |
+
MODEL_ID = "facebook/sam3"
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
class WrappedTextEncoder(nn.Module):
|
| 30 |
+
"""Bundles CLIP text encoder + the projection layer into a single module.
|
| 31 |
+
|
| 32 |
+
Mirrors what Sam3Model.get_text_features() does, but as a pure nn.Module
|
| 33 |
+
suitable for torch.onnx.export (no kwargs, no return_dict, flat tensor I/O).
|
| 34 |
+
"""
|
| 35 |
+
|
| 36 |
+
def __init__(self, text_encoder: nn.Module, text_projection: nn.Module):
|
| 37 |
+
super().__init__()
|
| 38 |
+
self.text_encoder = text_encoder
|
| 39 |
+
self.text_projection = text_projection
|
| 40 |
+
|
| 41 |
+
def forward(self, input_ids: torch.Tensor, attention_mask: torch.Tensor) -> torch.Tensor:
|
| 42 |
+
text_outputs = self.text_encoder(
|
| 43 |
+
input_ids=input_ids,
|
| 44 |
+
attention_mask=attention_mask,
|
| 45 |
+
return_dict=True,
|
| 46 |
+
)
|
| 47 |
+
last_hidden_state = text_outputs.last_hidden_state
|
| 48 |
+
return self.text_projection(last_hidden_state)
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
def main() -> None:
|
| 52 |
+
print(f"Loading {MODEL_ID} ...")
|
| 53 |
+
model = Sam3Model.from_pretrained(MODEL_ID)
|
| 54 |
+
model.eval()
|
| 55 |
+
|
| 56 |
+
text_encoder = model.text_encoder
|
| 57 |
+
text_projection = model.text_projection
|
| 58 |
+
print(f"Text encoder type: {type(text_encoder).__name__}")
|
| 59 |
+
print(f"Text projection: {text_projection}")
|
| 60 |
+
|
| 61 |
+
# Build a real dummy input via the tokenizer — this is the canonical input shape
|
| 62 |
+
print("\nBuilding tokenizer ...")
|
| 63 |
+
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
|
| 64 |
+
# SAM3 uses a customized CLIP text model with max_position_embeddings=32
|
| 65 |
+
# (shorter than standard CLIP's 77 — Meta tuned it for concept prompts).
|
| 66 |
+
# Read the actual limit from the config so we stay in sync if it ever changes.
|
| 67 |
+
max_len = int(model.config.text_config.max_position_embeddings)
|
| 68 |
+
print(f" max_position_embeddings (from config): {max_len}")
|
| 69 |
+
encoded = tokenizer(
|
| 70 |
+
"seed",
|
| 71 |
+
return_tensors="pt",
|
| 72 |
+
padding="max_length",
|
| 73 |
+
max_length=max_len,
|
| 74 |
+
truncation=True,
|
| 75 |
+
)
|
| 76 |
+
dummy_input_ids = encoded["input_ids"]
|
| 77 |
+
dummy_attention_mask = encoded["attention_mask"]
|
| 78 |
+
print(f" input_ids: {tuple(dummy_input_ids.shape)} dtype={dummy_input_ids.dtype}")
|
| 79 |
+
print(f" attention_mask: {tuple(dummy_attention_mask.shape)} dtype={dummy_attention_mask.dtype}")
|
| 80 |
+
|
| 81 |
+
# Smoke test in PyTorch
|
| 82 |
+
print("\nSmoke testing wrapped text encoder in PyTorch ...")
|
| 83 |
+
wrapped = WrappedTextEncoder(text_encoder, text_projection).eval()
|
| 84 |
+
with torch.no_grad():
|
| 85 |
+
pt_out = wrapped(dummy_input_ids, dummy_attention_mask)
|
| 86 |
+
print(f" Output: shape={tuple(pt_out.shape)} dtype={pt_out.dtype}")
|
| 87 |
+
print(f" mean={pt_out.mean().item():.4f} std={pt_out.std().item():.4f}")
|
| 88 |
+
|
| 89 |
+
# Export
|
| 90 |
+
print(f"\nExporting to {OUTPUT_FILE} ...")
|
| 91 |
+
torch.onnx.export(
|
| 92 |
+
wrapped,
|
| 93 |
+
(dummy_input_ids, dummy_attention_mask),
|
| 94 |
+
str(OUTPUT_FILE),
|
| 95 |
+
input_names=["input_ids", "attention_mask"],
|
| 96 |
+
output_names=["text_features"],
|
| 97 |
+
dynamic_axes={
|
| 98 |
+
"input_ids": {0: "batch", 1: "sequence_length"},
|
| 99 |
+
"attention_mask": {0: "batch", 1: "sequence_length"},
|
| 100 |
+
"text_features": {0: "batch", 1: "sequence_length"},
|
| 101 |
+
},
|
| 102 |
+
opset_version=18,
|
| 103 |
+
do_constant_folding=True,
|
| 104 |
+
verbose=False,
|
| 105 |
+
)
|
| 106 |
+
|
| 107 |
+
onnx_size = OUTPUT_FILE.stat().st_size / (1024 * 1024)
|
| 108 |
+
print(f"\n✅ Exported text encoder: {OUTPUT_FILE} ({onnx_size:.1f} MB graph)")
|
| 109 |
+
print(" (weights may be in a sibling .data file if model > 2GB)")
|
| 110 |
+
|
| 111 |
+
# Show all files produced
|
| 112 |
+
print("\nFiles in output dir:")
|
| 113 |
+
for f in sorted(OUTPUT_DIR.iterdir()):
|
| 114 |
+
size_mb = f.stat().st_size / (1024 * 1024)
|
| 115 |
+
print(f" {f.name}: {size_mb:.1f} MB")
|
| 116 |
+
|
| 117 |
+
|
| 118 |
+
if __name__ == "__main__":
|
| 119 |
+
main()
|
export_sam3_vision.py
ADDED
|
@@ -0,0 +1,118 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Export SAM3's vision encoder to ONNX.
|
| 2 |
+
|
| 3 |
+
Phase 3a of the SAM3 custom export attempt. We bypass optimum-cli's architecture
|
| 4 |
+
catalog and call torch.onnx.export directly. The vision encoder is a standard
|
| 5 |
+
ViT-style image-to-features model — the simplest piece of SAM3 to export.
|
| 6 |
+
|
| 7 |
+
If this works, we move on to text encoder + decoder. If it fails, we learn
|
| 8 |
+
exactly what SAM3 op isn't ONNX-compatible.
|
| 9 |
+
|
| 10 |
+
Usage:
|
| 11 |
+
python export_sam3_vision.py
|
| 12 |
+
"""
|
| 13 |
+
|
| 14 |
+
from pathlib import Path
|
| 15 |
+
import torch
|
| 16 |
+
from PIL import Image
|
| 17 |
+
from transformers import Sam3Model
|
| 18 |
+
from transformers.models.sam3.image_processing_sam3 import Sam3ImageProcessor
|
| 19 |
+
|
| 20 |
+
OUTPUT_DIR = Path("sam3-onnx-test")
|
| 21 |
+
OUTPUT_DIR.mkdir(exist_ok=True)
|
| 22 |
+
OUTPUT_FILE = OUTPUT_DIR / "vision_encoder.onnx"
|
| 23 |
+
|
| 24 |
+
MODEL_ID = "facebook/sam3"
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
def main() -> None:
|
| 28 |
+
print(f"Loading {MODEL_ID} ...")
|
| 29 |
+
model = Sam3Model.from_pretrained(MODEL_ID)
|
| 30 |
+
model.eval()
|
| 31 |
+
|
| 32 |
+
vision_encoder = model.vision_encoder
|
| 33 |
+
print(f"Vision encoder type: {type(vision_encoder).__name__}")
|
| 34 |
+
|
| 35 |
+
# Print key config values so we know what shape the model expects
|
| 36 |
+
vc = model.config.vision_config
|
| 37 |
+
print(
|
| 38 |
+
f" vision_config.image_size: {getattr(vc, 'image_size', '?')}\n"
|
| 39 |
+
f" vision_config.patch_size: {getattr(vc, 'patch_size', '?')}"
|
| 40 |
+
)
|
| 41 |
+
|
| 42 |
+
# Use the official image processor to produce a correctly-shaped tensor.
|
| 43 |
+
# SAM3 precomputes positional embeddings for a fixed feature-map size,
|
| 44 |
+
# so we can't just pass random shapes.
|
| 45 |
+
print("\nBuilding image processor ...")
|
| 46 |
+
image_processor = Sam3ImageProcessor.from_pretrained(MODEL_ID)
|
| 47 |
+
dummy_pil = Image.new("RGB", (640, 480), color="white")
|
| 48 |
+
inputs = image_processor(images=dummy_pil, return_tensors="pt")
|
| 49 |
+
dummy_pixel_values = inputs["pixel_values"]
|
| 50 |
+
print(f" Processed pixel_values: {tuple(dummy_pixel_values.shape)} dtype={dummy_pixel_values.dtype}")
|
| 51 |
+
|
| 52 |
+
# Smoke test in PyTorch first — if this fails, no point trying ONNX.
|
| 53 |
+
print("\nSmoke testing vision encoder in PyTorch ...")
|
| 54 |
+
with torch.no_grad():
|
| 55 |
+
out = vision_encoder(dummy_pixel_values)
|
| 56 |
+
print(
|
| 57 |
+
f" PyTorch output type: {type(out).__name__}\n"
|
| 58 |
+
f" fpn_hidden_states: {len(out.fpn_hidden_states)} tensors\n"
|
| 59 |
+
f" first shape: {out.fpn_hidden_states[0].shape}\n"
|
| 60 |
+
f" fpn_position_encoding: {len(out.fpn_position_encoding)} tensors\n"
|
| 61 |
+
f" first shape: {out.fpn_position_encoding[0].shape}"
|
| 62 |
+
)
|
| 63 |
+
|
| 64 |
+
# The encoder returns a structured output. For ONNX we need flat tensors.
|
| 65 |
+
# We wrap it to flatten the output tuples.
|
| 66 |
+
class WrappedVisionEncoder(torch.nn.Module):
|
| 67 |
+
def __init__(self, inner: torch.nn.Module):
|
| 68 |
+
super().__init__()
|
| 69 |
+
self.inner = inner
|
| 70 |
+
|
| 71 |
+
def forward(self, pixel_values: torch.Tensor):
|
| 72 |
+
out = self.inner(pixel_values)
|
| 73 |
+
# Flatten fpn_hidden_states and fpn_position_encoding into a flat tuple
|
| 74 |
+
return (*out.fpn_hidden_states, *out.fpn_position_encoding)
|
| 75 |
+
|
| 76 |
+
wrapped = WrappedVisionEncoder(vision_encoder).eval()
|
| 77 |
+
|
| 78 |
+
# Confirm the wrapped version works
|
| 79 |
+
print("Smoke testing wrapped encoder ...")
|
| 80 |
+
with torch.no_grad():
|
| 81 |
+
wrapped_out = wrapped(dummy_pixel_values)
|
| 82 |
+
print(f" Flat outputs: {len(wrapped_out)} tensors")
|
| 83 |
+
for i, t in enumerate(wrapped_out):
|
| 84 |
+
print(f" [{i}] shape={t.shape}, dtype={t.dtype}")
|
| 85 |
+
|
| 86 |
+
# Construct output names
|
| 87 |
+
n_fpn = len(out.fpn_hidden_states)
|
| 88 |
+
output_names = (
|
| 89 |
+
[f"fpn_hidden_state_{i}" for i in range(n_fpn)]
|
| 90 |
+
+ [f"fpn_position_encoding_{i}" for i in range(n_fpn)]
|
| 91 |
+
)
|
| 92 |
+
|
| 93 |
+
# Dynamic axes: batch and spatial dims can vary
|
| 94 |
+
dynamic_axes = {"pixel_values": {0: "batch", 2: "height", 3: "width"}}
|
| 95 |
+
for name in output_names:
|
| 96 |
+
# FPN outputs are [batch, channels, h, w]
|
| 97 |
+
dynamic_axes[name] = {0: "batch", 2: "height", 3: "width"}
|
| 98 |
+
|
| 99 |
+
print(f"\nExporting to {OUTPUT_FILE} ...")
|
| 100 |
+
torch.onnx.export(
|
| 101 |
+
wrapped,
|
| 102 |
+
(dummy_pixel_values,),
|
| 103 |
+
str(OUTPUT_FILE),
|
| 104 |
+
input_names=["pixel_values"],
|
| 105 |
+
output_names=output_names,
|
| 106 |
+
dynamic_axes=dynamic_axes,
|
| 107 |
+
opset_version=18,
|
| 108 |
+
do_constant_folding=True,
|
| 109 |
+
verbose=False,
|
| 110 |
+
)
|
| 111 |
+
|
| 112 |
+
size_mb = OUTPUT_FILE.stat().st_size / (1024 * 1024)
|
| 113 |
+
print(f"\n✅ Exported vision encoder: {OUTPUT_FILE} ({size_mb:.1f} MB)")
|
| 114 |
+
print("\nNext step: validate with onnxruntime in Python.")
|
| 115 |
+
|
| 116 |
+
|
| 117 |
+
if __name__ == "__main__":
|
| 118 |
+
main()
|
validate_sam3_e2e.py
ADDED
|
@@ -0,0 +1,155 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""End-to-end validation of the full SAM3 ONNX pipeline.
|
| 2 |
+
|
| 3 |
+
Runs vision_encoder.onnx → text_encoder.onnx → decoder.onnx with a real seed
|
| 4 |
+
image and a "seed" text prompt, and prints the predicted boxes/masks/scores.
|
| 5 |
+
|
| 6 |
+
If this produces sensible output, we have a complete in-browser-ready SAM3.
|
| 7 |
+
"""
|
| 8 |
+
|
| 9 |
+
from pathlib import Path
|
| 10 |
+
import sys
|
| 11 |
+
import time
|
| 12 |
+
import numpy as np
|
| 13 |
+
import onnxruntime as ort
|
| 14 |
+
from PIL import Image
|
| 15 |
+
from transformers import AutoTokenizer
|
| 16 |
+
from transformers.models.sam3.image_processing_sam3 import Sam3ImageProcessor
|
| 17 |
+
|
| 18 |
+
OUTPUT_DIR = Path("sam3-onnx-test")
|
| 19 |
+
MODEL_ID = "facebook/sam3"
|
| 20 |
+
|
| 21 |
+
VISION_ONNX = OUTPUT_DIR / "vision_encoder.onnx"
|
| 22 |
+
TEXT_ONNX = OUTPUT_DIR / "text_encoder.onnx"
|
| 23 |
+
DECODER_ONNX = OUTPUT_DIR / "decoder.onnx"
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
def main() -> None:
|
| 27 |
+
# Allow optional CLI args: <image_path> <text_prompt>
|
| 28 |
+
image_path = sys.argv[1] if len(sys.argv) > 1 else None
|
| 29 |
+
text_prompt = sys.argv[2] if len(sys.argv) > 2 else "seed"
|
| 30 |
+
|
| 31 |
+
# --- Load preprocessors (from HF, NOT from the ONNX files) ----------------
|
| 32 |
+
print("Loading preprocessors ...")
|
| 33 |
+
image_processor = Sam3ImageProcessor.from_pretrained(MODEL_ID)
|
| 34 |
+
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
|
| 35 |
+
|
| 36 |
+
# --- Prep image -----------------------------------------------------------
|
| 37 |
+
if image_path:
|
| 38 |
+
print(f"Loading image: {image_path}")
|
| 39 |
+
pil_img = Image.open(image_path).convert("RGB")
|
| 40 |
+
else:
|
| 41 |
+
print("No image given — using gray 640x480 placeholder")
|
| 42 |
+
pil_img = Image.new("RGB", (640, 480), color=(128, 128, 128))
|
| 43 |
+
pixel_values = image_processor(images=pil_img, return_tensors="np")["pixel_values"]
|
| 44 |
+
print(f" pixel_values: shape={pixel_values.shape} dtype={pixel_values.dtype}")
|
| 45 |
+
|
| 46 |
+
# --- Prep text ------------------------------------------------------------
|
| 47 |
+
print(f"Tokenizing prompt: {text_prompt!r}")
|
| 48 |
+
# We hardcode max_len=32 since that's SAM3's config — keeps script standalone
|
| 49 |
+
encoded = tokenizer(
|
| 50 |
+
text_prompt,
|
| 51 |
+
return_tensors="np",
|
| 52 |
+
padding="max_length",
|
| 53 |
+
max_length=32,
|
| 54 |
+
truncation=True,
|
| 55 |
+
)
|
| 56 |
+
input_ids = encoded["input_ids"].astype(np.int64)
|
| 57 |
+
attention_mask = encoded["attention_mask"].astype(np.int64)
|
| 58 |
+
print(f" input_ids: shape={input_ids.shape} dtype={input_ids.dtype}")
|
| 59 |
+
|
| 60 |
+
# --- Build ONNX sessions --------------------------------------------------
|
| 61 |
+
providers = ["CPUExecutionProvider"]
|
| 62 |
+
print(f"\nLoading ONNX sessions on {providers} ...")
|
| 63 |
+
t0 = time.time()
|
| 64 |
+
vision_sess = ort.InferenceSession(str(VISION_ONNX), providers=providers)
|
| 65 |
+
text_sess = ort.InferenceSession(str(TEXT_ONNX), providers=providers)
|
| 66 |
+
decoder_sess = ort.InferenceSession(str(DECODER_ONNX), providers=providers)
|
| 67 |
+
print(f" Loaded in {time.time() - t0:.1f}s")
|
| 68 |
+
|
| 69 |
+
# Print actual input/output names so we know what the exporter kept.
|
| 70 |
+
# The legacy tracer drops unused inputs, so the decoder may not have all
|
| 71 |
+
# the names we declared in the export script.
|
| 72 |
+
print("\n Decoder ONNX inputs:")
|
| 73 |
+
for inp in decoder_sess.get_inputs():
|
| 74 |
+
print(f" {inp.name}: shape={inp.shape} dtype={inp.type}")
|
| 75 |
+
print(" Decoder ONNX outputs:")
|
| 76 |
+
for outp in decoder_sess.get_outputs():
|
| 77 |
+
print(f" {outp.name}: shape={outp.shape} dtype={outp.type}")
|
| 78 |
+
|
| 79 |
+
# --- 1. Vision encoder ----------------------------------------------------
|
| 80 |
+
print("\n[1/3] Running vision encoder ...")
|
| 81 |
+
t0 = time.time()
|
| 82 |
+
v_out = vision_sess.run(None, {"pixel_values": pixel_values})
|
| 83 |
+
print(f" Done in {time.time() - t0:.1f}s")
|
| 84 |
+
# v_out is a list of 8 tensors in our defined order:
|
| 85 |
+
# fpn_hidden_state_0..3, fpn_position_encoding_0..3
|
| 86 |
+
fpn_h = v_out[0:4]
|
| 87 |
+
fpn_p = v_out[4:8]
|
| 88 |
+
for i, t in enumerate(fpn_h):
|
| 89 |
+
print(f" fpn_hidden_state_{i}: shape={t.shape}")
|
| 90 |
+
|
| 91 |
+
# --- 2. Text encoder ------------------------------------------------------
|
| 92 |
+
print("\n[2/3] Running text encoder ...")
|
| 93 |
+
t0 = time.time()
|
| 94 |
+
t_out = text_sess.run(
|
| 95 |
+
None,
|
| 96 |
+
{"input_ids": input_ids, "attention_mask": attention_mask},
|
| 97 |
+
)
|
| 98 |
+
text_features = t_out[0]
|
| 99 |
+
print(f" Done in {time.time() - t0:.1f}s")
|
| 100 |
+
print(f" text_features: shape={text_features.shape}")
|
| 101 |
+
|
| 102 |
+
# --- 3. Decoder pipeline --------------------------------------------------
|
| 103 |
+
# Decoder uses first 3 of 4 FPN levels (forward() does fpn_hidden_states[:-1])
|
| 104 |
+
print("\n[3/3] Running decoder ...")
|
| 105 |
+
t0 = time.time()
|
| 106 |
+
# Build feed dict dynamically — the legacy tracer may have dropped unused inputs
|
| 107 |
+
candidate_inputs = {
|
| 108 |
+
"fpn_hidden_state_0": fpn_h[0],
|
| 109 |
+
"fpn_hidden_state_1": fpn_h[1],
|
| 110 |
+
"fpn_hidden_state_2": fpn_h[2],
|
| 111 |
+
"fpn_position_encoding_0": fpn_p[0],
|
| 112 |
+
"fpn_position_encoding_1": fpn_p[1],
|
| 113 |
+
"fpn_position_encoding_2": fpn_p[2],
|
| 114 |
+
"text_features": text_features,
|
| 115 |
+
"attention_mask": attention_mask,
|
| 116 |
+
}
|
| 117 |
+
expected_input_names = {inp.name for inp in decoder_sess.get_inputs()}
|
| 118 |
+
feed = {k: v for k, v in candidate_inputs.items() if k in expected_input_names}
|
| 119 |
+
missing = expected_input_names - feed.keys()
|
| 120 |
+
if missing:
|
| 121 |
+
raise RuntimeError(f"Decoder expects inputs we didn't provide: {missing}")
|
| 122 |
+
dropped = candidate_inputs.keys() - feed.keys()
|
| 123 |
+
if dropped:
|
| 124 |
+
print(f" (note: tracer optimized out unused inputs: {sorted(dropped)})")
|
| 125 |
+
d_out = decoder_sess.run(None, feed)
|
| 126 |
+
pred_masks, pred_boxes, pred_logits = d_out
|
| 127 |
+
print(f" Done in {time.time() - t0:.1f}s")
|
| 128 |
+
print(f" pred_masks: shape={pred_masks.shape} dtype={pred_masks.dtype}")
|
| 129 |
+
print(f" pred_boxes: shape={pred_boxes.shape} dtype={pred_boxes.dtype}")
|
| 130 |
+
print(f" pred_logits: shape={pred_logits.shape} dtype={pred_logits.dtype}")
|
| 131 |
+
|
| 132 |
+
# --- Inspect top detections -----------------------------------------------
|
| 133 |
+
# Apply sigmoid to logits to get scores in [0, 1]
|
| 134 |
+
scores = 1.0 / (1.0 + np.exp(-pred_logits)) # sigmoid
|
| 135 |
+
scores_b0 = scores[0]
|
| 136 |
+
top_k = 10
|
| 137 |
+
top_idx = np.argsort(-scores_b0)[:top_k]
|
| 138 |
+
|
| 139 |
+
print(f"\nTop {top_k} detections by score:")
|
| 140 |
+
print(f" {'idx':>4} {'score':>7} {'box (xyxy normalized)':>30}")
|
| 141 |
+
for i in top_idx:
|
| 142 |
+
x1, y1, x2, y2 = pred_boxes[0, i]
|
| 143 |
+
print(f" {i:>4} {scores_b0[i]:>7.4f} ({x1:.3f}, {y1:.3f}, {x2:.3f}, {y2:.3f})")
|
| 144 |
+
|
| 145 |
+
# How many detections above a reasonable threshold?
|
| 146 |
+
for thresh in (0.5, 0.3, 0.1, 0.05):
|
| 147 |
+
n_above = (scores_b0 > thresh).sum()
|
| 148 |
+
print(f" detections with score > {thresh}: {n_above}")
|
| 149 |
+
|
| 150 |
+
print("\n✅ End-to-end ONNX pipeline ran without error.")
|
| 151 |
+
print(" If scores look reasonable for your test image, we're shipping.")
|
| 152 |
+
|
| 153 |
+
|
| 154 |
+
if __name__ == "__main__":
|
| 155 |
+
main()
|