import dataclasses
import json
import torch
from data import Vocab
from train_target import train_target_wants_text
from model import (
    ModelArgs,
    MusicalPositionEmbedTransformer as MusicalPositionEmbedTransformerForTraining,
)
from model_onnx import (
    MusicalPositionEmbedTransformer,
    MusicalPositionEmbedTransformerOneStep,
    MusicalPositionEmbedTransformerOneStepNoCross,
)
from tqdm import tqdm
from embedder import Embedder
import fire
from data_gen import DataGenerator, Dataset, Augmenter
import dataset_classes
import data_aug

MPET_LEN = 9


def load(vocab, config, model_file):
    state_dict = torch.load(model_file, map_location="cpu")
    params = dataclasses.replace(
        ModelArgs.from_config(vocab, config),
        enable_cross_attention=bool(config["enable_cross_attention"]),
        enable_flash=False,
    )
    print(params)
    model = MusicalPositionEmbedTransformer(vocab, params)
    one_step_class = (
        MusicalPositionEmbedTransformerOneStep
        if params.enable_cross_attention
        else MusicalPositionEmbedTransformerOneStepNoCross
    )
    model_one_step = one_step_class(vocab, params)
    model.load_state_dict(state_dict, strict=False)
    model_one_step.load_state_dict(state_dict, strict=False)
    model.eval()
    model_one_step.eval()
    return model, model_one_step


def load_training_model(vocab, config, model_file):
    state_dict = torch.load(model_file, map_location="cpu")
    train_target = config["train_target"]
    params = dataclasses.replace(
        ModelArgs.from_config(vocab, config),
        enable_cross_attention=train_target_wants_text(train_target),
        enable_flash=False,
    )
    model = MusicalPositionEmbedTransformerForTraining(
        vocab,
        params,
        encoder=(
            Embedder(load_pretrained_weights=False)
            if train_target_wants_text(train_target)
            else None
        ),
    )
    model.load_state_dict(state_dict)
    model.eval()
    return model


