from collections import defaultdict import json import math import os import random import re from dataclasses import dataclass from tqdm import tqdm import numpy as np import orjson from tokenizers import AddedToken import torch import torch.nn.functional as F from torch.utils.data import IterableDataset from transformers import PreTrainedTokenizerFast from oracle_dataset import get_sample_oracle_file_segment from modules.gpt import GPTConfig, GPTTrainConfig from utils.helpers import print_with_time_master from utils.bct import Block, BlockSequence, PackedBlockSequence, BlockType # to avoid: "The current process just got forked, after parallelism has already been used" os.environ["TOKENIZERS_PARALLELISM"] = "False" RELOAD_MEMMAP = False global tokenizer g_tokenizer = None TextBlockType = BlockType( name="text", is_causal=True, ) CausalSemanticBlockType = BlockType( name="semantic", is_causal=True, ) ArtistBlockType = BlockType( name="artist", is_causal=True, ) PlaylistBlockType = BlockType( name="playlist", is_causal=True, ) UnderpaintBlockType = BlockType( name="underpaint", is_causal=True, ) OverpaintBlockType = BlockType( name="overpaint", is_causal=True, ) CoverBlockType = BlockType( name="cover", is_causal=True, ) PrefixBlockType = BlockType( name="prefix", is_causal=True, ) SuffixBlockType = BlockType( name="suffix", is_causal=True, ) NonCausalSemanticBlockType = BlockType( name="non_causal_semantic", is_causal=False, ) # ++ Define a dataclass to hold sampling parameters ++ @dataclass class SamplingParams: # Flags controlling data augmentation and task types allow_infill: bool = False allow_artist: bool = False allow_cover: bool = False allow_overpaint: bool = False allow_underpaint: bool = False allow_skip: bool = False allow_playlist: bool = False text_loss: bool = False # Renamed from allow_text_loss for consistency if desired # probs prob_infill: float = 0.5 # Used when random.random() <= 0.25 prob_artist: float = 0.1 # Used when random.random() < p_use_artist (0.1 or 0.5) prob_cover: float = 0.5 # Used when random.random() < 0.5 prob_overpaint: float = 0.5 # Used when random.random() < 0.5 prob_underpaint: float = 0.5 # Used when random.random() < 0.5 prob_skip: float = 0.05 # Used when random.random() <= 0.1 prob_playlist: float = 0.2 # Used when random.random() < 0.1 # Flags controlling sampling mode inference: bool = False dummy_data: bool = False suppress_text: bool = False dataset_idx: int | None = None # Flags/params controlling batching and padding pack: bool = False mask_padding: bool = False min_text_offs: int | None = None # Add any other related parameters that are frequently passed together MAX_TAG_LEN = 1024 # characters, not tokens MAX_TOT_TAGS_LEN = 2048 # characters, not tokens MAX_N_TAGS = 10 # to avoid overfitting to an artist def _load_tokenizer(tokenizer_fp=None): global g_tokenizer if g_tokenizer is not None: return g_tokenizer assert os.path.exists(tokenizer_fp) g_tokenizer = PreTrainedTokenizerFast( tokenizer_file=tokenizer_fp, unk_token="[UNK]", pad_token="[PAD]", ) g_tokenizer.add_special_tokens({"additional_special_tokens": [AddedToken("\n")]}) return g_tokenizer def _space_repl(m): s = m.group() n_newline = s.count("\n") if n_newline >= 2: return "\n\n" elif n_newline == 1: return "\n" return " " def _simplify_whitespace(text, retain_newlines=True): """simplify while respecting up to 2 newlines""" if retain_newlines: text = re.sub(r"\s+", _space_repl, text).strip() else: text = re.sub(r"\s+", " ", text).strip() return text def tokenize_batch( text_list, max_tokens=None, pad_token_id=0, retain_newlines=True, tokenizer_fp=None, ): tokenizer = _load_tokenizer(tokenizer_fp) text_list = [_simplify_whitespace(s, retain_newlines=retain_newlines) for s in text_list] text_enc = tokenizer( text_list, add_special_tokens=False, truncation=True, max_length=max_tokens, padding="longest", return_tensors="pt", )["input_ids"].type(torch.long) text_enc[text_enc == tokenizer.pad_token_id] = pad_token_id return text_enc def read_jsonl(filepath, parse_idx_set=None, max_lines=None, progress_bar=False): data = [] with open(filepath) as f: line_idx = 0 for line in tqdm(f, total=max_lines, disable=not progress_bar, mininterval=5): line = line.strip() if len(line) == 0: continue if parse_idx_set is not None and line_idx not in parse_idx_set: data.append(None) line_idx += 1 continue m = json.loads(line) data.append(m) line_idx += 1 if max_lines is not None and line_idx >= max_lines: break return data def write_jsonl(data, filepath): with open(filepath, "w") as f: for d in data: f.write(json.dumps(d, ensure_ascii=False) + "\n") def load_dataset( load_kwargs: dict, data_dir: str, filename: str, info_filename: str, metas_filename: str, weights_multiplier_map: dict, is_finetune: bool, use_raw_audio: bool = False, ) -> tuple: dataset_names = [] data_idx_lists = [] # used to randomly sample from the dataset data_weights = [] data = None if not use_raw_audio: data = np.memmap(os.path.join(data_dir, filename), dtype=np.uint16, mode="r") if load_kwargs["semantic_n_codebooks"] + load_kwargs["coarse_n_codebooks"] > 1: data = data.reshape( -1, load_kwargs["semantic_n_codebooks"] + load_kwargs["coarse_n_codebooks"], ) assert ( data[:100, :, : load_kwargs["semantic_n_codebooks"]].max() <= load_kwargs["semantic_vocab_size"] ) assert ( data[:100, :, load_kwargs["semantic_n_codebooks"] :].max() <= load_kwargs["coarse_vocab_size"] ) elif load_kwargs["semantic_n_codebooks"] == 1 and load_kwargs["coarse_n_codebooks"] == 0: assert data[:100].max() <= load_kwargs["semantic_vocab_size"] with open(os.path.join(data_dir, info_filename)) as f: infos = json.load(f) # make sure we turn into int since keys in json get auto turned into strings for k in infos.keys(): if "idx_map" in infos[k]: infos[k]["idx_map"] = {int(k): v for k, v in infos[k]["idx_map"].items()} for use_condition, task_name in zip( ("allow_cover", "allow_overpaint", "allow_underpaint"), ("covers", "overpaint", "underpaint"), ): if not load_kwargs[use_condition]: # we can shard the data and cache data with cover even if we don't want to use covers # so we remove the covers from the infos task_keys = [ dset_name for dset_name in infos.keys() if infos[dset_name].get("task", "default") == task_name ] for task_key in task_keys: infos.pop(task_key) metas = read_jsonl(os.path.join(data_dir, metas_filename)) if use_raw_audio: assert all("s3_filepath" in m for m in metas) artist_to_songs = defaultdict(list) for i, m in enumerate(metas): if "artists" in m: dset_prefix = m["dataset"].split("_")[0] artist_to_songs[f"{dset_prefix}__{'__'.join(sorted(m['artists']))}"].append(i) artist_to_songs = {k: v for k, v in artist_to_songs.items() if len(v) > 1} if load_kwargs["allow_artist"]: assert len(artist_to_songs) > 0, "no artist data found" print_with_time_master(f"found {len(artist_to_songs):,} samples with matching artists found") playlist_to_songs = defaultdict(list) for i, m in enumerate(metas): for playlist_id in m.get("playlist_ids", []): playlist_to_songs[playlist_id].append(i) playlist_to_songs = {k: v for k, v in playlist_to_songs.items() if len(v) > 1} if load_kwargs["allow_playlist"]: assert len(playlist_to_songs) > 0, "no playlist data found" print_with_time_master(f"found {len(playlist_to_songs):,} samples with matching playlists found") idx_set = set() has_cover = False has_overpaint = False has_underpaint = False for dset_name, info in infos.items(): dataset_names.append(dset_name) if info.get("task", "default") == "default": idx_list = info["idx_list"][:] idx_set |= set(idx_list) elif info["task"] == "covers": if not load_kwargs["allow_cover"]: continue assert load_kwargs["pack"], "for now pack needs to be active to do covers" has_cover = True idx_list = [] n_covers = 0 for idx, child_idx_l in info["idx_map"].items(): dset_prefix = dset_name.split("_")[0] idx_list.append(idx) idx_set.add(idx) idx_set |= set(child_idx_l) n_covers += len(child_idx_l) info["idx_list"] = idx_list print_with_time_master(f"found {len(idx_list):,} samples with {n_covers:,} total covers") elif info["task"] == "overpaint": if not load_kwargs["allow_overpaint"]: continue assert load_kwargs["pack"], "for now pack needs to be active to do overpaint" has_overpaint = True idx_list = [] for idx, child_idx in info["idx_map"].items(): # full song to instrumental map dset_prefix = dset_name.split("_")[0] idx_list.append(idx) idx_set.add(idx) idx_set.add(child_idx) info["idx_list"] = idx_list print_with_time_master(f"found {len(idx_list):,} total overpaints") elif info["task"] == "underpaint": if not load_kwargs["allow_underpaint"]: continue assert load_kwargs["pack"], "for now pack needs to be active to do underpaint" has_underpaint = True idx_list = [] for idx, child_idx in info["idx_map"].items(): # full song to vocals map dset_prefix = dset_name.split("_")[0] idx_list.append(idx) idx_set.add(idx) idx_set.add(child_idx) info["idx_list"] = idx_list print_with_time_master(f"found {len(idx_list):,} total underpaint") else: raise ValueError(f"unknown task for {dset_name} in info file") random.shuffle(idx_list) data_idx_lists.append(idx_list) data_weights.append(len(idx_list) * weights_multiplier_map.get(dset_name, 1.0)) if load_kwargs["allow_cover"]: assert has_cover, "no cover data found" if load_kwargs["allow_overpaint"]: assert has_overpaint, "no overpaint data found" if load_kwargs["allow_underpaint"]: assert has_underpaint, "no underpaint data found" weights_norm = np.sum(data_weights) data_weights = [v / weights_norm for v in data_weights] if not is_finetune: print_with_time_master(f"indexed {len(idx_set) / len(metas) * 100:.1f}% of data") for k in weights_multiplier_map.keys(): assert k in dataset_names print_with_time_master(f"weight for {k}: {weights_multiplier_map[k]}") # some checks on data vs metas if not use_raw_audio: for m in metas: assert m["offset_idx"] <= len(data) assert 0.99 <= max([m["offset_idx"] for m in metas]) / len(data) <= 1.0 assert len(infos) == len(dataset_names) == len(data_weights) == len(data_idx_lists) shard_info = "" if load_kwargs["local_data_shard_dir"] is None else " (sharded)" print_with_time_master(f"{len(metas):,} lines of {metas_filename} loaded.{shard_info}") del idx_set return ( dataset_names, data_idx_lists, data_weights, data, metas, infos, artist_to_songs, playlist_to_songs, ) CASE_AUGMENT_FUNCS = [ str.upper, str.lower, str.capitalize, str.title, ] def _augment_tag(s): # case augment if random.random() >= 0.8: s = random.choice(CASE_AUGMENT_FUNCS)(s) # other misc formatting if random.random() >= 0.5: s = s.replace("-", " ").strip() return s def _get_control_tags(duration_s, sample_vocal_start_s, do_augment=True): control_tags = [] control_tags.append(f"duration:{int(round(duration_s))}") min_durations = [n * 60 for n in range(10) if n * 60 <= duration_s] max_durations = [n * 60 for n in range(10) if n * 60 >= duration_s] if len(min_durations) > 0: control_tags.append(f"min_duration:{int(random.choice(min_durations))}") if len(max_durations) > 0: control_tags.append(f"max_duration:{int(random.choice(max_durations))}") if sample_vocal_start_s is not None: if sample_vocal_start_s <= 5: control_tags.append("vocals:early") if sample_vocal_start_s <= 15: control_tags.append("vocals:normal") if 10 <= sample_vocal_start_s <= 20: control_tags.append("vocals:intro") if do_augment: if random.random() >= 0.5: random.shuffle(control_tags) control_tags = control_tags[: random.randint(0, len(control_tags))] if len(control_tags) == 0: return None return "{" + ";".join(control_tags) + "}" def mask_middle_padding(stream: np.ndarray, old_pad_id: int, new_pad_id=-1): B, C, T = stream.shape shift_left = stream[:, :, :-2] == old_pad_id shift_right = stream[:, :, 2:] == old_pad_id current = stream[:, :, 1:-1] == old_pad_id # The middle mask now checks if the current (center) token # and its immediate neighbors (left and right) are all padding tokens. # We use logical AND on shifted views of the 'stream' array. middle_mask = shift_left & current & shift_right # Apply the mask to a slice of the stream to avoid affecting the edges, # since the original mask does not include the first and last columns. stream[:, :, 1:-1][middle_mask] = new_pad_id return stream def mask_batch_padding(batch: np.ndarray, cfg): """batch has shape (B, C, T) where C is (optional text) + n_semantic + n_coarse streams""" B, C, T = batch.shape if C == 1 + cfg.semantic_n_codebooks + cfg.coarse_n_codebooks: text_offs = 1 mask_middle_padding(batch[:, :text_offs], cfg.text_pad_token) elif C == cfg.semantic_n_codebooks + cfg.coarse_n_codebooks: text_offs = 0 else: raise ValueError() mask_middle_padding( batch[:, text_offs : text_offs + cfg.semantic_n_codebooks], cfg.semantic_pad_token ) mask_middle_padding(batch[:, text_offs + cfg.semantic_n_codebooks :], cfg.coarse_pad_token) # mask semantic mask batch[:, text_offs : text_offs + cfg.semantic_n_codebooks][ batch[:, text_offs : text_offs + cfg.semantic_n_codebooks] == cfg.semantic_mask_token ] = -1 return batch def _augment_tags(tags): random.shuffle(tags) if random.random() <= 0.5: tags = tags[: random.randint(0, len(tags))] tags = [_augment_tag(tag) for tag in tags] return tags def _clean_tags(tags): return [ clean_tag[:MAX_TAG_LEN] for tag in tags if len(clean_tag := _simplify_whitespace(tag, retain_newlines=False)) > 0 ] def _clean_inline_tags(m): tags = m.group(2).split(";") tags = _clean_tags(tags) ts = ";".join(tags)[:MAX_TOT_TAGS_LEN] if len(ts) > 0: return f"[{m.group(1)}: {ts}]" return f"[{m.group(1)}]" def _augment_inline_tags(m): tags = m.group(2).split(";") tags = _augment_tags(tags) merge_char = random.choice([", ", " ", "; ", ",", ";", ". "]) ts = merge_char.join(tags) if len(ts) > 0: return f"[{m.group(1)}: {ts}]" return f"[{m.group(1)}]" def build_text( tags, text, duration_s, sample_vocal_start_s, inference, suppress_text, enable_control_tags=True, passin_control_tags=None, ): if suppress_text: return "" text_elements = [] # collect tags tags = _clean_tags(tags) if not inference: tags = _augment_tags(tags) merge_char = random.choice([", ", " ", "; ", ",", ";", ". "]) tags_str = f"{merge_char.join(tags[:MAX_N_TAGS])}"[:MAX_TOT_TAGS_LEN] else: # keep the non-inference behavoir tags_str = f"{','.join([tag[:MAX_TAG_LEN] for tag in tags])[:MAX_TOT_TAGS_LEN]}" if len(tags_str) > 0: text_elements.append(f"[{tags_str}]") # get lyrics if len(text) > 0 and not inference: text = re.sub(r"\[(.*?)\:(.*?)\]", _clean_inline_tags, text) # augment tags inside text: text = re.sub(r"\[(.*?)\:(.*?)\]", _augment_inline_tags, text) if random.random() >= 0.95: text = text.lower() if random.random() >= 0.95: text = re.sub(r"\n+", " ", text) if len(text) > 0: text_elements.append(text.strip()) # get control tags for n in range(len(text_elements)): text_elements[n] = text_elements[n].replace("{", "").replace("}", "") if (inference or random.random() >= 0.1) and enable_control_tags: control_tags = _get_control_tags(duration_s, sample_vocal_start_s, do_augment=not inference) if control_tags is not None: text_elements = [control_tags] + text_elements # this is hard coded control tags that are passed in by the user if passin_control_tags: text_elements = [passin_control_tags] + text_elements if inference or random.random() >= 0.5: text = "\n\n".join(text_elements) else: text = "" for t in text_elements: text += random.choice([" ", "\n", "\n\n"]) + t text = text.strip() return text def pad_audio_arr(arr, cfg): """Pads an arr of size (C, T) to (C, cfg.t_audio)""" C, T = arr.shape assert T <= cfg.t_audio assert C == cfg.semantic_n_codebooks + cfg.coarse_n_codebooks pad = np.empty((C, cfg.t_audio - T), dtype=np.int64) pad[: cfg.semantic_n_codebooks] = cfg.semantic_pad_token pad[cfg.semantic_n_codebooks :] = cfg.coarse_pad_token return np.concatenate([arr, pad], axis=-1) def pad_x_arr(x_arr, cfg, sz: int = None): """Pads an x_arr of size (C, T) to (C, cfg.block_size)""" C, T = x_arr.shape if sz is None: sz = cfg.block_size assert C == cfg.semantic_n_codebooks + cfg.coarse_n_codebooks + 1 pad = np.empty((C, sz - T), dtype=np.int64) pad[0] = cfg.text_pad_token pad[1 : 1 + cfg.semantic_n_codebooks] = cfg.semantic_pad_token pad[1 + cfg.semantic_n_codebooks :] = cfg.coarse_pad_token return np.concatenate([x_arr, pad], axis=-1) def build_audio_arr(data_row, cfg, semantic_infer_token=None, include_eos=True, skip_factor=1): """Builds an audio arr from a data row. Prepends the semantic and coarse streams with the infer token.""" data_row = data_row[:, ::skip_factor] audio_len = ( cfg.semantic_n_codebooks * cfg.semantic_shift_factor * min(1, cfg.coarse_n_codebooks) + max((cfg.coarse_n_codebooks - 1), 0) * cfg.coarse_shift_factor + data_row[-1].shape[-1] ) if cfg.coarse_n_codebooks == 0 and include_eos: audio_len += 1 # used as eos token # build semantic y_semantic_arr = np.full( (cfg.semantic_n_codebooks, audio_len), cfg.semantic_pad_token, dtype=np.int64 ) for n in range(cfg.semantic_n_codebooks): offs = n * cfg.semantic_shift_factor y_semantic_arr[n, offs : offs + data_row[n].shape[-1]] = data_row[n] # build coarse if cfg.coarse_n_codebooks > 0: y_coarse_arr = np.full((cfg.coarse_n_codebooks, audio_len), cfg.coarse_pad_token, dtype=np.int64) for n in range(cfg.coarse_n_codebooks): offs = cfg.semantic_n_codebooks * cfg.semantic_shift_factor + n * cfg.coarse_shift_factor n2 = cfg.semantic_n_codebooks + n y_coarse_arr[n, offs : offs + data_row[n2].shape[-1]] = data_row[n2] # combine audio and add x with infer token if semantic_infer_token is None: semantic_infer_token = cfg.semantic_infer_token y_audio_arr = y_semantic_arr if cfg.coarse_n_codebooks > 0: y_audio_arr = np.concatenate([y_audio_arr, y_coarse_arr], axis=0) x_audio_arr = y_audio_arr.copy() x_audio_arr = np.concatenate( [ np.array( [[semantic_infer_token]] * cfg.semantic_n_codebooks + [[cfg.coarse_infer_token]] * cfg.coarse_n_codebooks ).repeat(skip_factor, 1), x_audio_arr, ], axis=-1, ) return x_audio_arr def _get_rand_int(min_val, max_val): if max_val <= min_val: return min_val return random.randint(min_val, max_val) @dataclass class SampleData: data_row: np.ndarray data_meta: dict sampling_params: SamplingParams is_full_track: bool = False data_row_cover: np.ndarray | None = None artist_tracks: list[np.ndarray] | None = None playlist_tracks: list[np.ndarray] | None = None overpaint_track: np.ndarray | None = None underpaint_track: np.ndarray | None = None idx: int | None = None def get_sample_codes_or_audio( data_sampling_info, split, sampling_params: SamplingParams, dataset_idx=None, row_idx=None, # absolute, overrides dataset_idx ): data = data_sampling_info[split]["data"] metas = data_sampling_info[split]["metas"] idx_lists = data_sampling_info[split]["idx_lists"] infos = data_sampling_info[split]["infos"] weights = data_sampling_info[split]["weights"] dset_names = data_sampling_info[split]["names"] artist_to_songs = data_sampling_info[split]["artist_to_songs"] playlist_to_songs = data_sampling_info[split]["playlist_to_songs"] cfg = data_sampling_info["cfg"] if row_idx is not None: pass elif dataset_idx is not None: row_idx = random.choice(idx_lists[dataset_idx]) else: weights = weights[:] for allow_task, task_name in zip( ( sampling_params.allow_cover, sampling_params.allow_overpaint, sampling_params.allow_underpaint, ), ("covers", "overpaint", "underpaint"), ): if not allow_task: # set weight for cover dataset to 0 for n, dset_name in enumerate(dset_names): if infos.get(dset_name, {}).get("task", "default") == task_name: weights[n] = 0 assert sum(weights) > 0 dataset_idx = random.choices(list(range(len(weights))), weights=weights, k=1)[0] row_idx = random.choice(idx_lists[dataset_idx]) skip_factor = 1 if ( sampling_params.allow_skip and not sampling_params.inference and random.random() <= sampling_params.prob_skip ): skip_factor = 4 if random.random() <= 0.5 else 2 def load_data_row(row_idx, dummy_data=False): data_meta = metas[row_idx] offset_idx = None if data is None: # TODO: fix this once we encode and know after use_n_tokens = int(math.floor(data_meta["duration_s"] * cfg.semantic_rate_hz)) else: offset_idx = data_meta["offset_idx"] use_n_tokens = data_meta["n_tokens"] phase_offset = random.choice(list(range(0, skip_factor))) use_n_tokens = min(math.floor((use_n_tokens - phase_offset) / skip_factor), cfg.t_audio) is_full_track = use_n_tokens < cfg.t_audio if dummy_data: data_row = np.zeros( (cfg.semantic_n_codebooks + cfg.coarse_n_codebooks + 1, use_n_tokens), dtype=np.int64, ) else: if data is None: audio = get_sample_oracle_file_segment(data_meta["s3_filepath"]) audio = audio.convert(24000, 2, 1) data_row = audio.array_float else: data_row = ( data[ offset_idx + phase_offset : offset_idx + phase_offset + use_n_tokens * skip_factor : skip_factor ] .astype(np.int64) .reshape(cfg.semantic_n_codebooks + cfg.coarse_n_codebooks, -1) ) return data_row, data_meta, is_full_track data_row, data_meta, is_full_track = load_data_row(row_idx, dummy_data=sampling_params.dummy_data) if ( sampling_params.allow_cover and not sampling_params.inference and dataset_idx is not None and (dataset_name := data_sampling_info[split]["names"][dataset_idx]) and infos[dataset_name].get("task", "default") == "covers" ): # sample a random cover cover_list = infos[dataset_name]["idx_map"][row_idx] cover_idx = random.choice(cover_list) data_row_cover, _, _ = load_data_row(cover_idx, dummy_data=sampling_params.dummy_data) else: data_row_cover = None used_cover = data_row_cover is not None # Artist audio to audio dset_prefix = data_meta["dataset"].split("_")[0] p_use_artist = ( sampling_params.prob_artist * 5 if used_cover else sampling_params.prob_artist ) # increase chance for cover cause voice beautifier artist_tracks = [] if ( sampling_params.allow_artist and not sampling_params.inference and "artists" in data_meta and random.random() < p_use_artist and len(artist_to_songs.get(f"{dset_prefix}__{'__'.join(sorted(data_meta['artists']))}", [])) >= 2 ): # sample a random song from the artist track_idx_list = [ _idx for _idx in artist_to_songs[f"{dset_prefix}__{'__'.join(sorted(data_meta['artists']))}"] ] track_idx_list.remove(row_idx) if len(track_idx_list) > 0: # sample multiple artist segments n_max_tracks = 5 max_track_duration_s = 60 # assert max_track_duration_s * n_max_tracks * cfg.semantic_rate_hz <= cfg.t_audio for _ in range(random.randint(1, n_max_tracks)): track_idx = random.choice(list(track_idx_list)) data_row_track, _, _ = load_data_row(track_idx, dummy_data=sampling_params.dummy_data) # use only part of artist left_idx = random.randint(0, data_row_track.shape[-1]) use_track_duration_s = random.randint( 0, min(data_row_track.shape[-1] - left_idx, max_track_duration_s) ) right_idx = random.randint(left_idx, left_idx + use_track_duration_s) assert right_idx - left_idx <= max_track_duration_s, (right_idx, left_idx) data_row_track = data_row_track[:, left_idx:right_idx] if not 0 <= left_idx and left_idx < right_idx and right_idx <= data_row_track.shape[-1]: continue artist_tracks.append(data_row_track) # Playlist audio to audio playlist_tracks = [] if ( sampling_params.allow_playlist and not sampling_params.inference and "playlist_ids" in data_meta and random.random() < sampling_params.prob_playlist and len( [ True for playlist_id in data_meta["playlist_ids"] if len(playlist_to_songs.get(playlist_id, [])) >= 2 ] ) >= 1 ): # sample a random song from the artist playlist_ids = [ playlist_id for playlist_id in data_meta["playlist_ids"] if len(playlist_to_songs.get(playlist_id, [])) >= 2 ] playlist_id = random.choice(playlist_ids) track_idx_list = [_idx for _idx in playlist_to_songs[playlist_id]] track_idx_list.remove(row_idx) if len(track_idx_list) > 0: # sample multiple artist segments n_max_tracks = 5 max_track_duration_s = 60 # assert max_track_duration_s * n_max_tracks * cfg.semantic_rate_hz <= cfg.t_audio for _ in range(random.randint(1, n_max_tracks)): track_idx = random.choice(list(track_idx_list)) data_row_track, _, _ = load_data_row(track_idx, dummy_data=sampling_params.dummy_data) # use only part of artist left_idx = random.randint(0, data_row_track.shape[-1]) use_track_duration_s = random.randint( 0, min(data_row_track.shape[-1] - left_idx, max_track_duration_s) ) right_idx = random.randint(left_idx, left_idx + use_track_duration_s) assert right_idx - left_idx <= max_track_duration_s, (right_idx, left_idx) data_row_track = data_row_track[:, left_idx:right_idx] if not 0 <= left_idx and left_idx < right_idx and right_idx <= data_row_track.shape[-1]: continue playlist_tracks.append(data_row_track) # Overpaint overpaint_track = None if ( sampling_params.allow_overpaint and not sampling_params.inference and dataset_idx is not None and (dataset_name := data_sampling_info[split]["names"][dataset_idx]) and infos[dataset_name].get("task", "default") == "overpaint" ): # map is instrumental to full row_idx_child = infos[dataset_name]["idx_map"][row_idx] overpaint_track, _, _ = load_data_row(row_idx_child, dummy_data=sampling_params.dummy_data) # Underpaint underpaint_track = None if ( sampling_params.allow_underpaint and not sampling_params.inference and dataset_idx is not None and (dataset_name := data_sampling_info[split]["names"][dataset_idx]) and infos[dataset_name].get("task", "default") == "underpaint" ): # map is vocals to full song row_idx_child = infos[dataset_name]["idx_map"][row_idx] underpaint_track, _, _ = load_data_row(row_idx_child, dummy_data=sampling_params.dummy_data) sample_data = SampleData( data_row=data_row, data_meta=data_meta, sampling_params=sampling_params, skip_factor=skip_factor, is_full_track=is_full_track, data_row_cover=data_row_cover, artist_tracks=artist_tracks, playlist_tracks=playlist_tracks, overpaint_track=overpaint_track, underpaint_track=underpaint_track, idx=row_idx, ) return sample_data def get_sample_from_row(model_cfg: GPTConfig, tokenizer_fp: str, sample_data: SampleData): cfg = model_cfg sampling_params = sample_data.sampling_params data_row = sample_data.data_row data_meta = sample_data.data_meta def sample_skip_factor(): if sampling_params.allow_skip and random.random() <= sampling_params.prob_skip: return random.randint(2, 6) return 1 is_full_track = sample_data.is_full_track data_row_cover = sample_data.data_row_cover artist_tracks = sample_data.artist_tracks playlist_tracks = sample_data.playlist_tracks overpaint_track = sample_data.overpaint_track underpaint_track = sample_data.underpaint_track if data_row.shape[-1] > cfg.block_size - cfg.t_text: print( f"WARNING: data row is long {data_meta}: {data_row.shape[-1]} > {cfg.block_size - cfg.t_text}" ) data_row = data_row[:, : cfg.block_size - cfg.t_text] # build main audio array sample_tags = data_meta.get("tags", []) sample_vocal_start_s = None audio_blocks = [] if ( sampling_params.allow_infill and not sampling_params.inference and data_row.shape[-1] >= 1 and random.random() <= sampling_params.prob_infill ): # a_b_c make b, where a and c can both be empty to allow arbitrary handling t = data_row.shape[-1] b_duration = _get_rand_int(1, t) a_right_idx = _get_rand_int(0, t - b_duration) b_right_idx = a_right_idx + b_duration if random.random() <= 0.05: a_left_idx = 0 elif random.random() <= 0.05: a_left_idx = a_right_idx else: a_left_idx = _get_rand_int(0, a_right_idx) if random.random() <= 0.05: c_right_idx = t elif random.random() <= 0.05: c_right_idx = b_right_idx else: c_right_idx = _get_rand_int(b_right_idx, t) sample_text = "" if len(data_meta.get("text_lines", [])) > 0: text_lines = data_meta["text_lines"] text_left_idx = 0 text_right_idx = len(text_lines) for idx, m_line in enumerate(text_lines): if m_line["start_s"] <= a_right_idx / cfg.semantic_rate_hz: text_left_idx = idx if m_line["end_s"] >= b_right_idx / cfg.semantic_rate_hz: text_right_idx = idx + 1 break text_left_idx = _get_rand_int(0, text_left_idx) text_right_idx = _get_rand_int(text_right_idx, len(text_lines)) sample_text = "\n".join([m["text"] for m in text_lines[text_left_idx:text_right_idx]]) elif "text" in data_meta: sample_text = data_meta["text"] a_audio_arr = build_audio_arr( data_row[:, a_left_idx:a_right_idx], cfg, semantic_infer_token=cfg.semantic_history_token, include_eos=False, skip_factor=sample_skip_factor(), ) c_audio_arr = build_audio_arr( data_row[:, b_right_idx:c_right_idx], cfg, semantic_infer_token=cfg.semantic_future_token, include_eos=False, skip_factor=sample_skip_factor(), ) b_audio_arr = build_audio_arr( data_row[:, a_right_idx:b_right_idx], cfg, include_eos=True, skip_factor=sample_skip_factor(), ) x_audio_arr = np.concatenate([c_audio_arr, a_audio_arr, b_audio_arr], axis=-1) a_audio_arr = torch.from_numpy(a_audio_arr) b_audio_arr = torch.from_numpy(b_audio_arr) c_audio_arr = torch.from_numpy(c_audio_arr) audio_blocks.extend( [ Block( spec=SuffixBlockType, inputs={"semantic_input": c_audio_arr}, targets={"semantic_output": Block.shift_left(c_audio_arr, cfg.semantic_pad_token)}, ), Block( spec=PrefixBlockType, inputs={"semantic_input": a_audio_arr}, targets={"semantic_output": Block.shift_left(a_audio_arr, cfg.semantic_pad_token)}, ), Block( spec=CausalSemanticBlockType, inputs={"semantic_input": b_audio_arr}, targets={"semantic_output": Block.shift_left(b_audio_arr, cfg.semantic_pad_token)}, ), ] ) sample_duration_s = (c_right_idx - a_left_idx) / cfg.semantic_rate_hz text = build_text( sample_tags, sample_text, sample_duration_s, sample_vocal_start_s, sampling_params.inference, sampling_params.suppress_text, ) else: if "text_lines" in data_meta and len(data_meta["text_lines"]) > 0: sample_vocal_start_s = data_meta["text_lines"][0]["start_s"] x_audio_arr = build_audio_arr( data_row, cfg, include_eos=is_full_track, skip_factor=sample_skip_factor() ) x_audio_arr = torch.from_numpy(x_audio_arr) audio_blocks.append( Block( spec=CausalSemanticBlockType, inputs={"semantic_input": x_audio_arr}, targets={"semantic_output": Block.shift_left(x_audio_arr, cfg.semantic_pad_token)}, ) ) sample_text = data_meta.get("text", "") sample_duration_s = data_row.shape[-1] / cfg.semantic_rate_hz text = build_text( sample_tags, sample_text, sample_duration_s, sample_vocal_start_s, sampling_params.inference, sampling_params.suppress_text, ) # build text arr if not sampling_params.inference and random.random() >= 0.98: # sometimes do unconditional text = "" x_text_arr = tokenize_batch( [text], max_tokens=cfg.t_text, pad_token_id=cfg.text_pad_token, tokenizer_fp=tokenizer_fp, )[0] # (N_TEXT_TOKENS) if ( sampling_params.min_text_offs is not None and sampling_params.min_text_offs > x_text_arr.shape[-1] ): x_text_arr = F.pad( x_text_arr, (0, sampling_params.min_text_offs - x_text_arr.shape[-1]), "constant", cfg.text_pad_token, ) if sampling_params.text_loss: x_text_arr = F.pad( x_text_arr, (1, 0), "constant", cfg.text_infer_token, ) if x_text_arr[-1] != cfg.text_pad_token: x_text_arr = F.pad( x_text_arr, (0, 1), "constant", cfg.text_pad_token, ) # Covers if data_row_cover is not None: if random.random() <= 0.5: # randomly drop out some of the cover data_row_cover = data_row_cover[:, : random.randint(0, data_row_cover.shape[-1])] # prepend cover to x_audio_arr cover_audio_arr = build_audio_arr( data_row_cover, cfg, semantic_infer_token=cfg.semantic_cover_token, include_eos=False, skip_factor=sample_skip_factor(), ) cover_audio_arr = torch.from_numpy(cover_audio_arr) audio_blocks.insert( 0, Block( spec=CoverBlockType, inputs={"semantic_input": cover_audio_arr}, targets={"semantic_output": Block.shift_left(cover_audio_arr, cfg.semantic_pad_token)}, ), ) # Artist audio to audio if artist_tracks is not None: for data_row_track in artist_tracks: artist_audio_arr = build_audio_arr( data_row_track, cfg, semantic_infer_token=cfg.semantic_artist_token, include_eos=False, skip_factor=sample_skip_factor(), ) artist_audio_arr = torch.from_numpy(artist_audio_arr) audio_blocks.insert( 0, Block( spec=ArtistBlockType, inputs={"semantic_input": artist_audio_arr}, targets={ "semantic_output": Block.shift_left(artist_audio_arr, cfg.semantic_pad_token) }, ), ) # Playlist audio to audio if playlist_tracks is not None: for data_row_track in playlist_tracks: playlist_audio_arr = build_audio_arr( data_row_track, cfg, semantic_infer_token=cfg.semantic_playlist_token, include_eos=False, skip_factor=sample_skip_factor(), ) playlist_audio_arr = torch.from_numpy(playlist_audio_arr) audio_blocks.insert( 0, Block( spec=PlaylistBlockType, inputs={"semantic_input": playlist_audio_arr}, targets={ "semantic_output": Block.shift_left(playlist_audio_arr, cfg.semantic_pad_token) }, ), ) # Overpaint if overpaint_track is not None: # prepend instrumental to x_audio_arr overpaint_audio_arr = build_audio_arr( overpaint_track, cfg, semantic_infer_token=cfg.semantic_overpaint_token, include_eos=False, skip_factor=sample_skip_factor(), ) overpaint_audio_arr = torch.from_numpy(overpaint_audio_arr) audio_blocks.insert( 0, Block( spec=OverpaintBlockType, inputs={"semantic_input": overpaint_audio_arr}, targets={ "semantic_output": Block.shift_left(overpaint_audio_arr, cfg.semantic_pad_token) }, ), ) # Underpaint if underpaint_track is not None: # prepend vocals to x_audio_arr underpaint_audio_arr = build_audio_arr( underpaint_track, cfg, semantic_infer_token=cfg.semantic_underpaint_token, include_eos=False, skip_factor=sample_skip_factor(), ) underpaint_audio_arr = torch.from_numpy(underpaint_audio_arr) audio_blocks.insert( 0, Block( spec=UnderpaintBlockType, inputs={"semantic_input": underpaint_audio_arr}, targets={ "semantic_output": Block.shift_left(underpaint_audio_arr, cfg.semantic_pad_token) }, ), ) # build blocks # add semantic pad to text x_text_arr = x_text_arr.unsqueeze(0) text_block = Block( spec=TextBlockType, inputs={ "text_input": x_text_arr, }, ) text_is_post = ( True if sampling_params.text_loss and not sampling_params.inference and random.random() < 0.1 else False ) if text_is_post: blocks = audio_blocks + [text_block] else: blocks = [text_block] + audio_blocks block_sequence = BlockSequence(blocks) # print(block_sequence) sample_info = { "text": text, "n_tokens_text": len(x_text_arr), "text_is_post": text_is_post, "blocks": block_sequence, } return sample_info def get_batch( batch_size_tokens: int, cfg: GPTConfig, sample_generator, sampling_params: SamplingParams, return_idx=False, ): if return_idx: raise NotImplementedError("return_idx not implemented") packed_sequences = PackedBlockSequence([]) if sampling_params.mask_padding: raise NotImplementedError("mask_padding not implemented") if sampling_params.pack: # fill up block_size with samples cur_len = 0 n_loop = 0 for sample in tqdm(sample_generator, desc="Packing batch", disable=True): if cur_len >= batch_size_tokens: break packed_sequences.append(sample["blocks"]) cur_len += sample["blocks"].n_tokens n_loop += 1 # crop so total length is batch_size_tokens packed_sequences = packed_sequences.crop_to_max_tokens(batch_size_tokens) assert packed_sequences.n_tokens == batch_size_tokens extra_factor = 1 if not sampling_params.allow_skip else 4 if ( n_loop > int(round(batch_size_tokens / cfg.block_size)) * 16 * extra_factor ): # should almost never trigger on 8k blocks print(f"warning, {n_loop} loops in dataloading") return packed_sequences class CustomAudioDataset(IterableDataset): def __init__( self, data_sampling_info, split, dataset_idx=None, sampling_params: SamplingParams | None = None, ): self.data_sampling_info = data_sampling_info self.dataset_idx = dataset_idx self.split = split self.sampling_params = sampling_params if split == "val": assert self.sampling_params.inference def sample_generator(self): while True: try: sample_data = get_sample_codes_or_audio( self.data_sampling_info, self.split, self.sampling_params, dataset_idx=self.dataset_idx, ) packed_sequences = get_sample_from_row( self.data_sampling_info["cfg"], self.data_sampling_info["tokenizer_fp"], sample_data, ) except Exception as e: print(f"Error in sample_generator: {e}") continue yield packed_sequences def __iter__(self): return self def __next__(self): # Randomly sample from the dataset # ++ Pass the stored SamplingParams object to get_batch ++ batch = get_batch( self.data_sampling_info["batch_size_tokens"], self.data_sampling_info["cfg"], self.sample_generator(), self.sampling_params, ) return batch from suno_utils.tasks.mert_25 import ( preload_models as preload_semantic_models_, encode as encode_semantic, ) def preload_semantic_models(device): preload_semantic_models_( checkpoint_filepath="s3://suno-data/georg/models/semantic/mert_25.pt", centroids_filepath="s3://suno-data/georg/models/semantic/mert_25_2x4k.npy", device=device, ) class JSONLMemmap: def __init__(self, path, verbose=False): self.path = path self._index = self.get_line_index_map(verbose=verbose) def get_line_index_map(self, verbose=False): line_positions = [0] with open(self.path, "rb") as f: for line in tqdm(f, disable=not verbose, desc="Making memmap index"): line_positions.append(len(line) + line_positions[-1]) return np.array(line_positions[:-1]) def get_line_from_index(self, index: int): with open(self.path, "rb") as f: f.seek(self._index[index]) return json.loads(f.readline()) def __getitem__(self, index: int): return self.get_line_from_index(index) def __len__(self): return len(self._index) def __iter__(self): for i in range(len(self)): yield self[i] class OracleWavAudioDataset(IterableDataset): """Gets SampleData objects from the oracle dataset, which contain raw wav data""" def __init__( self, model_cfg: GPTConfig, train_cfg: GPTTrainConfig, metas_path: str, batch_size_tokens: int, tokenizer_fp: str, device: str, split="train", dataset_idx=None, sampling_params: SamplingParams | None = None, ): self.model_cfg = model_cfg self.train_cfg = train_cfg self.metas_path = metas_path self.batch_size_tokens = batch_size_tokens self.tokenizer_fp = tokenizer_fp self.device = device self.dataset_idx = dataset_idx self.split = split self.sampling_params = sampling_params if split == "val": assert self.sampling_params.inference # load metas as a memmap to save memory self.metas = JSONLMemmap(self.metas_path, verbose=False) # make essential dicts self.weights = [] self.id_to_index = {} self.artist_to_ids = defaultdict(list) self.playlist_to_ids = defaultdict(list) with open(self.metas_path, "r") as f: for i, line in tqdm( enumerate(f), disable=True, total=len(self.metas), desc="Loading metas", mininterval=10, ): meta = orjson.loads(line) self.weights.append(meta.get("weight", 0)) self.id_to_index[meta["id"]] = i for artist_id in meta.get("artist_ids", []): self.artist_to_ids[artist_id].append(meta["id"]) for playlist_id in meta.get("playlist_ids", []): self.playlist_to_ids[playlist_id].append(meta["id"]) self.tokenizer = _load_tokenizer(self.tokenizer_fp) self.random_cache = [] def __iter__(self): return self def _load_audio_for_mert(self, local_filepath, s3_filepath, start_s=0, max_duration_s=60 * 30): model_max_duration_s = self.model_cfg.block_size / self.model_cfg.semantic_rate_hz max_duration_s = min(max_duration_s, model_max_duration_s) mock = False try: audio = get_sample_oracle_file_segment( local_filepath=local_filepath, s3_filepath=s3_filepath, start_s=start_s, max_duration_s=max_duration_s, mock=mock, ) audio = audio.convert(24000, 2, 1) # MERT sample rate except Exception as e: print(f"Error loading audio for {s3_filepath}: {e}") print(f"start_s: {start_s}, max_duration_s: {max_duration_s}") print(f"audio: {audio.array_float.shape}") raise e return audio.array_float def _load_cover_audio(self, main_meta): covers = main_meta.get("cover_ids", []) if len(covers) > 0: cover_id = random.choice(covers) cover_meta = self.metas[self.id_to_index[cover_id]] data_row_cover = self._load_audio_for_mert( cover_meta["local_filepath"], cover_meta["s3_filepath"] ) else: data_row_cover = None return data_row_cover def _load_artist_audio(self, main_meta): artists = main_meta.get("artist_ids", []) if len(artists) == 0: return None artist_id = random.choice(artists) artist_song_ids = self.artist_to_ids[artist_id] # remove main_meta id from artist_song_ids artist_song_ids = [id for id in artist_song_ids if id != main_meta["id"]] if len(artist_song_ids) == 0: return None # sample multiple segments n_segments = random.randint(1, 5) min_segment_len = 5 max_segment_len = 120 artist_audio_arrs = [] for _ in range(n_segments): track_id = random.choice(artist_song_ids) track_meta = self.metas[self.id_to_index[track_id]] dur_s = random.random() * max_segment_len dur_s = min(max(min_segment_len, dur_s), track_meta["duration_s"]) start_s = random.random() * (track_meta["duration_s"] - dur_s) assert start_s >= 0 data_row_artist = self._load_audio_for_mert( track_meta["local_filepath"], track_meta["s3_filepath"], start_s=start_s, max_duration_s=dur_s, ) artist_audio_arrs.append(data_row_artist) return artist_audio_arrs def _load_playlist_audio(self, main_meta): playlists = main_meta.get("playlist_ids", []) if len(playlists) == 0: return None playlist_id = random.choice(playlists) playlist_song_ids = self.playlist_to_ids[playlist_id] # remove main_meta id from playlist_song_ids playlist_song_ids = [id for id in playlist_song_ids if id != main_meta["id"]] if len(playlist_song_ids) == 0: return None playlist_audio_arrs = [] n_segments = random.randint(1, 5) min_segment_len = 5 max_segment_len = 60 for _ in range(n_segments): track_id = random.choice(playlist_song_ids) track_meta = self.metas[self.id_to_index[track_id]] dur_s = random.random() * max_segment_len dur_s = min(max(min_segment_len, dur_s), track_meta["duration_s"]) start_s = random.random() * (track_meta["duration_s"] - dur_s) assert start_s >= 0 data_row_playlist = self._load_audio_for_mert( track_meta["local_filepath"], track_meta["s3_filepath"], start_s=start_s, max_duration_s=dur_s, ) playlist_audio_arrs.append(data_row_playlist) return playlist_audio_arrs def _load_overpaint_audio(self, main_meta): overpaint_id = main_meta.get("overpaint_id", None) if overpaint_id is None: return None overpaint_meta = self.metas[self.id_to_index[overpaint_id]] data_row_overpaint = self._load_audio_for_mert( overpaint_meta["local_filepath"], overpaint_meta["s3_filepath"] ) return data_row_overpaint def _load_underpaint_audio(self, main_meta): underpaint_id = main_meta.get("underpaint_id", None) if underpaint_id is None: return None underpaint_meta = self.metas[self.id_to_index[underpaint_id]] data_row_underpaint = self._load_audio_for_mert( underpaint_meta["local_filepath"], underpaint_meta["s3_filepath"] ) return data_row_underpaint def __next__(self): try: return self._next() except Exception as e: print(f"Error in __next__: {e}") # raise e return self.__next__() def _sample_meta(self): """ Randomly sample a meta from the metas list. Weighted sampling is slow for large weighted datasets so we sample 10k at a time. """ if len(self.random_cache) == 0: choices = random.choices(range(len(self.metas)), weights=self.weights, k=10000) self.random_cache.extend(choices) idx = self.random_cache.pop() return self.metas[idx] def _next(self): # randomly sample a meta main_meta = self._sample_meta() # load wav data_row = self._load_audio_for_mert(main_meta["local_filepath"], main_meta["s3_filepath"]) # TODO: add cover, artist, playlist, overpaint, underpaint if self.sampling_params.allow_cover and random.random() < self.sampling_params.prob_cover: data_row_cover = self._load_cover_audio(main_meta) else: data_row_cover = None p_use_artist = ( 0.5 if data_row_cover is not None else self.sampling_params.prob_artist ) # increase chance for cover cause voice beautifier if self.sampling_params.allow_artist and random.random() < p_use_artist: data_row_artist = self._load_artist_audio(main_meta) else: data_row_artist = None if self.sampling_params.allow_playlist and random.random() < self.sampling_params.prob_playlist: data_row_playlist = self._load_playlist_audio(main_meta) else: data_row_playlist = None if ( self.sampling_params.allow_overpaint and random.random() < self.sampling_params.prob_overpaint ): data_row_overpaint = self._load_overpaint_audio(main_meta) else: data_row_overpaint = None if ( self.sampling_params.allow_underpaint and random.random() < self.sampling_params.prob_underpaint ): data_row_underpaint = self._load_underpaint_audio(main_meta) else: data_row_underpaint = None sample_data = SampleData( data_row=data_row, data_meta=main_meta, sampling_params=self.sampling_params, is_full_track=True, data_row_cover=data_row_cover, artist_tracks=data_row_artist, playlist_tracks=data_row_playlist, overpaint_track=data_row_overpaint, underpaint_track=data_row_underpaint, ) return sample_data class RawAudioDataset(IterableDataset): def __init__( self, sample_data_dl, batch_size_tokens: int, cfg: GPTConfig, sampling_params: SamplingParams | None = None, device: str | None = None, tokenizer_fp: str | None = None, ): self.sample_data_dl = sample_data_dl self.batch_size_tokens = batch_size_tokens self.model_cfg = cfg self.sampling_params = sampling_params self.device = device self.tokenizer_fp = tokenizer_fp def sample_generator_fn(self, sample_data_dl): while True: sample_data = next(sample_data_dl) def encode_semantic_track(track, fuzz_edges=True): assert isinstance(track, np.ndarray) track = torch.from_numpy(track).unsqueeze(0).contiguous() codes = encode_semantic( track, pad_to_chunksize=True, device=self.device, batch_size=48 ).T[[0]] # (1, T) if fuzz_edges and random.random() < 0.2: # crop up to 5s from both sides so we aren't always on the chunk boundary l_cut = random.randint(0, 5 * self.model_cfg.semantic_rate_hz) r_cut = random.randint(1, 5 * self.model_cfg.semantic_rate_hz) # avoid -0 edge case fuzzed_codes = codes[:, l_cut:-r_cut] if fuzzed_codes.shape[1] > 0: # if we crop too much, dont fuzz codes = fuzzed_codes if codes.shape[1] > self.model_cfg.block_size - self.model_cfg.t_text: # TODO: why does this happen? codes = codes[:, : self.model_cfg.block_size - self.model_cfg.t_text] return codes # encode all data if sample_data.data_row is not None: sample_data.data_row = encode_semantic_track(sample_data.data_row) # (1, T) if sample_data.data_row_cover is not None: sample_data.data_row_cover = encode_semantic_track(sample_data.data_row_cover) # (1, T) if sample_data.artist_tracks is not None: sample_data.artist_tracks = [ encode_semantic_track(track) for track in sample_data.artist_tracks ] # (N, T) if sample_data.playlist_tracks is not None: sample_data.playlist_tracks = [ encode_semantic_track(track) for track in sample_data.playlist_tracks ] # (N, T) if sample_data.overpaint_track is not None: sample_data.overpaint_track = encode_semantic_track( sample_data.overpaint_track ) # (1, T) if sample_data.underpaint_track is not None: sample_data.underpaint_track = encode_semantic_track( sample_data.underpaint_track ) # (1, T) sample = get_sample_from_row(self.model_cfg, self.tokenizer_fp, sample_data) yield sample def __iter__(self): self.sample_generator = self.sample_generator_fn(self.sample_data_dl) return self def __next__(self): batch = get_batch( self.batch_size_tokens, self.model_cfg, self.sample_generator, self.sampling_params, ) return batch