ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

TensorRT 弱类型到强类型迁移自动化:trt-strong-typing-migration 辅助脚本实战指南

TensorRT 弱类型到强类型迁移自动化:trt-strong-typing-migration 辅助脚本实战指南 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 typingPython 构建代码自动化改写为强类型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_networkset_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 --checkdry-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_flagREFIT、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 模型上跑基线 → 用 ModelOptconvert_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(..., rtol5e-3, atol5e-3)校验输出一致性。这正是迁移工作流 Step 3 验证环节的完整模板。强类型 Python 构建器samples/python/network_api_pytorch_mnist/sample.py 是手写网络场景下的强类型参照。trtexec 强类型命令行samples/trtexec/README.md 的 Example 6 展示了./trtexec --onnxmodel.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留下的空条件壳。重建并验证用迁移后的流程重建引擎跑基线输入与弱类型基线对比atol5e-3, rtol5e-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),仅供参考
返回列表