SAM 3 ONNX TensorRT 导出

1. SAM 3 介绍

SAM 3(Segment Anything Model 3)由 Meta AI2025 年 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 骨干特征金字塔 neckSam3DualViTDetNeck)组成,在检测器和跟踪器之间共享。
  • 采用流式(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-encodertext-encoderdecoder 三个 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,因其轻量且受益于动态形状。

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