#!/usr/bin/env python3 """Export MERT (Music Understanding Transformer) to ONNX. Usage: pip install transformers torch onnx python export_mert_onnx.py --output models/mert-v1.onnx MERT is a pre-trained music understanding model from m-a-p/MERT-v1-330M. It handles 14+ MIR tasks: instrument, genre, mood, tempo, key, tags. """ import argparse import os def main(): parser = argparse.ArgumentParser(description="Export MERT to ONNX") parser.add_argument("--output", type=str, default="models/mert-v1.onnx") parser.add_argument("--model-name", type=str, default="m-a-p/MERT-v1-330M") args = parser.parse_args() os.makedirs(os.path.dirname(args.output) or ".", exist_ok=True) print(f"Exporting {args.model_name} to ONNX...") print(f"Output: {args.output}") try: import torch from transformers import AutoModel, AutoFeatureExtractor print("Loading MERT model...") model = AutoModel.from_pretrained(args.model_name, trust_remote_code=True) processor = AutoFeatureExtractor.from_pretrained(args.model_name, trust_remote_code=True) model.eval() # Create dummy input (16kHz, 5 seconds) dummy_input = torch.randn(1, 16000 * 5) print("Exporting to ONNX...") torch.onnx.export( model, dummy_input, args.output, input_names=["audio"], output_names=["embeddings"], opset_version=17, dynamic_axes={"audio": {1: "samples"}, "embeddings": {1: "frames"}}, ) file_size_mb = os.path.getsize(args.output) / (1024 * 1024) print(f"Export complete: {args.output} ({file_size_mb:.1f} MB)") except ImportError as e: print(f"Missing dependency: {e}") print("Install with: pip install transformers torch onnx") if __name__ == "__main__": main()