import time

import modal
from modal.cls import ClsMixin

from transformers import WavLMModel
from torchaudio.pipelines import HDEMUCS_HIGH_MUSDB_PLUS
from suno_utils.worker.generative_worker_v2 import ModelV2Worker
from suno_utils.worker.modal_base import get_modal_base_image, MODAL_MOUNTS
from suno_utils.worker.utils import recursive_ls_dir
from transformers import BertTokenizer
from encodec import EncodecModel

aws_secret = modal.Secret.from_name("studio-aws")


def download_model_wrapper_lg_2():
    ModelV2Worker.download_models(model_size="lg")

    WavLMModel.from_pretrained("microsoft/wavlm-large")
    BertTokenizer.from_pretrained("bert-base-multilingual-cased")
    HDEMUCS_HIGH_MUSDB_PLUS.get_model()
    EncodecModel.encodec_model_24khz()


image = get_modal_base_image().run_function(download_model_wrapper_lg_2, secret=aws_secret)
STUB_NAME = "bark-v2-lg-martin"
stub = modal.Stub(STUB_NAME, image=image)


@stub.cls(
    cpu=2.0,
    gpu=modal.gpu.A10G(count=1),
    secret=aws_secret,
    timeout=800,
    container_idle_timeout=300,
    mounts=MODAL_MOUNTS,
    retries=2,
    # keep_warm=1,
    # max_concurrency=4,
)
class ModelV2Stub(ClsMixin):
    def __enter__(self):
        import torch

        num_gpus = torch.cuda.device_count()
        print(f"Found {num_gpus} GPUs.")
        recursive_ls_dir("/suno/models")

        self.worker = ModelV2Worker(0, model_size="lg", bg_image="/suno/models/assets/wave-bg-2.png")
        self.worker.preload()

    @modal.method()
    def generate(self, queue_item: str):
        import json

        from suno_utils.worker.schema import QueueItem

        print(queue_item)

        start_time = time.time()
        item = QueueItem(**json.loads(queue_item))
        item_id = item.id
        ids = []
        try:
            audios = self.worker.process_item(item)
            for i, audio in enumerate(audios):
                new_id = f"{item_id}_{i}"
                ids.append(new_id)
                item.id = new_id
                self.worker._write_audio(item, audio, [], show_text=True)

            ok = True
        except Exception:
            import traceback

            traceback.print_exc()
            ok = False

        finish_time = time.time()

        self.worker.notify_finish(
            item,
            {
                "id": item_id,
                "model": "bark-v2",
                "n_audios": len(ids),
                "ok": 1 if ok else 0,
                "gen_duration": finish_time - start_time,
            },
            queue_name="results:q",
        )

        return item.id
