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

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


class MSDDataset(data.Dataset):
    def __init__(
            self, 
            data_path="/app/suno/data/audio_mono_24khz/msd/", 
            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
        ids = np.load(os.path.join(data_path, "splits", "long_ids.npy"))
        random.seed(142)
        random.shuffle(ids)
        if split == "train":
            self.ids = ids[:-500]
        elif split == "valid":
            self.ids = ids[-500:]
        print("There are %d audio files longer than 29 seconds." % len(self.ids))

    def __getitem__(self, index):
        # read data frame
        track_id = self.ids[index]

        # load audio
        audio = Audio.from_file(os.path.join(self.data_path, "audio", "%s.clip.wav" % track_id), sample_rate=self.fs)
        wav = audio.array_float

        # random crop
        input_length = int(self.fs * self.input_length_s)
        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.ids)
