CANN ops-math MatrixDiagV3 算子详解:对角线张量构建的原理、参数与图模式调用实战
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_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_LEFTalign超对角线superdiagonal次对角线subdiagonalLEFT_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、BOOLNDk输入公式中的k_l和k_u。可以是标量单条对角线也可以是长度为 2 的向量对角线带的下界和上界元素个数只能为 1 或 2且k[0] k[1]。INT32NDnum_rows输入输出矩阵的行数即公式中的num_rows。取值为-1时表示由k和x自动推导。INT32NDnum_cols输入输出矩阵的列数即公式中的num_cols。取值为-1时表示由k和x自动推导。INT32NDpadding_value输入公式中的padding_value用于填充不在指定对角线带内的位置数据类型与x一致且必须为单个元素标量。与x相同NDalign可选属性指定超对角线和次对角线的对齐方式。支持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 示例参数配置示例中的算子配置如下xshape 为{2}数据类型DT_COMPLEX128数据为{(1.0, 2.0), (3.0, 4.0)}kshape 为{2}向量值为{-1, -1}即只选择d -1这条次对角线num_rows标量3num_cols标量2padding_value标量(9.0, -1.0)alignRIGHT_LEFT输出yshape 为{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 初始化与图执行mapAscendString, 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_LEFTop_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解析k、ComputeDiagLenAndContentOffset处理 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_castint(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_INVALIDk[0] k[1]、NUM_ROWS_INVALID/NUM_COLS_INVALID尺寸过小、NUM_DIAGS_INVALID对角线条数不匹配、PADDING_VALUE_INVALIDpadding 非标量、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.hMatrixDiag 算子族共享的形状推导辅助函数。【免费下载链接】ops-math本项目是CANN提供的数学类基础计算算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-math创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考