Polygraphy 命令行参数组(Argument Groups)机制全解析:从 BaseArgs 接口到自定义 CLI 工具开发

原创2026-09-14 12:05:451,413 阅读
文章标签:人工智能推理引擎深度学习本地部署模型优化

Polygraphy 命令行参数组(Argument Groups)机制全解析:从 BaseArgs 接口到自定义 CLI 工具开发

Polygraphy 是 NVIDIA TensorRT 生态中面向模型调试与部署验证的 CLI 工具集,其命令行工具普遍共享大量相似功能(加载模型、构建引擎、运行推理、比较输出等)。为此,Polygraphy 将「命令行选项 + 解析逻辑 + 相关功能」打包为可复用的**参数组(Argument Group)**组件,并贯穿于全部 CLI 工具之中。本文以 tools/Polygraphy/docs/tool/args/toc.rst 为骨架,结合 polygraphy/tools/args 目录下的源码实现,系统讲解参数组的体系结构、各组职责、关键参数语义,以及如何基于该机制扩展或编写新的 CLI 工具。读完本文,你将掌握 Polygraphy 参数组的订阅、注册、解析与脚本生成全链路,能够读懂 polygraphy run 等工具的行为,并具备编写自定义参数组的能力。

一、什么是参数组(Argument Groups)

Polygraphy 的多个 CLI 工具(如 polygraphy run、polygraphy surgeon 等)之间存在大量重复需求:解析模型文件、配置 TensorRT Builder、选择推理后端、比较输出、控制日志输出。若每个工具各自实现一遍,会导致代码重复、选项行为不一致、维护成本急剧上升。

参数组正是为消除这种重复而设计的可复用组件。一个参数组将一组命令行选项、对应的解析逻辑以及与这些选项相关的功能实现绑定在一起,供任意工具按需"订阅"。Polygraphy 官方对该机制的描述与用法记录在 polygraphy/tools/args/README.md,而参数组的完整清单目录即本文所依据的 tools/Polygraphy/docs/tool/args/toc.rst。

从文档目录结构看,Polygraphy 的参数组被组织为五大部分:

文档章节 对应 RST 文件 核心内容
Base Interface base.rst 所有参数组的基类接口 BaseArgs
Backends backend/toc.rst 各类推理后端:ONNX、ONNX-Runtime、TensorFlow、TensorRT、PluginRef
Comparator comparator/toc.rst 推理执行与输出比较:Comparator.run、比较函数、数据加载、后处理
Logger logger/toc.rst 日志与调试输出控制
Model model.rst 模型文件路径、类型与输入形状

所有参数组统一归属于模块 polygraphy.tools.args,每个参数组本质上是该模块下的一个 Python 类。

二、核心接口:BaseArgs

所有参数组都继承自 base.py 中定义的 BaseArgs。它承担三类职责:向 argparse 解析器添加选项、从命令行参数中解析出所需信息、生成可执行的 Python 脚本代码。

2.1 三个核心方法

  • add_parser_args(self, parser):向 argparse.ArgumentParser 添加命令行选项。该方法保证在 register 之后被调用,因为某些参数组会根据"当前已有哪些参数组"来决定是否注册特定选项。
  • parse(self, args):从 argparse 生成的 args 命名空间(Namespace)中解析出本组关心的参数,并填充到参数组自身的属性中。对简单参数而言,这通常只是把 args 对象的某个属性赋给参数组属性;对复杂参数(如输入形状),则会涉及更复杂的解析逻辑。
  • add_to_script(self, script):向一段 Python 脚本中添加与本参数组功能相关的代码。例如负责加载 ONNX 模型的 OnnxLoadArgs,会在脚本中注入一个 OnnxFromPath 加载器。

2.2 为什么是"向脚本添加代码"而非直接执行

这是 Polygraphy 参数组设计中最具特色的地方。与其提供一个直接执行功能的 load_onnx() 方法,Polygraphy 选择把功能拼装成一段可编辑的 Python 脚本。原因在于:polygraphy run 等工具需要组合出复杂行为,而生成脚本后,用户可以拿到这段脚本进行编辑,再用于更高级的定制需求(例如在脚本中插入自定义数据处理、断点调试等)。

