import os
import random
import librosa
import numpy as np
import soundfile as sf
from torch.utils import data


class MTATDataset(data.Dataset):
    def __init__(
        self,
        data_path="/app/suno/minz/datasets/mtat",
        split="train",
        input_length_s=30.0,
        sample_rate=24000,
        num_samples=-1,
    ):
        assert split in ["train", "valid", "test"]
        self.data_path = data_path
        self.split = split
        self.input_length_s = input_length_s
        self.sample_rate = sample_rate
        self.num_samples = num_samples

        # load files
        self.filelist = np.load(os.path.join(data_path, "splits", "%s.npy" % split))
        self.binary = np.load(os.path.join(data_path, "splits", "binary.npy"))
        self.tags = np.load(os.path.join(data_path, "splits", "tags.npy"))
        print("%d files are available for %s set" % (len(self.filelist), split))

    def concatenate_tags(self, tag_binary):
        tags = self.tags[tag_binary > 0].tolist()
        if self.split == "train":
            random.shuffle(tags)
        concatenated_tags = ", ".join(tags)

        # add [CLS] token
        concatenated_tags = "[CLS]" + concatenated_tags
        return concatenated_tags

    def __getitem__(self, index):
        # read data
        ix, fn = self.filelist[index].split("\t")

        # load audio
        audio_path = os.path.join(self.data_path, "audio_24kHz", fn[:-3] + "wav")
        wav, _ = sf.read(audio_path)
        wav = wav[:int(self.sample_rate * 29.1)]

        # load tag labels
        tag_binary = self.binary[int(ix)]
        concatenated_tags = self.concatenate_tags(tag_binary)
        
        return wav, concatenated_tags

    def __len__(self):
        if self.num_samples > 0:
            return self.num_samples
        else:
            return len(self.filelist)


    