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)