import os
import glob
import random
import numpy as np
import pandas as pd

from torch.utils import data
from suno_utils.audio import Audio


class FMADataset(data.Dataset):
    def __init__(
            self, 
            data_path="/app/suno/data/audio_mono_24khz/fma/", 
            split="train", 
            input_length_s=29.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
        filenames = glob.glob(os.path.join(data_path, "audio", "*.wav"))
        random.seed(142)
        random.shuffle(filenames)
        if split == "train":
            self.filenames = filenames[:-500]
        elif split == "valid":
            self.filenames = filenames[-500:]
        print("There are %d audio files." % len(self.filenames))

    def __getitem__(self, index):
        # load audio
        audio = Audio.from_file(self.filenames[index], sample_rate=self.fs)
        wav = audio.array_float

        # padding
        input_length = int(self.fs * self.input_length_s)
        if len(wav) < input_length:
            pad_len = input_length - len(wav)
            wav = np.pad(wav, (0, pad_len), mode="constant", constant_values=0.0)

        # random crop
        if self.split == "train":
            random_ix = random.randint(0, len(wav) - 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.filenames)
