import os

os.environ["CUDA_VISIBLE_DEVICES"] = ""
os.environ["PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION"] = "python"

# 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

import math
import numpy as np
import tqdm
import time
import torch
import funcy
import json
import gc
import re
import random
import tempfile
import collections
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

import os

os.environ["OMP_NUM_THREADS"] = "1"

# Constants
DURATION_S = 30  # 30 sec
# DURATION_S = 120  # 2 min = 120 sec
# DURATION_S = 240  # 4 min = 240 sec
# DURATION_S = 360  # 6 min = 120 sec
SEMANTIC_CODEBOOK_SIZE = 4000
SEMANTIC_N_CODEBOOKS = 1
SEMANTIC_RATE_HZ = 25
SEMANTIC_N_TOKENS_MEMMAP = int(DURATION_S * SEMANTIC_RATE_HZ)
print(f"SEMANTIC_N_TOKENS_MEMMAP: {SEMANTIC_N_TOKENS_MEMMAP}")

CODEC_CODEBOOK_SIZE = 2048
CODEC_N_CODEBOOKS = 12

CODEC_DIM = 128
CODEC_N_CODEBOOKS = 12
CODEC_RATE_HZ = 25
CODEC_N_TOKENS_MEMMAP = int(DURATION_S * CODEC_RATE_HZ)
print(f"CODEC_N_TOKENS_MEMMAP: {CODEC_N_TOKENS_MEMMAP}")

HOOT_RATE_HZ = 12.5333333
HOOT_N_TOKENS_MEMMAP = int(DURATION_S * HOOT_RATE_HZ)
print(f"HOOT_N_TOKENS_MEMMAP: {HOOT_N_TOKENS_MEMMAP}")

# # TODO: not sure yet what's needed but vae is large so let's use small for now
# NJOBS = 32
# CHUNKSIZE = 32
NJOBS = 32
CHUNKSIZE = 32

SEMANTIC_EMBED_DIR = "mert_25"
# SEMANTIC_EMBED_DIR = "musicfm_v2_l6_5s"
# CODEC_EMBED_DIR = "dac_2c_4"
CODEC_EMBED_DIR = "convnext_vae_tuned_25hz"
# CODEC_EMBED_DIR = "dac_vae_100hz_peaq"  # new 100 Hz VAE with KL term
# CODEC_EMBED_DIR = "dac_vae_25hz_64_peaq"  # new 25 Hz VAE with KL term
# CODEC_EMBED_DIR = "dac_vae_25hz_peaq"  # new 25 Hz VAE with KL term
# HOOT_EMBED_DIR = "hoot_emb_512"  # new 512 dim hoot embeddings

CODEC_TYPE = "VAE"  # "RVQ" or "VAE"

# to use local cache run this: note that this is a large download and will require about 40TB of space
# mkdir /mnt/localdisk/datasets_cache
# cd /mnt/localdisk/datasets_cache
# aws s3 sync s3://suno-data/datasets/bundles/v1/genius_hq/mert_25_2x4k bundles/v1/genius_hq/mert_25_2x4k
# aws s3 sync s3://suno-data/datasets/bundles/v1/genius_hq/dac_vae_100hz_peaq bundles/v1/genius_hq/dac_vae_100hz_peaq
# aws s3 sync s3://suno-data/datasets/bundles/v1/youtube_music/mert_25_2x4k bundles/v1/youtube_music/mert_25_2x4k
# aws s3 sync s3://suno-data/datasets/bundles/v1/youtube_music/dac_vae_100hz_peaq bundles/v1/youtube_music/dac_vae_100hz_peaq


# aws s3 sync s3://suno-data/datasets/bundles/v3/diffusion_mix/mert_25 bundles/v3/diffusion_mix/mert_25
# aws s3 sync s3://suno-data/datasets/bundles/v3/diffusion_mix/dac_vae_25hz_peaq bundles/v3/diffusion_mix/dac_vae_25hz_peaq

