| 368 | |
| 369 | |
| 370 | def build_tensorrt_engine(onnx_path: Path, engine_path: Path, args: argparse.Namespace) -> dict[str, Any]: |
| 371 | engine_path.parent.mkdir(parents=True, exist_ok=True) |
| 372 | logger = trt.Logger(trt.Logger.INFO if args.verbose_trt else trt.Logger.WARNING) |
| 373 | builder = trt.Builder(logger) |
| 374 | network_flags = 0 |
| 375 | if hasattr(trt.NetworkDefinitionCreationFlag, "STRONGLY_TYPED"): |
| 376 | network_flags |= 1 << int(trt.NetworkDefinitionCreationFlag.STRONGLY_TYPED) |
| 377 | |
| 378 | start = time.perf_counter() |
| 379 | with builder.create_network(network_flags) as network, trt.OnnxParser(network, logger) as parser: |
| 380 | if not parser.parse(onnx_path.read_bytes()): |
| 381 | errors = [str(parser.get_error(i)) for i in range(parser.num_errors)] |
| 382 | return {"ok": False, "stage": "parse", "errors": errors} |
| 383 | |
| 384 | config = builder.create_builder_config() |
| 385 | config.set_memory_pool_limit(trt.MemoryPoolType.WORKSPACE, args.workspace_gb * (1 << 30)) |
| 386 | config.set_memory_pool_limit(trt.MemoryPoolType.TACTIC_DRAM, args.tactic_gb * (1 << 30)) |
| 387 | if hasattr(trt.BuilderFlag, "TF32") and args.tf32: |
| 388 | config.set_flag(trt.BuilderFlag.TF32) |
| 389 | |
| 390 | serialized_engine = builder.build_serialized_network(network, config) |
| 391 | elapsed = time.perf_counter() - start |
| 392 | if serialized_engine is None: |
| 393 | return { |
| 394 | "ok": False, |
| 395 | "stage": "build", |
| 396 | "build_s": elapsed, |
| 397 | "inputs": [ |
| 398 | {"name": network.get_input(i).name, "shape": list(network.get_input(i).shape)} |
| 399 | for i in range(network.num_inputs) |
| 400 | ], |
| 401 | "outputs": [ |
| 402 | {"name": network.get_output(i).name, "shape": list(network.get_output(i).shape)} |
| 403 | for i in range(network.num_outputs) |
| 404 | ], |
| 405 | "layers": network.num_layers, |
| 406 | } |
| 407 | engine_path.write_bytes(serialized_engine) |
| 408 | return { |
| 409 | "ok": True, |
| 410 | "path": str(engine_path), |
| 411 | "bytes": engine_path.stat().st_size, |
| 412 | "build_s": elapsed, |
| 413 | "inputs": [ |
| 414 | {"name": network.get_input(i).name, "shape": list(network.get_input(i).shape)} |
| 415 | for i in range(network.num_inputs) |
| 416 | ], |
| 417 | "outputs": [ |
| 418 | {"name": network.get_output(i).name, "shape": list(network.get_output(i).shape)} |
| 419 | for i in range(network.num_outputs) |
| 420 | ], |
| 421 | "layers": network.num_layers, |
| 422 | } |
| 423 | |
| 424 | |
| 425 | def run_tensorrt_engine( |