class OnnxTool:
    def __init__(self, config_path, model_path):
        self.model_path = model_path
        with open(config_path, "r") as f:
            self._config = json.load(f)

    def export(self, out_path, out_one_step_path):
        """
        Export the inference-optimized model to ONNX format.
        """
        print(f"Loading model from {self.model_path} for export")

        enable_cross_attention = bool(self._config["enable_cross_attention"])
        vocab = Vocab.from_config(self._config)
        model, model_one_step = load(
            vocab,
            self._config,
            self.model_path,
        )
        device = next(model.parameters()).device

        num_layers = model.params.n_layers
        dim = model.params.dim
        encoder_dim = model.params.cross_attention_embedding_dim
        nheads = model.params.n_heads
        head_dim = dim // nheads
        inference_batch_size = model.params.inference_batch_size

        # these don't really matter, just need something to test + trace the model with
        decoder_seq_len = 123
        encoder_seq_len = 32

        a_xks = torch.ones(
            decoder_seq_len,
            inference_batch_size,
            num_layers,
            nheads,
            head_dim,
            dtype=torch.float,
            device=device,
        )
        a_xvs = torch.ones(
            decoder_seq_len,
            inference_batch_size,
            num_layers,
            nheads,
            head_dim,
            dtype=torch.float,
            device=device,
        )
        c_xks = torch.ones(
            num_layers,
            inference_batch_size,
            nheads,
            head_dim,
            encoder_seq_len,
            dtype=torch.float,
            device=device,
        )
        c_xvs = torch.ones(
            num_layers,
            inference_batch_size,
            nheads,
            encoder_seq_len,
            head_dim,
            dtype=torch.float,
            device=device,
        )
        x = torch.ones(
            inference_batch_size,
            decoder_seq_len,
            MPET_LEN,
            dtype=torch.long,
            device=device,
        )
        xb = torch.ones(
            inference_batch_size, 1, MPET_LEN, dtype=torch.long, device=device
        )
        e = torch.ones(
            inference_batch_size,
            encoder_seq_len,
            encoder_dim,
            dtype=torch.float,
            device=device,
        )
        v = torch.ones(
            inference_batch_size, encoder_seq_len, dtype=torch.bool, device=device
        )

        # check the model actually runs first
        with torch.no_grad():
            model(x=x, encoder_out=e, encoder_valid=v)
            if enable_cross_attention:
                model_one_step(
                    x=xb,
                    a_xks=a_xks,
                    a_xvs=a_xvs,
                    c_xks=c_xks,
                    c_xvs=c_xvs,
                    encoder_valid=v,
                )
            else:
                model_one_step(x=xb, a_xks=a_xks, a_xvs=a_xvs)

        print(f"Exporting model to {out_path} and {out_one_step_path}")

        torch.onnx.export(
            model,
            (x, e, v),
            out_path,
            verbose=False,
            input_names=["x", "encoder_out", "encoder_valid"],
            output_names=[
                "y",
                "a_xks",
                "a_xvs",
                *(["c_xks", "c_xvs"] if enable_cross_attention else []),
                "embeddings",
            ],
            dynamic_axes={"x": [0, 1], "encoder_out": [0, 1], "encoder_valid": [0, 1]},
            do_constant_folding=True,
            opset_version=16,
        )
        torch.onnx.export(
            model_one_step,
            (
                (xb, a_xks, a_xvs, c_xks, c_xvs, v)
                if enable_cross_attention
                else (xb, a_xks, a_xvs)
            ),
            out_one_step_path,
            verbose=False,
            input_names=[
                "x",
                "a_xks",
                "a_xvs",
                *(
                    ["c_xks", "c_xvs", "encoder_valid"]
                    if enable_cross_attention
                    else []
                ),
            ],
            output_names=["y", "a_xk_news", "a_xv_news"],
            dynamic_axes=dict(
                {"x": [0], "a_xks": [0, 1], "a_xvs": [0, 1]},
                **(
                    {"c_xks": [1, 4], "c_xvs": [1, 3], "encoder_valid": [0, 1]}
                    if enable_cross_attention
                    else {}
                ),
            ),
            do_constant_folding=True,
            opset_version=16,
        )

    def export_calibration_data(self, out_path, batch_size=3, batches=4096):
        """
        Export calibration data for the inference-optimized model.
        Data will be a list of dicts, each dict containing

        tgt: the target sequences, as a BxLx9 tensor
        encoded: the encoder output, as a BxMx512 tensor
        encoded_valid: the encoder valid mask, as a BxM tensor
        ref_y: the reference output, as a BxLxN tensor
        """

        vocab = Vocab.from_config(self._config)
        model = load_training_model(vocab, self._config, self.model_path)
        seq_len = int(self._config["seq_len"])
        seq_len_min = int(self._config["seq_len_min"])
        train_target = self._config["train_target"]
        train_split_idx = int(self._config["train_split_idx"])

        dataset = Dataset.dynamic_from_config(self._config, "dataset_class")
        auger = Augmenter.dynamic_from_config(vocab, self._config, "augmenter_class")

        data_gen = DataGenerator(
            dataset,
            vocab,
            split=train_split_idx,
            ranksize=(0, 1),
            batch_size=batch_size,
            seq_len=seq_len,
            seq_len_min=seq_len_min,
            seq_len_max=None,
            train_target=train_target,
            parallelism=1,
            to_device="cpu",
            text_tokenize=model.encoder.tokenize,
            example_continuous=False,
            pack_batch=False,
            augmenter=auger,
        )

        calibs = []

        for batch in tqdm(
            data_gen.generate(
                force_batches=batches,
                order="random",
            ),
            total=batches,
        ):
            with torch.no_grad():
                tgt = batch.tgt[
                    :,
                    : torch.max(
                        torch.nonzero(batch.tgt[:, :, 0] != model.params.padding_idx)
                    ).item(),
                ]
                encoded = model.encoder(
                    batch.encoder_input_ids, batch.encoder_attention_mask
                )
                ref_y = model.forward(
                    tgt,
                    start_pos=0,
                    encoder_input_ids=batch.encoder_input_ids,
                    encoder_attention_mask=batch.encoder_attention_mask,
                )
            calibs.append(
                {
                    "tgt": tgt,
                    "encoded": encoded,
                    "encoded_valid": batch.encoder_attention_mask.to(dtype=torch.bool),
                    "ref_y": ref_y,
                }
            )

        torch.save(calibs, out_path)


if __name__ == "__main__":
    fire.Fire(OnnxTool)
