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]>
60 lines
1.8 KiB
Python
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()
|