import dataclasses
import json
from math import inf, ceil
import random
import time
import numpy as np
from tqdm import tqdm
from api import (
    ComposerAPI,
    SearchCtx,
    ComposerAPIBasePerfRecorder,
    GenerationRequest,
)
from data import Vocab, midi2wavtool, VocabError
from data_gen import (
    stream_examples,
    augment_example,
)
from train_target import train_target_wants_text, WITH_TEXT
from data_aug import Augmenter, TextAugmenter, AugmentError
import sys
import torch
from embedder import Embedder
from loader import load_config
from logitbias_json import json2logitbias
from model import ModelArgs, MusicalPositionEmbedTransformer
from dump_training_examples import clip_props, separate_drums
from model_accelerated import MusicalPositionEmbedTransformerAccelerated
import itertools


def generate_eval_data(config, vocab):
    seq_len = int(config["seq_len"])
    seq_len_min = int(config["seq_len_min"])
    eval_split = int(config["eval_split_idx"])
    auger = Augmenter.from_config(vocab, config["aug"])
    text_auger = TextAugmenter(is_eval=False)

    for (
        id,
        descs,
        stripe,
        example,
        accomps,
        example_tags,
        accomps_tagss,
    ) in stream_examples(
        split=eval_split,
        minlen=seq_len_min,
        maxlen=seq_len,
        order="random",
        aug_limit=10,
        train_target=WITH_TEXT,
    ):
        try:
            main, combined_accomp, _, text = augment_example(
                stripe,
                vocab,
                auger,
                text_auger,
                descs,
                example,
                None,
                accomps,
                seq_len_min,
                seq_len,
                example_tags,
                accomps_tagss,
            )
        except (VocabError, AugmentError):
            continue
        yield main, combined_accomp, text