# aws s3 sync s3://suno-data/datasets/bundles/v3/diffusion_mix_fix/mert_25 bundles/v3/diffusion_mix_fix/mert_25
# aws s3 sync s3://suno-data/datasets/bundles/v3/diffusion_mix_fix/dac_vae_25hz_peaq bundles/v3/diffusion_mix_fix/dac_vae_25hz_peaq
# aws s3 sync s3://suno-data/datasets/bundles/v3/diffusion_mix_fix/dac_vae_100hz_peaq bundles/v3/diffusion_mix_fix/dac_vae_100hz_peaq


# aws s3 cp genius_alignments_v11.jsonl s3://suno-data/datasets/metadata/alignments/genius_alignments_v11.jsonl
# aws s3 cp ytm_alignments_v11.jsonl s3://suno-data/datasets/metadata/alignments/ytm_alignments_v11.jsonl


# BUNDLES_DIR = "/mnt/localdisk/datasets_cache/bundles"  # for use with local disk
BUNDLES_DIR = "s3://suno-data/datasets/bundles"  # for use with S3

# METAS_DIR = "/app/suno/data/chirp_v4/metadata_30s"
METAS_DIR = "/app/suno/data/diffusion_mix/metadata"
OUT_DATA_DIR = "/app/suno/data/diffusion_mix/convnext_vae_tuned_25hz"
os.makedirs(OUT_DATA_DIR, exist_ok=True)

# v4 100 Hz VAE 30 sec chunks
# v5 100 Hz VAE 360 sec chunks


def _trim_to_common(arr_1, arr_2, arr_3):
    common_len = min(len(arr_1), len(arr_2), len(arr_3))
    arr_1 = arr_1[:common_len]
    arr_2 = arr_2[:common_len]
    arr_3 = arr_3[:common_len]
    return arr_1, arr_2, arr_3


