必须经过onnx中间格式,因为tensorrt不解析pytorch的torch.nn.module或.pt文件,仅支持onnx、uff(已弃用)等输入;onnx能保留模型结构、权重和算子语义,且pytorch导出支持度高,是当前最稳定兼容的桥梁。

不能直接用 torch.save 或 torch.load 转换;PyTorch 模型必须先导出为 ONNX,再由 TensorRT 解析生成引擎。
为什么必须经过 ONNX 中间格式
TensorRT 不解析 PyTorch 的 torch.nn.Module 或 .pt 文件,它只支持自己的序列化格式(.engine)或 ONNX、UFF(已弃用)、TensorFlow SavedModel 等输入。ONNX 是目前最稳定、兼容性最好的桥梁 —— 它能保留模型结构、权重和算子语义,且 PyTorch 的 torch.onnx.export 支持度高。
常见错误现象:RuntimeError: Exporting the operator 'aten::xxx' to ONNX opset version xxx is not supported,通常因模型用了动态控制流(如 Python if、for)、自定义算子或未 trace 的 torch.nn.functional 调用导致。
- 确保模型处于
eval()模式,关闭 dropout/batch norm 更新 - 输入 tensor 需固定 shape(如
torch.randn(1, 3, 224, 224)),避免 dynamic axes 除非必要 - 优先使用
opset_version=11或12(17对较新 PyTorch 更友好,但需 TensorRT ≥ 8.6) - 若含自定义算子,需提前注册 ONNX symbolic function 或改用等效原生算子
如何用 trtexec 命令行工具生成 engine(推荐快速验证)
trtexec 是 NVIDIA 提供的轻量级命令行工具,无需写 C++/Python API,适合调试导出是否成功、检查精度/性能瓶颈。
典型命令:
trtexec --onnx=model.onnx --saveEngine=model.engine --fp16 --workspace=2048 --shapes=input:1x3x224x224
-
--fp16启用半精度(绝大多数场景提速明显,精度损失可控);加--int8需额外校准,不建议首次尝试 -
--workspace=2048单位 MB,显存不足时会报Out of memory,可逐步调小(最低约 256) -
--shapes必须与 ONNX 中指定的 dynamic axes 匹配;若 ONNX 是静态 shape,此项可省略 - 失败时看日志末尾的
[E] Error行,常因算子不支持(如torch.nn.Softmax2d在旧 opset 中无对应 ONNX node)
用 Python API 构建 engine 并做推理(生产部署常用)
当需要动态 shape、多 batch 输入、或集成到现有 Python 服务中时,得用 tensorrt Python 包(需 pip install nvidia-tensorrt)。
关键步骤不是“加载模型”,而是“构建 builder → 创建 network → 解析 ONNX → 构建 engine”:
import tensorrt as trt logger = trt.Logger(trt.Logger.WARNING) builder = trt.Builder(logger) network = builder.create_network(1
-
EXPLICIT_BATCH是必须 flag(TensorRT ≥ 7.0),否则 parse 会静默失败 -
builder.max_batch_size已废弃,shape 信息全靠 network input tensor 的dynamic_range或opt_profile控制 - 若 ONNX 有多个输入,需在 parser 后手动设置每个 input 的
shape和dynamic_range - 构建耗时长(尤其大模型),建议保存
engine.serialize()到文件,后续直接runtime.deserialize_cuda_engine()
常见精度/性能异常的排查点
生成的 engine 推理结果和 PyTorch 不一致,或速度没提升,大概率卡在这些地方:
- ONNX 导出时用了
training=True或未设do_constant_folding=True(默认 True,但显式写出更稳妥) - TensorRT 版本太低(如 7.2 不支持
GroupNorm),查 官方支持矩阵 确认算子兼容性 - 输入预处理不一致:PyTorch 用
torchvision.transforms归一化,TRT engine 却直接喂原始像素,或 channel order(RGB vs BGR)搞反 - 没启用
builder.fp16_mode或config.set_flag(trt.BuilderFlag.FP16),仍以 FP32 运行 - GPU 上下文未绑定(多卡时
cudaSetDevice()缺失),engine 在默认卡上构建却在另一卡运行
真正麻烦的从来不是“怎么跑通”,而是“为什么输出差 0.02”或者“batch=4 时快,batch=1 时反而慢”——这些细节藏在 ONNX shape 推导、TensorRT profile 选择、以及 GPU kernel launch overhead 里,没法跳过验证环节。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











