CANN opbase 算子 Shape 广播关系校验:CheckBroadcastShape 使用指南与源码解析

发布时间:2026/9/18 14:19:45
CANN opbase 算子 Shape 广播关系校验:CheckBroadcastShape 使用指南与源码解析
CANN opbase 算子 Shape 广播关系校验CheckBroadcastShape 使用指南与源码解析【免费下载链接】opbase本项目是CANN算子库的基础框架库为算子提供公共依赖文件和基础调度能力。项目地址: https://gitcode.com/cann/opbase在 CANN 算子库基础框架库opbase中算子开发者经常需要判断两个张量Tensor的 shape 之间是否满足广播Broadcast关系例如元素级Elementwise算子的输入 shape 校验与输出 shape 推导。op::CheckBroadcastShape正是 opdev 对外提供的 shape 工具函数之一用于快速校验两组 shape 是否满足 NumPy 风格的广播规则。本文将结合函数原型、参数说明、调用示例深入 源码实现 与单元/系统测试用例帮助你掌握该接口的语义、边界行为及在算子开发中的典型用法。功能说明CheckBroadcastShape用于校验两个 shape 之间是否满足广播broadcast关系。广播规则与 NumPy 的广播规则一致即从两个 shape 的最右侧维度开始向左逐维对齐比较两个维度相等则该维度可以广播其中一个维度为 1则该维度也可以广播维度为 1 的一方会被拉伸到与另一方相同两个维度都不相等且都不为 1则不满足广播关系。例如[2, 1]与[2, 10]最右侧维度1与10一方为 1可广播次右侧维度2与2相等满足广播关系[2, 2]与[2, 10]最右侧维度2与10不相等且均不为 1不满足广播关系。当两个 shape 的维度数不同时维度数较少的 shape 在其左侧补 1即“右对齐”后再逐维比较这也是广播判断与推导的核心前提。函数原型bool CheckBroadcastShape(const op::Shape self, const op::Shape other);接口声明位于 include/nnopbase/opdev/shape_utils.h实现在 src/nnopbase/common/utils/shape_utils.cpp均位于op命名空间内。其中op::Shape是gert::Shape的别名op::ShapeVector是FVectorint64_t, MAX_DIM_NUM的别名参见 include/nnopbase/opdev/common_types.h算子侧可直接使用gert::Shape对象作为实参传入。参数说明参数输入/输出说明self输入第一组 shape。other输入第二组 shape。两个参数均为const op::Shape 只读引用函数不会修改传入的 shape 对象可放心传入复用中的 Shape 实例。返回值说明当self与other满足广播关系时返回true否则返回false。需要注意该函数只做“是否满足广播关系”的布尔判断并不产出广播后的目标 shape。如果需要同时得到广播后的 shape可配合同文件的op::BroadcastInferShape使用详见下文“与 BroadcastInferShape 的配合”小节。约束说明无。self与other的维度数、各维度取值任意组合均可安全调用函数内部对维度数不同的情况做了右对齐处理不会越界访问。实现原理右对齐逐维比较从 源码实现 可以看出CheckBroadcastShape的判定逻辑分为三步确定长短 shape 与维度差比较self与other的GetDimNum()维度多的一方记为largerDimShape维度少的一方记为smallerDimShape两者维度差为lenSub右对齐逐维比较从smallerDimNum向 1 递减遍历即从最右侧维度向左取largerDimShape.GetDim(lenSub i - 1)与smallerDimShape.GetDim(i - 1)进行单维广播判断单维判断BroadcastDim私有辅助函数 BroadcastDim 的规则为——若两维度相等则直接通过若两者均不为 1 则失败否则将维度为 1 的一方扩展为另一方的维度后通过。其中BroadcastDim的判定逻辑在源码中以矩阵形式注释dim1为列、dim2为行dim 0 1 d2 0 0 0 E 1 0 1 d2 d1 E d1 E矩阵中0表示维度为 1、d1/d2表示大于 1 的维度、E表示不满足广播关系Error。可见只要存在一个维度的组合是(非1, 非1)且不相等整个判断就立即返回false因此该实现是短路判定的具备 O(min(dimNum)) 的时间复杂度。此外函数入口处会调用OP_LOGD打印参与广播判断的两个 shape通过op::ToString序列化为[d0, d1, ...]形式便于在调试日志中定位 shape 校验问题。调用示例以下示例生成 shape 为[2, 1]与[2, 10]的两个Shape对象校验两者是否满足广播关系// 生成shape为[2 1]和[2, 10]的两个Shape对象校验两个shape是否满足broadcast关系。 void Func() { gert::Shape shapeA; shapeA.AppendDim(1); shapeA.AppendDim(2); gert::Shape shapeB; shapeB.AppendDim(10); shapeB.AppendDim(2); bool isBrc CheckBroadcastShape(shapeA, shapeB); }示例中shapeA通过AppendDim依次追加维度得到[2, 1]shapeB得到[2, 10]由于最右侧维度1与10满足“一方为 1”的广播条件最终isBrc为true。需要注意的是该示例中的注释写作[21]与[2,10]实际追加顺序对应 shape 为[2, 1]与[2, 10]广播判断与维度追加顺序无关仅与最终的维度序列有关。测试用例验证广播判定的覆盖场景CheckBroadcastShape在单元测试与系统测试中均有覆盖测试文件分别为 tests/nnopbase/ut/composite_op/test_shape_utils.cpp 与 tests/nnopbase/st/composite_op/test_shape_utils.cpp两处用例完全一致TEST_F(TestShapeUtils, TestCheckBroadcastShape) { op::Shape shape1({2, 2}); op::Shape shape2({2}); op::Shape shape3({2, 1}); op::Shape shape4({2, 1}); op::Shape shape5({2, 5}); op::Shape shape6({2, 2, 5}); op::Shape shape7({2, 1, 5}); EXPECT_TRUE(op::CheckBroadcastShape(shape1, shape2)); // [2,2] 与 [2] - 右侧对齐可广播 EXPECT_TRUE(op::CheckBroadcastShape(shape2, shape1)); // 顺序交换结果一致对称性 EXPECT_TRUE(op::CheckBroadcastShape(shape2, shape3)); // [2] 与 [2,1] - 左侧补1 EXPECT_FALSE(op::CheckBroadcastShape(shape1, shape5)); // [2,2] 与 [2,5] - 2与5均非1不满足 EXPECT_TRUE(op::CheckBroadcastShape(shape3, shape4)); // [2,1] 与 [2,1] - 完全相同 EXPECT_TRUE(op::CheckBroadcastShape(shape6, shape7)); // [2,2,5] 与 [2,1,5] - 中间维一方为1 EXPECT_FALSE(op::CheckBroadcastShape(shape1, shape6)); // [2,2] 与 [2,2,5] - 维数不同且不满足 }这 7 组断言覆盖了广播判定的全部关键场景维度数不同shape1vsshape2、shape2vsshape6验证右对齐补 1 后的比较逻辑顺序无关对称性shape1vsshape2与shape2vsshape1验证self、other交换后结果不变一方维度为 1shape6vsshape7验证单维拉伸两方均大于 1 且不等shape1vsshape5验证失败分支的短路返回维度数不同且不满足shape1vsshape6验证最坏情况下返回false。与 BroadcastInferShape 的配合从“能否广播”到“广播成什么”CheckBroadcastShape只回答“能否广播”的问题当校验通过后若还需推导广播后的实际 shape应使用同一头文件中声明的 op::BroadcastInferShape参考 BroadcastInferShape 接口文档bool BroadcastInferShape(const op::Shape self, const op::Shape other, op::Shape broadcastShape);两者共享同一套BroadcastDim单维判定逻辑src/nnopbase/common/utils/shape_utils.cpp区别在于CheckBroadcastShape仅返回布尔结果适用于形如“校验输入是否合法”的防御性检查BroadcastInferShape除返回布尔结果外还会把广播后的 shape 写入输出参数broadcastShape并在失败时通过OP_LOGE_FOR_INVALID_ARGUMENT_TENSOR_INPUT_SHAPE上报带详细原因的非法输入日志包含两侧 shape 字符串与冲突维度值适用于算子的输出 shape 推导流程。从源码结构看src/nnopbase/common/utils/shape_utils.cppBroadcastInferShape在广播失败时会给出形如The tensor whose shape is [...] and the tensor whose shape is [...] do not meet the broadcast condition的错误信息开发者可将两者组合使用先用CheckBroadcastShape做轻量预检再调用BroadcastInferShape获取目标 shape或直接依赖BroadcastInferShape一步完成“校验 推导”。在算子开发中的典型应用场景作为 opdev shape 工具族shape_utils 索引的一员CheckBroadcastShape的典型使用场景包括Elementwise 类算子的输入校验在 infershape 或算子入口处对两个输入张量的 shape 做广播关系预检不满足时提前返回失败避免后续计算阶段出现维度错位权重/偏置广播场景如矩阵运算中偏置项 shape如[1]、[C]与主输入 shape 的兼容性判断动态 shape 场景op::Shape由gert::Shape承载天然支持动态维度场景下的维度序列描述可直接对运行时推导出的 shape 调用该校验函数。需要留意的是该接口仅基于维度序列做纯数学判定不感知gert::Shape中可能携带的动态/静态标记语义也不涉及具体张量数据的连续性或格式Format判断如需验证连续 strides 等布局信息应配合ToContiguousStrides等工具使用。【免费下载链接】opbase本项目是CANN算子库的基础框架库为算子提供公共依赖文件和基础调度能力。项目地址: https://gitcode.com/cann/opbase创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考