尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

SAM2模型Python+ONNX部署实战:从PyTorch到高效推理全流程

SAM2模型Python+ONNX部署实战:从PyTorch到高效推理全流程 简介模型部署是深度学习从研究走向应用的关键环节其核心目标在于将训练好的模型高效、稳定地集成到生产环境中。ONNX作为一种开放的模型表示格式配合ONNX Runtime高性能推理引擎能够实现跨平台、跨硬件的统一部署有效解决框架锁定和性能瓶颈问题。这一技术组合在计算机视觉、自然语言处理等领域具有极高的工程价值尤其适用于需要低延迟、高吞吐的实时应用场景。本文聚焦于Segment Anything 2这一前沿的零样本图像分割模型针对其庞大的参数量和复杂的动态输入结构详细拆解了如何通过模型拆分、动态轴设置和缓存优化等策略将其从PyTorch框架成功转换为轻量化的ONNX格式并封装为易于集成的Python API为视觉大模型的落地提供了完整的工程实践路径。1. 项目缘起为什么选择PythonONNX来部署SAM2最近在做一个需要高精度图像分割的项目从SAM到SAM2Meta的Segment Anything系列模型确实在零样本分割能力上带来了质的飞跃。但模型能力越强部署的“甜蜜烦恼”就越多。原始的PyTorch模型动辄几个G推理速度也受限于框架本身直接用在生产环境或者集成到客户端应用里无论是资源占用还是响应延迟都让人头疼。这时候ONNX Runtime就成了一个非常自然的选择。它就像一个“模型翻译官”和“高效执行引擎”能把PyTorch、TensorFlow等框架训练出来的模型转换成一种中间表示格式ONNX然后在一个高度优化的运行时里执行。这样做的好处显而易见模型体积显著减小、推理速度大幅提升、并且能跨平台Windows/Linux/macOS和跨硬件CPU/GPU运行。对于SAM2这种视觉大模型部署的轻量化和加速是刚需。所以这个项目的核心目标就很明确了将庞大的SAM2 PyTorch模型通过ONNX转换和优化变成一个轻量、快速、易于集成的推理模块并用Python封装成一套清晰可用的API。整个过程会涉及到模型导出、动态轴处理、后处理优化等一系列实战细节这也是很多朋友从研究走向落地时最容易卡壳的地方。接下来我就把完整的流程、踩过的坑以及优化技巧毫无保留地分享出来。2. 环境准备与模型获取搭建稳定可复现的基础万事开头难一个稳定、干净的环境是后续所有操作成功的前提。这里我强烈建议使用Conda来管理Python环境它能很好地解决不同项目间依赖冲突的问题。2.1 创建并激活专用Conda环境首先我们创建一个名为sam2_onnx的Python 3.9环境3.8-3.10都是比较稳定的选择。conda create -n sam2_onnx python3.9 -y conda activate sam2_onnx2.2 安装核心依赖库激活环境后安装PyTorch。请务必根据你是否有CUDA以及CUDA版本去 PyTorch官网 获取正确的安装命令。这里以CUDA 11.8为例pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118接着安装ONNX相关的工具链和SAM2的官方库pip install onnx onnxruntime-gpu # 如果只用CPU则安装 onnxruntime pip install githttps://github.com/facebookresearch/segment-anything-2.git pip install opencv-python pillow matplotlib注意onnxruntime-gpu和onnxruntime是互斥的只能安装一个。如果你的机器有NVIDIA GPU且配置好了CUDA安装-gpu版本会获得巨大的加速。可以通过import onnxruntime as ort; print(ort.get_device())来验证是否成功识别到GPU。2.3 下载SAM2预训练模型SAM2提供了多个规模的模型。对于部署我们通常从sam2_hiera_large.yaml配置对应的模型开始它在精度和速度上有一个较好的平衡。模型文件可以从Meta官方或Hugging Face等镜像站获取。假设我们下载好的模型权重文件为sam2_hiera_large.pt。import torch from sam2.build_sam import build_sam2 # 加载模型配置和权重 model_config path/to/sam2_hiera_large.yaml checkpoint path/to/sam2_hiera_large.pt sam2_model build_sam2(model_config, checkpoint) sam2_model.eval() # 切换到评估模式至此基础环境就搭建好了我们手里也有了一个可以运行的PyTorch版SAM2。3. 核心攻坚将PyTorch模型转换为ONNX格式这是整个流程中最关键、也最容易出错的一步。SAM2模型结构复杂包含Vision Transformer (ViT)作为图像编码器和一个轻量级的掩码解码器输入输出都是动态的可变尺寸的图像和可变数量的提示点。3.1 理解SAM2的输入输出结构在转换前必须彻底弄清楚模型的前向传播需要什么输出什么。SAM2的典型推理流程是图像编码器输入一张图片输出对应的图像嵌入Image Embedding。这部分计算量大但对于同一张图片嵌入只需计算一次并可以缓存。提示编码器 掩码解码器输入图像嵌入、以及用户提供的提示如点、框、文本输出分割掩码、关联的置信度分数等。为了部署高效一个常见的策略是将图像编码器和掩码解码器分开导出为两个ONNX模型。这样当需要对同一张图片进行多次交互式分割时昂贵的图像嵌入只需计算并输入一次。3.2 导出图像编码器为ONNX图像编码器的输入是经过预处理的图像张量输出是固定维度的图像嵌入。import torch.onnx # 假设我们有一个图像预处理函数 preprocess_image def preprocess_image(image_path): import cv2 import numpy as np from torchvision import transforms # 读取图像调整大小归一化等 img cv2.imread(image_path) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) # SAM2通常需要将图像填充到特定尺寸如1024x1024 # 这里简化处理实际需要根据模型配置调整 input_tensor transform(img).unsqueeze(0) # [1, 3, H, W] return input_tensor # 获取一个示例输入 dummy_image torch.randn(1, 3, 1024, 1024).to(cuda) # 示例张量 sam2_model.image_encoder.to(cuda) # 导出图像编码器 onnx_encoder_path sam2_image_encoder.onnx torch.onnx.export( sam2_model.image_encoder, # 要导出的模型子模块 dummy_image, # 模型输入示例 onnx_encoder_path, # 输出文件路径 input_names[image], # 输入节点名 output_names[image_embedding], # 输出节点名 dynamic_axes{ image: {2: height, 3: width}, # 动态轴高度和宽度可变 image_embedding: {2: embed_height, 3: embed_width} }, opset_version14, # ONNX算子集版本建议14 do_constant_foldingTrue # 优化常量折叠 ) print(f图像编码器已导出至: {onnx_encoder_path})关键点解析dynamic_axes这是处理可变尺寸输入的核心。我们声明了输入图像的高度(2)和宽度(3)维度是动态的。ONNX模型将能接受任意符合[1, 3, H, W]格式的输入。opset_versionONNX算子集版本。版本越高支持的算子越多但也要考虑部署环境如某些移动端推理引擎的支持情况。14是一个广泛兼容的版本。do_constant_folding启用常量折叠优化可以简化计算图提升推理速度。3.3 导出掩码解码器为ONNX掩码解码器的输入更复杂图像嵌入、提示点坐标、提示点标签前景/背景、以及可选的掩码输入。# 准备掩码解码器的示例输入 image_embedding sam2_model.image_encoder(dummy_image) # 获取真实的图像嵌入作为示例 dummy_image_embedding image_embedding # 假设有2个提示点 (x, y) 坐标归一化到[0,1] dummy_point_coords torch.tensor([[[0.3, 0.5], [0.7, 0.5]]], dtypetorch.float32).to(cuda) # 对应的标签1表示前景点0表示背景点 dummy_point_labels torch.tensor([[1, 1]], dtypetorch.float32).to(cuda) # 初始掩码输入可选通常为None dummy_mask_input torch.zeros(1, 1, 256, 256).to(cuda) # 示例 # 是否有掩码输入的标志 dummy_has_mask_input torch.tensor([0], dtypetorch.float32).to(cuda) # 将掩码解码器移到GPU并导出 sam2_model.mask_decoder.to(cuda) onnx_decoder_path sam2_mask_decoder.onnx # 注意这里需要根据SAM2实际的前向函数签名来组织输入。 # SAM2的mask_decoder可能是一个更复杂的调用。以下是一个示意实际需要查阅源码。 # 假设其forward函数签名是forward(image_embedding, point_coords, point_labels, mask_input, has_mask_input) input_tuple (dummy_image_embedding, dummy_point_coords, dummy_point_labels, dummy_mask_input, dummy_has_mask_input) torch.onnx.export( sam2_model.mask_decoder, input_tuple, onnx_decoder_path, input_names[image_embedding, point_coords, point_labels, mask_input, has_mask_input], output_names[masks, iou_predictions, low_res_masks], dynamic_axes{ image_embedding: {2: embed_h, 3: embed_w}, point_coords: {1: num_points}, # 提示点数量可变 point_labels: {1: num_points}, # mask_input 和 has_mask_input 通常是固定的 masks: {1: num_output_masks}, # 输出掩码数量可变与提示相关 }, opset_version14, do_constant_foldingTrue ) print(f掩码解码器已导出至: {onnx_decoder_path})踩坑实录动态轴与输入对齐这里最大的坑在于dynamic_axes的定义和实际输入数据的维度必须严格对齐。例如point_coords的维度是[batch, num_points, 2]那么动态轴{1: num_points}就表示第1维从0开始计数num_points是可变的。在后续用ONNX Runtime推理时你传入的point_coords的num_points值可以是任意正整数但batch和最后的2必须固定。如果定义错了推理时会报维度不匹配的错误。4. ONNX模型优化与验证确保转换正确性与性能导出的ONNX模型可能包含冗余算子或未被优化的结构。直接使用可能效率不高甚至存在错误。4.1 使用ONNX Runtime工具进行优化ONNX Runtime提供了一个强大的图形优化工具onnxruntime.tools.optimize_onnx_model。import onnx from onnxruntime.tools import optimize_onnx_model # 加载原始模型 onnx_encoder_model onnx.load(onnx_encoder_path) # 进行优化包括常量折叠、算子融合、冗余节点消除等 optimized_encoder_model optimize_onnx_model(onnx_encoder_model) # 保存优化后的模型 optimized_encoder_path sam2_image_encoder_optimized.onnx onnx.save(optimized_encoder_model, optimized_encoder_path) # 对解码器执行同样的操作 onnx_decoder_model onnx.load(onnx_decoder_path) optimized_decoder_model optimize_onnx_model(onnx_decoder_model) optimized_decoder_path sam2_mask_decoder_optimized.onnx onnx.save(optimized_decoder_model, optimized_decoder_path)4.2 验证模型正确性ONNX Runtime推理比对这是至关重要的一步确保ONNX模型和原始PyTorch模型的输出在数值上基本一致允许微小的浮点误差。import onnxruntime as ort import numpy as np # 1. 创建ONNX Runtime推理会话 ort_session_encoder ort.InferenceSession(optimized_encoder_path, providers[CUDAExecutionProvider, CPUExecutionProvider]) ort_session_decoder ort.InferenceSession(optimized_decoder_path, providers[CUDAExecutionProvider, CPUExecutionProvider]) # 2. 准备输入数据 (numpy格式) dummy_image_np dummy_image.cpu().numpy() dummy_image_embedding_np dummy_image_embedding.cpu().detach().numpy() dummy_point_coords_np dummy_point_coords.cpu().numpy() dummy_point_labels_np dummy_point_labels.cpu().numpy() dummy_mask_input_np dummy_mask_input.cpu().numpy() dummy_has_mask_input_np dummy_has_mask_input.cpu().numpy() # 3. 运行ONNX编码器推理 ort_inputs_encoder {ort_session_encoder.get_inputs()[0].name: dummy_image_np} ort_outs_encoder ort_session_encoder.run(None, ort_inputs_encoder) onnx_image_embedding ort_outs_encoder[0] # 4. 运行ONNX解码器推理 ort_inputs_decoder { ort_session_decoder.get_inputs()[0].name: dummy_image_embedding_np, ort_session_decoder.get_inputs()[1].name: dummy_point_coords_np, ort_session_decoder.get_inputs()[2].name: dummy_point_labels_np, ort_session_decoder.get_inputs()[3].name: dummy_mask_input_np, ort_session_decoder.get_inputs()[4].name: dummy_has_mask_input_np, } ort_outs_decoder ort_session_decoder.run(None, ort_inputs_decoder) onnx_masks, onnx_iou, onnx_low_res ort_outs_decoder # 5. 运行PyTorch推理进行比对 (使用之前已经计算过的image_embedding) with torch.no_grad(): # 注意这里需要调用SAM2完整的推理流程来获取可对比的输出 # 假设有一个函数 predict_masks 封装了PyTorch模型的完整调用 # pytorch_masks, pytorch_iou sam2_model.predict(...) # 此处简化用解码器前向代替 pytorch_outputs sam2_model.mask_decoder(dummy_image_embedding, dummy_point_coords, dummy_point_labels, dummy_mask_input, dummy_has_mask_input) pytorch_masks pytorch_outputs[0].cpu().numpy() # 6. 比较输出 (以掩码为例) def compare_tensors(name, onnx_tensor, pytorch_tensor, rtol1e-03, atol1e-05): onnx_flat onnx_tensor.flatten() pytorch_flat pytorch_tensor.flatten() diff np.abs(onnx_flat - pytorch_flat) max_diff diff.max() mean_diff diff.mean() print(f{name} - 最大差异: {max_diff:.6f}, 平均差异: {mean_diff:.6f}) if np.allclose(onnx_tensor, pytorch_tensor, rtolrtol, atolatol): print(f ✅ {name} 输出匹配) return True else: print(f ❌ {name} 输出存在显著差异) return False compare_tensors(掩码输出, onnx_masks, pytorch_masks)如果验证通过恭喜你ONNX模型转换成功如果差异过大需要回头检查导出时的动态轴设置、输入数据预处理、以及模型是否处于eval()模式。5. 构建完整的Python部署管道有了优化验证好的ONNX模型我们就可以构建一个完整的、用户友好的推理管道了。这个管道要处理图像预处理、提示编码、ONNX Runtime会话管理、后处理等所有环节。5.1 设计推理类SAM2ONNXInference我们将核心功能封装到一个类里这样使用起来更清晰。import cv2 import numpy as np import torch import onnxruntime as ort from typing import List, Tuple, Optional, Union class SAM2ONNXInference: def __init__(self, encoder_model_path: str, decoder_model_path: str, device: str cuda): 初始化SAM2 ONNX推理器。 Args: encoder_model_path: 图像编码器ONNX模型路径 decoder_model_path: 掩码解码器ONNX模型路径 device: 推理设备cuda 或 cpu self.device device providers [CUDAExecutionProvider, CPUExecutionProvider] if device cuda else [CPUExecutionProvider] # 创建ONNX Runtime会话 self.encoder_session ort.InferenceSession(encoder_model_path, providersproviders) self.decoder_session ort.InferenceSession(decoder_model_path, providersproviders) # 获取输入输出信息 self.encoder_input_name self.encoder_session.get_inputs()[0].name self.decoder_input_names [inp.name for inp in self.decoder_session.get_inputs()] # 图像预处理参数 (需与训练时一致) self.img_size 1024 # SAM2通常的输入尺寸 self.pixel_mean np.array([123.675, 116.28, 103.53]) self.pixel_std np.array([58.395, 57.12, 57.375]) self._cached_image_embedding None self._cached_original_size None def preprocess_image(self, image: np.ndarray) - Tuple[np.ndarray, dict]: 预处理输入图像返回模型输入张量和元信息。 Args: image: BGR格式的numpy数组 (H, W, 3) Returns: input_tensor: 预处理后的张量 [1, 3, self.img_size, self.img_size] transform_info: 包含缩放、填充等信息的字典用于后处理时坐标映射 # 转换颜色空间 img_rgb cv2.cvtColor(image, cv2.COLOR_BGR2RGB) original_h, original_w img_rgb.shape[:2] # 计算缩放比例保持长宽比 scale self.img_size / max(original_h, original_w) new_h, new_w int(original_h * scale), int(original_w * scale) img_resized cv2.resize(img_rgb, (new_w, new_h), interpolationcv2.INTER_LINEAR) # 填充到正方形 top (self.img_size - new_h) // 2 bottom self.img_size - new_h - top left (self.img_size - new_w) // 2 right self.img_size - new_w - left img_padded cv2.copyMakeBorder(img_resized, top, bottom, left, right, cv2.BORDER_CONSTANT, value0) # 归一化 (减去均值除以标准差) img_normalized (img_padded - self.pixel_mean) / self.pixel_std # 调整维度顺序为 [C, H, W] 并添加批次维度 - [1, C, H, W] input_tensor img_normalized.transpose(2, 0, 1)[np.newaxis, ...].astype(np.float32) transform_info { original_size: (original_h, original_w), resized_size: (new_h, new_w), padding: (top, bottom, left, right), # (上下左右) scale: scale } return input_tensor, transform_info def encode_image(self, image: np.ndarray) - np.ndarray: 编码单张图像返回图像嵌入并缓存。 input_tensor, transform_info self.preprocess_image(image) # ONNX推理 ort_inputs {self.encoder_input_name: input_tensor} image_embedding self.encoder_session.run(None, ort_inputs)[0] self._cached_image_embedding image_embedding self._cached_transform_info transform_info return image_embedding def predict_masks( self, point_coords: Optional[List[List[float]]] None, point_labels: Optional[List[int]] None, box: Optional[List[float]] None, # [x1, y1, x2, y2] 归一化坐标 image_embedding: Optional[np.ndarray] None, transform_info: Optional[dict] None ) - Tuple[np.ndarray, np.ndarray]: 根据提示预测分割掩码。 Args: point_coords: 提示点坐标列表每个坐标是[x, y] (归一化到[0,1]相对于原始图像) point_labels: 对应点的标签1为前景0为背景 box: 提示框归一化坐标 [x1, y1, x2, y2] image_embedding: 可选的图像嵌入若为None则使用缓存的 transform_info: 可选的变换信息若为None则使用缓存的 Returns: masks: 预测的掩码 [N, H, W], N为输出掩码数量 scores: 对应的IoU置信度分数 [N] if image_embedding is None: if self._cached_image_embedding is None: raise ValueError(未提供image_embedding且无缓存请先调用encode_image。) image_embedding self._cached_image_embedding transform_info self._cached_transform_info # 1. 准备提示输入 # 将原始图像坐标转换为模型输入空间坐标 all_point_coords [] all_point_labels [] if point_coords is not None and point_labels is not None: for (x, y), label in zip(point_coords, point_labels): # 坐标转换原始图像 - 填充后图像 - 模型输入空间 # 首先根据缩放和填充将原始坐标映射到1024x1024输入上的坐标 scale transform_info[scale] top, bottom, left, right transform_info[padding] # 原始坐标归一化到[0,1]先乘原始尺寸得到像素坐标再应用变换 orig_h, orig_w transform_info[original_size] px, py x * orig_w, y * orig_h # 转为像素坐标 # 缩放 px_scaled, py_scaled px * scale, py * scale # 加上填充偏移 px_final, py_final px_scaled left, py_scaled top # 归一化到模型输入空间 [0, 1] x_final, y_final px_final / self.img_size, py_final / self.img_size all_point_coords.append([x_final, y_final]) all_point_labels.append(label) # 处理框提示将框转换为4个角点前景点 if box is not None: x1, y1, x2, y2 box # 将框的左上角和右下角作为两个点加入也可以加入四个角点这里简化 all_point_coords.extend([[x1, y1], [x2, y2]]) all_point_labels.extend([2, 3]) # 使用特殊标签表示框点SAM2解码器能识别 if not all_point_coords: raise ValueError(必须提供至少一种提示点或框。) # 转换为numpy数组并调整维度 point_coords_np np.array([all_point_coords], dtypenp.float32) # [1, N, 2] point_labels_np np.array([all_point_labels], dtypenp.float32) # [1, N] # 2. 准备其他解码器输入掩码输入和标志位这里用默认值 mask_input_np np.zeros((1, 1, 256, 256), dtypenp.float32) has_mask_input_np np.array([0], dtypenp.float32) # 3. ONNX解码器推理 ort_inputs { self.decoder_input_names[0]: image_embedding, self.decoder_input_names[1]: point_coords_np, self.decoder_input_names[2]: point_labels_np, self.decoder_input_names[3]: mask_input_np, self.decoder_input_names[4]: has_mask_input_np, } ort_outs self.decoder_session.run(None, ort_inputs) masks, iou_predictions, low_res_masks ort_outs # 4. 后处理选择最佳掩码并映射回原始图像尺寸 # 通常选择置信度最高的掩码 best_idx np.argmax(iou_predictions) best_mask masks[0, best_idx] # masks shape: [1, N, 256, 256] best_score iou_predictions[0, best_idx] # 将256x256的掩码上采样并去除填充区域映射回原始图像尺寸 final_mask self._postprocess_mask(best_mask, transform_info) return final_mask, best_score def _postprocess_mask(self, mask: np.ndarray, transform_info: dict) - np.ndarray: 将模型输出的掩码后处理为原始图像尺寸的二值掩码。 # 1. 上采样到模型输入尺寸 (1024x1024) mask_1024 cv2.resize(mask, (self.img_size, self.img_size), interpolationcv2.INTER_LINEAR) # 2. 裁剪掉填充区域 top, bottom, left, right transform_info[padding] if top bottom left right 0: mask_cropped mask_1024[top:self.img_size-bottom, left:self.img_size-right] else: mask_cropped mask_1024 # 3. 缩放到原始图像尺寸 original_h, original_w transform_info[original_size] mask_original cv2.resize(mask_cropped, (original_w, original_h), interpolationcv2.INTER_LINEAR) # 4. 二值化 (通常以0.0为阈值) binary_mask (mask_original 0.0).astype(np.uint8) * 255 return binary_mask这个类封装了从图像预处理到掩码后处理的完整流程。使用方式非常直观# 初始化推理器 inferencer SAM2ONNXInference( encoder_model_pathsam2_image_encoder_optimized.onnx, decoder_model_pathsam2_mask_decoder_optimized.onnx, devicecuda ) # 读取图像 image cv2.imread(test_image.jpg) # 编码图像只需一次 image_embedding inferencer.encode_image(image) # 定义提示点例如用户点击了图像上的某个物体 # 坐标是归一化的 [x, y] 范围[0, 1] point_coords [[0.5, 0.5]] # 图像中心点 point_labels [1] # 1表示前景点 # 预测掩码 mask, score inferencer.predict_masks(point_coordspoint_coords, point_labelspoint_labels) print(f预测掩码的IoU分数: {score:.3f}) # 可视化结果 cv2.imshow(Original Image, image) cv2.imshow(Predicted Mask, mask) cv2.waitKey(0)5.2 性能优化技巧与注意事项在实际部署中性能至关重要。这里分享几个关键优化点会话Session复用与预热ort.InferenceSession的创建开销较大。务必在应用初始化时创建一次然后在整个生命周期内复用。在正式处理请求前可以用一个小的虚拟输入运行一次推理进行“预热”触发图优化和内核初始化。图像嵌入缓存这是SAM系列模型部署的核心优化。对于同一张图片的多次交互如用户多次点击调整image_embedding只需计算一次。我们的SAM2ONNXInference类内部做了缓存。批处理支持上述示例是单张图片推理。ONNX模型本身可以支持批处理但需要确保导出模型时dynamic_axes的第0维批次维也是动态的如{image: {0: batch_size}}。在构建输入时将多张图片堆叠成一个批次张量传入可以显著提升吞吐量。IO绑定与计算绑定对于视频流或实时应用可以将图像预处理、后处理等CPU操作与ONNX Runtime的GPU推理异步进行使用生产者-消费者队列模式最大化GPU利用率。模型量化如果对精度损失有一定容忍度例如从FP32到INT8可以使用ONNX Runtime的量化工具对模型进行量化能进一步减少模型体积和提升推理速度尤其有利于边缘设备部署。6. 项目源码结构与使用指南一个完整的可交付项目除了核心代码清晰的目录结构和说明文档必不可少。以下是一个建议的项目结构sam2_onnx_deployment/ ├── README.md # 项目总说明快速开始指南 ├── requirements.txt # Python依赖列表 ├── configs/ # 配置文件目录 │ └── sam2_config.yaml # SAM2模型配置文件 (从官方仓库复制) ├── models/ # 模型文件目录 │ ├── sam2_hiera_large.pt # PyTorch预训练权重 (需自行下载) │ ├── sam2_image_encoder.onnx │ ├── sam2_image_encoder_optimized.onnx │ ├── sam2_mask_decoder.onnx │ └── sam2_mask_decoder_optimized.onnx ├── scripts/ # 实用脚本 │ ├── export_onnx.py # 模型导出脚本 │ ├── optimize_onnx.py # 模型优化脚本 │ └── benchmark.py # 性能基准测试脚本 ├── src/ # 核心源代码 │ ├── __init__.py │ ├── inference.py # 包含 SAM2ONNXInference 类 │ ├── preprocess.py # 图像预处理函数 │ └── postprocess.py # 掩码后处理函数 ├── examples/ # 使用示例 │ ├── simple_inference.py # 基础推理示例 │ ├── interactive_demo.py # 交互式演示 (使用OpenCV鼠标回调) │ └── video_processing.py # 视频流处理示例 └── tests/ # 单元测试 └── test_inference.py在README.md中需要详细说明环境配置如何安装Conda、CUDA、依赖库。模型准备在哪里下载预训练权重如何运行导出脚本得到ONNX模型。快速开始提供一个最简单的代码示例让用户能在5分钟内跑通第一个分割结果。API文档对SAM2ONNXInference类的主要方法进行说明。高级用法如何集成到Web服务如FastAPI、如何处理视频、如何进行批量预测等。常见问题列出部署过程中可能遇到的典型错误和解决方案。7. 进阶话题动态提示、批处理与多模态扩展掌握了基础部署后可以探索更高级的应用场景。7.1 处理动态数量的提示我们的解码器模型已经通过dynamic_axes支持了可变数量的提示点。在实际的交互式应用中用户可能点击多个点前景和背景。我们的predict_masks方法已经能够处理任意长度的point_coords列表。关键在于坐标转换逻辑要正确确保每个从原始图像空间点击的坐标都能准确映射到模型输入的空间。7.2 实现批处理推理对于需要处理大量图片或视频帧的场景批处理能极大提升效率。需要对代码进行以下改造模型导出确保导出ONNX时图像编码器和解码器的批次维通常是第0维也被设置为动态轴。预处理preprocess_image函数需要支持批量处理输入从单张图片的[H, W, 3]变为批量的[B, H, W, 3]或一个图片路径列表。推理类修改encode_image和predict_masks使其能接受批量图像和批量提示。注意批处理时每张图片的提示数量可能不同这需要更精细的数据结构如列表的列表或填充padding来处理。后处理相应的后处理函数也需要支持批量操作。7.3 集成文本提示如果SAM2支持如果未来SAM2的官方实现或社区版本集成了强大的CLIP等文本编码器我们的部署管道也可以扩展。核心思路是增加一个文本编码器ONNX模型将文本提示如“一只狗”编码成与图像嵌入对齐的向量然后将这个文本嵌入向量作为额外的输入传递给掩码解码器。这需要修改模型导出步骤和推理类的输入接口。整个PythonONNX部署SAM2的旅程从环境搭建、模型转换、优化验证到封装成健壮的推理管道最后考虑性能优化和功能扩展每一步都充满了工程实践的细节。这套方案不仅适用于SAM2其方法论——模型拆分、动态轴处理、缓存优化、管道封装——对于部署其他复杂的视觉Transformer模型也具有很强的参考价值。最关键的是通过ONNX Runtime我们成功地将一个研究级的模型变成了一个可以在生产环境中高效、稳定运行的服务组件。本文还有配套的精品资源点击获取
返回列表