import os import random import json import librosa import numpy as np import soundfile as sf from torch.utils import data from suno_utils.audio import Audio from suno_utils.utils.text import read_jsonl class ArtistYTMSDDataset(data.Dataset): def __init__( self, data_path="/app/suno/data/audio_mono_24khz/artist_ytmsd", split="train", input_length_s=30.0, sample_rate=24000, num_samples=-1, ): 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.num_samples = num_samples # load files self.metadata = read_jsonl(os.path.join(data_path, "%s.jsonl" % split)) print("%d files are available for %s set" % (len(self.metadata), split)) def concatenate_tags(self, tags): if self.split == "train": random.shuffle(tags) tags = [tag.lower() for tag in tags if len(tag.split(" ")) < 5] concatenated_tags = ", ".join(tags) # add [CLS] token concatenated_tags = "[CLS]" + concatenated_tags return concatenated_tags def __getitem__(self, index): # read data metadata = self.metadata[index] # sample two songs index_a, index_b = random.sample(metadata["indices"], 2) metadata_a = self.metadata[index_a] metadata_b = self.metadata[index_b] tags_a = metadata_a["clean_tags"] tags_b = metadata_b["clean_tags"] # load audio audio_path_a = metadata_a["filepath"] audio_path_b = metadata_b["filepath"] audio_a = Audio.from_file(audio_path_a, sample_rate=self.sample_rate) audio_b = Audio.from_file(audio_path_b, sample_rate=self.sample_rate) if self.split == "train": # random crop try: start_ms_a = random.randint( 0, (audio_a.duration_ms - int(1000 * self.input_length_s) - 1) ) start_s_a = start_ms_a / 1000 start_ms_b = random.randint( 0, (audio_b.duration_ms - int(1000 * self.input_length_s) - 1) ) start_s_b = start_ms_b / 1000 except ValueError as e: start_s_a = 0.0 start_s_b = 0.0 wav_a = audio_a.get_segment( from_s=start_s_a, to_s=start_s_a + self.input_length_s ).array_float wav_b = audio_b.get_segment( from_s=start_s_b, to_s=start_s_b + self.input_length_s ).array_float elif self.split == "valid": # crop first 30s wav_a = audio_a.get_segment( from_s=0.0, to_s=self.input_length_s ).array_float wav_b = audio_b.get_segment( from_s=0.0, to_s=self.input_length_s ).array_float if len(wav_a) < int(self.sample_rate * self.input_length_s): pad = int(self.sample_rate * self.input_length_s) - len(wav_a) wav_a = np.pad(wav_a, (0, pad), mode="constant", constant_values=0) if len(wav_b) < int(self.sample_rate * self.input_length_s): pad = int(self.sample_rate * self.input_length_s) - len(wav_b) wav_b = np.pad(wav_b, (0, pad), mode="constant", constant_values=0) # load tag labels concatenated_tags_a = self.concatenate_tags(tags_a) concatenated_tags_b = self.concatenate_tags(tags_b) return wav_a, wav_b, concatenated_tags_a, concatenated_tags_b def __len__(self): if self.num_samples > 0: return self.num_samples else: return len(self.metadata)