CANN PTO TMATMUL_ACC 指令详解:带累加器输入的融合矩阵乘法
发布时间:2026/9/20 3:44:52来源:尧图网络
人工智能指令集算子库CANNAscend【免费下载链接】pto-isaParallel Tile Operation (PTO) is a virtual instruction set architecture designed by Ascend CANN, focusing on tile-level operations. This repository offers high-performance, cross-platform tile operations across Ascend platforms.项目地址https://gitcode.com/cann/pto-isa点击查看免费下载导读本文详细讲解 CANN PTOParallel Tile Operation虚拟指令集中的TMATMUL_ACC指令——一种将外部累加器输入与矩阵乘法融合执行的 Tile 级运算。该指令在算子开发中用于 GEMM 的 K 维分块累加、残差融合、bias 累加等场景是实现高性能矩阵乘法的核心指令之一。读完本文你将掌握 TMATMUL_ACC 的数学语义、SSA/DPS 两级汇编形式、三种 C 内建接口重载、约束条件与实操示例并能结合仓库源码理解其底层实现原理与浮点累加顺序的工程影响。指令定位与 TMATMUL 的关系TMATMUL_ACC是 TMATMUL矩阵乘法 GEMM生成累加器/输出 Tile的融合累加变体。两者的区别在于TMATMUL执行C A · B输出由硬件/实现从零开始累积TMATMUL_ACC执行C1 C0 A · B其中C0是外部传入的初始累加值乘法结果直接叠加在C0上并非先独立计算点积再在末尾加上C0。这一差异在跨 K tile 累加时会体现为浮点舍入顺序的不同融合乘加FMA序列中C0作为初始累加值参与每一步累加与先算完整点积再加C0两种顺序可能产生不同的浮点结果。因此不能仅从数学公式上把两者混为一谈。数学语义设M aMatrix.GetValidRow()K aMatrix.GetValidCol()N bMatrix.GetValidCol()对于0 i M和0 j N$$ \mathrm{C1}{i,j} \mathrm{C0}{i,j} \sum_{k0}^{K-1} \mathrm{A}{i,k} \cdot \mathrm{B}{k,j} $$其中C0是融合乘加序列的初始累加值。值得注意的是m/k/n三个维度全部取自矩阵 Tile 的Valid有效尺寸行数m与 K 维k来自左矩阵aMatrix列数n来自右矩阵bMatrix而非 Tile 的物理分配尺寸Rows/Cols。汇编语法TMATMUL_ACC 在 PTO 汇编中呈现为两级形式与 TMATMUL 保持一致的语法骨架区别在于多出一个累加器输入操作数。同步形式PTO 汇编%acc1 tmatmul.acc %acc0, %a, %b : (!pto.tile..., !pto.tile..., !pto.tile...) - !pto.tile...AS Level 1SSA%c_out pto.tmatmul.acc %c_in, %a, %b : (!pto.tile..., !pto.tile..., !pto.tile...) - !pto.tile...AS Level 2DPSpto.tmatmul.acc ins(%c_in, %a, %b : !pto.tile_buf..., !pto.tile_buf..., !pto.tile_buf...) outs(%c_out : !pto.tile_buf...)其中%c_in为初始累加值输入%c_out为累加结果输出。DPS 形式通过ins(...)/outs(...)显式区分输入与输出 buffer。C 内建接口TMATMUL_ACC 的 C 内建接口声明于 include/pto/common/pto_instr.hpp公共包含头为pto/pto-inst.hpp。接口提供三种重载// 重载 1显式输入/输出无 Phase 模板参数 template typename TileRes, typename TileLeft, typename TileRight, typename... WaitEvents PTO_INST RecordEvent TMATMUL_ACC(TileRes cOutMatrix, TileRes cInMatrix, TileLeft aMatrix, TileRight bMatrix, WaitEvents ... events); // 重载 2显式输入/输出带 AccPhase 模板参数UF 感知 template AccPhase Phase, typename TileRes, typename TileLeft, typename TileRight, typename... WaitEvents PTO_INST RecordEvent TMATMUL_ACC(TileRes cOutMatrix, TileRes cInMatrix, TileLeft aMatrix, TileRight bMatrix, WaitEvents ... events); // 重载 3原地模式cMatrix 同时作为累加器输入与输出 template AccPhase Phase AccPhase::Unspecified, typename TileRes, typename TileLeft, typename TileRight, typename... WaitEvents PTO_INST RecordEvent TMATMUL_ACC(TileRes cMatrix, TileLeft aMatrix, TileRight bMatrix, WaitEvents ... events);三种重载均返回RecordEvent尾部WaitEvents ... events为可选的事件等待参数接口内部通过detail::PtoWaitEvents(events...)完成同步后再发射指令。重载 3 的语义要点该重载将cMatrix同时作为累加器输入与输出原地操作读取其现有值作为C0写回C1 C0 A·B既不清零也不覆盖。这在连续 K tile 循环累加场景中最为常用——同一块累加 Tile 反复作为输入与输出传递避免显式区分两个 buffer。从源码结构看重载 2/3 最终都路由到TMATMUL_ACC_IMPLPhase(...)而重载 1 通过MAP_INSTR_IMPL宏分发到目标实现的TMATMUL_ACC_IMPL且三种重载的模板形参AccPhase Phase默认值为AccPhase::Unspecified详见 include/pto/common/pto_instr.hpp。约束TMATMUL_ACC 的约束分为两部分继承自 TMATMUL 的约束所有来自TMATMUL的约束都适用于(cOutMatrix, aMatrix, bMatrix)三元组。以 Atlas A2/A3 训练系列产品/Atlas A2/A3 推理系列产品为例TMATMUL 文档 明确支持的(CType, AType, BType)三元组(int32_t, int8_t, int8_t)、(float, half, half)、(float, float, float)、(float, bfloat16_t, bfloat16_t)静态形状约束TileLeft::Rows TileRes::Rows、TileLeft::Cols TileRight::Rows、TileRight::Cols TileRes::ColsTile 位置TileLeft::Loc Left、TileRight::Loc Right、TileRes::Loc Acc运行时m/k/n必须在[1, 4095]范围内。实现说明Atlas A2/A3 训练系列产品/Atlas A2/A3 推理系列产品/Ascend 950PR/Ascend 950DTTMATMUL_ACC_IMPL使用aMatrix.GetValidRow()、aMatrix.GetValidCol()和bMatrix.GetValidCol()作为m/k/ncInMatrix在当前实现中不通过显式断言进行验证目标定义的行为——即累加器输入的有效性由编程者保证编译器/运行时不做强制检查。源码中的数据类型与分形校验在 Ascend 950PR/Ascend 950DTA5实现 include/pto/npu/a5/TMatmul.hpp 中CheckMadValid编译期断言了更严格的约束累加类型CType必须是int32_t或float若CType int32_t则AType int8_t且BType int8_t若CType float支持half/half、bfloat16_t/bfloat16_t、float/float、选定的 fp8 组合float8_e4m3_t与float8_e5m2_t的四种配对以及hifloat8_t/hifloat8_t目标定义分形/布局约束Left Tile 需Loc Left、非行主序、SFractal RowMajorRight Tile 需Loc Right、行主序、SFractal ColMajorAcc Tile 需Loc Acc、非行主序、SFractal RowMajor另有MadAccStrideCompatibleTileRes, MadRows断言检查累加 Tile 的 pitch 与mad写回间距一致mad以ceil16(m)行为块列间距写回累加 Tile 的Rows必须与其匹配。示例自动Auto模式自动模式下由编译器/运行时负责 Tile 的资源放置与调度直接调用接口即可#include pto/pto-inst.hpp using namespace pto; void example_auto() { using A TileLefthalf, 16, 16; using B TileRighthalf, 16, 16; using C TileAccfloat, 16, 16; A a; B b; C c0, c1; TMATMUL_ACC(c1, c0, a, b); }手动Manual模式手动模式下需先用TASSIGN显式绑定 Tile 资源地址再发射指令#include pto/pto-inst.hpp using namespace pto; void example_manual() { using A TileLefthalf, 16, 16; using B TileRighthalf, 16, 16; using C TileAccfloat, 16, 16; A a; B b; C c0, c1; TASSIGN(a, 0x1000); TASSIGN(b, 0x2000); TASSIGN(c0, 0x3000); TASSIGN(c1, 0x4000); TMATMUL_ACC(c1, c0, a, b); }两个示例均使用half输入、float累加器的典型组合符合上节列出的支持类型三元组。手动模式下0x1000等地址由开发者自行规划 Tile 内存布局。汇编示例ASM自动模式# 自动模式由编译器/运行时负责资源放置与调度。 %c_out pto.tmatmul.acc %c_in, %a, %b : (!pto.tile..., !pto.tile..., !pto.tile...) - !pto.tile...手动模式# 手动模式先显式绑定资源再发射指令。 # 可选当该指令包含 tile 操作数时 # pto.tassign %arg0, tile(0x1000) # pto.tassign %arg1, tile(0x2000) %c_out pto.tmatmul.acc %c_in, %a, %b : (!pto.tile..., !pto.tile..., !pto.tile...) - !pto.tile...PTO 汇编形式同步 DPS 对照%acc1 tmatmul.acc %acc0, %a, %b : (!pto.tile..., !pto.tile..., !pto.tile...) - !pto.tile... # AS Level 2 (DPS) pto.tmatmul.acc ins(%c_in, %a, %b : !pto.tile_buf..., !pto.tile_buf..., !pto.tile_buf...) outs(%c_out : !pto.tile_buf...)源码级实现解析累加是如何发生的从源码实现可以清楚看到TMATMUL_ACC与TMATMUL在底层的关键差异——累加器初始化标志。A5Ascend 950PR/950DT实现在 include/pto/npu/a5/TMatmul.hpp 中// TMATMULcmatrixInitVal 为 1表示 C 矩阵初值视为 0 TMatmulPhase, TileRes, TileLeft, TileRight, false, true, true( cMatrix.data(), aMatrix.data(), bMatrix.data(), m, k, n); // TMATMUL_ACCcmatrixInitVal 为 0表示使用 C 矩阵中的真实数值作为初值 TMatmulPhase, TileRes, TileLeft, TileRight, false, false, true( cOutMatrix.data(), aMatrix.data(), bMatrix.data(), m, k, n);源码注释给出了该标志的语义cmatrixInitValIndicates the initial matrix, 1: the number in C matrix is 0, 0: use the real number in C matrix。TMATMUL_ACC_IMPL传入false即保留 C 矩阵中已有的真实数值作为累加初值而TMATMUL传入true即硬件将 C 初值视为 0。这正是两条指令数学语义差异的硬件落地。此外TMATMUL_ACC_IMPL与TMATMUL_IMPL一样会先执行CheckMadValid静态校验与CheckDynamicMmad(m, k, n)动态范围检查且TMATMUL_ACC_IMPL的原地重载直接展开为template AccPhase Phase, typename TileRes, typename TileLeft, typename TileRight PTO_INTERNAL void TMATMUL_ACC_IMPL(TileRes cMatrix, TileLeft aMatrix, TileRight bMatrix) { TMATMUL_ACC_IMPLPhase(cMatrix, cMatrix, aMatrix, bMatrix); }即把同一个 Tile 同时作为输入累加器与输出累加器传入印证了 C 接口重载 3 的原地语义。CPU 仿真实现在 include/pto/cpu/TMatmul.hpp 中CPU 侧实现更直接template typename TileAcc, typename TileLeft, typename TileRight PTO_INTERNAL void TMATMUL_ACC_IMPL(TileAcc cOutMatrix, TileAcc cInMatrix, TileLeft aMatrix, TileRight bMatrix) { TMatmulNzZn(cOutMatrix, cInMatrix, aMatrix, bMatrix); }TMatmulNzZn通过非空指针cInMatrix显式接收累加输入与TMATMUL传nullptr形成对照。这为开发者在 CPU 仿真环境验证 K tile 累加行为提供了可读的实现参照。测试用例验证仓库 NPU 侧测试直接覆盖了原地累加用法。在 tests/npu/a5/src/st/testcase/tmatmul/tmatmul_kernel.cpp 中if (is_first_k_tile) { TMATMUL(cTile, aTile, bTile); // K 维第一个分块从零开始 } else { TMATMUL_ACC(cTile, cTile, aTile, bTile); // 后续分块原地累加 }这是TMATMUL_ACC最典型的生产用法——GEMM K 维循环累加首个 K tile 用TMATMUL完成初始化后续每个 K tile 用TMATMUL_ACC的原地重载在同一块cTile上持续累加配合PIPE_M/PIPE_MTE2事件同步set_flag/wait_flag保证流水线正确性。同样的模式也出现在tests/cpu/st/testcase/tmatmul/tmatmul_kernel.cpp等 CPU 测试中验证了跨平台语义一致性。实战要点小结K tile 累加K 维分块时首个分块用TMATMUL清零初始化后续分块用TMATMUL_ACC(cTile, cTile, a, b)原地累加避免额外的清零与拷贝开销浮点顺序C0作为融合乘加序列的初始累加值参与每一步计算与先算独立点积再加C0的舍入结果可能不同跨 K tile 累加时各分块的累加顺序会影响最终精度需结合算子精度需求权衡分块策略资源绑定手动模式下c0/c1或原地模式下的cMatrix需通过TASSIGN显式分配地址地址规划需避开a/b的 Tile 空间类型约束累加 Tile 类型与输入类型需满足目标平台的(CType, AType, BType)三元组约束如(float, half, half)、(int32_t, int8_t, int8_t)以及 Left/Right/Acc 的分形与布局要求m/k/n 取值维度取自GetValidRow()/GetValidCol()有效尺寸运行时需落在[1, 4095]范围内cInMatrix不做显式断言校验正确性由编程者保证。更多关联指令可参考 TMATMUL、TMATMUL_BIAS、TGEMV_ACC 等文档以及 PTO-Virtual-ISA-Manual 获取完整的虚拟指令集体系说明。赞分享人工智能指令集算子库CANNAscend【免费下载链接】pto-isaParallel Tile Operation (PTO) is a virtual instruction set architecture designed by Ascend CANN, focusing on tile-level operations. This repository offers high-performance, cross-platform tile operations across Ascend platforms.项目地址https://gitcode.com/cann/pto-isa点击查看免费下载相关推荐CANN PTO-ISA 指令详解带偏置融合的矩阵乘法 TMATMUL_BIAS 的语义、约束与多后端实现CANN PTO ISA 指令详解带偏置融合的矩阵乘法 TMATMUL_BIAS 的语义、约束与多后端实现 TMATMUL_BIAS 是 CANN PTOP人工智能指令集算子库CANNAscend如何快速掌握Netgen3D四面体网格生成的终极入门指南如何快速掌握Netgen3D四面体网格生成的终极入门指南 想要在有限元分析中获得高质量的网格吗Netgen作为一款强大的开源3D四面体网格生成器能够帮助您人工智能指令集算子库CANNAscendCANN PTO-ISA 指令详解TMATMUL_BIAS 带偏置的 Tile 级矩阵乘法实现与使用指南CANN PTO ISA 指令详解TMATMUL_BIAS 带偏置的 Tile 级矩阵乘法实现与使用指南 TMATMUL_BIAS 是 CANN pto is人工智能指令集算子库CANNAscend创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
网站建设高端定制企业官网