"""Build a GPU-specific TensorRT FP16 engine from the raw ONNX export.""" from __future__ import annotations import argparse import hashlib import json from pathlib import Path import time import torch # Loads the CUDA libraries supplied by the pinned PyTorch package. import tensorrt as trt def build(args: argparse.Namespace) -> None: onnx_path = Path(args.onnx).resolve() engine_path = Path(args.output).resolve() engine_path.parent.mkdir(parents=True, exist_ok=True) logger = trt.Logger(trt.Logger.INFO if args.verbose else trt.Logger.WARNING) trt.init_libnvinfer_plugins(logger, "") builder = trt.Builder(logger) # TensorRT 10 supports weak typing: choose fast FP16 tactics while native # normalization accumulates in FP32. TRT 11 requires explicit ONNX types. if not trt.__version__.startswith("10."): raise RuntimeError(f"Expected TensorRT 10.x, found {trt.__version__}") network = builder.create_network(0) parser = trt.OnnxParser(network, logger) if not parser.parse_from_file(str(onnx_path)): errors = "\n".join(str(parser.get_error(i)) for i in range(parser.num_errors)) raise RuntimeError(f"TensorRT could not parse {onnx_path}:\n{errors}") config = builder.create_builder_config() config.set_flag(trt.BuilderFlag.FP16) config.clear_flag(trt.BuilderFlag.TF32) config.set_memory_pool_limit(trt.MemoryPoolType.WORKSPACE, args.workspace_gib << 30) config.builder_optimization_level = args.optimization_level # A single stream is easier to capture in CUDA graphs and avoids concurrent # auxiliary work competing with batch-one latency. config.max_aux_streams = 0 config.profiling_verbosity = trt.ProfilingVerbosity.DETAILED normalization_layers = [] for index in range(network.num_layers): layer = network.get_layer(index) if layer.type == trt.LayerType.NORMALIZATION: # Native normalization defaults to FP32 accumulation independently # of BuilderFlag.FP16. Keep that default and allow FP16 I/O/fusions. normalization_layers.append(layer.name) profile = builder.create_optimization_profile() minimum = (1, 2) optimum = (args.opt_batch, args.opt_sequence) maximum = (args.max_batch, args.max_sequence) for index in range(network.num_inputs): inp = network.get_input(index) if len(inp.shape) != 2: raise ValueError(f"Expected rank-two NER input; {inp.name}: {inp.shape}") profile.set_shape(inp.name, minimum, optimum, maximum) config.add_optimization_profile(profile) cache_path = engine_path.with_suffix(".timing.cache") cache_bytes = cache_path.read_bytes() if cache_path.exists() else b"" timing_cache = config.create_timing_cache(cache_bytes) config.set_timing_cache(timing_cache, ignore_mismatch=False) started = time.perf_counter() serialized = builder.build_serialized_network(network, config) if serialized is None: raise RuntimeError("TensorRT engine build failed; inspect builder messages above") engine_path.write_bytes(bytes(serialized)) cache_path.write_bytes(bytes(config.get_timing_cache().serialize())) metadata = { "tensorrt": trt.__version__, "torch": torch.__version__, "cuda": torch.version.cuda, "gpu": torch.cuda.get_device_name(), "compute_capability": list(torch.cuda.get_device_capability()), "onnx": str(onnx_path), "onnx_sha256": hashlib.file_digest(onnx_path.open("rb"), "sha256").hexdigest(), "precision": "FP16 with FP32 normalization accumulation; TF32 disabled", "normalization_layers": len(normalization_layers), "profile": {"min": minimum, "opt": optimum, "max": maximum}, "workspace_gib": args.workspace_gib, "optimization_level": args.optimization_level, "build_seconds": time.perf_counter() - started, "engine_bytes": engine_path.stat().st_size, } engine_path.with_suffix(".json").write_text(json.dumps(metadata, indent=2) + "\n") print(json.dumps(metadata, indent=2), flush=True) if __name__ == "__main__": parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--onnx", default="artifacts/model.onnx") parser.add_argument("--output", default="artifacts/model.fp16.engine") parser.add_argument("--opt-batch", type=int, default=1) parser.add_argument("--opt-sequence", type=int, default=128) parser.add_argument("--max-batch", type=int, default=32) parser.add_argument("--max-sequence", type=int, default=512) parser.add_argument("--workspace-gib", type=int, default=4) parser.add_argument("--optimization-level", type=int, default=3) parser.add_argument("--verbose", action="store_true") build(parser.parse_args())