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)