# 7b data prep import numpy as np import tqdm import funcy 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 = 4160 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 = 4288 N_TOKENS_TEXT = 1152 N_TOKENS_AUDIO = 3008 # max 120s of audio # make sure we have enough space for shift 10 assert BLOCK_SIZE >= ( N_TOKENS_TEXT + N_TOKENS_AUDIO + 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 _verify_stuff(dset_name, meta_info): if dset_name == "genius_hq": assert ( "private_text_segments" in meta_info or "private_text" in meta_info or "text_segments" in meta_info or "text" in meta_info ) 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(dset_name, meta_info, semantic_arr, coarse_arr, enable_random_start=True): _verify_stuff(dset_name, meta_info) # prep segment metas (use semantic for timekeeping) segments_info = [] # first check if we have known segments if "text_segments" in meta_info: for m in meta_info["text_segments"]: 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"], m.get("private_text"), True, 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, ) ) else: # randomize offset to not get only multiples if no text available offs = 0 if enable_random_start and "text" not in meta_info and random.random() > 0.9: # don't randomize offsets if we have lyrics offs = random.randint(int(30 * SEMANTIC_RATE_HZ), int(60 * SEMANTIC_RATE_HZ)) offs = min(offs, N_TOKENS_AUDIO - 1) segments_info.append((0, min(len(semantic_arr), offs), None, None, False, None, None)) total_steps = int(np.ceil((len(semantic_arr) - offs) / N_TOKENS_AUDIO)) for n in range(total_steps): start_idx = offs + n * N_TOKENS_AUDIO end_idx = min(len(semantic_arr), offs + (n + 1) * N_TOKENS_AUDIO) 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 meta_info.get("private_text") if n == 0 else None, # add text only to first piece False, None, None, ) ) if dset_name == "musescore": # break after first piece only text-audio pairs useful here # TODO: lang here is a hack meta_info["lang"] = "en" break arr_list = [] for ( sem_start_idx, sem_end_idx, text, private_text, is_aligned, vocal_start_idx, vocal_end_idx, ) in segments_info: if ( sem_end_idx - sem_start_idx > N_TOKENS_AUDIO 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 if len(arr_c) < N_TOKENS_AUDIO: arr_c = np.pad( arr_c, ((0, N_TOKENS_AUDIO - len(arr_c)), (0, 0)), constant_values=COARSE_PAD_TOKEN, mode="constant", ) arr_s = np.pad( arr_s, ((0, N_TOKENS_AUDIO - 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_AUDIO, 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, } if text is not None: new_meta["text"] = text if private_text is not None and private_text != text: new_meta["text_private"] = private_text new_meta["text_lang"] = meta_info.get("lang", "") new_meta["text_aligned"] = is_aligned new_meta["dset_suffix"] = "lyrics" if new_meta["text_lang"] == "en" else "lyrics_foreign" if "tags" in meta_info: new_meta["tags"] = meta_info["tags"] if "private_tags" in meta_info: if "tags" in meta_info: # verify that superset assert len(set(meta_info["tags"]) - set(meta_info["private_tags"])) == 0 if "tags" not in meta_info or meta_info["private_tags"] != meta_info["tags"]: new_meta["tags_private"] = meta_info["private_tags"] if "original_id" in meta_info: new_meta["original_id"] = meta_info["original_id"] if "views" in meta_info: new_meta["views"] = meta_info["views"] arr_list.append((arr, new_meta)) del arr_s, arr_c return arr_list def _process_archives( 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): continue for k, v in read_from_s3(s3_semantic_archive_filepath, read_f=np.load).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): continue for k, v in read_from_s3(s3_coarse_archive_filepath, read_f=np.load).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(dset_name, relevant_metas[uid], semantic_arr, coarse_arr)) del semantic_archive, coarse_archive return arr_list def _collect_uids( s3_semantic_metas_filepaths, s3_coarse_metas_filepaths, ): semantic_uids = [] for fp in s3_semantic_metas_filepaths: semantic_uids.extend([m["id"] for m in read_from_s3(fp, read_f=read_jsonl)]) coarse_uids = [] for fp in s3_coarse_metas_filepaths: coarse_uids.extend([m["id"] for m in read_from_s3(fp, read_f=read_jsonl)]) return set(semantic_uids) & set(coarse_uids) def _prep_data( dataset, out_data_dir, meta_info_map, meta_cutoff_freq, # TODO: This arg should be removed in the longer run... njobs=5, chunksize=10, is_val=False, n_offs=0, ): dset_name, dset_version, (start_idx, end_idx), n_sem, n_coarse = 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 ) 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)( 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 ) # print(len(encoded_arrays_list)) add_metas = [] for encoded_arrays in encoded_arrays_list: # print(len(encoded_arrays)) 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, "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, } # a fill of the cutoff freq if found (only exist for genius as of Jan 18) if add_meta["id"] in meta_cutoff_freq: add_meta["cutoff_freq"] = meta_cutoff_freq[add_meta["id"]] if "original_id" in arr_meta: add_meta["original_id"] = arr_meta["original_id"] if "tags" in arr_meta: add_meta["tags"] = arr_meta["tags"] if "tags_private" in arr_meta: add_meta["tags_private"] = arr_meta["tags_private"] if "text" in arr_meta: add_meta["text"] = arr_meta["text"].strip() if "text_private" in arr_meta: add_meta["text_private"] = arr_meta["text_private"].strip() if "text_lang" in arr_meta: add_meta["text_lang"] = arr_meta["text_lang"] if "text_aligned" in arr_meta: add_meta["text_aligned"] = arr_meta["text_aligned"] if "views" in arr_meta: add_meta["views"] = arr_meta["views"] 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 write_jsonl( add_metas, os.path.join(out_metas_filepath), do_append=bool(n_offs != 0), ) del encoded_arrays_list 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, meta_cutoff_freq, # TODO: This arg should be removed in the longer run..., 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, meta_cutoff_freq, 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": []} datasets_info[m["dataset"]]["idx_list"].append(n) n += 1 write_json(datasets_info, out_info_filepath)