Polygraphy 命令行参数组(Argument Groups)机制全解析:从 BaseArgs 接口到自定义 CLI 工具开发
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,几乎每个工具都会订阅它。它负责三件事:
- 位置参数
model_file:模型文件路径(nargs="?"表示可省略,是否必填由构造参数model_opt_required控制)。 --model-type:显式指定模型类型。ModelArgs.ModelType枚举定义了全部合法取值(model.py):
| 类型 | 含义 |
|---|---|
frozen / keras / ckpt |
TensorFlow 冻结图 / Keras 模型 / checkpoint 目录 |
onnx |
ONNX 模型 |
engine / uff / trt-network-script |
TensorRT 引擎 / UFF 文件(已弃用)/ 网络脚本 |
caffe |
Caffe prototxt(已弃用) |
--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
- ONNX(docs/tool/args/backend/onnx/toc.rst)仅含 Loader 一类,源码位于 backend/onnx/loader.py,定义了
OnnxInferShapesArgs(形状推断)、OnnxSaveArgs(保存模型)、OnnxLoadArgs(加载模型,可注入OnnxFromPath加载器)、OnnxFromTfArgs(从 TensorFlow 转换)等参数组。 - ONNX-Runtime(docs/tool/args/backend/onnxrt/toc.rst)同时提供 loader.rst(
OnnxrtSessionArgs,对应 backend/onnxrt/loader.py)与 runner.rst(OnnxrtRunnerArgs,对应 backend/onnxrt/runner.py)。 - TensorFlow(docs/tool/args/backend/tf/toc.rst)除 loader.rst(
TfLoadArgs、TfTrtArgs)与 runner.rst(TfRunnerArgs)外,还有配置类TfConfigArgs(backend/tf/config.py)。 - PluginRef(docs/tool/args/backend/pluginref/runner.rst)的
PluginRefRunnerArgs用于以参考实现运行插件,便于校验 TensorRT 插件输出的正确性。
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,其装配流程完整呈现了参数组的生命周期:
- 订阅(subscription):工具通过实现
get_subscriptions_impl()返回List<a href="https://link.gitcode.com/i/2eff8ffedccb831275794b523a3a46f2" target="_blank">BaseArgs]声明要使用的参数组;基类get_subscriptions()提供默认空实现([tool.py)。 - 实例化与注册:
setup_parser中,LoggerArgs永远最先被实例化,随后是订阅列表中的各组;所有组存入self.arg_groups(一个ArgGroups类型字典,键为参数组类型),并调用每个组的register(self.arg_groups),使参数组能够互相访问(例如ModelArgs可读取RunnerSelectArgs的结果)。这也是某些参数组能"按上下文条件化注册选项"的前提(tool.py)。 - 汇总缩写设置:仅当所有参数组都允许缩写时才启用 argparse 的
allow_abbrev(tool.py)。 - 添加参数:对每个参数组调用
add_parser_args(parser),随后调用工具自身的add_parser_args(可能再挂载子工具 subparsers)(tool.py)。 - 解析与运行:
parse(args)遍历所有参数组调用其parse;随后工具run时即可通过self.arg_groups按类型取出任意参数组及其解析结果(tool.py)。
这套机制的官方示例位于 examples/dev/01_writing_cli_tools(如其中的 gen-data 示例工具),展示了最小化订阅与运行流程。
八、实践:编写自定义参数组
基于上述机制,向 Polygraphy 添加新参数组或新工具的标准步骤如下:
- 继承
BaseArgs(Runner 类则继承BaseRunnerArgs),并在 docstring 中按契约书写Section Header: Description与Depends on:依赖列表——这是 help 输出与依赖声明的唯一来源。 - 实现
add_parser_args_impl:在self.group(基类已自动创建好的 argparse argument group)上添加选项;若本组不添加任何选项,基类会自动将其从 help 中移除。 - 实现
parse_impl:从args命名空间中读取选项并填充属性;复杂参数(形状、容差字典)可借助polygraphy.tools.args.util的解析辅助函数(如parse_meta、parse_arglist_to_dict)。 - 实现
add_to_script_impl:通过polygraphy.tools.script提供的make_invocable、safe、inline等工具将功能写入脚本;BaseRunnerArgs额外要求实现get_name_opt_impl以声明 Runner 名称与--<opt>开关。 - 在工具类中订阅:在
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 中的逐类实现。