tensorflow模型需通过tf.experimental.tensorrt.converter在tf上下文中原位优化,而非直接编译;必须指定静态shape、调用convert()后执行build()触发引擎构建,int8校准需真实数据,验证需查日志、耗时与显存。

TensorFlow模型不能直接编译为TensorRT引擎
TensorRT不接受原生TensorFlow SavedModel或Frozen Graph作为输入。必须先将模型转换为ONNX或UFF(已弃用),再经TensorRT解析器导入——但更主流、稳定的做法是:**通过TensorFlow-TensorRT(TF-TRT)集成,在TensorFlow图内做原位优化和编译**。
这意味着你不是“导出再编译”,而是用tf.experimental.tensorrt.Converter在TensorFlow上下文中封装并优化计算图,最终生成可执行的ConcreteFunction,其底层已绑定TensorRT引擎。
使用tf.experimental.tensorrt.Converter进行动态编译
这是目前最可靠、与TensorFlow 2.x兼容的方式,适用于SavedModel格式模型。关键点在于:它只优化支持的子图(subgraph),其余部分仍走TF CPU/GPU路径,属于混合执行模式。
- 确保安装匹配版本:
tensorrt>=8.6+tensorflow>=2.10(TF 2.15+已内置TF-TRT,无需额外pip install) - 模型必须是静态shape:输入张量需指定
batch_size和具体height/width,不能含None维度(如[None, 224, 224, 3]不行,得用[1, 224, 224, 3]) - 调用
converter.convert()后,必须调用converter.build(input_fn)触发实际引擎构建——仅convert()不会生成TRT引擎,只会标记可优化节点 -
input_fn必须返回一个tf.data.Dataset或可迭代对象,每次yield一个符合签名的tf.Tensor;若跳过build(),推理时会fallback到原TF图
示例片段:
快速生成专业的 Python 脚本和应用代码。一键创建完整项目结构,支持CLI、API、爬虫、Bot、Django等多种项目类型,包含完整的项目结构、配置文件、依赖管理、测试、README和文档。
import tensorflow as tf
converter = tf.experimental.tensorrt.Converter(
input_saved_model_dir="my_model",
precision_mode="FP16", # 或"INT8"(需calibration)
maximum_batch_size=16
)
converter.convert()
def input_fn():
for _ in range(10):
yield tf.random.normal([1, 224, 224, 3])
converter.build(input_fn=input_fn) # ← 这步不可省!
converter.save("trt_saved_model")
INT8校准必须显式提供校准数据集
启用precision_mode="INT8"时,TensorRT需要真实输入分布来确定激活值范围。TF-TRT不会自动采样,你必须传入有代表性的校准样本(通常200–1000张图),且input_fn必须能重复遍历(推荐用tf.data.Dataset.cache())。
- 校准数据应覆盖模型典型输入:尺寸、亮度、对比度、类别分布都要贴近线上流量
- 不要用
tf.random.*生成假数据——TRT会据此算出错误的scale,导致精度崩塌 - 若校准后精度不达标,可尝试增大
use_calibration=True(默认True)并调整minimum_segment_size(默认3),避免小op段被跳过
验证是否真正启用了TensorRT引擎
最直接的办法是看日志和运行时行为,而非只信“保存成功”:
- 启用
TF_CPP_MIN_LOG_LEVEL=0,启动时搜"TensorRT optimization enabled"和"Created TensorRT engine" - 推理耗时对比:同一GPU上,TRT优化后的
ConcreteFunction应比原SavedModel快1.5–3倍(视模型而定),若无差异,大概率没生效 - 检查生成目录:
trt_saved_model/variables/为空,trt_saved_model/saved_model.pb体积显著变小(TRT引擎序列化进assets/或variables/下二进制文件) - 用
nvidia-smi观察GPU内存占用:TRT引擎常驻显存,首次推理后显存不回落,而纯TF模型可能波动更大
真正麻烦的是混合执行边界——比如模型里嵌了tf.py_function或自定义OP,整个子图会被隔离出TRT流程。这种细节不会报错,但性能就卡在那儿了。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!










