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、COMPLEX128 | ND |
| y | 输出 | 公式中的 y。 | INT64、INT32、FLOAT、FLOAT16、DOUBLE、BF16、COMPLEX64、COMPLEX128 | ND |
几点需要特别留意:
- 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 与动态 Shape;DynamicFormatFlag(false):不支持动态 Format;ExtendCfgInfo("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、COMPLEX128);diag_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 结构体中; - 对空 Tensor(nSize=0)单独走
TilingEmptyTensor分支,申请 16MB 系统 workspace 并设置单核执行(diag_tiling_arch35.cpp)。
Tiling 阶段还会通过平台接口获取 AIV 核数与 UB 内存大小(TilingGetCompileInfo),并据此计算每个核可处理的最大循环元素数。
4. Kernel:SIMT 并行写对角元素
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_cast<uint64_t>(xIdx) * (nSize + 1); y[yIdx] = x[xIdx]; }其巧妙之处在于:输出对角线元素在展平后的y内存中的下标恰好是xIdx * (nSize + 1),因此无需逐元素判断即可直接定位写入位置。而输出矩阵中的非对角线元素则通过ResetUnifiedBufferZero先以Duplicate指令将整段 Unified Buffer 清零,再经DataCopyPad搬运到 Global Memory(diag_simt.h),随后通过 MTE3 事件同步(SetFlag/WaitFlag)保证"先清零、后写对角"的顺序,最终由asc_vf_call拉起向量线程完成对角元素写入。
配置与测试佐证
二进制 Kernel 配置
ascend950/diag_binary.json 为每个支持的数据类型(bfloat16、int64、float16、int32、double、complex64、float32)分别生成了对应的二进制 Kernel(bin_filename),输入输出 Shape 均以-2表示动态维度;ascend950/diag_simplified_key.ini 配置[Diag] default=0,用于控制 opc 工具编译二进制 Kernel 时的--simplified_key_mode取值。ascend350平台下存在对应配置目录,二者共同支撑上表所列产品的算子部署。
单算子测试用例
仓库自带的 ST(System Test)用例位于 ttk_kernel_diag_st.csv,覆盖了不同规模的 int32 输入:
| 用例名 | 输入 Shape | 输出 Shape | 精度容差 |
|---|---|---|---|
| diag_performance_003 | (4,) | (4, 4) | 0.0001 / 0.001 |
| diag_performance_004 | (8,) | (8, 8) | 0.0001 / 0.001 |
| diag_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),仅供参考