Skip to content

Latest commit

 

History

4 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

SAM2 PyTorch to ONNX Converter

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.

Requirements

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)

Basic Installation

pip install -r requirements.txt

About SAM2

The 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.

Original SAM2 Copyright Notice

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.

Key Modifications

The following modifications were made to the original SAM2 implementation to ensure TensorRT compatibility:

  1. In mask_decoder.py, replaced torch.repeat_interleave with torch.tile in 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.

Supported Models

  • sam2.1_hiera_tiny
  • sam2.1_hiera_small
  • sam2.1_hiera_large
  • sam2.1_hiera_base_plus

Usage

Basic Usage

python export_sam2_onnx.py <model_type> <checkpoint_path> [options]

The checkpoint file can be downloaded from the original SAM2 repository.

Simplify ONNX Model (Optional)

# Install onnxsim if not already installed
pip install onnxsim

# Simplify the exported models
onnxsim encoder.onnx encoder.onnx
onnxsim decoder.onnx decoder.onnx

Required Arguments

  • model_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

Optional Arguments

  • --output-dir: Directory to save exported ONNX models (default: ./output)

Examples

# 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_models

Output Files

The 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

Features

  • Optimized for TensorRT conversion and deployment
  • Supports dynamic batch size for both encoder and decoder

License

This project is licensed under the Apache License, Version 2.0. See the LICENSE file for details.

About

tools for converting onnx models from pytorch

Resources

Stars

7 stars

Watchers

1 watching

Forks

Releases

Packages

Used by

Contributors

Languages