import os import random import re import numpy as np from tokenizers import AddedToken import torch import torch.nn.functional as F from transformers import PreTrainedTokenizerFast # to avoid: "The current process just got forked, after parallelism has already been used" os.environ["TOKENIZERS_PARALLELISM"] = "False" global tokenizer g_tokenizer = None def _load_tokenizer(tokenizer_fp=None): global g_tokenizer if g_tokenizer is not None: return g_tokenizer # tokenizer = BertTokenizerFast.from_pretrained( # "bert-base-multilingual-cased", # model_max_length=512*4, # ) 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 _clean_tag(tag): return re.sub(r"\s+", " ", tag).strip() def _augment_tag(s): if random.random() >= 0.95: s = s.upper() elif random.random() >= 0.95: s = s.capitalize() elif random.random() >= 0.9: s = s.title() elif random.random() >= 0.9: s = s.lower() if random.random() >= 0.5: s = s.replace("-", " ").strip() return s # Structure: # {start;vocals:start} # is song start & vocals start within 8s of actual start ## tags go here ## lyrics go here # {start;vocals:end} # is song end & vocals end within 8s of actual end def _get_start_control_tags(data_meta): start_s = data_meta.get("start_s") vocal_start_s = data_meta.get("vocal_start_s") control_tags = [] if start_s is not None and start_s <= 0.5: control_tags.append("start") if vocal_start_s is not None and vocal_start_s <= 8: control_tags.append("vocals:start") if len(control_tags) == 0: return None return "{" + ";".join(control_tags) + "}" def _get_end_control_tags(data_meta): end_s = data_meta.get("end_s") vocal_end_s = data_meta.get("vocal_end_s") original_duration_s = data_meta.get("original_duration_s") control_tags = [] if end_s is not None and original_duration_s - end_s <= 0.5: control_tags.append("end") if vocal_end_s is not None and end_s - vocal_end_s <= 10: control_tags.append("vocals:end") if len(control_tags) == 0: return None return "{" + ";".join(control_tags) + "}" def get_computed_tags(data_meta): computed_tags = [] cutoff_freq = data_meta.get("cutoff_freq") if cutoff_freq is None: return computed_tags if cutoff_freq <= 16_000: computed_tags.append("low rolloff") if cutoff_freq >= 18_000: computed_tags.append("high rolloff") return computed_tags def shift_codebooks(cfg, data_row, array_width=None): if array_width is None: array_width = ( data_row.shape[-1] + cfg.semantic_n_codebooks * cfg.semantic_shift_factor + (cfg.coarse_n_codebooks - 1) * cfg.coarse_shift_factor ) # build semantic y_semantic_arr = np.full( (cfg.semantic_n_codebooks, array_width), 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 y_coarse_arr = np.full((cfg.coarse_n_codebooks, array_width), 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 y_audio_arr = np.concatenate([y_semantic_arr, y_coarse_arr], axis=0) return y_audio_arr def get_sample( data_sampling_info, split, dataset_idx=None, rel_row_idx=None, use_private=False, inference=False, suppress_text=False, dummy_data=False, return_idx=False, ): if dummy_data: cfg = data_sampling_info["cfg"] x_audio_arr = np.zeros( ( cfg.semantic_n_codebooks + cfg.coarse_n_codebooks, cfg.block_size - cfg.t_text, ), dtype=np.int64, ) y_audio_arr = np.zeros( ( cfg.semantic_n_codebooks + cfg.coarse_n_codebooks, cfg.block_size - cfg.t_text, ), dtype=np.int64, ) return "", x_audio_arr, y_audio_arr data = data_sampling_info[split]["data"] metas = data_sampling_info[split]["metas"] if rel_row_idx is None: if dataset_idx is None: weights = data_sampling_info[split]["weights"] dataset_idx = random.choices(list(range(len(weights))), weights=weights, k=1)[0] idx_lists = data_sampling_info[split]["idx_lists"] rel_row_idx = random.choice(list(range(len(idx_lists[dataset_idx])))) rel_row_idx = rel_row_idx % len(idx_lists[dataset_idx]) row_idx = idx_lists[dataset_idx][rel_row_idx] # names = data_sampling_info[split]["names"] # dataset_name = names[dataset_idx] data_row = data[row_idx].astype(np.int64) data_meta = metas[row_idx] cfg = data_sampling_info["cfg"] # TODO: change the transpose here data_row = data_row.T # split into chunks here n_tokens = np.where(data_row[0] != cfg.semantic_pad_token)[0][-1] + 1 # TODO: this is a hack to not have to increase blocksize n_tokens = min(n_tokens, 3008 - 250) # abc -> acb # - predict only b # - a can be 0 # take duration_s for calculation # 25% chance a is 0 (pre-painting) # max = min(duration_s-0.1, 59.9) # then pick a 0.01-max # then pick c 0.01-max # b is duration_s - a - c assert n_tokens >= 3 c_len = random.randrange(1, n_tokens - 1) if random.random() <= -1: a_len = 0 else: a_len = random.randrange(1, n_tokens - c_len) b_len = n_tokens - c_len - a_len assert a_len >= 0 and b_len > 0 and c_len > 0 a_arr = data_row[:, :a_len] b_arr = data_row[:, a_len : a_len + b_len] c_arr = data_row[:, a_len + b_len : a_len + b_len + c_len] a_y_audio_arr = shift_codebooks(cfg, a_arr) b_y_audio_arr = shift_codebooks(cfg, b_arr) c_y_audio_arr = shift_codebooks(cfg, c_arr) a_len_post = a_y_audio_arr.shape[-1] b_len_post = b_y_audio_arr.shape[-1] c_len_post = c_y_audio_arr.shape[-1] assert a_len_post + c_len_post + b_len_post <= cfg.t_audio y_audio_arr = np.full( (cfg.semantic_n_codebooks + cfg.coarse_n_codebooks, cfg.t_audio), cfg.semantic_pad_token, dtype=np.int64, ) y_audio_arr[cfg.semantic_n_codebooks :, :] = cfg.coarse_pad_token x_audio_arr = np.full( (cfg.semantic_n_codebooks + cfg.coarse_n_codebooks, cfg.t_audio), cfg.semantic_pad_token, dtype=np.int64, ) x_audio_arr[cfg.semantic_n_codebooks :, :] = cfg.coarse_pad_token # assemble y y_audio_arr[:, :a_len_post] = a_y_audio_arr y_audio_arr[:, a_len_post + 1 : a_len_post + c_len_post + 1] = c_y_audio_arr y_audio_arr[:, a_len_post + c_len_post + 2 : a_len_post + c_len_post + b_len_post + 2] = ( b_y_audio_arr ) # make x and add infer x_audio_arr[: cfg.semantic_n_codebooks, 0:1] = cfg.semantic_infer_token x_audio_arr[cfg.semantic_n_codebooks :, 0:1] = cfg.coarse_infer_token x_audio_arr[:, 1 : 1 + a_len_post] = a_y_audio_arr x_audio_arr[: cfg.semantic_n_codebooks, 1 + a_len_post : 2 + a_len_post] = cfg.semantic_infer_token x_audio_arr[cfg.semantic_n_codebooks :, 1 + a_len_post : 2 + a_len_post] = cfg.coarse_infer_token x_audio_arr[:, 2 + a_len_post : 2 + a_len_post + c_len_post] = c_y_audio_arr x_audio_arr[ : cfg.semantic_n_codebooks, 2 + a_len_post + c_len_post : 3 + a_len_post + c_len_post ] = cfg.semantic_infer_token x_audio_arr[ cfg.semantic_n_codebooks :, 2 + a_len_post + c_len_post : 3 + a_len_post + c_len_post ] = cfg.coarse_infer_token x_audio_arr[:, 3 + a_len_post + c_len_post : 3 + a_len_post + c_len_post + b_len_post] = ( b_y_audio_arr ) # use text as indicator variable for abc # how to (optionally) encode duration # - fill in 100 tokens # - fill in arbitrary amount x_text_indicator = np.full((1, cfg.t_audio), cfg.text_pad_token, dtype=np.int64) x_text_indicator[:, : 1 + a_len_post] = cfg.text_pad_token + 1 x_text_indicator[:, 1 + a_len_post : 2 + a_len_post + c_len_post] = cfg.text_pad_token + 3 x_text_indicator[:, 2 + a_len_post + c_len_post : 3 + a_len_post + c_len_post + b_len_post] = ( cfg.text_pad_token + 2 ) assert x_audio_arr.shape[-1] == cfg.block_size - cfg.t_text assert y_audio_arr.shape[-1] == cfg.block_size - cfg.t_text # build text text = "" # collect tags if use_private: tags = data_meta.get("tags_private", data_meta.get("tags", [])) else: tags = data_meta.get("tags", []) # add computed tags computed_tags = get_computed_tags(data_meta) if len(computed_tags) > 0 and random.random() >= 0.1: tags.extend(computed_tags) # for tags remove newlines, empty tags, and case augment tags = [clean_tag for tag in tags if len(clean_tag := _clean_tag(tag)) > 0] if len(tags) > 0 and (inference or random.random() >= 0.25): if inference: text += f"[{', '.join(tags)[:128]}]\n\n" else: random.shuffle(tags) tags = tags[: random.randint(1, len(tags))] tags = [_augment_tag(tag) for tag in tags] tag_str = random.choice([", ", " ", "; "]).join(tags) text += f"[{tag_str[:128]}]" # pretty arbitrary max len for now text += random.choice([" ", "\n", "\n\n"]) # for tts_text sometimes remove metas, newlines and lower-case augment if use_private: tts_text = data_meta.get("text_private", data_meta.get("text", "")) else: tts_text = data_meta.get("text", "") if len(tts_text) > 0 and (inference or random.random() >= 0.1): if not inference: if random.random() >= 0.9: tts_text = tts_text.lower() if random.random() >= 0.9: tts_text = re.sub(r"\n+", " ", tts_text) text += tts_text.strip() # get control tags text = text.replace("{", "").replace("}", "") if not inference and random.random() >= 0.1: control_tags_start = _get_start_control_tags(data_meta) control_tags_end = _get_end_control_tags(data_meta) if control_tags_start is not None: text = control_tags_start + random.choice([" ", "\n", "\n\n"]) + text if control_tags_end is not None: text = text + random.choice([" ", "\n", "\n\n"]) + control_tags_end text = text.strip() if suppress_text: text = "" if return_idx: return row_idx, text, x_text_indicator, x_audio_arr, y_audio_arr return text, x_text_indicator, x_audio_arr, y_audio_arr def get_batch( data_sampling_info, split, dataset_idx=None, row_idx=None, use_private=False, inference=False, min_text_offs=None, suppress_text=False, dummy_data=False, return_idx=False, n_offs=None, ): batch_size = data_sampling_info["batch_size"] device = data_sampling_info["device"] device_type = data_sampling_info["device_type"] tokenizer_fp = data_sampling_info.get("tokenizer_fp") cfg = data_sampling_info["cfg"] if not isinstance(dataset_idx, list): dataset_idx = [dataset_idx] * batch_size if not isinstance(row_idx, list): row_idx = [row_idx] * batch_size if n_offs is not None: row_idx = list(range(n_offs * batch_size, (n_offs + 1) * batch_size)) x_text_list = [] x_text_indicator_list = [] x_audio_list = [] y_list = [] idx_list = [] for n in range(batch_size): out = get_sample( data_sampling_info, split, dataset_idx=dataset_idx[n], rel_row_idx=row_idx[n], use_private=use_private, inference=inference, suppress_text=suppress_text, dummy_data=dummy_data, return_idx=return_idx, ) if return_idx: idx, x_text, x_text_indicator, x_audio, y = out idx_list.append(idx) else: x_text, x_text_indicator, x_audio, y = out x_text_list.append(x_text) x_text_indicator_list.append(x_text_indicator) x_audio_list.append(torch.from_numpy(x_audio)) y_list.append(torch.from_numpy(y)) x_text = tokenize_batch( x_text_list, max_tokens=cfg.t_text, pad_token_id=cfg.text_pad_token, tokenizer_fp=tokenizer_fp, ) if min_text_offs is not None and min_text_offs > x_text.shape[-1]: x_text = F.pad( x_text, (0, min_text_offs - x_text.shape[-1]), "constant", cfg.text_pad_token, ) x_audio = torch.stack(x_audio_list) y = torch.stack(y_list) # combine all x and pad as much as needed x = torch.concatenate( [ F.pad( x_audio[:, : cfg.semantic_n_codebooks], ( x_text.shape[-1], cfg.block_size - x_text.shape[-1] - x_audio.shape[-1], ), "constant", cfg.semantic_pad_token, ), F.pad( x_audio[:, cfg.semantic_n_codebooks :], ( x_text.shape[-1], cfg.block_size - x_text.shape[-1] - x_audio.shape[-1], ), "constant", cfg.coarse_pad_token, ), ], dim=1, ) x = torch.concatenate( [ F.pad( x_text[:, None], (0, cfg.block_size - x_text.shape[-1]), "constant", cfg.text_pad_token, ), x, ], dim=1, ) text_offset = x_text.shape[-1] # add text indicator info for n, x_text_indicator in enumerate(x_text_indicator_list): x[n, :1, text_offset : text_offset + x_text_indicator.shape[-1]] = torch.from_numpy( x_text_indicator ) assert x.shape == ( batch_size, 1 + cfg.semantic_n_codebooks + cfg.coarse_n_codebooks, cfg.block_size, ) if device_type == "cuda": # pin arrays x,y, which allows us to move them to GPU asynchronously (non_blocking=True) x, y = ( x.pin_memory().to(device, non_blocking=True), y.pin_memory().to(device, non_blocking=True), ) else: x, y = x.to(device), y.to(device) del x_text_list, x_audio_list, y_list, x_text, x_audio if return_idx: return idx_list, text_offset, x, y return text_offset, x, y