danilobukvic commited on
Commit
3a8f269
·
verified ·
1 Parent(s): 07fd415

Upload 6 files

Browse files
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()