ONNX 算子形状推断(Shape Inference)实现指南:从 TypeAndShapeInferenceFunction 到测试验证

ONNX 算子形状推断(Shape Inference)实现指南:从 TypeAndShapeInferenceFunction 到测试验证 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.hPython 测试tests/python/shape_inference_test.py注册方式是在算子 schema 上调用OpSchema OpSchema::TypeAndShapeInferenceFunction(InferenceFunction inferenceFunction);InferenceFunction与核心接口结构体InferenceContext均定义于 onnx/defs/shape_inference.hInferenceContext是传入推断函数的上下文负责读取算子输入信息并写入推断结果。图级入口是shape_inference::InferShapes(ModelProto m, const ISchemaRegistry* schema_registry)它在原模型上原地标注形状信息C APIPython 侧则有对应的 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::vectorint64_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当多个输入维度期望相同、或输入维度需传播到特定输出维度时使用参考RoiAlignunifyInputShape/unifyInputShapePrefix在unifyInputDim之上构建的声明式高层工具一次调用统一输入的全部或前缀维度适合简单场景复杂场景仍需逐个unifyInputDimhasInputShape(ctx, n)单输入形状检查hasNInputShapes的基础。这些工具都对缺失的形状/维度做了安全处理。从源码看 hasNInputShapes 的语义onnx/defs/shape_inference.h 的实现显示hasInputShape需要同时满足三个条件ctx.getNumInputs() n、ctx.getInputType(n)非空、且该类型支持 tensor/sparse tensor/sequence/optional 递归确实携带形状。hasNInputShapes则对前 n 个输入逐一检查。而propagateShapeonnx/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_nameattr_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 还覆盖了 Cunk__*物化问题针对自由维度的已知坑。测试覆盖清单SKILL 文档明确要求测试覆盖以下类别已知形状输入维度全部已知验证精确输出形状部分形状None部分维度未知时的行为秩推断至少推断出正确的输出维数错误场景非法属性值、秩不匹配等广播二元算子广播规则的推断结果属性依赖的形状如perm、axis等属性对输出形状的影响。仓库测试中test_transpose_preexisting_incorrect_shape、test_transpose_preexisting_incorrect_type、test_transpose_incorrect_repeated_permtests/python/shape_inference_test.py正是错误场景类别的实例分别验证既有但错误的形状/类型声明会被纠正、重复的perm值会抛出推断错误。八、健壮性五条军规实现健壮推断的规则SKILL 文档核心结论与官方文档常见错误规避一节相互印证始终先调用hasNInputShapes(ctx, n)再访问形状——输入形状可能缺失缺失时应按秩未知的动态张量处理源码中getInputShape在形状缺失时会直接fail_shape_inferenceonnx/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),仅供参考