MCPcopy Create free account
hub / github.com/Ar-Ray-code/lingbot-depth-trt / build_tensorrt_engine

Function build_tensorrt_engine

tools/export_trt.py:370–422  ·  view source on GitHub ↗
(onnx_path: Path, engine_path: Path, args: argparse.Namespace)

Source from the content-addressed store, hash-verified

368
369
370def 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
425def run_tensorrt_engine(

Callers 1

mainFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected