import os import random import numpy as np import pandas as pd from torch.utils import data from suno_utils.audio import Audio class KeyTencyDataset(data.Dataset): def __init__( self, data_path="/app/suno/christian_c/datasets/tency", split="train", input_length_s=20.0, sample_rate=24000, ): assert split in ["train", "valid"] self.data_path = data_path self.split = split self.input_length_s = input_length_s self.sample_rate = sample_rate self.sample_length = int(input_length_s * sample_rate) self.keys = [ "A major", "A minor", "Ab major", "G# minor", "B major", "B minor", "Bb major", "Bb minor", "C major", "C minor", "D major", "D minor", "Db major", "C# minor", "E major", "E minor", "Eb major", "D# minor", "F major", "F minor", "G major", "G minor", "F# major", "F# minor" ] # load files self.df = pd.read_csv(os.path.join(data_path, f'tency_{self.split}.csv'), index_col='ID') self.track_ids = list(self.df.index) print("%d files are available for %s set" % (len(self.track_ids), split)) def __getitem__(self, index): # read data track_id = self.track_ids[index] entry = self.df.loc[track_id] key = entry['key'] key_ix = self.keys.index(key) # load audio audio_path = os.path.join(self.data_path, f'{track_id}.mp3') wav = Audio.from_file(audio_path, sample_rate=self.sample_rate).array_float # crop audio # clip first t seconds if wav.shape[0] > self.sample_length: wav = wav[: self.sample_length] elif wav.shape[0] < self.sample_length: # pad by repeating the signal if shorter than window pad_size = self.sample_length - wav.shape[0] wav = np.pad(wav, (0, pad_size), 'wrap') return wav, key_ix def __len__(self): return len(self.track_ids)