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)