def _parse_arrays(dset_name, meta_info, semantic_arr, vae_arr):
    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))
            # print(m["start_s"], m["end_s"], m["text"], selected_line_start_s)
            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
            # i don't know why we do this. I think this adds all lyrics as one example
            # let's exclude for now as we want to focus on aligned segments
            if "text" in meta_info and False:
                segments_info.append(
                    (
                        0,
                        min(len(semantic_arr), SEMANTIC_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, SEMANTIC_N_TOKENS_MEMMAP - 1)
        #    segments_info.append(
        #        (0, min(len(semantic_arr), offs), None, None, None, None)
        #    )
        # total_steps = 1  # only use the first chunk
        total_steps = len(semantic_arr) // SEMANTIC_N_TOKENS_MEMMAP
        # print(f"total_steps: {total_steps}")
        for n in range(total_steps):
            start_idx = offs + n * SEMANTIC_N_TOKENS_MEMMAP
            end_idx = min(len(semantic_arr), offs + (n + 1) * SEMANTIC_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"),
                    None,
                    None,
                    None,
                )
            )
    arr_list = []
    # print("segments_info", len(segments_info))
    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 > SEMANTIC_N_TOKENS_MEMMAP
        #    or sem_end_idx - sem_start_idx < SEMANTIC_RATE_HZ  # arbitrary
        # ):
        #    continue

        # in this mode, we use the exact start and end times
        # vae_start_idx = int(round(sem_start_idx * VAE_RATE_HZ / SEMANTIC_RATE_HZ))
        # vae_end_idx = int(round(sem_end_idx * VAE_RATE_HZ / SEMANTIC_RATE_HZ))

        # in this mode, we use the start time, but the end is the memmap length
        # this is like padding the segment with extra part of the song
        # this could cause slight errors in alignment but should be fine
        sem_end_idx = sem_start_idx + SEMANTIC_N_TOKENS_MEMMAP
        vae_start_idx = int(round(sem_start_idx * CODEC_RATE_HZ / SEMANTIC_RATE_HZ))
        vae_end_idx = vae_start_idx + CODEC_N_TOKENS_MEMMAP

        assert sem_end_idx >= 0 and vae_start_idx >= 0

        # we will allow padding so don't need to check if end is out of bounds
        # if sem_end_idx > len(semantic_arr) or vae_end_idx > len(vae_arr):
        #    continue

        # get array segments
        arr_s = semantic_arr[sem_start_idx:sem_end_idx, :SEMANTIC_N_CODEBOOKS].copy()
        arr_v = vae_arr[vae_start_idx:vae_end_idx, :].copy()

        # print(semantic_arr.shape, vae_arr.shape)
        # print(arr_s.shape, arr_v.shape)
        # print()

        # adjust end indices to match the actual length
        sem_end_idx = sem_start_idx + arr_s.shape[0]
        vae_end_idx = vae_start_idx + arr_v.shape[0]

        # check the size before padding
        n_sem_tokens = arr_s.shape[0]
        n_vae_tokens = arr_v.shape[0]

        # this is for padding them to the proper length
        # but should not be needed in this setup when using alignments
        if arr_s.shape[0] < SEMANTIC_N_TOKENS_MEMMAP:
            arr_s = np.pad(
                arr_s,
                ((0, SEMANTIC_N_TOKENS_MEMMAP - arr_s.shape[0]), (0, 0)),
                mode="constant",
                constant_values=SEMANTIC_CODEBOOK_SIZE,
            )
        if CODEC_TYPE == "VAE":
            if arr_v.shape[0] < CODEC_N_TOKENS_MEMMAP:
                arr_v = np.pad(
                    arr_v,
                    ((0, CODEC_N_TOKENS_MEMMAP - arr_v.shape[0]), (0, 0)),
                    mode="constant",
                    constant_values=0,
                )
        elif CODEC_TYPE == "RVQ":
            if arr_v.shape[0] < CODEC_N_TOKENS_MEMMAP:
                arr_v = np.pad(
                    arr_v,
                    ((0, CODEC_N_TOKENS_MEMMAP - arr_v.shape[0]), (0, 0)),
                    mode="constant",
                    constant_values=CODEC_CODEBOOK_SIZE,
                )

        # fix any alignment mistakes
        # arr_s, arr_c = _trim_to_common(arr_s, arr_c)
        # don't do this for now

        # assert len(arr_s) == len(arr_c)
        # they wont be the same length

        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,
            "n_vae_tokens": n_vae_tokens,
            "n_semantic_tokens": n_sem_tokens,
        }
        if text is not None:
            new_meta["text"] = text
            new_meta["text_lang"] = meta_info.get("lang", None)
            if task_name == "default":
                new_meta["dset_suffix"] = (
                    "lyrics"
                    if meta_info.get("lang", None) == "en"
                    else "lyrics_foreign"
                )
        if "phonemized_text" in meta_info:
            new_meta["phonemized_text"] = meta_info["phonemized_text"]
        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"]

        # check if we have alignments
        if new_meta["id"] in alignments_map:
            alignments = alignments_map[new_meta["id"]]
            start_s = new_meta["start_s"]
            end_s = new_meta["end_s"]
            for alignment in alignments:
                if start_s == alignment["start_s"] and end_s == alignment["end_s"]:
                    new_meta["text_aligned"] = alignment["text"]

        # 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_s, arr_v, new_meta))
        del arr_s, arr_v
    return arr_list


