ARTICLE DETAIL

资讯详情

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

CANN ops-transformer GroupedMatmul 算子设计详解:场景划分、分组策略与 tiling 分核方案

CANN ops-transformer GroupedMatmul 算子设计详解:场景划分、分组策略与 tiling 分核方案 CANN ops-transformer GroupedMatmul 算子设计详解场景划分、分组策略与 tiling 分核方案【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformerGroupedMatmul 是 CANN ops-transformer 仓库中面向 MoEMixture of Experts等场景的核心融合算子它把按专家分组的 token 一次性打包完成所有分组的矩阵乘计算每组维度可各不相同。本文以 gmm/grouped_matmul/docs/GroupedMatmul算子设计介绍.md 为骨架结合仓库中的 tiling 定义、kernel 模板选择与 infershape 校验源码系统讲解算子的场景划分、m/k 轴分组、多 tensor/单 tensor 支持以及 UB buffer 分配、基本块分核与对角线分核等核心设计帮助读者理解并复现该算子在 NPU 上的加速计算方案。1 算子概述GroupedMatmul 的功能是进行分组矩阵乘计算每组矩阵乘的维度大小可以不同。其数学定义如下$$ y_i[m_i,n_i]x_i[m_i,k_i] \times weight_i[k_i,n_i],\quad i1...g $$其中 $g$ 为分组个数$m_i$、$k_i$、$n_i$ 为对应分组的 shape。若考虑 bias基础计算公式为 $y_ix_i\times weight_i bias_i$。在 MoE 网络如moe_gating_top_k → moe_init_routing → grouped_matmul(gate_up_proj) → swiglu → grouped_matmul(down_proj) → moe_finalize_routing见 gmm/grouped_matmul/README.md中GroupedMatmul 通常被调用两次第一次对输入 token 与各专家的上投影gate_up_proj权重做分组矩阵乘得到各专家的中间结果第二次将中间结果与各专家的下投影down_proj权重做分组矩阵乘得到各专家的输出。算子实现时还需考虑以下两个方面场景与模板的多样化需要支持不同的参数与数据类型如有无 bias、不同激活函数类型以及非量化、量化、伪量化等不同场景。不同场景计算流程不同、性能优化方法不同因此实现上被划分为不同的模板各有独立的模板参数AiCore 内存与流水约束硬件上 AiCore 内存大小有限一般完成一个算子的计算需要对数据进行切分并对数据搬运和计算过程进行流水并行排布该过程对算子性能影响非常大也是性能优化阶段的主要调整对象。host 上的 tiling 函数正是为完成该切分与流水而做的参数计算。2 场景划分从功能角度GroupedMatmul 可分为非量化、量化、伪量化三种场景代码层面通过以下三种方式选择具体模板选择方式说明代码载体编译宏通过 x 和 weight 的数据类型编译的宏ORIG_DTYPE_X、ORIG_DTYPE_WEIGHTgrouped_matmul_utils.htilingKey如 x/weight 是否转置TRANS_A/TRANS_Bgrouped_matmul_tiling_key.htilingData如 tiling 中isPerTokenQuant表示是否为 per token 量化grouped_matmul_tiling.h三种场景的含义如下非量化x、weight、y 均为浮点数类型如 float16/bfloat16/float32为纯 cube 场景计算过程由 matmul 高阶 API 实现量化x 和 weight 为低精度整数类型。GroupedMatmul 支持 A8W8 场景包括 per-tensor per-channel 量化和 per-token per-channel 量化简称 per-token 量化伪量化x 为浮点数类型weight 为低精度整数类型。GroupedMatmul 支持 A16W8 和 A16W4 场景。从源码看这些场景由编译宏直接推导。在 grouped_matmul_utils.h 中ORIG_DTYPE_X ORIG_DTYPE_WEIGHT时依据数据类型区分DT_INT8且输出为DT_BF16/DT_FLOAT16定义GMM_QUANT_BF16/GMM_QUANT_FLOAT16量化 A8W8O16DT_INT8且输出为DT_INT32定义GMM_QUANT_INT32量化 A8W8O32输出为DT_INT8定义GMM_QUANT_INT8量化 A8W8O8DT_INT4定义GMM_A4W4量化 A4W4其余浮点类型定义GMM_FLOAT非量化。而当ORIG_DTYPE_X ! ORIG_DTYPE_WEIGHT时定义GMM_ANTI_QUANT伪量化若 x 为DT_INT8、weight 为DT_INT4则进一步定义GMM_ANTI_QUANT_A8W4_MSD与GMM_ANTI_QUANT_A8W4并据此决定MM_DTYPE_Y与输出类型。由此可见tilingKey 侧的模板参数D_T_A、D_T_B、D_T_Y即由这些编译宏推导出的场景枚举GMM_TPL_FLOAT16/GMM_TPL_INT8/GMM_TPL_INT4等见 grouped_matmul_tiling_key.h确定。2.1 per token 量化场景的算法流程以 per token 量化为例GroupedMatmul 的计算过程为matmul(int32) → 反量化(fp32) → mul(fp32) → 激活函数(fp32)(可选) → cast(fp16/bf16)其中mul(fp32)的输入perTokenScale还需要从 shape (m) broadcast 成 (m, n)。整体流程即上图GroupedMatmul 量化场景计算流程图。2.2 分组方式m 轴分组与 k 轴分组针对不同场景GroupedMatmul 可分为m 轴分组切 M与k 轴分组切 K。正向训练过程对 m 轴进行分组反向计算梯度时则需要对 k 轴进行分组。m 轴分组groupType 0$k_i$ 各组相同$weight_i/y_i$ 可以在 $n_i$ 上拼接。m 轴分组可用于非量化正向训练场景、量化场景和伪量化场景。k 轴分组groupType 2$k_i$ 各不相同但 $m_i/n_i$ 每组相同此时 $x_i/weight_i$ 可以在 $k_i$ 上拼接。k 轴分组仅用于非量化训练场景用于求损失关于 weight 的梯度。由于求 weight 梯度时需要对 x 进行转置因此转置后 x 就从 m 轴分组变为 k 轴分组。该约束在 infershape 阶段被显式校验在 grouped_matmul_infershape.cpp 的CheckGroupType中groupType仅允许取GMM_NO_SPLIT-1、GMM_SPLIT_M0、GMM_SPLIT_K2而沿 N 轴分组GMM_SPLIT_N当前不支持会直接报错返回。2.3 多 tensor / 单 tensor 支持GroupedMatmul 算子支持输入输出为多 tensor 或单 tensor。单 tensor指一个 tensor list 中所有分组的 tensor 在 groupType 指定的分组轴上合并为 1 个否则为多 tensor。下表介绍不同方案 tensor 支持的 shape其中单表示单 tensor多表示多 tensor表示顺序为 x、weight、y例如单多单表示支持 x 为单 tensor、weight 为多 tensor、y 为单 tensor 的场景group_typesupported scenariox shapeweight shapey shapeoptional-dynamic inputs shape if neededgroup_list shape if passedper_token_scale shape if passed-1多多多[(M1,K1),(M2,K2),...][(K1,N1),(K2,N2),...][(M1,N1),(M2,N2),...][(N1),(N2),...]not supportnot support0单单单[(M,K)][(G,K,N)][(M,N)][(G,N)](G)(M)0单多单[(M,K)][(K,N),(K,N),...][(M,N)][(N),(N),...](G)not support0多多单[(M1,K1),(M2,K2),...][(K1,N),(K2,N),...][(M,N)][(N),(N),...](G)not support例如在多多多场景x shape{{4,16},{12,16},{16,16}}weight shape{{16,8},{16,8},{16,8}}y shape{{4,8},{12,8},{16,8}}如果要在单单单场景进行相同的计算则 x shape{32,16}weight shape{3,16,8}y shape{32,8}groupList{4,12,16}。单 tensor/多 tensor 的组合在 tilingData 的GMMBaseParams中通过singleWeight、singleX、singleY三个字段显式标记见 grouped_matmul_tiling.h而groupListType字段则记录 groupList 的分组方式0累积和、1各组大小、2[组索引,组大小] 对。3 tiling 设计3.1 tilingData 设计GroupedMatmul 的 tilingData 由三个结构体组合而成定义如下见 grouped_matmul_tiling.hBEGIN_TILING_DATA_DEF(GMMTilingData) TILING_DATA_FIELD_DEF_STRUCT(GMMBaseParams, gmmBaseParams); TILING_DATA_FIELD_DEF_STRUCT(GMMArray, gmmArray); TILING_DATA_FIELD_DEF_STRUCT(TCubeTiling, mmTilingData); END_TILING_DATA_DEF;tilingData 主要包含上述结构里的三个部分GMMBaseParamsGroupedMatmul 的分组数量、AiCore 核数、UB tiling 参数、matmul 分核 tiling 参数等基础 tiling 参数以及是否有激活函数、量化类型per token 或 per tensor、激活函数类型等功能参数。字段示例包括groupNum、coreNum、activeType、ubBaseK、ubBaseN、ubCalSize、singleWeight/singleX/singleY、groupType、quantParam量化场景下 per token 为 1伪量化场景下表示 per-group size等GMMArray当输入是多 tensor 时通过三个数组mList、kList、nList长度 128对应MAX_TENSOR_CONT记录每组 matmul 的 shapekernel 通过GlobalTensor::GetValue()的方式获取。当输入为全单 tensor 时m/k/n 中的一个值在 group_list 中另外两个值在所有的 group 中都相同此时 kernel 只需要访问数组中的第一个值TCubeTilingmatmul 高阶 API 对应的 tilingData。kernel 中为了避免在栈空间中申请GMMArray中的 3 个数组以及避免拷贝这些数组采用GET_TILING_DATA_MEMBER接口只拷贝除GMMArray之外的结构体。从 grouped_matmul.cpp 的GMM_IMP宏可以看到kernel 侧通过GET_TILING_DATA_MEMBER(GMMTilingData, gmmBaseParams, ...)、GET_TILING_DATA_MEMBER(GMMTilingData, mmTilingData, ...)拷贝前两个结构体再通过GET_TILING_DATA_MEMBER_ADDR(GMMTilingData, gmmArray, gmmArrayAddr_, ...)直接以地址方式访问 GMMArray避免大数组的栈拷贝。3.2 UB buffer 分配在初始化阶段GroupedMatmul 需要确定 UB buffer 的分配复用情况。非量化、伪量化的计算过程简单没有 UB buffer 的复用而量化场景多、计算复杂UB buffer 需要复用以提高单次计算的数据量。定义每份 buffer 分配的字节大小比上处理的数据个数baseM × baseN记为ubCalcSize此处 baseM/baseN 为 vector 计算的参数为该 buffer 的份数。以 per token 量化为例假设 baseM 24、baseN 256则ubCalcSize baseM * baseN 6kb。UB buffer 分配如下vector 计算的输入即 matmul 的输出类型为 int32在开启 doubleBuffer 之后输入需要sizeof(int32) * 2即8 份buffervector 计算的输出为 fp16/bf16其需要sizeof(fp16/bf16) * 2即4 份buffervector 计算的中间计算过程需要申请 tmpBuffer其中 broadcast 的输出需要sizeof(fp32)块 buffer反量化的输出sizeof(fp32)块 buffer还需申请 sharedBuffer 进行 buffer 复用包含 broadcast 临时空间、反量化临时空间和 Mul 输出大小为 8 块 buffer。共需要申请的 tmpBuffer 大小为16 块buffer。因此总共需要分配的 UB buffer 为(8 4 16) * 6kb 28 * 6kb 168kb。该分配逻辑在 host tiling 侧落地。在 grouped_matmul_tiling.cpp 的GMMCalUbSize中首先根据场景为ubDivideBlkNum_赋值如UB_STATIC_QUANT_BLOCK_NUM_FP16/BF16、UB_A4W4_BLOCK_NUM、UB_A16W8_BLOCK_NUM_*等见同文件第 1367-1422 行再计算ubCalSize ubSize / ubDivideBlkNum_并按ubBlockAlign_对齐、剩余空间按 32BUB_BLOCK_UNIT_SIZE对齐最终计算出ubBaseK/ubBaseN等 tiling 参数。3.3 基本块分核方案GroupedMatmul 实现时需要处理输入为多个 tensor 的情况即每组 matmul 的 shape 可能各不相同而 kernel 侧不能为每组 matmul 单独配置对应的 matmul 高阶 API 接口实例tiling 结构体和 core 栈空间大小均不允许。为了适配不同 shape 的 matmul 计算GroupedMatmul 采用基本块方式横向分核以 baseM、baseN 为基本块进行分核计算此处 baseM/baseN 为 matmul 的参数。从 kernel 侧看grouped_matmul_utils.h 中定义了静态 tiling 模板的基本块尺寸BASIC_BLOCK_SIZE_128 128、BASIC_BLOCK_SIZE_256 256以及STATIC_TILING_DEPTH_A1_B1 8、STATIC_TILING_STEP_KA_KB 4、DOUBLE_BUFFER_L0A_L0B 2等流水参数grouped_matmul.cpp 中的GetGmmMatmulApiTiling则通过depthA1/depthB1、stepM/stepN/stepKa/stepKb、dbL0A/dbL0B/dbL0C等字段配置 matmul 高阶 API 的静态 tiling保证不同 shape 的分组都能以统一的基本块流水执行。3.4 对角线分核方案按基本块方案分核容易存在同地址访问问题。例如当基本块方案中nDim coreNum时同一时间所有核都在访问左矩阵的相同地址对性能影响较大。因此当基本块数量超过 coreNum时没超过 coreNum 时对角线方案无法解决同地址访问问题可以采用对角线方案同一时间不同核尽量错开对数据的访问。下图中每个方块代表一个输出的基本块数字代表基本块遍历顺序横向分核为原始基本块分核方案。对于适合对角线优化的场景在基本块方案上输入数据会存在多次访问。当 nDim/mDim 很大时对角线不能直接任意往下延伸否则在 k 值比较大的场景下对角线上对应的输入数据均为不同的数据已经加载过一次的数据被新数据从 L2 cache 中替换掉后续需要加载时还是从 DDR 内存中加载导致性能劣化。因此需要限制对角线范围以充分利用 L2 cache 中缓存的数据。将对角线遍历按阈值进行分组。为尽量避免同地址访问两个方向的阈值最好都不小于实际的物理核数因此有$$ \min(T_m, T_n) \geq numCore $$为充分利用 cache一个分组块对应的 $X$、$Weight$、$Y$ 的总数据量最好不超过设备的 L2 cache 大小因此有$$ sizeof(dtype) \cdot (T_m \cdot singleM \cdot K K \cdot T_n \cdot singleN T_m \cdot singleM \cdot T_n \cdot singleN) \leq L2_{size} $$若两个条件不能同时满足须根据实际情况做取舍。从 tiling 实现看grouped_matmul_tiling.cpp 中通过ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::L2, ...)获取设备的 L2 容量用于相关容量判断同时同文件第 754-772 行的IsOutputDisableL2Cache还根据输出大小是否超过 L2 cache、后续算子是否可能复用输出来决定输出是否走 L2 cache这同样是围绕 L2 cache 命中率做的性能取舍。4 模板选择机制的源码级印证场景划分与模板选择在 kernel 编译期以编译宏 → tilingKey → 模板参数的链路完成核心证据如下场景推导宏见 grouped_matmul_utils.h由ORIG_DTYPE_X/ORIG_DTYPE_WEIGHT/ORIG_DTYPE_Y推导出GMM_FLOAT非量化、GMM_QUANT_BF16/GMM_QUANT_FLOAT16/GMM_QUANT_INT32/GMM_QUANT_INT8A8W8 量化、GMM_A4W4A4W4 量化、GMM_ANTI_QUANTA16W8/A16W4 伪量化等宏tilingKey 模板参数见 grouped_matmul_tiling_key.h模板参数包括数据类型D_T_A/D_T_B/D_T_Y、是否转置TRANS_A/TRANS_B、GROUP_LIST_TYPE累积和/各组大小/稀疏、A8W4_KERNEL_TEMPLATEMSD API 反量化、MSD vector 反量化、per-channel 伪量化、per-group 伪量化、autotiling 等 5 类模板、A16W8_KERNEL_TEMPLATE、AIV_AIC_RATIO纯 cube、AIV:AIC1:1、1:2以及IS_ENABLE_FIXED_AXIS等kernel 分发宏见 grouped_matmul.cppGMM_IMP、GMM_CUBE_IMP、GMM_CV_SPLIT_IMP、GMM_A4W4_IMP、GMM_CV_SPLIT_IMP_A8W4_MSD、GMM_CV_SPLIT_IMP_A8W4等宏分别对应非量化纯 cube、量化cubevector 混合、伪量化A8W4 MSD/A8W4、A4W4 等不同计算链路的实例化入口并通过ASCEND_IS_AIV/ASCEND_IS_AIC区分 vector 核与 cube 核上的不同职责如 A8W4 场景下 AIV 核先执行GMMA8W4FakeQuantPreProcess的权重预处理。5 相关文档与代码导航设计文档本文主体 GroupedMatmul算子设计介绍.md另有 GroupedMatmulTransFusionPass.md、aclnnGroupedMatmulV5.md 等接口与融合 pass 文档算子总览与参数说明gmm/grouped_matmul/README.md含 x/weight/bias/scale/offset/antiquantScale/perTokenScale/groupList 等全部输入输出参数、数据类型与各产品系列支持情况tiling 定义op_host/op_tiling/grouped_matmul_tiling.h 与 op_host/op_tiling/grouped_matmul_tiling.cpp含 arch35 下 no_quant、quant、weight_quant 等专用 tiling 目录kernel 实现op_kernel/grouped_matmul.cpp、op_kernel/grouped_matmul_tiling_key.h、op_kernel/grouped_matmul_utils.h以及op_kernel/a16w4_msd/、op_kernel/arch35/下的分场景 kernel 模板infershape 校验op_host/grouped_matmul_infershape.cpp示例与测试examples/含 A8W8、A16W4、MX 量化、weight NZ 等场景样例、tests/ut/op_host/test_grouped_matmul_tiling.cpp、tests/ut/op_kernel/test_grouped_matmul.cpp。综上GroupedMatmul 通过编译宏 tilingKey tilingData三层机制完成场景与模板的静态选择以 m 轴/k 轴分组与多 tensor/单 tensor 组合覆盖 MoE 前向与反向计算并在 tiling 侧通过 UB buffer 复用、基本块分核与带 L2 cache 约束的对角线分核方案在有限 AiCore 内存与 cache 容量下最大化数据吞吐是理解 CANN 融合算子host tiling 切分 kernel 模板化计算设计范式的重要范例。【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表