ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

CANN ops-math Diag 算子深度解析:从 1D 输入到 2D 对角矩阵的 NPU 实现

CANN ops-math Diag 算子深度解析:从 1D 输入到 2D 对角矩阵的 NPU 实现 CANN ops-math Diag 算子深度解析从 1D 输入到 2D 对角矩阵的 NPU 实现【免费下载链接】ops-math本项目是CANN提供的数学类基础计算算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-math导读Diag 是 CANN ops-math 数学算子库中一个典型的张量转换conversion算子它将输入 Tensor 展平视为 1D 向量后取对角线元素构造成一个 2D 对角矩阵。本文以 conversion/diag/README.md 为核心完整梳理 Diag 算子的功能定义、产品支持范围、参数与约束并结合仓库中的算子注册、InferShape、Tiling 与 SIMT 内核源码深入剖析其在 NPU 上从构图到 Kernel 执行的完整实现链路帮助读者掌握该算子在昇腾场景下的使用方式与底层运行原理。功能说明与数学定义Diag 算子的核心功能为将输入 Tensor展平视为 1D的对角线元素展开为 2D 对角矩阵。它本质上完成了向量 → 对角矩阵的张量变换是数学与深度学习场景中常用的基础算子例如在矩阵分解、梯度计算、张量重建等流程中用于构造对角结构。设输入向量 $\mathbf{x} \in \mathbb{R}^n$输出对角矩阵 $\mathbf{y} \in \mathbb{R}^{n \times n}$其计算规则为$$ \mathbf{y}_{i,i} x_i $$$$ \mathbf{y}_{i,j} 0 \quad (i \ne j) $$即输出矩阵的非对角线元素全部置零对角线元素按序取自输入向量。从源码实现看这一规则在 Kernel 层被精确落实内核中每个线程将输入元素写入输出矩阵的等间隔位置见下文Kernel 实现并通过先写零再写对角元素的方式完成整矩阵构造。产品支持情况Diag 算子在当前仓库中的支持范围覆盖昇腾主流训练与推理产品线具体如下表所示产品是否支持Ascend 950PR/Ascend 950DT√Atlas A3 训练系列产品/Atlas A3 推理系列产品√Atlas A2 训练系列产品/Atlas A2 推理系列产品√Atlas 200I/500 A2 推理产品√Atlas 推理系列产品√Atlas 训练系列产品√从算子配置看diag_def.cpp 中为ascend950与ascend350两个平台分别注册了 AICore 配置与上表中的支持范围保持一致对应的二进制 Kernel 配置diag_binary.json与编译选项diag_simplified_key.ini也分别存放在 config/ascend950 与 config/ascend350 目录下。参数说明Diag 算子仅包含一个输入与一个输出无额外的属性参数参数名输入/输出/属性描述数据类型数据格式x输入公式中的 x。INT64、INT32、FLOAT、FLOAT16、DOUBLE、BF16、COMPLEX64、COMPLEX128NDy输出公式中的 y。INT64、INT32、FLOAT、FLOAT16、DOUBLE、BF16、COMPLEX64、COMPLEX128ND几点需要特别留意BF16 平台限制在 Atlas 训练系列产品、Atlas 推理系列产品、Atlas 200I/500 A2 推理产品、Atlas A2 训练/推理系列产品、Atlas A3 训练/推理系列产品上不支持 BF16。BF16 仅支持于 Ascend 950PR/Ascend 950DT这一限制在算子 IR 注册中同样有体现diag_proto.h 明确注明 bfloat16 is only supported on Ascend950PR/Ascend950DT。数据类型一致性输出 y 与输入 x 保持同类型这在算子定义中由输入/输出的TensorType列表一一对应保证见 diag_proto.h。格式约束输入输出统一使用ND格式且从diag_binary.json中的format_match_mode: FormatAgnostic可知该算子对 Format 匹配采用格式无关策略即设备侧格式以实际下发为准。另外算子定义diag_def.cpp中还声明了以下能力开关实际生效于编译与图调度阶段DynamicCompileStaticFlag(true)支持动态编译DynamicRankSupportFlag(true)与DynamicShapeSupportFlag(true)支持动态 Rank 与动态 ShapeDynamicFormatFlag(false)不支持动态 FormatExtendCfgInfo(opFile.value, diag_apt)指定 Kernel 实现文件为diag_apt。约束说明Diag 只支持输入 Shape 维度为1-4 维对应输出为2-8 维不支持标量输入0 维输入。该约束在 Tiling 阶段会被显式校验diag_tiling_arch35.cpp中的TilingCheckInputParams通过MIN_INPUT_DIM 1、MAX_INPUT_DIM 4检查输入维度范围超出范围直接返回失败并记录错误日志见 diag_tiling_arch35.cpp。调用说明调用方式调用样例说明图模式调用NA通过 算子IR 构图方式调用 Diag 算子当前仓库中 Diag 算子的调用形态为图模式构图调用在构建计算图时通过算子 IR 注册的REG_OP(Diag)声明diag_proto.h创建 Diag 节点。该 IR 定义与 TensorFlow 的Diag算子保持兼容源码注释明确 Compatible with the TensorFlow operator Diag同时 diag_tf_plugin.cpp 中以REGISTER_CUSTOM_OP(Diag)的方式将自定义算子注册到 TensorFlow 框架并通过AutoMappingByOpFn自动映射参数ImplyType::TVM声明其实现类型。这意味着上游框架产生的 Diag 算子可以无缝下沉到该实现执行。源码级实现剖析1. 算子注册与 IR 声明算子的对外接口统一由 diag_proto.h 声明输入x与输出y均支持 8 种数据类型FLOAT16、FLOAT、DOUBLE、BF16、INT32、INT64、COMPLEX64、COMPLEX128diag_def.cpp 则通过OpDef机制在算子注册层补充了参数类型、Format、动态能力开关与平台 AICore 配置。2. InferShape输出维度翻倍diag_infershape.cpp 实现了形状推导对于输入维度数为x_dim_num的 x输出 y 的维度数被设置为输入的两倍且每一维大小与输入对应维完全一致即把输入 shape 整体重复两遍后拼接。这与输入 1-4 维 → 输出 2-8 维的约束完全对应例如输入(4,)推导出输出(4, 4)。3. Tiling多核切分策略Diag 属于批量为 1、单算子多核的 SIMT 型 Kernel。diag_tiling_arch35.cpp 中的CalcSimtTiling负责生成 Tiling 参数以HALF_VL_LEN (128) / dtypeSize作为边长度量因子计算所需的 block 数实际使用核数取min(平台AIV核数, blockNum)将 n输入元素个数均分到各核产生mainBlockCount主核数、mainBlockFactor主核元素数、tailBlockFactor尾核元素数等字段记录在 DiagSimtTilingData 结构体中对空 TensornSize0单独走TilingEmptyTensor分支申请 16MB 系统 workspace 并设置单核执行diag_tiling_arch35.cpp。Tiling 阶段还会通过平台接口获取 AIV 核数与 UB 内存大小TilingGetCompileInfo并据此计算每个核可处理的最大循环元素数。4. KernelSIMT 并行写对角元素Kernel 入口位于 diag_apt.cpp通过TILING_KEY_SIMT分发到 SIMT 路径真正的计算核心在 diag_simt.h 的DiagSIMTCompute中for (uint32_t idx threadIdx.x; idx curCoreElements; idx blockDim.x) { uint32_t xIdx xBaseIdx idx; uint64_t yIdx static_castuint64_t(xIdx) * (nSize 1); y[yIdx] x[xIdx]; }其巧妙之处在于输出对角线元素在展平后的y内存中的下标恰好是xIdx * (nSize 1)因此无需逐元素判断即可直接定位写入位置。而输出矩阵中的非对角线元素则通过ResetUnifiedBufferZero先以Duplicate指令将整段 Unified Buffer 清零再经DataCopyPad搬运到 Global Memorydiag_simt.h随后通过 MTE3 事件同步SetFlag/WaitFlag保证先清零、后写对角的顺序最终由asc_vf_call拉起向量线程完成对角元素写入。配置与测试佐证二进制 Kernel 配置ascend950/diag_binary.json 为每个支持的数据类型bfloat16、int64、float16、int32、double、complex64、float32分别生成了对应的二进制 Kernelbin_filename输入输出 Shape 均以-2表示动态维度ascend950/diag_simplified_key.ini 配置[Diag] default0用于控制 opc 工具编译二进制 Kernel 时的--simplified_key_mode取值。ascend350平台下存在对应配置目录二者共同支撑上表所列产品的算子部署。单算子测试用例仓库自带的 STSystem Test用例位于 ttk_kernel_diag_st.csv覆盖了不同规模的 int32 输入用例名输入 Shape输出 Shape精度容差diag_performance_003(4,)(4, 4)0.0001 / 0.001diag_performance_004(8,)(8, 8)0.0001 / 0.001diag_performance_007(64,)(64, 64)0.0001 / 0.001测试用例的输入输出均声明为 ND 格式输入数据范围设置为(1, None)验证了不同 n 值下向量展平 → 对角矩阵的结果正确性同时 golden.py 提供了基准结果的生成逻辑可供离线比对。总结Diag 算子是 CANN ops-math 中张量结构转换类算子的代表数学语义直观向量 → 对角矩阵但 NPU 实现涉及算子注册、形状推导、多核 Tiling、SIMT 内核与事件同步等多个环节。理解其实现对在昇腾平台上进行张量级结构变换类算子的开发与调优具有直接的参考价值。若需进一步掌握算子构图方式可结合 diag_proto.h 与 README.md 继续深入。【免费下载链接】ops-math本项目是CANN提供的数学类基础计算算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-math创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进