资讯详情

资讯详情

SAM2图像分割ONNX部署实战:从PyTorch导出到推理优化

简介面向算法部署场景这套 Python ONNX 实战项目聚焦 SAM2 图像分割算法提供完整源码与流程教程适合算法工程师、深度学习者以及有模型落地需求的开发者。内容覆盖自动驾驶、医学影像分析、视频监控等应用方向从算法原理到环境配置、模型转换和推理部署均有可操作指引。压缩包共 12 个文件以 Python 脚本为主辅以 txt 说明与 md 文档并含示意图、演示图片等包体约 10.37MB目录简洁便于按模块查阅。已有 395 人学习/下载。通过该套件可掌握 SAM2 原理与典型场景、Python 图像处理和部署基础以及 ONNX 格式转换与硬件适配优化教程还就模型精度保持、部署效率等常见问题提供排错思路适合在实践修改中沉淀可复用的部署经验。1. SAM2 图像分割落地为什么卡在 ONNX先说结论再做选型SAM2 图像分割在开源圈热度很高但真正敢把它推到生产环境的团队并不多。原因不在模型精度而在部署链路PyTorch 里跑得通不代表换到 Python Onnx Runtime 也能跑得通。SAM2 不是那种“输入一张图、输出一张图”的单次前向模型它由图像编码器和掩码解码器组成交互式点提示要分多次进解码器这就让 ONNX 导出和推理脚本比 YOLO 那一类复杂不少。本文适合正在做图像分割算法部署、想用 ONNX 托管 SAM2 的工程师也适合做交互式抠图、标注工具、医学影像辅助判读的团队参考。我会把导出、推理、后处理、量化验证和一线的坑从头讲一遍给出能直接抄的脚本和参数。2. PyTorch 转 ONNXSAM2 导出的三个动态轴与固定输出名2.1 SAM2 能用 ONNX 部署吗先看清 torch 里跑的是什么用 ONNX 部署 SAM2 前先理解它在 PyTorch 里的运行时结构。官方的 SAM2ImagePredictor 内部是两段式image encoder 把图片编码成 image_embed 和三层 high_res_featsprompt encoder 把点击坐标、框、掩码编码成 prompt embedding最后由 mask decoder 输出低分辨率掩码。交互分割时图片只编码一次点击点和掩码要多次送入 decoder。这个结构对 ONNX 非常不友好因为动态轴不只 H、W还有提示点数量 N。所以部署方案几乎只有一种合理拆法把 image encoder 单独导出一个 ONNX把 mask decoder 单独导出一个 ONNX推理进程里分别建 session。这跟 YOLO 转 ONNX 导出完全不是一回事YOLO 一次前向搞定SAM2 必须拆成“一次性编码、循环解码”。我见过有人硬把整个 predictor 合成一个 ONNX结果 prompt 变成固定张量失去交互能力这是方向性错误。2.2 pytorch 转 onnx 的最小导出脚本归一化写进模型里导出 image encoder 时我习惯把归一化直接写进 ONNX 图里而不是留给部署端做。SAM2 的预处理用 ImageNet 均值方差PyTorch 推理时归一化发生在 predictor 内部导出时不包进去部署端就多一个容易写错的步骤而且 RGB 和 BGR 顺序颠倒是高频事故。把归一化做成一个包装模块随 encoder 一起导出是最省事的import torch import torch.nn as nn class NormalizedEncoder(nn.Module): def __init__(self, image_encoder): super().__init__() self.encoder image_encoder self.register_buffer( mean, torch.tensor([123.675, 116.28, 103.53]).view(1, 3, 1, 1) ) self.register_buffer( std, torch.tensor([58.395, 57.12, 57.375]).view(1, 3, 1, 1) ) def forward(self, x): x (x - self.mean) / self.std feats self.encoder(x) return ( feats[image_embed], feats[high_res_feats][0], feats[high_res_feats][1], feats[high_res_feats][2], )predictor 加载完成后把 image_encoder 塞进这个包装模块predictor build_sam2_image_predictor(config, ckpt_path) wrapped NormalizedEncoder(predictor.model.image_encoder.eval()) dummy torch.randn(1, 3, 1024, 1024, dtypetorch.float32) torch.onnx.export( wrapped, dummy, sam2_tiny_image_encoder.onnx, input_names[input_image], output_names[ image_embed, high_res_feats_0, high_res_feats_1, high_res_feats_2, ], dynamic_axes{ input_image: {2: H, 3: W}, image_embed: {2: H_feat, 3: W_feat}, high_res_feats_0: {2: H_feat0, 3: W_feat0}, high_res_feats_1: {2: H_feat1, 3: W_feat1}, high_res_feats_2: {2: H_feat2, 3: W_feat2}, }, opset_version17, do_constant_foldingTrue, trainingtorch.onnx.TrainingMode.EVAL, )这里的 register_buffer 是关键mean 和 std 会被固化进 ONNX 的常量节点部署端只要喂 0 到 255 的 RGB 图像就行。high_res_feats 三层分别对应 1/4、1/8、1/16 分辨率后续 mask decoder 上采样时要拿这三层特征做融合不能省。dynamic_axes 只放开 H、W通道维和 batch 维保持静态这样可以避免很多 ONNX Runtime 的动态 shape 重编译问题。2.3 导出时的两个关键参数opset 版本与动态轴取舍opset 版本我建议直接选 17。SAM2 里有 aten::pad、双线性插值、一些较新的算子opset 太低会触发算子回退opset 太高在旧版 ONNX Runtime 上又可能找不到实现。选 17 是当前 CPU 部署最稳的折中方案配合 ONNX Runtime 1.15 以上版本基本不会遇到 NotImplementedError。动态轴要不要全放开这是踩出来的经验。图像尺寸如果业务场景固定比如只处理 1024×1024 的输入直接不设 dynamic_axes静态图性能最好。如果必须支持任意尺寸只开放 H、W 两个轴千万别把 batch 和通道也放开。mask decoder 那边的动态轴更复杂prompt 数量 N 是核心动态轴下面导出 wrapper 时专门说明。我建议默认 encoder 固定 1024×1024点提示数量动态这是交互分割最常见的部署形态。2.4 mask decoder 导出把 image_pe 固化为常量mask decoder 导出是最容易翻车的一步。torch 里 decoder 的输入包括 image_embed、image_pe、sparse_prompt_embeddings、dense_prompt_embeddings 等直接导的话部署端得自己实现 prompt encoder得不偿失。常见做法是写一个包装模块把 prompt encoder 和 mask decoder 一起包进去外部只暴露 image_embed、点坐标、点标签和上一轮掩码class MaskDecoderOnnx(nn.Module): def __init__(self, predictor): super().__init__() self.prompt_encoder predictor.model.prompt_encoder self.mask_decoder predictor.model.mask_decoder self.image_pe nn.Parameter( predictor.model.image_pe.detach().clone(), requires_gradFalse, ) def forward( self, image_embed, high_res_feats_0, high_res_feats_1, high_res_feats_2, point_coords, point_labels, mask_input, ): sparse_embed, dense_embed self.prompt_encoder( points(point_coords, point_labels), boxesNone, masksmask_input, ) low_res_masks self.mask_decoder( image_embedimage_embed, image_peself.image_pe, sparse_prompt_embeddingssparse_embed, dense_prompt_embeddingsdense_embed, multimask_outputFalse, repeat_imageFalse, high_res_features[ high_res_feats_0, high_res_feats_1, high_res_feats_2, ], ) return low_res_masks导出时 point_coords 的形状是 (B, N, 2)point_labels 是 (B, N)mask_input 是上一轮的 low_res mask形状 (B, 1, 256, 256)。没有上一轮掩码时传全零张量。image_pe 用 nn.Parameter 固化后部署端就不需要再生成位置编码这一步能省掉很多逻辑wrapper MaskDecoderOnnx(predictor).eval() dummy_embed torch.randn(1, 256, 64, 64) dummy_hr0 torch.randn(1, 256, 256, 256) dummy_hr1 torch.randn(1, 256, 128, 128) dummy_hr2 torch.randn(1, 256, 64, 64) dummy_coords torch.randn(1, 2, 2, dtypetorch.float32) dummy_labels torch.tensor([[[1, 1]]], dtypetorch.float32) dummy_mask torch.zeros(1, 1, 256, 256, dtypetorch.float32) torch.onnx.export( wrapper, (dummy_embed, dummy_hr0, dummy_hr1, dummy_hr2, dummy_coords, dummy_labels, dummy_mask), sam2_tiny_mask_decoder.onnx, input_names[ image_embed, high_res_feats_0, high_res_feats_1, high_res_feats_2, point_coords, point_labels, mask_input, ], output_names[low_res_masks], dynamic_axes{ point_coords: {0: B, 1: num_prompts}, point_labels: {0: B, 1: num_prompts}, mask_input: {0: B}, low_res_masks: {0: B}, }, opset_version17, do_constant_foldingTrue, trainingtorch.onnx.TrainingMode.EVAL, )注意 boxes 参数这里直接传 None业务端把框转成两个点左上角点的 label 记 2右下角点的 label 记 3和 SAM2 的 prompt encoder 内部约定一致。这样部署接口统一成“点列表 标签列表”不用单独处理框的空维度。3. 用 ONNX Runtime 跑通 SAM2 图像分割最小推理脚本与参数说明3.1 拆分会话encoder 与 mask decoder 为什么必须分开跑很多第一次接触 SAM2 部署的人会问能不能合成一个 ONNX一次 run 出结果能做但千万别做。交互场景里图片通常只传一次点击会来很多次合并成一个模型意味着每次点击都要重新跑一遍沉重的高分辨率编码器延迟和算力直接翻倍。而且 mask decoder 的输入输出动态轴很多两个模型合并后 ONNX Runtime 算子融合效率也差。正确的做法是两个 sessionencoder session 跑一次结果缓存每次点击只跑 decoder session。我一般把 encoder 输出放在内存里点击循环里只更新 mask_input 和 point_coords。这样用户点的响应时间基本就是 mask decoder 的推理时间几百毫秒内可以完成一次交互。3.2 最小推理脚本Python onnxruntime 跑通单点分割下面这个脚本是完整的最小闭环从读图到输出二值掩码不含任何多余封装import cv2 import numpy as np import onnxruntime as ort img_bgr cv2.imread(demo.jpg) h_orig, w_orig img_bgr.shape[:2] input_img cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB) input_img cv2.resize(input_img, (1024, 1024), interpolationcv2.INTER_LINEAR) input_img input_img[None].astype(np.float32).transpose(0, 3, 1, 2) ort.set_default_logger_severity(3) enc ort.InferenceSession( sam2_tiny_image_encoder.onnx, providers[CPUExecutionProvider], ) dec ort.InferenceSession( sam2_tiny_mask_decoder.onnx, providers[CPUExecutionProvider], ) image_embed, hr0, hr1, hr2 enc.run( None, {input_image: input_img} ) scale_x 1024.0 / w_orig scale_y 1024.0 / h_orig point np.array([[[520 * scale_x, 300 * scale_y]]], dtypenp.float32) label np.array([[1]], dtypenp.float32) mask_input np.zeros((1, 1, 256, 256), dtypenp.float32) low_res dec.run( None, { image_embed: image_embed, high_res_feats_0: hr0, high_res_feats_1: hr1, high_res_feats_2: hr2, point_coords: point, point_labels: label, mask_input: mask_input, }, )[0] print(low_res_masks shape:, low_res.shape)这里 point_coords 传入的是 1024 坐标空间下的坐标不是原图坐标。原图里点 (520, 300)经过 scale_x 和 scale_y 换算后才能送进模型。mask_input 初始为零张量代表没有上一轮掩码做细化。第一次 run 如果发现图片尺寸很大resize 用 INTER_LINEAR 就好别用 INTER_AREASAM2 训练时用的是双线性推理要和训练对齐。3.3 推理参数说明256 维提示编码与 256×256 低分辨率掩码输出 low_res_masks 的形状是 (1, 1, 256, 256)这是带 logits 的低分辨率掩码不是 sigmoid 之后的结果。SAM2 的 mask decoder 内部会把掩码上采样到 1024×1024 再做输出但 ONNX 导出时为了算子稳定性常见做法是导到 256×256 低分辨率把上采样留给部署端的 cv2.resize。这样做的另一个好处是多轮点击时 mask_input 可以直接把上一轮的 low_res 传回去不用反复缩放。point_labels 的取值规则要背下来0 表示背景点1 表示前景点2 和 3 表示框的左上角和右下角。点提示的 embedding 维度是 256每一维对应 prompt encoder 输出的一个 token多个点会拼接成 (B, N, 256)。mask decoder 里的 cross-attention 会让不同点之间互相影响所以一次传入多个点比多次单点调用效果好尤其是负样本点。4. 提示工程与后处理点的编码、掩码融合与阈值4.1 点的输入构造从原图坐标到 1024 坐标交互分割的核心是把用户点击转换成模型能吃的点。常见做法是前端把点击坐标和类型传过来后端只维护一个数组。每次新增点击整个数组重新组装一次 run 解码器。注意坐标换算必须用浮点运算避免整数除法误差尤其是缩放比例不是整数时clicks [] clicks.append((520, 300, 1)) # 前景点 clicks.append((350, 480, 0)) # 背景点 clicks.append((610, 180, 2)) # 框左上角 clicks.append((780, 420, 3)) # 框右下角 coords [] labels [] for x, y, label in clicks: coords.append([x * scale_x, y * scale_y]) labels.append(label) coords np.array([coords], dtypenp.float32) labels np.array([labels], dtypenp.float32) low_res dec.run(None, { image_embed: image_embed, high_res_feats_0: hr0, high_res_feats_1: hr1, high_res_feats_2: hr2, point_coords: coords, point_labels: labels, mask_input: mask_input, })[0]有个容易踩的坑框的坐标也是原图坐标不能把框当成独立输入要转成两个点并分别标 label 2 和 label 3。如果业务里既有框又有普通点那就全部塞进同一个 coords 数组labels 用 0 到 3 区分类型。prompt encoder 会按 label 索引不同的 embedding 表不要自己改 label 含义。4.2 掩码融合与后处理为什么阈值是 0 而不是 0.5decoder 输出的是 logits不是概率所以阈值用 0不是 0.5。logits 大于 0 表示前景小于 0 表示背景。这个细节很容易让人困惑很多人习惯性做 sigmoid 再阈值 0.5结果其实一样但多一步计算没必要。后处理代码如下def postprocess(low_res_logits, orig_h, orig_w): mask cv2.resize( low_res_logits[0, 0], (1024, 1024), interpolationcv2.INTER_LINEAR, ) mask cv2.resize( mask, (orig_w, orig_h), interpolationcv2.INTER_LINEAR, ) return (mask 0.0).astype(np.uint8)从 256 到 1024 再到原图尺寸两次 resize 是必要的。第一次把低分辨率 logits 还原到模型输入分辨率第二次还原到原图。如果直接一次 resize 到原图会出现边缘锯齿尤其是小目标。后处理里还可以加一个连通域过滤去掉面积小于 50 像素的孤立点这个阈值按业务调。掩码融合方面如果用户连续点多个前景点模型内部已经做了 attention 融合不需要在外部做 max 或 or 操作直接信任 decoder 输出即可。4.3 交互分割循环点加、点减与 mask_input 回填多轮交互时一个提升效果非常明显的技巧是把上一轮的 low_res_masks 回填给 mask_input。SAM2 的 mask decoder 支持迭代细化上一轮掩码作为 dense prompt 输入配合新的点击点能修正边缘和漏检区域。回填的 mask 必须保持 256×256 低分辨率不能用原图掩码或 1024 分辨率的掩码prev_mask np.zeros((1, 1, 256, 256), dtypenp.float32) for click in click_stream: coords, labels build_prompt(click) low_res dec.run(None, { image_embed: image_embed, high_res_feats_0: hr0, high_res_feats_1: hr1, high_res_feats_2: hr2, point_coords: coords, point_labels: labels, mask_input: prev_mask, })[0] prev_mask low_res final_mask postprocess(low_res, h_orig, w_orig)这里要注意 mask_input 和 point_coords 的 batch 维必须一致。如果同时处理多张图mask_input 的形状是 (B, 1, 256, 256)point_coords 的形状是 (B, N, 2)B 对齐即可。多轮交互里如果用户把框删了重新拉框必须把 mask_input 清零否则旧掩码会干扰新框。5. 部署避坑清单导出黑匣子、算子回退与版本玄学5.1 现象模型能导出但 ONNX Runtime 一 run 就报 NotImplementedError第一次跑的时候经常遇到这种情况torch.onnx.export 成功onnxsim 也过了结果 session.run 直接抛异常错误信息指向某个自定义算子或 aten:: 开头的节点。原因是 SAM2 的一些高分辨率特征融合用到了较新的 PyTorch 算子ONNX Runtime 的 CPU kernel 没有实现或者导出图里带了 flash attention 之类的自定义 kernel。解决方法是三步导出时用 EVAL 模式并设置 trainingtorch.onnx.TrainingMode.EVAL把 opset 固定到 17导出后用 onnxsim 常量折叠一遍把不需要的动态分支和辅助输出删掉。如果还有 aten:: 算子残留优先检查 high_res_feats 融合部分换成最朴素的 interpolation 和卷积组合。5.2 现象第一次推理特别慢后续正常很多人以为这是模型性能问题其实多半是 ONNX Runtime 的线程池初始化和内存分配器预热。第一次 run 会把整个图加载、分配临时 buffer、初始化线程池耗时可能到几十秒。解决方法是服务启动后先用一张 1024×1024 的黑色图跑一次 warmup再进入正式请求。如果 warmup 之后还是慢才是真正的性能问题。检查 intra_op_num_threads 是否设置合理CPU 核心数多时先设 4 到 8再往上加收益不大。另外确保没有在点击循环里反复创建 InferenceSessionsession 创建成本很高应该常驻内存。常见做法是服务启动时创建两个全局 session用锁保护并发访问。5.3 现象ONNX 结果和 PyTorch 对不上掩码形状完全不同这类问题九成出在预处理。PyTorch 里 SAM2 的归一化在 predictor 内部完成导出的 encoder 虽然包了归一化但如果部署端自己又做了一遍等于归一化了两次。另一个高频原因是 cv2 读图是 BGRPyTorch 训练时用的是 RGB漏掉 cvtColor 会导致颜色通道错乱掩码结果完全不可信。对拍方法很简单用同一张图同一个点坐标分别跑 PyTorch 和 ONNX对比 low_res_masks 的数值。误差在 1e-3 级别算正常超过 0.1 就要检查预处理。我一般会写一个对拍脚本把两张图 mask 的 IoU 算出来低于 0.95 就报警。5.4 现象点击同一个位置多次返回的 mask 不一样这种情况通常是动态 shape 导致的重新编译。ONNX Runtime 对动态 H、W 会缓存编译结果但如果输入 shape 频繁变化比如一会 1024 一会 800每次都要重新优化结果也可能出现微小差异。另一个原因是并发访问同一个 session多个线程同时 run虽然 ONNX Runtime 声称 session 线程安全但交互场景里最好还是加锁或者每个线程独立创建 session。解决方法是把图像输入固定为 1024×1024点数量动态轴设上限比如最多 32 个点。超出就截断提示用户减少点击。这样 model 的输入 shape 变化范围可控重编译次数大幅下降。固定 shape 后我实测 CPU 推理延迟能稳定在几百毫秒内。5.5 现象量化后发现 mask decoder 反而变慢onnx 量化 int8 看起来是万能优化但对 SAM2 这种结构未必适用。image encoder 是卷积加 attention 的混合结构int8 权重量化后 CPU 推理通常能快 1.5 到 2 倍。但 mask decoder 很小动态 prompt 路径上有大量 reshape、concat、gatherint8 量化后反而要频繁做量化反量化开销比省掉的算力还大。我目前的落地配置是image encoder 用 int8 动态量化mask decoder 保持 FP32。实测整体延迟下降主要来自编码器解码器的几百毫秒本来就是交互操作可接受的。别贪全模型量化收益不大还引入精度风险。6. 进阶量化与验证——从“能跑”到“敢上线”6.1 用 onnx 量化 int8 给 image encoder“降温”CPU 部署时image encoder 往往是最大耗电点1024×1024 输入在 Hiera backbone 上跑一次FP32 可能要两秒以上。动态量化是最省事的优化手段不需要校准集直接对权重做 int8from onnxruntime.quantization import quantize_dynamic, QuantType quantize_dynamic( sam2_tiny_image_encoder.onnx, sam2_tiny_image_encoder_int8.onnx, weight_typeQuantType.QUInt8, op_types_to_quantize[Conv, MatMul], )量化后第一次 run 前建议做一次 warmupint8 模型的权重反量化也有初始化开销。如果追求更高加速可以上静态量化但需要准备几百张业务分布接近的校准图否则边缘场景可能出现掩码塌陷。我踩过静态量化在遥感图上丢小目标的坑后来退回动态量化损失约 1 到 2 个点的 IoU换取可控的稳定性。6.2 上线前的点位命中率验证逐点对拍不能省量化完了不能只看几张样例图就上线。交互分割的验证重点不是 mIoU而是“用户随便点一个位置能不能得到合理掩码”。我会写一个自动化对拍脚本准备十张覆盖不同场景的测试图每张图预设十个人工点分别跑 FP32 模型和 int8 模型对比掩码 IoUdef validate(onnx_path_a, onnx_path_b, images, points): for img, pts in zip(images, points): for x, y in pts: mask_a predict(onnx_path_a, img, x, y) mask_b predict(onnx_path_b, img, x, y) iou compute_iou(mask_a, mask_b) assert iou 0.85, fpoint ({x},{y}) iou too low: {iou:.3f}这条规则里 IoU 阈值 0.85 是我试出来的经验值低于这个值说明量化失效或某个算子回退。如果发现只有特定尺寸的图出问题优先查 resize 分支和动态轴是否在编译缓存之外。上线前再跑一次压测确认并发点击下 session 锁不会成为瓶颈。整个方案做完SAM2 的 ONNX 部署才算是闭环。这些坑看起来都是细节但每个都真实消耗过我的调试时间。希望帮到你。本文还有配套的精品资源点击获取
觉得有用,分享给同行:

为您的企业打造数字门面

稳重轻奢商务风格,端正雅致视觉,长效耐看不易过时。

立即咨询 →