Files
rustytorch/scripts/export_mert_onnx.py
T
osobhandClaude Opus 4.6 85b77d49f2 Add audio neural layers and model architectures for ClawSample integration
New nn layers:
- ConvTranspose1d with stride, padding, output_padding (9 tests)
- LSTM/BiLSTM with multi-layer support and hidden state (10 tests)

Audio source separation:
- Demucs ONNX inference with segmented overlap-add processing
- Native HtDemucs architecture (encoder/decoder with BiLSTM bottleneck)
- StemType enum: vocals, drums, bass, other, piano, guitar

Audio generation:
- Stable Audio Open ONNX inference scaffold
- GenerationParams (prompt, duration, steps, cfg_scale, seed)

ONNX export scripts:
- export_demucs_onnx.py — Demucs v4 to ONNX with segment chunking
- export_stable_audio_onnx.py — Stable Audio Open components
- export_mert_onnx.py — MERT music understanding transformer

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-04-17 12:27:12 -07:00

60 lines
1.8 KiB
Python

#!/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()