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, read_json class YTMDataset(data.Dataset): def __init__( self, data_path="/app/suno/data/audio_mono_24khz/ytm_tagged", split="train", input_length_s=30.0, sample_rate=24000, num_iteration=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_iteration = num_iteration # load files self.metadata = read_jsonl( os.path.join(data_path, "metadata", "%s_balanced_tagging.jsonl" % split) ) self.all_tags = np.load(os.path.join(data_path, "metadata", "tags.npy")) self.tag_to_index = {tag: i for i, tag in enumerate(self.all_tags)} if split == "train": self.tag_to_ids = read_json( os.path.join(data_path, "metadata", "tag_to_train_ids.json") ) self.id_to_ix = {line["id"]: ix for ix, line in enumerate(self.metadata)} print("%d files are available for %s set" % (len(self.metadata), split)) def tag_to_binary(self, tags): binary = np.zeros(len(self.all_tags)) for tag in tags: binary[self.tag_to_index[tag]] = 1 return binary def __getitem__(self, index): # balanced sample in training if self.split == "train": _tag = self.all_tags[index % len(self.all_tags)] _id = random.choice(self.tag_to_ids[_tag]) index = self.id_to_ix[_id] # read data metadata = self.metadata[index] track_id = metadata["id"] tags = metadata["tags"] # load audio audio_path = os.path.join(self.data_path, "audio", track_id + ".wav") audio = Audio.from_file(audio_path, sample_rate=self.sample_rate) if self.split == "train": # random crop try: start_ms = random.randint( 0, (audio.duration_ms - int(1000 * self.input_length_s) - 1) ) start_s = start_ms / 1000 except ValueError as e: start_s = 0.0 wav = audio.get_segment( from_s=start_s, to_s=start_s + self.input_length_s ).array_float elif self.split == "valid": # crop first 30s wav = audio.get_segment(from_s=0.0, to_s=self.input_length_s).array_float # padding if len(wav) < int(self.sample_rate * self.input_length_s): pad = int(self.sample_rate * self.input_length_s) - len(wav) wav = np.pad(wav, (0, pad), mode="constant", constant_values=0) # tag to binary binary = self.tag_to_binary(tags) return wav, binary def __len__(self): if self.split == "train": return len(self.all_tags) * self.num_iteration else: return len(self.metadata)