"""Audio feature extraction on modal. Clip -> features """ import json import logging import modal import time import traceback from typing import Optional, Dict, Any from suno_utils.worker.schema import QueueItem from suno_utils.utils.clip import SunoClip from suno_utils.worker.utils import retry_decorator from suno_utils.worker.modal_base import get_modal_base_image from suno_utils.audio import Audio import tempfile import boto3 import torch from modal.experimental import stop_fetching_inputs ############## CHANGE THESE ############## DEPLOYMENT_TYPE = "dev" # dev, prod ########################################## aws_secret = modal.Secret.from_name("studio-aws") SECRETS = [ aws_secret, modal.Secret.from_name("openai-secret"), modal.Secret.from_name("api-callback-token"), ] image = ( get_modal_base_image() .pip_install( "git+https://github.com/marc-suno/madmom.git@bf7d502#egg=madmom", "git+https://github.com/CPJKU/beat_this@117ff34", "cvxpy==1.6.5", ) .add_local_python_source("suno_utils", copy=False) ) logger = logging.getLogger(__name__) APP_NAME = f"audio-features-{DEPLOYMENT_TYPE}" app = modal.App(APP_NAME, image=image) UPLOADS_BUCKET = "suno-data-uploads" UPLOADS_KEY_PREFIX = "studio/uploads" @app.cls( secrets=SECRETS, gpu="T4", timeout=60, scaledown_window=480, memory=3000, cpu=2, retries=modal.Retries( max_retries=1, backoff_coefficient=2.0, initial_delay=5.0, ), max_containers=1000, region="us-east", buffer_containers=0 if DEPLOYMENT_TYPE == "dev" else 1, min_containers=1, ) @modal.concurrent(max_inputs=2) class AudioFeaturesStub: def __init__(self): from suno_utils.tasks.audio_features.beat_this_downbeat import BeatThisDownbeatExtractor from suno_utils.tasks.audio_features.key import KeyExtractor from suno_utils.tasks.audio_features.instrument import InstrumentExtractor print("Initializing AudioFeaturesStub") self.modal_f_convert_to_wav = modal.Cls.from_name( f"cycle-{DEPLOYMENT_TYPE}", "WavCycleStub", )().convert_to_wav self.downbeat_extractor = BeatThisDownbeatExtractor( device="cuda", model_path="s3://suno-data/m4burns/beat_this_beta.pt" ) self.key_extractor = KeyExtractor() self.instrument_extractor = InstrumentExtractor() self.s3_client = boto3.client("s3") self.retry_s3_download = retry_decorator(3, wait_seconds=5)(self.s3_client.download_fileobj) torch.backends.cuda.cufft_plan_cache[0].max_size = 0 # warmup warmup_audio = self._get_aligned_audio("40ec1fc7-0b0a-4f01-89de-6a6a6db65475") self.downbeat_extractor.extract(warmup_audio) def _get_aligned_audio(self, gen_id: str) -> Audio: # try grabbing the opus from s3 key = f"{UPLOADS_KEY_PREFIX}/{gen_id}.opus" with tempfile.NamedTemporaryFile(suffix=".opus") as f: try: self.s3_client.download_fileobj(Bucket=UPLOADS_BUCKET, Key=key, Fileobj=f) f.flush() return Audio.from_file(f.name, n_channels=1, sample_rate=48000) except (self.s3_client.exceptions.NoSuchKey, self.s3_client.exceptions.ClientError): print(f"Opus not found for {gen_id}, trying to generate...") except Exception as e: print(f"Error getting opus for {gen_id}: {e}") raise e self.modal_f_convert_to_wav.remote( f'{{"id": "{gen_id}", "metadata": {{"convert_to_opus": true, "user_id": 0, "clip_user_id": 0}}}}' ) try: f.truncate(0) self.retry_s3_download(Bucket=UPLOADS_BUCKET, Key=key, Fileobj=f) f.flush() return Audio.from_file(f.name, n_channels=1, sample_rate=48000) except Exception as e: print(f"Error getting opus for {gen_id}: {e}") raise e @modal.method() def extract_downbeats_callback( self, queue_item_str: str, return_result: bool = False, ) -> Optional[Dict[str, Any]]: queue_item = QueueItem(**json.loads(queue_item_str)) print("extract_downbeats_callback", queue_item.model_dump_json()) try: audio = self._get_aligned_audio(queue_item.id) result_dict = self.downbeat_extractor.extract(audio) queue_item.notify_progress( { "id": queue_item.id, "metadata": queue_item.metadata, **result_dict, }, ) except torch.cuda.OutOfMemoryError: # container is likely to be in a bad state # stop fetching inputs and raise to trigger a restart print(f"Out of memory for {queue_item.id} - exiting") stop_fetching_inputs() raise except Exception as e: print(f"Error extracting downbeats for {queue_item.id}: {e}") traceback.print_exc() queue_item.notify_progress( {"id": queue_item.id, "error": str(e), "metadata": queue_item.metadata}, ) return None if return_result: return result_dict @modal.method() def extract_key_callback( self, queue_item_str: str, return_result: bool = False, ) -> Optional[str]: queue_item = QueueItem(**json.loads(queue_item_str)) try: key = self.key_extractor.extract(SunoClip(queue_item.id)) except Exception as e: print(f"Error extracting key for {queue_item.id}: {e}") traceback.print_exc() queue_item.notify_progress( {"id": queue_item.id, "error": str(e), "metadata": queue_item.metadata}, ) return queue_item.notify_progress( {"id": queue_item.id, "key": key, "metadata": queue_item.metadata}, ) if return_result: return key @modal.method() def extract_instruments_callback( self, queue_item_str: str, return_result: bool = False, ) -> Optional[Dict[str, Any]]: queue_item = QueueItem(**json.loads(queue_item_str)) try: _, _, group_class, inst_class = self.instrument_extractor.extract(SunoClip(queue_item.id)) result = {"groups": group_class, "instruments": inst_class} except Exception as e: print(f"Error extracting instruments for {queue_item.id}: {e}") traceback.print_exc() queue_item.notify_progress( {"id": queue_item.id, "error": str(e), "metadata": queue_item.metadata}, ) return queue_item.notify_progress( {"id": queue_item.id, "result": result, "metadata": queue_item.metadata}, ) if return_result: return result @app.local_entrypoint() def main(): # tests model = AudioFeaturesStub() for gen_id in [ "3fabec37-fe9e-489d-b107-f0517feb3cf1", "40ec1fc7-0b0a-4f01-89de-6a6a6db65475", "879c3b07-87a1-4926-bdea-97fa77cc2d65", "ba6e3678-8fa6-4d9f-87f0-72530ccf3245", "a16c1c28-aaf8-4cb2-8ba9-4855f8b7054c", "a67cda87-cc19-49a2-90bc-bcfa211c29aa", ]: start = time.time() q = model.extract_downbeats_callback.remote( f'{{"id": "{gen_id}", "metadata": {{}}}}', return_result=True ) assert q is not None json.dumps(q) print(f"Time taken for downbeats: {time.time() - start}") start = time.time() q = model.extract_key_callback.remote( f'{{"id": "{gen_id}", "metadata": {{}}}}', return_result=True ) assert q is not None json.dumps(q) print(f"Time taken for modulation: {time.time() - start}") start = time.time() q = model.extract_instruments_callback.remote( f'{{"id": "{gen_id}", "metadata": {{}}}}', return_result=True ) assert q is not None json.dumps(q) print(f"Time taken for instruments: {time.time() - start}") print("Done")