资讯详情

CANN opbase 广播形状推导:BroadcastInferShape 接口深度解析与实践

📅 2026/9/19 14:35:22 | 华诺云谱 👁 阅读
CANN opbase 广播形状推导:BroadcastInferShape 接口深度解析与实践
CANN opbase 广播形状推导BroadcastInferShape 接口深度解析与实践【免费下载链接】opbase本项目是CANN算子库的基础框架库为算子提供公共依赖文件和基础调度能力。项目地址: https://gitcode.com/cann/opbaseBroadcastInferShape是 CANN opbase 算子基础框架库cann/opbase中op::shape_utils工具集提供的广播形状推导接口用于在算子开发尤其是 shape 推导 / infershape 阶段中根据 NumPy 广播规则从两个输入 shape 推导出广播后的输出 shape。本文基于官方接口文档 BroadcastInferShape.md结合 接口声明、核心实现 与 单测/系统测试完整讲解其原型、广播判定算法、调用方式与工程实践帮助算子开发者快速掌握在多输入算子中正确使用广播形状推导的方法。1. 接口定位shape_utils 工具集中的广播能力BroadcastInferShape位于op::shape_utils工具集与 shape_utils.md 中列出的ToShape、ToShapeVector、ToContiguousStrides、CheckBroadcastShape等接口共同构成算子 shape 处理的公共工具层。它与同族的CheckBroadcastShape职责互补CheckBroadcastShape只做合法性校验判断两个 shape 是否满足广播关系不产出广播结果BroadcastInferShape在校验的同时推导广播结果满足条件时把广播后的 shape 写入broadcastShape输出参数不满足时返回false。在算子 shape 推导infershape场景中凡是涉及两个及以上输入按元素对齐element-wise的算子如 Add、Mul、Sub、Div 等都可以直接调用本接口生成输出 shape无需自行实现逐维比较逻辑。2. 函数原型与参数说明接口在 include/nnopbase/opdev/shape_utils.h 中声明于op命名空间头文件第 27 行原型如下namespace op { bool BroadcastInferShape(const op::Shape self, const op::Shape other, op::Shape broadcastShape); } // namespace op各参数含义如下表所示参数输入/输出说明self输入第一组 shape。other输入第二组 shape。broadcastShape输出self和other经过 broadcast 后推导出的 shape结果写入该参数。参数类型op::Shape是算子开发中的通用 shape 表示头文件通过#include exe_graph/runtime/shape.h引入底层 shape 实现官方文档示例中出现的gert::Shape即与之对应的 shape 类型。op::Shape提供GetDimNum()维数、GetDim(i)取第 i 维长度、SetDim(i, v)/SetDimNum(n)写入维度、AppendDim(v)追加维度等操作详情可参考同目录 ToShape 等接口文档。3. 返回值与错误处理当self与other满足广播关系时返回truebroadcastShape中为推导出的广播 shape当两者不满足广播关系时返回false。需要特别注意的是广播失败时实现还会通过OP_LOGE_FOR_INVALID_ARGUMENT_TENSOR_INPUT_SHAPE记录一条带完整 shape 信息的错误日志对应 error_code 文档 EZ1007-Invalid_Argument_Tensor_Input_Shape.md日志中会带上较大 shape、较小 shape 以及发生冲突的维度值便于在算子推理失败时快速定位是哪两个输入、哪一维不满足广播条件。调用方拿到false后应向上传播错误例如在 infershape 回调中返回失败不要继续使用broadcastShape中的内容。4. 广播判定算法从源码看逐维比对过程BroadcastInferShape的完整实现在 src/nnopbase/common/utils/shape_utils.cpp 第 118149 行。其核心逻辑与 NumPy 广播规则一致可归纳为以下三步4.1 确定较大 shapeconst auto largerDimShape selfShapeLen otherShapeLen ? self : other; const auto smallerDimShape largerDimShape self ? other : self; auto lenSub largerDimNum - smallerDimNum;秩维数更大的 shape 作为largerDimShape另一个作为smallerDimShapelenSub是两者秩的差值。由于广播是从最后一个维度最右侧开始向前对齐因此实现选择从尾部逐维比较。4.2 单维广播判定BroadcastDimstatic bool BroadcastDim(int64_t dim1, const int64_t dim2) { if (dim1 dim2) { return true; } if ((dim1 ! 1) (dim2 ! 1)) { return false; } dim1 (dim1 1) ? dim2 : dim1; return true; }源码注释中用一张真值表完整刻画了单维广播规则行是dim1列是dim2表中值为broadcast(dim1, dim2)的结果E表示不合法即返回 falsedim1 \ dim201d2000E101d2d1Ed1E即两维相等 → 直接通过某一维为 1 → 广播为另一维的值两维都既不等又不为 1 → 广播失败。注意该函数会原地改写dim1使其变为广播后的维度这正是BroadcastInferShape能直接产出结果的机制。4.3 写入广播结果broadcastShape.SetDimNum(largerDimNum); for (size_t i smallerDimNum; i 0; i--) { // 从尾部逐维调用 BroadcastDim 并写回 broadcastShape ... broadcastShape.SetDim(lenSub i - 1, dim1); } for (size_t i 0; i lenSub; i) { broadcastShape.SetDim(i, largerDimShape.GetDim(i)); }广播结果的秩取两者中的较大值SetDimNum(largerDimNum)尾部smallerDimNum维经过BroadcastDim逐维判定后写入结果对应位置较大 shape 多出来的前lenSub维即两者秩不同时左侧多出的维直接原样拷贝到结果中。4.4 典型推导示例按文档给出的示例[1, 10]与[2, 1]广播尾部对齐后10 vs 1 → 101 vs 2 → 2广播结果为[2, 10]。这与源码逐维比较、维度为 1 的一方被拉伸为另一方维度的过程完全对应。5. 调用示例完整可运行代码以下示例来自官方文档BroadcastInferShape.md构造 shape 为[2, 1]和[2, 10]的两个 Shape 对象并推导广播 shape// Generate two shape objects with shapes [2, 1] and [2, 10] to obtain the broadcast shape. void Func() { gert::Shape shapeA; shapeA.AppendDim(1); shapeA.AppendDim(2); gert::Shape shapeB; shapeB.AppendDim(10); shapeB.AppendDim(2); gert::Shape shapeBrc; bool isBrc BroadcastInferShape(shapeA, shapeB, shapeBrc); }注意AppendDim逐次追加维度后shapeA实际为[2, 1]先 1 后 2shapeB实际为[2, 10]先 10 后 2。isBrc为trueshapeBrc为[2, 10]。建议在实际工程中使用op::Shape并采用聚合初始化写法见下文测试用例例如#include opdev/shape_utils.h #include op_ctx_def.h op::Shape selfShape({2, 1}); op::Shape otherShape({2, 10}); op::Shape broadcastShape; bool ok op::BroadcastInferShape(selfShape, otherShape, broadcastShape); // ok true, broadcastShape [2, 10]6. 测试用例验证规则覆盖与边界场景仓库在 UT 与 ST 两层对BroadcastInferShape都有完整覆盖可作为接口行为的事实依据单元测试 tests/nnopbase/ut/composite_op/test_shape_utils.cpp 的TestBroadcastInferShape第 56 行起系统测试 tests/nnopbase/st/composite_op/test_shape_utils.cpp 的TestBroadcastInferShape第 41 行起。从 UT 用例可以归纳出接口支持的全部典型场景输入 shape输入 shape广播结果 / 返回值覆盖要点[2][2, 2][2, 2]true秩不同1 维 vs 2 维左侧补维[2][2, 1][2, 2]true尾部对齐 维度 1 广播[2, 2][2, 5]false尾部两维均不为 1 且不等[2, 1][2, 1][2, 1]true完全相等[2, 2, 5][2, 1, 5][2, 2, 5]true3 维场景中间维 1 广播[2, 2][2, 2, 5]false秩不同且对齐维冲突[2][2, 2, 3, 2][2, 2, 3, 2]true大秩 尾部对齐[2]与[3, 2]对齐[2, 3, 2][2, 2, 3, 2][2, 2, 3, 2]true秩不同、左侧补维 尾部对齐测试中每个用例在调用前都会把outShape的每一维先置为-3见ClearShape辅助函数调用成功后用op::ToString(outShape)与期望 shape 的字符串比对从而验证broadcastShape确实被完整、正确地写回。这对使用者的启示是调用接口前无需手动初始化输出 shape实现内部会先SetDimNum再逐维SetDim保证输出被完整覆盖。7. 使用约束与工程建议约束说明官方文档对BroadcastInferShape标注的约束为“无”Restrictions: None即没有参数类型、维数上限等额外限制广播规则与 NumPy 完全对齐可放心用于动态 shape 场景。仅用于判断推导若只需要判断两个 shape 能否广播而不关心结果可使用开销更小的 CheckBroadcastShape需要结果时优先使用本接口避免“先校验再手动推导”的重复实现。失败即返回接口返回false时已记录含 shape 明细的错误日志调用方在 infershape 回调中应直接向上返回失败不要继续使用输出参数。头文件包含使用前需包含 include/nnopbase/opdev/shape_utils.h接口实现在src/nnopbase/common/utils/shape_utils.cpp中随libopbase基础库一同编译链接见 src/nnopbase/CMakeLists.txt。8. 相关接口与延伸阅读shape_utils.mdshape 工具集总览含ToShape、ToShapeVector、ToContiguousStrides等配套接口CheckBroadcastShape.md只校验不推导的同族接口src/nnopbase/common/utils/shape_utils.cppBroadcastInferShape与CheckBroadcastShape的完整实现二者共用BroadcastDim单维判定逻辑tests/nnopbase/ut/composite_op/test_shape_utils.cppUT 用例覆盖上表全部广播场景广播失败对应错误码说明见 EZ1007-Invalid_Argument_Tensor_Input_Shape.md。以上接口、实现与测试均为当前仓库 cann/opbase 中真实存在的内容读者可直接在仓库中对照源码与用例深入学习。【免费下载链接】opbase本项目是CANN算子库的基础框架库为算子提供公共依赖文件和基础调度能力。项目地址: https://gitcode.com/cann/opbase创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
📝

华诺云谱内容团队

资深建站顾问 · 行业研究员

10年+企业数字化服务经验,专注智能建站、SEO优化与品牌营销,持续输出建站技巧、行业洞察与营销干货,已帮助5000+企业实现数字化增长。

你可能需要的服务

订阅华诺云谱资讯周报

每周一封,精选建站技巧、SEO与营销干货,直达邮箱。已有 8,000+ 企业主订阅,助你少走弯路。