与此同时,多数参数组也提供立即求值的辅助方法(如 OnnxLoadArgs.load_onnx()),供那些无需生成脚本的工具直接调用。这些辅助方法通常通过 polygraphy.tools.args.util 中的 run_script 方法复用 add_to_script 的逻辑,从而保证两种调用路径行为一致。

2.3 参数组的 docstring 契约

add_parser_args 在实现上有一个隐含契约(见 base.py):子类必须在其 docstring 首行写明 Section Header: Description 格式的标题与描述,例如:

TensorRT Engine: loading TensorRT engines.

Depends on:
    - ModelArgs
    - TrtLoadPluginsArgs
    - TrtLoadNetworkArgs: if building engines
    ...

其中「标题 + 描述」会被用来填充工具 --help 输出的参数组标题(格式为 Options related to <desc>);Depends on: 部分则声明本参数组的依赖项。如果 docstring 格式不正确,add_parser_args 会通过 G_LOGGER.internal_error 直接抛出内部错误(base.py)。

2.4 其他关键能力

  • allows_abbreviation():是否允许使用选项前缀缩写。默认返回 True,此时 --iterations 可简写为 --iter。但缩写会破坏 argparse.REMAINDER,因此任何使用该特殊参数的参数组应将其禁用(base.py)。
  • 空组自动剔除:add_parser_args 会记录调用前后的 parser._actions 数量;若某参数组没有添加任何选项,会从 help 输出中移除该空组(base.py)。
  • BaseRunnerArgs:BaseArgs 的专用子类,面向「Runner(运行器)」类参数组。它新增 get_name_opt(),返回一个二元组:人类可读的 Runner 名称 + 用于选择该 Runner 的命令行选项名(不含前导 --),例如 ("TensorRT", "trt")(base.py)。这一约定是下方 Runner 选择机制的基础。

三、Model 参数组:模型、类型与输入形状

model.rst 对应的实现是 model.py 中的 ModelArgs,几乎每个工具都会订阅它。它负责三件事:

  1. 位置参数 model_file:模型文件路径(nargs="?" 表示可省略,是否必填由构造参数 model_opt_required 控制)。
  2. --model-type:显式指定模型类型。ModelArgs.ModelType 枚举定义了全部合法取值(model.py):
类型 含义
frozen / keras / ckpt TensorFlow 冻结图 / Keras 模型 / checkpoint 目录
onnx ONNX 模型
engine / uff / trt-network-script TensorRT 引擎 / UFF 文件(已弃用)/ 网络脚本
caffe Caffe prototxt(已弃用)
  1. --input-shapes(别名 --inputs):以 <name>:<shape> 格式声明推理输入形状,例如 --input-shapes image:[1,3,224,224] other_input:[10]。该信息用于确定生成推理输入数据时的张量形状。

当未显式指定 --model-type 时,ModelArgs 会依据扩展名自动推断类型。其 EXT_MODEL_TYPE_MAPPING(model.py)为:

  • .hdf5 → keras;.uff → uff;.prototxt → caffe;.onnx → onnx
  • .engine / .plan → engine;.graphdef → frozen;.py → trt-network-script

其中 trt-network-script 是一种特殊类型:模型文件是一个定义了 load_network 函数(无参,返回 TensorRT Builder、Network 及可选的 Parser)的 Python 脚本。函数名可通过冒号附加在文件路径后,如 my_custom_script.py:my_func。若启用了 guess_model_type_from_runners(默认 False),模型类型还可结合已选 Runner 来推断——此时要求 RunnerSelectArgs 必须先于 ModelArgs 完成解析,否则会触发内部错误(model.py)。

四、Backend 参数组:六大推理后端

backend/toc.rst 汇总了与推理后端相关的参数组,分为四类:ONNX、ONNX-Runtime(onnxrt)、TensorFlow(tf)、TensorRT(trt),外加 PluginRef。每个后端下又细分为 Loader(加载器) 与 Runner(运行器) 两类参数组。

4.1 Runner 选择机制:RunnerSelectArgs

