ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

ONNX Runtime跑通SAM2:从导出到量化的完整部署指南

ONNX Runtime跑通SAM2:从导出到量化的完整部署指南 简介面向深度学习与图像分割开发者这是一套将 Segment Anything 2SAM2模型转换为 ONNX 格式并提供推理能力的 Python 脚本资源解决 SAM2 在多种平台和设备上的跨框架部署与高效分割问题。压缩包共 14 个文件其中 4 个 Python 脚本覆盖模型转换与调用入口4 个 ONNX 模型文件对应图像编码与分割解码环节另含 4 张示例图片和 README 文档整体约 591.77MB目录按功能拆分便于对照学习与修改。资源发布后已有 199 人学习浏览。通过这套脚本用户既能把 SAM2 导出为通用 ONNX 模型也能直接调用模型完成图像分割编码器与解码器文件相互配合示例图片可快速验证分割效果通过简单调用即可完成预处理、推理与后处理流程。整体上降低了 SAM2 的应用门槛适合需要跨平台推理、边缘部署或做定制分割的研究者与工程师。1. 拿到 ONNX-SAM2 压缩包之后先别急着解压跑模型拿到ONNX-SAM2-Segment-Anything.zip的时候大多数人第一反应是解压、装环境、跑个 demo但实际情况往往是第一步就翻车。SAM2Segment Anything 2官方仓库给的只有 PyTorch 权重和推理脚本ONNX 版本基本是社区二次导出的产物解压之后你可能发现模型文件、脚本、权重散落得到处都是甚至根本没有权重文件。这篇文章要解决的问题很直接把这个压缩包里的 Python 脚本用起来把 ONNX Runtime 推理链路完整跑通包括导出、预处理、后处理、量化以及那些会让 mask 全黑、脚本崩溃的常见坑。适合两类人一类是要在 CPU 或 GPU 服务里快速集成 SAM2 分割能力的后端工程师另一类是刚接触 ONNX 部署、拿着 SAM2 练手但不想在环境依赖上耗太久的算法工程师。这里不聊论文创新只聊怎么把模型跑起来。2. 为什么用 ONNX Runtime 跑 SAM2模型拆解与导出路径选择2.1 SAM2 不是一个大模型而是三个模块的组合先说一个最容易误解的地方SAM2 不是一个黑匣子单模型而是 image encoder、mask decoder、memory attention 三块拼起来的。官方 PyTorch 代码里Sam2Predictor对外暴露的是一整套推理接口内部会在set_image里跑一遍 image encoder把结果缓存下来然后在predict时把 prompt 和缓存特征一起喂给 mask decoder。如果只是做单张图片的分割memory attention 根本不会参与。这对 ONNX 部署影响非常大。社区里流传的onnx-sam2脚本绝大多数会把模型拆成两个或三个 ONNX 文件一个image_encoder.onnx负责把 1024×1024 的输入图变成特征图一个mask_decoder.onnx负责把 prompt 和特征图变成 mask logits。如果脚本里还带了视频分割能力那还会多一个memory_attention.onnx用来在帧间传递状态。所以拿到压缩包第一步是先看清楚里面拆成了几个 ONNX 文件而不是急着跑。我的经验是只做静态图片分割找image_encoder加mask_decoder两个文件就够了强行把三个模块拼成一个超大 ONNX不仅导出容易碰到算子不支持推理速度也会被拖垮。压缩包里如果只有单个sam2.onnx反而要警惕那多半是把 image encoder 整个打包进去了输入输出维度非常怪新手用起来很容易在 shape 上卡住。2.2 PyTorch 转 ONNX 的三种常见路线torch.onnx.export 与 onnx-simplifier如果压缩包里的脚本自带pytorch2onnx.py那可以直接用。如果没有最常见的做法是拿官方 SAM2 权重自己导出。导出这一环的坑主要集中在动态 shape 和算子兼容性上我一般按下面这个最小脚本走import torch model load_sam2_model(sam2.1_hiera_large.pt) # 换成你手里的权重 model.eval() # 构造一组有代表性的输入shape 必须以实际推理时的 shape 为准 image torch.randn(1, 3, 1024, 1024) prompt_points torch.randn(1, 5, 2) # 一个 batch5 个点 prompt_labels torch.randint(0, 2, (1, 5)).float() torch.onnx.export( model, (image, prompt_points, prompt_labels), sam2_mask_decoder.onnx, opset_version17, input_names[image, points, labels], output_names[low_res_logits, iou_predictions], dynamic_axes{ image: {0: batch}, points: {0: batch, 1: num_points}, }, do_constant_foldingTrue, )导出后强烈建议再用onnxsim过一遍把常量折叠后残留的多余 reshape 和 gather 清掉。这一步在 SAM2 上几乎是必须的因为 mask decoder 内部有大量基于num_points的动态索引导出时 torch 会留下很多冗余算子直接推理不会报错但 CPU 上延迟能差出 30% 以上。onnxsim安装后一行命令python -m onnxsim sam2_mask_decoder.onnx sam2_mask_decoder_sim.onnx参数层面只有一个opset_version值得注意。我固定用 17原因是在更低的 opset 下torch.flatten和torch.split导出的Split算子对动态轴支持不全容易在 ONNX Runtime 里报 shape 推导错误。opset 不是越高越好18 以上部分算子对 ORT 老版本反而兼容性差。2.3 导出后先用 Netron 和 onnxruntime 做一次健康检查导出这一步最容易出现“明明代码没报错推理结果却不对”的玄学问题。我现在每导出一个 ONNX都固定做三件事用 Netron 打开看输入输出节点确认输入 shape 和dynamic_axes是否按预期生效用onnxruntime跑一次随机输入确认输出形状不是(1,1,256,256)这种被固定死的最后用同一张图分别在 PyTorch 和 ONNX Runtime 里跑对比输出 logits 的余弦相似度。import onnxruntime as ort import numpy as np session ort.InferenceSession(sam2_mask_decoder_sim.onnx) out session.run(None, { image: np.random.randn(1, 3, 1024, 1024).astype(np.float32), points: np.random.randn(1, 5, 2).astype(np.float32), labels: np.random.randint(0, 2, (1, 5)).astype(np.float32), }) print([o.shape for o in out])如果输出 shape 第一个维度不是 1而是和num_points绑定的动态值说明dynamic_axes没设对。这里有个血泪教训mask decoder 的输入点数一改变ONNX Runtime 会重新做一次 session optimize第一次调用可能慢到几百毫秒后续才恢复正常。后面第 5 章专门讲这个问题。3. 跑通 ONNX-SAM2 最小脚本加载模型、前处理、推理、后处理3.1 环境准备与压缩包内容核对解压ONNX-SAM2-Segment-Anything.zip后先别管代码先确认三样东西有没有.onnx文件有没有.pt/.pth权重文件有没有requirements.txt。很多这类压缩包只放了脚本和 README权重让你自己去 HuggingFace 下载模型文件上百 MB文章中不会打包。如果你看到的包里只有脚本去 README 里找权重下载链接路径配置一般在脚本头部的MODEL_PATH常量里。环境方面Python 版本压到 3.93.11 之间最稳。ONNX Runtime 对 3.12 的支持虽然没问题但 SAM2 脚本里经常依赖旧版fvcore、hydra-core这类库这些库在新 Python 上容易编译失败。注意一个细节在 VSCode 里配置 Python 环境时一定要确认python命令指向的是虚拟环境而不是系统默认解释器否则onnxruntime装进了.venvpython却指向全局 Python导入时报错会非常绕。依赖安装我一般这样处理pip install onnxruntime-gpu1.17.1 opencv-python numpy pillow版本随机器而变但不要装最新版 onnxruntime。1.18 之后内部结构改动较大部分 SAM2 脚本里用到的registry.register接口会失效。这个版本问题不写到 README 里得自己拍板。3.2 onnxruntime 会话初始化检查 provider 而不是直接 run初始化 ONNX Runtime 会话时不建议直接写死providers[CUDAExecutionProvider]因为换一台没 GPU 的机器脚本就崩。更稳的写法是获取可用 provider 后做选择import onnxruntime as ort print(ort.get_available_providers()) session ort.InferenceSession( sam2_mask_decoder.onnx, providers[CUDAExecutionProvider, CPUExecutionProvider], )这段逻辑说明get_available_providers()返回当前环境实际可用的执行提供程序InferenceSession会按顺序尝试加载CUDA 不可用会自动回退到 CPU。参数上不需要额外设置但如果要用 GPU提前确认onnxruntime-gpu装的是 CUDA 12 还是 CUDA 11 的构建版本和本机驱动不匹配时会出现第 5 章讲的段错误。3.3 前处理图像缩放、padding 和归一化要和导出时严格对齐SAM2 的前处理是大多数脚本跑出全黑 mask 的第一嫌疑犯。官方 SAM2 的 image encoder 输入是 1024×1024但输入图片通常不是正方形所以要先做等比例缩放再在短边补 padding 到 1024×1024。这还没完归一化方式的坑更大官方在导出 ONNX 时有些脚本把归一化写进了模型内部有些则留在外面。我的做法是每一次都在 Python 侧做归一化而不信任 ONNX 模型内部自带预处理因为模型一旦被量化内部 BN 层常常被折叠输入分布会发生变化。import cv2 import numpy as np def preprocess(image: np.ndarray, target_size: int 1024): h, w image.shape[:2] scale target_size / max(h, w) new_h, new_w int(round(h * scale)), int(round(w * scale)) resized cv2.resize(image, (new_w, new_h), interpolationcv2.INTER_LINEAR) pad_h target_size - new_h pad_w target_size - new_w padded np.zeros((target_size, target_size, 3), dtypenp.float32) padded[:new_h, :new_w, :] resized / 255.0 # 归一化到 [0,1] # 转 CHW并加 batch 维度 tensor padded.transpose(2, 0, 1)[None, ...].astype(np.float32) return tensor, scale, (new_h, new_w)这段代码里必须记住的值是scale和(new_h, new_w)后处理要把 padding 区域裁掉并缩放回原图尺寸。很多脚本只做了 padding 没记录scale结果 mask 和原图叠加时位置完全错位。注意这里用的是cv2.resize而不是 PIL 的resize两者在非整数倍缩放时的像素插值算法不同容易导致输出有 1-2 像素的整体偏移。如果你手里的脚本是用 PIL 写的建议全套保持一致混用必出问题。3.4 完整推理脚本从加载模型到输出 mask前面几步拼起来就是一个可以直接跑的 SAM2 推理脚本。下面这段代码我把 image encoder 和 mask decoder 拆成两个 session 来处理这是社区脚本里最常见的结构import cv2 import numpy as np import onnxruntime as ort image_encoder ort.InferenceSession(sam2_image_encoder.onnx) mask_decoder ort.InferenceSession(sam2_mask_decoder.onnx) # 1. 前处理 image cv2.imread(demo.jpg) tensor, scale, (new_h, new_w) preprocess(image, 1024) # 2. image encoder 输出特征 features image_encoder.run(None, {image: tensor}) # features[0] 的 shape 通常是 (1, 64, 256, 256)具体以模型为准 # 3. prompt 构造以单点为例 points np.array([[[500, 300]]], dtypenp.float32) # 原图坐标不是缩放后坐标 labels np.array([[1]], dtypenp.float32) # 4. mask decoder 推理 logits, iou mask_decoder.run(None, { image_features: features[0], points: points, labels: labels, }) # 5. 后处理 mask 1 / (1 np.exp(-logits[0, 0])) # sigmoid mask_binary (mask 0.5).astype(np.uint8) * 255 mask_binary mask_binary[:new_h, :new_w] # 裁掉 padding mask_resized cv2.resize(mask_binary, (image.shape[1], image.shape[0]))这里有一个必须强调的参数细节points给的是原始图像上的绝对坐标而 image encoder 的输入是缩放并 padding 后的 1024×1024 张量所以 ONNX 模型内部已经做了坐标映射的假设。如果压缩包脚本没有帮你处理坐标你需要手动把点坐标按照scale一起缩放points_scaled points * scale这个点不处理标出来的位置和 mask 永远对不上。另外mask_decoder.run里传的image_features必须是image_encoder输出中的特征张量有些脚本里 image encoder 会同时输出多个尺度的特征mask decoder 需要的是其中某几个拼错位置会导致输出 logits 数值全是 NaN 或固定值。4. 量化与性能调优int8 量化、FP16 与动态轴的取舍4.1 先别量化先确认你的瓶颈在哪给 SAM2 做 int8 量化之前先想清楚部署目标是 CPU 还是 GPU。GPU 上 int8 的收益远不如 CPU 明显而且 GPU 推理主要瓶颈往往在 image encoder 的密集卷积上ONNX Runtime 的 CUDA int8 支持并不完整很多算子会回退到 FP32收益打折。CPU 机器做 int8 静态量化才是 SAM2 最有性价比的优化路径尤其是把 image encoder 从 FP32 压到 int8模型体积大约减小到四分之一推理延迟能降一半左右。如果目标是 GPU 且只是想让单次推理快点优先用 FP16 半精度而不是 int8。FP16 转换在代码上最简单风险也最小唯一要注意的是 ONNX 里如果有LayerNormalization这类对精度敏感的算子FP16 下可能出现轻微抖动肉眼未必看得出来的 mask 边界毛糙。4.2 静态 int8 量化实操校准数据决定了精度底线int8 量化在 ONNX Runtime 里的 API 相对稳定核心是quantize_static但它是数据依赖型你必须提供一个校准数据读取器让 ORT 统计每一层激活值的动态范围。下面是 SAM2 image encoder 的一个最小量化实现片段from onnxruntime.quantization import ( CalibrationDataReader, QuantFormat, QuantType, quantize_static, ) class ImageCalibReader(CalibrationDataReader): def __init__(self, image_paths, input_nameimage): self.data [] for path in image_paths: img cv2.imread(path) tensor, _, _ preprocess(img) self.data.append({input_name: tensor}) self.iter iter(self.data) def get_next(self): return next(self.iter, None) calib_reader ImageCalibReader([calib_1.jpg, calib_2.jpg, calib_3.jpg]) quantize_static( sam2_image_encoder.onnx, sam2_image_encoder_int8.onnx, calib_reader, quant_formatQuantFormat.QDQ, per_channelTrue, weight_typeQuantType.QInt8, )这段代码里值得调的参数有三个。第一个是per_channelTrue卷积权重按输出通道单独算缩放系数比 per-tensor 精度高很多尤其对 Hiera backbone 这种带残差连接的结构per-tensor 量化后 mask 边缘容易出现方块。第二个是QuantFormat.QDQ它保留 DeQuantize 节点允许部分算子以混合精度运行比直接QOperator格式更稳。第三个是校准数据量我实测 SAM2 image encoder 至少需要 20 张不同场景的图片太少了激活值范围统计不准量化后的 mask 会出现大块黑色区域。校准数据最好从实际业务数据里采样而不是拿 COCO 的图片糊弄否则上线后遇到分布外的图像精度崩得厉害。不要忽略一个常见坑如果你想转出的 int8 模型还要再转成 RKNN 或 NCNN 部署建议先确认目标工具链对 QDQ 格式的支持。很多边缘端工具链只认纯 QOperator 格式这时要把QuantFormat切换成QOperator并关掉per_channel用稍低精度换取兼容性。4.3 用固定分辨率换取成倍吞吐动态 shape 是 SAM2 ONNX 部署的性能杀手。前面 2.3 提过输入分辨率一变ONNX Runtime 会重新优化内部执行计划这个开销在最坏情况下能到几百毫秒比你实际推理时间还长。如果你在服务里接的是摄像头流每帧分辨率一致问题不大但如果是用户上传图片分辨率参差不齐每来一张图就重新 plan吞吐直接崩掉。我一般会在服务启动时固定 2 到 3 个分辨率池比如 1024、768、512每个分辨率各创建一个InferenceSession请求进来按图缩放到最近的分辨率然后从 session 池里取出对应会话推理。这样既保留了动态 shape 的能力又避开了反复 re-plan 的代价。代价是显存和内存会翻倍但在实际部署中完全值得。之所以不用单一固定分辨率是因为从 2K 图直接缩到 512 会让小目标丢失而固定三个档位能把召回率保持在一个可接受范围内。4.4 量化后必须做一个回归验证很多人在量化后只看一眼输出 mask 是否还是轮廓就拍板上线了这基本属于自欺欺人。量化模型在边缘区域的表现力下降是必然的你需要一个客观指标来回答“精度到底掉了多少”。用一个自己的测试集跑量化前后模型的平均 IoU记录两个数字量化前 FP32 模型的 IoU量化后 int8 模型的 IoU。两者差值超过 0.05 就要考虑回退或者只量化 image encoder 而保持 mask decoder 为 FP32。因为 mask decoder 的输入输出维度小量化省不了多少内存却对精度极其敏感。5. 部署避坑5 个让 SAM2 脚本崩溃或输出全黑 mask 的常见问题5.1 一运行就段错误onnxruntime-gpu 和 CUDA 版本失配现象脚本刚执行到InferenceSession(...)就崩溃报segmentation fault没有任何 Python traceback。这个问题我在多个机器上遇到过尤其在服务器上装过多个 CUDA 版本时最容易触发。原因onnxruntime-gpu是为特定 CUDA 版本预编译的如果你装的是 CUDA 11.8 的构建机器上却是 CUDA 12 的环境底层动态库加载直接失败。问题在于很多部署机同时存在多个 CUDA 版本环境变量LD_LIBRARY_PATH指向了错误的那一个。解决先确定显卡驱动支持的 CUDA 版本然后严格匹配onnxruntime-gpu的构建号。一个快速验证方法是python -c import onnxruntime as ort; print(ort.__version__) pip list | grep onnxruntime然后用ldd查一下onnxruntime_pybind11_state.so链接到了哪个libcudart。如果不对卸载重装匹配版本。5.2 mask 全黑或全白输入归一化方式和导出时不一致现象模型能跑通logits 输出数值也正常但 sigmoid 后所有像素都低于 0.5mask 变成全黑或者全部高于 0.5变成全白。原因最常见的两个来源一是前处理里做了减均值除方差而导出时模型内部已经做了归一化等于归一化了两次把输入分布完全破坏二是用了 PIL 读取图像通道顺序是 RGB但模型训练时用的是 BGR输入图像整体色偏特征完全错乱。解决先放弃猜测直接用一张图分别跑 PyTorch 官方脚本和 ONNX 脚本打印 logits 的均值、方差、最大最小值多组对比基本能定位问题。如果是归一化问题把 Python 侧的归一化删掉保留x / 255.0。如果是通道顺序问题把读取改成cv2.imread得到 BGR或者把 PIL 读到的rgb_image[:, :, ::-1]翻转回去。提示对比 PyTorch 和 ONNX 输出时PyTorch 侧的输入也必须经过同样的预处理不要直接用官方 demo 的输出做基准否则你对比的是两个不同预处理链的差异。5.3 多个 prompt 时 mask 错乱batch 维度组合错误现象单点 prompt 一切正常一旦传入多个点时输出的 mask 数量和位置完全不符合预期有时是全部 mask 都集中于第一个点有时输出通道数比点数多。原因mask decoder 内部的num_points不仅影响 prompt 张量的 shape还决定了输出张量在batch维度上的语义。不同版本 SAM2 脚本对 prompt 维度的组织方式不一样有的接受(1, N, 2)有的要求(N, 1, 2)有的内部要先expand到 batch 维度。解决在跑服务和做测试时我给 prompt 处理单独写一个适配函数集中处理维度问题而不是在业务代码里到处传递原始数组。这里最容易踩坑的是一次给一张图传 5 个点模型返回 5 个 mask但每个 mask 是 5 个 prompt 共同作用的结果而不是各自独立的。如果需要每个点独立分割就要循环调用 5 次而不是一次性传 5 个点。5.4 动态分辨率导致延迟暴增第一次推理总是特别慢现象推理服务上线后发现第一张图延迟 800ms之后每张只要 100ms。用户量一大每来一种新分辨率都有人撞上首次推理的延迟惩罚平均响应时间很难看。原因这个是 ONNX Runtime 动态 shape 的固有行为每次遇到新的输入 shape重新推导图结构并生成执行计划代价从百毫秒到秒级不等。SAM2 的 image encoder 输入是动态轴这个问题几乎必然出现。解决我在 4.3 写的 session 池方案是行之有效的但还需要一个额外动作在服务预热阶段把池子里每个 session 都用该档位分辨率的一张随机图跑一次让 ORT 提前完成 plan。我第一次没有做预热上线后预热请求的延迟一样打到了用户身上。别省这一步。5.5 导出时报 UnsupportedOperator控制流和动态索引是元凶现象torch.onnx.export过程中报UnsupportedOperator或者导出成功但 ONNX Runtime 推理时报NotImplemented算子集中在unique、nonzero、sort、masked_fill这类。原因SAM2 的 mask decoder 里有大量基于点数的动态控制流比如对不同 prompt 数量做条件分支这些逻辑在导出时会被 torch 转换成一些 ONNX 不完全支持的算子。还有一种情况是脚本里混入了torch.backends.cuda.flash_attention相关逻辑导出的注意力算子带有fP16的专用实现。解决优先把涉及控制流的部分用手工方式展开例如把动态循环改成固定最大点数的张量操作。如果算子问题只集中在某个特定 op可以试着把opset_version从 17 降到 14有时高 opset 引入的复合算子反而是 ORT 不支持的。最后一个办法是把模型的 PyTorch 版本和 onnxruntime 版本对齐升级老版本的 torch 导出的算子模式在 ORT 新版本里有时候反而不兼容这属于版本黑匣子只能靠组合测试试出来。6. 一个 IoU 回归脚本量化之后用它护体量化、改前处理、换 onnx 版本之后你对模型“还有没有以前准”的判断不能靠肉眼。我推荐你写一个极简的 mask 一致性验证脚本固定在一个测试集上跑每次改动之后先过这个脚本再谈上线。测试集不用大二三十张带标注的图就够关键是场景分布要接近真实业务。核心逻辑是量化前后模型输出 mask 与人工标注之间的 IoU 对比以及量化模型和 FP32 模型输出之间的一致性。import numpy as np def compute_iou(pred, gt): pred (pred 0.5) gt (gt 0.5) intersection np.logical_and(pred, gt).sum() union np.logical_or(pred, gt).sum() return intersection / (union 1e-6) results [] for img_path, gt_path in test_set: pred run_onnx_sam2(img_path, session_int8) gt cv2.imread(gt_path, cv2.IMREAD_GRAYSCALE) results.append(compute_iou(pred, gt)) print(int8 IoU:, np.mean(results))这个脚本里的threshold0.5本身就是一个可以调的参数。SAM2 输出的 logits 经过 sigmoid 后0.5 是默认阈值但在某些场景里 0.3 或 0.7 反而更合适。比如医学薄壁结构边缘对比度很低0.5 会把细枝末节全砍掉降到 0.3 能多保留一点召回。这类调节如果不记录下来下次换个人来维护就会一脸懵。我的习惯是把阈值连同模型版本、量化配置一起写进配置文件里并且让 IoU 回归脚本在跑的时候自动读取保证测试结果可复现。最后说一个实际教训我之前在一个目标检测服务里集成 SAM2 做分割后处理当时没有跑量化回归直接上线结果 int8 模型在夜间低光图上的 mask 面积掉了 15%肉眼根本看不出来是下游统计模块先发现数据异常才追回来的。从那之后凡是动过模型文件、量化配置、前处理代码我都会先跑一遍上面的脚本确认 IoU 掉点在可接受范围内再进 CI。这个 IoU 回归脚本一百行不到却是整个 SAM2 部署流程里我最后悔没有早点写的东西。希望帮到你。本文还有配套的精品资源点击获取
返回列表