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)