在各后端参数组之上,runner_select.py 中的 RunnerSelectArgs 扮演调度角色:它遍历当前工具订阅的所有参数组,凡属于 BaseRunnerArgs 的,都会自动获得一个 --<opt> 开关(如 --trt、--onnxrt),用于在命令行中按出现顺序选择要执行的 Runner(runner_select.py)。解析后得到的 runners 属性是形如 [("trt", "TensorRT"), ("onnxrt", "ONNX-Runtime")] 的有序列表。

4.2 ONNX 与 ONNX-Runtime

4.3 TensorRT 参数组:从网络加载到引擎保存

TensorRT 是 Polygraphy 的核心后端,其参数组最丰富,文档见 docs/tool/args/backend/trt/toc.rst(loader 与 runner 两篇),源码见 backend/trt/loader.py、backend/trt/config.py 与 backend/trt/runner.py。按其职责可归纳为五类:

(1)加载网络:TrtLoadNetworkArgs(loader.py)提供的关键选项包括:

  • --trt-outputs:指定 TensorRT 输出张量名称;传 mark all 表示将所有张量作为输出。
  • --trt-exclude-outputs(实验性):取消标记某些输出张量。
  • --layer-precisions <layer>:<precision>:按层指定计算精度,如 --layer-precisions example_layer:float16 other_layer:int8,精度取值来自 TensorRT 数据类型别名(float32/float16/int8/bool 等);使用该选项时应同时配合 --precision-constraints prefer|obey。
  • --tensor-dtypes(别名 --tensor-datatypes)<tensor>:<dtype>:按张量指定网络 I/O 数据类型。
  • --tensor-formats <tensor>:[<format>,...]:按张量指定允许的格式(取值来自 trt.TensorFormat 枚举、大小写不敏感),例如 --tensor-formats example_tensor:[linear,chw4]。
  • --strongly-typed:将网络标记为强类型(strongly typed)。
  • --mark-debug:将指定张量标记为调试张量。
  • --trt-network-postprocess-script(别名 --trt-npps):指定对解析后网络做后处理的脚本,可携带函数名(如 process.py:do_something),默认调用 postprocess(network=...),多个脚本按给定顺序执行。

(2)加载插件:TrtLoadPluginsArgs(loader.py)用于加载自定义 TensorRT 插件。

(3)构建配置:TrtConfigArgs(config.py)负责生成 TensorRT BuilderConfig,其 docstring 按契约声明依赖 ModelArgs(用于获取输入形状以构造 profile)。动态形状 profile 由 parse_profile_shapes 解析(config.py):它基于默认输入形状,叠加各 profile 的 min/opt/max 形状参数,生成 {input_name: (min, opt, max)} 形式的 profile 列表;若三者中任一数量多于其他,以最大数量为准补齐;还会对动态维度做覆盖处理并给出告警。若 min/opt/max 的输入名集合不一致,会直接 critical 报错。

(4)保存/加载引擎:TrtSaveEngineBytesArgs、TrtSaveEngineArgs、TrtLoadEngineBytesArgs、TrtLoadEngineArgs(loader.py)分别管理引擎字节流与引擎文件的保存、加载。

(5)运行器:TrtRunnerArgs(backend/trt/runner.py)作为 BaseRunnerArgs 子类,通过 get_name_opt 声明自身选项为 trt,从而被 RunnerSelectArgs 自动注册为 --trt 开关。

五、Comparator 参数组:推理执行与输出比较

comparator/toc.rst 汇总了与「比较」相关的四类参数组,是 polygraphy run 这类工具的核心。对应源码位于 backend 之外的 comparator 目录 下。

5.1 ComparatorRunArgs:运行推理

实现于 comparator.py,对应 Comparator.run(),依赖 DataLoaderArgs。核心选项:

  • --warm-up NUM:正式计时前执行的预热次数。
  • --use-subprocess:在独立子进程中运行各 Runner(不能与调试器同时使用)。
  • --save-inputs(别名 --save-input-data):将推理输入保存为 JSON 编码的 List[Dict[str, numpy.ndarray]]。
  • --save-outputs(别名 --save-results):将各 Runner 结果(RunResults)保存为 JSON。

