# vocab: # 0-60_000 text # 1x0-3999 semantic # 12x0-2047 coarse # 4000 semantic pad token # 4001 semantic infer token # 2048 coarse pad token # 2049 coarse infer token # Memmaps: # Nx9x3584 for audio tokens # Jsons: # N*Dict with meta keys # "dataset" # "original_id", "original_duration_s", # "start_s", "end_s", # "text_segments", "private_text_segments", # "text", "private_text", # "tags", "private_tags", # "views", # Dict with meta keys {"dataset": ["idx_list"]} # Bundles (mert_25_2x4k & dac_2c_25_12): # s3://suno-data/datasets/bundles/ # v1/youtube_music # v1/genius_hq # v1/jamendo # v1/imslp # v2/pond5_music # v2/deezer # v2/ytm_tagged # v3/discogs import os import gc import re import json import math import tqdm import torch import funcy import random import tempfile import collections import numpy as np from collections import defaultdict from joblib import Parallel, delayed from transformers import BertTokenizer from suno_utils.utils.text import ( write_jsonl, read_jsonl, write_json, read_json, normalize_whitespace, ) from suno_utils.utils.s3 import read_from_s3, check_s3_file_exists # ------------ global configuration ------------ 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 ) 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 # why do we do this? if "text" in meta_info and False: # turn off for now 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"] # add extra metas if "audio_filepath" in meta_info: new_meta["audio_filepath"] = meta_info["audio_filepath"] if "s3_filepath" in meta_info: new_meta["s3_filepath"] = meta_info["s3_filepath"] if "audio_quality" in meta_info: new_meta["audio_quality"] = meta_info["audio_quality"] 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"] # add extra metas if "audio_filepath" in arr_meta: add_meta["s3_filepath"] = arr_meta["audio_filepath"] if "s3_filepath" in arr_meta: add_meta["s3_filepath"] = arr_meta["s3_filepath"] if "audio_quality" in arr_meta: add_meta["audio_quality"] = arr_meta["audio_quality"] 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) if __name__ == "__main__": NJOBS = 32 CHUNKSIZE = 32 SEMANTIC_EMBED_DIR = "mert_25_2x4k" CODEC_EMBED_DIR = "dac_2c_25_12" # TODO: hopefully fast enough for multicore write METAS_DIR = "/app/suno/data/chirp_v4_genius_hq_filtered/metadata" OUT_DATA_DIR = "/app/suno/data/chirp_v4_genius_hq_filtered/base" os.makedirs(OUT_DATA_DIR, exist_ok=True) # load manifests of IDs and text and tags etc meta_info_map = { "genius_hq": { m["id"]: m for m in read_jsonl(os.path.join(METAS_DIR, "genius_hq_v6.jsonl")) }, # "youtube_music": {m["id"]: m for m in read_jsonl(os.path.join(METAS_DIR, "youtube_music.jsonl"))}, # "jamendo": { # m["id"]: m for m in read_jsonl(os.path.join(METAS_DIR, "jamendo.jsonl")) # }, # "imslp": { # m["id"]: m # for m in read_jsonl(os.path.join(METAS_DIR, "imslp_filtered.jsonl")) # }, # "pond5_music": {m["id"]: m for m in read_jsonl(os.path.join(METAS_DIR, "pond5_music.jsonl"))}, # "deezer": {m["id"]: m for m in read_jsonl(os.path.join(METAS_DIR, "deezer.jsonl"))}, # "ytm_tagged": {m["id"]: m for m in read_jsonl(os.path.join(METAS_DIR, "ytm_tagged.jsonl"))}, # "discogs": { # m["id"]: m for m in read_jsonl(os.path.join(METAS_DIR, "discogs.jsonl")) # }, } print(meta_info_map.keys()) # load audio quality and merge with meta_info_map with open( "/home/christian/code/christian/metadata/genius_hq_metas_quality.json", "r" ) as f: quality_metas_map = json.load(f) for meta_id, v in meta_info_map["genius_hq"].items(): if meta_id in quality_metas_map: meta_info_map["genius_hq"][meta_id].update(quality_metas_map[meta_id]) datasets = [ # ("youtube_music", "v1", (0, 1), 1, 1), # ("genius_hq", "v1", (1, 4302), 1, 1, "default"), # train ("genius_hq", "v1", (0, 1), 1, 1, "default"), # val # ("jamendo", "v1", (0, 1), 1, 1), # ("imslp", "v1", (0, 1), 1, 1), # ("pond5_music", "v2", (0, 1), 1, 1), # ("deezer", "v2", (0, 1), 1, 1), # ("ytm_tagged", "v2", (0, 1), 1, 1), # ("discogs", "v3", (0, 1), 1, 1), ] prep_data( datasets, OUT_DATA_DIR, meta_info_map, is_val=True, njobs=NJOBS, chunksize=CHUNKSIZE, )