import os
import random
import pandas as pd
import soundfile as sf
from torch.utils import data


class MERTDataset(data.Dataset):
    def __init__(
        self,
        data_path="/app/suno/data/mert_25hz_long",
        split="train",
        input_length_s=30.0,
        num_samples=-1,
    ):
        assert split in ["train", "valid"]
        self.data_path = data_path
        self.split = split
        self.input_length_s = input_length_s
        self.num_samples = num_samples
        self.fs = 24000

        # get dataframe
        self.df = self.get_df()

    def get_df(self):
        df_path = os.path.join(self.data_path, "audio_tsv", self.split + ".tsv")
        df = pd.read_csv(df_path, sep="\t", names=["fn", "length"], skiprows=1)
        filtered_df = df[df.length > int(self.fs * self.input_length_s)]

        # get stats
        num_files = len(filtered_df)
        len_hours = filtered_df.length.sum() / self.fs / 60 / 60
        print(
            "There are %d audio files longer than %.1f seconds. The dataset is %d hours in total."
            % (num_files, self.input_length_s, len_hours)
        )
        return filtered_df

    def __getitem__(self, index):
        # read data frame
        filename = self.df.iloc[index].fn
        length = self.df.iloc[index].length

        # load audio
        wav, _ = sf.read(os.path.join(self.data_path, "audio", filename))

        # random crop
        input_length = int(self.fs * self.input_length_s)
        if self.split == "train":
            random_ix = random.randint(0, length - input_length)
            wav = wav[random_ix : random_ix + input_length]
        elif self.split == "valid":
            wav = wav[:input_length]

        return wav.astype("float32")

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