ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

ONNX 符号形状推断与部分数据传播提案解析:从静态 Shape Inference 到符号维推理的演进

ONNX 符号形状推断与部分数据传播提案解析:从静态 Shape Inference 到符号维推理的演进 人工智能深度学习机器学习【免费下载链接】onnxOpen standard for machine learning interoperability项目地址https://gitcode.com/gh_mirrors/onn/onnx点击查看免费下载本文围绕 ONNX 官方技术提案 docs/proposals/0005-SymbolicShapeInfProposal.md 展开该提案2021-09-16 发起、状态 accepted已在 ONNX 1.10 中落地实现对应 PR 3518、3551、3593、3580。它系统性地分析了 ONNX 静态形状推断Shape Inference在动态形状、符号维场景下的三大局限并给出符号生成与传播 部分数据计算与传播两套解决方案。读完本文你将理解 ONNX 图级形状推断的内部工作方式、DataPropagationContext与PartialDataPropagationFunction的接口设计、SymbolTable的符号生成机制以及如何通过 Python API 的data_prop参数开启数据传播从而让 Reshape 等动态形状算子不再阻断后续节点的形状推断。背景ONNX Shape Inference 是什么为什么要改进它ONNX 在 onnx/shape_inference/implementation.cc 中提供了针对 ONNX 图的形状推断实现。形状推断由算子级的 shape inference 函数定义在算子 schema 上驱动其价值在于无需真正在会话session中启动模型执行即可获取张量的形状信息。这种静态形状推断可用于在运行时之前捕获明显的错误消除那些必然通过的运行时检查改进运行时的静态内存规划static memory planning提升模型可视化的体验。对于 PyTorch exporter 以及 Nuphar 这类基于编译器的执行提供方形状推断甚至是硬性要求——rank inference秩推断是最低要求它们无法处理未知形状。Pre ONNX 1.10 时代 Shape Inference 的三大局限原提案明确指出ONNX 1.10 之前的形状推断不保证完备。虽然尽可能回退到 rank inference但仍有若干场景连 rank 都无法推断导致推断中断动态行为阻断推断流某些动态行为会阻塞形状推断的传递推断随之停止。典型例子是Reshape 到一个动态计算的形状——只有真正执行计算才能知道目标形状静态推断到此卡死。只支持常量与简单变量不支持符号算术形状推断只能处理常量如(5, 2)和简单变量既无法表达含变量的算术表达式也无法生成符号。例如对形状为(5, 2)与(7, 2)的张量做 Concat可以推断出(12, 2)但对(5, 2)与(N, 2)做 Concat结果只能是(?, 2)——这里的?表示既没有 dim value 也没有 dim param的维而不是N5这样的表达式也不会生成新符号(M, 2)。此时形状传播停止。并非所有算子都有 shape inference 实现遇到未实现 shape inference 的算子时推断停止也存在没有做 rank inference 回退的情况。提案注明这是持续修复中的问题本文档不聚焦此局限。目标与非目标目标修复两类核心场景下的形状推断缺口——分支中进行的形状计算对应局限 1存在符号维的场景对应局限 2。通过修复期望达成解除 PyTorch exporter 因缺少形状信息而无法导出模型的阻塞改善运行时的静态内存规划支持在运行时之外预分配输出缓冲区使其生命周期可由调用方自行管理。非目标不把符号表达式加入 ONNX 标准虽然符号表达式能显著减少引入的符号数量、在某些特殊场景提供更确定的形状计算但其代价是增加复杂度因此本阶段不做留待未来迭代考虑不为旧算子集开启数据计算与传播详见提案正文说明。提案同时注明该工作同样惠及 Nuphar但当时没有让 Nuphar 迁移到该方案的计划。术语节点级与图级形状推断形状推断可拆分为两个层次节点级形状推断Node level shape inference算子专属的 shape inference 函数随算子 schema 一并定义。它负责根据当前节点的输入类型/形状推断输出类型/形状。图级形状推断Graph-level shape inference更上层的逻辑遍历整张图从节点级 shape inference 函数取得推断结果再决定如何将推断形状与已有形状合并使其可供下游节点使用。这一分层在实现中也清晰可见节点级逻辑位于各算子的defs中而图级遍历与合并逻辑集中在 onnx/shape_inference/implementation.cc核心类ShapeInferenceImplBase中。提案核心三方面扩展提案提出扩展现有 shape inference使其支持符号生成与传播Symbol generation and propagation部分数据计算与传播Partial data computation and propagation扩展 Shape 算子使其能生成形状的切片slice of the shape以简化形状计算。符号生成与传播让 ? 变成可判等的 K核心问题两个 ? 不能判等以 Concat 为例若输入形状为[M]与[N]当前形状推断返回[?]。假设 Concat 的输出X依次经过两个一元算子Op1()、Op2()得到Y、Z则?会一路传播Y、Z的推断形状都是[?]。但我们无法推断出 X、Y、Z 具有相同形状——因为两个 ? 不能被判定为相等。这直接导致运行时无法利用这三个张量形状相同的事实去复用内存。解决方案图级符号表按提案推断形状中的 ? 将由图级形状推断替换为新的唯一符号。继续上面的例子Concat 产出[?]后图级推断将其替换为[K]随后下游推断即可得出 X、Y、Z 具有相同形状[K]。运行时据此可对这组张量复用内存。这一机制在源码中已有完整落地。图级推断维护一个符号表SymbolTable接口定义见 onnx/defs/shape_inference.hSymbolTable提供addFromGraph(const GraphProto g)将主图或子图中已存在的符号加入符号表避免新符号与旧符号冲突以及createNew()默认前缀unk__来生成不与任何已有符号重复的新符号具体实现SymbolTableImpl位于 onnx/shape_inference/implementation.h在 onnx/shape_inference/implementation.cc 中GenerateSymbolicShape()遍历推断形状的每一维凡既无dim_value也无dim_param的维就调用symbol_table.createNew()为其设置dim_paramMaterializeSymbolicShape()则递归处理 tensor、sparse tensor、sequence、optional、map 等各类类型的形状。图级处理开始时还会通过TraverseGraphsToAddExistingSymbols先把图中已有符号登记进符号表。需要特别强调的是符号生成发生在图级形状推断层因此所有模型无论旧算子集还是最新算子集版本都能从中受益。部分数据计算与传播让动态形状不再阻断推断动机当形状输入是动态计算得到时例如 Reshape 的目标形状来自上游某个算子的输出Reshape 节点之后的形状推断会停止。提案的解法是在形状推断期间把参与形状计算的数据实际计算出来并交给 Reshape 节点使用。之所以叫部分partial是因为这种数据计算只服务于形状计算并非完整的算子 kernel 实现相应地只会为有限的一组算子实现数据传播函数。虽然未来会逐步扩大覆盖面但有些算子如 LSTM、卷积类、池化类算子等永远不会添加数据传播函数——它们本就不参与形状计算。第一阶段算子清单提案给出第一批实现数据传播的算子这些算子通常用于形状计算OpsAddSubMulCastConcatGatherReshapeShapeSliceSizeSqueezeUnSqueeze接口设计扩展 OpSchema提案将OpSchema类扩展加入一个可选的PartialDataPropagationFunction与既有的TypeAndShapeInferenceFunction并列。该函数为算子提供数据计算逻辑随后由图级形状推断将计算结果传播给下游算子。调用时机上PartialDataPropagationFunction会在节点的TypeAndShapeInference执行之后被图级形状推断调用——因为部分数据计算需要先拿到输出形状。提案同时新增了DataPropagationContext接口供PartialDataPropagationFunction访问给定节点数据传播所需的全部信息并写入计算得到的数据。提案给出了如下接口草图示意using DataPropagationFunction std::functionvoid(DataPropagationContext) class OpSchema final { public: OpSchema PartialDataPropagationFunction(DataPropagationFunction dataPropagationFunction) { partial_data_propagation_function_ std::move(dataPropagationFunction); return *this; } DataPropagationFunction GetDataPropagationFunction() const { return partial_data_propagation_function_ ? partial_data_propagation_function_ : dummyDataPropagator; } }; // Operator schema example ONNX_OPERATOR_SET_SCHEMA( Shape, 13, OpSchema() .SetDoc() .Input(0, data, An input tensor., T, ...) .Output(0, shape, Shape of the input tensor, T1, ...) .TypeConstraint(T, OpSchema::all_tensor_types()) .TypeConstraint(T1, {tensor(int64)}) .TypeAndShapeInferenceFunction([](InferenceContext ctx) { // ... }) .PartialDataPropagationFunction([](DataPropagationContext ctx) { TensorShapeProto tp; // compute output data for shape operator // add computed data to DataPropagationContext for propagating it downstream ctx.addOutputData(0, std::move(tp)); }));新旧算子集的差异符号生成在图级形状推断层完成因此旧算子集模型同样受益数据计算与传播绑定在OpSchema上发生在节点级初期只添加到最新的算子 schema。旧 schema 可在后续按需逐个扩展以支持高优先级场景——这意味着一开始旧算子集模型不会因该增强而获得形状推断改进。源码中的实际落地情况1. 选项开关ShapeInferenceOptionsonnx/defs/shape_inference.h 中定义了ShapeInferenceOptions三个字段分别是check_type是否检查类型约束error_modestrict_mode错误处理模式enable_data_propagation是否开启数据传播默认关闭。2. 数据传播的实际调用链在 onnx/shape_inference/implementation.cc 中图级推断对每个节点依次执行TypeAndShapeInferenceFunction→ 更新输出类型/形状 →ProcessConstant跟踪常量值以辅助后续节点→若enable_data_propagation开启且该算子 schema 声明了数据传播函数则构造DataPropagationContextImpl并调用schema-GetDataPropagationFunction()(data_propagation_ctx)。也就是说数据传播严格发生在类型与形状推断之后与提案设计一致。3. schema 侧接口onnx/defs/schema.h 中PartialDataPropagationFunction(DataPropagationFunction)注册数据传播函数GetDataPropagationFunction()获取函数若未注册则返回dummyDataPropagationFunction空操作has_data_propagation_function()判断 schema 是否注册了数据传播函数。4. 数据传播的工具函数onnx/defs/data_propagators.h 提供了若干可复用的数据传播工具例如appendDimToTensorShapeProto从输入数据中按索引支持负索引取一维追加到输出形状 protoaxisIsZero判定axis属性是否为 0支持负 axis需借助输入秩信息PropagateShapeDataFromInputToOutput把输入的形状数据直接复制传播到输出Identity 类语义GatherOp13DataPropagatorGather-13 的数据传播实现——仅当 axis 为 0 且输入数据、索引数据均为已知常量时按索引逐个取出对应的维并写入输出数据。这些函数被 onnx/defs/tensor/defs.cc、onnx/defs/math/defs.cc 等算子定义文件中的PartialDataPropagationFunction使用。5. Python 侧 APIPython 侧入口为 onnx/shape_inference.pyonnx.shape_inference.infer_shapes(model, check_typeFalse, strict_modeFalse, data_propFalse, strict_mode_valueNone)对 ModelProto 做形状推断onnx.shape_inference.infer_shapes_path(model_path, output_path, check_typeFalse, strict_modeFalse, data_propFalse)直接对模型文件做推断并写回输出文件。其中data_prop参数即对应 C 层的enable_data_propagationdata_propTrue时为有限的算子开启数据传播以完成形状计算默认False。strict_mode为严格模式开启后遇错会直接抛异常。相应地onnx/onnx_cpp2py_export/shape_inference.pyi 中给出了对应的类型签名。6. 测试佐证仓库中提供了专门的测试来验证这两项能力tests/python/data_propagation_test.py数据传播专项测试tests/python/shape_inference_test.py形状推断综合测试含符号维相关用例tests/cpp/shape_inference_test.ccC 侧形状推断测试。特殊情形Edge Cases的处理策略广播与符号维当两个未知维 M 与 N 之间发生广播时不能推断 M 与 N 必然相等——运行时语义允许其中一个符号取值为 1另一个取非 1 值。因此将 M 与 N 合并视为同一值是潜在不健全unsound的。此时策略是为输出形状生成一个新符号形状推断继续进行。推断形状与已有形状不匹配推断形状与已有形状可能不一致。虽然此时令形状推断失败看起来是正确做法但并非总是实际可行。默认行为遇到此类情况时形状推断失败但调用方可以选择用推断类型覆盖已有类型——启用该选项后推断将以推断类型继续执行。符号维 数据传播当形状含有符号维时会尽量将其传播到下游但在对符号维执行了某些算术运算的场景下会创建新符号并传播新符号对应提案中不引入符号表达式的设计取舍。输出形状依赖输入数据某些节点如NonZero的输出形状取决于输入数据本身此时无法完整推断形状策略是基于推断出的 rank 创建新的符号形状形状推断继续而不是中断。总结0005-SymbolicShapeInfProposal.md是 ONNX 形状推断能力演进的关键里程碑它以图级符号生成 节点级部分数据传播的组合绕开了在 ONNX 标准中引入符号表达式的高复杂度路径使动态形状尤其是动态 Reshape场景下的形状推断得以继续为 PyTorch exporter、静态内存规划、输出缓冲预分配等下游能力扫清了障碍。该提案在 ONNX 1.10 中完整落地其核心机制——SymbolTable符号生成、PartialDataPropagationFunction/DataPropagationContext接口、enable_data_propagation选项——至今仍可在 onnx/defs/shape_inference.h、onnx/defs/schema.h、onnx/defs/data_propagators.h 与 onnx/shape_inference/implementation.cc 中直接查阅验证。对于希望在自有推理引擎或导出工具链中复用这套能力的开发者可直接调用onnx.shape_inference.infer_shapes(..., data_propTrue)并在 C 侧通过ShapeInferenceOptions的enable_data_propagation字段开启对应能力。赞分享人工智能深度学习机器学习【免费下载链接】onnxOpen standard for machine learning interoperability项目地址https://gitcode.com/gh_mirrors/onn/onnx点击查看免费下载相关推荐tsParticles 无限符号∞形状插件全解析tsparticles/shape-infinity 从版本演进、配置到 Canvas 绘制源码tsParticles 无限符号∞形状插件全解析tsparticles/shape infinity 从版本演进、配置到 Canvas 绘制源码 tsP前端Mole命令行工具入门教程30秒快速上手、命令速查表与参数详解Mole命令行工具入门教程30秒快速上手、命令速查表与参数详解 你是不是也遇到过这种时刻打开 Finder磁盘空间快不足的提示红得刺眼翻了一圈却找不CLI开发工具运维观测PyPTO Tensor 构造函数完全指南从静态形状到动态符号化形状的创建实战PyPTO Tensor 构造函数完全指南从静态形状到动态符号化形状的创建实战 Tensor 是 PyPTOParallel Tensor/Tile Ope人工智能编译器模型编译高性能计算深度学习CANN上一篇5分钟上手Medical SAM Adapter从环境搭建到首次分割完整教程下一篇解决llava-calm2-siglip常见问题新手必看的故障排除指南创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表