# 30b data prep
import numpy as np
import tqdm
import funcy
import gc
import json
import os
import random
from collections import defaultdict
from joblib import Parallel, delayed

from suno_utils.utils.text import write_jsonl, read_jsonl, write_json
from suno_utils.utils.s3 import read_from_s3, check_s3_file_exists

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
)

SEMANTIC_EMBED_DIR = "mert_25_2x4k"
CODEC_EMBED_DIR = "dac_2c_25_12"


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
            if "text" in meta_info:
                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"]
        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"]
                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)
