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 LyricSimDataset(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["lyric_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 lyric_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_lyrics(self, index): lyrics = self.metadata[index]["lyrics"] if self.split == "train": if len(lyrics) > self.input_text_length: start_ix = random.randint(0, len(lyrics) - self.input_text_length - 1) lyrics = lyrics[start_ix : start_ix + self.input_text_length] elif self.split == "valid": lyrics = lyrics[: self.input_text_length] lyrics = "[CLS]" + "[Lyrics]" + lyrics return lyrics def __getitem__(self, index): # read data if self.split == "train": random_ix = random.choice(self.indices) audio_path = self.metadata[random_ix]["filepath"] lyrics = self.get_lyrics(random_ix) elif self.split == "valid": audio_path = self.metadata[self.indices[index]]["filepath"] lyrics = self.get_lyrics(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, lyrics def __len__(self): if self.num_samples > 0: return self.num_samples else: return len(self.indices)