import os
import random
import numpy as np
import pandas as pd
import soundfile as sf
from torch.utils import data
from suno_utils.audio import Audio


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

        # get filelist
        self.fl = np.load(os.path.join(data_path, "metadata", "train_filelist.npy"))
        print("There are %d audio files from Genius data" % len(self.fl))

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

        # random crop
        input_length = int(self.fs * self.input_length_s)
        if len(wav) < input_length:  # zero padding
            wav = np.pad(
                wav, (0, input_length - len(wav)), mode="constant", constant_values=0.0
            )
        if self.split == "train":
            random_ix = random.randint(0, len(wav) - input_length)
            wav = wav[random_ix : random_ix + input_length]

        return wav.astype("float32")

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