import random from torch.utils.data import Dataset, Sampler from suno_utils.utils.text import read_json, read_jsonl from ditto_v2.data_loaders.self_sim import SelfSimDataset from ditto_v2.data_loaders.self_vox_sim import SelfVoxSimDataset from ditto_v2.data_loaders.self_lyric_sim import SelfLyricSimDataset from ditto_v2.data_loaders.artist_sim import ArtistSimDataset from ditto_v2.data_loaders.artist_vox_sim import ArtistVoxSimDataset from ditto_v2.data_loaders.album_sim import AlbumSimDataset from ditto_v2.data_loaders.genre_sim import GenreSimDataset from ditto_v2.data_loaders.lyric_sim import LyricSimDataset class MultiTaskDataset(Dataset): def __init__( self, split="train", self_sim_input_length_s=6.0, self_vox_sim_input_length_s=6.0, self_lyric_sim_input_length=150, artist_sim_input_length_s=15.0, artist_vox_sim_input_length_s=10.0, album_sim_input_length_s=15.0, genre_sim_input_length_s=15.0, genre_tag_input_length=200, lyric_sim_input_length_s=15.0, lyric_sim_input_text_length=200, sample_rate=24000, num_samples=1000, num_self_sim_samples=-1, num_vox_sim_samples=-1, num_self_lyric_sim_samples=-1, num_artist_sim_samples=-1, num_artist_vox_sim_samples=-1, num_album_sim_samples=-1, num_genre_sim_samples=-1, num_lyric_sim_samples=-1, sample_ratio=[1, 1, 1, 1, 1, 1, 1, 1], ): assert split in ["train", "valid"] # read data metadata = read_jsonl("/app/suno/data/v2_audio/metadata/metas_%s.jsonl" % split) task_indices = read_json("/app/suno/data/v2_audio/metadata/task_indices.json") self.sample_ratio = sample_ratio self.self_sim_dataset = SelfSimDataset( metadata, task_indices, split, self_sim_input_length_s, sample_rate, num_self_sim_samples, ) self.self_vox_sim_dataset = SelfVoxSimDataset( metadata, task_indices, split, self_vox_sim_input_length_s, sample_rate, num_vox_sim_samples, ) self.self_lyric_sim_dataset = SelfLyricSimDataset( split, input_length=self_lyric_sim_input_length, num_samples=num_self_lyric_sim_samples, ) self.artist_sim_dataset = ArtistSimDataset( metadata, task_indices, split, artist_sim_input_length_s, sample_rate, num_artist_sim_samples, ) self.artist_vox_sim_dataset = ArtistVoxSimDataset( metadata, task_indices, split, artist_vox_sim_input_length_s, sample_rate, num_artist_vox_sim_samples, ) self.album_sim_dataset = AlbumSimDataset( metadata, task_indices, split, album_sim_input_length_s, sample_rate, num_album_sim_samples, ) self.genre_sim_dataset = GenreSimDataset( metadata, task_indices, split, genre_sim_input_length_s, genre_tag_input_length, sample_rate, num_genre_sim_samples, ) self.lyric_sim_dataset = LyricSimDataset( metadata, task_indices, split, lyric_sim_input_length_s, lyric_sim_input_text_length, sample_rate, num_lyric_sim_samples, ) self.split = split self.num_samples = num_samples def __getitem__(self, index): if self.split == "train": task_choice = random.choices( [ "self_sim", "self_vox_sim", "self_lyric_sim", "artist_sim", "artist_vox_sim", "album_sim", "genre_sim", "lyric_sim", ], weights=self.sample_ratio, k=1, )[0] if task_choice == "self_sim": return self.self_sim_dataset[index], task_choice elif task_choice == "self_vox_sim": return self.self_vox_sim_dataset[index], task_choice elif task_choice == "self_lyric_sim": return self.self_lyric_sim_dataset[index], task_choice elif task_choice == "artist_sim": return self.artist_sim_dataset[index], task_choice elif task_choice == "artist_vox_sim": return self.artist_vox_sim_dataset[index], task_choice elif task_choice == "album_sim": return self.album_sim_dataset[index], task_choice elif task_choice == "genre_sim": return self.genre_sim_dataset[index], task_choice elif task_choice == "lyric_sim": return self.lyric_sim_dataset[index], task_choice elif self.split == "valid": if index < len(self.self_sim_dataset): return self.self_sim_dataset[index], "self_sim" elif index < len(self.self_sim_dataset) + len(self.self_vox_sim_dataset): return self.self_vox_sim_dataset[ index - len(self.self_sim_dataset) ], "self_vox_sim" elif index < len(self.self_sim_dataset) + len( self.self_vox_sim_dataset ) + len(self.self_lyric_sim_dataset): return self.self_lyric_sim_dataset[ index - len(self.self_sim_dataset) - len(self.self_vox_sim_dataset) ], "self_lyric_sim" elif index < len(self.self_sim_dataset) + len( self.self_vox_sim_dataset ) + len(self.self_lyric_sim_dataset) + len(self.artist_sim_dataset): return self.artist_sim_dataset[ index - len(self.self_sim_dataset) - len(self.self_vox_sim_dataset) - len(self.self_lyric_sim_dataset) ], "artist_sim" elif index < len(self.self_sim_dataset) + len( self.self_vox_sim_dataset ) + len(self.self_lyric_sim_dataset) + len(self.artist_sim_dataset) + len( self.artist_vox_sim_dataset ): return self.artist_vox_sim_dataset[ index - len(self.self_sim_dataset) - len(self.self_vox_sim_dataset) - len(self.self_lyric_sim_dataset) - len(self.artist_sim_dataset) ], "artist_vox_sim" elif index < len(self.self_sim_dataset) + len( self.self_vox_sim_dataset ) + len(self.self_lyric_sim_dataset) + len(self.artist_sim_dataset) + len( self.artist_vox_sim_dataset ) + len(self.album_sim_dataset): return self.album_sim_dataset[ index - len(self.self_sim_dataset) - len(self.self_vox_sim_dataset) - len(self.self_lyric_sim_dataset) - len(self.artist_sim_dataset) - len(self.artist_vox_sim_dataset) ], "album_sim" elif index < len(self.self_sim_dataset) + len( self.self_vox_sim_dataset ) + len(self.self_lyric_sim_dataset) + len(self.artist_sim_dataset) + len( self.artist_vox_sim_dataset ) + len(self.album_sim_dataset) + len(self.genre_sim_dataset): return self.genre_sim_dataset[ index - len(self.self_sim_dataset) - len(self.self_vox_sim_dataset) - len(self.self_lyric_sim_dataset) - len(self.artist_sim_dataset) - len(self.artist_vox_sim_dataset) - len(self.album_sim_dataset) ], "genre_sim" elif index < len(self.self_sim_dataset) + len( self.self_vox_sim_dataset ) + len(self.self_lyric_sim_dataset) + len(self.artist_sim_dataset) + len( self.artist_vox_sim_dataset ) + len(self.album_sim_dataset) + len(self.genre_sim_dataset) + len( self.lyric_sim_dataset ): return self.lyric_sim_dataset[ index - len(self.self_sim_dataset) - len(self.self_vox_sim_dataset) - len(self.self_lyric_sim_dataset) - len(self.artist_sim_dataset) - len(self.artist_vox_sim_dataset) - len(self.album_sim_dataset) - len(self.genre_sim_dataset) ], "lyric_sim" def __len__(self): if self.split == "train": return self.num_samples elif self.split == "valid": return ( len(self.self_sim_dataset) * (self.sample_ratio[0] > 0) + len(self.self_vox_sim_dataset) * (self.sample_ratio[1] > 0) + len(self.self_lyric_sim_dataset) * (self.sample_ratio[2] > 0) + len(self.artist_sim_dataset) * (self.sample_ratio[3] > 0) + len(self.artist_vox_sim_dataset) * (self.sample_ratio[4] > 0) + len(self.album_sim_dataset) * (self.sample_ratio[5] > 0) + len(self.genre_sim_dataset) * (self.sample_ratio[6] > 0) )