import glob import os import numpy as np from tqdm import tqdm from suno_utils.utils.s3 import read_from_s3 from suno_utils.utils.text import read_jsonl, write_jsonl if __name__ == "__main__": npz_dir = "/home/christian/data/dpo/genius_hq_corrupt_dpo_25hz_30_v2_npz" out_dir = "/home/christian/data/dpo/genius_hq_corrupt_dpo_25hz_30_v5" os.makedirs(out_dir, exist_ok=True) # now we create a train and val memmap for these latents with new metas base_metas = ( "/home/christian/code/christian/metadata/genius_hq_metas_filtered.jsonl" ) metas = read_jsonl(base_metas) print(f"loaded {len(metas):,} metas") # also load the lyric alignments genius_alignments_filepath = ( "/home/tony/Work/tony/hoot/tmp/genius_hq_alignments_t30_v1.jsonl" ) if os.path.exists(genius_alignments_filepath): print("loading genius alignments") genius_alignments = read_jsonl(genius_alignments_filepath, progress=False) genius_alignments_map = {a[0]: a[1] for a in genius_alignments} else: genius_alignments_map = {} # merge alignments into single map alignments_map = {**genius_alignments_map} print(f"loaded {len(alignments_map):,} alignments") # create metas map metas_map = {meta["id"]: meta for meta in metas} print(f"loaded {len(metas_map):,} metas") # find all npz files in the out_dir npz_files = glob.glob(os.path.join(npz_dir, "*.npz")) print(len(npz_files)) SEMANTIC_MEMMAP_SIZE = 750 VAE_MEMMAP_SIZE = 750 VAE_DIM = 128 # split into train and val npz_files_train = npz_files[: int(len(npz_files) * 0.9)] npz_files_val = npz_files[int(len(npz_files) * 0.9) :] print( f"found {len(npz_files):,} npz files, {len(npz_files_train):,} train, {len(npz_files_val):,} val" ) for dset, npz_files in [("tr", npz_files_train), ("val", npz_files_val)]: out_mm_vae_filepath = os.path.join(out_dir, f"data_vae_{dset}.bin") out_mm_semantic_filepath = os.path.join(out_dir, f"data_semantic_{dset}.bin") out_metas_filepath = os.path.join(out_dir, f"metas_{dset}.jsonl") new_metas = [] # 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 to_write_len_s = SEMANTIC_MEMMAP_SIZE * len(npz_files) * 2 to_write_len_v = VAE_MEMMAP_SIZE * VAE_DIM * len(npz_files) * 2 print(to_write_len_s, to_write_len_v) 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,), ) for npz_file in tqdm(npz_files): data = np.load(npz_file) meta_id = os.path.basename(npz_file).replace(".npz", "") meta = metas_map[meta_id] # construct text aligned if meta_id in alignments_map: alignments = alignments_map[meta_id] else: alignments = [] text_aligned = "" for text_segment in alignments: start_s = text_segment["start_s"] end_s = text_segment["end_s"] text_aligned = text_segment["text"] # construct new meta new_meta = meta.copy() new_meta["text"] = meta.get("lyrics", "") new_meta["text_aligned"] = text_aligned new_meta["tags"] = meta.get("tags_text", []) new_meta["n_vae_tokens"] = VAE_MEMMAP_SIZE new_meta.pop("lyrics", None) # append it twince, once for original and once for corrupted new_metas.append(new_meta) new_metas.append(new_meta) # write to memmap arr_s = data["semantic_codes"] arr_v = data["vae_latents_original"].astype(np.float16) arr_v = np.swapaxes(arr_v, 0, 1) 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 arr_s = data["semantic_codes_corrupted"] arr_v = data["vae_latents_corrupted"].astype(np.float16) arr_v = np.swapaxes(arr_v, 0, 1) 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)