import random
from torch.utils.data import Dataset, Sampler
from suno_utils.utils.text import read_json, read_jsonl
from ditto_v2.data_loaders.self_sim import SelfSimDataset
from ditto_v2.data_loaders.self_vox_sim import SelfVoxSimDataset
from ditto_v2.data_loaders.self_lyric_sim import SelfLyricSimDataset
from ditto_v2.data_loaders.artist_sim import ArtistSimDataset
from ditto_v2.data_loaders.artist_vox_sim import ArtistVoxSimDataset
from ditto_v2.data_loaders.album_sim import AlbumSimDataset
from ditto_v2.data_loaders.genre_sim import GenreSimDataset
from ditto_v2.data_loaders.lyric_sim import LyricSimDataset


class MultiTaskDataset(Dataset):
    def __init__(
        self,
        split="train",
        self_sim_input_length_s=6.0,
        self_vox_sim_input_length_s=6.0,
        self_lyric_sim_input_length=150,
        artist_sim_input_length_s=15.0,
        artist_vox_sim_input_length_s=10.0,
        album_sim_input_length_s=15.0,
        genre_sim_input_length_s=15.0,
        genre_tag_input_length=200,
        lyric_sim_input_length_s=15.0,
        lyric_sim_input_text_length=200,
        sample_rate=24000,
        num_samples=1000,
        num_self_sim_samples=-1,
        num_vox_sim_samples=-1,
        num_self_lyric_sim_samples=-1,
        num_artist_sim_samples=-1,
        num_artist_vox_sim_samples=-1,
        num_album_sim_samples=-1,
        num_genre_sim_samples=-1,
        num_lyric_sim_samples=-1,
        sample_ratio=[1, 1, 1, 1, 1, 1, 1, 1],
    ):
        assert split in ["train", "valid"]

        # read data
        metadata = read_jsonl("/app/suno/data/v2_audio/metadata/metas_%s.jsonl" % split)
        task_indices = read_json("/app/suno/data/v2_audio/metadata/task_indices.json")

        self.sample_ratio = sample_ratio
        self.self_sim_dataset = SelfSimDataset(
            metadata,
            task_indices,
            split,
            self_sim_input_length_s,
            sample_rate,
            num_self_sim_samples,
        )
        self.self_vox_sim_dataset = SelfVoxSimDataset(
            metadata,
            task_indices,
            split,
            self_vox_sim_input_length_s,
            sample_rate,
            num_vox_sim_samples,
        )
        self.self_lyric_sim_dataset = SelfLyricSimDataset(
            split,
            input_length=self_lyric_sim_input_length,
            num_samples=num_self_lyric_sim_samples,
        )
        self.artist_sim_dataset = ArtistSimDataset(
            metadata,
            task_indices,
            split,
            artist_sim_input_length_s,
            sample_rate,
            num_artist_sim_samples,
        )
        self.artist_vox_sim_dataset = ArtistVoxSimDataset(
            metadata,
            task_indices,
            split,
            artist_vox_sim_input_length_s,
            sample_rate,
            num_artist_vox_sim_samples,
        )
        self.album_sim_dataset = AlbumSimDataset(
            metadata,
            task_indices,
            split,
            album_sim_input_length_s,
            sample_rate,
            num_album_sim_samples,
        )
        self.genre_sim_dataset = GenreSimDataset(
            metadata,
            task_indices,
            split,
            genre_sim_input_length_s,
            genre_tag_input_length,
            sample_rate,
            num_genre_sim_samples,
        )
        self.lyric_sim_dataset = LyricSimDataset(
            metadata,
            task_indices,
            split,
            lyric_sim_input_length_s,
            lyric_sim_input_text_length,
            sample_rate,
            num_lyric_sim_samples,
        )
        self.split = split
        self.num_samples = num_samples

    def __getitem__(self, index):
        if self.split == "train":
            task_choice = random.choices(
                [
                    "self_sim",
                    "self_vox_sim",
                    "self_lyric_sim",
                    "artist_sim",
                    "artist_vox_sim",
                    "album_sim",
                    "genre_sim",
                    "lyric_sim",
                ],
                weights=self.sample_ratio,
                k=1,
            )[0]
            if task_choice == "self_sim":
                return self.self_sim_dataset[index], task_choice
            elif task_choice == "self_vox_sim":
                return self.self_vox_sim_dataset[index], task_choice
            elif task_choice == "self_lyric_sim":
                return self.self_lyric_sim_dataset[index], task_choice
            elif task_choice == "artist_sim":
                return self.artist_sim_dataset[index], task_choice
            elif task_choice == "artist_vox_sim":
                return self.artist_vox_sim_dataset[index], task_choice
            elif task_choice == "album_sim":
                return self.album_sim_dataset[index], task_choice
            elif task_choice == "genre_sim":
                return self.genre_sim_dataset[index], task_choice
            elif task_choice == "lyric_sim":
                return self.lyric_sim_dataset[index], task_choice
        elif self.split == "valid":
            if index < len(self.self_sim_dataset):
                return self.self_sim_dataset[index], "self_sim"
            elif index < len(self.self_sim_dataset) + len(self.self_vox_sim_dataset):
                return self.self_vox_sim_dataset[
                    index - len(self.self_sim_dataset)
                ], "self_vox_sim"
            elif index < len(self.self_sim_dataset) + len(
                self.self_vox_sim_dataset
            ) + len(self.self_lyric_sim_dataset):
                return self.self_lyric_sim_dataset[
                    index - len(self.self_sim_dataset) - len(self.self_vox_sim_dataset)
                ], "self_lyric_sim"
            elif index < len(self.self_sim_dataset) + len(
                self.self_vox_sim_dataset
            ) + len(self.self_lyric_sim_dataset) + len(self.artist_sim_dataset):
                return self.artist_sim_dataset[
                    index
                    - len(self.self_sim_dataset)
                    - len(self.self_vox_sim_dataset)
                    - len(self.self_lyric_sim_dataset)
                ], "artist_sim"
            elif index < len(self.self_sim_dataset) + len(
                self.self_vox_sim_dataset
            ) + len(self.self_lyric_sim_dataset) + len(self.artist_sim_dataset) + len(
                self.artist_vox_sim_dataset
            ):
                return self.artist_vox_sim_dataset[
                    index
                    - len(self.self_sim_dataset)
                    - len(self.self_vox_sim_dataset)
                    - len(self.self_lyric_sim_dataset)
                    - len(self.artist_sim_dataset)
                ], "artist_vox_sim"
            elif index < len(self.self_sim_dataset) + len(
                self.self_vox_sim_dataset
            ) + len(self.self_lyric_sim_dataset) + len(self.artist_sim_dataset) + len(
                self.artist_vox_sim_dataset
            ) + len(self.album_sim_dataset):
                return self.album_sim_dataset[
                    index
                    - len(self.self_sim_dataset)
                    - len(self.self_vox_sim_dataset)
                    - len(self.self_lyric_sim_dataset)
                    - len(self.artist_sim_dataset)
                    - len(self.artist_vox_sim_dataset)
                ], "album_sim"
            elif index < len(self.self_sim_dataset) + len(
                self.self_vox_sim_dataset
            ) + len(self.self_lyric_sim_dataset) + len(self.artist_sim_dataset) + len(
                self.artist_vox_sim_dataset
            ) + len(self.album_sim_dataset) + len(self.genre_sim_dataset):
                return self.genre_sim_dataset[
                    index
                    - len(self.self_sim_dataset)
                    - len(self.self_vox_sim_dataset)
                    - len(self.self_lyric_sim_dataset)
                    - len(self.artist_sim_dataset)
                    - len(self.artist_vox_sim_dataset)
                    - len(self.album_sim_dataset)
                ], "genre_sim"
            elif index < len(self.self_sim_dataset) + len(
                self.self_vox_sim_dataset
            ) + len(self.self_lyric_sim_dataset) + len(self.artist_sim_dataset) + len(
                self.artist_vox_sim_dataset
            ) + len(self.album_sim_dataset) + len(self.genre_sim_dataset) + len(
                self.lyric_sim_dataset
            ):
                return self.lyric_sim_dataset[
                    index
                    - len(self.self_sim_dataset)
                    - len(self.self_vox_sim_dataset)
                    - len(self.self_lyric_sim_dataset)
                    - len(self.artist_sim_dataset)
                    - len(self.artist_vox_sim_dataset)
                    - len(self.album_sim_dataset)
                    - len(self.genre_sim_dataset)
                ], "lyric_sim"

    def __len__(self):
        if self.split == "train":
            return self.num_samples
        elif self.split == "valid":
            return (
                len(self.self_sim_dataset) * (self.sample_ratio[0] > 0)
                + len(self.self_vox_sim_dataset) * (self.sample_ratio[1] > 0)
                + len(self.self_lyric_sim_dataset) * (self.sample_ratio[2] > 0)
                + len(self.artist_sim_dataset) * (self.sample_ratio[3] > 0)
                + len(self.artist_vox_sim_dataset) * (self.sample_ratio[4] > 0)
                + len(self.album_sim_dataset) * (self.sample_ratio[5] > 0)
                + len(self.genre_sim_dataset) * (self.sample_ratio[6] > 0)
            )