def gen_critic_data(model_prefix):
    config = load_config(f"{model_prefix}/config.json")
    vocab = Vocab.from_config(config)
    encoder = Embedder() if train_target_wants_text(config["train_target"]) else None
    model = MusicalPositionEmbedTransformer(
        vocab,
        dataclasses.replace(
            ModelArgs.from_config(vocab, config),
            max_batch_size=2,
            max_seq_len=config["inference"]["max_decoder_seq_len"],
            cache=False,
            enable_flash=True,
        ),
        encoder,
    )
    model.load_state_dict(torch.load(f"{model_prefix}/model.pt", map_location="cpu"))
    model.eval()
    model = model.cuda()

    model_accelerated = MusicalPositionEmbedTransformerAccelerated(
        model.params,
        f"{model_prefix}/model.trt",
        f"{model_prefix}/model_one_step.trt",
        model.encoder,
    )

    api = ComposerAPI(vocab, model_accelerated)

    class_embeds = []
    nonclass_embeds = []
    class_examples = []
    nonclass_examples = []

    limit = 1000

    for main, combined_accomp, text in tqdm(
        itertools.islice(generate_eval_data(config, vocab), limit), total=limit
    ):
        encoder_input_ids, encoder_attention_mask = encoder.tokenize([text])

        if len(main) == 0:
            continue

        is_drums = main[0]["note"] >= 1000

        main_len = min(max(n["offBeat"] for n in main), 39)
        main = [n for n in main if n["offBeat"] <= main_len]
        prefix_len = random.random() * main_len * 0.9
        prefix = [n for n in main if n["onBeat"] < prefix_len]

        try:
            main_vec, _ = vocab.midi_to_tensor(
                main,
                accompany=combined_accomp,
                strict=True,
            )
            prefix_vec, _ = vocab.midi_to_tensor(
                prefix,
                accompany=combined_accomp,
                strict=True,
                rest_end=True,
            )
        except VocabError:
            continue

        main_vec = main_vec[:-1]  # remove end symbol

        fake_max_len = min(int(main_vec.shape[0] * 1.5), api.max_seq_len)
        fake_ys = np.zeros((api.inference_batch_size, fake_max_len, 9), dtype=np.int32)
        fake_ys[0, : prefix_vec.shape[0]] = prefix_vec
        dists = np.zeros(
            (api.inference_batch_size, fake_max_len, vocab.N), dtype=np.float32
        )
        logit_masks = np.zeros((api.inference_batch_size, vocab.N), dtype=np.float32)

        req = GenerationRequest(0.3, f"error || beat >= {main_len}")
        bias = json2logitbias(
            {("pitch < 1000" if is_drums else "pitch >= 1000"): -1000.0}
        )
        search_ctx = SearchCtx(
            request_ctxts=[
                api.setup_request_context(
                    0, req, 0.0, 0.0, bias, fake_ys, prefix_vec.shape[0]
                )
            ],
            encoder_input_ids=encoder_input_ids,
            encoder_attention_mask=encoder_attention_mask,
            start=prefix_vec.numpy(),
            logit_masks=logit_masks,
            ys=fake_ys,
            dists=dists,
            pos=prefix_vec.shape[0],
            recorder=ComposerAPIBasePerfRecorder(),
            start_time=time.time(),
            deadline=inf,
        )

        api.generate_ctx(search_ctx)

        # convert to midi to cut in the same way as the real example
        fake_full_clip = vocab.tensor_to_midi_using_embeds(
            fake_ys[0][: search_ctx.request_ctxts[0].stop_pos]
        )[-1]
        fake_bounded_clip = [n for n in fake_full_clip if n["offBeat"] <= main_len]
        try:
            fake_vec, _ = vocab.midi_to_tensor(
                fake_bounded_clip,
                accompany=combined_accomp,
                strict=True,
            )
            fake_vec = fake_vec[:-1]  # remove end symbol
        except VocabError:
            continue

        inp = torch.zeros(
            (2, max(main_vec.shape[0], fake_vec.shape[0]), 9), dtype=torch.long
        )

        if inp.shape[1] > model.params.max_seq_len:
            continue

        inp[:, :, 0] = vocab.pad.index
        inp[0, : main_vec.shape[0]] = main_vec
        inp[1, : fake_vec.shape[0]] = fake_vec

        emb_real, emb_fake = model.forward(
            inp,
            return_embeddings=True,
            encoder_input_ids=encoder_input_ids,
            encoder_attention_mask=encoder_attention_mask,
        ).unbind(0)

        if torch.allclose(emb_real, emb_fake):
            continue

        class_embeds.append([emb_real])
        nonclass_embeds.append([emb_fake])
        class_examples.append(main_vec)
        nonclass_examples.append(fake_vec)

        # print("real:")
        # print(vocab.repr_tensor(main_vec))
        # print("fake:")
        # print(vocab.repr_tensor(fake_vec))

        # print("---")

    torch.save(
        {
            "class_embeds": class_embeds,
            "nonclass_embeds": nonclass_embeds,
        },
        "data.pt",
    )
    dump_examples(vocab, class_examples, nonclass_examples)


def dump_examples(vocab, class_examples, nonclass_examples):
    clipss = [[] for i in range(4)]
    beat_accum = 0
    for i, c_nc in enumerate(zip(class_examples, nonclass_examples)):
        clips = [vocab.tensor_to_midi_using_embeds(c)[-1] for c in c_nc]
        example_beats = ceil(max(n["offBeat"] for clip in clips for n in clip))
        separated_clips = []
        for clip in clips:
            separated_clips.extend(list(separate_drums(clip)))
        wavtool_clips = [
            dict(
                midi2wavtool(c),
                **clip_props,
                loopEnd=example_beats,
                timelineStart=beat_accum,
                timelineEnd=beat_accum + example_beats,
                name=f"{n}",
            )
            for c, n in zip(
                separated_clips,
                [
                    "TH |",
                    "TD |",
                    "FH |",
                    "FD |",
                ],
            )
        ]

        if all(len(clip["notes"]) == 0 for clip in wavtool_clips):
            continue

        beat_accum += example_beats

        for i, clip in enumerate(wavtool_clips):
            if len(clip["notes"]) > 0:
                clipss[i].append(clip)
    obj = {
        "content": [{"clips": cs, "automationPoints": []} for cs in clipss],
        "length": beat_accum,
    }
    from pprint import pprint

    with open("out.json", "w") as f:
        json.dump(obj, f)


if __name__ == "__main__":
    from util import configure_logging

    configure_logging()
    with torch.no_grad():
        gen_critic_data(sys.argv[1])
