import os import subprocess import time import traceback import torch from suno_utils.tasks.encodec import decode as codec_decode from suno_utils.tasks.encodec import encode as codec_encode from suno_utils.tasks.nano.gpt import ( generate_coarse, generate_fine, generate_semantic, generate_text_semantic, ) from suno_utils.tasks.nano.gpt import preload_models as preload_gpt_models from suno_utils.tasks.wavlm import ( encode as wavlm_encode, ) from suno_utils.tasks.wavlm import ( preload_models as preload_wavlm_models, ) from suno_utils.worker.schema import QueueItem from suno_utils.tasks import real_or_fake from suno_utils.worker.loader import S3Loader def _clear_cuda_cache(): if torch.cuda.is_available(): torch.cuda.empty_cache() torch.cuda.synchronize() class ModelV1Worker(S3Loader): gpu_id: int def __init__(self, gpu_id): super().__init__() self.gpu_id = gpu_id def preload_wavlm(self): centroids_filepath = "/suno/models/wavlm_v0/10k_centroids_wavlm.npy" _ = preload_wavlm_models(device=f"cuda:{self.gpu_id}", centroids_filepath=centroids_filepath) print("Preloaded WavLM models.") def preload(self): """Preload models onto GPU.""" # _clear_cuda_cache() start_time = time.time() semantic_ckpt_path = "/suno/models/wavlm_v0/lg_semantic.pt" coarse_ckpt_path = "/suno/models/wavlm_v0/lg_coarse.pt" fine_ckpt_path = "/suno/models/md_fine.pt" text_ckpt_path = "/suno/models/wavlm_v0/md_text_semantic.pt" centroids_filepath = "/suno/models/wavlm_v0/10k_centroids_wavlm.npy" _ = preload_wavlm_models(device=f"cuda:{self.gpu_id}", centroids_filepath=centroids_filepath) print("Preloaded WavLM models.") preload_gpt_models( semantic_ckpt_path=semantic_ckpt_path, coarse_ckpt_path=coarse_ckpt_path, fine_ckpt_path=fine_ckpt_path, text_ckpt_path=text_ckpt_path, ) real_or_fake.preload_model( "/suno/models/wavlm_v0/real_or_fake/model.pt", device=f"cuda:{self.gpu_id}" ) preload_time = time.time() - start_time print(f"Finished preloading models in {preload_time}s.") def generate_saved(self, item: QueueItem): """Use a saved history prompt (NPZ) and text to generate.""" history_prompt = self._load_history_prompt(item) prompt_semantic_arr = history_prompt["semantic_prompt"] prompt_coarse_arr = history_prompt["coarse_prompt"] prompt_fine_arr = history_prompt["fine_prompt"] text_prompt = item.prompt_text gen_semantic_arr = generate_text_semantic( text_prompt, temp=item.metadata.get("gen_semantic_temp", 0.8), semantic_history=prompt_semantic_arr, ) gen_coarse_arr = generate_coarse( gen_semantic_arr, x_history=(prompt_semantic_arr, prompt_coarse_arr), temp=item.metadata.get("gen_coarse_temp", 0.9), ) gen_fine_arr = generate_fine( gen_coarse_arr, x_fine_history=prompt_fine_arr, temp=item.metadata.get("gen_fine_temp", 0.3), ) audio = codec_decode(gen_fine_arr, n_codebooks=8) self._write_audio(item, audio, [gen_semantic_arr, gen_coarse_arr, gen_fine_arr]) def render_history_prompt(self, id: str): history_prompt = self._decode_history_prompt(id) prompt_fine_arr = history_prompt["fine_prompt"] audio = codec_decode(prompt_fine_arr, n_codebooks=8) self._write_audio(QueueItem(id=id, metadata={}), audio, []) def generate_prompt(self, item: QueueItem): prompt_audio = self._load_audio_prompt(item) # make embeddings prompt_semantic_arr = wavlm_encode(prompt_audio, device=f"cuda:{self.gpu_id}") prompt_coarse_arr = codec_encode(prompt_audio, n_codebooks=2) prompt_fine_arr = codec_encode(prompt_audio, n_codebooks=8) text_prompt = item.prompt_text gen_semantic_arr = generate_text_semantic( text_prompt, temp=item.metadata.get("gen_semantic_temp", 0.8), semantic_history=prompt_semantic_arr, ) gen_coarse_arr = generate_coarse( gen_semantic_arr, x_history=(prompt_semantic_arr, prompt_coarse_arr), temp=item.metadata.get("gen_coarse_temp", 0.9), ) gen_fine_arr = generate_fine( gen_coarse_arr, x_fine_history=prompt_fine_arr, temp=item.metadata.get("gen_fine_temp", 0.3), ) audio = codec_decode(gen_fine_arr, n_codebooks=8) self._write_audio(item, audio, [gen_semantic_arr, gen_coarse_arr, gen_fine_arr]) def generate_continuation(self, item: QueueItem): prompt_audio = self._load_audio_prompt(item) # make embeddings prompt_semantic_arr = wavlm_encode(prompt_audio, device=f"cuda:{self.gpu_id}") prompt_coarse_arr = codec_encode(prompt_audio, n_codebooks=2) prompt_fine_arr = codec_encode(prompt_audio, n_codebooks=8) gen_semantic_arr = generate_semantic( x=prompt_semantic_arr, gen_duration_s=item.gen_duration, temp=item.metadata.get("gen_semantic_temp", 0.8), ) gen_coarse_arr = generate_coarse( gen_semantic_arr, x_history=(prompt_semantic_arr, prompt_coarse_arr), temp=item.metadata.get("gen_coarse_temp", 0.9), ) gen_fine_arr = generate_fine( gen_coarse_arr, x_fine_history=prompt_fine_arr, temp=item.metadata.get("gen_fine_temp", 0.3), ) audio = codec_decode(gen_fine_arr, n_codebooks=8) self._write_audio(item, audio, [gen_semantic_arr, gen_coarse_arr, gen_fine_arr]) def generate_text(self, item: QueueItem): gen_semantic_arrs = [ generate_text_semantic( item.prompt_text, temp=item.metadata.get("gen_semantic_temp", 0.8), ) for i in range(3) ] scored = list( zip( real_or_fake.score_audio(gen_semantic_arrs, device=f"cuda:{self.gpu_id}"), gen_semantic_arrs, ) ) scored = sorted(scored, key=lambda x: x[0]) gen_semantic_arr = scored[-1][1] gen_coarse_arr = generate_coarse( gen_semantic_arr, temp=item.metadata.get("gen_coarse_temp", 0.9), ) gen_fine_arr = generate_fine( gen_coarse_arr, temp=item.metadata.get("gen_fine_temp", 0.3), ) audio = codec_decode(gen_fine_arr, n_codebooks=8) self._write_audio(item, audio, [gen_semantic_arr, gen_coarse_arr, gen_fine_arr]) def generate_unconditional(self, item: QueueItem): gen_semantic_arr = generate_semantic( gen_duration_s=item.gen_duration, temp=item.metadata.get("gen_semantic_temp", 0.8), ) gen_coarse_arr = generate_coarse( gen_semantic_arr, temp=item.metadata.get("gen_coarse_temp", 0.9), ) gen_fine_arr = generate_fine( gen_coarse_arr, temp=item.metadata.get("gen_fine_temp", 0.3), ) audio = codec_decode(gen_fine_arr, n_codebooks=8) self._write_audio(item, audio, [gen_semantic_arr, gen_coarse_arr, gen_fine_arr]) def process_item(self, item): try: if item.prompt_npz and item.prompt_text: self.generate_saved(item) elif not item.prompt_text and not item.prompt_audio: self.generate_unconditional(item) elif not item.prompt_text: self.generate_continuation(item) elif not item.prompt_audio: self.generate_text(item) else: self.generate_prompt(item) ok = True except: traceback.print_exc() print("Errored", item) ok = False return ok @staticmethod def download_models(): """Use AWS CLI to download models if they don't exist.""" dir_path = "/suno/models/" files = [ "wavlm_v0/md_text_semantic.pt", "wavlm_v0/lg_coarse.pt", "wavlm_v0/lg_semantic.pt", "wavlm_v0/10k_centroids_wavlm.npy", "md_fine.pt", "wavlm_v0/real_or_fake/model.pt", ] for f in files: full_path = os.path.abspath(os.path.join(dir_path, f)) if not os.path.exists(full_path): parent_dir = os.path.dirname(full_path) os.makedirs(parent_dir, exist_ok=True) print("Downloading", full_path) cmd = [ "aws", "s3", "cp", f"s3://suno-data/georg/checkpoints/{f}", full_path, ] subprocess.run(cmd) else: print("Skipping", f)