CANN ops-math MatrixDiagV3 算子详解:对角线张量构建的原理、参数与图模式调用实战
2026/9/18 17:35:57 网站建设 项目流程

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_rowsnum_colsk = [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、BOOLND
k输入公式中的k_lk_u。可以是标量(单条对角线),也可以是长度为 2 的向量(对角线带的下界和上界);元素个数只能为 1 或 2,且k[0] <= k[1]INT32ND
num_rows输入输出矩阵的行数,即公式中的num_rows。取值为-1时表示由kx自动推导。INT32ND
num_cols输入输出矩阵的列数,即公式中的num_cols。取值为-1时表示由kx自动推导。INT32ND
padding_value输入公式中的padding_value,用于填充不在指定对角线带内的位置,数据类型与x一致,且必须为单个元素(标量)。x相同ND
align可选属性指定超对角线和次对角线的对齐方式。支持RIGHT_LEFTLEFT_RIGHTLEFT_LEFTRIGHT_RIGHT,默认值为RIGHT_LEFTSTRING-
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 的约束归纳如下:

  1. x的秩至少为 1(单对角线);当k表示对角线带时,x的秩至少为 2(见 matrix_diag_v3_infershape.cpp 中kDiagBandMinRank = 2的检查)。
  2. k的元素个数只能为 1 或 2;当k为 2 个元素时,必须满足k[0] <= k[1]
  3. k表示对角线带时,x的倒数第二维长度必须等于k_u - k_l + 1,最后一维长度必须等于max_diag_len;否则内核报 "k parameter implies [N] diagonals, but diagonal data contains [M] diagonals"。
  4. padding_value必须为标量(单元素)。
  5. y的数据类型必须与x一致,否则内核报参数非法。
  6. num_rowsnum_cols若显式给出,不能小于由xk推导出的最小行数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_rowsnum_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,其中IsRankInvalidIsRankAboveLimitIsShapeFullyDefined等复现了源码侧 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)
  • alignRIGHT_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 个输入(xknum_rowsnum_colspadding_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.cppAICPU 内核实现:参数校验、GetDiagIndex解析kComputeDiagLenAndContentOffset处理 align 对齐、SetResult逐元素写结果
op_kernel_aicpu/matrix_diag_v3_aicpu_def.cppAICPU 算子定义注册(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.cppInferShape 单测:覆盖显式尺寸、自动推导、对角线带、动态 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_valueDoCompute按 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 = -1num_rows = 3num_cols = 2padding = 9的输出矩阵。
  • 批量场景BATCH_SUCCESS用例输入xshape 为{2, 3},输出{2, 3, 3},验证 batch 维度逐批写入主对角线。
  • 异常路径ALIGN_INVALID(非法 align)、K_RANGE_INVALIDk[0] > k[1])、NUM_ROWS_INVALID/NUM_COLS_INVALID(尺寸过小)、NUM_DIAGS_INVALID(对角线条数不匹配)、PADDING_VALUE_INVALID(padding 非标量)、OUTPUT_DTYPE_MISMATCH(输出类型与 x 不一致)等用例均断言返回KERNEL_STATUS_PARAM_INVALID

七、常见问题与使用建议

  1. 输出秩的变化:单对角线(k为标量或k[0] == k[1])时yx多一维;对角线带时秩不变。构图前若对输出 shape 有硬编码,需按此规则调整。
  2. align 与对角线带:当对角线带中各条对角线长度不一致时,align决定短对角线在max_diag_len内的补齐方向。理解LEFT/RIGHT分别作用于超/次对角线的组合方式(见 2.3 节表),可避免数据错位。
  3. num_rows/num_cols 显式值过小:内核与 InferShape 都会校验其不小于xk推导出的最小值,报错信息形如 "The number of rows is too small"。
  4. 动态形状:若k为运行期变量(非常量),InferShape 将输出降级为 1 维未知形状,需确保下游节点支持动态 shape。
  5. 数据类型一致性padding_valuey必须与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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询