import time from suno_utils.audio import Audio from suno_utils.tasks.gpt.generation import ( preload_models as preload_gpt_models, ) from suno_utils.tasks.wavlm import ( preload_models as preload_wavlm_models, ) from suno_utils.tasks.gpt.generation import ( preload_models as preload_gpt_models, ) from suno_utils.tasks.gpt import bark_v2 from suno_utils.worker.loader import S3Loader from suno_utils.worker.utils import download_models_to_dir class ModelV2Worker(S3Loader): gpu_id: int model_size: str = "xl" def __init__(self, gpu_id, model_size="xl"): super().__init__() self.gpu_id = gpu_id self.model_size = model_size def preload(self): start_time = time.time() semantic_ckpt_path = "/suno/models/georg/trained_models/bark_v2/lg_semantic.pt" coarse_ckpt_path = f"/suno/models/georg/trained_models/bark_v2/{self.model_size}_coarse.pt" print("Preloading GPT") preload_gpt_models( semantic_ckpt_path=semantic_ckpt_path, coarse_ckpt_path=coarse_ckpt_path, text_tokenizer_path="bert-base-multilingual-cased", ) print("Preloading MERT") centroids_filepath = "/suno/models/georg/trained_models/bark_v2/8x10k_centroids_wavlm.npy" preload_wavlm_models( centroids_filepath=centroids_filepath, ) finish_time = time.time() print(f"Preloading took {finish_time - start_time}s") @staticmethod def download_models(dir_path="/suno/models", model_size="xl"): """Use AWS CLI to download models if they don't exist.""" files = [ f"georg/trained_models/bark_v2/{model_size}_coarse.pt", "georg/trained_models/bark_v2/lg_semantic.pt", "georg/trained_models/bark_v2/8x10k_centroids_wavlm.npy", ] download_models_to_dir(files, dir_path) def process_item(self, item) -> list[Audio]: N_SAMPLES = 2 history_audio = ( self._load_audio_prompt(item) if item.prompt_audio or item.metadata.get("audio_url", None) else None ) if item.prompt_text: audios = bark_v2.text_to_audio( history_audio=history_audio, text=item.prompt_text, n_batch=N_SAMPLES, cfg_gamma=1.8, max_history_duration_s=12, ) else: audios = bark_v2.generate_audio(history_audio=history_audio, n_batch=N_SAMPLES) # if item.metadata.get("voice_only"): # stems = demucs.encode(audios) # audios = [Audio.from_array_float(a[3], demucs.SAMPLE_RATE) for a in stems] return audios