ONNX 算子形状推断(Shape Inference)实现指南:从 TypeAndShapeInferenceFunction 到测试验证
【免费下载链接】onnxOpen standard for machine learning interoperability项目地址: https://gitcode.com/gh_mirrors/onn/onnx
本篇指南以 ONNX 仓库中.agents/skills/add-shape-inference/SKILL.md为骨架,结合 docs/ShapeInference.md 官方文档与 onnx/defs/shape_inference.h、onnx/defs/tensor/defs.cc、tests/python/shape_inference_test.py 等源码,系统讲解如何为 ONNX 算子实现类型与形状推断(TypeAndShapeInferenceFunction)、传播维度信息、处理广播逻辑并编写测试。读完本文,你将掌握 ONNX 形状推断的完整实现套路:推断函数注册位置、类型推断与形状推断的职责划分、常用工具函数矩阵、维度算术、测试写法以及健壮性规则,能够独立为自定义算子或既有算子补充/修复形状推断逻辑。
一、形状推断在 ONNX 中的定位
ONNX 提供了一套可选的图级形状推断实现,覆盖每个核心算子,并暴露了扩展接口:你可以直接调用现成的推断能力,也可以为自定义算子定义推断实现。推断函数以OpSchema成员(TypeAndShapeInferenceFunction)的形式存储,随算子 schema 一同注册。
关于形状推断的能力边界,docs/ShapeInference.md 明确了以下事实:
- 推断并非保证完备:例如
Reshape到动态提供的形状会阻断推断流;且并非所有算子都要求实现推断函数。 - 推断只处理常量与简单变量:
Concat对(5, 2)与(7, 2)可以推断出(12, 2),但对(5, 2)与(N, 2)只能得到未知符号(M, 2)——M 代表与其他出现处相同的未知量(符号维度dim_param会被传播)。这些是当前实现的属性而非根本性约束。
静态张量形状由TensorShapeProto表示,与运行时形状相区别:
shape字段未定义 ⇒ 秩未知的张量;shape已定义 ⇒ 秩已知;- 每个
Dimension的dim_value表示已知整数值,dim_param表示符号标识,两者均未设置则为匿名未知值。
二、推断函数的文件位置与注册方式
从.agents/skills/add-shape-inference/SKILL.md的文件位置表出发,推断相关代码分布在三个层次:
| 组件 | 位置 |
|---|---|
| 推断函数本体 | onnx/defs/<domain>/defs.cc(内联在 schema 定义处) |
| 工具函数与核心接口 | onnx/defs/shape_inference.h |
| Python 测试 | tests/python/shape_inference_test.py |
注册方式是在算子 schema 上调用:
OpSchema& OpSchema::TypeAndShapeInferenceFunction(InferenceFunction inferenceFunction);InferenceFunction与核心接口结构体InferenceContext均定义于 onnx/defs/shape_inference.h:InferenceContext是传入推断函数的上下文,负责读取算子输入信息并写入推断结果。图级入口是shape_inference::InferShapes(ModelProto& m, const ISchemaRegistry* schema_registry),它在原模型上原地标注形状信息(C++ API);Python 侧则有对应的 shape inference API(见 docs/PythonAPIOverview.md)。
命名函数优于内联 lambda
代码规范要求将推断函数定义为独立命名函数而非内联 lambda:宏展开后内联 lambda 上的断点不可靠。短单行实现(如直接引用propagateShapeAndTypeFromFirstInput)则可以直接引用。以Transpose的 schema 为例(onnx/defs/tensor/defs.cc),其推断函数就是直接挂在TypeAndShapeInferenceFunction上。
三、类型推断与形状推断的职责划分
类型推断(元素类型)通常由 schema 的类型约束(type constraints)自动完成:当类型约束变量(如"T")同时出现在输入与输出上时,框架会自动将输入元素类型传播到输出,无需显式推断代码。但许多既有算子仍显式调用propagateElemTypeFromInputToOutput作为健壮性最佳实践——这在类型约束已覆盖的情况下是无害的,且能保证无论推断如何被调用都行为正确。
只有以下场景才需要在TypeAndShapeInferenceFunction中写显式类型推断逻辑:
- 输出类型由属性决定(如
Cast的to属性指定输出元素类型); - 输出类型与所有输入类型都不同,且无法用共享类型约束变量表达;
- 算子使用异构(heterogeneous)变参输入/输出。
同构与异构变参
同构/异构标志只适用于变参(repeated)输入或输出:
- 同构(默认):所有重复参数类型必须相同,类型约束变量约束它们一致,框架自动强制并传播;
- 异构:每个重复参数可以类型不同,类型约束变量只描述"允许的类型集合"。
Loop、Scan等算子使用该模式(其携带状态变量可混合类型)。
使用异构变参时,推断函数必须为每个参数显式传播类型,框架无法自动完成。
形状推断则几乎总是需要显式逻辑,因为输出形状通常取决于输入形状、属性或两者共同决定。
四、三种核心推断模式(含完整代码)
4.1 一元逐元素算子
.TypeAndShapeInferenceFunction(propagateShapeAndTypeFromFirstInput)propagateShapeAndTypeFromFirstInput的实现(onnx/defs/shape_inference.h)会先propagateElemTypeFromInputToOutput(ctx, 0, 0)复制元素类型,再在hasNInputShapes(ctx, 1)通过时复制整个形状。
4.2 带广播的二元算子
static void InferShapeForBinaryOp(InferenceContext& ctx) { propagateElemTypeFromInputToOutput(ctx, 0, 0); if (hasNInputShapes(ctx, 2)) bidirectionalBroadcastShapeInference( ctx.getInputType(0)->tensor_type().shape(), ctx.getInputType(1)->tensor_type().shape(), *ctx.getOutputType(0)->mutable_tensor_type()->mutable_shape()); }bidirectionalBroadcastShapeInference(L, R, out)(onnx/defs/shape_inference.h)实现 Numpy 风格的双向广播规则,且会安全处理缺失的维度。
4.3 改变形状的算子(以 Transpose 为例)
SKILL 文档给出的模板与仓库实际实现一致。Transpose的真实推断函数(onnx/defs/tensor/defs.cc)除按perm重排维度外,还包含属性合法性校验:
static void InferShapeForTranspose(InferenceContext& ctx) { propagateElemTypeFromInputToOutput(ctx, 0, 0); if (!hasNInputShapes(ctx, 1)) return; auto input_shape = ctx.getInputType(0)->tensor_type().shape(); int rank = input_shape.dim_size(); std::vector<int64_t> perm; getRepeatedAttribute(ctx, "perm", perm); auto* output_shape = getOutputShape(ctx, 0); for (int i = 0; i < rank; ++i) { *output_shape->add_dim() = input_shape.dim(perm[i]); } }仓库实现还额外做了两类校验(可视为该模式的完整版):若未提供perm则默认反转维度(perm.reserve(shape.dim_size())后从高到低填入索引);若提供了perm则检查每个索引在[0, rank-1]范围内且不重复,否则调用fail_type_inference报错。测试 tests/python/shape_inference_test.py 覆盖了perm=[1, 0, 2]下(2, 3, 4) → (3, 2, 4)的完整推断、标量输入、部分形状等场景。
五、核心工具函数速查表
SKILL 文档整理了推断函数中最常用的工具函数:
| 函数 | 用途 |
|---|---|
propagateElemTypeFromInputToOutput(ctx, in, out) | 复制元素类型 |
propagateShapeFromInputToOutput(ctx, in, out) | 复制整个形状 |
propagateShapeAndTypeFromFirstInput(ctx) | 从输入 0 复制类型与形状 |
hasNInputShapes(ctx, n) | 检查前 n 个输入是否有形状 |
getOutputShape(ctx, out) | 获取可变的输出形状 |
bidirectionalBroadcastShapeInference(L, R, out) | Numpy 风格广播 |
getRepeatedAttribute(ctx, "name", vec) | 读取重复属性值 |
getAttribute(ctx, "name", default) | 读取单个属性值 |
mergeInDimensionInfo(src, dst, dim_idx) | 合并维度信息 |
fail_shape_inference("msg") | 抛出推断错误 |
官方文档 docs/ShapeInference.md 还补充了一组更高层的工具:
checkInputRank(ctx, n, rank):校验输入必须是固定秩(参考RoiAlign的推断实现);unifyInputDim/unifyDim/updateOutputShape:当多个输入维度期望相同、或输入维度需传播到特定输出维度时使用(参考RoiAlign);unifyInputShape/unifyInputShapePrefix:在unifyInputDim之上构建的声明式高层工具,一次调用统一输入的全部(或前缀)维度,适合简单场景;复杂场景仍需逐个unifyInputDim;hasInputShape(ctx, n):单输入形状检查(hasNInputShapes的基础)。
这些工具都对缺失的形状/维度做了安全处理。
从源码看 hasNInputShapes 的语义
onnx/defs/shape_inference.h 的实现显示,hasInputShape需要同时满足三个条件:ctx.getNumInputs() > n、ctx.getInputType(n)非空、且该类型(支持 tensor/sparse tensor/sequence/optional 递归)确实携带形状。hasNInputShapes则对前 n 个输入逐一检查。而propagateShape(onnx/defs/shape_inference.h)展示了形状传播的细节:当输入形状"未知"时,输出也保持未知(即不给输出赋值任何形状),并支持 tensor、sparse tensor、sequence、optional、map 等多种类型的递归传播。
使用unifyInputShape的声明式写法
官方文档给出的矩阵乘法例子展示了两种等价写法。显式写法:
checkInputRank(ctx, 0, 2); // 输入 0 秩为 2(若其秩已知) checkInputRank(ctx, 1, 2); // 输入 1 秩为 2(若其秩已知) Dim M, K, N; unifyInputDim(ctx, 0, 0, M); unifyInputDim(ctx, 0, 1, K); unifyInputDim(ctx, 1, 0, K); unifyInputDim(ctx, 1, 1, N); updateOutputShape(ctx, 0, {M, N});更简洁的声明式写法:
Dim M, K, N; unifyInputShape(ctx, 0, {M, K}); unifyInputShape(ctx, 1, {K, N}); updateOutputShape(ctx, 0, {M, N});六、维度算术(Dimension Arithmetic)
当输出维度由输入维度通过算术计算得出时,可使用符号维度(Dim)重载运算符:
Dim operator*(const Dim& a, const Dim& b); // 维度相乘 Dim operator*(const Dim& a, int64_t val); // 维度乘常量 Dim operator/(const Dim& a, int64_t divisor); // 维度整除 Dim multiplyDims(const TensorShapeProto& shape, int from, int upto); // 区间维度连乘官方文档提示可参考SpaceToDepth的推断实现:*与/可安全作用于符号维度。Dim算术天然兼容dim_value与dim_param两种维度,并在 onnx/defs/shape_inference.h 中对整数溢出、除零等异常抛出fail_shape_inference错误。
七、编写形状推断测试
参数化测试:_make_graph+_assert_inferred
对需要跨算子版本(opset)回归的场景,使用_make_graph/_assert_inferred辅助函数做参数化扫描(tests/python/shape_inference_test.py 的test_transpose是典型范例):
@pytest.mark.parametrize("version", all_versions_for("OpName")) def test_opname(self, version) -> None: graph = self._make_graph( [("X", TensorProto.FLOAT, (2, 3, 4))], [make_node("OpName", ["X"], ["Y"], attr_name=attr_value)], [], ) self._assert_inferred( graph, [make_tensor_value_info("Y", TensorProto.FLOAT, expected_shape)], opset_imports=[helper.make_opsetid(ONNX_DOMAIN, version)], )单次固件测试:优先使用 onnxtxt
对于一次性固件——凡是带属性、子图(body subgraphs)或非平凡类型信息的场景——优先采用 onnxtxt skill 提供的 parser 式固件。该 skill 还覆盖了 C++unk__*物化问题(针对自由维度的已知坑)。
测试覆盖清单
SKILL 文档明确要求测试覆盖以下类别:
- 已知形状:输入维度全部已知,验证精确输出形状;
- 部分形状(
None):部分维度未知时的行为; - 秩推断:至少推断出正确的输出维数;
- 错误场景:非法属性值、秩不匹配等;
- 广播:二元算子广播规则的推断结果;
- 属性依赖的形状:如
perm、axis等属性对输出形状的影响。
仓库测试中test_transpose_preexisting_incorrect_shape、test_transpose_preexisting_incorrect_type、test_transpose_incorrect_repeated_perm(tests/python/shape_inference_test.py)正是"错误场景"类别的实例:分别验证既有但错误的形状/类型声明会被纠正、重复的perm值会抛出推断错误。
八、健壮性五条军规
实现健壮推断的规则(SKILL 文档核心结论,与官方文档"常见错误规避"一节相互印证):
- 始终先调用
hasNInputShapes(ctx, n)再访问形状——输入形状可能缺失,缺失时应按"秩未知的动态张量"处理;源码中getInputShape在形状缺失时会直接fail_shape_inference(onnx/defs/shape_inference.h),因此先检查是避免误报的前提。 - 使用
dim_value()前必须检查has_dim_value()——维度可能没有静态已知值。 - 优雅处理未知维度——保持不设置(leave unset),而不是失败。
- 至少提供秩推断——即使无法给出精确维度,也要给出正确的输出维数。
- 尽可能传播符号维度(
dim_param)——保持未知符号的一致性,使M与其他M保持同一含义。
九、改动完成后的验证流程
SKILL 文档给出的收尾流程,与仓库实际的 CI 与文档生成链路一致:
# 运行新增/修改的推断测试(-k 按测试名过滤,-x 失败即停) pytest tests/python/shape_inference_test.py -k "test_opname" -x # 重新生成算子文档(推断相关的文档字符串变化会反映到 docs/Operators.md) python onnx/defs/gen_doc.py # 运行全量 lint 检查(--output oneline 输出单行摘要) lintrunner -a --output oneline十、实战小结:一次完整实现路径
将以上内容串成一条可直接照做的实现路径:
- 在
onnx/defs/<domain>/defs.cc对应算子的 schema 上,通过.TypeAndShapeInferenceFunction(...)注册推断函数;复杂逻辑定义为命名静态函数。 - 先处理类型:能由类型约束自动传播就不写代码;需要显式时用
propagateElemTypeFromInputToOutput;输出类型由属性决定时用propagateElemTypeFromAttributeToOutput或propagateElemTypeFromDtypeToOutput(见 onnx/defs/shape_inference.h 附近)。 - 再处理形状:按算子类别选择模式——一元逐元素用
propagateShapeAndTypeFromFirstInput,二元广播用bidirectionalBroadcastShapeInference,变形状算子用getRepeatedAttribute读取属性并逐维写入getOutputShape,涉及维度统一时用unifyInputShape/updateOutputShape,需要算术时用Dim重载运算符。 - 遵守健壮性五条军规,尤其保证未知形状/维度不崩溃、至少给出秩推断。
- 在 tests/python/shape_inference_test.py 中补充参数化测试与错误用例,覆盖已知/部分/秩/错误/广播/属性依赖六类场景。
- 运行
pytest、python onnx/defs/gen_doc.py、lintrunner -a收尾验证。
通过这条路径,你既能看懂仓库内每个既有算子推断函数的写法(从一元到Loop/Scan这类异构变参复杂算子),也能为自定义算子写出同样健壮、可测试的类型与形状推断实现。
【免费下载链接】onnxOpen standard for machine learning interoperability项目地址: https://gitcode.com/gh_mirrors/onn/onnx
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考