SAM2转ONNX部署实战:图像编码器与掩码解码器的推理方案
发布时间:2026/9/27 23:38:42来源:尧图网络
简介面向需要在ONNX Runtime环境中运行Segment Anything 2SAM2模型的深度学习开发者这份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 Runtimex86 / NVIDIA GPU官方导出的算子基本全覆盖INT8 静态、动态量化最省心服务端首选TensorRTNVIDIA 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 python3.10 -y conda activate sam2_onnx pip install torch2.1.2 torchvision --index-url https://download.pytorch.org/whl/cu118 pip install onnx onnxruntime1.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 # 固定输入尺寸 1024batch 设置为可动态变化 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_version17, input_names[input_image], output_names[image_embedding], dynamic_axes{ input_image: {0: batch}, image_embedding: {0: batch}, }, do_constant_foldingTrue, verboseFalse, ) onnx.checker.check_model(sam2_image_encoder.onnx) print(image encoder export done)逻辑说明build_sam2只负责把 checkpoint 加载成模型对象真正导出的是model.image_encoder这一步很关键别顺手把整个 model 导出否则输入接口会牵扯到 prompt 参数导出阶段就会报错。opset_version17是我试下来对 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), boxesNone, masksmask_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_embeddingsimage_embedding, image_peimage_pe, sparse_prompt_embeddingssparse_embeddings, dense_prompt_embeddingsdense_embeddings, multimask_outputTrue, ) 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, dtypenp.float32) emb img_sess.run(None, {input_image: fake_image})[0] fake_coords np.array([[[0.5, 0.5], [0.2, 0.8]]], dtypenp.float32) fake_labels np.array([[1, 0]], dtypenp.float32) fake_mask np.zeros((1, 1, 256, 256), dtypenp.float32) fake_has_mask np.array([0], dtypenp.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), interpolationcv2.INTER_LINEAR) img img.astype(np.float32) / 255.0 mean np.array([0.485, 0.456, 0.406], dtypenp.float32) std np.array([0.229, 0.224, 0.225], dtypenp.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_maskNone): # coords/labels 归一化到 [0,1] 范围 has_mask np.array([1 if prior_mask is not None else 0], dtypenp.float32) if prior_mask is None: prior_mask np.zeros((1, 1, 256, 256), dtypenp.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), interpolationcv2.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 维度。结果推理时换成batch2直接报维度错误或者换成 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]] ], dtypenp.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_arenaFalse并限制图优化级别二是把 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 量级内再进入下一阶段。希望这些经验能帮你少踩几个坑。本文还有配套的精品资源点击获取
网站建设高端定制企业官网