SAM 3 — LiteRT (on-device, GPU)

SAM 3 on-device: "person"

SAM 3 (Meta, facebook/sam3) running fully on-device with the LiteRT CompiledModel API: open-vocabulary, text-prompted detection + instance segmentation. Type "wheel", "paper bag", "person" — any phrase — and get every match's box, score, and 288×288 instance mask. No server, no cloud.

  • ViT-L/14 trunk @1008² + tri-neck → CLIP-L text encodertext-conditioned DETR head (200 queries, presence token) — ~830 M params.
  • Verified on a Pixel 8a (vision GPU 9.2 s / text CPU 0.5 s / head GPU 1.4 s; re-prompt 1.9 s; kept-set equal to PyTorch fp32, mask IoU ≥ 0.98) and on Apple Metal (exact with enforce_f32).
  • Every graph re-authoring is exact (corr 1.0 vs PyTorch) — no approximations.

SAM 3 on-device: "car"

Files

File Size Role Accelerator
sam3_vision.tflite 930 MB (fp16) image [1,3,1008,1008]fpn288 | fpn144 | fpn72 GPU
sam3_text.tflite 607 MB (fp16) token embeddings [1,32,1024] → text memory [32·256] CPU (see note)
sam3_head.tflite 68 MB (fp16) [fpn×3 | text_mem | pad] → 200 logits + boxes cxcywh + presence + 200×288² mask logits GPU
sam3_token_embed.bin 101 MB fp16 [49408×1024] token-embedding table (host lookup) host
tokenizer/vocab.json, tokenizer/merges.txt 2 MB CLIP byte-level BPE (ctx 32, BOS 49406, EOT 49407, zero-pad) host
tracker/*.tflite 1.0 GB video tracker graphs (shared trunk + memory attention + decoders), see below GPU

Requires LiteRT ≥ 2.2.0 (2.1.5 mis-executes the head graph on Android GPU). Why text runs on CPU: the CLIP-L residual stream reaches |x|≈1.2e3; fp16 GPU execution corrupts some prompt embeddings. CPU is exact and takes ~0.5 s (on Apple Metal you can use enforce_f32 instead).

Usage (Kotlin, CompiledModel)

val vision = CompiledModel.create(visionPath, CompiledModel.Options(Accelerator.GPU), null)
val text = CompiledModel.create(textPath, CompiledModel.Options(Accelerator.CPU), null)
val head = CompiledModel.create(headPath, CompiledModel.Options(Accelerator.GPU), null)

// image -> features (once per image; cache to re-prompt instantly)
visIn[0].writeFloat(preprocess(bitmap))          // 1008x1008, (x/255-0.5)/0.5, NCHW
vision.run(visIn, visOut)
val feats = visOut[0].readFloat()

// prompt -> text memory (host BPE + fp16 table lookup, then the text graph)
val ids = bpe.encode("wheel")                    // [BOS, ..., EOT, 0-pad] x32
textIn[0].writeFloat(lookupEmbeddings(ids))      // [1,32,1024]
text.run(textIn, textOut)

// features + text -> detections
headIn[0].writeFloat(feats + textOut[0].readFloat() + padMask(ids))
head.run(headIn, headOut)
val y = headOut[0].readFloat()
val presence = sigmoid(y[1000])
// query q kept if sigmoid(y[q]) * presence > 0.5; box y[200+4q..], mask y[1001+q*288*288..]

Usage (Python, CompiledModel)

import numpy as np
from ai_edge_litert.compiled_model import CompiledModel
from ai_edge_litert.hardware_accelerator import HardwareAccelerator

vision = CompiledModel.from_file("sam3_vision.tflite", HardwareAccelerator.GPU)
text = CompiledModel.from_file("sam3_text.tflite", HardwareAccelerator.CPU)
head = CompiledModel.from_file("sam3_head.tflite", HardwareAccelerator.GPU)

def run(model, x, n_out):
    ib, ob = model.create_input_buffers(0), model.create_output_buffers(0)
    ib[0].write(np.ascontiguousarray(x, np.float32).ravel())
    model.run_by_index(0, ib, ob)
    return np.array(ob[0].read(n_out, np.float32))

feats = run(vision, image_1008, 256 * (288**2 + 144**2 + 72**2))   # (x/255-0.5)/0.5, NCHW
table = np.fromfile("sam3_token_embed.bin", np.float16).reshape(-1, 1024)
emb = table[token_ids].astype(np.float32)                          # CLIP BPE, ctx 32, 0-pad
mem = run(text, emb[None], 32 * 256)
pad = (np.array(token_ids) == 0).astype(np.float32)
y = run(head, np.concatenate([feats, mem, pad]), 1001 + 200 * 288 * 288)
prob = 1 / (1 + np.exp(-y[:200])) / (1 + np.exp(-y[1000]))
keep = np.where(prob > 0.5)[0]                                     # boxes y[200:1000], masks y[1001:]

Video tracker (stage 2)

tracker/ holds the Object-Multiplex tracker graphs sharing one trunk pass: sam3_vision_tri.tflite (trunk + detector/interactive/propagation necks), trk_memattn_n7.tflite (memory attention, 7 spatial slots + 16 pointer frames), trk_maskdec.tflite (16-object multiplex decoder), trk_memenc.tflite, trk_initdec.tflite. The host state machine (detection↔track association, hotstart, recondition, memory bank + temporal pos-enc) is ported and verified against the official model (ids identical, mask IoU ≥ 0.992) — the executable spec and the Kotlin/Swift ports live in the LiteRT-Models zoo sample (sam3/), together with all conversion scripts.

Conversion notes

Converted with litert-torch; every GPU-compatibility re-authoring is exact: the >4-D ViT attention (silently mis-lowered otherwise) is rebuilt in ≤4-D with the interleaved RoPE baked into the qkv weights; SafeLayerNorm handles the |x|≈300 residual stream; the DETR decoder is batch-first rank-4 end-to-end (rank-3 [1,N,C] fan-outs mis-execute on mobile GPUs); masked softmax uses the delegate-safe form; ConvTranspose necks are zero-stuff + Conv2d. Details and the on-device debugging record ship with the sample app.

License

SAM Materials, © Meta Platforms — redistributed under the SAM License (see LICENSE, provided with these materials as the license requires). Built with SAM.

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

Model tree for mlboydaisuke/SAM3-LiteRT

Base model

facebook/sam3
Finetuned
(22)
this model