import tensorrt as trt
import json
import fire
from pathlib import Path
from contextlib import contextmanager

EXPLICIT_BATCH = 1 << (int)(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH)
MPET_LEN = 9


@contextmanager
def trt_timing_cache(builder_config, cache_file_path):
    cache_path = Path(cache_file_path)
    cache_data = b""
    if cache_path.exists():
        with cache_path.open("rb") as cache_file:
            cache_data = cache_file.read()

    timing_cache = builder_config.create_timing_cache(cache_data)
    builder_config.set_timing_cache(timing_cache, ignore_mismatch=False)
    yield
    with cache_path.open("wb") as cache_file:
        cache_file.write(timing_cache.serialize())


class TrtTool:
    def __init__(self, config_path, timing_cache_path=None):
        with open(config_path, "r") as f:
            self._config = json.load(f)
        self._dim = int(self._config["dim"])
        self._encoder_dim = 1472  # TODO fix
        self._layers = int(self._config["layers"])
        self._nheads = int(self._config["heads"])
        self._head_dim = self._dim // self._nheads
        self._enable_cross_attention = bool(self._config["enable_cross_attention"])
        inference_config = self._config["inference"]
        self._batch_size = int(inference_config["batch_size"])
        self._build_one_step_batch_sizes = list(
            sorted(set(inference_config["build_one_step_batch_sizes"]))
        )
        assert all(int(x) and x > 0 for x in self._build_one_step_batch_sizes)
        if self._batch_size not in self._build_one_step_batch_sizes:
            print(
                f"Warning: batch size {self._batch_size} not in build_one_step_batch_sizes {self._build_one_step_batch_sizes}"
            )
        self._opt_decoder_seq_len = int(inference_config["opt_decoder_seq_len"])
        self._max_decoder_seq_len = int(inference_config["max_decoder_seq_len"])
        self._opt_encoder_seq_len = int(inference_config["opt_encoder_seq_len"])
        self._max_encoder_seq_len = int(inference_config["max_encoder_seq_len"])

        self._logger = trt.Logger(trt.Logger.VERBOSE)
        self._builder = trt.Builder(self._logger)
        self._timing_cache_path = timing_cache_path

    def _config_common(self, config):
        config.set_flag(trt.BuilderFlag.TF32)
        config.builder_optimization_level = 4

    def _setup(self):
        network = self._builder.create_network(EXPLICIT_BATCH)
        parser = trt.OnnxParser(network, self._logger)
        config = self._builder.create_builder_config()
        self._config_common(config)
        return network, parser, config

    def _load_onnx(self, parser, model_path):
        with open(model_path, "rb") as model:
            if not parser.parse(model.read()):
                for error in range(parser.num_errors):
                    print(parser.get_error(error))
                raise RuntimeError(f"Failed to parse ONNX model {model_path}")

    def _write_serialized_engine(self, serialized_engine, out_path):
        with open(out_path, "wb") as f:
            f.write(serialized_engine)

    @contextmanager
    def _maybe_timing_cache(self, config):
        if self._timing_cache_path is not None:
            with trt_timing_cache(config, self._timing_cache_path):
                yield
        else:
            yield

    def build_trt(self, model_path, out_path):
        network, parser, config = self._setup()

        with self._maybe_timing_cache(config):
            for static_batch_size in [1, 2]:  # without and with CFG
                profile_static = self._builder.create_optimization_profile()
                profile_static.set_shape(
                    "x",
                    (static_batch_size, 1, MPET_LEN),
                    (static_batch_size, self._opt_decoder_seq_len, MPET_LEN),
                    (static_batch_size, self._max_decoder_seq_len, MPET_LEN),
                )
                profile_static.set_shape(
                    "encoder_out",
                    (static_batch_size, 1, self._encoder_dim),
                    (static_batch_size, self._opt_encoder_seq_len, self._encoder_dim),
                    (static_batch_size, self._max_encoder_seq_len, self._encoder_dim),
                )
                profile_static.set_shape(
                    "encoder_valid",
                    (static_batch_size, 1),
                    (static_batch_size, self._opt_encoder_seq_len),
                    (static_batch_size, self._max_encoder_seq_len),
                )
                config.add_optimization_profile(profile_static)

            max_osbs = max(self._build_one_step_batch_sizes)
            profile_multi_batch = self._builder.create_optimization_profile()
            profile_multi_batch.set_shape(
                "x",
                (1, 1, MPET_LEN),
                (max_osbs, self._opt_decoder_seq_len, MPET_LEN),
                (max_osbs, self._max_decoder_seq_len, MPET_LEN),
            )
            profile_multi_batch.set_shape(
                "encoder_out",
                (1, 1, self._encoder_dim),
                (max_osbs, self._opt_encoder_seq_len, self._encoder_dim),
                (max_osbs, self._max_encoder_seq_len, self._encoder_dim),
            )
            profile_multi_batch.set_shape(
                "encoder_valid",
                (1, 1),
                (max_osbs, self._opt_encoder_seq_len),
                (max_osbs, self._max_encoder_seq_len),
            )
            config.add_optimization_profile(profile_multi_batch)

            self._load_onnx(parser, model_path)

            serialized_engine = self._builder.build_serialized_network(network, config)
            self._write_serialized_engine(serialized_engine, out_path)

    def build_trt_one_step(self, model_path, out_path):
        network, parser, config = self._setup()

        with self._maybe_timing_cache(config):
            print(
                f"Building one-step model with {self._layers} layers for batch sizes {self._build_one_step_batch_sizes}"
            )

            for batch_size in self._build_one_step_batch_sizes:
                profile = self._builder.create_optimization_profile()
                profile.set_shape(
                    "x",
                    (batch_size, 1, MPET_LEN),
                    (batch_size, 1, MPET_LEN),
                    (batch_size, 1, MPET_LEN),
                )
                for kv in ("a_xks", "a_xvs"):
                    profile.set_shape(
                        kv,
                        (1, batch_size, self._layers, self._nheads, self._head_dim),
                        (
                            self._opt_decoder_seq_len,
                            batch_size,
                            self._layers,
                            self._nheads,
                            self._head_dim,
                        ),
                        (
                            self._max_decoder_seq_len,
                            batch_size,
                            self._layers,
                            self._nheads,
                            self._head_dim,
                        ),
                    )
                if self._enable_cross_attention:
                    profile.set_shape(
                        "c_xks",
                        (self._layers, batch_size, self._nheads, self._head_dim, 1),
                        (
                            self._layers,
                            batch_size,
                            self._nheads,
                            self._head_dim,
                            self._opt_encoder_seq_len,
                        ),
                        (
                            self._layers,
                            batch_size,
                            self._nheads,
                            self._head_dim,
                            self._max_encoder_seq_len,
                        ),
                    )
                    profile.set_shape(
                        "c_xvs",
                        (self._layers, batch_size, self._nheads, 1, self._head_dim),
                        (
                            self._layers,
                            batch_size,
                            self._nheads,
                            self._opt_encoder_seq_len,
                            self._head_dim,
                        ),
                        (
                            self._layers,
                            batch_size,
                            self._nheads,
                            self._max_encoder_seq_len,
                            self._head_dim,
                        ),
                    )
                    profile.set_shape(
                        "encoder_valid",
                        (batch_size, 1),
                        (batch_size, self._opt_encoder_seq_len),
                        (batch_size, self._max_encoder_seq_len),
                    )
                config.add_optimization_profile(profile)

            self._load_onnx(parser, model_path)

            serialized_engine = self._builder.build_serialized_network(network, config)
            self._write_serialized_engine(serialized_engine, out_path)


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