# 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,
    )
