import os import glob import random import torch import numpy as np from torch.utils import data from suno_utils.audio import Audio from torchaudio_augmentations import ( RandomApply, Noise, Gain, PitchShift, Compose, ) class GenreSimDataset(data.Dataset): def __init__( self, metadata, task_indices, split="train", input_length_s=30.0, input_text_length=300, sample_rate=24000, num_samples=-1, ): assert split in ["train", "valid"] self.metadata = metadata self.indices = task_indices["genre_sim"][split] self.input_length_s = input_length_s self.input_text_length = input_text_length self.sample_rate = sample_rate self.num_samples = num_samples self.split = split # get augmentation if split == "train": self._get_augmentations() random.seed(207) random.shuffle(self.indices) print(f"{len(self.indices)} files are available for genre_sim {split} set") def _get_augmentations(self): # Stochastic data augmentation transforms = [ RandomApply([Noise(min_snr=0.1, max_snr=0.5)], p=0.3), RandomApply([Gain()], p=0.2), RandomApply( [ PitchShift( n_samples=int(self.input_length_s * self.sample_rate), sample_rate=self.sample_rate, pitch_shift_min=-3.0, pitch_shift_max=3.0, ) ], p=0.4, ), ] self.augmentation = Compose(transforms=transforms) def get_tags(self, index): tags = self.metadata[index]["genres"] if "country" in self.metadata[index]: tags.append(self.metadata[index]["country"]) if "decade" in self.metadata[index]: if random.random() < 0.5: decade = self.metadata[index]["decade"][2:] # e.g., 80s, 90s else: decade = self.metadata[index]["decade"] tags.append(decade) if self.split == "train": random.shuffle(tags) num_tags = random.choice(range(1, len(tags) + 1)) tags = tags[:num_tags] # Convert list of tags to a single string tag_string = ", ".join(tags) # Prepend with "[CLS][Tag]" and limit to 200 characters tag_string = "[CLS][Tag]" + tag_string.lower() tag_string = tag_string[: self.input_text_length] return tag_string def __getitem__(self, index): # read data if self.split == "train": random_ix = random.choice(self.indices) audio_path = self.metadata[random_ix]["filepath"] tags = self.get_tags(random_ix) elif self.split == "valid": audio_path = self.metadata[self.indices[index]]["filepath"] tags = self.get_tags(self.indices[index]) 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: start_s = 0.0 wav = audio.get_segment( from_s=start_s, to_s=start_s + self.input_length_s ).array_float # augmentation wav = ( self.augmentation(torch.from_numpy(wav).unsqueeze(0)).squeeze(0).numpy() ) elif self.split == "valid": # crop first and last 30s start_s = 0.0 wav = audio.get_segment( from_s=start_s, to_s=start_s + self.input_length_s ).array_float # zero 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) return wav, tags def __len__(self): if self.num_samples > 0: return self.num_samples else: return len(self.indices)