# 30b data prep import numpy as np import tqdm import funcy import gc import json import os import random from collections import defaultdict from joblib import Parallel, delayed from suno_utils.utils.text import write_jsonl, read_jsonl, write_json from suno_utils.utils.s3 import read_from_s3, check_s3_file_exists TEXT_CODEBOOK_SIZE = 60_001 TEXT_PAD_TOKEN = TEXT_CODEBOOK_SIZE TEXT_VOCAB_SIZE = 60_032 SEMANTIC_CODEBOOK_SIZE = 4000 SEMANTIC_N_CODEBOOKS = 1 SEMANTIC_PAD_TOKEN = SEMANTIC_CODEBOOK_SIZE SEMANTIC_INFER_TOKEN = SEMANTIC_CODEBOOK_SIZE + 1 SEMANTIC_VOCAB_SIZE = 4032 SEMANTIC_RATE_HZ = 25 SEMANTIC_SHIFT_FACTOR = 50 assert SEMANTIC_VOCAB_SIZE == (np.floor(SEMANTIC_CODEBOOK_SIZE // 64) + 1) * 64 COARSE_CODEBOOK_SIZE = 2048 COARSE_N_CODEBOOKS = 12 COARSE_PAD_TOKEN = COARSE_CODEBOOK_SIZE COARSE_INFER_TOKEN = COARSE_CODEBOOK_SIZE + 1 COARSE_VOCAB_SIZE = 2112 COARSE_RATE_HZ = 25 COARSE_SHIFT_FACTOR = 5 assert COARSE_CODEBOOK_SIZE + 3 < COARSE_VOCAB_SIZE assert COARSE_VOCAB_SIZE % 64 == 0 assert SEMANTIC_RATE_HZ == COARSE_RATE_HZ BLOCK_SIZE = 8704 N_TOKENS_TEXT = 2560 N_TOKENS_MEMMAP = 6016 # max 240s of audio # make sure we have enough space for shift 10 assert BLOCK_SIZE >= ( N_TOKENS_TEXT + N_TOKENS_MEMMAP + SEMANTIC_N_CODEBOOKS * SEMANTIC_SHIFT_FACTOR + (COARSE_N_CODEBOOKS - 1) * COARSE_SHIFT_FACTOR ) SEMANTIC_EMBED_DIR = "mert_25_2x4k" CODEC_EMBED_DIR = "dac_2c_25_12" def _trim_to_common(arr_1, arr_2): common_len = min(len(arr_1), len(arr_2)) arr_1 = arr_1[:common_len] arr_2 = arr_2[:common_len] return arr_1, arr_2 def _parse_arrays(task_name, dset_name, meta_info, semantic_arr, coarse_arr): if task_name is None: task_name = "default" # prep segment metas (use semantic for timekeeping) if "lyrics" in meta_info and "text" not in meta_info: # hotfix incase it's called differently meta_info["text"] = meta_info["lyrics"] segments_info = [] # each segment infomration is: # start_idx, end_idx, text, vocal_start_idx, vocal_end_idx, line_start_times # first check if we have known segments if "text_segments" in meta_info: for m in meta_info["text_segments"]: # we keep the relative start time selected_line_start_s = [] for line_start_s in m.get("line_start_s", []): if line_start_s is not None: selected_line_start_s.append(round(line_start_s - m["start_s"], 2)) segments_info.append( ( int(round(m["start_s"] * SEMANTIC_RATE_HZ)), min(len(semantic_arr), int(round(m["end_s"] * SEMANTIC_RATE_HZ))), m["text"], int(round(m["vocal_start_s"] * SEMANTIC_RATE_HZ)) if m["vocal_start_s"] is not None else None, min( len(semantic_arr), int(round(m["vocal_end_s"] * SEMANTIC_RATE_HZ)), ) if m["vocal_end_s"] is not None else None, selected_line_start_s, ) ) # if "text" in segmented text then add twice if "text" in meta_info: segments_info.append( ( 0, min(len(semantic_arr), N_TOKENS_MEMMAP), meta_info["text"], # add text only to first piece None, None, None, ) ) else: # randomize offset to not get only multiples if no text available offs = 0 if task_name == "default" and "text" not in meta_info and random.random() > 0.5: # randomize if we don't have lyrics and not special task offs = random.randint(10 * SEMANTIC_RATE_HZ, N_TOKENS_MEMMAP - 1) segments_info.append((0, min(len(semantic_arr), offs), None, None, None, None)) total_steps = int(np.ceil((len(semantic_arr) - offs) / N_TOKENS_MEMMAP)) for n in range(total_steps): start_idx = offs + n * N_TOKENS_MEMMAP end_idx = min(len(semantic_arr), offs + (n + 1) * N_TOKENS_MEMMAP) if end_idx - start_idx < SEMANTIC_RATE_HZ: # might as well skip mini ones continue segments_info.append( ( start_idx, end_idx, meta_info.get("text") if n == 0 else None, # add text only to first piece None, None, None, ) ) arr_list = [] for ( sem_start_idx, sem_end_idx, text, vocal_start_idx, vocal_end_idx, line_start_times, ) in segments_info: if ( sem_end_idx - sem_start_idx > N_TOKENS_MEMMAP or sem_end_idx - sem_start_idx < SEMANTIC_RATE_HZ # arbitrary ): continue coarse_start_idx = int(round(sem_start_idx * COARSE_RATE_HZ / SEMANTIC_RATE_HZ)) coarse_end_idx = int(round(sem_end_idx * COARSE_RATE_HZ / SEMANTIC_RATE_HZ)) assert sem_end_idx >= 0 and coarse_start_idx >= 0 if sem_end_idx > len(semantic_arr) or coarse_end_idx > len(coarse_arr): continue # get array segments arr_s = semantic_arr[sem_start_idx:sem_end_idx, :SEMANTIC_N_CODEBOOKS].copy() arr_c = coarse_arr[coarse_start_idx:coarse_end_idx, :COARSE_N_CODEBOOKS].copy() # fix any alignment mistakes arr_s, arr_c = _trim_to_common(arr_s, arr_c) assert len(arr_s) == len(arr_c) # concat and stack n_tokens = len(arr_c) if len(arr_c) < N_TOKENS_MEMMAP: arr_c = np.pad( arr_c, ((0, N_TOKENS_MEMMAP - len(arr_c)), (0, 0)), constant_values=COARSE_PAD_TOKEN, mode="constant", ) arr_s = np.pad( arr_s, ((0, N_TOKENS_MEMMAP - len(arr_s)), (0, 0)), constant_values=SEMANTIC_PAD_TOKEN, mode="constant", ) arr = np.concatenate([arr_s, arr_c], axis=-1) arr = arr.astype(np.uint16) assert arr.shape == (N_TOKENS_MEMMAP, SEMANTIC_N_CODEBOOKS + COARSE_N_CODEBOOKS) new_meta = { "id": meta_info["id"], "start_s": round(sem_start_idx / SEMANTIC_RATE_HZ, 2), "end_s": round(sem_end_idx / SEMANTIC_RATE_HZ, 2), "original_duration_s": round(len(semantic_arr) / SEMANTIC_RATE_HZ, 2), "vocal_start_s": ( round(vocal_start_idx / SEMANTIC_RATE_HZ, 2) if vocal_start_idx is not None else None ), "vocal_end_s": ( round(vocal_end_idx / SEMANTIC_RATE_HZ, 2) if vocal_end_idx is not None else None ), "line_start_s": line_start_times if line_start_times else None, "n_tokens": n_tokens, } if text is not None: new_meta["text"] = text new_meta["text_lang"] = meta_info.get("lang") if task_name == "default": new_meta["dset_suffix"] = "lyrics" if meta_info["lang"] == "en" else "lyrics_foreign" if "tags" in meta_info: new_meta["tags"] = meta_info["tags"] if "original_id" in meta_info: new_meta["original_id"] = meta_info["original_id"] if "parent_id" in meta_info: new_meta["parent_id"] = meta_info["parent_id"] if "artist" in meta_info: new_meta["artist"] = meta_info["artist"] arr_list.append((arr, new_meta)) del arr_s, arr_c return arr_list def _process_archives( task_name, dset_name, s3_semantic_archive_filepaths, s3_coarse_archive_filepaths, relevant_metas, ): semantic_archive = {} s3_semantic_archive_filepaths = set(s3_semantic_archive_filepaths) # print(len(s3_semantic_archive_filepaths)) for s3_semantic_archive_filepath in s3_semantic_archive_filepaths: if not check_s3_file_exists(s3_semantic_archive_filepath): print(f"missing {s3_semantic_archive_filepath}") continue try: archive = { k: v for k, v in read_from_s3(s3_semantic_archive_filepath, read_f=np.load).items() } except: # corrupt archive print(f"corrupt {s3_semantic_archive_filepath}") continue for k, v in archive.items(): semantic_archive[k] = v coarse_archive = {} s3_coarse_archive_filepaths = set(s3_coarse_archive_filepaths) # print(len(s3_coarse_archive_filepaths)) for s3_coarse_archive_filepath in s3_coarse_archive_filepaths: if not check_s3_file_exists(s3_coarse_archive_filepath): print(f"missing {s3_coarse_archive_filepath}") continue try: archive = {k: v for k, v in read_from_s3(s3_coarse_archive_filepath, read_f=np.load).items()} except: # corrupt archive print(f"corrupt {s3_coarse_archive_filepath}") continue for k, v in archive.items(): coarse_archive[k] = v semantic_uids, coarse_uids = set(semantic_archive.keys()), set(coarse_archive.keys()) # removing this for now just incase (eg imslp) # assert (len(semantic_uids) < 10 and len(coarse_uids) < 10) or ( # len(semantic_uids & coarse_uids) / (len(semantic_uids) + len(coarse_uids)) > 0.1 # ) # assert len(semantic_uids & coarse_uids) > 0 arr_list = [] for uid in semantic_uids & coarse_uids: if uid not in relevant_metas: continue semantic_arr = semantic_archive[uid] coarse_arr = coarse_archive[uid] if np.abs(len(coarse_arr) / COARSE_RATE_HZ - len(semantic_arr) / SEMANTIC_RATE_HZ) > 0.1: # skip if embeddings not roughly the same duration continue semantic_arr, coarse_arr = _trim_to_common(semantic_arr, coarse_arr) assert len(coarse_arr) == len(semantic_arr) * COARSE_RATE_HZ / SEMANTIC_RATE_HZ arr_list.extend( _parse_arrays(task_name, dset_name, relevant_metas[uid], semantic_arr, coarse_arr) ) del semantic_archive, coarse_archive gc.collect() return arr_list def _collect_uids( s3_semantic_metas_filepaths, s3_coarse_metas_filepaths, ): semantic_uids = [] for fp in s3_semantic_metas_filepaths: try: metas = read_from_s3(fp, read_f=read_jsonl) except: print(f"failed on metas for fp: {fp}") continue semantic_uids.extend([m["id"] for m in metas]) coarse_uids = [] for fp in s3_coarse_metas_filepaths: try: metas = read_from_s3(fp, read_f=read_jsonl) except: print(f"failed on metas for fp: {fp}") continue coarse_uids.extend([m["id"] for m in metas]) return set(semantic_uids) & set(coarse_uids) def _prep_data( dataset, out_data_dir, meta_info_map, njobs=5, chunksize=10, is_val=False, n_offs=0, ): dset_name, dset_version, (start_idx, end_idx), n_sem, n_coarse, task_name = dataset dset_type = "val" if is_val else "tr" out_mm_filepath = os.path.join(out_data_dir, f"data_{dset_type}.bin") out_metas_filepath = os.path.join(out_data_dir, f"metas_{dset_type}.jsonl") tot_duration_dict = defaultdict(int) n_chunks = int(np.ceil((end_idx - start_idx) / chunksize)) for idx_chunk in tqdm.tqdm(funcy.chunks(chunksize, list(range(start_idx, end_idx))), total=n_chunks): n_jobs = np.min([njobs, chunksize, len(idx_chunk)]) # collect relevant parts of meta file to avoid copying all to subprocesses tmp_uid_chunks = Parallel(n_jobs=n_jobs, prefer="threads")( delayed(_collect_uids)( [ f"s3://suno-data/datasets/bundles/{dset_version}/{dset_name}/{SEMANTIC_EMBED_DIR}/" + f"metas/part_{idx_idx}.jsonl" for idx_idx in range(idx * n_sem, (idx + 1) * n_sem) ], [ f"s3://suno-data/datasets/bundles/{dset_version}/{dset_name}/{CODEC_EMBED_DIR}/" + f"metas/part_{idx_idx}.jsonl" for idx_idx in range(idx * n_coarse, (idx + 1) * n_coarse) ], ) for idx in idx_chunk ) ## PART A: takes ~40% of loop time uids_per_part = {idx: tmp_uid_chunks[n] for n, idx in enumerate(idx_chunk)} # print(len(uids_per_part)) # collect data encoded_arrays_list = Parallel(n_jobs=n_jobs, prefer="processes")( delayed(_process_archives)( task_name, dset_name, [ f"s3://suno-data/datasets/bundles/{dset_version}/{dset_name}/{SEMANTIC_EMBED_DIR}/" + f"part_{idx_idx}.npz" for idx_idx in range(idx * n_sem, (idx + 1) * n_sem) ], [ f"s3://suno-data/datasets/bundles/{dset_version}/{dset_name}/{CODEC_EMBED_DIR}/" + f"part_{idx_idx}.npz" for idx_idx in range(idx * n_coarse, (idx + 1) * n_coarse) ], { uid: meta_info_map[dset_name][uid] for uid in uids_per_part[idx] if uid in meta_info_map[dset_name] }, ) for idx in idx_chunk ) ## end Part A ## PART B: takes ~40% of loop time add_metas = [] for encoded_arrays in encoded_arrays_list: to_write_len = np.sum([arr.size for arr, _ in encoded_arrays]) if to_write_len == 0: continue out_mm = np.memmap( out_mm_filepath, dtype=np.uint16, mode="r+", shape=(n_offs + to_write_len,), ) for arr, arr_meta in encoded_arrays: out_mm[n_offs : n_offs + arr.size] = arr.reshape( -1, ) n_offs += arr.size dataset_str = dset_name if "dset_suffix" in arr_meta: dataset_str += f"_{arr_meta['dset_suffix']}" add_meta = { "dataset": dataset_str, "task": task_name if task_name is not None else "default", "id": arr_meta["id"], "start_s": round(arr_meta["start_s"], 2), "end_s": round(arr_meta["end_s"], 2), "original_duration_s": arr_meta["original_duration_s"], "vocal_start_s": round(arr_meta["vocal_start_s"], 2) if arr_meta["vocal_start_s"] is not None else None, "vocal_end_s": round(arr_meta["vocal_end_s"], 2) if arr_meta["vocal_end_s"] is not None else None, "line_start_s": arr_meta["line_start_s"] if arr_meta.get("line_start_s") else None, } if "original_id" in arr_meta: add_meta["original_id"] = arr_meta["original_id"] if "parent_id" in arr_meta: add_meta["parent_id"] = arr_meta["parent_id"] if "artist" in arr_meta: add_meta["artist"] = arr_meta["artist"] if "tags" in arr_meta: add_meta["tags"] = arr_meta["tags"] if "text" in arr_meta: add_meta["text"] = arr_meta["text"] if "text_lang" in arr_meta: add_meta["text_lang"] = arr_meta["text_lang"] if "n_tokens" in arr_meta: add_meta["n_tokens"] = arr_meta["n_tokens"] tot_duration_dict[dataset_str] += arr_meta["end_s"] - arr_meta["start_s"] add_metas.append(add_meta) # write it once out_mm.flush() del out_mm ## end Part B write_jsonl( add_metas, os.path.join(out_metas_filepath), do_append=bool(n_offs != 0), ) del encoded_arrays_list # TODO: this gc collect takes super long (>50% of loop if on) let's try to disable # gc.collect() for k, v in tot_duration_dict.items(): print(f"{round(v / 60 / 60):,} hours of {k}") return n_offs def prep_data( datasets, out_data_dir, meta_info_map, is_val=False, njobs=5, chunksize=10, ): n_offs = 0 dset_type = "val" if is_val else "tr" out_mm_filepath = os.path.join(out_data_dir, f"data_{dset_type}.bin") out_metas_filepath = os.path.join(out_data_dir, f"metas_{dset_type}.jsonl") out_info_filepath = os.path.join(out_data_dir, f"info_{dset_type}.json") _ = np.memmap(out_mm_filepath, dtype=np.uint16, mode="w+", shape=(1,)) with open(out_metas_filepath, "w") as f: f.write("") print("start prepare data") for dataset in datasets: n_offs = _prep_data( dataset, out_data_dir, meta_info_map, njobs=njobs, chunksize=chunksize, is_val=is_val, n_offs=n_offs, ) datasets_info = {} with open(out_metas_filepath) as f: n = 0 for line in f: line = line.strip() if len(line) == 0: continue m = json.loads(line) if m["dataset"] not in datasets_info: datasets_info[m["dataset"]] = { "idx_list": [], "task": m["task"], } datasets_info[m["dataset"]]["idx_list"].append(n) n += 1 write_json(datasets_info, out_info_filepath)