litert-community/SAM2.1-Hiera-Tiny-Image-Encoder
SAM 2.1 (Hiera-Tiny) image encoder — LiteRT GPU
On-device LiteRT / TFLite conversion of the image encoder of **SAM 2.1 Hiera-Tiny** (Meta, Apache-2.0), running fully on the mobile GPU via the LiteRT CompiledModel API (ML Drift / LITERT_CL delegate). The whole graph is GPU-resident — no CPU/XNNPACK fallback ops.
This is the heavy backbone of the Segment Anything 2 image path: it turns an RGB image into the multi-scale feature pyramid that a (small) prompt-encoder + mask-decoder then query per click/box.
Preprocessing (must match)
resize to 1024x1024 (bilinear) -> x/255 -> (x - mean) / std
mean = [0.485, 0.456, 0.406], std = [0.229, 0.224, 0.225] # ImageNet, RGB, NCHWGPU-clean conversion (what was re-authored)
Converted with litert-torch. SAM 2's Hiera encoder is not GPU-clean out of the box; these exact, weights-faithful rewrites were applied (model-side only — no converter patch):
- `window_partition` / `window_unpartition`: the 6-D
view+permutewindow reshape rejected by the GPU delegate (>4-D) is re-expressed as a sequence of ≤4-Dreshape/transposeops (numerically exact, verified vs the original). - `Sam2MultiScaleAttention`: the 5-D fused-QKV reshape is decomposed into separate q/k/v, and attention runs as a 3-D batched SDPA (
[B*heads, N, d]). A 4-D SDPA makes the delegate emit a[C,C]->[nW,ws,C,C]BROADCAST_TOon every windowed block; the 3-D form removes all 9. - Windowed positional embedding: the bicubic-interpolate + tile of the constant
pos_embedis baked to a buffer (add only) — removes a runtime interpolate of a constant. - Neck: the (constant, shape-only) sine FPN position encodings are dropped from the graph (compute them host-side) — removes the remaining
BROADCAST_TOops. - Overflow-safe LayerNorm (scale-before-square) as an fp16 safety margin for the deep stages.
Net: banned ops = NONE, >4-D tensors = 0, full GPU residency.
Fidelity (honest)
Eager re-authoring is numerically exact (cos = 1.000, mae = 0). On-device GPU output vs the CPU reference, per FPN level:
The deepest 64×64 feature drifts slightly on the GPU. This is not LayerNorm overflow (scale-before-square LayerNorm doesn't change it, and the CPU fp16 model matches PyTorch fp32 at corr 0.999999) — it is the mobile GPU computing the deep-stage global attention (64×64 = 4096 tokens) in true fp16, where the CPU path upcasts to fp32. The high-resolution features that carry mask boundaries are near-exact, so mask quality is preserved in practice.
Minimal usage
Android (Kotlin, CompiledModel GPU)
val model = CompiledModel.create(context.assets, "sam2_tiny_image_encoder_fp16.tflite",
CompiledModel.Options(Accelerator.GPU), null)
val inputs = model.createInputBuffers()
val outputs = model.createOutputBuffers()
inputs[0].writeFloat(chw) // [1,3,1024,1024] ImageNet-normalized, NCHW
model.run(inputs, outputs)
// FPN maps: [1,256,256,256], [1,256,128,128], [1,256,64,64] -> SAM 2 prompt/mask decoderPython (desktop verification)
MEAN = np.array([0.485, 0.456, 0.406], np.float32)
STD = np.array([0.229, 0.224, 0.225], np.float32)
import numpy as np
from PIL import Image
from ai_edge_litert.interpreter import Interpreter
img = Image.open("photo.jpg").convert("RGB").resize((1024, 1024))
x = ((np.asarray(img, np.float32) / 255 - MEAN) / STD).transpose(2, 0, 1)[None]
it = Interpreter(model_path="sam2_tiny_image_encoder_fp16.tflite"); it.allocate_tensors()
it.set_tensor(it.get_input_details()[0]["index"], x); it.invoke()
o = {tuple(d["shape"]): it.get_tensor(d["index"]) for d in it.get_output_details()}
fpn0, fpn1, fpn2 = o[(1,256,256,256)], o[(1,256,128,128)], o[(1,256,64,64)]
# feed to the SAM 2.1 Hiera-Tiny mask decoder (companion repo) for tap-to-segment;
# the v2 file emits decoder-ready features directly (see the variant note below)Training data & PII
SAM 2 was trained by Meta on SA-1B (licensed photos) and SA-V (licensed videos) with model-in-the-loop mask annotation. No new training was performed for this conversion — it is a weights-faithful format change of the public facebook/sam2.1-hiera-tiny checkpoint. Because the source data is real-world imagery, it may incidentally contain people, faces, vehicles, signage and other PII; no PII was deliberately collected and this conversion adds none. Apply your own content/PII filtering as appropriate. See the SAM 2 release and paper for full dataset details.
Performance
Measured on a Pixel 8a (Tensor G3, Android 16) with the standard TFLite `benchmark_model` tool — 10 warm-up runs then 50 timed runs, reported as the tool's mean.
Any on-device figure recorded when this model shipped came from a different runtime. It was taken through LiteRT's own CompiledModel accelerator (logcat reports it as LITERT_CL), which is the path the Kotlin sample app and the LiteRT API use, and it appears elsewhere on this card. The rows above are the classic TFLite OpenCL delegate, measured with a tool anyone can download and re-run. The two are not comparable, so read the rows above as a reproducible floor rather than as this model's speed on LiteRT.
Pixel 8a — LiteRT CompiledModel
The GPU takes the whole graph through LiteRT's own accelerator — Replacing 862 out of 862 node(s) with delegate (LITERT_CL) for the base file, 867 / 867 for the decoder-ready variant. That is the path the Kotlin snippet above uses, and it is the path the classic delegate in the previous section could not run at all on this phone.
Measured on a Pixel 8a (Tensor G3, Android 16) with LiteRT CompiledModel 2.2.0, one accelerator per process, 5 warm-up runs then N=50 timed runs, median reported. Every row above held thermal status NONE throughout.
On this phone the GPU is 18x faster than the CPU for the decoder-ready file (576.7 ms against 10545 ms). A Galaxy S26 runs that same file at 178.0 ms on its GPU (see below), about a 3x device gap on the same code path.
There is no CPU row for `sam2_tiny_image_encoder_fp16.tflite`, and that absence is itself the result. Four attempts each ran the phone out of its thermal window — N=50 at roughly 11 s a call is ten minutes of unbroken CPU load — and this sweep only accepts a run that begins and ends at thermal status NONE. Read it as: this graph is not something to put on the CPU of a mid-range phone.
The CPU row here and the benchmark_model CPU rows in the previous section come from different runtimes and sit about 2x apart. Neither is wrong; read each against its own tool.
Snapdragon NPU (Hexagon)
sam2_tiny_image_encoder_fp16.tflite— the GPU is faster: 208.0 ms against 280.4 ms on the NPU, a factor of 1.35. The NPU still loads 4.99x faster (457 ms against 2278 ms).sam2_tiny_image_encoder_v2_fp16.tflite— the GPU is faster: 178.0 ms against 260.2 ms on the NPU, a factor of 1.46. The NPU still loads 5.12x faster (438 ms against 2241 ms).
Measured on a Samsung Galaxy S26 (Snapdragon 8 Elite Gen 5 / SM8850, Hexagon v81, Android 16) with LiteRT CompiledModel 2.2.0, one accelerator per process, 5 warm-up runs then N=50 timed runs, median reported. Every run held thermal status NONE throughout. Headroom 0.75–0.83, where 1.0 is the throttling threshold.
The NPU rows ran the published file unchanged. LiteRT compiled it for the Hexagon on the device at first load. Those first compiles took 18.1 min to 18.3 min here. The load column above is the cached load every later run pays. Recipe and the runtime libraries it needs: NPU guide.
GPU wiring: GPU guide.
License
Apache-2.0, inherited from the upstream SAM 2.1. This is a format conversion; all credit to the original authors (Meta AI).
Variant: decoder-ready (sam2_tiny_image_encoder_v2_fp16.tflite)
A second file in this repo, sam2_tiny_image_encoder_v2_fp16.tflite, additionally folds the SAM 2 mask decoder's conv_s0 (256→32) / conv_s1 (256→64) projections and the no_memory embedding into the graph, so it directly emits decoder-ready features: image_embeddings [1,256,64,64], feat_s1 [1,64,128,128], feat_s0 [1,32,256,256]. Pair it with the **SAM 2.1 Hiera-Tiny mask decoder** for promptable "tap to segment" (see the LiteRT interactive_segmentation sample). Same GPU-clean re-authoring and fidelity as the base encoder above; FP16, ~80 MB, full LITERT_CL residency (867/867).
