SAM3 ONNX TensorRT 导出
SAM 3 ONNX TensorRT 导出
1. SAM 3 介绍
SAM 3(Segment Anything Model 3)由 Meta AI 于 2025 年 11 月 19 日发布,是 Segment Anything 系列的第三代模型。它最核心的突破是把任务从前两代的 PVS(Promptable Visual Segmentation,可提示视觉分割)升级为 PCS(Promptable Concept Segmentation,可提示概念分割)。
关键能力
-
开放词汇概念分割:与 SAM 1/2 只能根据点/框/掩码分割单个对象不同,SAM 3 能检测、分割并跟踪由文本提示、图像示例(exemplar)或两者共同指定的某个视觉概念的所有实例。例如输入
"striped cat"(条纹猫)或"red apple",模型会找出图像/视频中所有匹配的目标。 -
文本 + 视觉混合提示:为聚焦于识别原子级视觉概念,文本被约束为简单名词短语(NP),如
"red apple"或"striped cat";SAM 3 不擅长长指代表达或需要推理的查询,但可与多模态大模型(MLLM)结合处理复杂语言提示。 -
图像与视频统一:在视频中,检测器在每帧发现新的概念实例,跟踪器从前序帧传播 masklets(时序对象片段),再通过匹配与更新阶段融合,保证遮挡或重现时的一致性。
-
规模与性能:模型约 8.4 亿参数(约 3.4 GB),面向 GPU 推理;在 NVIDIA H200 上 100+ 目标约 30ms/图,在 16GB 显存 GPU 上可运行典型负载。
2. SAM 3 结构介绍
SAM 3 采用解耦的检测器–跟踪器(detector–tracker)双分支架构,共享一个视觉编码器。848M 参数被组织成一个检测器和一个跟踪器,二者共享视觉编码器。
(1) 视觉骨干 / Perception Encoder(PE)
SAM 3 的文本与图像编码器来自 Meta Perception Encoder(2025 年 4 月开源);相比以往编码器选择带来显著性能提升。
- 视觉编码器由 ViT 骨干和特征金字塔 neck(
Sam3DualViTDetNeck)组成,在检测器和跟踪器之间共享。 - 采用流式(streaming)设计,Perception Encoder 对每个输入视频帧只编码一次,得到的 “unconditioned tokens” 作为整个系统的唯一特征来源,避免在检测与跟踪间切换时重复编码。
(2) 检测器(Sam3Image,DETR-based)
检测器是基于 DETR 的模型,以文本、几何和图像示例为条件。
其关键创新是 Presence Head(存在头):一个学习得到的全局 token,用于预测目标概念是否存在于图像/帧中,通过把识别(recognition)与定位(localization)分离来改善检测——存在头全局预测概念是否存在,而 proposal queries 只专注于定位,避免目标冲突。这在用困难负样本短语训练时尤其有效。
(3) 跟踪器(Sam3TrackerPredictor,SAM 2-style)
跟踪器继承 SAM 2 的 transformer 编码器–解码器架构,支持视频分割与交互式细化。其**记忆库(memory bank)**和记忆编码器以 SAM 2 为基础。
(4) 解耦设计的动机
检测与跟踪解耦是为了避免任务冲突:
- 检测器需要身份无关(identity-agnostic)
- 跟踪器的主要目标是在视频中区分不同身份
3. ONNX 导出 & TensorRT
核心约束:必须拆分,且视频记忆机制无法整体转 TRT
完整模型不能直接 torch.onnx.export()。SAM 3 拥有包含视觉编码器、文本编码器和多个输出头的复杂多模态架构,导出并不像调用 torch.onnx.export() 那样简单。
更重要的是视频跟踪部分的动态性问题:Sam3VideoModel 的记忆机制依赖动态数据结构——记忆库维护变长的历史帧特征列表,记忆注意力引用变长历史信息;而 TensorRT 面向静态计算图,无法原生处理这些动态结构,使得记忆机制的完整 TensorRT 转换不切实际。
推荐拆分方式
典型的 ONNX 子模块拆分(仅检测器部分):
SAM3 TensorRT (Detector Only)
├── vision-encoder.onnx → Vision Encoder (PE/ViT)
├── text-encoder.onnx → Text Encoder (CLIP-style)
└── decoder.onnx → Mask Decoder
该方案仅使用 Sam3Model(检测器部分)的 ONNX 模型,拆分为 vision-encoder、text-encoder、decoder 三个 ONNX,然后用 ByteTrack 做 ID 关联。
| 方案 | 性能 |
|---|---|
| PyTorch 实时 | ~5 FPS |
| TensorRT + ByteTrack | 30+ FPS,并保持持久性 |
整体流程概览
加载模型 → 修复导出兼容性 → 预处理图像/文本 → 分模块导出 ONNX → onnxruntime 验证
导出脚本 如下,该脚本导出的ONNX 模型是具有 mask 部分的
import os
import cv2
import torch
import torch.nn.functional as F
from torchvision.transforms import v2
from sam3.model.data_misc import FindStage
from sam3.model.geometry_encoders import Prompt
import sam3.model.vitdet as vitdet
from sam3.model_builder import build_sam3_image_model
device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"Using device: {device}")
weights_path = "sam3_weights/sam3.pt"
output_dir = "models/onnx_segment"
prompt = "ear"
image_url = "./1.jpg"
resolution = 1008
def use_exportable_mlp_activation():
"""Avoid the fused bf16-only MLP path when tracing/exporting ONNX."""
def exportable_addmm_act(activation, linear, mat1):
x = linear(mat1)
if activation in (F.relu, torch.nn.ReLU):
return F.relu(x)
if activation in (F.gelu, torch.nn.GELU):
return F.gelu(x)
raise ValueError(f"Unexpected activation {activation}")
vitdet.addmm_act = exportable_addmm_act
def use_exportable_rope(module):
"""Use real-valued RoPE math because ONNX cannot export complex tensors."""
for submodule in module.modules():
if (
isinstance(submodule, vitdet.Attention)
and submodule.use_rope
and not submodule.use_ve_rope
):
submodule.use_rope_real = True
if not hasattr(submodule, "freqs_cis_real"):
submodule.register_buffer("freqs_cis_real", submodule.freqs_cis.real)
if not hasattr(submodule, "freqs_cis_imag"):
submodule.register_buffer("freqs_cis_imag", submodule.freqs_cis.imag)
def use_text_only_geometry_prompt(sam3):
"""Skip geometric-prompt pooling branches unused by this text-only export."""
sam3.geometry_encoder.points_pool_project = None
sam3.geometry_encoder.boxes_pool_project = None
def disable_segmentation_head_checkpointing(sam3):
"""Disable activation checkpointing inside the segmentation head for ONNX tracing."""
seg_head = sam3.segmentation_head
if seg_head is None:
return
seg_head.act_ckpt = False
for submodule in seg_head.modules():
if hasattr(submodule, "act_ckpt"):
submodule.act_ckpt = False
def move_decoder_coordinate_cache(decoder, target_device):
if decoder.compilable_cord_cache is not None:
decoder.compilable_cord_cache = tuple(
value.to(target_device) for value in decoder.compilable_cord_cache
)
decoder.coord_cache = {
key: tuple(value.to(target_device) for value in values)
for key, values in decoder.coord_cache.items()
}
use_exportable_mlp_activation()
model = build_sam3_image_model(
checkpoint_path=weights_path,
load_from_HF=False,
device=device,
eval_mode=True,
)
model.float().eval()
use_exportable_rope(model)
use_text_only_geometry_prompt(model)
disable_segmentation_head_checkpointing(model)
model.use_act_checkpoint_seg_head = False
os.makedirs(output_dir, exist_ok=True)
image = cv2.imread(image_url)
image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
transform = v2.Compose(
[
v2.ToDtype(torch.uint8, scale=True),
v2.Resize(size=(resolution, resolution)),
v2.ToDtype(torch.float32, scale=True),
v2.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]),
]
)
pixel_values = transform(v2.functional.to_image(image).to(device)).unsqueeze(0)
class Sam3ImageEncoderONNXWrapper(torch.nn.Module):
def __init__(self, sam3):
super().__init__()
self.backbone = sam3.backbone
def forward(self, pixel_values):
backbone_out = self.backbone.forward_image(pixel_values)
return tuple(backbone_out["backbone_fpn"] + backbone_out["vision_pos_enc"])
class Sam3TextEncoderONNXWrapper(torch.nn.Module):
def __init__(self, sam3):
super().__init__()
self.text_encoder = sam3.backbone.language_backbone
def forward(self, input_ids):
attention_mask = (input_ids != 0).bool()
inputs_embeds = self.text_encoder.encoder.token_embedding(input_ids)
_, text_memory = self.text_encoder.encoder(input_ids)
# Match VETextEncoder.forward output layout used by backbone.forward_text.
language_mask = attention_mask.ne(1)
language_features = self.text_encoder.resizer(text_memory.transpose(0, 1))
language_embeds = inputs_embeds.transpose(0, 1)
return language_features, language_mask, language_embeds
class Sam3DecoderONNXWrapper(torch.nn.Module):
"""Run grounding decoder + segmentation head and return boxes/logits/masks."""
def __init__(self, sam3, num_fpn_features, num_pos_features):
super().__init__()
self.sam3 = sam3
self.num_fpn_features = num_fpn_features
self.num_pos_features = num_pos_features
def forward(self, *inputs):
fpn_end = self.num_fpn_features
pos_end = fpn_end + self.num_pos_features
backbone_fpn = list(inputs[:fpn_end])
vision_pos_enc = list(inputs[fpn_end:pos_end])
language_features, language_mask = inputs[pos_end : pos_end + 2]
export_device = backbone_fpn[0].device
geometric_prompt = Prompt(
box_embeddings=torch.zeros(1, 1, 4, device=export_device),
box_mask=torch.ones(1, 1, device=export_device, dtype=torch.bool),
box_labels=torch.ones(1, 1, device=export_device, dtype=torch.long),
point_embeddings=torch.zeros(1, 1, 2, device=export_device),
point_mask=torch.ones(1, 1, device=export_device, dtype=torch.bool),
point_labels=torch.ones(1, 1, device=export_device, dtype=torch.long),
)
find_stage = FindStage(
img_ids=torch.tensor([0], device=export_device, dtype=torch.long),
text_ids=torch.tensor([0], device=export_device, dtype=torch.long),
input_boxes=None,
input_boxes_mask=None,
input_boxes_label=None,
input_points=None,
input_points_mask=None,
)
backbone_out = {
"backbone_fpn": backbone_fpn,
"vision_pos_enc": vision_pos_enc,
"language_features": language_features,
"language_mask": language_mask,
}
outputs = self.sam3.forward_grounding(
backbone_out=backbone_out,
find_input=find_stage,
geometric_prompt=geometric_prompt,
find_target=None,
)
pred_boxes = outputs["pred_boxes"]
pred_logits = outputs["pred_logits"]
pred_masks = outputs["pred_masks"]
presence_logit_dec = outputs.get("presence_logit_dec", None)
if presence_logit_dec is None:
presence_logit_dec = torch.zeros(
(pred_logits.shape[0], pred_logits.shape[1]),
device=pred_logits.device,
dtype=pred_logits.dtype,
)
return pred_boxes, pred_logits, pred_masks, presence_logit_dec
def export_onnx(
module, args, output_file, input_names, output_names, use_dynamo=False
):
torch.onnx.export(
module,
args,
output_file,
input_names=input_names,
output_names=output_names,
opset_version=17,
external_data=True,
dynamo=use_dynamo,
optimize=False,
do_constant_folding=False,
)
print(f"Exported: {output_file}")
@torch.inference_mode()
def main():
image_encoder = Sam3ImageEncoderONNXWrapper(model).to(device).eval()
text_encoder = Sam3TextEncoderONNXWrapper(model).to(device).eval()
sample_backbone_out = model.backbone.forward_image(pixel_values)
image_features = tuple(
sample_backbone_out["backbone_fpn"] + sample_backbone_out["vision_pos_enc"]
)
num_fpn_features = len(sample_backbone_out["backbone_fpn"])
num_pos_features = len(image_features) - num_fpn_features
fpn_names = [f"backbone_fpn_{idx}" for idx in range(num_fpn_features)]
pos_names = [f"vision_pos_enc_{idx}" for idx in range(num_pos_features)]
tokenizer = model.backbone.language_backbone.tokenizer
context_length = model.backbone.language_backbone.context_length
input_ids = tokenizer([prompt], context_length=context_length).to(device)
language_features, language_mask, language_embeds = text_encoder(input_ids)
print(
"text outputs:",
language_features.shape,
language_mask.shape,
language_embeds.shape,
)
export_onnx(
image_encoder,
(pixel_values,),
os.path.join(output_dir, "sam3_image_encoder.onnx"),
["pixel_values"],
fpn_names + pos_names,
)
export_onnx(
text_encoder,
(input_ids,),
os.path.join(output_dir, "sam3_text_encoder.onnx"),
["input_ids"],
["language_features", "language_mask", "language_embeds"],
)
decoder = Sam3DecoderONNXWrapper(
model,
num_fpn_features=num_fpn_features,
num_pos_features=num_pos_features,
).to("cpu").eval()
move_decoder_coordinate_cache(decoder.sam3.transformer.decoder, "cpu")
decoder_inputs = (
*image_features[:num_fpn_features],
*image_features[num_fpn_features:],
language_features,
language_mask,
)
decoder_inputs = tuple(
value.cpu() if isinstance(value, torch.Tensor) else value
for value in decoder_inputs
)
export_onnx(
decoder,
decoder_inputs,
os.path.join(output_dir, "sam3_decoder.onnx"),
fpn_names + pos_names + ["language_features", "language_mask"],
["pred_boxes", "pred_logits", "pred_masks", "presence_logit_dec"],
use_dynamo=True,
)
try:
import onnxruntime as ort
except ImportError:
return
for filename in (
"sam3_image_encoder.onnx",
"sam3_text_encoder.onnx",
"sam3_decoder.onnx",
):
output_file = os.path.join(output_dir, filename)
print(f"\n{filename}")
try:
session = ort.InferenceSession(
output_file,
providers=["CUDAExecutionProvider", "CPUExecutionProvider"],
)
except Exception as exc:
print(f"onnxruntime load failed: {exc}")
continue
print("input_info:")
for info in session.get_inputs():
print(info.name, info.shape, info.type)
print("output_info:")
for info in session.get_outputs():
print(info.name, info.shape, info.type)
if __name__ == "__main__":
main()
更详细的代码脚本,以及部署的注意事项,在我的下一篇博客 https://blog.csdn.net/dream_of_studies/article/details/161505707?spm=1011.2124.3001.6209
注:本文内容基于作者个人实际应用过程的总结与记录,旨在技术分享与学习交流之用。如内容中涉及任何版权问题或存在争议,欢迎联系作者进行处理或删除。
💡 常见折中:prompt encoder 与 mask decoder 保留在 PyTorch,因其轻量且受益于动态形状。
更多推荐




所有评论(0)