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=15.0, self_vox_sim_input_length_s=15.0, self_lyric_sim_input_length=300, artist_sim_input_length_s=15.0, artist_vox_sim_input_length_s=15.0, album_sim_input_length_s=15.0, genre_sim_input_length_s=15.0, genre_tag_input_length=300, lyric_sim_input_length_s=15.0, lyric_sim_input_text_length=300, 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, ): 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.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 self.accumulated_lengths = self.get_accumulated_lengths() def get_accumulated_lengths(self): accumulated_lengths = [] current_sum = 0 for dataset in [ self.self_sim_dataset, self.self_vox_sim_dataset, self.self_lyric_sim_dataset, self.artist_sim_dataset, self.artist_vox_sim_dataset, self.album_sim_dataset, self.genre_sim_dataset, self.lyric_sim_dataset, ]: current_sum += len(dataset) accumulated_lengths.append(current_sum) return accumulated_lengths def __getitem__(self, index): if index < self.accumulated_lengths[0]: return self.self_sim_dataset[index], "self_sim" elif index < self.accumulated_lengths[1]: return self.self_vox_sim_dataset[ index - self.accumulated_lengths[0] ], "self_vox_sim" elif index < self.accumulated_lengths[2]: return self.self_lyric_sim_dataset[ index - self.accumulated_lengths[1] ], "self_lyric_sim" elif index < self.accumulated_lengths[3]: return self.artist_sim_dataset[ index - self.accumulated_lengths[2] ], "artist_sim" elif index < self.accumulated_lengths[4]: return self.artist_vox_sim_dataset[ index - self.accumulated_lengths[3] ], "artist_vox_sim" elif index < self.accumulated_lengths[5]: return self.album_sim_dataset[ index - self.accumulated_lengths[4] ], "album_sim" elif index < self.accumulated_lengths[6]: return self.genre_sim_dataset[ index - self.accumulated_lengths[5] ], "genre_sim" else: return self.lyric_sim_dataset[ index - self.accumulated_lengths[6] ], "lyric_sim" def __len__(self): return self.accumulated_lengths[-1] class MultiTaskBatchSampler(Sampler): def __init__(self, dataset, batch_size): self.dataset = dataset self.batch_size = batch_size self.task_sizes = [ len(dataset.self_sim_dataset), len(dataset.self_vox_sim_dataset), len(dataset.self_lyric_sim_dataset), len(dataset.artist_sim_dataset), len(dataset.artist_vox_sim_dataset), len(dataset.album_sim_dataset), len(dataset.genre_sim_dataset), len(dataset.lyric_sim_dataset), ] self.task_indices = list(range(len(self.task_sizes))) def __iter__(self): if self.dataset.split == "train": while True: # Randomly choose a task task = random.choice(self.task_indices) # Calculate the start and end indices for the chosen task start_idx = sum(self.task_sizes[:task]) end_idx = start_idx + self.task_sizes[task] # Generate a batch of indices for the chosen task batch_indices = random.sample( range(start_idx, end_idx), min(self.batch_size, self.task_sizes[task]), ) yield batch_indices else: # valid task_index = 0 while True: # Calculate the start and end indices for the current task start_idx = sum(self.task_sizes[:task_index]) end_idx = start_idx + self.task_sizes[task_index] # Generate batches for the current task for i in range(start_idx, end_idx, self.batch_size): batch_indices = list(range(i, min(i + self.batch_size, end_idx))) yield batch_indices # Move to the next task task_index = (task_index + 1) % len(self.task_indices) def __len__(self): return sum(self.task_sizes) // self.batch_size