import os import numpy as np import shutil import random from tqdm import tqdm from suno_utils.utils.text import read_jsonl, write_jsonl from concurrent.futures import ThreadPoolExecutor, as_completed if __name__ == "__main__": NPZ_DIR = "/app/suno/christian/data/genius_hq_filtered_raw_10s_1920_npz" OUT_DATA_DIR = "/app/suno/christian/data/genius_hq_filtered_raw_10s_1920_memmap" METAS_PATH = ( "/home/christian/code/christian/metadata/genius_hq_metas_filtered.jsonl" ) VAL_SIZE = 10_000 VAL_ONLY = False DURATION_S = 10.0 VAE_FRAME_RATE = 25 SEMANTIC_FRAME_RATE = 25 CHUNK_SIZE = 100 SEMANTIC_MEMMAP_SIZE = int(DURATION_S * SEMANTIC_FRAME_RATE) VAE_MEMMAP_SIZE = int(DURATION_S * VAE_FRAME_RATE * 2) VAE_DIM = 1920 print(f"VAE_MEMMAP_SIZE: {VAE_MEMMAP_SIZE}") print(f"SEMANTIC_MEMMAP_SIZE: {SEMANTIC_MEMMAP_SIZE}") # load base metas base_metas = read_jsonl(METAS_PATH) print(f"Found {len(base_metas)} metas") # delete the out dir if it exists if os.path.exists(OUT_DATA_DIR): shutil.rmtree(OUT_DATA_DIR) os.makedirs(OUT_DATA_DIR, exist_ok=True) # shuffle the metas random.seed(42) random.shuffle(base_metas) # split into train and val train_metas = base_metas[:-VAL_SIZE] val_metas = base_metas[-VAL_SIZE:] # now iterate over the val, then train metas if VAL_ONLY: dset_types = ["val"] else: dset_types = ["val", "tr"] for dset_type in dset_types: dset_metas = val_metas if dset_type == "val" else train_metas out_mm_vae_filepath = os.path.join(OUT_DATA_DIR, f"data_vae_{dset_type}.bin") out_metas_filepath = os.path.join(OUT_DATA_DIR, f"metas_{dset_type}.jsonl") out_mm_semantic_filepath = os.path.join( OUT_DATA_DIR, f"data_semantic_{dset_type}.bin" ) # initial write out_mm_semantic = np.memmap( out_mm_semantic_filepath, dtype=np.uint16, mode="w+", shape=(1), ) out_mm_vae = np.memmap( out_mm_vae_filepath, dtype=np.float16, mode="w+", shape=(1), ) n_offs_s = 0 n_offs_v = 0 # split valid metas into chunks of CHUNK_SIZE dset_metas_chunks = [ dset_metas[i : i + CHUNK_SIZE] for i in range(0, len(dset_metas), CHUNK_SIZE) ] print("total chunks: ", len(dset_metas_chunks)) for chunk_idx, dset_meta_chunk in enumerate(tqdm(dset_metas_chunks)): arr_s_list = [] arr_v_list = [] new_metas = [] def process_meta(meta): # try to load from npz npz_path = os.path.join(NPZ_DIR, f"{meta['id']}.npz") if os.path.exists(npz_path): data = np.load(npz_path) arr_s = data["semantic_codes"] arr_v = data["frames"] else: return None # check for nan in arr_v or arr_s if not np.all(np.isfinite(arr_v)) or not np.all(np.isfinite(arr_s)): return None try: if arr_s.size < SEMANTIC_MEMMAP_SIZE: return None if arr_v.size < VAE_MEMMAP_SIZE * VAE_DIM: return None result = (arr_s, arr_v, meta) return result except Exception as e: print(f"error loading {meta['id']}: {e}") return None with ThreadPoolExecutor(max_workers=16) as executor: futures = [ executor.submit(process_meta, meta) for meta in dset_meta_chunk ] for future in as_completed(futures): result = future.result() if result: arr_s, arr_v, meta = result arr_s_list.append(arr_s) arr_v_list.append(arr_v) new_meta = meta.copy() new_meta["text"] = meta["lyrics"] new_meta["tags"] = meta["tags_text"] new_meta["n_vae_tokens"] = VAE_MEMMAP_SIZE new_metas.append(new_meta) print(len(arr_s_list), len(arr_v_list), len(new_metas)) assert len(arr_s_list) == len(arr_v_list) == len(new_metas) # now write to the memmap # get a list of all the ids in the id_to_s3_paths to_write_len_s = SEMANTIC_MEMMAP_SIZE * len(arr_v_list) to_write_len_v = VAE_MEMMAP_SIZE * VAE_DIM * len(arr_v_list) out_mm_semantic = np.memmap( out_mm_semantic_filepath, dtype=np.uint16, mode="r+", shape=(n_offs_s + to_write_len_s,), ) out_mm_vae = np.memmap( out_mm_vae_filepath, dtype=np.float16, mode="r+", shape=(n_offs_v + to_write_len_v,), ) # write to memmap (has to happen sequentially) for new_meta, arr_s, arr_v in zip(new_metas, arr_s_list, arr_v_list): out_mm_semantic[n_offs_s : n_offs_s + arr_s.size] = arr_s.reshape( -1, ) out_mm_vae[n_offs_v : n_offs_v + arr_v.size] = arr_v.reshape( -1, ) n_offs_s += arr_s.size n_offs_v += arr_v.size # write it once out_mm_semantic.flush() out_mm_vae.flush() del out_mm_semantic, out_mm_vae write_jsonl(new_metas, out_metas_filepath, do_append=True)