
简介本资源是一套面向深度学习开发者与计算机视觉工程师的ONNX格式SAM2图像分割工具脚本聚焦于将Segment Anything 2SAM2模型高效部署至多平台环境解决原生PyTorch模型跨框架兼容性差、边缘设备部署难等实际问题。压缩包共14个文件含4个核心Python脚本如inference.py、sam2.py、4个ONNX模型文件含sam_vit_b.onnx、sam2_hiera_base_plus_encoder.onnx等、4张示例图像dog.jpg、result.jpg等及README说明文档整体体积591.77MB结构清晰模块分离明确——模型导出、推理调用、可视化结果各司其职。已有202人学习下载适合具备Python基础与ONNX使用经验的中高级开发者快速集成SAM2能力。读者可直接运行脚本完成ONNX模型导出、图像实例分割全流程复用预置模型权重与推理逻辑并基于源码进行轻量化适配或边缘端部署优化。1. 为什么把 SAM2 模型导出成 ONNX 不是“一键转换”而是场硬仗从 PyTorch 黑匣子到可部署推理引擎的实操真相Segment Anything Model 2SAM2发布后一线算法工程师和嵌入式部署同学几乎同时陷入两难一边是论文里惊艳的视频分割能力一边是官方只提供 PyTorch 原生权重、无 ONNX 导出支持、无量化说明、无跨平台推理验证。你手头那个ONNX-SAM2-Segment-Anything.zip文件绝不是“下载解压就能跑”的玩具包——它本质是一套经过反复踩坑、手动补全算子、绕过动态 shape 陷阱、重写 prompt encoder 接口后才勉强落地的最小可行部署链路。这个脚本真正解决的不是“能不能转 ONNX”而是“怎么让 SAM2 在 CPU 环境下稳定输出 mask、不崩 shape、不丢精度、不卡帧率”。适合三类人需要在边缘设备Jetson/树莓派/工控机跑实时视频分割的嵌入式工程师被客户要求把 SAM2 集成进已有 C/C# 生产系统的算法交付工程师以及正在为模型服务化FastAPI/Triton做准备、但被 ONNX 动态输入搞到失眠的 MLOps 同学。别信“PyTorch 转 ONNX 只需一行 torch.onnx.export”SAM2 的 prompt encoder 是个带 condition 控制流的黑匣子mask decoder 里藏着多尺度特征融合的隐式循环——这些才是 zip 包里 Python 脚本真正要啃的硬骨头。2. 从源码出发SAM2 官方模型结构拆解与 ONNX 兼容性断点定位SAM2 的核心架构不是简单堆叠 CNN 或 Transformer而是由Prompt Encoder Memory Attention Mask Decoder三大部分构成且存在强时序依赖video mode 下 memory token 需跨帧累积。官方 GitHub 仓库facebookresearch/sam2中sam2/modeling/sam2.py定义了主干但关键问题藏在细节里Sam2ImagePredictor和Sam2VideoPredictor的 forward 流程完全不同prompt encoder 中get_dense_pe()返回的 position embedding 是动态生成的mask decoder 的predict_masks()内部调用self._process_mask_decoder_outputs()时会根据 prompt 类型point / box / mask触发不同分支而 ONNX export 对 control flow 支持极弱。我们不能直接对Sam2VideoPredictor整体调用torch.onnx.export——它会报错RuntimeError: Exporting the operator aten::new_empty to ONNX opset version 17 is not supported因为内部用了torch.empty_like()构造动态 shape tensor。2.1 SAM2 模型的三个 ONNX 敏感区哪部分必须重写、哪部分可冻结、哪部分得阉割模块是否可直接导出原因实操策略Image EncoderViT✅ 可直接导出结构规整无条件分支输入固定1024×1024输出 shape 确定H×W×C使用torch.jit.tracetorch.onnx.exportopset17dynamic_axes{}留空Prompt Encoder❌ 必须重写get_dense_pe()依赖self.pe_layer的forward()内部含torch.arangeunsqueeze生成 shape 与输入分辨率强耦合point prompt 处理使用torch.whereONNX 不支持动态索引提前预计算 PE 并存为常量将 point prompt 编码逻辑抽离为独立函数用torch.nn.functional.grid_sample替代原始坐标映射Mask Decoder⚠️ 部分重写predict_masks()中self._process_mask_decoder_outputs()根据is_mask_from_points切换路径ONNX 不支持布尔控制流memory attention 的self.memory_attention在 video mode 下需传入历史 memory tokensshape 动态变化强制拆分为两个子模型mask_decoder_point仅处理点 prompt和mask_decoder_box仅处理框 prompt禁用混合 prompt 输入memory tokens 输入改为固定长度如 max_frames5padding 后截断提示不要试图用torch.onnx.export(..., dynamic_axes{input_points: {0: batch, 1: num_points}})让 point prompt 支持变长——SAM2 的 point 数量直接影响 decoder 内部 attention mask 构建逻辑ONNX runtime 无法解析这种嵌套动态依赖。真实做法是固定最大点数如 32不足则 zero-pad超出则 trunc并在 Python 脚本中做前置校验。2.2 ONNX-SAM2-Segment-Anything.zip 的核心文件结构与作用链解压ONNX-SAM2-Segment-Anything.zip后你会看到以下关键文件非官方发布是社区硬核适配产物ONNX-SAM2-Segment-Anything/ ├── sam2_onnx_exporter.py # 主导出脚本加载 PyTorch checkpoint → 构建 wrapper → 调用 onnx.export ├── sam2_onnx_wrapper.py # 核心封装类继承 torch.nn.Module重写 forward屏蔽原始 control flow ├── sam2_onnx_inference.py # 推理入口加载 .onnx → 预处理 → run → 后处理包括 mask threshold filter ├── sam2_config.yaml # 模型配置指定 image_size、max_points、max_boxes、onnx_opset、quantize_int8 开关 ├── models/ │ ├── sam2_hiera_tiny.pt # PyTorch 原始权重tiny 版便于调试 │ └── sam2_hiera_large.pt # PyTorch 原始权重large 版生产推荐 ├── onnx/ │ ├── sam2_tiny_image_encoder.onnx # 已导出的 encoder静态 shape │ ├── sam2_tiny_prompt_encoder.onnx # 重写后的 prompt encoder固定 point 数 │ └── sam2_tiny_mask_decoder.onnx # 拆分后的 decoderpoint-only branch └── utils/ ├── preprocess.py # 图像 resize normalize适配 ONNX 输入要求 └── postprocess.py # mask 解码 sigmoid thresholdONNX 输出是 logits非概率注意该 zip不包含 Triton 模型仓库或 WebAssembly 编译产物它专注解决“本地 Python 环境下 ONNX 可运行”这一最小闭环。所有.onnx文件均通过onnxruntime.InferenceSession加载未使用 TensorRT 或 OpenVINO——这是为了保证跨平台一致性Windows/Linux/macOS 均可跑通。3. 手把手跑通用 sam2_onnx_exporter.py 导出你的第一个 SAM2 ONNX 模型导出不是执行一个命令就完事而是要先理解 SAM2 的 checkpoint 加载机制、再构造符合 ONNX 约束的 wrapper、最后用特定参数调用 export。整个过程必须严格遵循sam2_onnx_wrapper.py定义的接口契约否则导出的模型在推理时会 shape mismatch 或 output missing。3.1 环境准备与依赖锁定为什么必须用 torch2.1.2 onnx1.15.0SAM2 官方要求 PyTorch ≥ 2.0.1但 ONNX 导出稳定性在 2.1.2 达到峰值。更高版本如 2.2引入了torch.compile默认启用会干扰 trace 过程更低版本如 2.0.1对torch.nn.MultiheadAttention的 ONNX 映射不完整。同样onnx1.15.0是最后一个兼容opset_version17且不强制要求onnxscript的版本——而 SAM2 的 ViT encoder 中大量使用torch.nn.functional.scaled_dot_product_attention该算子在 ONNX opset 17 中才被正式支持。# 创建干净环境推荐 conda conda create -n sam2-onnx python3.9 conda activate sam2-onnx pip install torch2.1.2cu118 torchvision0.16.2cu118 --extra-index-url https://download.pytorch.org/whl/cu118 pip install onnx1.15.0 onnxruntime-gpu1.17.1 opencv-python4.8.1 numpy1.23.5 PyYAML6.0.1 # 安装 SAM2 官方库注意必须从源码安装pip install sam2 会缺失 video predictor git clone https://github.com/facebookresearch/sam2.git cd sam2 pip install -e .注意onnxruntime-gpu1.17.1是关键——它支持 CUDA Graph 加速且对MultiHeadAttention的 kernel 优化成熟。若用 CPU 推理请换onnxruntime1.17.1非 -gpu 后缀否则会报CUDA initialization failed。3.2 修改 sam2_onnx_wrapper.py重写 Prompt Encoder 的三大硬编码点打开sam2_onnx_wrapper.py找到class SAM2ONNXWrapper(torch.nn.Module)。其__init__中已加载原始 SAM2 模型但forward方法必须重写。重点修改三处PE 预计算替代动态生成原始代码dense_pe self.prompt_encoder.get_dense_pe()→ 触发torch.arange→ ONNX 不支持改为# 在 __init__ 中预计算并注册为 buffer self.register_buffer(pe_1024, self.prompt_encoder.get_dense_pe(), persistentFalse) # 在 forward 中直接复用 dense_pe self.pe_1024 # shape: [1, 256, 64, 64]Point prompt 编码去 control flow原始代码根据input_labels值分支处理点/框/掩码改为强制只接受input_pointsN×2 tensorinput_labels固定为全 1表示 foregroundinput_boxes设为 None# 在 forward 中 if input_points is not None: sparse_embeddings self.prompt_encoder(points(input_points, input_labels)) else: sparse_embeddings torch.zeros(1, 0, 256, deviceinput_images.device) # dummyMask decoder 输出标准化原始predict_masks()返回(masks, iou_preds, low_res_masks)其中masksshape 为[B, N, H, W]但 ONNX 要求明确维度名改为# 在 wrapper.forward 最终 return return { masks: masks, # [1, 3, 256, 256] —— 固定 3 个 mask 输出 iou_preds: iou_preds, # [1, 3] low_res_masks: low_res_masks # [1, 3, 256, 256] }3.3 执行导出sam2_onnx_exporter.py 的最小可运行命令与参数含义python sam2_onnx_exporter.py \ --model-type sam2_hiera_t \ --checkpoint models/sam2_hiera_tiny.pt \ --output-dir onnx/ \ --image-size 1024 \ --max-points 32 \ --opset-version 17 \ --quantize-int8 False参数详解--model-type必须与 checkpoint 匹配可选sam2_hiera_ttiny、sam2_hiera_ssmall、sam2_hiera_bbase、sam2_hiera_llarge。tiny 版本 encoder 仅 12M 参数适合快速验证。--image-size必须等于训练时的输入尺寸SAM2 官方全部为 1024否则 PE buffer shape 错误。--max-points决定 prompt encoder 输入张量的第二维。设为 32 意味着你最多传入 32 个点不足则 pad超则 trunc。--opset-version 17强制指定opset 16 不支持scaled_dot_product_attentionopset 18 在某些 onnxruntime 版本中不稳定。--quantize-int8 False先确保 FP32 模型能跑通再开启量化。INT8 量化需额外安装onnxruntime-tools且会损失约 1.2% mIoU在 COCO-Val 上测试。导出成功后onnx/目录下会生成三个文件sam2_tiny_image_encoder.onnx约 18MBsam2_tiny_prompt_encoder.onnx约 2.1MBsam2_tiny_mask_decoder.onnx约 4.7MB提示不要尝试合并这三个 ONNX 文件SAM2 的三段式设计天然适合 pipeline 推理——encoder 输出 feature map → prompt encoder 输出 sparse embedding → decoder 融合二者输出 mask。强行合并会导致输入输出接口混乱且无法单独替换某一段如只升级 decoder。4. 避坑指南ONNX-SAM2-Segment-Anything.zip 中最常翻车的 5 个血泪现场ONNX-SAM2-Segment-Anything.zip 的作者不是神仙他踩过的坑都明明白白写在注释里。但新手往往忽略这些 warning直到 inference 时 mask 全黑、iou 为 nan、或者 runtime 直接 segfault。以下是我在 3 个工业项目中复现并验证过的 5 个致命坑每一条都附带现象、根因和可复制的修复命令。4.1 现象onnxruntime.capi.onnxruntime_pybind11_state.InvalidArgument: Input with name input_points has invalid shape原因input_points输入张量 shape 应为[1, max_points, 2]但用户传入[N, 2]N 为实际点数ONNX runtime 拒绝 reshape。SAM2 的 wrapper 未做自动 pad而是严格校验。解决在sam2_onnx_inference.py的run_inference()函数开头插入# 确保 input_points 是 [1, max_points, 2] if input_points.ndim 2 and input_points.shape[1] 2: input_points input_points.unsqueeze(0) # [N,2] - [1,N,2] pad_len self.max_points - input_points.shape[1] if pad_len 0: pad_tensor torch.zeros(1, pad_len, 2, dtypeinput_points.dtype, deviceinput_points.device) input_points torch.cat([input_points, pad_tensor], dim1) elif pad_len 0: input_points input_points[:, :self.max_points, :]4.2 现象mask 输出全为 0 或全为 1sigmoid 后无中间值原因ONNX 导出时未冻结 batch norm导致BatchNorm2d在 eval 模式下仍使用 running_mean/std而 ONNX runtime 不执行 BN 的统计更新逻辑输出 logits 偏移。解决在sam2_onnx_wrapper.py的__init__中对所有 BN 层显式设置for m in self.modules(): if isinstance(m, torch.nn.BatchNorm2d): m.eval() # 强制 eval m.weight.requires_grad False m.bias.requires_grad False4.3 现象RuntimeError: Expected all tensors to be on the same device原因sam2_onnx_inference.py中ort_session.run()返回的 numpy array 默认在 CPU但后续cv2.resize或torch.tensor()操作未指定 device导致 tensor 混布。解决统一在推理后转 torch tensor 并指定 device# 在 run_inference() 中 outputs ort_session.run(None, inputs) masks torch.from_numpy(outputs[0]).to(devicecpu) # 显式指定 cpu masks torch.sigmoid(masks) # 避免在 gpu 上做 sigmoid 再转 cpu4.4 现象onnxruntime.capi.onnxruntime_pybind11_state.Fail: Non-zero status code returned while running Split node原因ONNX opset 17 中Split算子要求split属性必须是 int list但某些 PyTorch 版本导出时写成了 float list如[1.0, 1.0, 1.0]。解决用onnx.utils.polish_model()修复需安装onnx-simplifierpip install onnx-simplifier python -m onnxsim onnx/sam2_tiny_mask_decoder.onnx onnx/sam2_tiny_mask_decoder_sim.onnx然后在sam2_onnx_inference.py中加载sam2_tiny_mask_decoder_sim.onnx。4.5 现象视频模式下 memory token 积累错误第 3 帧开始 mask 消失原因sam2_onnx_wrapper.py中未实现 memory token 的跨帧传递逻辑每次forward都重置 memory导致 decoder 无法利用历史信息。解决在 wrapper 中添加self.memory_tokens属性并在forward中# 若是 video modememory_tokens 作为额外输入 if memory_tokens is not None: self.memory_tokens memory_tokens # [1, T, C] else: # 第一帧初始化为空 self.memory_tokens torch.zeros(1, 0, 256, deviceinput_images.device) # 在 decoder 调用时传入 mask_outputs self.mask_decoder( image_embeddingsimage_embeddings, image_pedense_pe, sparse_prompt_embeddingssparse_embeddings, dense_prompt_embeddingsdense_embeddings, multimask_outputTrue, memory_tokensself.memory_tokens # 关键 )5. 推理加速实战如何用 ONNX Runtime 的 Execution Provider 和 Session Options 挤出最后 15% 性能导出 ONNX 模型只是第一步真正影响落地效果的是推理时的 session 配置。sam2_onnx_inference.py默认用 CPU provider但在 Jetson Orin 或 RTX 4090 上不启用 CUDA EP 就是浪费硬件。更关键的是SAM2 的 decoder 有大量 small tensor ops如Add,Mul,Sigmoid默认 session 会频繁 host-device copy必须用 graph optimization 和 memory pattern 调优。5.1 必开的 4 个 Session Options让 ONNX Runtime 不再“傻跑”在sam2_onnx_inference.py初始化InferenceSession时必须传入SessionOptionsimport onnxruntime as ort so ort.SessionOptions() so.graph_optimization_level ort.GraphOptimizationLevel.ORT_ENABLE_EXTENDED # 启用所有图优化 so.intra_op_num_threads 0 # 0 表示使用系统逻辑核数非线程数避免线程争抢 so.execution_mode ort.ExecutionMode.ORT_SEQUENTIAL # SAM2 是串行 pipeline不用 PARALLEL so.add_session_config_entry(session.use_env_allocator, 1) # 启用内存池减少 malloc/free # 关键启用 memory pattern对 SAM2 这种固定 shape 模型提升显著 so.add_session_config_entry(session.allow_mem_pattern, 1)提示allow_mem_pattern1是 ONNX Runtime 1.16 新增选项它会记录第一次 run 的内存分配 pattern后续 run 复用同一块 memory避免重复 allocation。SAM2 的输入 shape 固定1024×1024 图像 32 点 prompt此选项可降低 12%~18% latency。5.2 Execution Provider 选择策略CUDA vs CPU vs TensorRT何时该切场景推荐 EP理由验证命令RTX 4090 / A100LinuxCUDAExecutionProviderFP16 自动启用decoder 中的 MatMul 可加速 3.2×ort_session ort.InferenceSession(model_path, providers[CUDAExecutionProvider])Jetson OrinUbuntu 20.04CUDAExecutionProviderTensorrtExecutionProvider混合TRT 对 Conv/BN 优化强但对 SAM2 的 Attention kernel 支持不全混合模式最稳providers[(TensorrtExecutionProvider, {trt_fp16_enable: True}), (CUDAExecutionProvider, {})]Windows 11 i7-12800HCPUExecutionProviderOpenVINOExecutionProviderOpenVINO 对 ViT encoder 的 patch embedding 有专项优化pip install openvino然后providers[OpenVINOExecutionProvider]验证 EP 是否生效print(ort_session.get_providers()) # 应输出 [CUDAExecutionProvider] 而非 [CPUExecutionProvider] print(ort_session.get_inputs()[0].shape) # 应显示 [1, 3, 1024, 1024]而非 [-1, ...]5.3 INT8 量化实操不掉点的量化参数配置表SAM2 的 ONNX INT8 量化不是onnxruntime.quantization.quantize_static一行搞定。由于 decoder 输出 logits 范围窄-5 ~ 5直接 per-channel 量化会丢失细节。必须用per-tensor asymmetric custom calibration data。参数推荐值说明calibrate_methodMinMaxCalibrater不用EntropyCalibraterSAM2 的 logits 分布偏态严重activation_typeQuantType.QUInt8激活用无符号避免负数截断weight_typeQuantType.QInt8权重用有符号保留方向性extra_options{WeightSymmetric: False, ActivationSymmetric: False}必须关闭对称量化否则 sigmoid 前的 logits 截断严重calibration_dataset自定义 100 张 COCO val 图像 随机点 prompt不能用 imagenet必须用分割任务数据分布量化命令需onnxruntime-toolspython -m onnxruntime.quantization.calibrate \ --input onnx/sam2_tiny_mask_decoder.onnx \ --output onnx/sam2_tiny_mask_decoder_quant.onnx \ --calibrate_dataset ./calib_data/ \ --quant_format QDQ \ --per_channel False \ --symmetric False量化后实测FP32 推理 124ms → INT8 推理 78msRTX 4090mIoU 下降 0.9%COCO-Val完全可接受。6. 验证你的 ONNX-SAM2 模型是否真的“能用”三步验证法与工业级交付 checklist很多工程师导出 ONNX 后只跑了一张图、看 mask 有输出就认为成功。但工业交付要的是1000 帧视频连续跑不崩、不同光照条件 mask 稳定、客户给的任意尺寸图像自动 resize 不失真、CPU 占用率低于 70%。我用这套三步验证法在 3 个交付项目中零返工——它不依赖 fancy 指标只问最朴素的问题它敢不敢上产线6.1 Step 1单帧压力测试Stress Test——用 100 张图连续跑看内存是否泄漏写一个stress_test.py加载 ONNX 模型后循环推理 100 次监控内存增长import psutil import time process psutil.Process() start_mem process.memory_info().rss / 1024 / 1024 # MB for i in range(100): # 随机生成 1024x1024 图像 16 个点 img np.random.randint(0, 256, (1024, 1024, 3), dtypenp.uint8) points np.random.randint(0, 1024, (16, 2)) masks inference_engine.run_inference(img, points) if i % 10 0: curr_mem process.memory_info().rss / 1024 / 1024 print(fiter {i}: mem {curr_mem:.1f} MB (delta: {curr_mem - start_mem:.1f})) # ✅ 通过标准100 次后内存 delta 5 MB # ❌ 翻车信号delta 20 MB → 说明 tensor 未释放检查 onnxruntime session 是否复用、numpy array 是否 detach6.2 Step 2跨尺寸鲁棒性测试Robustness Test——验证 resize 逻辑是否真“自适应”SAM2 官方要求输入 1024×1024但客户现场图像是 1920×1080、400×300、甚至 3000×2000。你的preprocess.py必须做到长边缩放到 1024短边等比缩放不 crop 不 stretchpadding 用cv2.copyMakeBorder填充 0非 mean因为 SAM2 encoder 的 normalization 是ImageNetmean/std填 0 不影响point prompt 坐标必须按相同比例缩放并 round 到整数ONNX 输入要求 int64测试脚本robustness_test.pytest_sizes [(1920, 1080), (400, 300), (3000, 2000), (100, 100)] for h, w in test_sizes: img np.random.randint(0, 256, (h, w, 3), dtypenp.uint8) orig_points np.array([[w//2, h//2], [w//4, h//4]]) # 原图坐标 # 调用你的 preprocess 函数 proc_img, proc_points, (pad_h, pad_w) preprocess_image_and_points(img, orig_points) # 断言proc_img.shape (1024, 1024, 3) # 断言proc_points 应在 [0,1024) 范围内且保持相对位置 assert proc_img.shape (1024, 1024, 3) assert proc_points.min() 0 and proc_points.max() 10246.3 Step 3视频流时延测试Latency Test——用 cv2.VideoCapture 模拟真实 pipeline这才是交付验收的终极考题。写video_latency_test.py用cv2.VideoCapture(0)读 webcam每帧加 3 个随机点测 end-to-end latencycap cv2.VideoCapture(0) cap.set(cv2.CAP_PROP_FRAME_WIDTH, 1280) cap.set(cv2.CAP_PROP_FRAME_HEIGHT, 720) latencies [] for i in range(200): # 200 帧 ret, frame cap.read() if not ret: break start_time time.time() # 随机选 3 个点模拟用户点击 h, w frame.shape[:2] points np.array([[np.random.randint(0,w), np.random.randint(0,h)] for _ in range(3)]) # 推理 masks inference_engine.run_inference(frame, points) end_time time.time() latencies.append(end_time - start_time) cap.release() avg_latency np.mean(latencies) * 1000 # ms fps 1000 / avg_latency print(fAVG Latency: {avg_latency:.1f}ms ({fps:.1f} FPS)) # ✅ 交付标准Jetson Orin 上 avg_latency 180ms5.5 FPSRTX 4090 上 85ms11.7 FPS我的习惯每次交付前把这三步测试写成ci_test.sh放进 GitLab CIpush 一次就自动跑。曾经有个项目开发说“模型没问题”CI 却在 Step 3 报latency 300ms一查发现他把onnxruntime.InferenceSession放在 for 循环里重建——session 初始化耗时 120ms占了大头。真正的交付不是“能跑”而是“敢压测、敢连跑、敢上视频流”。希望帮到你。本文还有配套的精品资源点击获取