ONNX 算子形状推断(Shape Inference)实现指南:从 TypeAndShapeInferenceFunction 到测试验证
2026/9/20 23:15:55 网站建设 项目流程

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已定义 ⇒ 秩已知;
  • 每个Dimensiondim_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中写显式类型推断逻辑:

  • 输出类型由属性决定(如Castto属性指定输出元素类型);
  • 输出类型与所有输入类型都不同,且无法用共享类型约束变量表达;
  • 算子使用异构(heterogeneous)变参输入/输出。

同构与异构变参

同构/异构标志只适用于变参(repeated)输入或输出:

  • 同构(默认):所有重复参数类型必须相同,类型约束变量约束它们一致,框架自动强制并传播;
  • 异构:每个重复参数可以类型不同,类型约束变量只描述"允许的类型集合"。LoopScan等算子使用该模式(其携带状态变量可混合类型)。

使用异构变参时,推断函数必须为每个参数显式传播类型,框架无法自动完成。

形状推断则几乎总是需要显式逻辑,因为输出形状通常取决于输入形状、属性或两者共同决定。

四、三种核心推断模式(含完整代码)

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() > nctx.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_valuedim_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:部分维度未知时的行为;
  • 秩推断:至少推断出正确的输出维数;
  • 错误场景:非法属性值、秩不匹配等;
  • 广播:二元算子广播规则的推断结果;
  • 属性依赖的形状:如permaxis等属性对输出形状的影响。

仓库测试中test_transpose_preexisting_incorrect_shapetest_transpose_preexisting_incorrect_typetest_transpose_incorrect_repeated_perm(tests/python/shape_inference_test.py)正是"错误场景"类别的实例:分别验证既有但错误的形状/类型声明会被纠正、重复的perm值会抛出推断错误。

八、健壮性五条军规

实现健壮推断的规则(SKILL 文档核心结论,与官方文档"常见错误规避"一节相互印证):

  1. 始终先调用hasNInputShapes(ctx, n)再访问形状——输入形状可能缺失,缺失时应按"秩未知的动态张量"处理;源码中getInputShape在形状缺失时会直接fail_shape_inference(onnx/defs/shape_inference.h),因此先检查是避免误报的前提。
  2. 使用dim_value()前必须检查has_dim_value()——维度可能没有静态已知值。
  3. 优雅处理未知维度——保持不设置(leave unset),而不是失败。
  4. 至少提供秩推断——即使无法给出精确维度,也要给出正确的输出维数。
  5. 尽可能传播符号维度(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

十、实战小结:一次完整实现路径

将以上内容串成一条可直接照做的实现路径:

  1. onnx/defs/<domain>/defs.cc对应算子的 schema 上,通过.TypeAndShapeInferenceFunction(...)注册推断函数;复杂逻辑定义为命名静态函数。
  2. 先处理类型:能由类型约束自动传播就不写代码;需要显式时用propagateElemTypeFromInputToOutput;输出类型由属性决定时用propagateElemTypeFromAttributeToOutputpropagateElemTypeFromDtypeToOutput(见 onnx/defs/shape_inference.h 附近)。
  3. 再处理形状:按算子类别选择模式——一元逐元素用propagateShapeAndTypeFromFirstInput,二元广播用bidirectionalBroadcastShapeInference,变形状算子用getRepeatedAttribute读取属性并逐维写入getOutputShape,涉及维度统一时用unifyInputShape/updateOutputShape,需要算术时用Dim重载运算符。
  4. 遵守健壮性五条军规,尤其保证未知形状/维度不崩溃、至少给出秩推断。
  5. 在 tests/python/shape_inference_test.py 中补充参数化测试与错误用例,覆盖已知/部分/秩/错误/广播/属性依赖六类场景。
  6. 运行pytestpython onnx/defs/gen_doc.pylintrunner -a收尾验证。

通过这条路径,你既能看懂仓库内每个既有算子推断函数的写法(从一元到Loop/Scan这类异构变参复杂算子),也能为自定义算子写出同样健壮、可测试的类型与形状推断实现。

【免费下载链接】onnxOpen standard for machine learning interoperability项目地址: https://gitcode.com/gh_mirrors/onn/onnx

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询