ARTICLE DETAIL

资讯详情

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

ops-math 中的 broadcast(广播)关系:NPU 算子 Shape 兼容三大规则与特殊类型限制详解

ops-math 中的 broadcast(广播)关系:NPU 算子 Shape 兼容三大规则与特殊类型限制详解 ops-math 中的 broadcast广播关系NPU 算子 Shape 兼容三大规则与特殊类型限制详解【免费下载链接】ops-math本项目是CANN提供的数学类基础计算算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-math在 CANN 数学算子库 ops-math 中绝大多数二元算子如加法、比较、广播填充类的输入参数都支持广播broadcast两个 shape 不同的张量可以按规则自动对齐后参与元素级运算。本文以仓库文档 broadcast关系 为核心完整讲解广播概念、三大广播规则与特殊数据类型的广播限制条件并结合 math/add 等算子的源码实现说明这些规则在算子 Host 侧形状推导中的落地方式帮助你在调用算子 API 前正确设计输入 shape、避免广播报错。一、什么是 broadcast 关系broadcast广播描述了算子在运算期间如何处理不同形状的张量或数组。大部分情况下允许不同形状的张量或数组在进行元素操作时自动扩展其形状使其维度相互兼容通常较小的张量或数组会“广播”为较大的张量或数组。在 ops-math 中许多算子 API 参数的 shape 支持广播这样做可以适当提高计算效率无需用户显式构造与另一输入同形状的副本由算子直接按广播语义计算减少内存占用尤其在大模型训练中的大规模数据场景如 bias 加到[B, H, W, C]特征图上、scale/bias 加到[B, C, H, W]张量上避免了物化一份与较大张量同形状的中间数据。广播的技术基础与 NumPy 的广播语义一致ops-math 算子文档中也建议读者参考 NumPy 官方文档的 broadcasting 章节理解细节。本文聚焦算子开发视角下的三条规则与限制。二、三大广播规则一般进行广播计算时需要理解以下三条规则。规则1维度数不一致时向最长形状看齐并在左侧补 1如果数组间维度数不一致所有数组向最长形状的数组看齐形状不足的部分在左侧填充 1直至维度数相同。说明举例1维度数Number of Dimensions是指张量或数组对应 shape 的维数比如x.shape(1, 1, 2, 4)维度数是 4。举例2比如计算ab其中a.shape(2, 2, 3)、b.shape(2, 3)那么数组b将被 broadcast 为b.shape(1, 2, 3)。这里的关键点是“左侧对齐”补维只补在最高维一侧而不是尾部。例如(2, 3)补成(1, 2, 3)而不是(2, 3, 1)这是很多 shape 设计错误的根源。规则2同维度位置为 1 的数组被拉伸匹配另一方如果数组间维度数一致且某个数组的某一维度为 1则该维度为 1 的数组将被拉伸以匹配另一个数组对应维度形状。说明 本场景下只需保证在某一维度做 broadcast 即可。比如计算ab其中a.shape(1, 3)、b.shape(3, 1)那么两个数组会 broadcast 为a.shape(3, 3)、b.shape(3, 3)。规则3维度既不一致又不为 1则报错如果数组间在同一个维度上既不相等、又不为 1即无法通过规则1、规则2对齐则会报错。这是调用算子前需要重点自查的一条逐维检查两个 shape任何一维ne且! 1即非法。一个完整示例先按规则1扩维再按规则2拉伸基于上述规则广播过程一般先按规则1进行扩维再按规则2进行形状拉伸。以ab为例假设a.shape(2,2,3)取值形如: [[[1 2 3],[4 5 6]], [[1 2 3],[4 5 6]]] 假设b.shape(2,3)取值形如: [[1 2 3], [-1 -2 -3]] 根据规则1扩展维度b.shape(1,2,3)取值如下 [[[1 2 3], [-1 -2 -3]]] 根据规则2拉伸形状b.shape(2,2,3)取值如下 [[[1 2 3],[-1 -2 -3]], [[1 2 3],[-1 -2 -3]]] 计算ab实际结果如下: [[[2 4 6],[3 3 3]], [[2 4 6],[3 3 3]]]注意结果中第 2、4 行[3 3 3]它来自规则2把b的第二维值为 -1、-2、-3 的行沿 batch 方向复制叠加到a上验证了“先扩维、再拉伸”的顺序。三、限制特殊数据类型的广播轴合并要求三条规则只描述“形状能否对齐”而 ops-math 还对特定数据类型施加了额外限制。当满足 broadcast 关系的两个输入a和b的数据类型或推导后的数据类型在COMPLEX64、COMPLEX128、DOUBLE、INT16、UINT16、UINT64中时除了满足上述广播规则还需满足如下条件否则广播会失败导致算子执行报错条件连续的需要广播的轴和连续的不需要广播的轴合并之后的维度要求小于 6。也就是说把相邻的广播轴、相邻的非广播轴分别归并为若干“轴段”归并后轴段总数必须小于 6。举例当a.shape(5, 1, 5, 1, 5, 1)b.shape(5, 5, 5, 5, 5, 5)时6 个轴全部是“需要广播的轴”且彼此被非 1 轴分隔没有需要合并的轴段最后轴段维度为 6广播报错。当a.shape(5, 1, 5, 5, 1, 1)b.shape(5, 5, 5, 5, 5, 5)时第 2 维和第 3 维都不需要广播可合并为一段第 4、5 维都需要广播分别连续合并合并后的轴段维度为 4广播成功。从 math/add 的算子文档可以看到aclnnAdd的self/other支持的数据类型恰好覆盖了 FLOAT、DOUBLE、INT16、COMPLEX64、COMPLEX128 等上述受限类型完整列表为 FLOAT、FLOAT16、DOUBLE、INT32、INT64、INT16、INT8、UINT8、BOOL、COMPLEX128、COMPLEX64、BFLOAT16。这意味着对aclnnAdd而言只要输入或推导类型命中受限集合就必须在设计 shape 时额外验证“轴段合并后小于 6”这一条件而 FLOAT/FLOAT16/BFLOAT16/INT32 等类型则不受该轴段数约束。四、源码级佐证ops-math 如何落地广播推导结合仓库源码可以看到广播关系在 ops-math 中的实现位置与调用方式便于你定位报错来源。1. Host 侧形状推导直接复用统一的广播工具以加法算子为例add_infershape.cpp 中注册了 Add 的形状推导函数static ge::graphStatus InferShapeForAdd(gert::InferShapeContext* context) { OP_LOGI(Begin InferShapeForAdd); return Ops::Base::InferShape4Broadcast(context); } IMPL_OP_INFERSHAPE(Add).InferShape(InferShapeForAdd);从源码结构看Add 这类二元算子并不各自实现广播逻辑而是直接调用公共的Ops::Base::InferShape4Broadcast头文件见infershape_broadcast_util.h由统一的广播工具完成“扩维 拉伸 冲突检测”推导出输出 shape。这解释了为什么不同算子对同一对输入 shape 会给出一致的对齐语义也说明规则3的报错是在 Host 侧推导阶段就产生的。2. 算子文档把“满足 broadcast 关系”写成参数的前置约束aclnnAdd 与 aclnnInplaceAdd 的接口文档 对self参数的使用说明明确写着数据类型与other的数据类型需满足数据类型推导规则shape 需要与other满足 broadcast 关系。即广播约束是参数级契约在写 ACLNN 调用代码前就应保证两个输入 shape 可广播而不是依赖运行期报错。3. 两段式 API 下广播推导发生在 GetWorkspaceSize 阶段ops-math 的 ACLNN 接口为两段式参见 两段式接口先调用aclnnAddGetWorkspaceSize获取 workspace 大小与执行器再调用aclnnAdd执行计算接口定义见 aclnn_add.h、实现见 aclnn_add.cpp。从流程上可以推断shape 推导含广播校验在第一段GetWorkspaceSize时即已执行workspace 大小依赖推导后的输出形状因此广播不合法的调用通常在这一段就返回错误状态码而不是真正跑 kernel 时才暴露。若你的调用失败优先检查该段返回的状态与两端 shape。五、实操自查清单结合以上规则与限制编写或调试使用广播的算子调用时可按以下顺序自查逐维对齐把两个 shape 左侧对齐后逐维比较确认每一维满足“相等或至少一方为 1”规则1 规则2任何一维违反即触发规则3报错检查数据类型确认输入类型或推导后的类型是否落在 COMPLEX64、COMPLEX128、DOUBLE、INT16、UINT16、UINT64 受限集合内注意推导类型可能与输入类型不同若命中受限类型统计“连续广播轴段 连续非广播轴段”的合并总数要求小于 6若达到 6典型如a形如(x, 1, x, 1, x, 1)对(x, x, x, x, x, x)的六维逐轴交替需要调整 shape 布局例如合并相邻的 1 维、转置使广播轴相邻使其满足合并条件确认报错阶段若错误出现在xxxGetWorkspaceSize段基本可定位为 shape/类型推导问题回到第 1~3 步排查。参考docs/zh/context/broadcast_relationship.md广播概念、三大规则与限制的原始文档本文主体来源docs/zh/context/basic_concept.md算子基本概念导航broadcast 关系为其子主题之一math/add/docs/aclnnAddaclnnInplaceAdd.mdaclnnAdd/aclnnInplaceAdd 接口文档展示 broadcast 关系作为参数契约的写法math/add/op_host/add_infershape.cppAdd 算子的 Host 侧广播形状推导实现math/add/op_api/aclnn_add.h、math/add/op_api/aclnn_add.cpp两段式 ACLNN 接口定义与实现【免费下载链接】ops-math本项目是CANN提供的数学类基础计算算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-math创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表