"""Audio to MIDI transcription application on modal. Audio --> Model --> MIDI """ import modal import tempfile import time import os from uuid import uuid4 import pathlib from typing import Optional from suno_utils.audio import Audio from suno_utils.worker.settings import s3_client from suno_utils.tasks.midi_transcription import MidiTranscriber from suno_utils.worker.modal_base import get_modal_base_image_with_flash_attention from suno_utils.worker.modal_model_volume import ( MODEL_STORE_VOLUME_DIR, model_store_volume, ) base_image = ( get_modal_base_image_with_flash_attention() .pip_install_from_pyproject( str(pathlib.Path(__file__).parent.parent.parent / "pyproject.toml"), ) .pip_install( "tqdm", "transformers==4.46.1", "pretty_midi", "midi_player", "miditok==3.0.6", "tokenizers==0.20.2", ) .add_local_python_source("suno_utils", copy=False) ) APP_NAME = "midi-transcription-dev" app = modal.App(APP_NAME, image=base_image) aws_secret = modal.Secret.from_name("studio-aws") @app.cls( gpu="A10G", secrets=[aws_secret], timeout=600, # 10 minutes volumes={MODEL_STORE_VOLUME_DIR: model_store_volume}, min_containers=0, max_containers=20, keep_warm=2, scaledown_window=60 * 3, ) class MidiTranscriptionStub: def __init__(self): print("Loading MidiTranscriber model...") self.transcriber = MidiTranscriber("s3://suno-data/victor/checkpoints/midi_transcription_v0.pt") print("Model loaded!") # warmup with a sample audio print("Warming up model...") self.transcriber.transcribe( # chosen because it's a short audio Audio.from_s3( "s3://suno-data-uploads/studio/uploads/c64e0f95-b907-4bf1-aa74-e07c25f23ded.mp3" ) ) print("Warmup complete!") @modal.method() def transcribe(self, s3_path: str, midi_id: Optional[str] = None) -> str: """Takes an S3 path to an audio file and transcribes it to MIDI. Returns the S3 path to the generated MIDI file. """ start_time = time.time() print(f"Processing audio from: {s3_path}") if not s3_path.startswith("s3://"): s3_path = f"s3://suno-data-uploads/studio/uploads/{s3_path}.mp3" print(f"Using default S3 path: {s3_path}") with tempfile.TemporaryDirectory() as td: # Download and process audio # The transcriber handles stereo to mono conversion. audio = Audio.from_s3(s3_path) # Generate unique ID for the output file output_id = midi_id if midi_id else str(uuid4()) midi_path = os.path.join(td, f"{output_id}.mid") # Transcribe audio to MIDI print("Transcribing audio to MIDI...") midi = self.transcriber.transcribe(audio, use_tqdm=False) # Save MIDI to file midi.write(midi_path) print(f"MIDI file saved locally to {midi_path}") # Upload MIDI to S3 output_s3_path = f"s3://suno-data-uploads/studio/uploads/{output_id}.mid" s3_client.upload_file( midi_path, "suno-data-uploads", f"studio/uploads/{output_id}.mid", ExtraArgs={ "ContentType": "audio/midi", }, ) print( f"Uploaded MIDI file to {output_s3_path}. Took {round(time.time() - start_time, 2)} seconds" ) return output_s3_path @app.local_entrypoint() def main(): model = MidiTranscriptionStub() test_sample_s3_path = ( "s3://suno-data-uploads/studio/uploads/a5e2198a-f352-4abb-9a24-7f81b143ded3.mp3" # stone ) print(f"Transcribing {test_sample_s3_path}...") output_path = model.transcribe.remote(test_sample_s3_path) print(f"MIDI file for test sample saved to: {output_path}")