CANN ops-math MatrixDiagV3 算子详解:对角线张量构建的原理、参数与图模式调用实战
【免费下载链接】ops-math本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-math
本文围绕 CANN ops-math 仓库中 conversion/matrix_diag_v3/README.md 的核心内容展开,系统讲解 MatrixDiagV3 算子的数学语义、输入输出参数、对齐方式(align)语义与约束条件,并结合仓库中的算子原型、InferShape、AICPU 内核实现与单测代码,说明如何在昇腾 NPU 上通过图模式(GEIR)完成该算子的构图与调用。读完本文,你将能够理解 MatrixDiagV3 的完整行为模型,并能参照示例工程独立编写基于算子 IR 的调用程序。
一、算子概述与应用场景
MatrixDiagV3 是一个"由对角线值构造矩阵"的算子:它根据输入x(单条或多条对角线上的元素值),在输出矩阵的指定对角线带上写入这些值,对角线带之外的位置统一用padding_value填充。
该算子与 TensorFlow 的tf.linalg.diag系列算子(MatrixDiag / MatrixDiagV2 / MatrixDiagV3)语义对齐,是矩阵构造、带状矩阵生成、注意力掩码构建等场景的基础算子。在 CANN ops-math 项目中,它被归入 conversion(数据格式与结构转换)类算子目录,产品形态同时提供图模式(GE 构图)与 AICPU 内核两种落地路径。
产品支持情况
| 产品 | 是否支持 |
|---|---|
| Ascend 950PR/Ascend 950DT | √ |
| Atlas A3 训练系列产品/Atlas A3 推理系列产品 | √ |
| Atlas A2 训练系列产品/Atlas A2 推理系列产品 | √ |
| Atlas 200I/500 A2 推理产品 | × |
| Atlas 推理系列产品 | √ |
| Atlas 训练系列产品 | √ |
从支持矩阵可见,MatrixDiagV3 覆盖当前主流训练与推理产品,仅 Atlas 200I/500 A2 推理产品暂不支持。
二、功能说明与数学语义
2.1 输出元素与对角线的对应关系
设输出张量最后两维大小分别为num_rows和num_cols,k = [k_l, k_u]表示待写入的对角线范围;单对角线场景下k_l = k_u。令d = j - i表示位置(i, j)所在的对角线编号,则最大对角线长度为:
max_diag_len = min(num_rows + min(k_u, 0), num_cols - max(k_l, 0))输出元素满足:
y_{..., i, j} = x_{..., k_u - d, p(i, j)} 若 k_l ≤ d ≤ k_u padding_value 其他情况其中p(i, j)表示对角线元素在输入x最后一维中的位置,具体由属性align控制左右对齐方式。
- 当
k为单个整数或k_l == k_u时,x表示单条对角线; - 当
k_l < k_u时,x表示对角线带,x的倒数第二维保存对角线条数,最后一维保存各对角线按align补齐后的数据。
2.2 对角线编号约定
在 matrix_diag_v3_proto.h 的算子注释中对k的语义有明确说明:
k为正数表示超对角线(superdiagonal,位于主对角线右上方);k = 0表示主对角线;k为负数表示次对角线(subdiagonal,位于主对角线左下方);k[0]不得大于k[1]。
2.3 align 对齐语义
align决定超对角线与次对角线在max_diag_len长度内的对齐方向,支持四种取值,默认RIGHT_LEFT:
| align | 超对角线(superdiagonal) | 次对角线(subdiagonal) |
|---|---|---|
LEFT_LEFT | 左对齐 | 左对齐 |
LEFT_RIGHT | 左对齐 | 右对齐 |
RIGHT_LEFT(默认) | 右对齐 | 左对齐 |
RIGHT_RIGHT | 右对齐 | 右对齐 |
在 AICPU 内核 matrix_diag_v3_aicpu.cpp 的CheckParam中,align 被解析为两个布尔标志位:
left_align_superdiagonal_ = align == "LEFT_LEFT" || align == "LEFT_RIGHT"; left_align_subdiagonal_ = align == "LEFT_LEFT" || align == "RIGHT_LEFT";随后在ComputeDiagLenAndContentOffset中,content_offset的计算为:左对齐时偏移为 0,右对齐时偏移为max_diag_len - diag_len,即把较短的对角线向行尾(右端)补齐。
三、参数说明
下表完整继承 README 的参数定义,并补充了来自 matrix_diag_v3_proto.h 与 matrix_diag_v3_aicpu_def.cpp 的实现细节。
| 参数名 | 输入/输出/属性 | 描述 | 数据类型 | 数据格式 |
|---|---|---|---|---|
| x | 输入 | 公式中的x。当k表示单条对角线时,x的最后一维保存该对角线的数据;当k表示对角线带时,x的倒数第二维保存对角线条数,最后一维保存各对角线按align补齐后的数据。秩至少为 1。 | DOUBLE、FLOAT、FLOAT16、INT8、INT16、INT32、INT64、UINT8、UINT16、UINT32、UINT64、COMPLEX64、COMPLEX128、BOOL | ND |
| k | 输入 | 公式中的k_l和k_u。可以是标量(单条对角线),也可以是长度为 2 的向量(对角线带的下界和上界);元素个数只能为 1 或 2,且k[0] <= k[1]。 | INT32 | ND |
| num_rows | 输入 | 输出矩阵的行数,即公式中的num_rows。取值为-1时表示由k和x自动推导。 | INT32 | ND |
| num_cols | 输入 | 输出矩阵的列数,即公式中的num_cols。取值为-1时表示由k和x自动推导。 | INT32 | ND |
| padding_value | 输入 | 公式中的padding_value,用于填充不在指定对角线带内的位置,数据类型与x一致,且必须为单个元素(标量)。 | 与x相同 | ND |
| align | 可选属性 | 指定超对角线和次对角线的对齐方式。支持RIGHT_LEFT、LEFT_RIGHT、LEFT_LEFT、RIGHT_RIGHT,默认值为RIGHT_LEFT。 | STRING | - |
| y | 输出 | 公式中的y,生成后的矩阵张量,数据类型与x一致。当k为单个元素或k[0] == k[1]时,y的秩为x的秩加 1;否则y的秩与x一致。 | 与x相同 | ND |
需要特别注意的是,padding_value在 AICPU 内核中通过padding_value_num == 1强校验(见CheckParam),传入多个元素会直接报KERNEL_STATUS_PARAM_INVALID;单测用例PADDING_VALUE_INVALID对该行为做了覆盖验证。
四、约束说明
结合 README 与 InferShape/内核实现,MatrixDiagV3 的约束归纳如下:
x的秩至少为 1(单对角线);当k表示对角线带时,x的秩至少为 2(见 matrix_diag_v3_infershape.cpp 中kDiagBandMinRank = 2的检查)。k的元素个数只能为 1 或 2;当k为 2 个元素时,必须满足k[0] <= k[1]。- 当
k表示对角线带时,x的倒数第二维长度必须等于k_u - k_l + 1,最后一维长度必须等于max_diag_len;否则内核报 "k parameter implies [N] diagonals, but diagonal data contains [M] diagonals"。 padding_value必须为标量(单元素)。y的数据类型必须与x一致,否则内核报参数非法。num_rows、num_cols若显式给出,不能小于由x与k推导出的最小行数min_num_rows = max_diag_len - min(k_u, 0)与最小列数min_num_cols = max_diag_len + max(k_l, 0)。
4.1 自动推导规则(num_rows / num_cols = -1)
当num_rows与num_cols均为-1时,输出退化为方阵,取max(min_num_rows, min_num_cols);当只有一个为-1时,按上面对应的最小值补齐(见内核AdjustRowsAndCols)。InferShape 侧的推导逻辑与内核一致,保证编译期形状推导与运行期实际计算语义对齐。
4.2 动态场景的降级处理
从 matrix_diag_v3_infershape.cpp 可以看到,当k不是编译期常量、或x秩未知、或k的形状未完全确定时,InferShape 无法推导具体输出形状,会将输出降级为 1 维未知形状(SetUnknownShape),交由后续动态 shape 流程处理。该设计保证了合法的动态图不会被拒绝(相关秩检查辅助函数见 matrix_diag_infershape_common.h,其中IsRankInvalid、IsRankAboveLimit、IsShapeFullyDefined等复现了源码侧 WithRank / WithRankAtMost / FullyDefined 的容错语义)。
五、调用说明:图模式(GEIR)构图调用
MatrixDiagV3 支持图模式调用,即通过算子 IR 构图后交给 GE 引擎执行。仓库在 examples/test_geir_matrix_diag_v3.cpp 中给出了完整可运行的示例,其完整调用链为:
创建 MatrixDiagV3 算子节点 → 构造 x / k / num_rows / num_cols / padding_value 五个输入(Data 或 Const 节点) → 设置 align 属性 → 声明输出 desc → GEInitialize 初始化 GE → 构建 Graph 并 SetInputs/SetOutputs → 创建 Session 并 AddGraph → RunGraph 执行并校验输出5.1 示例参数配置
示例中的算子配置如下:
x:shape 为{2},数据类型DT_COMPLEX128,数据为{(1.0, 2.0), (3.0, 4.0)};k:shape 为{2}(向量),值为{-1, -1},即只选择d = -1这条次对角线;num_rows:标量3;num_cols:标量2;padding_value:标量(9.0, -1.0);align:RIGHT_LEFT;- 输出
y:shape 为{3, 2},数据类型DT_COMPLEX128。
由于k[0] == k[1](单条对角线),输出y的秩等于x的秩加 1,即 1 维输入变为 2 维矩阵。程序期望输出为:
{(9.0, -1.0), (9.0, -1.0)}, {(1.0, 2.0), (9.0, -1.0)}, {(9.0, -1.0), (3.0, 4.0)}即次对角线(1,0)与(2,1)位置写入x的元素,其余位置写入padding_value。
5.2 关键代码片段
算子节点创建与属性设置:
auto matrix_diag_v3 = op::MatrixDiagV3("matrix_diag_v3"); matrix_diag_v3.set_attr_align("RIGHT_LEFT");输入输出通过set_input_*与update_input_desc_*/update_output_desc_y进行绑定:
matrix_diag_v3.set_input_x(data); matrix_diag_v3.update_input_desc_x(desc); // ... 依次绑定 k、num_rows、num_cols、padding_value TensorDesc output_desc(ge::Shape({3, 2}), FORMAT_ND, DT_COMPLEX128); matrix_diag_v3.update_output_desc_y(output_desc);GE 初始化与图执行:
map<AscendString, AscendString> global_options = {{"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "1"}}; ge::GEInitialize(global_options); // ... 构图、SetInputs/SetOutputs Session *session = new Session(build_options); session->AddGraph(graph_id, graph, graph_options); session->RunGraph(graph_id, input, output);5.3 输出校验
示例在运行后会对输出与期望值逐元素比对(CompareTensor),比对失败即打印 "Output validation failed" 并以非 0 码退出,成功则打印 "MatrixDiagV3 example passed"。这种"构图-执行-比对"三步式的写法可以直接复用到其他输入配置的验证中。
六、算子实现结构(源码导读)
MatrixDiagV3 在仓库中的实现横跨算子定义、形状推导、图推断与内核计算四层,目录为 conversion/matrix_diag_v3:
| 文件 | 职责 |
|---|---|
| op_graph/matrix_diag_v3_proto.h | 使用REG_OP注册算子原型:5 个输入(x、k、num_rows、num_cols、padding_value)、1 个输出y、1 个字符串属性align(默认RIGHT_LEFT) |
| op_graph/matrix_diag_v3_graph_infer.cpp | 图推断:通过InferDataTypeOutputSameAsInput声明输出数据类型与输入x一致 |
| op_host/matrix_diag_v3_infershape.cpp | 形状推导:根据k的常量值与x的最后一维推导输出[num_rows, num_cols],单对角线秩加 1,对角线带秩不变 |
| op_kernel_aicpu/matrix_diag_v3_aicpu.cpp | AICPU 内核实现:参数校验、GetDiagIndex解析k、ComputeDiagLenAndContentOffset处理 align 对齐、SetResult逐元素写结果 |
| op_kernel_aicpu/matrix_diag_v3_aicpu_def.cpp | AICPU 算子定义注册(OP_ADD(MatrixDiagV3)),明确各输入输出允许的数据类型清单 |
| examples/test_geir_matrix_diag_v3.cpp | 图模式调用示例(见第五节) |
| tests/ut/op_kernel_aicpu/test_matrix_diag_v3.cpp | 内核单测:覆盖全部 14 种数据类型 + 批量场景 + 非法参数用例 |
| tests/ut/op_host/test_matrix_diag_v3_infershape.cpp | InferShape 单测:覆盖显式尺寸、自动推导、对角线带、动态 k 等场景 |
6.1 内核计算核心逻辑
SetResult中针对输出矩阵逐元素计算其所在对角线编号与源数据下标,其核心索引关系为:
const int diag_index = static_cast<int>(j - i); // 当前元素所在对角线编号 d const int diag_index_in_input = upper_diag_index_ - diag_index; // 在输入 x 的第几行 const int index_in_the_diagonal = (j - max(diag_index, 0)) + content_offset; // 对角线内偏移当lower_diag_index_ <= diag_index <= upper_diag_index_时从x取数,否则写入padding_value。DoCompute按 batch 循环(num_batches = num_elements / (num_rows * num_cols)),支持带 batch 维的输入。
6.2 测试覆盖情况
- 数据类型全覆盖:
MATRIX_DIAG_V3_BASIC_CASE宏为 INT32/INT64/FLOAT/DOUBLE/FLOAT16/INT8/UINT8/UINT16/UINT32/UINT64 生成基础用例,另有 COMPLEX64、COMPLEX128、BOOL 专用用例,均验证k = -1、num_rows = 3、num_cols = 2、padding = 9的输出矩阵。 - 批量场景:
BATCH_SUCCESS用例输入xshape 为{2, 3},输出{2, 3, 3},验证 batch 维度逐批写入主对角线。 - 异常路径:
ALIGN_INVALID(非法 align)、K_RANGE_INVALID(k[0] > k[1])、NUM_ROWS_INVALID/NUM_COLS_INVALID(尺寸过小)、NUM_DIAGS_INVALID(对角线条数不匹配)、PADDING_VALUE_INVALID(padding 非标量)、OUTPUT_DTYPE_MISMATCH(输出类型与 x 不一致)等用例均断言返回KERNEL_STATUS_PARAM_INVALID。
七、常见问题与使用建议
- 输出秩的变化:单对角线(
k为标量或k[0] == k[1])时y比x多一维;对角线带时秩不变。构图前若对输出 shape 有硬编码,需按此规则调整。 - align 与对角线带:当对角线带中各条对角线长度不一致时,
align决定短对角线在max_diag_len内的补齐方向。理解LEFT/RIGHT分别作用于超/次对角线的组合方式(见 2.3 节表),可避免数据错位。 - num_rows/num_cols 显式值过小:内核与 InferShape 都会校验其不小于
x与k推导出的最小值,报错信息形如 "The number of rows is too small"。 - 动态形状:若
k为运行期变量(非常量),InferShape 将输出降级为 1 维未知形状,需确保下游节点支持动态 shape。 - 数据类型一致性:
padding_value与y必须与x保持同一数据类型,BOOL 与 COMPLEX 系列同样适用。
如需深入阅读算子全量实现与测试,可继续查看 conversion/matrix_diag_v3 目录下的源码文件,以及公共的 matrix_diag_infershape_common.h(MatrixDiag 算子族共享的形状推导辅助函数)。
【免费下载链接】ops-math本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-math
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考