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是FVector<int64_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。需要注意的是,该示例中的注释写作[2,1]与[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交换后结果不变; - 一方维度为 1(
shape6vsshape7):验证单维拉伸; - 两方均大于 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.cpp),BroadcastInferShape在广播失败时会给出形如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),仅供参考