import time from typing import Union from pydantic import BaseModel from torchaudio.pipelines import HDEMUCS_HIGH_MUSDB_PLUS from suno_utils.audio import Audio from suno_utils.tasks.gpt import chirp_v0 from suno_utils.tasks import mert from encodec import EncodecModel from transformers import AutoModel from .loader import S3Loader from .schema import QueueItem from .utils import download_models_to_dir class ChirpQueueItem(BaseModel): id: str prompt_audio: Union[str, None] = None prompt_npz: Union[str, None] = None prompt_text: Union[str, None] = None metadata: dict gen_duration: Union[int, None] = 12 CHIRP_MODELS = dict( semantic_ckpt_path="georg/trained_models/chirp_v0/xl_semantic.pt", coarse_ckpt_path="georg/trained_models/chirp_v0/lg_coarse.pt", text_tokenizer_path="georg/trained_models/chirp_v0/tokenizer.json", centroids_filepath="georg/trained_models/chirp_v0/4x1k_centroids_mert.npy", ) MOUNT_PATH = "/suno/models" class ChirpV0Worker(S3Loader): gpu_id: int def __init__(self, gpu_id): super().__init__() self.gpu_id = gpu_id def preload(self): start_time = time.time() chirp_v0.preload_models(**{k: f"{MOUNT_PATH}/{v}" for k, v in CHIRP_MODELS.items()}) finish_time = time.time() print(f"Preloading took {finish_time - start_time}s") @staticmethod def download_models(dir_path=MOUNT_PATH): """Use AWS CLI to download models if they don't exist.""" files = list(CHIRP_MODELS.values()) download_models_to_dir(files, dir_path) HDEMUCS_HIGH_MUSDB_PLUS.get_model() EncodecModel.encodec_model_24khz() AutoModel.from_pretrained( mert.DEFAULT_MODEL_NAME, trust_remote_code=True, revision=mert.DEFAULT_REVISION, ) def process_item(self, item: QueueItem) -> list[Audio]: N_BATCH = 2 history_audio = self._load_audio_prompt(item) if item.prompt_audio else None if item.prompt_text: audios = chirp_v0.text_to_audio( item.prompt_text + "\n\n-", history_audio=history_audio, n_batch=N_BATCH, max_gen_duration_s=30 if not history_audio else 20, cfg_gamma=1.8, # semantic_top_k=200, max_history_duration_s=10, ) else: audios = chirp_v0.generate_audio( history_audio=history_audio, n_batch=N_BATCH, gen_duration_s=20, max_history_duration_s=15, ) return audios