其 add_to_script_impl 展示了脚本生成的实际形态:注入 Comparator 导入,拼接 results = Comparator.run(...) 调用,并在指定保存路径时追加 results.save(...) 代码(comparator.py)。

5.2 CompareFuncSimpleArgs:简单比较函数

实现于 compare.py,对应 CompareFunc.simple。它定义了精度校验的核心参数,其中 --rtol、--atol、--check-error-stat、--error-quantile 均支持按输出张量分别指定,格式为 [<out_name>:]<value>:

选项 语义
--no-shape-check 关闭输出形状严格一致性检查
--rtol / --rel-tol 相对容差,以第二组输出值的百分比表示,如 0.01 表示 1% 以内
--atol / --abs-tol 绝对容差
--check-error-stat 检查的误差统计量(如 max、mean、median)
--infinities-compare-equal 匹配的正负无穷视为 absdiff 0(默认视为 NaN)
--error-quantile 比较的误差分位数([0,1] 内的浮点数)
--save-heatmaps / --show-heatmaps(实验性) 保存 / 显示绝对与相对误差热力图
--save-error-metrics-plot / --show-error-metrics-plot(实验性) 保存 / 显示误差指标图,用于分析误差趋势、判断是否仅个别离群点失准

源码提示:默认容差对 FP32 通常适用,但对 FP16、INT8 等低精度可能过严(compare.py)。

5.3 DataLoaderArgs:输入数据加载与生成

实现于 data_loader.py。它负责为推理生成或加载输入数据:

  • --seed SEED:随机输入的数据种子。
  • --val-range:生成数据的取值范围,支持按输入名分别指定,如 --val-range [0,1] inp0:[2,50] inp1:[3.0,4.6];未显式指定的输入使用默认范围。
  • --int-min / --int-max / --float-min / --float-max:已弃用,统一改用 --val-range。
  • --iterations(别名 --iters)NUM:默认数据加载器供数的推理迭代次数。
  • --data-loader-backend-module:生成输入数组所用模块,可选 numpy、torch。
  • --load-inputs(别名 --load-input-data):从 JSON 文件加载输入(List[Dict[str, numpy.ndarray]]),使用后其余数据加载参数全部忽略。
  • --data-loader-script:指定定义数据加载函数的 Python 脚本,函数无参、返回生成输入的生成器/可迭代对象;默认查找 load_data,可用冒号指定其他函数名(my_custom_script.py:my_func)。

上述 --load-inputs 与 --data-loader-script 处于互斥组中(data_loader.py)。

5.4 其他 Comparator 参数组

comparator.rst 还包含 ComparatorCompareArgs(比较配置,CompareFuncSimpleArgs 与 CompareFuncIndicesArgs 的依赖项)以及 postprocess.rst 对应的 ComparatorPostprocessArgs(推理输出后处理)。

六、Logger 参数组:日志与调试输出

logger/toc.rst 对应的 LoggerArgs 实现于 logger/logger.py,它是每个工具默认必订阅的参数组(见下文工具基类流程)。其选项构成一套完整的日志控制体系:

  • -v / --verbose 与 -q / --quiet:计数型选项,可重复指定以逐级提高/降低日志详细程度。
  • --verbosity:优先级高于 -v/-q,且支持按路径控制日志级别。取值来自 Logger 类的日志级别(大小写不敏感),如 --verbosity INFO;按路径格式为 <path>:<verbosity>,例如 --verbosity backend/trt:INFO backend/trt/loader.py:VERBOSE。路径相对于 polygraphy/ 目录(即 polygraphy/backend 写作 backend),匹配时取最贴近的路径级别(例如同时存在 warning、backend:info、backend/trt:verbose 时,comparator 下文件用 WARNING,backend/onnx 用 INFO,backend/trt 用 VERBOSE)。
  • --silent:关闭所有输出。
  • --log-format:日志格式,可选 timestamp(含时间戳)、line-info(含文件与行号)、no-colors(禁用颜色),可组合指定。
  • --log-file:将 Polygraphy 日志写入文件(不包含 TensorRT、ONNX-Runtime 等依赖库自身的日志)。

