from trt_execution import TRTModel
import tensorrt as trt
import torch
from onnx_export import load
import json
from data import Vocab
import time

model_prefix = "model_2024-04-21"
config_path = f"{model_prefix}/config.json"
model_path = f"{model_prefix}/model.pt"


def timing_test_trt():
    with open(config_path, "r") as f:
        config = json.load(f)

    vocab = Vocab.from_config(config)
    ref_model, ref_model_one_step = load(vocab, config, model_path)
    ref_model.eval()
    ref_model_one_step.eval()
    ref_model = ref_model.cuda()
    ref_model_one_step = ref_model_one_step.cuda()

    TRT_LOGGER = trt.Logger(trt.Logger.VERBOSE)
    runtime = trt.Runtime(TRT_LOGGER)

    stream = torch.cuda.Stream()

    model = TRTModel(runtime, f"{model_prefix}/model.trt")
    context = model.create_execution_context(stream=stream)

    model_one_step = TRTModel(runtime, f"{model_prefix}/model_one_step.trt")
    context_one_step = model_one_step.create_execution_context(stream=stream)

    eval_data = torch.load(f"{model_prefix}/calibration.pt", map_location="cuda")

    with torch.no_grad(), torch.cuda.stream(stream):
        context.set_optimization_profile_index(
            model.get_profile_index_for_dim_constraint(
                "x", 0, eval_data[0]["tgt"].shape[0]
            )
        )
        context_one_step.set_optimization_profile_index(
            model_one_step.get_profile_index_for_dim_constraint(
                "x", 0, eval_data[0]["tgt"].shape[0]
            )
        )
        for obj in eval_data[:100]:
            tgt = obj["tgt"]
            encoded = obj["encoded"]
            encoded_valid = obj["encoded_valid"]
            ref_y = obj["ref_y"]

            start_point = (
                torch.min(
                    torch.count_nonzero(tgt[:, :, 0] != vocab.pad.index, dim=-1)
                ).item()
                - 3
            )

            t0 = time.time()
            ref_out = ref_model(tgt, encoded, encoded_valid)
            ref_dt = time.time() - t0

            delta = ref_out[0] - ref_y
            delta[tgt[:, :, 0] == 0] = 0

            t0 = time.time()
            out = context.eval(
                x=tgt.type(torch.int32),
                encoder_out=encoded,
                encoder_valid=encoded_valid,
            )
            dt = time.time() - t0

            delta_out = out["y"] - ref_y
            delta_out[tgt[:, :, 0] == 0] = 0

            # test one step ref matches full ref
            a_xks0 = ref_out[1].transpose(0, 1).contiguous()
            a_xvs0 = ref_out[2].transpose(0, 1).contiguous()
            c_xks = ref_out[3].clone()
            c_xvs = ref_out[4].clone()
            y0 = ref_out[0].clone()

            t0 = time.time()
            y1, a_xks1, a_xvs1 = ref_model_one_step(
                tgt[:, start_point - 1 : start_point],
                a_xks0[: start_point - 1],
                a_xvs0[: start_point - 1],
                c_xks,
                c_xvs,
                encoded_valid,
            )
            ref_step_dt = time.time() - t0

            delta_step_ref = y0[:, start_point - 1] - y1[:, 0]
            delta_cache_k_ref = a_xks0[start_point - 1] - a_xks1
            delta_cache_v_ref = a_xvs0[start_point - 1] - a_xvs1

            # test one step matches full ref
            out = context.eval(
                x=tgt[:, :start_point].type(torch.int32),
                encoder_out=encoded,
                encoder_valid=encoded_valid,
            )
            a_xks0 = out["a_xks"].transpose(0, 1).contiguous()
            a_xvs0 = out["a_xvs"].transpose(0, 1).contiguous()
            c_xks = out.get("c_xks", None).clone()
            c_xvs = out.get("c_xvs", None).clone()
            y0 = out["y"].clone()

            t0 = time.time()
            out_step = context_one_step.eval(
                x=tgt[:, start_point - 1 : start_point].type(torch.int32),
                a_xks=a_xks0[:-1],
                a_xvs=a_xvs0[:-1],
                c_xks=c_xks,
                c_xvs=c_xvs,
                encoder_valid=encoded_valid,
            )
            step_dt = time.time() - t0

            y1 = out_step["y"]
            a_xks1 = out_step["a_xk_news"]
            a_xvs1 = out_step["a_xv_news"]

            delta_step = y0[:, -1] - y1[:, 0]

            delta_cache_k = a_xks0[-1] - a_xks1
            delta_cache_v = a_xvs0[-1] - a_xvs1

            print(
                f"{ref_out[0].abs().max().item():.2f}",
                f"{delta.abs().max().item():.4f}",
                f"({ref_dt * 1000:.2f} ms)",
                f"{delta_out.abs().max().item():.4f}",
                f"({dt * 1000:.2f} ms)",
                f"r[{delta_step_ref.abs().max().item():.4f}",
                f"{delta_cache_k_ref.abs().max().item():.4f}",
                f"{delta_cache_v_ref.abs().max().item():.4f}]",
                f"({ref_step_dt * 1000:.2f} ms)",
                f"t[{delta_step.abs().max().item():.4f}",
                f"{delta_cache_k.abs().max().item():.4f}",
                f"{delta_cache_v.abs().max().item():.4f}]",
                f"({step_dt * 1000:.2f} ms)",
            )


if __name__ == "__main__":
    timing_test_trt()
