# 7b data prep
import numpy as np
import tqdm
import funcy
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 = 4160
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 = 4288
N_TOKENS_TEXT = 1152
N_TOKENS_AUDIO = 3008  # max 120s of audio
# make sure we have enough space for shift 10
assert BLOCK_SIZE >= (
    N_TOKENS_TEXT
    + N_TOKENS_AUDIO
    + 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 _verify_stuff(dset_name, meta_info):
    if dset_name == "genius_hq":
        assert (
            "private_text_segments" in meta_info
            or "private_text" in meta_info
            or "text_segments" in meta_info
            or "text" in meta_info
        )


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(dset_name, meta_info, semantic_arr, coarse_arr, enable_random_start=True):
    _verify_stuff(dset_name, meta_info)
    # prep segment metas (use semantic for timekeeping)
    segments_info = []
    # first check if we have known segments
    if "text_segments" in meta_info:
        for m in meta_info["text_segments"]:
            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"],
                    m.get("private_text"),
                    True,
                    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,
                )
            )
    else:
        # randomize offset to not get only multiples if no text available
        offs = 0
        if enable_random_start and "text" not in meta_info and random.random() > 0.9:
            # don't randomize offsets if we have lyrics
            offs = random.randint(int(30 * SEMANTIC_RATE_HZ), int(60 * SEMANTIC_RATE_HZ))
            offs = min(offs, N_TOKENS_AUDIO - 1)
            segments_info.append((0, min(len(semantic_arr), offs), None, None, False, None, None))
        total_steps = int(np.ceil((len(semantic_arr) - offs) / N_TOKENS_AUDIO))
        for n in range(total_steps):
            start_idx = offs + n * N_TOKENS_AUDIO
            end_idx = min(len(semantic_arr), offs + (n + 1) * N_TOKENS_AUDIO)
            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
                    meta_info.get("private_text") if n == 0 else None,  # add text only to first piece
                    False,
                    None,
                    None,
                )
            )
            if dset_name == "musescore":
                # break after first piece only text-audio pairs useful here
                # TODO: lang here is a hack
                meta_info["lang"] = "en"
                break
    arr_list = []
    for (
        sem_start_idx,
        sem_end_idx,
        text,
        private_text,
        is_aligned,
        vocal_start_idx,
        vocal_end_idx,
    ) in segments_info:
        if (
            sem_end_idx - sem_start_idx > N_TOKENS_AUDIO
            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
        if len(arr_c) < N_TOKENS_AUDIO:
            arr_c = np.pad(
                arr_c,
                ((0, N_TOKENS_AUDIO - len(arr_c)), (0, 0)),
                constant_values=COARSE_PAD_TOKEN,
                mode="constant",
            )
            arr_s = np.pad(
                arr_s,
                ((0, N_TOKENS_AUDIO - 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_AUDIO, 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,
        }
        if text is not None:
            new_meta["text"] = text
            if private_text is not None and private_text != text:
                new_meta["text_private"] = private_text
            new_meta["text_lang"] = meta_info.get("lang", "")
            new_meta["text_aligned"] = is_aligned
            new_meta["dset_suffix"] = "lyrics" if new_meta["text_lang"] == "en" else "lyrics_foreign"
        if "tags" in meta_info:
            new_meta["tags"] = meta_info["tags"]
        if "private_tags" in meta_info:
            if "tags" in meta_info:
                # verify that superset
                assert len(set(meta_info["tags"]) - set(meta_info["private_tags"])) == 0
            if "tags" not in meta_info or meta_info["private_tags"] != meta_info["tags"]:
                new_meta["tags_private"] = meta_info["private_tags"]
        if "original_id" in meta_info:
            new_meta["original_id"] = meta_info["original_id"]
        if "views" in meta_info:
            new_meta["views"] = meta_info["views"]
        arr_list.append((arr, new_meta))
        del arr_s, arr_c
    return arr_list


def _process_archives(
    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):
            continue
        for k, v in read_from_s3(s3_semantic_archive_filepath, read_f=np.load).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):
            continue
        for k, v in read_from_s3(s3_coarse_archive_filepath, read_f=np.load).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(dset_name, relevant_metas[uid], semantic_arr, coarse_arr))
    del semantic_archive, coarse_archive
    return arr_list


def _collect_uids(
    s3_semantic_metas_filepaths,
    s3_coarse_metas_filepaths,
):
    semantic_uids = []
    for fp in s3_semantic_metas_filepaths:
        semantic_uids.extend([m["id"] for m in read_from_s3(fp, read_f=read_jsonl)])
    coarse_uids = []
    for fp in s3_coarse_metas_filepaths:
        coarse_uids.extend([m["id"] for m in read_from_s3(fp, read_f=read_jsonl)])
    return set(semantic_uids) & set(coarse_uids)


def _prep_data(
    dataset,
    out_data_dir,
    meta_info_map,
    meta_cutoff_freq,  # TODO: This arg should be removed in the longer run...
    njobs=5,
    chunksize=10,
    is_val=False,
    n_offs=0,
):
    dset_name, dset_version, (start_idx, end_idx), n_sem, n_coarse = 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
        )
        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"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
        )
        # print(len(encoded_arrays_list))
        add_metas = []
        for encoded_arrays in encoded_arrays_list:
            # print(len(encoded_arrays))
            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,
                    "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,
                }
                # a fill of the cutoff freq if found (only exist for genius as of Jan 18)
                if add_meta["id"] in meta_cutoff_freq:
                    add_meta["cutoff_freq"] = meta_cutoff_freq[add_meta["id"]]
                if "original_id" in arr_meta:
                    add_meta["original_id"] = arr_meta["original_id"]
                if "tags" in arr_meta:
                    add_meta["tags"] = arr_meta["tags"]
                if "tags_private" in arr_meta:
                    add_meta["tags_private"] = arr_meta["tags_private"]
                if "text" in arr_meta:
                    add_meta["text"] = arr_meta["text"].strip()
                if "text_private" in arr_meta:
                    add_meta["text_private"] = arr_meta["text_private"].strip()
                if "text_lang" in arr_meta:
                    add_meta["text_lang"] = arr_meta["text_lang"]
                if "text_aligned" in arr_meta:
                    add_meta["text_aligned"] = arr_meta["text_aligned"]
                if "views" in arr_meta:
                    add_meta["views"] = arr_meta["views"]
                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
        write_jsonl(
            add_metas,
            os.path.join(out_metas_filepath),
            do_append=bool(n_offs != 0),
        )
        del encoded_arrays_list
    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,
    meta_cutoff_freq,  # TODO: This arg should be removed in the longer run...,
    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,
            meta_cutoff_freq,
            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": []}
            datasets_info[m["dataset"]]["idx_list"].append(n)
            n += 1
    write_json(datasets_info, out_info_filepath)
