简介:面向需要在ONNX Runtime环境中运行Segment Anything 2(SAM2)模型的深度学习开发者,这份Python脚本资源提供了将SAM2导出为ONNX格式并执行图像分割的完整工作流与推理方案。资源包共14个文件,主要包括4个Python源码(负责模型转换与推理逻辑)、4个ONNX模型文件(涵盖编码器与解码器)、4张测试图片以及1个Markdown说明文档,整体约591.77MB。目录采用主模块与核心脚本分离的结构,代码组织清晰,便于理解导出流程并快速迁移到自有项目。已有199人学习浏览,适合具备一定模型部署基础、希望跨框架使用SAM2的工程师与研究人员。通过该脚本,可省去手工搭建转换环境的步骤,直接获得可运行的ONNX版SAM2,并利用ONNX的跨平台兼容性在边缘设备或服务端高效完成实例分割;配合内置的示例图片可快速验证分割效果,说明文档中对转换流程与常见问题作了补充,Python源码也便于按需修改,为后续二次开发和性能优化留足空间。
1. 用于ONNX的SAM2 Python脚本到底解决了什么问题
拿到「ONNX-SAM2-Segment-Anything.zip」这个压缩包,我的第一反应是:官方 SAM2 代码跑起来太吃显存了,一张 1024 分辨率的图推理一次,PyTorch 环境下轻松吃掉 10G+ 显存。而这个标题所指向的脚本方案,本质上是把 Meta 的 SAM2 从 PyTorch 的温室里搬到 ONNX 这个跨平台中间格式里,再用 ONNX Runtime 做推理。它能解决的痛点是三件事:脱离 Python 生态之外的部署环境、更可控的内存占用、以及把同一个模型送到 RKNN、ncnn 等端侧工具链继续压榨的能力。这套 Python 脚本适合两类人——一类是被 GPU 显存和依赖环境卡住的服务端开发者,另一类是准备在嵌入式设备上跑分割但不想从头训模型的算法工程师。动手之前先把一件事想清楚:SAM2 不是一张静态图,它有清晰的模块边界,哪些能导出、哪些必须留在外围,这是整个方案成立的前提。
2. 为什么 SAM2 值得转 ONNX:从模型结构到部署选型
2.1 SAM2 的模型结构:图像编码、提示编码与掩码解码的分工
熟悉 SAM 系列的人都知道,模型被拆成三个相对独立的部分。图像编码器 backbone 用的是 Hiera 结构,一个基于 Vision Transformer 的层级特征提取网络。输入一张 1024x1024 的图,经过一个卷积 stem 和若干带窗口注意力(window attention)的 stage,最终输出的是一个空间分辨率较低、通道数很高的图像特征,SAM2 里通常叫 itm_embedding 或 image_embedding,实际尺寸取决于 checkpoint,常见是 256 通道、64x64 的空间规格。这里要特别注意:Hiera 的窗口注意力在 ONNX 导出时会有不少 reshape 和 transpose 操作,这些算子恰恰是后续最容易出幺蛾子的地方。
提示编码器负责把你点击的点、画的框或者涂的 mask 转成向量。它分两条路:稀疏提示(点、框)走的是位置编码加 MLP,稠密提示(mask)走的是卷积降采样。这个模块输出两组东西,稀疏嵌入和稠密嵌入。掩码解码器拿到图像特征和这两组提示特征,通过一个轻量的 transformer decoder 加上上采样模块,输出低分辨率的 mask 和对应的 IoU 预测分数。理解这个结构的作用在于:导出 ONNX 时不可能把三个模块一次性打包成一个大黑匣子,常见的做法是拆成两个模型文件,image encoder 一个,prompt encoder 加 mask decoder 合并成另一个。这样前向一次只需要把图像编码跑一遍,多轮提示词交互时复用 embedding,不会重复计算卷积层。
2.2 导出的边界:静态子图能出国,循环和状态留在原地
SAM2 与 SAM1 最大的不同是加入了 memory bank 机制,专门服务于视频分割。帧与帧之间会保留一组历史特征,用于指导当前帧的 mask 预测。这组 memory 是不断迭代更新的,实现在模型外部,由 Python 的循环逻辑控制。你可以把当前帧的图像编码和提示编码喂给 mask decoder,但 memory 的读取、追加、淘汰算法没法在 torch.onnx.export 里作为一个固定的计算图导出。原因很简单,ONNX 的图是静态的,循环长度不定、状态集合动态变化,这些恰好是 ONNX 最不擅长表达的。
所以实践里要建立一条明确的边界:图像编码器和掩码解码器做静态导出,memory bank 的读写逻辑用 Python 在推理代码里手写。你在这个压缩包里看到的脚本,大多数也遵循这个拆法——一个导出脚本负责把可静态化子图转成 ONNX,另一个推理脚本负责在运行时维护状态。另一个容易忽略的边界是输入尺寸。SAM2 官方训练用 1024x1024,但你的业务图可能是长条、可能是竖图。ONNX 模型如果只在固定尺寸下导出,长宽比一变,模型不会报错,但分割精度会明显退化。常见做法是保持训练尺寸,输入时做 letterbox 或直接 resize,而不是依赖导出脚本去支持动态宽高。动态 batch 倒是可以留,代价是部分优化算子会被关闭,推理变慢一点点,一般按需取舍。
2.3 ONNX Runtime、RKNN 与 ncnn:三选一还是先 ONNX 再降级
ONNX 不是终点,它更像是模型生态里的一个中转站。服务器端直接用 ONNX Runtime 跑,开 CUDA 和 TensorRT 执行提供者,性能比原始 PyTorch 推理快且稳定得多。端侧场景则常见顺手再走一次模型转换,比如瑞芯微的 RKNN、手机端的 ncnn。规划部署路径时,我用下面这张表来评估:
| 推理后端 | 适用硬件 | 算子兼容性 | 量化支持 | 落地成熟度 |
|---|---|---|---|---|
| ONNX Runtime | x86 / NVIDIA GPU | 官方导出的算子基本全覆盖 | INT8 静态、动态量化 | 最省心,服务端首选 |
| TensorRT | NVIDIA GPU | 部分算子需插件手写 | FP16、INT8 | 性能上限高,容器集成略繁琐 |
| RKNN | 瑞芯微 NPU | 依赖 onnx 转 rknn 的映射表 | INT8 为主 | 端侧可用,算子覆盖率看版本 |
| ncnn | 手机 CPU / GPU | 转前建议先做 ONNX 简化 | FP16、INT8 | 移动端生态成熟,需逐层检查 |
我的习惯是无论最终跑在哪,都先拿 ONNX 版本做基准测试。原因很朴素:ONNX 模型的可调试性最好,onnxruntime 的日志能精确告诉你哪个算子崩了,而 RKNN、ncnn 工具链报错经常是模糊的。先确认 PyTorch 模型转换后精度不掉,再往端侧转,排查范围会窄很多。这里也劝一句,网上那种「onnx 转 rknn 在线网站」很省事,但生产环境别依赖在线服务,模型文件和中间文件都有泄露风险,本地装工具链也就十几分钟。
3. 跑通 ONNX-SAM2 脚本:环境组合与最小导出步骤
3.1 环境准备:Python、PyTorch 与 ONNX Runtime 的版本搭配
拿到脚本的第一步不是改代码,是把环境复原到作者当初的开发状态。我用 conda 新建虚拟环境,Python 选 3.10,这个版本对 PyTorch 2.x 和 onnxruntime 的兼容性最平衡。PyTorch 版本我固定在 2.1 以上,因为 SAM2 官方代码用到了 torch.nn.functional.scaled_dot_product_attention,这个算子在高版本导出 ONNX 时才有稳定的映射。onnxruntime 选 1.17 之后的版本,之前版本的 CPU 执行提供者对部分 transformer 算子的实现有性能缺陷。
conda create -n sam2_onnx python=3.10 -y conda activate sam2_onnx pip install torch==2.1.2 torchvision --index-url https://download.pytorch.org/whl/cu118 pip install onnx onnxruntime==1.17.1 opencv-python numpy参数说明:torch 和 torchvision 的版本必须配套,指定 cu118 的 index-url 是为了让 CUDA 11.8 工具链一致,如果你机器是纯 CPU 推理,把这两行换成pip install torch torchvision就行。onnx 是转换和检查用的,onnxruntime 是推理用的,两个包别混成一团。opencv 用来读图和做预处理,numpy 负责张量搬运。装完后用一行命令验证:python -c "import torch, onnx, onnxruntime; print(torch.__version__, onnx.__version__, onnxruntime.__version__)",版本正确再往下走。
3.2 导出脚本一:把图像编码器单独导出
SAM2 的图像编码器是模型里计算量最大的部分,也是相对容易导出的部分。它没有循环、没有条件控制,只要把输入输出 shape 定清楚就行。下面这段是核心导出代码的骨架:
import torch import onnx from sam2.build_sam import build_sam2 # 以官方源码包为例 # 加载 checkpoint 并切到 eval 模式 model = build_sam2("sam2_hiera_large", "checkpoints/sam2_hiera_large.pt") model.eval() # 只拿图像编码器,避免把提示编码和 mask 解码一起带进图里 image_encoder = model.image_encoder # 固定输入尺寸 1024,batch 设置为可动态变化 dummy_input = torch.randn(1, 3, 1024, 1024).float() with torch.no_grad(): torch.onnx.export( image_encoder, dummy_input, "sam2_image_encoder.onnx", opset_version=17, input_names=["input_image"], output_names=["image_embedding"], dynamic_axes={ "input_image": {0: "batch"}, "image_embedding": {0: "batch"}, }, do_constant_folding=True, verbose=False, ) onnx.checker.check_model("sam2_image_encoder.onnx") print("image encoder export done")逻辑说明:build_sam2只负责把 checkpoint 加载成模型对象,真正导出的是model.image_encoder,这一步很关键,别顺手把整个 model 导出,否则输入接口会牵扯到 prompt 参数,导出阶段就会报错。opset_version=17是我试下来对 Hiera 窗口注意力最稳妥的版本,低于 16 会出现 ScaledDotProductAttention 映射缺失。dynamic_axes刻意只放开 batch 维度,宽高锁死 1024,这是有意为之的取舍——动态宽高会让 onnxruntime 放弃大量图优化,速度反而亏了。
3.3 导出脚本二:提示编码器与掩码解码器合并导出
这部分比图像编码器麻烦,因为输入输出里既有序列长度的动态维度,又有点位坐标的归一化问题。SAM2 官方代码在 mask decoder 前会先做一次 prompt 编码,导出的模型需要同时接收原始提示词参数,并缓存解码器要用的位置编码。我一般用包装模块来做:
import torch from torch import nn class PromptMaskExportWrapper(nn.Module): def __init__(self, prompt_encoder, mask_decoder): super().__init__() self.prompt_encoder = prompt_encoder self.mask_decoder = mask_decoder def forward( self, image_embedding: torch.Tensor, point_coords: torch.Tensor, point_labels: torch.Tensor, mask_input: torch.Tensor, has_mask_input: torch.Tensor, ): # 提示编码:稀疏提示走坐标编码,稠密提示按 has_mask 判断是否参与 sparse_embeddings, dense_embeddings = self.prompt_encoder( points=(point_coords, point_labels), boxes=None, masks=mask_input if has_mask_input[0] > 0 else None, ) # 位置编码从 prompt encoder 里拿,解码器直接消费 image_pe = self.prompt_encoder.get_dense_pe() low_res_masks, iou_predictions = self.mask_decoder( image_embeddings=image_embedding, image_pe=image_pe, sparse_prompt_embeddings=sparse_embeddings, dense_prompt_embeddings=dense_embeddings, multimask_output=True, ) return low_res_masks, iou_predictions参数说明:point_coords的形状是[batch, num_points, 2],坐标必须是归一化到 0 到 1 的相对坐标;point_labels是[batch, num_points],1 表示正前景点,0 表示背景点;mask_input是[batch, 1, 256, 256]的先前 mask,用来做精细化迭代;has_mask_input是一个标量张量,告诉模型这轮有没有提供 mask 提示。导出时dynamic_axes要单独指定num_points维度可变化,否则推理时换个点数就会报输入维度不匹配。
3.4 验证导出结果:用 onnx.checker 和 onnxruntime 做一次冒烟测试
导出完成不等于模型能用,我见过太多导出来是一回事、跑起来是另一回事的情况。这一步先用 checker 做结构校验,再用 onnxruntime 加载并和 PyTorch 原始输出做逐像素对比。冒烟测试的代码很简单:
import numpy as np import onnxruntime as ort # 两个会话分别加载 img_sess = ort.InferenceSession("sam2_image_encoder.onnx") dec_sess = ort.InferenceSession("sam2_mask_decoder.onnx") # 构造一张全灰图做形状冒烟 fake_image = np.full((1, 3, 1024, 1024), 128, dtype=np.float32) emb = img_sess.run(None, {"input_image": fake_image})[0] fake_coords = np.array([[[0.5, 0.5], [0.2, 0.8]]], dtype=np.float32) fake_labels = np.array([[1, 0]], dtype=np.float32) fake_mask = np.zeros((1, 1, 256, 256), dtype=np.float32) fake_has_mask = np.array([0], dtype=np.float32) low_res, iou = dec_sess.run(None, { "image_embedding": emb, "point_coords": fake_coords, "point_labels": fake_labels, "mask_input": fake_mask, "has_mask_input": fake_has_mask, }) print("low_res shape:", low_res.shape, "iou shape:", iou.shape)这段验证的意义在于提前暴露两个经典错误:一是图像编码器输出和 mask 解码器期望的 embedding 通道数不一致,报错信息会直接告诉你维度对不上;二是动态轴的坐标没有正确归一化,导致 mask 解码器输出的掩码整体偏移。冒烟测试通过后,再拿着真实图片与 PyTorch 原始输出对比,允许的误差一般在 1e-3 量级,超过这个量级就得回头查导出的模型哪里被算错了。
4. 用 ONNX Runtime 跑通 SAM2 推理:核心代码与参数设置
4.1 数据预处理:归一化、尺寸、坐标系的坑
很多人在推理阶段翻车,问题全出在预处理和原始 PyTorch 代码不一致。SAM2 的训练数据是 ImageNet 归一化方式:像素值除以 255 后,按通道减去均值 [0.485, 0.456, 0.406],再除以方差 [0.229, 0.224, 0.225]。注意顺序是减均值再除方差,不是先减再缩放到 -1 到 1。图像通道顺序是 RGB,如果你用 OpenCV 读图,默认是 BGR,必须转一遍,否则模型输出的掩码会像色偏照片一样,边缘是对的,内部类别全错。
import cv2 import numpy as np def preprocess_image(image_path: str, target_size: int = 1024): img = cv2.imread(image_path) img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img = cv2.resize(img, (target_size, target_size), interpolation=cv2.INTER_LINEAR) img = img.astype(np.float32) / 255.0 mean = np.array([0.485, 0.456, 0.406], dtype=np.float32) std = np.array([0.229, 0.224, 0.225], dtype=np.float32) img = (img - mean) / std # 调整为 NCHW 格式 img = np.transpose(img, (2, 0, 1))[None, ...] return img参数说明:target_size和导出时锁定的 1024 保持一致,不要随输入尺寸变化;resize 用线性插值,最近邻会让边缘出现锯齿。整张图的等比缩放问题我建议直接交给 call 方处理,业务方如果不想拉伸,可以先做 letterbox 再传进来。这段预处理要和导出模型的输入命名对齐,onnxruntime 的run方法里,字典的 key 必须和导出时的input_names一致,拼错一个字母就是 KeyError。
4.2 推理主流程:embedding 提取与 mask 生成
推理的核心思路是:图像只编码一次,提示词可以多次变化。第一步把 image encoder 的输出存起来,第二步每次点击、画框后只跑 decoder。这套流程和交互式分割产品的交互逻辑天然匹配,用户点一下,你只需要跑一次轻量的 decoder,而不是重算整个 backbone。下面是完整推理函数:
import onnxruntime as ort class SAM2ONNXInference: def __init__(self, encoder_path, decoder_path): # providers 顺序调整:优先走 CUDA,然后再回退 CPU self.enc_sess = ort.InferenceSession( encoder_path, providers=["CUDAExecutionProvider", "CPUExecutionProvider"], ) self.dec_sess = ort.InferenceSession( decoder_path, providers=["CUDAExecutionProvider", "CPUExecutionProvider"], ) self.image_embedding = None def set_image(self, image_tensor: np.ndarray): # 只跑一次图像编码器 self.image_embedding = self.enc_sess.run( None, {"input_image": image_tensor} )[0] def predict(self, coords, labels, prior_mask=None): # coords/labels 归一化到 [0,1] 范围 has_mask = np.array([1 if prior_mask is not None else 0], dtype=np.float32) if prior_mask is None: prior_mask = np.zeros((1, 1, 256, 256), dtype=np.float32) outputs = self.dec_sess.run( None, { "image_embedding": self.image_embedding, "point_coords": coords.astype(np.float32), "point_labels": labels.astype(np.float32), "mask_input": prior_mask.astype(np.float32), "has_mask_input": has_mask, }, ) low_res_masks, iou_scores = outputs[0], outputs[1] return low_res_masks, iou_scores逻辑说明:providers的列表顺序有优先级含义,onnxruntime 会按顺序尝试加载执行提供者,CUDA 不可用时就自动落到 CPU。这个机制在服务端很实用,不用在代码里写if device == "cuda"之类的分支。set_image和predict的分离,是这套部署方案最关键的优化点——一次图像编码、多次掩码解码,显存占用和延迟都是可控的。
4.3 后处理:从 low_res 到原图分辨率,以及多 mask 选择
mask decoder 输出的low_res_masks形状是[batch, num_mask_candidates, 256, 256],其中num_mask_candidates取决于导出时的multimask_output,设置 True 时通常输出 4 个候选。这不是 4 个类别,而是同一个目标点的 4 种不同分割粒度,最终选哪个由iou_scores说了算。取最高分的候选 mask,再上采样回原图分辨率,才算完成一次完整推理。
import cv2 def postprocess_mask(low_res_masks, iou_scores, original_shape): # 按 IoU 分数挑出最佳候选 best_idx = int(np.argmax(iou_scores[0])) best_mask = low_res_masks[0, best_idx] # shape: [256, 256] # sigmoid 激活后转到 0-255 best_mask = 1 / (1 + np.exp(-best_mask)) best_mask = (best_mask * 255).astype(np.uint8) # 双线性上采样回原图尺寸 h, w = original_shape[:2] mask_resized = cv2.resize( best_mask, (w, h), interpolation=cv2.INTER_LINEAR ) return mask_resized > 127 # 二值化阈值可调参数说明:iou_scores是模型对每个候选 mask 与真实目标一致性的自信度估计,直接取最大值的做法在绝大多数情况下是对的。但边界贴合要求高的场景,也可以把 4 个候选都返回给前端,让用户手动选。二值化阈值 127 是经验值,实际业务里 150 到 180 更常用,因为低置信度区域会在边缘产生浅灰色过渡带,阈值略高一点能滤掉毛刺。
5. SAM2 转 ONNX 的避坑手册:5 条被反复问到的翻车记录
5.1 现象:导出的模型输入尺寸写死,batch 和 mask 数量全被锁住
有人导出后把dummy_input设置成(1, 3, 1024, 1024),但dynamic_axes里漏写了 batch 维度。结果推理时换成batch=2直接报维度错误,或者换成 512 的分辨率输入,模型不报错但输出 mask 全是噪声。原因基本是只抄了导出代码,没理解dynamic_axes的语义。解决方法是回到导出脚本,确认三处动态维度都设了:batch、点序列长度、以及低分辨率 mask 的 mask 数量维度。检查方法很简单,用onnx.load之后print(model.graph.input),看到维度里出现batch、num_points这些符号名而不是固定数字,才算真正生效。
5.2 现象:点击一个点后分割出来的 mask 完全错位
这不是模型坏了,是坐标归一化方式不对。用户点击的坐标是原始像素坐标,比如一张 1920x1080 的图,你点在 (960, 540)。但 SAM2 期望的输入是除以输入尺寸后的相对坐标,也就是 (0.5, 0.5),而且这个归一化必须在 resize 之后按新尺寸算。如果你直接把原图像素坐标传进模型,mask 会定位到一个小得不成比例的角落区域。原因是推理代码里用了预处理后的图,却忘了把坐标也同步缩放。解决方案是在predict函数里多传一个缩放系数:
# 假设原图经过 resize 到 1024x1024 scale_x = 1024 / original_width scale_y = 1024 / original_height normalized_coords = np.array([ [[pt_x * scale_x / 1024, pt_y * scale_y / 1024]] ], dtype=np.float32)5.3 现象:GPU 上正常、CPU 上跑出的 mask 边缘毛刺明显
同一份 ONNX 文件,在 CUDA 执行提供者下结果完美,切到 CPU 后掩码边缘出现大量孤立噪点。排查时先确认是精度问题还是算子实现差异。我用ort.get_available_providers()查看可用的执行提供者,再把enable_cpu_mem_arena关掉,用纯 CPU 的默认内核跑一遍。常见原因是 LayerNorm 在 CPU 内核的求均值实现上用了不同的归约顺序,累积误差被放大。解决方式有两个:一是导出时把opset_version提到 17,新算子的 CPU 实现更成熟;二是在初始化 session 时设置session.set_optimization_level(ort.ORT_ENABLE_BASIC),关闭高级图优化,排除融合算子引发的数值偏差。
5.4 现象:int8 量化后 mask 糊成一团,目标边界完全丢失
很多人在部署时图省事,直接对导出的 ONNX 做全模型静态 int8 量化,结果精度崩得不能看。原因是 SAM2 这类分割模型对数值范围非常敏感,尤其是 mask decoder 里的上采样卷积,低比特量化会让跨层的信息传递损失过大。我试过的相对稳的方案是混合精度:用 Quantization 工具先跑一遍离线校准,统计每层的激活分布,然后只量化图像编码器里计算量最大的 attention 模块的 matmul 算子,mask decoder 整体保持 float 精度。这样模型体积能压缩一半,推理速度提升明显,但 mask 的边界精度几乎不掉。具体实现可以用 onnxruntime 的quantize_static配合nodes_to_quantize参数,手动指定算子集。
5.5 现象:显存占用依旧爆炸,两个 session 叠加吃满显卡
导出了 ONNX 并用 ONNX Runtime 跑,显存还是居高不下,原因是同时加载了 image encoder 和 mask decoder 两个 session,各自开启了独立的 CUDA 上下文。如果业务是多路并发,每个进程都来一份,显卡很快被榨干。解决思路有三条,按性价比排序:一是给InferenceSession传入sess_options,设置enable_cpu_mem_arena=False并限制图优化级别;二是把 image encoder 和 mask decoder 分进程部署,image encoder 常驻,mask decoder 按请求拉起;三是如果场景支持批处理,把多路图像的编码合并进同一个set_image调用,让 onnxruntime 内部复用显存。
6. 进阶用法:视频分割的 embedding 复用、量化提速与精度验证
视频分割场景是最能体现这套 ONNX 方案价值的地方。SAM2 的记忆机制在 ONNX 里没法导出,但 image encoder 的复用逻辑不变:同一帧画面只需要在首帧做一次图像编码,后续帧直接用第一帧的 embedding 配合新增提示点做 mask 更新。实际项目里我会缓存上一帧的分割结果,叠加进当前帧的 mask_input 参数中,形成粗略的时序一致性。这样即使不用完整 memory bank,连续帧的 mask 抖动也会明显减少。
量化的推荐路径是先做动态量化看精度,再做静态量化看加速。动态量化只压缩权重,不改变激活计算,对 SAM2 这类模型比较友好,体积能减一半但推理提速有限。静态量化需要准备 100 到 200 张覆盖阴影、强光、低对比度的真实图像做校准集,校准数据太少会出现其他环境精度骤降。验证量化模型是否可用的办法不是只盯着 mIoU,关键看边缘像素的连续性:拉一条穿过目标边界的横线,统计像素从 0 到 1 的过渡带宽度,量化后过渡带变宽超过 2 个像素,就得回退图层级精度检查。
精度验证的最后一环是在你自己业务数据上跑满 50 张图,覆盖不同光照和遮挡情况,逐张对比 PyTorch 原始模型和 ONNX 模型的输出,记录最大像素差和 mask 面积差。我最早跑通这套脚本后的教训是:不要拿一张标准测试图验证就上线,图像编码器在暗光场景下输出特征的微小差异,会被 mask decoder 放大成明显的边界偏移。我现在的习惯是验证脚本里固定了一个种子、一套对比代码,每次环境变化后跑一遍回归,确认输出差异稳定在 1e-3 量级内再进入下一阶段。希望这些经验能帮你少踩几个坑。
本文还有配套的精品资源,点击获取