Polygraphy Plugin 工具实战:基于 ONNX 图模式匹配与替换实现子图到 Plugin 的自动置换
Polygraphy Plugin 工具实战基于 ONNX 图模式匹配与替换实现子图到 Plugin 的自动置换【免费下载链接】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导读在 NVIDIA TensorRT 生态中将 ONNX 模型中的特定算子组合子图替换为自定义 Plugin是优化推理性能、支持自定义算子语义的常见手段。Polygraphy 的plugin工具正是为这一流程提供的开箱即用解决方案它基于 onnx-graphsurgeon 的图模式匹配能力通过匹配 → 人工审核 → 替换三步流程把用户从手写 onnx-graphsurgeon 脚本的繁琐中解放出来。读完本文你将掌握polygraphy plugin match / list / replace三个子工具的组合用法、pattern.py图模式描述文件的编写规范、config.yaml中间文件的审核与编辑技巧并能用polygraphy run完成替换前后的行为一致性校验。本文以仓库中的完整示例 01_match_and_replace_plugin 为主线展开。一、plugin工具的设计动机三步式子图替换流程plugin工具的目标很明确在 ONNX 模型中查找并替换子图。整个替换被拆解为三个彼此解耦、可人工介入的步骤见 README.mdMatch匹配基于插件提供的图模式描述pattern.py在模型中查找匹配的子图并把所有潜在替换点写入一个用户可编辑的中间文件config.yamlReview审核/编辑人工检查并编辑config.yaml例如当模型中有 2 处匹配子图但只想替换 1 处时可直接从文件中删除对应条目——该文件本质上是一份替换待办清单TODO listReplace替换基于config.yaml中的清单把列出的子图从图中移除替换为单个代表 Plugin 的节点输出新的 ONNX 文件默认名为replaced.onnx原始文件保持不变。这一流程可以用文档中的管道图直观表示original.onnx ------- match ------- config.yaml ------- replace ------- replaced.onnx plugins ----------------^ usr input---^ plugins--------^从实现上看plugin是注册在 plugin.py 中的一个 Polygraphy 工具通过get_subtools_impl()注册了三个子工具Match对应match、ListPlugins对应list、Replace对应replace。因此命令行形态为polygraphy plugin match|list|replace model.onnx ...。二、Match 阶段图模式描述与匹配机制2.1 什么是pattern.py匹配的入口是图模式描述——由插件作者提供的pattern.py文件。该文件需要包含三部分信息见 plugin_base.py 中的调用逻辑图拓扑与约束描述目标子图的节点连接关系以及对节点的附加约束条件如属性取值范围属性计算方式根据匹配到的子图计算 Plugin 节点属性attributes的规则插件元数据插件的name目录名对应的插件名与op替换后 ONNX 节点的op_type。只有提供了pattern.py的插件才会被纳入匹配候选。plugin_base.py中通过glob.glob(os.path.join(plugin_dir, *, pattern.py))扫描--plugin-dir下的每一个子目录凡是包含pattern.py的子目录名都会被识别为一个候选插件随后对每个候选插件依次执行匹配。此外还支持--include name...与--exclude name...两个互斥参数用于精确圈定只匹配/跳过哪些插件。2.2 用 GraphPattern 描述子图拓扑模式描述底层依赖 onnx-graphsurgeon 的GraphPatternAPI实现见 graph_pattern.py。示例中的 toyPlugin 模式描述了一个经典菱形子图见 pattern.pyfrom polygraphy import mod gs mod.lazy_import(onnx_graphsurgeon0.5.0) from typing import List, Dict def get_plugin_pattern(): Toy plugin pattern: A B \ / C, attrs[x] 2.0 / \ D E pattern gs.GraphPattern() in_0 pattern.variable() in_1 pattern.variable() a_out pattern.add(Anode, A, inputs[in_0]) b_out pattern.add(Bnode, B, inputs[in_1]) check_function lambda node : node.attrs[x] 2.0 c_out pattern.add(Cnode, C, inputs[a_out, b_out], check_funccheck_function) d_out pattern.add(Dnode, D, inputs[c_out]) e_out pattern.add(Enode, E, inputs[c_out]) pattern.set_output_tensors([d_out, e_out]) return pattern逐行拆解这段模式描述pattern.variable()声明一个变量张量它是该图模式的输入张量不绑定任何产生它的节点见 graph_pattern.pypattern.add(name, op, inputs[...], check_func...)向模式中添加一个节点。name是该模式节点的逻辑名后续用于取回匹配结果op是 ONNX 算子类型如A、B、Cinputs是该节点的输入张量 id 列表check_func是可选的单节点附加匹配回调见 graph_pattern.pycheck_function lambda node : node.attrs[x] 2.0对Cnode施加约束——只有x属性小于 2.0 的C算子才满足匹配这体现了拓扑 属性约束的组合匹配能力pattern.set_output_tensors([d_out, e_out])声明该模式的两个输出张量set_output_tensors内部会断言输出张量确实有输入节点见 graph_pattern.py。2.3 从匹配结果计算 Plugin 属性get_matching_subgraphs(graph)是pattern.py中第二个必须实现的函数它负责对模型执行匹配并把结果转成结构化的输入/输出/属性字典def get_matching_subgraphs(graph) - List[Dict[str,str]]: gp get_plugin_pattern() matches gp.match_all(graph) ans [] for m in matches: # save the input and output tensor names of the matching subgraph(s) input_tensors list(set([ip_tensor.name for ip_tensor in m.inputs])) output_tensors list(set([op_tensor.name for op_tensor in m.outputs])) attrs {ToyX: int(m.get(Cnode).attrs[x]) * 2} ioa { inputs:input_tensors, outputs:output_tensors, attributes:attrs } ans.append(ioa) return ans关键点gp.match_all(graph)返回图中所有匹配实例其定义见 graph_pattern.py 附近位于同文件的match_all方法每个匹配实例m的inputs/outputs是真实的张量对象这里提取其name去重后作为子图的输入/输出张量名属性计算是插件作者的发挥空间示例把匹配到的Cnode的x属性取出、乘以 2 后作为 Plugin 属性ToyX。这意味着替换时插件节点可以携带与原始子图语义一致的、甚至经过变换的参数m.get(Cnode)即按模式节点名取回匹配到的实际节点映射 API 见 graph_pattern.py。最后get_plugin_metadata()返回插件元数据其中name必须与--plugin-dir下的子目录名一致这里为toyPluginop是替换后 ONNX 节点的算子类型这里为CustomToyPlugindef get_plugin_metadata() - Dict[str,str]: return {name:toyPlugin, op:CustomToyPlugin, }2.4 执行匹配并生成 config.yaml对示例网络执行匹配polygraphy plugin match toy_subgraph.onnx \ --plugin-dir ./plugins -o config.yaml运行时会逐插件打印匹配过程日志例如checking toyPlugin in model [I] Start a subgraph matching... [I] Checking node: n1 against pattern node: Anode. [I] No match because: Op did not match. Node op was: O but pattern op was: A. [I] Start a subgraph matching... [I] Found a matched subgraph! [I] Start a subgraph matching...日志直观展示了模式匹配的试探过程每个候选节点都会与模式中的首节点Anode比对Op did not match表示算子类型不一致、跳过该起点直到找到一个完整满足拓扑与约束的起点并报告Found a matched subgraph!。生成的config.yaml结构如下-o未指定时默认输出到模型所在目录下的config.yaml见 match.py 与 plugin_base.pyname: toyPlugin instances: - inputs: - i1 - i1 outputs: - o1 - o2 attributes: ToyX: 2字段含义name插件名replace阶段据此在--plugin-dir/name/pattern.py定位替换逻辑instances匹配实例列表每项包含inputs子图输入张量名、outputs子图输出张量名、attributes插件节点属性来自get_matching_subgraphs的计算该文件可被yaml.safe_load_all加载为多个 YAML 文档见 replace.py因此一个文件里可容纳多个插件的替换清单。注意示例输出中的inputs为i1, i1是因为匹配实例的输入张量可能被多个模式输入变量命中实际生成时建议在get_matching_subgraphs里如示例代码那样用set去重。测试 test_plugin.py 对 config.yaml 的断言len(plugin[instances]) 1、attributes[ToyX] 2等给出了该文件的规范形态。2.5 干跑预览plugin listplugin list子工具是match的预览/干跑版本只打印每个插件的命中次数不生成config.yamlpolygraphy plugin list toy_subgraph.onnx \ --plugin-dir ./plugins输出形如checking toyPlugin in model [I] Start a subgraph matching... [I] Checking node: n1 against pattern node: Anode. [I] No match because: Op did not match. Node op was: O but pattern op was: A. [I] Start a subgraph matching... ... [I] Found a matched subgraph! [I] Start a subgraph matching... [I] Checking node: n6 against pattern node: Anode. [I] No match because: Op did not match. Node op was: E but pattern op was: A. the following plugins would be used: {toyPlugin: 1}末尾的{toyPlugin: 1}即将使用的插件及命中次数统计。从实现上看ListPlugins继承自 plugin_base.py 中的PluginBase构造时传入list_pluginsTruematch_plugin()在统计完plugin_frequency后若list_pluginsTrue则直接返回、跳过写文件见 plugin_base.py。对应测试见 test_plugin.py。三、Replace 阶段把子图替换为 Plugin 节点3.1 替换命令polygraphy plugin replace toy_subgraph.onnx \ --plugin-dir ./plugins --config config.yaml -o replaced.onnx运行时会打印加载模型等日志最终输出replaced.onnx——示例网络中匹配的子图被替换为toyPlugin节点原始toy_subgraph.onnx保持不变。若省略-o默认输出到模型同目录下的replaced.onnx见 replace.py。replace子工具的完整参数见 replace.py参数是否必填说明model.onnx位置参数是待替换的 ONNX 模型路径--plugin-dir dir是插件目录其下每个含pattern.py的子目录视为一个插件--config path否config.yaml路径缺省时默认读取模型同目录下的config.yaml-o, --output path否输出 ONNX 文件路径缺省为模型同目录下的replaced.onnx3.2 替换的内部实现Replace工具本身不是PluginBase的子类而是独立的Tool见 replace.py其核心流程在replace_plugin()中用 onnx-graphsurgeon 加载模型为graph并建立tensor_map graph.tensors()张量名索引读取config.yamlyaml.safe_load_all支持多插件清单对每个插件尝试从--plugin-dir/plugin_name/pattern.py导入名为replace_with_plugin的自定义替换函数若插件未提供则回退到内置的default_replace_with_plugin见 replace.py逐实例执行替换统计成功替换数量若replace_cnt ! len(plugin[instances])会打印警告用onnx.save(gs.export_onnx(graph), output_onnx)写出结果。内置的default_replace_with_plugin完成了三类图操作见 replace.py断开输入张量到子图内部节点的边对每个输入张量找到其输出节点中输入恰好是子图输入子集的节点即子图内部节点从in_tensor.outputs中摘除断开子图内部节点到输出张量的边对每个输出张量摘除其所有输入节点即子图内部节点插入 Plugin 节点graph.layer(opop, inputsinput_tensors, outputsoutput_tensors, attrsattrs)创建单一节点随后graph.cleanup().toposort()清理孤立节点并重排拓扑。3.3 自定义替换函数如果内置替换逻辑不满足需求例如需要特殊的属性变换或更复杂的图编辑插件作者可以在pattern.py中额外定义replace_with_plugin(graph, input_tensors, output_tensors, attrsNone, opNone)函数其签名与default_replace_with_plugin保持一致replace.py会优先导入并使用它见 replace.py。替换后的模型行为可通过测试窥见一斑test_plugin.py 断言替换后的模型只剩 2 个节点保留n1其余n2~n6全部消失且新节点op_type CustomToyPlugin、携带属性ToyX 2。四、Compare 阶段替换前后行为一致性校验替换是否等价必须用数据说话。文档给出的校验方案是复用 Polygraphy 的run工具先跑原始模型并保存输出再跑替换后模型并比对polygraphy run original.onnx --trt --save-outputs model_output.json polygraphy run replaced.onnx --trt --load-outputs model_output.json--trt指定用 TensorRT 作为后端第一行把原始模型的输出保存到model_output.json第二行加载该输出作为比对基准run会对替换后模型的输出执行数值比对从而验证 Plugin 替换没有改变模型行为。该思路与 Polygraphy 通用的保存输出 → 加载输出比对工作流一致同样适用于 ONNX Runtime 等其他后端。五、配套测试与目录结构本示例在仓库中的配套资源如下均可对照阅读示例入口与完整流程说明01_match_and_replace_plugin/README.md待替换的示例网络toy_subgraph.onnx插件模式描述plugins/toyPlugin/pattern.py工具实现polygraphy/tools/plugin/plugin.py、subtool/match.py、subtool/list_plugins.py、subtool/replace.py、subtool/plugin_base.py图模式匹配引擎graph_pattern.py自动化测试tests/tools/test_plugin.py覆盖 match / list / replace 三条链路该示例同时被 tests/test_examples.py 注册为端到端示例测试产物为config.yaml与replaced.onnx六、应用要点与最佳实践小结目录即插件--plugin-dir下每个含pattern.py的子目录被自动识别为一个插件子目录名必须与get_plugin_metadata()[name]一致模式约束越精确误匹配越少check_func是对节点属性乃至任意可由节点对象计算的谓词施加约束的入口示例中用node.attrs[x] 2.0演示了属性门控复杂场景还可叠加pattern.constant()常量张量匹配能力见 graph_pattern.pyconfig.yaml 是人工把关点匹配结果不是一键替换而是先落盘成可读的 YAML多实例场景可删减条目后再执行替换避免误伤先用 list 干跑不确定命中情况时先polygraphy plugin list看命中数与分布零成本验证模式写得对不对替换后必须做数值校验polygraphy run --save-outputs/--load-outputs是成本最低的一致性验证手段尤其当check_func涉及属性条件时更要确认替换节点携带的属性与原子图语义等价默认替换逻辑可被覆盖在pattern.py中提供replace_with_plugin即可接管替换行为满足定制化图编辑需求。七、结语Polygraphy 的plugin工具把ONNX 子图 → Plugin这条在 TensorRT 优化中高频出现的工作流收敛为match → 审阅 config.yaml → replace → run 校验四个清晰步骤。其核心价值在于模式描述pattern.py与替换执行replace的职责分离让插件作者只关心匹配什么与如何计算属性而把图编辑的脏活交给default_replace_with_plugin与 onnx-graphsurgeon。本文示例虽小菱形子图但 A→C、B→C 的分支汇合、双输出拆分、属性约束与属性变换等要素一应俱全完全可以作为你接入自定义 TensorRT 插件的起点模板。【免费下载链接】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),仅供参考