import time import pathlib import modal import torch from suno_utils.audio import Audio from suno_utils.worker.loader import S3Loader from suno_utils.worker.modal_base import MODAL_MOUNTS from suno_utils.worker.utils import print_gpu_memory_usage from suno_utils.gpt import chirp_v2_5 as chirp_v3 from suno_utils.models.ditto.ditto import Ditto ############## CHANGE THESE ############## DEPLOYMENT_TYPE = "dev" # dev, prod ########################################## ENCODER_CONCURRENCY_LIMITS = { "dev": 50, "prod": 250, } KEEP_WARM = { "dev": 1, "prod": 1, } # set number of cpus. ENCODER_MAX_INPUT = 10 # encoder can handle much more traffic, but let's be conservative... VERBOSE_MESSAGE = DEPLOYMENT_TYPE == "dev" MOUNT_PATH = "/suno/models" aws_secret = modal.Secret.from_name("studio-aws") SECRETS = [ aws_secret, modal.Secret.from_dict( { "SUNO_ASSETS_PATH": "/suno/models/assets", "XDG_CACHE_HOME": "/suno/models/", } ), modal.Secret.from_name("openai-secret"), ] base_image = ( modal.Image.debian_slim() .apt_install("curl", "ffmpeg", "sox", "unzip", "libsox-fmt-mp3") .run_commands( [ 'curl "https://awscli.amazonaws.com/awscli-exe-linux-x86_64.zip" -o "awscliv2.zip"', "unzip -q awscliv2.zip", "./aws/install", ] ) .pip_install( "torch==2.2.0.+cu118", "torchaudio==2.2.0+cu118", index_url="https://download.pytorch.org/whl/cu118", ) .pip_install_private_repos( "github.com/suno-ai/glockenspiel.git@f05e2f251#subdirectory=descript-audio-codec&egg=descript-audio-codec", git_user="mcamac", secrets=[modal.Secret.from_name("victor-modal-github-token")], ) .pip_install_private_repos( "github.com/suno-ai/hoot.git@1ad12a3", git_user="mcamac", secrets=[modal.Secret.from_name("victor-modal-github-token")], ) .pip_install( "boto3", "tokenizers", "encodec", "ctc_segmentation", "psutil", "pydantic", "nnAudio", "biopython>=1.81", # TODO: don't love this depdendency, for hoot "turbopuffer", ) .pip_install_from_pyproject( str(pathlib.Path(__file__).parent.parent.parent / "pyproject.toml"), ) .run_commands( "FLASH_ATTENTION_SKIP_CUDA_BUILD=TRUE pip install flash-attn==2.5.2 --no-build-isolation", ) ) DITTO_S3_PATH = "s3://suno-data/victor/checkpoints/ditto/step_370k.pt" class DittoWorker(S3Loader): def __init__(self): S3Loader.__init__(self) print("Start loading models") self.ditto_path = chirp_v3._get_model_if_needed(DITTO_S3_PATH, cache_dir=MOUNT_PATH) self.music_encoder_path = chirp_v3._get_model_if_needed( "s3://suno-data/victor/checkpoints/ditto/musicfm_concat_epoch=51.pt", cache_dir=MOUNT_PATH, ) print(f"Downloaded model to {self.ditto_path}") self.ditto = Ditto( music_encoder_name="musicfm_concat", latent_dim=128, model_path=self.ditto_path, music_encoder_path=self.music_encoder_path, is_flash=False, ) self.ditto = self.ditto.eval().cuda() print("Finish loading models") @staticmethod def download_models(dir_path=MOUNT_PATH): print("Start downloading models") dl_path = chirp_v3._get_model_if_needed(DITTO_S3_PATH, cache_dir=dir_path) music_encoder_path = chirp_v3._get_model_if_needed( "s3://suno-data/victor/checkpoints/ditto/musicfm_concat_epoch=51.pt", cache_dir=dir_path, ) print(f"Downloaded model to {dl_path}") print(f"Downloaded model to {music_encoder_path}") print("Finish downloading models") def download_model_wrapper_e(): # this print is necessary to have modal rerun this when MODEL changes # Modal tracks referenced global variables # Change the name of the function to force a rerun print("Downloading model", DITTO_S3_PATH) DittoWorker.download_models() image = base_image.run_function(download_model_wrapper_e, secrets=SECRETS) STUB_NAME = f"ditto-{DEPLOYMENT_TYPE}" stub = modal.Stub(STUB_NAME, image=image) @stub.cls( gpu=modal.gpu.T4(count=1), secrets=SECRETS, timeout=100, container_idle_timeout=400, mounts=MODAL_MOUNTS, retries=modal.Retries( max_retries=1, backoff_coefficient=2.0, initial_delay=5.0, ), keep_warm=KEEP_WARM[DEPLOYMENT_TYPE], concurrency_limit=ENCODER_CONCURRENCY_LIMITS[DEPLOYMENT_TYPE], allow_concurrent_inputs=ENCODER_MAX_INPUT, ) class DittoWorkerStub: def __enter__(self): import torch num_gpus = torch.cuda.device_count() print(f"Found {num_gpus} GPUs.") self.worker = DittoWorker() @modal.method() def encode_audio( self, id: str = None, s3_url: str = None, start: float = 0, dur: float = 30, callback_url=None ) -> None: if VERBOSE_MESSAGE: print_gpu_memory_usage(self.__class__.__name__) torch.cuda.reset_max_memory_allocated() if dur > 30: raise ValueError("Duration must be less than 30 seconds") assert id or s3_url, "Either id or s3_url must be provided" if id is not None: s3_url = f"s3://suno-data-uploads/studio/uploads/{id}.mp3" audio = Audio.from_s3(s3_url, n_channels=1, sample_rate=24000) audio = audio.get_segment(from_s=start, to_s=start + dur) wav = torch.tensor(audio.array_float).unsqueeze(0).cuda() emb = self.worker.ditto.music_to_latent(wav)[0].detach().cpu().numpy() if id and callback_url: import requests requests.post( callback_url, json={"id": id, "vector": emb.tolist()}, headers={ "Authentication": "Bearer 562a512f-0dce-4acd-bf23-ad9bb8d8a084", }, ) return emb @modal.method() def encode_audio_and_upload( self, id: str = None, s3_url: str = None, start: float = 0, dur: float = 30 ) -> None: """Use this method to backfill existing clips with embeddings. Should be used for one-off tasks (e.g. from your notebook) only. Please note that the DEPLOYMENT_TYPE should be set properly. For exmaple, you can't upload staging clips to prod namespace""" if VERBOSE_MESSAGE: print_gpu_memory_usage(self.__class__.__name__) torch.cuda.reset_max_memory_allocated() if dur > 30: raise ValueError("Duration must be less than 30 seconds") assert id or s3_url, "Either id or s3_url must be provided" if id is not None: s3_url = f"s3://suno-data-uploads/studio/uploads/{id}.mp3" audio = Audio.from_s3(s3_url, n_channels=1, sample_rate=24000) audio = audio.get_segment(from_s=start, to_s=start + dur) wav = torch.tensor(audio.array_float).unsqueeze(0).cuda() emb = self.worker.ditto.music_to_latent(wav)[0].detach().cpu().numpy() import turbopuffer as tpuf tpuf.api_key = "C6FVrjLLHJ65WwP8WbDJDut8NSHXT1rj" ns = tpuf.Namespace(f"energy-{DEPLOYMENT_TYPE}") ns.upsert( ids=[id], vectors=[emb.tolist()], distance_metric="cosine_distance", ) return emb @modal.method() def encode_text(self, text: str) -> None: if VERBOSE_MESSAGE: print_gpu_memory_usage(self.__class__.__name__) torch.cuda.reset_max_memory_allocated() te = self.worker.ditto.text_to_latent("[CLS]" + text)[0].detach().cpu().numpy() return te @stub.local_entrypoint() def main(): ditto_worker = DittoWorkerStub() print(ditto_worker.encode_audio.remote(id="4a77dea7-19f3-46d2-8b0a-b2b7e9ea9a05")) print( ditto_worker.encode_audio.remote( s3_url="s3://suno-data-uploads/studio/uploads/4a77dea7-19f3-46d2-8b0a-b2b7e9ea9a05.mp3" ) ) print(ditto_worker.encode_text.remote("jazz")) for i in range(30): ditto_worker.encode_audio.spawn("4a77dea7-19f3-46d2-8b0a-b2b7e9ea9a05") time.sleep(100)