TensorRT 弱类型到强类型迁移自动化:trt-strong-typing-migration 辅助脚本实战指南
【免费下载链接】TensorRTNVIDIA® TensorRT™ is an SDK for high-performance deep learning inference on NVIDIA GPUs. This repository contains the open source components of TensorRT.项目地址: https://gitcode.com/GitHub_Trending/tens/TensorRT
本篇技术指南围绕 TensorRT 开源仓库(tensorrt-oss)中.agents/skills/trt-strong-typing-migration/scripts/目录下的两个辅助脚本展开:基于 Python AST 的自动重写器migrate.py与端到端验证脚本verify.sh。它们用于把 TensorRT 10.12 起被弃用、11.0 被移除的"弱类型(weak typing)"Python 构建代码,自动化改写为强类型(strong typing)形式。读完本文,你将掌握这两个脚本的完整用法、AST 重写的底层原理、它们对哪些代码动手与哪些代码刻意放过,以及如何在真实代码库上安全地跑完"dry-run 审查 → 原地改写 → 验证 → 重建"的迁移流程。
为什么需要自动化迁移脚本
TensorRT 11.x 中强类型(strong typing)是唯一的构建模式:网络中每个张量的数据类型都由网络自身的类型推断规则、输入类型与算子注解决定,不再允许通过构建配置"提示"精度。弱类型在 10.12 被标记为弃用并在 11.0 移除。从 include/NvInfer.h 的枚举定义可以看到,NetworkDefinitionCreationFlag::kSTRONGLY_TYPED在 11.0 已被标记为 deprecated 并恒为默认生效(值为 0,保留仅为 API 兼容):
enum class NetworkDefinitionCreationFlag : int32_t { //! Mark the network to be strongly typed. ... Deprecated in TensorRT 11.0. //! Strongly typed mode is always enabled. This flag is retained for API compatibility but is ignored. kSTRONGLY_TYPED TRT_DEPRECATED_ENUM = 0, ... };迁移本身高度机械化:把create_network的建网标志换成STRONGLY_TYPED、删除set_flag(FP16/INT8/...)这类精度提示、删除layer.precision与set_output_type赋值。这些替换规则明确、重复度高,正适合交给脚本自动化处理——这正是migrate.py存在的意义。完整的迁移背景、三步迁移路径(Python 构建器 / trtexec / C++ 构建器)与 ModelOpt AutoCast 前置步骤,参见技能主文档 SKILL.md,本文聚焦于其中配套的自动化工具。
migrate.py:用法速览
migrate.py是一个面向 Python TensorRT 构建代码的 AST 重写器。它识别弱类型构建模式(create_network+set_flag(BuilderFlag.FP16/...)),并将其改写为强类型形式。脚本位于 .agents/skills/trt-strong-typing-migration/scripts/migrate.py。
三种典型调用方式:
# 仅展示将要发生的变化(dry-run):打印 unified diff;有待迁移变更时退出码为 1,无需变更时退出码为 0 python3 migrate.py path/to/build.py # 原地重写文件: python3 migrate.py path/to/build.py --write # 递归处理目录树下的所有 .py 文件: python3 migrate.py path/to/project/ --write关键行为说明:
- 默认 dry-run:只打印统一 diff,不落盘。源码中 main() 的退出码设计借鉴了
black --check:dry-run 模式下一旦检测到待变更内容即返回 1,方便接入 CI 或作为"是否已迁移"的哨兵;--write应用变更后返回 0。 - 目录递归:通过
_iter_files对传入的目录执行rglob("*.py"),仅处理.py后缀文件,其余文件自动跳过。 - AST 而非正则:脚本基于
ast.NodeTransformer理解调用形态而非纯文本匹配,因此能正确处理经别名导入访问的set_flag、create_network中的多标志构造,以及逐层(per-layer)的precision/set_output_type赋值。无关逻辑与 docstring 保持原样。 - 格式与注释的代价:由于重写经
ast.unparse往返,普通的#注释与原始格式不会保留。务必在--write前审查 dry-run diff,之后按需重新补充注释。
migrate.py 的 AST 重写原理
要理解脚本的行为边界,需要看它内部如何判定"该改什么"。从源码看,其核心是三类精确匹配:
1. 精度标志白名单与网络标志映射
migrate.py 第 31-36 行定义了两个关键集合:
# Precision-hint BuilderFlag attribute names that must be removed. NOTE: TF32 is # deliberately NOT here — kTF32 is kept in TRT 11 (orthogonal to typing), like REFIT. PRECISION_FLAGS = frozenset({"FP16", "BF16", "INT8", "FP8"}) # NetworkDefinitionCreationFlag attribute names that map to STRONGLY_TYPED. WEAK_NETWORK_FLAGS = frozenset({"EXPLICIT_BATCH"})注意TF32被刻意排除在PRECISION_FLAGS之外——kTF32在 TRT 11 中仍然保留,与类型化正交,属于需要保留下来的标志。
2. 属性链匹配:不关心 import 别名
_is_attr辅助函数从node.attr -> node.value.attr -> ...反向走属性链,只校验尾部链条而不校验最前端的模块/对象名。这意味着trt.BuilderFlag.FP16、tensorrt.BuilderFlag.FP16乃至任何as别名导入(如import tensorrt as t; t.BuilderFlag.FP16)都能被同一套逻辑命中——这正是"grep 字符串匹配做不到、AST 匹配做得到"的地方。
3. 三组 NodeTransformer 访问器
visit_Call:拦截所有create_network(...)调用,将单个位置参数替换为1 << int(trt.NetworkDefinitionCreationFlag.STRONGLY_TYPED)。它区分了几种形态:- 无位置参数但带
flags=关键字:直接替换关键字值; - 无任何参数的空调用
create_network()(默认即弱类型):追加强类型参数; - 已有
STRONGLY_TYPED参数:幂等跳过(不改写已是强类型的文件); - 无法识别的参数形态:记入
skipped列表并提示"人工审查"。
- 无位置参数但带
visit_Expr:删除作为表达式语句出现的config.set_flag(trt.BuilderFlag.{FP16,BF16,INT8,FP8})调用(set_flag命中PRECISION_FLAGS即整条语句删除),并无条件删除layer.set_output_type(...)调用。visit_Assign:删除something.precision = ...形式的赋值语句。
4. 死代码清理
visit_If会顺带清理if builder.platform_has_fast_fp16 / platform_has_fast_int8 / platform_has_fast_bf16 / platform_has_fast_fp8:这类条件块:当set_flag是块内唯一语句、删除后整个if变空且没有else分支时,整块if一并删除;若else分支存在则用else块替换。
5. 变更判定基于"变换计数"而非文本差异
一个容易踩的细节:_process_file判定"文件是否需要迁移"用的是重写器四个计数器(rewrote_create_network + removed_flag_calls + removed_precision_assigns + removed_set_output_type)之和是否大于 0,而不是比较新旧文本。因为ast.unparse每次往返都会重排格式、剥离注释,如果按文本相等性判断,一个本不需要迁移的文件也会被误判为"已变更"并被剥掉注释。
6. 保守原则:宁可不改,不可乱改
凡是无法确认的形态,脚本一律跳过并在 stderr 打印[skip]/[note]信息,交由人工处理。例如create_network同时出现其他未知关键字、getattr(trt.BuilderFlag, 'FP16')这类动态访问等,均不在重写范围内。
migrate.py 改了什么、没改什么
What it changes
| 迁移前 | 迁移后 |
|---|---|
builder.create_network(0) | builder.create_network(1 << int(trt.NetworkDefinitionCreationFlag.STRONGLY_TYPED)) |
builder.create_network(1 << int(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH)) | builder.create_network(1 << int(trt.NetworkDefinitionCreationFlag.STRONGLY_TYPED)) |
config.set_flag(trt.BuilderFlag.FP16 / BF16 / INT8 / FP8) | (整行删除) |
layer.precision = trt.float16 | (语句删除) |
layer.set_output_type(0, trt.float16) | (语句删除) |
替换后的网络创建形式1 << int(trt.NetworkDefinitionCreationFlag.STRONGLY_TYPED)与仓库内已采用强类型的样例完全一致——例如 samples/python/network_api_pytorch_mnist/sample.py 中的 MNIST 构建器:
def build_engine(weights): builder = trt.Builder(TRT_LOGGER) network = builder.create_network(1 << int(trt.NetworkDefinitionCreationFlag.STRONGLY_TYPED)) config = builder.create_builder_config() runtime = trt.Runtime(TRT_LOGGER) config.set_memory_pool_limit(trt.MemoryPoolType.WORKSPACE, common.GiB(1)) populate_network(network, weights) plan = builder.build_serialized_network(network, config) return runtime.deserialize_cuda_engine(plan)注意该样例中已无任何BuilderFlag.FP16调用——这正是迁移目标形态的参照。
What it does NOT change
- 非精度提示类
set_flag:REFIT、SPARSE_WEIGHTS、DISABLE_TIMING_CACHE、TF32。这些与类型化正交,kTF32在 TRT 11 中被保留。 - 校准配置:
set_calibration_profile、IInt8Calibrator。必须人工删除,因为正确的替代方案(ONNX 中的 Q/DQ 节点)位于构建脚本之外。 platform_has_fast_fp16/platform_has_fast_int8条件逻辑:如前所述,若条件体内还残留其他语句,条件块本身不会被删除(删除后变空且无 else 的才会被清理),残留空壳需人工整理。
verify.sh:端到端验证脚本
verify.sh的作用是端到端确认migrate.py的行为符合预期。它会向临时目录写入一个代表性的弱类型样例,对其执行migrate.py --write,然后断言重写后的文件满足强类型契约:STRONGLY_TYPED标志存在、精度标志消失、且是合法 Python。
bash verify.sh # 运行并清理临时工作区 bash verify.sh --keep # 保留临时工作区以供检查成功退出码为 0,任一断言失败则退出码为 1。官方建议在把migrate.py应用到真实代码库之前先运行它确认工具本身可靠。
从 verify.sh 源码看,其验证逻辑包含三个层次:
- 覆盖全部变换路径的样例:内嵌的
sample_build.py刻意包含EXPLICIT_BATCH网络标志、FP16/INT8(应删除)与 TF32/REFIT(应保留)的set_flag混合、逐层precision/set_output_type覆盖,以及platform_has_fast_fp16门控——基本覆盖了 migrate.py 期望处理的每一条变换路径。 - 正反双向断言:
- 必须存在:
NetworkDefinitionCreationFlag.STRONGLY_TYPED、BuilderFlag.TF32、BuilderFlag.REFIT(非精度标志必须存活); - 必须消失:
NetworkDefinitionCreationFlag.EXPLICIT_BATCH、BuilderFlag.FP16、BuilderFlag.INT8、.precision = trt.float16、set_output_type(0, trt.float16)。
- 必须存在:
- 语法有效性:对重写结果执行
ast.parse,确保输出仍是合法 Python。
与仓库样例的对应:从脚本到真实流水线
migrate.py/verify.sh自动化的是"改代码"环节,而完整迁移链路还包含精度入图。仓库提供了可对照的参考实现:
- AutoCast → 强类型 INT8 全流程:samples/python/strongly_type_autocast/sample.py 演示了完整三阶段:ONNX Runtime 在 FP32 模型上跑基线 → 用 ModelOpt
convert_to_mixed_precision生成 FP32/FP16 混合精度 ONNX(含nodes_to_exclude、op_types_to_exclude、keep_io_types等参数,见其convert_model方法)→ 以create_network(1 << int(trt.NetworkDefinitionCreationFlag.STRONGLY_TYPED))构建强类型引擎,并用np.allclose(..., rtol=5e-3, atol=5e-3)校验输出一致性。这正是迁移工作流 Step 3 验证环节的完整模板。 - 强类型 Python 构建器:samples/python/network_api_pytorch_mnist/sample.py 是手写网络场景下的强类型参照。
- trtexec 强类型命令行:samples/trtexec/README.md 的 Example 6 展示了
./trtexec --onnx=model.onnx --stronglyTyped用法(trtexec 路径不在migrate.py覆盖范围内,因为那是命令行而非 Python 源码)。
推荐实操工作流
把以上工具串成一个安全、可回退的迁移流程:
- 先验证工具:运行
bash verify.sh,确认 migrate.py 行为符合预期(官方建议在动真实代码前执行)。 - 建立基线:迁移前先用已知输入集记录弱类型引擎的输出。强类型更严格,弱类型此前悄悄做的精度替换会变得可见。
- dry-run 审查:对目标文件运行
python3 migrate.py path/to/build.py,逐行审查 unified diff——尤其确认非精度标志(TF32/REFIT)确实被保留、被删除的都是真正的精度提示。注意 diff 中出现的"注释丢失"属预期行为。 - 原地改写:审查无误后运行
python3 migrate.py path/to/build.py --write,再按需补回重要注释。 - 人工收尾:处理脚本刻意不动的部分——删除
set_calibration_profile/IInt8Calibrator校准设置;若源模型是 FP32 ONNX 且需要混合精度,先经 ModelOpt AutoCast 把精度(Cast 节点、FP16 初始化器)写入图;清理由platform_has_fast_fp16留下的空条件壳。 - 重建并验证:用迁移后的流程重建引擎,跑基线输入,与弱类型基线对比(
atol=5e-3, rtol=5e-3内一致即完成;超出容差则排查遗留的set_flag、--best或 AutoCast 引入的精度抖动——详见 SKILL.md 的 Common Errors 一节)。
小结
migrate.py与verify.sh把弱类型 → 强类型迁移中最机械、最容易遗漏的代码改写环节自动化了:前者基于 AST 精确识别并重写调用形态(对别名导入免疫、对非精度标志手下留情、对无法确认的形态保守跳过),后者用正反断言 + 语法检查守住重写质量底线。二者的组合让开发者可以放心地对整棵代码树执行迁移,再把精力集中在脚本刻意留给人工的精度入图环节(AutoCast、Q/DQ、校准清理)上,从而平稳跨过 TensorRT 10.12 → 11.x 的强类型门槛。
【免费下载链接】TensorRTNVIDIA® TensorRT™ is an SDK for high-performance deep learning inference on NVIDIA GPUs. This repository contains the open source components of TensorRT.项目地址: https://gitcode.com/GitHub_Trending/tens/TensorRT
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考