import os import random import pandas as pd import soundfile as sf from torch.utils import data class MERTDataset(data.Dataset): def __init__( self, data_path="/app/suno/data/mert_25hz_long", split="train", input_length_s=30.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 self.df = self.get_df() def get_df(self): df_path = os.path.join(self.data_path, "audio_tsv", self.split + ".tsv") df = pd.read_csv(df_path, sep="\t", names=["fn", "length"], skiprows=1) filtered_df = df[df.length > int(self.fs * self.input_length_s)] # get stats num_files = len(filtered_df) len_hours = filtered_df.length.sum() / self.fs / 60 / 60 print( "There are %d audio files longer than %.1f seconds. The dataset is %d hours in total." % (num_files, self.input_length_s, len_hours) ) return filtered_df def __getitem__(self, index): # read data frame filename = self.df.iloc[index].fn length = self.df.iloc[index].length # load audio wav, _ = sf.read(os.path.join(self.data_path, "audio", filename)) # random crop input_length = int(self.fs * self.input_length_s) if self.split == "train": random_ix = random.randint(0, length - 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.df)