"""Audio to MIDI transcription application on modal. Audio --> BasicPitch --> MIDI """ import modal import tempfile import time import os from uuid import uuid4 import pathlib from suno_utils.audio import Audio from suno_utils.worker.settings import s3_client base_image = ( modal.Image.from_registry( "tensorflow/tensorflow:2.15.0-gpu", add_python="3.10", ) .apt_install("curl", "ffmpeg", "sox", "unzip", "libsox-fmt-mp3", "git", "clang") .pip_install_from_pyproject( str(pathlib.Path(__file__).parent.parent.parent / "pyproject.toml"), ) .pip_install( "basic-pitch[tf]", "boto3", "midiutil", ) ) APP_NAME = "midi-transcription-dev" app = modal.App(APP_NAME, image=base_image) aws_secret = modal.Secret.from_name("aws-bucket") @app.cls( cpu=4, gpu="A10G", secrets=[aws_secret], timeout=240, memory=8000, allow_concurrent_inputs=4, min_containers=1, ) class MidiTranscriptionStub: def __init__(self): from basic_pitch import ICASSP_2022_MODEL_PATH from basic_pitch.inference import Model print("Loading BasicPitch model...") self.model = Model(ICASSP_2022_MODEL_PATH) print("Model loaded!") @modal.method() def transcribe(self, s3_path: str, midi_id: 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 local_audio_path = os.path.join(td, "input.wav") audio = Audio.from_s3(s3_path, n_channels=1) audio.to_wav(local_audio_path) # Generate unique ID for the output file output_id = midi_id if midi_id else str(uuid4()) midi_path = os.path.join(td, "input_basic_pitch.mid") # Use basic_pitch's predict_and_save to save MIDI and other outputs from basic_pitch.inference import predict_and_save # Save MIDI file using the model output predict_and_save( [local_audio_path], # List of input audio paths td, # Output directory save_midi=True, # Save MIDI file sonify_midi=False, # Don't save audio rendering of MIDI save_model_outputs=False, # Don't save raw model outputs save_notes=False, # Don't save note events as CSV model_or_model_path=self.model, ) # 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 transcription with multiple instrument samples test_samples = { "stone_vocals": { "audio_path": "s3://suno-data-uploads/studio/uploads/45f04b32-8df3-4428-ab7a-47e561ff3847.mp3", "instrument": 53, }, "bass": { "audio_path": "s3://suno-data-uploads/studio/uploads/ec626427-771d-4b2e-8cd0-7bd2a7f1cc19.mp3", "instrument": 32, }, "piano": { "audio_path": "s3://suno-data-uploads/studio/uploads/fefebefd-a98b-4167-83ae-0722f8891824.mp3", "instrument": 1, }, "electric_guitar": { "audio_path": "s3://suno-data-uploads/studio/uploads/bac51035-7732-47b7-b180-7a759edb6bd0.mp3", "instrument": 25, }, "flute": { "audio_path": "s3://suno-data-uploads/studio/uploads/16dcfe3e-7760-40f5-9d6c-e598010383c9.mp3", "instrument": 73, }, "violin": { "audio_path": "s3://suno-data-uploads/studio/uploads/6c957b14-66f1-4219-8b83-6001860055d7.mp3", "instrument": 40, }, } for name, sample in test_samples.items(): print(f"Transcribing {name}...") id = sample["audio_path"].split("/")[-1].split(".")[0] output_path = model.transcribe.remote(sample["audio_path"], midi_id=id) print(f"MIDI file for {name} saved to: {output_path}")