def _process_archives(
    dset_name,
    s3_semantic_archive_filepaths,
    s3_vae_archive_filepaths,
    relevant_metas,
):
    #     print(len(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 "s3://" in s3_semantic_archive_filepath:
            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
        else:
            # check if file exists
            if not os.path.exists(s3_semantic_archive_filepath):
                print(f"missing {s3_semantic_archive_filepath}")
                continue
            try:
                archive = {
                    k: v for k, v in np.load(s3_semantic_archive_filepath).items()
                }
            except:
                # corrupt archive
                print(f"corrupt {s3_semantic_archive_filepath}")
                continue
        for k, v in archive.items():
            semantic_archive[k] = v

    vae_archive = {}
    s3_vae_archive_filepaths = set(s3_vae_archive_filepaths)
    for s3_vae_archive_filepath in s3_vae_archive_filepaths:
        if "s3://" in s3_vae_archive_filepath:
            if not check_s3_file_exists(s3_vae_archive_filepath):
                print(f"missing {s3_vae_archive_filepath}")
                continue
            try:
                archive = {
                    k: v
                    for k, v in read_from_s3(
                        s3_vae_archive_filepath, read_f=np.load
                    ).items()
                }
            except Exception as e:
                # corrupt archive
                print(e)
                print(f"corrupt {s3_vae_archive_filepath}")
                continue
        else:
            # check if file exists
            if not os.path.exists(s3_vae_archive_filepath):
                print(f"missing {s3_vae_archive_filepath}")
                continue
            try:
                archive = {k: v for k, v in np.load(s3_vae_archive_filepath).items()}
            except:
                # corrupt archive
                print(f"corrupt {s3_vae_archive_filepath}")
                continue
        for k, v in archive.items():
            vae_archive[k] = v

    semantic_uids, vae_uids = (
        set(semantic_archive.keys()),
        set(vae_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 & vae_uids:
        if uid not in relevant_metas:
            continue
        semantic_arr = semantic_archive[uid]
        vae_arr = vae_archive[uid]

        # if (
        #    np.abs(len(hoot_arr) / HOOT_RATE_HZ - len(semantic_arr) / SEMANTIC_RATE_HZ)
        #    > 0.1
        #    or np.abs(len(hoot_arr) / HOOT_RATE_HZ - len(vae_arr) / VAE_RATE_HZ) > 0.1
        # ):
        #    # skip if embeddings not roughly the same duration
        #    continue

        # trim to common length (do not need this for now but might be useful later)
        # semantic_arr, hoot_arr, vae_arr = _trim_to_common(
        #    semantic_arr, hoot_arr, vae_arr
        # )
        # assert len(codec_arr) == len(semantic_arr) == len(vae_arr)

        arr_list.extend(
            _parse_arrays(dset_name, relevant_metas[uid], semantic_arr, vae_arr)
        )
    del semantic_archive, vae_archive
    gc.collect()
    return arr_list


def _collect_uids(
    s3_semantic_metas_filepaths,
    s3_vae_metas_filepaths,
):
    semantic_uids = []
    for fp in s3_semantic_metas_filepaths:
        try:
            if "s3://" in fp:
                metas = read_from_s3(fp, read_f=read_jsonl)
            else:
                metas = read_jsonl(fp, progress=False)
        except:
            print(f"failed on metas for fp: {fp}")
            continue
        semantic_uids.extend([m["id"] for m in metas])

    vae_uids = []
    for fp in s3_vae_metas_filepaths:
        try:
            if "s3://" in fp:
                metas = read_from_s3(fp, read_f=read_jsonl)
            else:
                metas = read_jsonl(fp, progress=False)
        except:
            print(f"failed on metas for fp: {fp}")
            continue
        vae_uids.extend([m["id"] for m in metas])
    return set(semantic_uids) & set(vae_uids)


def _prep_data(
    dataset,
    njobs=5,
    chunksize=10,
    is_val=False,
    n_offs_s=0,
    n_offs_v=0,
):
    dset_name, dset_version, (start_idx, end_idx), n_sem, n_vae = dataset
    dset_type = "val" if is_val else "tr"
    out_mm_semantic_filepath = os.path.join(
        OUT_DATA_DIR, f"data_semantic_{dset_type}.bin"
    )
    # out_mm_hoot_filepath = os.path.join(OUT_DATA_DIR, f"data_hoot_{dset_type}.bin")
    out_mm_vae_filepath = os.path.join(OUT_DATA_DIR, f"data_vae_{dset_type}.bin")
    out_mm_codec_filepath = os.path.join(OUT_DATA_DIR, f"data_codec_{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="processes")(
            delayed(_collect_uids)(
                [
                    f"{BUNDLES_DIR}/{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"{BUNDLES_DIR}/{dset_version}/{dset_name}/{CODEC_EMBED_DIR}/"
                    + f"metas/part_{idx_idx}.jsonl"
                    for idx_idx in range(idx * n_vae, (idx + 1) * n_vae)
                ],
            )
            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)(
                dset_name,
                [
                    f"{BUNDLES_DIR}/{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"{BUNDLES_DIR}/{dset_version}/{dset_name}/{CODEC_EMBED_DIR}/"
                    + f"part_{idx_idx}.npz"
                    for idx_idx in range(idx * n_vae, (idx + 1) * n_vae)
                ],
                {
                    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_s = np.sum([arr.size for arr, _, _ in encoded_arrays])
            to_write_len_v = np.sum([arr.size for _, arr, _ in encoded_arrays])
            if to_write_len_s == 0 or to_write_len_v == 0:
                continue
            out_mm_semantic = np.memmap(
                out_mm_semantic_filepath,
                dtype=np.uint16,
                mode="r+",
                shape=(n_offs_s + to_write_len_s,),
            )
            # out_mm_hoot = np.memmap(
            #    out_mm_hoot_filepath,
            #    dtype=np.float16,
            #    mode="r+",
            #    shape=(n_offs_h + to_write_len_h,),
            # )
            if CODEC_TYPE == "VAE":
                out_mm_vae = np.memmap(
                    out_mm_vae_filepath,
                    dtype=np.float16,
                    mode="r+",
                    shape=(n_offs_v + to_write_len_v,),
                )
            elif CODEC_TYPE == "RVQ":
                out_mm_codec = np.memmap(
                    out_mm_codec_filepath,
                    dtype=np.uint16,
                    mode="r+",
                    shape=(n_offs_v + to_write_len_v,),
                )
            for arr_s, arr_v, arr_meta in encoded_arrays:
                out_mm_semantic[n_offs_s : n_offs_s + arr_s.size] = arr_s.reshape(
                    -1,
                )
                if CODEC_TYPE == "VAE":
                    out_mm_vae[n_offs_v : n_offs_v + arr_v.size] = arr_v.reshape(
                        -1,
                    )
                elif CODEC_TYPE == "RVQ":
                    out_mm_codec[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
                dataset_str = dset_name
                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"],
                    "n_vae_tokens": arr_meta["n_vae_tokens"],
                    "n_semantic_tokens": arr_meta["n_semantic_tokens"],
                }
                if "original_id" in arr_meta:
                    add_meta["original_id"] = arr_meta["original_id"]

                tot_duration_dict[dataset_str] += (
                    arr_meta["end_s"] - arr_meta["start_s"]
                )

                if "text" in arr_meta:
                    add_meta["text"] = arr_meta["text"]
                    add_meta["text_lang"] = arr_meta.get("text_lang")
                    add_meta["dset_suffix"] = arr_meta.get("dset_suffix")

                if "phonemized_text" in arr_meta:
                    add_meta["phonemized_text"] = arr_meta["phonemized_text"]

                if "text_aligned" in arr_meta:
                    add_meta["text_aligned"] = arr_meta["text_aligned"]

                if "tags" in arr_meta:
                    add_meta["tags"] = arr_meta["tags"]

                if "audio_filepath" in arr_meta:
                    add_meta["audio_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"]

                add_metas.append(add_meta)
            # write it once
            out_mm_semantic.flush()
            if CODEC_TYPE == "VAE":
                out_mm_vae.flush()
                del out_mm_vae
            elif CODEC_TYPE == "RVQ":
                out_mm_codec.flush()
                del out_mm_codec
            del out_mm_semantic
        ## end Part B
        write_jsonl(
            add_metas,
            os.path.join(out_metas_filepath),
            do_append=bool(n_offs_s != 0),
        )
        del encoded_arrays_list
    # TODO: this gc collect takes super long but maybe ok outside of loop. somehow needed sometimes
    gc.collect()
    for k, v in tot_duration_dict.items():
        print(f"{round(v / 60 / 60):,} hours of {k}")
    return n_offs_s, n_offs_v


def prep_data(
    datasets,
    is_val=False,
    njobs=5,
    chunksize=10,
):
    n_offs_s = 0
    n_offs_v = 0
    dset_type = "val" if is_val else "tr"
    out_mm_semantic_filepath = os.path.join(
        OUT_DATA_DIR, f"data_semantic_{dset_type}.bin"
    )
    if CODEC_TYPE == "RVQ":
        out_mm_codec_filepath = os.path.join(
            OUT_DATA_DIR, f"data_codec_{dset_type}.bin"
        )
        out_mm_codec = np.memmap(
            out_mm_codec_filepath, dtype=np.uint16, mode="w+", shape=(1,)
        )
    elif CODEC_TYPE == "VAE":
        out_mm_vae_filepath = os.path.join(OUT_DATA_DIR, f"data_vae_{dset_type}.bin")
        out_mm_vae = np.memmap(
            out_mm_vae_filepath, dtype=np.float16, mode="w+", shape=(1,)
        )

    out_metas_filepath = os.path.join(OUT_DATA_DIR, f"metas_{dset_type}.jsonl")
    out_mm_semantic = np.memmap(
        out_mm_semantic_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_s, n_offs_v = _prep_data(
            dataset,
            njobs=njobs,
            chunksize=chunksize,
            is_val=is_val,
            n_offs_s=n_offs_s,
            n_offs_v=n_offs_v,
        )


if __name__ == "__main__":
    # load manifests of IDs and text and tags etc
    meta_info_map = {
        "diffusion_mix_fix": {
            m["id"]: m
            for m in read_jsonl(os.path.join(METAS_DIR, "metas.jsonl"), progress=True)
        }
    }

    # load alignemnts
    ytm_alignments_filepath = (
        "/home/tony/Work/tony/hoot/tmp/ytm_hq_alignments_t30_v1.jsonl"
    )
    genius_alignments_filepath = (
        "/home/tony/Work/tony/hoot/tmp/genius_hq_alignments_t30_v1.jsonl"
    )
    if os.path.exists(ytm_alignments_filepath):
        print("loading youtube music alignments")
        ytm_alignments = read_jsonl(ytm_alignments_filepath, progress=False)
        ytm_alignments_map = {a[0]: a[1] for a in ytm_alignments}
    else:
        ytm_alignments_map = {}
    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 = {**ytm_alignments_map, **genius_alignments_map}
    print(f"loaded {len(alignments_map):,} alignments")

    # first do validiation set
    # (start_idx, end_idx), n_archives_semantic, n_archives_vae
    if False:
        datasets = [
            ("diffusion_mix_fix", "v3", (0, 1), 1, 1),
        ]
        prep_data(
            datasets,
            is_val=True,
            njobs=NJOBS,
            chunksize=CHUNKSIZE,
        )
        # 28 hours of youtube_music
        # 18 hours of genius_hq

    if True:  # no training for now
        # then do training set
        # (start_idx, end_idx), n_archives_semantic, n_archives_vae
        datasets = [
            ("diffusion_mix_fix", "v3", (1, 2921), 1, 1),
        ]
        prep_data(
            datasets,
            is_val=False,
            njobs=NJOBS,
            chunksize=CHUNKSIZE,
        )
        # youtube_music: ~45m prep
        #     88,538 hours of youtube_music
        #     21,141 hours of youtube_music_lyrics
        #     18,351 hours of youtube_music_lyrics_foreign
        # genius_hq_lyrics: ~50m prep
        #     90,948 hours of genius_hq_lyrics
        #     49,259 hours of genius_hq_lyrics_foreign
