"""Studio session bounce on modal. Studio session JSON -> WAV audio """ import asyncio import json import logging import modal import time import traceback import tempfile import subprocess import os from concurrent.futures import ThreadPoolExecutor from suno_utils.worker.schema import QueueItem from suno_utils.worker.modal_base import get_modal_base_image from suno_utils.audio import Audio from suno_utils.worker.loader import S3Loader, retry_s3_download from suno_utils.worker.settings import s3_client from suno_utils.worker.modal_model_configs import VAEVersion ############## CHANGE THESE ############## DEPLOYMENT_TYPE = "dev" # dev, prod ########################################## UPLOADS_S3_BUCKET = "suno-data-uploads" UPLOADS_S3_PREFIX = "studio/uploads" 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() .env( { "NVM_DIR": "/root/.nvm", "PATH": "$PATH:/root/.nvm/versions/node/v23.11.0/bin", } ) .run_commands( [ "curl -o- https://raw.githubusercontent.com/nvm-sh/nvm/v0.40.2/install.sh | bash", "bash -c 'source $NVM_DIR/nvm.sh && nvm install 23.11.0 && nvm use 23.11.0 && npm install -g yarn'", ] ) .add_local_dir( "../dsp-engine", "/dsp-engine", copy=True, ignore=["**/node_modules", "**/build", "**/dist"] ) .run_commands( ["cd /dsp-engine/buildtool && yarn install --frozen-lockfile && node ./index.js pull"], secrets=[aws_secret], ) .add_local_python_source("suno_utils", copy=False) ) logger = logging.getLogger(__name__) APP_NAME = f"studio-bounce-{DEPLOYMENT_TYPE}" app = modal.App(APP_NAME, image=image) @app.cls( secrets=SECRETS, timeout=400, scaledown_window=480, memory=15000, cpu=4, 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=4) class StudioBounceStub(S3Loader): def __init__(self): super().__init__() self.modal_f_encode_audio = modal.Cls.from_name( f"cycle-{DEPLOYMENT_TYPE}", "CycleStub", )().encode_audio self.modal_f_convert_to_wav = modal.Cls.from_name( f"cycle-{DEPLOYMENT_TYPE}", "WavCycleStub", )().convert_to_wav def _get_s3_ids_from_studio_state(self, studio_state_json): return set( filter( None, ( clip.get("clipId") for track in studio_state_json.get("tracks", []) for clip in track.get("clips", []) ), ) ) def _ensure_opus(self, s3_id): try: s3_client.head_object(Bucket=UPLOADS_S3_BUCKET, Key=f"{UPLOADS_S3_PREFIX}/{s3_id}.opus") return except (s3_client.exceptions.NoSuchKey, s3_client.exceptions.ClientError): print(f"Opus file for {s3_id} not found in S3, will try to convert") self.modal_f_convert_to_wav.remote( f'{{"id": "{s3_id}", "metadata": {{"convert_to_opus": true, "user_id": 0, "clip_user_id": 0}}}}' ) @modal.method() def bounce_callback( self, queue_item: str, ) -> None: local_queue_item = QueueItem(**json.loads(queue_item)) with ( tempfile.NamedTemporaryFile(delete=False, suffix=".json", mode="w") as studio_state_file, tempfile.NamedTemporaryFile(delete=False, suffix=".wav", mode="w") as output_file, ): output_file.close() try: if studio_state_json := local_queue_item.metadata.get("studio_state"): json.dump(studio_state_json, studio_state_file) studio_state_file.close() else: studio_state_file.close() with open(studio_state_file.name, "wb") as f: retry_s3_download( UPLOADS_S3_BUCKET, f"{UPLOADS_S3_PREFIX}/{local_queue_item.id}.json", f ) with open(studio_state_file.name, "r") as f: studio_state_json = json.load(f) with ThreadPoolExecutor(max_workers=4) as executor: futures = [ executor.submit( self._ensure_opus, s3_id, ) for s3_id in self._get_s3_ids_from_studio_state(studio_state_json) ] for future in futures: future.result() subprocess.run( [ "node", "/dsp-engine/bounce/dist/index.js", studio_state_file.name, "--start", str(float(local_queue_item.metadata.get("start_beats", "0"))), "--end", str(float(local_queue_item.metadata.get("end_beats", "100"))), "--output", output_file.name, ], check=True, ) audio = Audio.from_file(output_file.name, n_channels=2) asyncio.run( self._write_audio_only_async( item=local_queue_item, audio=audio, s3_bucket=UPLOADS_S3_BUCKET, s3_folder=UPLOADS_S3_PREFIX, ) ) # TODO: this operation should be blocking cause untils this is finished other operations can't be done self.modal_f_encode_audio.spawn( audio=f"s3://{UPLOADS_S3_BUCKET}/{UPLOADS_S3_PREFIX}/{local_queue_item.id}.opus", s3_npz_id=local_queue_item.id, encode_vae_version=VAEVersion.V_VAE_25_TUNED_2.value, ) local_queue_item.notify_progress( { "id": local_queue_item.id, "success": True, "metadata": { "duration": audio.duration_s, }, "is_remix": local_queue_item.metadata.get("is_remix", True), "parent_relationship_type": local_queue_item.metadata.get( "parent_relationship_type", "" ), "is_user_direct_parent_owner": local_queue_item.metadata.get( "is_user_direct_parent_owner", None ), }, ) print(f"{local_queue_item.id}: Notified progress for bounce.") except Exception as e: print(f"Error bouncing for {local_queue_item.id}: {e}") traceback.print_exc() local_queue_item.notify_progress( {"id": local_queue_item.id, "error": str(e)}, ) finally: if os.path.exists(studio_state_file.name): os.remove(studio_state_file.name) if os.path.exists(output_file.name): os.remove(output_file.name) @app.local_entrypoint() def main(): # tests model = StudioBounceStub() for gen_id in ["71e291bf-9d2f-43c3-a787-172a7ac8f897"]: for i in range(2): start = time.time() model.bounce_callback.remote( f'{{"id": "{gen_id}", "metadata": {{"start_beats": 0, "end_beats": 100}}}}' ) print(f"Time taken for bounce: {time.time() - start}") print("Done")