十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

CANN PTO-ISA TMATMUL_MX 指令详解:基于缩放 Tile 的混合精度/量化 GEMM 编程指南

CANN PTO-ISA TMATMUL_MX 指令详解:基于缩放 Tile 的混合精度/量化 GEMM 编程指南 人工智能指令集算子库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点击查看免费下载导读TMATMUL_MX是 CANN PTO-ISAParallel Tile Operation并行 Tile 虚拟指令集中面向混合精度与量化矩阵乘GEMM的核心指令。它在普通TMATMUL的基础上额外引入aScaleMatrix/bScaleMatrix两个缩放 Tile由硬件在 Cube 单元内部完成mad_mx的 scale 应用从而支撑 FP8、FP4e1m2/e2m1以及 A6 平台上的 HiF4 三级缩放矩阵乘。读完本文你将掌握TMATMUL_MX的数学语义、三级汇编形式同步/SSA/DPS、C Intrinsic 的全部重载形态、A5/A6 的数据类型约束与合法性检查并能依据仓库中的 Auto/Manual 示例与测试用例编写可运行的低比特 GEMM Kernel。TMATMUL_MX 是什么带缩放 Tile 的矩阵乘TMATMUL_MX在指令体系中的定位见 docs/isa/manifest.yamlcategory: Matrix Multiply其官方定义为在支持的目标上执行带额外缩放 Tile 的矩阵乘GEMM用于混合精度 / 量化 matmul。其核心形态为输入左矩阵aMatrixTileLeft、左缩放aScaleMatrixTileLeftScale、右矩阵bMatrixTileRight、右缩放bScaleMatrixTileRightScale输出累加 TilecMatrixTileAcc类型恒为float可选扩展acc形式支持累加输入cInMatrix与输出cOutMatrix分离bias形式在 GEMM 结果上融合一行 bias。从实现看该指令在当前仓库中落地于两个平台A5include/pto/npu/a5/TMatmul.hpp与include/pto/npu/a5/TMatmulImpls.hppA6include/pto/npu/a6/TMatmul.hpp。其中 A5 支持 FP8/FP4 组合A6 在继承这些组合的基础上进一步支持 HiF4详见 TMATMUL_MX_HIF4 变体。数学语义设M aMatrix.GetValidRow()K aMatrix.GetValidCol()N bMatrix.GetValidCol()则结果在有效 matmul 域0 i M0 j N上对应如下矩阵乘$$\mathrm{C}{i,j} \sum{k0}^{K-1} \mathrm{A}{i,k} \cdot \mathrm{B}{k,j}$$与普通TMATMUL不同这里的缩放 TileaScaleMatrix/bScaleMatrix用于配置实现定义的混合精度行为dequant/quant 语义由目标平台决定。也就是说标准公式描述的是概念上的乘累加结果而 FP8/FP4/HiF4 数据在进入 Cube 前/后的实际缩放处理由硬件mad_mx路径完成。MX 格式的维度含义从测试用例 tests/cpu/st/testcase/tmatmul_mx/tmatmul_mx_kernel.cpp 可以看到 MX 缩放 Tile 的典型尺寸约定对M×K的左矩阵aScale为M × kMX其中kMX ceil(kAlign / 32)对K×N的右矩阵bScale为kMX × N。缩放矩阵沿 K 维按 32 元素为一组进行细分这是 MXMicroscaling格式按组共享缩放因子的典型布局。汇编语法同步形式概念级%c tmatmul.mx %a, %a_scale, %b, %b_scale : (!pto.tile..., !pto.tile..., !pto.tile..., !pto.tile...) - !pto.tile... %c_out tmatmul.mx.acc %c_in, %a, %a_scale, %b, %b_scale : (!pto.tile..., !pto.tile..., !pto.tile..., !pto.tile..., !pto.tile...) - !pto.tile... %c tmatmul.mx.bias %a, %a_scale, %b, %b_scale, %bias : (!pto.tile..., !pto.tile..., !pto.tile..., !pto.tile..., !pto.tile...) - !pto.tile...AS Level 1SSA 形式%c pto.tmatmul.mx %a, %a_scale, %b, %b_scale : (!pto.tile..., !pto.tile..., !pto.tile..., !pto.tile...) - !pto.tile... %c_out pto.tmatmul.mx.acc %c_in, %a, %a_scale, %b, %b_scale : (!pto.tile..., !pto.tile..., !pto.tile..., !pto.tile..., !pto.tile...) - !pto.tile... %c pto.tmatmul.mx.bias %a, %a_scale, %b, %b_scale, %bias : (!pto.tile..., !pto.tile..., !pto.tile..., !pto.tile..., !pto.tile...) - !pto.tile...AS Level 2DPS 形式DPSData Parallel Scheduling形式将操作数划分为ins输入与outs输出两组操作数类型从!pto.tile变为!pto.tile_bufpto.tmatmul.mx ins(%a, %a_scale, %b, %b_scale : !pto.tile_buf..., !pto.tile_buf..., !pto.tile_buf..., !pto.tile_buf...) outs(%c : !pto.tile_buf...) pto.tmatmul.mx.acc ins(%c_in, %a, %a_scale, %b, %b_scale : !pto.tile_buf..., !pto.tile_buf..., !pto.tile_buf..., !pto.tile_buf..., !pto.tile_buf...) outs(%c_out : !pto.tile_buf...) pto.tmatmul.mx.bias ins(%a, %a_scale, %b, %b_scale, %bias : !pto.tile_buf..., !pto.tile_buf..., !pto.tile_buf..., !pto.tile_buf..., !pto.tile_buf...) outs(%c : !pto.tile_buf...)C Intrinsic 声明TMATMUL_MX的 C 接口声明在 include/pto/common/pto_instr.hpp共 6 个模板重载全部返回RecordEvent并支持可变参数WaitEvents调用前通过detail::PtoWaitEvents(events...)等待依赖事件// 基本形式c a * b含两侧缩放 template typename TileRes, typename TileLeft, typename TileLeftScale, typename TileRight, typename TileRightScale, typename... WaitEvents PTO_INST RecordEvent TMATMUL_MX(TileRes cMatrix, TileLeft aMatrix, TileLeftScale aScaleMatrix, TileRight bMatrix, TileRightScale bScaleMatrix, WaitEvents ... events); // 带 AccPhase 的基本形式可通过 AccPhase 选择单位标志/累加相位 template AccPhase Phase, typename TileRes, typename TileLeft, typename TileLeftScale, typename TileRight, typename TileRightScale, typename... WaitEvents PTO_INST RecordEvent TMATMUL_MX(TileRes cMatrix, TileLeft aMatrix, TileLeftScale aScaleMatrix, TileRight bMatrix, TileRightScale bScaleMatrix, WaitEvents ... events); // 累加形式cOut cIn a * b template typename TileRes, typename TileLeft, typename TileLeftScale, typename TileRight, typename TileRightScale, typename... WaitEvents PTO_INST RecordEvent TMATMUL_MX(TileRes cOutMatrix, TileRes cInMatrix, TileLeft aMatrix, TileLeftScale aScaleMatrix, TileRight bMatrix, TileRightScale bScaleMatrix, WaitEvents ... events); // 带 AccPhase 的累加形式 template AccPhase Phase, typename TileRes, typename TileLeft, typename TileLeftScale, typename TileRight, typename TileRightScale, typename... WaitEvents PTO_INST RecordEvent TMATMUL_MX(TileRes cOutMatrix, TileRes cInMatrix, TileLeft aMatrix, TileLeftScale aScaleMatrix, TileRight bMatrix, TileRightScale bScaleMatrix, WaitEvents ... events); // 带 Bias 形式c a * b bias template typename TileRes, typename TileLeft, typename TileLeftScale, typename TileRight, typename TileRightScale, typename TileBias, typename... WaitEvents PTO_INST RecordEvent TMATMUL_MX(TileRes cMatrix, TileLeft aMatrix, TileLeftScale aScaleMatrix, TileRight bMatrix, TileRightScale bScaleMatrix, TileBias biasData, WaitEvents ... events); // 带 AccPhase 与 Bias 的形式 template AccPhase Phase, typename TileRes, typename TileLeft, typename TileLeftScale, typename TileRight, typename TileRightScale, typename TileBias, typename... WaitEvents PTO_INST RecordEvent TMATMUL_MX(TileRes cMatrix, TileLeft aMatrix, TileLeftScale aScaleMatrix, TileRight bMatrix, TileRightScale bScaleMatrix, TileBias biasData, WaitEvents ... events);无AccPhase的重载通过MAP_INSTR_IMPL(TMATMUL_MX, ...)分发到TMATMUL_MX_IMPL带AccPhase的重载则直接调用TMATMUL_MX_IMPLPhase(...)从而在保留TMATMUL_MX指令名的同时支持累加单元标志UF选择。底层实现与合法性检查TMATMUL_MX_IMPL在 A5 侧位于 include/pto/npu/a5/TMatmulImpls.hpp其执行序列为从 Tile 取出运行时有效形状m aMatrix.GetValidRow()、k aMatrix.GetValidCol()、n bMatrix.GetValidCol()调用CheckDynamicMmad(m, k, n)做动态形状检查k、n的有效范围受MMAD_MAX_SUPPORT_LENGTH约束即[1, 4095]调用CheckMadMxValidTileRes, TileLeft, TileLeftScale, TileRight, TileRightScale()做编译期合法性检查调用TMatmulMxPhase, ...内部落到硬件mad_mx内建函数完成计算。Bias 形式在检查前还会对TileBias做两个static_assertTileBias::DType必须为float且TileBias::Loc TileType::Bias、Rows 1单行 bias。累加形式TMATMUL_MX_IMPL则传入cmatrixInitVal false语义即复用真实 C 矩阵初值参与累加。约束条件A5 静态检查详解CheckMadMxValid在 include/pto/npu/a5/TMatmul.hpp 中实现包含以下静态断言数据类型组合C恒为float缩放 Tile 恒为float8_e8m0_t支持的四类 (C, A, B) 组合由isSupportedFp4Combo/isSupportedFp8Combo定义见 include/pto/npu/a5/TMatmulCommon.hppFP8float累加(float, float8_e4m3_t, float8_e4m3_t)(float, float8_e4m3_t, float8_e5m2_t)(float, float8_e5m2_t, float8_e4m3_t)(float, float8_e5m2_t, float8_e5m2_t)FP4float累加(float, float4_e1m2x2_t, float4_e1m2x2_t)(float, float4_e1m2x2_t, float4_e2m1x2_t)(float, float4_e2m1x2_t, float4_e2m1x2_t)(float, float4_e2m1x2_t, float4_e1m2x2_t)形状约束TileLeft::Cols % 64 0aMatrixCol必须是 64 的倍数BASEK 64FP4 类型额外要求TileLeft::Cols % 2 0。Fractal 布局约束左矩阵Loc TileType::Left、非行主序、SFractal SLayout::RowMajor右矩阵Loc TileType::Right、行主序、SFractal SLayout::ColMajor累加 TileLoc TileType::Acc、非行主序、SFractal SLayout::RowMajor。L0C 容量约束static_assert(accBytes PTO_L0C_SIZE_BYTES, TMatmulMX:accumulator (Rows*Cols*sizeof(out)) exceeds L0C capacity.);即Rows × Cols × sizeof(float)不得超过片上 L0C 容量PTO_L0C_SIZE_BYTES。代码示例Auto 模式与 Manual 模式Auto 模式资源由编译器/运行时管理#include pto/pto-inst.hpp using namespace pto; void example_auto() { using A TileLeftfloat8_e5m2_t, 16, 64; using B TileRightfloat8_e5m2_t, 64, 32; using ScaleA TileLeftScalefloat8_e8m0_t, 16, 2; using ScaleB TileRightScalefloat8_e8m0_t, 2, 32; using Bias TileTileType::Bias, float, 1, 32; using C TileAccfloat, 16, 32; A a; B b; ScaleA scaleA; ScaleB scaleB; Bias bias; C c; TMATMUL_MX(c, a, scaleA, b, scaleB, bias); }注意本例的 Scale 维度ScaleA为16×2对应kMX 64/32 2ScaleB为2×32与 MX 格式按 32 元素一组共享 scale 的约定一致。Manual 模式资源需显式绑定#include pto/pto-inst.hpp using namespace pto; void example_manual() { using A TileLeftfloat8_e5m2_t, 16, 64; using B TileRightfloat8_e5m2_t, 64, 32; using ScaleA TileLeftScalefloat8_e8m0_t, 16, 2; using ScaleB TileRightScalefloat8_e8m0_t, 2, 32; using Bias TileTileType::Bias, float, 1, 32; using C TileAccfloat, 16, 32; A a; B b; ScaleA scaleA; ScaleB scaleB; Bias bias; C c; TASSIGN(a, 0x1000); TASSIGN(b, 0x2000); TASSIGN(scaleA, GetScaleAddr(a.data())); TASSIGN(scaleB, GetScaleAddr(b.data())); TASSIGN(bias, 0x3000); TASSIGN(c, 0x4000); TMATMUL_MX(c, a, scaleA, b, scaleB, bias); }Manual 模式中GetScaleAddr用于根据数据 Tile 的地址推导其 MX 缩放 Tile 的地址。这一用法在真实 Kernel 中同样可见CPU 参考实现 tests/cpu/st/testcase/tmatmul_mx/tmatmul_mx_kernel.cpp 中通过TASSIGN显式编排 GM 地址后以uint64_t scaleAAddr GetScaleAddr(aTile.data());取得 scale 地址并赋给 scale Tile随后执行TLOAD加载后调用TMATMUL_MX(cTile, aTile, aScaleTile, bTile, bScaleTile[, biasTile])__PTO_AUTO__宏控制走 Auto 还是 Manual 路径。汇编示例Auto 模式# Auto mode: compiler/runtime-managed placement and scheduling. %c pto.tmatmul.mx %a, %a_scale, %b, %b_scale : (!pto.tile..., !pto.tile..., !pto.tile..., !pto.tile...)Manual 模式# Manual mode: resources must be bound explicitly before issuing the instruction. # Optional for tile operands: # pto.tassign %arg0, tile(0x1000) # pto.tassign %arg1, tile(0x2000) %c pto.tmatmul.mx %a, %a_scale, %b, %b_scale : (!pto.tile..., !pto.tile..., !pto.tile..., !pto.tile...)PTO 汇编形式%c pto.tmatmul.mx %a, %a_scale, %b, %b_scale : (!pto.tile..., !pto.tile..., !pto.tile..., !pto.tile...) # AS Level 2 (DPS) pto.tmatmul.mx ins(%a, %a_scale, %b, %b_scale : !pto.tile_buf..., !pto.tile_buf..., !pto.tile_buf..., !pto.tile_buf...)TMATMUL_MX_HIF4 变体A6 HiFloat4 Cube 矩阵乘A6dav-920r1平台上的hifloat4x2_t重载构成 HiF4 Cube 矩阵乘A、B 操作数均为 HiF44-bit 打包数据并各自携带三级 HiF4 缩放Ea/Eb/Ec结果在 L0C 中以 FP32 累加通常在存储时经 FIXPIPE 转为 BF16。完整说明见 docs/isa/TMATMUL_MX_HIF4.md。典型流水线GM (BF16) ──TLOAD──▶ L1 ──TEXTRACT──▶ L0A/L0B L0AMX/L0BMX ──TMATMUL_MX──▶ L0C (FP32) ──TSTORE──▶ GM (BF16)TLOAD加载 HiF4 数据GM→L1与缩放字节GM→L1采用HIF4_A_ZZ/HIF4_B_NN布局数据使用copy_gm_to_cbuf_multi_nd2nz完成 ND→NZ fractal 变换TEXTRACT数据经load_c_buf_to_ca_s4从 L1→L0A/L0B缩放经load_c_buf_to_ca_mx从 L1→L0AMX/L0BMXTMATMUL_MXCube 单元的mad_mx配合hifloat4x2_t类型标签执行三级 Ea/Eb/Ec 缩放由硬件内部施加TSTOREFIXPIPE 将 FP32→BF16 后写回 GM。缩放布局[16,4] 64B cell每个 HiF4 缩放 patch 覆盖一个 M 或 N fractal × 一个 K-group64 元素共 64 字节bytes 0..31: [Ea(g0), Eb(g0), Ea(g1), Eb(g1), ... × 16 groups] (EaEb half) bytes 32..63: [Ec_lo(g0), Ec_hi(g0), ... × 16 groups] (Ec half)Ea8-bit e6m2 指数每 64 元素一组Eb8-bit 打包每 8 元素子组一个8 个 bit 全部保留Ec16-bit 打包每 4 元素子组一个。CCE 的TQuant通过pstupredicate → align 寄存器vstasalign → UB写入 Eb在输出频率下打包 predicate从而避免旧DS_B8降采样丢失 Eb 的 b4–b7 位对应实现见 include/pto/npu/a6/TQuant.hpp。L0C 容量与 N 方向分片L0C 为 256 KB累加精度 FP324 B/元素。当M × N × 4 256 KB时Kernel 需沿 N 分片tileN floor(L0C_SIZE / (M × 4))并向 64 取整满足 TEXTRACT 列对齐A 侧在 N 分片循环前只加载 抽取一次跨 N tile 共享每次迭代抽取宽度为tileN的 B 列切片执行一次TMATMUL_MX将M × tileN结果块按offset j × tileN、stride N写回 GM。累加类型约束L0C 累加器恒为float4 B由CheckMadMxValid保证static_assert(Rows × Cols × sizeof(float) PTO_L0C_SIZE_BYTES)。A6 数据类型组合扩展A6 的CheckMadMxValidinclude/pto/npu/a6/TMatmul.hpp在 A5 的 FP8/FP4 基础上扩展出更多组合包括(float8_e4m3_t, float4_e2m1x2_t)等 FP8×FP4 混合isSupportedFp8Fp4Combo(half, float4_e2m1x2_t)、(bfloat16_t, float4_e2m1x2_t)等 FP16/BF16×FP4isSupportedFp16Fp4Combo/isSupportedBf16Fp4Combo(float8_e4m3_t, hifloat4x2_t)、(half, hifloat4x2_t)、(bfloat16_t, hifloat4x2_t)等 FP8/FP16/BF16×HiF4isSupportedFp8Hif4Combo/isSupportedFp16Hif4Combo/isSupportedBf16Hif4Combo(hifloat4x2_t, hifloat4x2_t)isSupportedHif4Combo。上述 HiF4 相关组合仅在PTO_NPU_ARCH_A6宏下启用其他架构下对应 trait 恒为falseHiF4/FP4 组合同样要求TileLeft::Cols % 2 0。A6 测试 Kernel tests/npu/a6/src/st/testcase/tmatmul_mx/tmatmul_mx_kernel.cpp 展示了bfloat16_t × hifloat4x2_t的用法其中 scale Tile 经GetScaleAddr(al0.data())绑定到 L0A/L0B 对应地址。测试用例一览平台用例路径覆盖内容CPU STtests/cpu/st/testcase/tmatmul_mx/MX GEMM 的 CPU 参考实现含 FP8/FP4、Bias/非 Bias、ND/DN/ZZ 多种 GM 布局tmatmul_mx_nddn目录覆盖 ND/DN 布局A5 NPU STtests/npu/a5/src/st/testcase/tmatmul_mx/A5 平台 MX 矩阵乘端到端验证A6 NPU STtests/npu/a6/src/st/testcase/tmatmul_mx/A6 平台 MX 矩阵乘含 HiF4 用例HiF4 相关测试docs/isa/TMATMUL_MX_HIF4.md 中记录包括tmatmul_mx_hif4覆盖 128×128×128、128×256×128、256×128×128、64×64×64、256×256×256、128×512×128、512×128×512、128×128×256、256×128×512 等形状的 HiF4 Cube 端到端验证与tmatmul_mx_e1m2128×128×128 的 e1m2 MX oracle复用同一流水线但使用普通 MX scale。相关指令与进一步阅读基础矩阵乘TMATMUL、累加形式TMATMUL_ACC、带 Bias 的TMATMUL_BIAS见 docs/isa/TMATMUL.md、docs/isa/TMATMUL_ACC.md、docs/isa/TMATMUL_BIAS.mdHiF4 量化指令见 docs/isa/TQUANT_HIF4.mdCCE 实现在include/pto/npu/a6/TQuant.hpp指令目录与分类索引见 docs/isa/README.md 与 docs/isa/manifest.yaml操作数搬运配套指令TLOAD见 docs/isa/TLOAD.md、TEXTRACT见 docs/isa/TEXTRACT.md、TSTORE见 docs/isa/TSTORE.md。总结TMATMUL_MX是 PTO-ISA 面向量化/混合精度矩阵乘的关键指令其设计要点可归纳为双缩放 TileaScaleMatrix/bScaleMatrix将 MX 格式的按组缩放因子随数据一同送入 Cubemad_mx在硬件内部完成 scale 应用软件侧无需显式反量化三套语法层级概念同步形式、AS Level 1SSA、AS Level 2DPS配合 C Intrinsic 的 6 个重载覆盖基本 / 累加 / Bias 三种语义编译期强约束通过CheckMadMxValid的static_assert在编译期锁定数据类型组合A5FP8/FP4A6扩展 FP4 混合与 HiF4、fractal 布局、64 对齐的 K 维与 L0C 容量非法组合在编译期即被拦截平台差异化A6 的 HiF4 变体具备 TLOAD→TEXTRACT→TMATMUL_MX→TSTORE 完整流水线、64B scale cell 布局以及 L0C 容量驱动的 N 方向分片策略。开发者可按 Auto/Manual 两种模式接入Auto 模式由编译器托管资源与调度Manual 模式通过TASSIGN显式绑定地址数据、scale 与 bias并可使用GetScaleAddr从数据 Tile 地址推导 MX scale 地址。参考 tests/cpu/st/testcase/tmatmul_mx/tmatmul_mx_kernel.cpp 可快速获得一份完整可运行的最小实现。赞分享人工智能指令集算子库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 指令详解TGEMV_MX——带缩放 Tile 的混合精度 GEMV 指令CANN PTO ISA 指令详解TGEMV_MX——带缩放 Tile 的混合精度 GEMV 指令 导读 TGEMV_MX 是 CANN pto isaPa人工智能指令集算子库CANNAscendCANN PTO-ISA TLRELU 指令详解基于标量斜率的 Leaky ReLU Tile 运算CANN PTO ISA TLRELU 指令详解基于标量斜率的 Leaky ReLU Tile 运算 TLRELU 是 CANN PTOParallel T人工智能指令集算子库CANNAscendCANN pto-isa TSELS 指令详解基于 Mask 的 Tile 与标量逐元素选择CANN pto isa TSELS 指令详解基于 Mask 的 Tile 与标量逐元素选择 TSELSTile Select Scalar是 CANN人工智能指令集算子库CANNAscend创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表