ONNX GraphSurgeon 动态 Batch Size 实战用几行 Python 把静态 ONNX 模型改为任意批量推理【免费下载链接】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导读在生产部署中模型输入往往需要支持变化的批量大小batch size以适配在线服务、视频批处理等不同负载场景。本文以 NVIDIA TensorRT 仓库中 ONNX GraphSurgeon 工具链的官方示例 10_dynamic_batch_size 为核心完整讲解如何用几行 Python 代码把静态 batch size 的 ONNX 模型改造成支持任意 batch size 的动态模型先掌握输入张量维度符号化的核心写法再解决静态Reshape层等内部结构对动态化的阻碍并通过仓库自带的 generate.py 与 modify.py 源码逐步拆解实现原理。读完本文你将能够独立完成任意 ONNX 模型的动态 batch size 改造并理解N动态符号与-1自动推导维度背后的 ONNX GraphSurgeon IR 机制。为什么需要动态 Batch SizeONNX 模型的输入张量形状可以有两种形态静态形状如(1, 3, 28, 28)batch size 被硬编码为1。模型只能以 batch 1 的固定输入进行推理无法一次处理多张图也无法在不同批量间复用同一份引擎。动态形状如(N, 3, 28, 28)batch size 是一个可变的符号维度。模型可以接受任意批量部署灵活性和吞吐利用率显著提升。官方示例 README 开门见山地说明了本示例的任务首先生成一个输入 batch size 固定的基础模型然后将其修改为支持动态 batch size 的模型。这正是 ONNX GraphSurgeon 最典型的应用场景之一——在导出模型后、构建 TensorRT 引擎前对模型图进行结构性修正。第一步生成一个静态 Batch Size 的模型运行方式在 examples/10_dynamic_batch_size 目录下执行python3 generate.py脚本会在当前目录生成model.onnx——一个输入形状固定为(1, 3, 28, 28)、包含若干算子节点其中包含静态Reshape的基础模型。源码拆解用 Layer API 构建含静态 Reshape 的模型generate.py 完整展示了如何使用 ONNX GraphSurgeon 的 Layer API 从零构建模型。它首先通过gs.Graph.register()注册了三个自定义图方法把图层构建封装成语义化函数gs.Graph.register() def conv(self, inp, weights, dilations, group, strides): out self.layer( opConv, inputs[inp, weights], outputs[conv_out], attrs{ dilations: dilations, group: group, kernel_shape: weights.shape[2:], strides: strides, }, )[0] out.dtype inp.dtype return outGraph.layer()是 ONNX GraphSurgeon 的核心建图 API它接受算子类型op、输入输出张量列表与属性字典attrs一次性创建节点并返回输出张量。类似的reshape与matmul方法分别封装了Reshape与MatMul算子。随后脚本构造完整的计算图X gs.Variable(nameinput_1, dtypenp.float32, shape(1, 3, 28, 28)) graph gs.Graph(inputs[X], ir_version10) conv_out graph.conv( X, weightsnp.ones(shape(32, 3, 3, 3), dtypenp.float32), dilations[1, 1], group1, strides[1, 1], ) reshape_out graph.reshape(conv_out, np.array([1, 21632], dtypenp.int64)) matmul_out graph.matmul(reshape_out, np.ones(shape(21632, 10), dtypenp.float32)) graph.outputs [matmul_out] model onnx.shape_inference.infer_shapes(gs.export_onnx(graph)) onnx.save(model, model.onnx)从源码结构看这个模型的计算流是输入(1,3,28,28)→ Conv32 个 3×3 卷积核→ Reshape 成(1, 21632)→ MatMul与21632×10权重相乘→ 输出(1, 10)。这里的21632 32 × 26 × 263×3 卷积、无 padding、stride1 后 28×28 变为 26×26。关键点在于reshape这一步Reshape 的目标形状常量是np.array([1, 21632])其中第一个元素1就是被硬编码的 batch size。这就是示例特意构造的陷阱——一个简单的输入维度修改无法覆盖的静态内部结构。最后使用onnx.shape_inference.infer_shapes()做形状推断让每个中间张量都带有精确的形状信息便于后续步骤与可视化观察。下图是该静态模型的可视化结果由示例仓库使用 Netron 生成第二步把静态模型修改为动态 Batch Size运行方式生成model.onnx后在同一个目录下执行python3 modify.py脚本会输出修改后的dynamic.onnx。核心代码输入维度的符号化官方 README 明确指出将静态 ONNX 模型转成动态模型的核心代码极其简短graph gs.import_onnx(onnx_model) for input in graph.inputs: input.shape[0] N逐行拆解这段代码背后的机制gs.import_onnx(onnx_model)是 ONNX GraphSurgeon 的顶层导入 API其实现位于 onnx_importer.py作用是把onnx.ModelProto转换为 ONNX GraphSurgeon 的 IR 图对象Graphgraph.inputs返回图中所有输入张量Tensor的列表input.shape[0] N把输入张量第 0 维batch 维度从整数1改为字符串符号N。ONNX GraphSurgeon 的 IR 设计中张量的shape属性是一个自由修改的列表——从 tensor.py 的源码可以看到Variable的shape在构造时直接赋值对应源码中self.shape shape的赋值逻辑支持整数与字符串混用。因此把 batch 维度写成N后ONNX 导出器会把它编码为动态维度符号onnx.TensorShapeProto.Dimension的dim_param字段从而在 ONNX 语义层面声明该维度可变。示例 modify.py 中实际使用的代码与 README 完全一致graph gs.import_onnx(onnx.load(model.onnx)) # Update input shape for input in graph.inputs: input.shape[0] N为什么还不够静态 Reshape 的问题官方 README 特别强调上述代码对于简单模型已经足够但部分模型可能还需要更新内部层例如静态Reshape层。这是动态化改造中最常见的坑。以本示例模型为例输入维度改成N之后输入张量变为(N, 3, 28, 28)Conv 输出变为(N, 32, 26, 26)但紧接着的Reshape节点目标形状常量仍是[1, 21632]。这个1是硬编码的旧 batch size若推理时 batch 1恰好还能对上一旦 batch ≠ 1比如 N4Reshape会把(4, 32, 26, 26)强行 reshape 成(1, 21632)张量元素总数不匹配直接报错。因此凡是目标形状或shape输入常量中硬编码了 batch size 的Reshape节点都必须一并修正。修正方案把 Reshape 形状改为 -1ONNXReshape的 shape 输入中-1表示该维度大小由张量元素总数与其余维度自动推导。示例 modify.py 的处理逻辑如下# Update Reshape nodes (if they exist) reshape_nodes [node for node in graph.nodes if node.op Reshape] for node in reshape_nodes: # The batch dimension in the input shape is hard-coded to a static value in the original model. # To make the model work with our dynamic batch size, we can use a -1, which indicates that the # dimension should be automatically determined. node.inputs[1].values[0] -1关键点在于最后一行node.inputs[1].values[0] -1node.inputs[1]是Reshape节点的第二个输入张量即形状常量第一个输入是被 reshape 的数据.values是 ONNX GraphSurgeon 中Constant张量暴露的 NumPy 数组属性。从 tensor.py 源码看Constant的values属性按需从底层 ONNX TensorProto 惰性加载LazyValues一旦访问即转为可读写的 NumPy 数组修改后导出时会同步回模型values[0] -1把形状常量[1, 21632]的第一位改成-1即[-1, 21632]。这样无论 batch 是几Reshape都会自动计算出正确的(N, 21632)。由于本示例特意构造了Reshape 形状常量的第 0 位即 batch 维度这一规律values[0] -1即可精准命中。在实际项目中若 Reshape 形状常量里 batch 出现在其他位置、或使用了多个-1之外的特殊布局需要按相同思路定位到 batch 维度的具体下标再修改但原理完全一致把静态 batch 位改为-1让 ONNX 在推理时自动推导。导出动态模型最后一步是把修改后的 IR 导回 ONNX 模型onnx.save(gs.export_onnx(graph), dynamic.onnx)gs.export_onnx(graph)的实现在 onnx_exporter.py它把 ONNX GraphSurgeon 的Graph反向序列化为onnx.ModelProto再由标准onnx.save落盘。改造完成后的模型结构如下图所示对比两张图可以清晰看到改动效果张量/节点静态模型model.onnx动态模型dynamic.onnx输入input_1(1, 3, 28, 28)(N, 3, 28, 28)Conv 输出(1, 32, 26, 26)(N, 32, 26, 26)Reshape 形状常量[1, 21632][-1, 21632]MatMul 输出(1, 10)(N, 10)完整改造流程速览综合 README 与两个脚本整个动态化改造的完整流程是分析模型找出所有输入张量以及所有目标形状中硬编码了 batch 维度的Reshape节点符号化输入遍历graph.inputs把 batch 维度shape[0]替换为动态符号N修正内部结构遍历graph.nodes中所有op Reshape的节点把形状常量中的 batch 位改为-1导出验证gs.export_onnx(graph)导出并通过onnx.shape_inference或可视化工具检查形状传播是否正确。一个可直接套用的通用模板如下import onnx import onnx_graphsurgeon as gs graph gs.import_onnx(onnx.load(model.onnx)) # 1) 输入 batch 维度符号化 for input in graph.inputs: input.shape[0] N # 2) 修正静态 Reshape形状常量首元素为 batch 的场景 for node in graph.nodes: if node.op Reshape and len(node.inputs) 1 and node.inputs[1].values is not None: node.inputs[1].values[0] -1 # 3) 导出 onnx.save(gs.export_onnx(graph), dynamic.onnx)改造后得到的dynamic.onnx可以直接交给 TensorRT 的 ONNX 解析器构建引擎配合引擎构建时的动态形状配置optimization profile即可实现运行时任意 batch 的推理。从源码理解动态化的底层支撑Importers / IR / Exporters 三段式架构本示例依赖的 ONNX GraphSurgeon 库本身是 NVIDIA TensorRT 仓库中的独立工具链组件位于 tools/onnx-graphsurgeon。根据其 README它由三大组件构成Importers把 ONNX 模型导入为 ONNX GraphSurgeon IR接口定义在 base_importer.py高层 APIgs.import_onnx()是其最常用入口IR一切图修改发生的中间表示包含TensorVariable/Constant两个子类、Node、Graph三类对象。IR 的对象模型位于 ir/ 目录下tensor.py、node.py、graph.py。本示例的两处修改——改input.shape[0]与改node.inputs[1].values[0]——都发生在这个 IR 层面Exporters把 IR 导出回 ONNX 等格式接口定义在 base_exporter.py高层 API 为gs.export_onnx()。这种导入—修改—导出三段式正是 ONNX GraphSurgeon 的通用工作范式动态 batch size 改造只是其中的一个典型应用。两个关键 IR 语义Variable.shape支持混合类型从 tensor.py 源码看Variable的shape属性在构造时直接存储传入值允许整数维度与字符串符号共存。本示例利用这一特性把 batch 维写成N导出器会自动映射为 ONNX 的动态维度。Constant.values惰性加载、可读写Constant张量的值通过LazyValues按需从底层 ONNX 张量加载为 NumPy 数组源码中load()方法使用onnx.numpy_helper.to_array实现。node.inputs[1].values[0] -1正是通过这条路径完成形状常量的就地修改。与示例库其他示例的关联本示例属于 ONNX GraphSurgeon 示例集合examples中的第 10 号示例。与其直接相关的还有01_creating_a_model 与 07_creating_a_model_with_the_layer_api本示例 generate.py 使用的Graph.layer()建图方式即来自这一主题04_modifying_a_model讲解通用的模型修改范式动态 batch size 改造正是其方法论的一个具体实例09_shape_operations_with_the_layer_api与形状相关的 Layer API 操作可帮助理解 Reshape 等形状算子在 IR 中的表示方式。常见问题与注意事项只有输入改了Reshape 没改会怎样推理时 batch ≠ 1 会因元素总数不匹配而报错Reshape张量尺寸冲突这正是本示例专门演示静态Reshape场景的原因。为什么用-1而不是直接写NReshape的 shape 输入是常量张量其值在推理期确定-1是 ONNX 规定的自动推导该维度语义让 ONNX Runtime / TensorRT 根据实际输入动态算出 batch 维大小。把常量里的维度写成字符串符号在 ONNX 语义上并不适用。是不是所有模型都要改 Reshape不是。若模型内部没有把 batch 硬编码进常量/属性的结构官方 README 称之为简单模型仅修改输入形状即可。但卷积、池化、全连接等大部分算子都能自然传播 batch 维度而Reshape、Flatten、Squeeze/Unsqueeze、Gather涉及维度索引等形状相关算子需要重点排查。动态 batch 与 TensorRT 的衔接ONNX 层面声明动态 batch 后在 TensorRT 构建引擎时还需通过 optimization profile 指定min/opt/max批次数值范围引擎才会真正为动态 batch 生成优化方案ONNX GraphSurgeon 负责的是模型层面的动态化两者是前后衔接的两步。验证建议修改后可用onnx.shape_inference.infer_shapes()重新做形状推断或用 Netron 可视化确认 batch 维是否已传播到所有下游节点后续在推理框架中分别以 batch1、4、32 等数值实测验证动态化正确性。小结动态 batch size 改造是 ONNX 模型部署前的常见预处理步骤。通过本示例可以看出借助 ONNX GraphSurgeon这一改造只需三步import_onnx导入、把输入 batch 维改为符号N、把静态Reshape形状常量中的 batch 位改为-1最后export_onnx导出即可。官方 README 提供了这一核心思路而 generate.py 与 modify.py 则给出了可直接运行、直接复用的完整实现——这套改输入 修内部的组合拳是解决实际模型动态化问题的最短路径。【免费下载链接】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),仅供参考