import os import random import librosa import numpy as np import soundfile as sf from torch.utils import data class MTATDataset(data.Dataset): def __init__( self, data_path="/app/suno/minz/datasets/mtat", split="train", input_length_s=30.0, sample_rate=24000, num_samples=-1, ): assert split in ["train", "valid", "test"] 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.filelist = np.load(os.path.join(data_path, "splits", "%s.npy" % split)) self.binary = np.load(os.path.join(data_path, "splits", "binary.npy")) self.tags = np.load(os.path.join(data_path, "splits", "tags.npy")) print("%d files are available for %s set" % (len(self.filelist), split)) def concatenate_tags(self, tag_binary): tags = self.tags[tag_binary > 0].tolist() if self.split == "train": random.shuffle(tags) concatenated_tags = ", ".join(tags) # add [CLS] token concatenated_tags = "[CLS]" + concatenated_tags return concatenated_tags def __getitem__(self, index): # read data ix, fn = self.filelist[index].split("\t") # load audio audio_path = os.path.join(self.data_path, "audio_24kHz", fn[:-3] + "wav") wav, _ = sf.read(audio_path) wav = wav[:int(self.sample_rate * 29.1)] # load tag labels tag_binary = self.binary[int(ix)] concatenated_tags = self.concatenate_tags(tag_binary) return wav, concatenated_tags def __len__(self): if self.num_samples > 0: return self.num_samples else: return len(self.filelist)