import time from typing import Union, Any from pydantic import BaseModel from torchaudio.pipelines import HDEMUCS_HIGH_MUSDB_PLUS from suno_utils.audio import Audio from suno_utils.audio import Audio from suno_utils.tasks.gpt.generation import * from suno_utils.tasks.gpt_v2 import chirp_v1 from suno_utils.tasks import mert_v2 from transformers import AutoModel, Wav2Vec2FeatureExtractor from suno_utils.worker.loader import S3Loader from suno_utils.worker.schema import QueueItem from suno_utils.worker.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_V1_MODELS = dict( combo_path="georg/trained_models/chirp_v1/xl_2.pt", mert_path="georg/trained_models/chirp_v1/semantic_centroids.npy", codec_path="georg/trained_models/chirp_v1/codec.pt", genre_tags_paths="georg/trained_models/chirp_v1/common_genre_tags.json", fasttext_path="georg/trained_models/chirp_v1/lid.176.bin", ) MOUNT_PATH = "/suno/models" class ChirpV1Worker(S3Loader): gpu_id: int def __init__(self, gpu_id): super().__init__() self.gpu_id = gpu_id def preload(self): start_time = time.time() chirp_v1.preload_models( centroids_filepath=f"{MOUNT_PATH}/{CHIRP_V1_MODELS['mert_path']}", gpt_ckpt_path=f"{MOUNT_PATH}/{CHIRP_V1_MODELS['combo_path']}", codec_ckpt_path=f"{MOUNT_PATH}/{CHIRP_V1_MODELS['codec_path']}", common_genre_tags_path=f"{MOUNT_PATH}/{CHIRP_V1_MODELS['genre_tags_paths']}", fasttext_ckpt_path=f"{MOUNT_PATH}/{CHIRP_V1_MODELS['fasttext_path']}", local_whisper="/suno/models/whisper", ) 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_V1_MODELS.values()) download_models_to_dir(files, dir_path) HDEMUCS_HIGH_MUSDB_PLUS.get_model() AutoModel.from_pretrained( mert_v2.DEFAULT_MODEL_NAME, trust_remote_code=True, revision=mert_v2.DEFAULT_REVISION, ) Wav2Vec2FeatureExtractor.from_pretrained( mert_v2.DEFAULT_MODEL_NAME, trust_remote_code=True, revision=mert_v2.DEFAULT_REVISION, ) import whisper whisper.load_model("small.en") whisper.load_model("small") def process_item(self, item: QueueItem) -> tuple[list[Audio], list[Any], list[Any]]: history_audio = self._load_audio_prompt(item, "1.0.0.0") if item.prompt_audio else None options = item.metadata.get("options", {}) or {} n_batch = 2 chaos = options.get("chaos", 0) chaos = float(chaos) # cfg = float(options.get("lyrics_strength", 1.25)) # cfg_coef_tags = float(options.get("tags_strength", 1.75)) text = item.prompt_text text_tags = item.metadata.get("tags", None) if text_tags == "random": text_tags = None _, raw_arrays, audios, text_tags = chirp_v1.generate_audio( text=text, history_text_guess=item.metadata.get("continued_from_prompt", text), history_audio=history_audio, text_tags=text_tags or None, n_batch=n_batch, chaos_level=chaos, max_gen_duration_s=40 if history_audio is None else 30, max_history_duration_s=10.01, min_eos_p=0.25, return_raw_arrays=True, allow_genre_randomization=text_tags != "-", ) return audios, raw_arrays