七、参数组的订阅与装配:工具基类全流程

参数组最终服务于 CLI 工具。工具基类 Tool 位于 tools/base/tool.py,其装配流程完整呈现了参数组的生命周期:

  1. 订阅(subscription):工具通过实现 get_subscriptions_impl() 返回 List<a href="https://link.gitcode.com/i/2eff8ffedccb831275794b523a3a46f2" target="_blank">BaseArgs] 声明要使用的参数组;基类 get_subscriptions() 提供默认空实现([tool.py)。
  2. 实例化与注册:setup_parser 中,LoggerArgs 永远最先被实例化,随后是订阅列表中的各组;所有组存入 self.arg_groups(一个 ArgGroups 类型字典,键为参数组类型),并调用每个组的 register(self.arg_groups),使参数组能够互相访问(例如 ModelArgs 可读取 RunnerSelectArgs 的结果)。这也是某些参数组能"按上下文条件化注册选项"的前提(tool.py)。
  3. 汇总缩写设置:仅当所有参数组都允许缩写时才启用 argparse 的 allow_abbrev(tool.py)。
  4. 添加参数:对每个参数组调用 add_parser_args(parser),随后调用工具自身的 add_parser_args(可能再挂载子工具 subparsers)(tool.py)。
  5. 解析与运行:parse(args) 遍历所有参数组调用其 parse;随后工具 run 时即可通过 self.arg_groups 按类型取出任意参数组及其解析结果(tool.py)。

这套机制的官方示例位于 examples/dev/01_writing_cli_tools(如其中的 gen-data 示例工具),展示了最小化订阅与运行流程。

八、实践:编写自定义参数组

基于上述机制,向 Polygraphy 添加新参数组或新工具的标准步骤如下:

  1. 继承 BaseArgs(Runner 类则继承 BaseRunnerArgs),并在 docstring 中按契约书写 Section Header: Description 与 Depends on: 依赖列表——这是 help 输出与依赖声明的唯一来源。
  2. 实现 add_parser_args_impl:在 self.group(基类已自动创建好的 argparse argument group)上添加选项;若本组不添加任何选项,基类会自动将其从 help 中移除。
  3. 实现 parse_impl:从 args 命名空间中读取选项并填充属性;复杂参数(形状、容差字典)可借助 polygraphy.tools.args.util 的解析辅助函数(如 parse_meta、parse_arglist_to_dict)。
  4. 实现 add_to_script_impl:通过 polygraphy.tools.script 提供的 make_invocable、safe、inline 等工具将功能写入脚本;BaseRunnerArgs 额外要求实现 get_name_opt_impl 以声明 Runner 名称与 --<opt> 开关。
  5. 在工具类中订阅:在 get_subscriptions_impl() 中返回该参数组的实例,基类即会自动完成注册、加参、解析与脚本生成的全流程。

对 Runner 类参数组而言,只要实现了 get_name_opt_impl,RunnerSelectArgs 就会自动为它生成选择开关,无需任何额外注册代码——这正是参数组「可复用、可组合」设计意图的集中体现。

九、总结

Polygraphy 的参数组机制以 BaseArgs 为统一接口,将「argparse 选项注册、命令行解析、脚本代码生成」三件事封装为可订阅、可组合、可扩展的组件:ModelArgs 统一模型入口,六个 Backend 参数组(ONNX / ONNX-Runtime / TF / TRT / PluginRef)按 Loader 与 Runner 分工,Comparator 系参数组覆盖「数据加载 → 推理运行 → 输出比较 → 后处理」全链路,LoggerArgs 提供全局日志控制,而工具基类通过 get_subscriptions 与 arg_groups 完成装配。理解这一体系,不仅能让你熟练运用 polygraphy run 等工具的全部参数语义,也为基于该框架编写自己的 CLI 工具提供了清晰的扩展路径。更详细的参数组清单可继续查阅 docs/tool/args/toc.rst 下各子文档,以及源码目录 polygraphy/tools/args 中的逐类实现。

登录后查看全文
TensorRT