import os

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

import gc
import tqdm
import funcy
import shutil
import numpy as np

from collections import defaultdict
from joblib import Parallel, delayed
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


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

# SEMANTIC_EMBED_DIR = "mert_25_2x4k"
# 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
CODEC_EMBED_DIR = "dac_vae_tuned_25hz"  # new codec for v4.5 (no shimmer)
CODEC_TYPE = "VAE"  # "RVQ" or "VAE"

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

# use same metas as gpt
METAS_DIR = "/app/suno/tmp"
OUT_DATA_DIR = "/app/suno/data/splice_samples_30s/v0"
os.makedirs(OUT_DATA_DIR, exist_ok=True)
tokenizer_fp = os.path.join(OUT_DATA_DIR, "tokenizer_60k.json")
if not os.path.exists(tokenizer_fp):
    shutil.copy("/app/suno/data/chirp_v5/v1/tokenizer_60k.json", tokenizer_fp)

# v0 uses dac_vae_tuned_25hz codec, and genius, discogs_subset, imslp, pond5


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, 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(vae_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(vae_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,
):
    print(dataset)
    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)
    tot_alignments_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_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_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"]
                    tot_alignments_dict[dataset_str] += 1

                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}")

    # for k, v in tot_alignments_dict.items():
    #    print(f"{v} alignments for {k}")

    return 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
    if False:
        meta_info_map = {
            "discogs_subset": {
                m["id"]: m
                for m in read_jsonl(
                    os.path.join(METAS_DIR, "clean_discogs_subset_v0_metas.jsonl")
                )
            },
            "genius": {
                m["id"]: m
                for m in read_jsonl(
                    os.path.join(METAS_DIR, "clean_genius_v0_metas.jsonl")
                )
            },
            "pond5": {
                m["id"]: m
                for m in read_jsonl(
                    os.path.join(METAS_DIR, "clean_pond5_v0_metas.jsonl")
                )
            },
            "imslp": {
                m["id"]: m
                for m in read_jsonl(
                    os.path.join(METAS_DIR, "clean_imslp_v0_metas.jsonl")
                )
            },
        }
    else:
        meta_info_map = {
            "splice_samples_30s": {
                m["id"]: m
                for m in read_jsonl(
                    "/home/christian/code/christian/metadata/bundles/splice_samples_30s/metas.jsonl"
                )
            },
        }

    # genius_alignments_filepath = "/home/christian/code/christian/metadata/genius_hq_alignments_t30_v1_yt_ids.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")

    # first do validiation set
    # (start_idx, end_idx), n_archives_semantic, n_archives_vae
    # dset_name, dset_version, (start_idx, end_idx), n_sem, n_vae

    if True:
        datasets = [
            ("splice_samples_30s", "v0", (0, 1), 1, 1),
        ]
        prep_data(
            datasets,
            is_val=True,
            njobs=NJOBS,
            chunksize=CHUNKSIZE,
        )

    if False:
        if False:
            datasets = [
                ("discogs_subset", "v4", (0, 1), 1, 1),
                ("genius", "v4", (0, 1), 1, 1),
                ("pond5", "v4", (0, 1), 1, 1),
                ("imslp", "v4", (0, 1), 1, 1),
            ]
            prep_data(
                datasets,
                is_val=True,
                njobs=NJOBS,
                chunksize=CHUNKSIZE,
            )

        if True:
            # then do training set
            # (start_idx, end_idx), n_archives_semantic, n_archives_vae
            datasets = [
                ("discogs_subset", "v4", (1, 5834), 1, 1),
                ("genius", "v4", (1, 4181), 1, 1),
                ("pond5", "v4", (1, 4138), 1, 1),
                ("imslp", "v4", (1, 549), 1, 1),
            ]
            prep_data(
                datasets,
                is_val=False,
                njobs=NJOBS,
                chunksize=CHUNKSIZE,
            )
