This tool converts SAM2 (Segment Anything Model 2) PyTorch models to ONNX format specifically optimized for TensorRT deployment. The exported ONNX models are compatible with TensorRT conversion and deployment.
The project requires the following dependencies:
- PyTorch==2.3.0
- hydra-core>=1.3.2
- iopath>=0.1.10
- onnx>=1.14.0
- onnxruntime>=1.15.0
- numpy>=1.24.0
- typing-extensions>=4.5.0
- onnxsim>=0.4.33(optional)
pip install -r requirements.txtThe sam2 folder in this repository is directly copied from the original SAM2 repository with modifications to improve ONNX-to-TensorRT export compatibility. The original SAM2 implementation is licensed under the Apache 2.0 license.
Copyright (c) Meta Platforms, Inc. and affiliates.
All rights reserved.
This source code is licensed under the license found in the
LICENSE file in the root directory of this source tree.
The following modifications were made to the original SAM2 implementation to ensure TensorRT compatibility:
- In
mask_decoder.py, replacedtorch.repeat_interleavewithtorch.tilein two locations:# Original implementation # src = torch.repeat_interleave(image_embeddings, tokens.shape[0], dim=0) # pos_src = torch.repeat_interleave(image_pe, tokens.shape[0], dim=0) # Modified implementation for TensorRT compatibility src = torch.tile(image_embeddings, (tokens.shape[0], 1, 1, 1)) pos_src = torch.tile(image_pe, (tokens.shape[0], 1, 1, 1))
This modification was necessary because torch.repeat_interleave operations can cause issues during TensorRT conversion. The torch.tile operation provides equivalent functionality while maintaining better compatibility with TensorRT.
- sam2.1_hiera_tiny
- sam2.1_hiera_small
- sam2.1_hiera_large
- sam2.1_hiera_base_plus
python export_sam2_onnx.py <model_type> <checkpoint_path> [options]The checkpoint file can be downloaded from the original SAM2 repository.
# Install onnxsim if not already installed
pip install onnxsim
# Simplify the exported models
onnxsim encoder.onnx encoder.onnx
onnxsim decoder.onnx decoder.onnxmodel_type: Type of SAM2 model to export- Choices: sam2.1_hiera_tiny, sam2.1_hiera_small, sam2.1_hiera_large, sam2.1_hiera_base_plus
checkpoint_path: Path to the PyTorch model checkpoint file
--output-dir: Directory to save exported ONNX models (default: ./output)
# Basic usage
python export_sam2_onnx.py sam2.1_hiera_base_plus /path/to/checkpoint.pt
# With custom output directory
python export_sam2_onnx.py sam2.1_hiera_base_plus /path/to/checkpoint.pt --output-dir ./onnx_modelsThe converter will generate two ONNX files in the specified output directory:
<model_type>_encoder.onnx: Encoder model for image feature extraction<model_type>_decoder.onnx: Decoder model for mask prediction
- Optimized for TensorRT conversion and deployment
- Supports dynamic batch size for both encoder and decoder
This project is licensed under the Apache License, Version 2.0. See the LICENSE file for details.