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
