import os import re import random import numpy as np import pandas as pd import soundfile as sf from torch.utils import data from suno_utils.audio import Audio from suno_utils.utils.text import normalize_whitespace from suno_utils.utils.lyrics import remove_speakers from suno_utils.utils.text import read_jsonl def clean_text(text: str) -> str: """General text cleaning. A bit tight but makes the content very clean. Returns a cleaned text string that is expected to be recongizable by hoot. """ text = "\n" + text text = text.replace("’", "'").lower() text = text.replace('"', "").lower() text = re.sub(r"\[.+?\]", " ", text) # tags text = re.sub(r"\n.+?\:", " ", text) # new line ends with : text = re.sub(r"\n.+?\:", " ", text) # new line ends with : text = re.sub(r"\n\(.+?\)", " ", text) # new line with () text = re.sub(r"[\d]", " ", text) # digits text = re.sub(r"▁", "", text) # special stuff text = remove_speakers(text) text = re.sub(r"[^\w\'\s]", " ", text) # keep only the words text = normalize_whitespace(text) return text class GeniusLyricDataset(data.Dataset): def __init__( self, data_path="/app/suno/data/audio_mono_24khz/genius_hq", split="train", input_length_s=30.0, num_samples=-1, is_concat=False, ): assert split in ["train", "valid"] self.data_path = data_path self.split = split self.input_length_s = input_length_s self.num_samples = num_samples self.fs = 24000 self.is_concat = is_concat # get filelist if split == "train": self.ids = np.load( os.path.join(data_path, "metadata", "lyric_0.4", "train_ids.npy") ) self.metadata = read_jsonl( os.path.join(data_path, "metadata", "lyric_0.4", "train_metadata.jsonl") ) elif split == "valid": self.ids = np.load( os.path.join(data_path, "metadata", "lyric_0.4", "test_ids.npy") ) self.metadata = read_jsonl( os.path.join(data_path, "metadata", "lyric_0.4", "test_metadata.jsonl") ) print("There are %d audio files from Genius data" % len(self.ids)) def __getitem__(self, index): # load audio track_id = self.ids[index] audio_fn = os.path.join(self.data_path, "audio", track_id + ".wav") audio = Audio.from_file(audio_fn, sample_rate=self.fs) wav = audio.array_float # get lyrics if self.split == "train": metadata = random.choice(self.metadata[index][1]) elif self.split == "valid": metadata = self.metadata[index][1][0] lyrics = metadata["text"] start_s = metadata["start_s"] end_s = metadata["end_s"] start_ix = int(start_s * self.fs) end_ix = int(end_s * self.fs) # clearn lyrics lyrics = clean_text(lyrics) # append special tokens lyrics = "[CLS]" + "[Lyrics]" + lyrics # audio crop input_length = int(self.fs * self.input_length_s) wav = wav[start_ix:end_ix] if len(wav) < input_length: # zero padding wav = np.pad( wav, (0, input_length - len(wav)), mode="constant", constant_values=0.0 ) wav = wav[:input_length] if self.is_concat: return wav.astype("float32"), lyrics else: return wav.astype("float32"), lyrics, track_id, start_s def __len__(self): if self.num_samples > 0: return self.num_samples else: return len(self.ids)