ATB 算子深度解析:gmm_deq_swiglu_quant_gmm_deq 四阶段量化融合 Pipeline 的实现与使用
【免费下载链接】ascend-transformer-boost本项目是CANN提供的是一款高效、可靠的Transformer加速库,基于华为Ascend AI处理器,提供Transformer定制化场景的高性能融合算子。项目地址: https://gitcode.com/cann/ascend-transformer-boost
本文以 CANN ascend-transformer-boost 仓库中的
gmm_deq_swiglu_quant_gmm_deq算子知识条目为主线,结合 算子声明与校验源码、OpsRunner 实现、参数定义 以及 kernel 实现 展开,系统讲解该算子的复合 Pipeline、dtype 流转、输入输出契约、参数约束与运行时架构,帮助读者理解其设计原理并掌握在 Atlas 800I A2 推理产品上的调用方式。
一、算子定位:一个 XL 规模的 activation 类复合算子
gmm_deq_swiglu_quant_gmm_deq是 ATB(Ascend Transformer Boost)算子体系中的一个composite(复合)类型算子,归类于activation类别,规模等级为tier: XL。它并非一个单一数学运算,而是把 MoE(Mixture of Experts)类模型中一段典型的量化推理计算链完整融合进单个算子,从而减少中间张量的搬运与 kernel 启动开销。
从知识条目与源码可以确认如下事实:
- 类别:activation,类型 composite(src/ops/ops_infer/gmm_deq_swiglu_quant_gmm_deq/);
- 源码构成:
src/ops/ops_infer/下共 4 个文件,即 2 个 operation 文件 + 2 个 ops_runner 文件; - 功能语义(参数注释):
GroupedMatmul1 + Dequant1 + Swiglu + Quant + GroupedMatmul2 + Dequant2融合算子,即两个带反量化的分组矩阵乘中间夹着 SwiGLU 激活与再量化。
从仓库结构看,该算子的 Kernel 侧位于 src/kernels/mixkernels/gmm_deq_swiglu_quant_gmm_deq/,包含op_kernel/(N128/N256 两个 kernel 变体)、tiling/(tiling 实现)与公共头文件,属于 mixkernels 融合 kernel 体系。
二、复合 Pipeline:四阶段串联的计算骨架
该算子的核心是一条 4 阶段串联流水线:
GMM(Dequant) → SwiGLU → Quant → GMM(Dequant)逐段展开为:
- 阶段 1:GMM1 + Dequant1。第一个分组矩阵乘接收 int8 激活 x1 与 int8 权重 weight1,累加得到 fp32 中间结果后,使用 scale1(per-channel/权重侧尺度)与 perTokenScale1(per-token 尺度)执行反量化,输出 fp16;
- 阶段 2:SwiGLU。对 fp16 结果施加 SwiGLU 门控激活(即
silu(x) * y形式的 gated linear unit),中间仍为 fp16; - 阶段 3:Quant。将 fp16 激活再次量化为 int8,作为第二个矩阵乘的输入,同时生成对应的量化尺度;
- 阶段 4:GMM2 + Dequant2。第二个分组矩阵乘接收 int8 激活与 int8 权重 weight2,累加后使用 scale2 反量化,最终输出 fp16。
整条链路把两次量化矩阵乘、两次反量化、一次 SwiGLU、一次再量化全部融合在单一算子内。从 InferShape 实现 看,输出形状由输入 x1 的 M 维直接决定,即[m, SUPPORTED_N2],其中SUPPORTED_N2 = 7168。
三、dtype 流转与两个 dtype 断点
知识条目明确指出该算子的数据类型流转为:
int8 → fp16 → fp16 → int8 → fp16对照源码可以逐段印证:
| 阶段 | 数据类型 | 依据 |
|---|---|---|
| GMM1 输入(x1、weight1) | int8 | kernel 校验:x1.dtype == TENSOR_DTYPE_INT8、weight1.dtype == TENSOR_DTYPE_INT8 |
| GMM1 Dequant 之后(SwiGLU 输入/输出) | fp16 | 知识条目定义 |
| Quant 之后(GMM2 输入) | int8 | kernel 校验:weight2.dtype == TENSOR_DTYPE_INT8 |
| 最终输出 | fp16 | InferShape:outDesc.dtype = ACL_FLOAT16 |
这条链路上存在2 个 dtype 断点:
- 断点 1(int8 → fp16):位于 GMM1 累加与反量化之后。int8 激活与 int8 权重在矩阵乘单元中累加为 fp32,再经 scale1 × perTokenScale1 反量化后落到 fp16,进入 SwiGLU;
- 断点 2(fp16 → int8):位于 Quant 阶段。SwiGLU 的 fp16 输出被再量化为 int8,作为 GMM2 的激活输入。
正是这两个断点让算子可以全程使用 int8 矩阵乘的高吞吐计算单元(GMM1、GMM2 均为 int8 输入),同时在精度敏感的反量化与激活环节保持 fp16,兼顾性能与精度。两个尺度张量(scale1 与 scale2)与 per-token 尺度(perTokenScale1)均为 fp32(kernel 校验),保证了反量化计算的精度。
四、输入输出契约:7 输入 + 1 输出
该算子固定为7 个输入、1 个输出(源码常量INPUT_NUM = 7、OUTPUT_NUM = 1,见 operation.cpp)。输入按顺序为:
| 序号 | 名称 | 语义 | 维度 | dtype | format |
|---|---|---|---|---|---|
| 0 | x1 | 第一个矩阵乘的激活 | 2 维[m, K1] | int8 | ND |
| 1 | weight1 | 第一个矩阵乘的权重 | 3 维[groupCount, K1, N1] | int8 | FRACTAL_NZ |
| 2 | scale1 | 第一个矩阵乘的反量化尺度 | 2 维[groupCount, N1] | fp32 | ND |
| 3 | perTokenScale1 | per-token 反量化尺度 | 1 维[m] | fp32 | ND |
| 4 | groupList | 分组边界列表(cumsum 形式) | 1 维[groupCount] | int64 | ND |
| 5 | weight2 | 第二个矩阵乘的权重 | 3 维[groupCount, N2, K2](转置语义) | int8 | FRACTAL_NZ |
| 6 | scale2 | 第二个矩阵乘的反量化尺度 | 2 维[groupCount, N2] | fp32 | ND |
输出为 1 个 fp16、ND 格式、形状[m, 7168]的张量。
上述 dtype 与 format 约束在 kernel 侧校验函数 中逐一强制检查(CheckX1、CheckWeight1、CheckScale1、CheckPerTokenScale1、CheckGroupList、CheckWeight2、CheckScale2),任何一项不满足都会直接返回ERROR_INFERSHAPE_ERROR。
值得注意的转置语义:两个权重的维度解读不同。weight1 使用非转置语义(K 在 dim1、N 在 dim2),而 weight2 使用转置语义(IsTrans = true,N 在 dim1、K 在 dim2),这与参数transposeWeightUp = false、transposeWeightDown = true的约束一一对应(见下节)。
五、参数定义与强制约束
参数结构体GmmDeqSwigluQuantGmmDeqParam定义于 include/atb/infer_op_params.h,包含三个枚举与两个布尔开关:
枚举定义与默认值
| 参数 | 可选值 | 默认值 | 说明 |
|---|---|---|---|
outputType | OUTPUT_FLOAT16 = 0、OUTPUT_BFLOAT16、OUTPUT_INVALID | OUTPUT_FLOAT16 | 输出数据类型 |
groupListType | GROUP_LIST_CUMSUM = 0、GROUP_LIST_SINGLE、GROUP_LIST_INVALID | GROUP_LIST_CUMSUM | groupList 的编码形式 |
weightUpPermuteType | PERMUTE_N256 = 0、PERMUTE_N128、PERMUTE_INVALID | PERMUTE_N256 | weight1 与 scale1 的重排方式 |
transposeWeightUp | bool | false | weight1 是否转置 |
transposeWeightDown | bool | true | weight2 是否转置 |
当前版本下的强制约束(ParamCheck 实现):
- 平台限定:仅支持Atlas 800I A2 推理产品(
Config::Is910B()校验),其他平台直接报错拒绝运行; outputType仅支持OUTPUT_FLOAT16;groupListType仅支持GROUP_LIST_CUMSUM(groupList 以 cumsum 前缀和形式给出各分组边界);weightUpPermuteType不能为PERMUTE_INVALID,且必须与 kernel 变体匹配(N256/N128);transposeWeightUp仅支持false,transposeWeightDown仅支持true。
这些约束在 operation 层与 kernel 层被重复校验(kernel 侧 CheckGmmDeqSwigluQuantGmmDeq),保证两层校验语义一致。
形状约束(CheckInTensorsShape)同样严格:
| 约束项 | 取值 |
|---|---|
| M(x1 的 batch/token 维) | ≤ 131072(MAX_M) |
| 分组数(groupCount) | ≤ 32(MAX_GROUP_COUNT) |
| x1 的 K 维 | 固定 7168(SUPPORTED_K1) |
| weight1 的 K/N | 固定 7168 / 4096(SUPPORTED_K1/SUPPORTED_N1) |
| weight2 的 N/K | 固定 7168 / 2048(SUPPORTED_N2/SUPPORTED_K2) |
| 输出 N 维 | 固定 7168(SUPPORTED_N2) |
也就是说,这是一个面向特定模型形状深度定制的融合算子,形状不匹配时会在InferShapeCheckImpl与SetupCheckImpl阶段(operation.cpp)返回ERROR_INVALID_TENSOR_DIM/ERROR_INVALID_TENSOR_DIM_NUM等错误码。
六、运行时架构:单一 OpsRunner,无 ACLNN 路径
知识条目强调该算子使用单一 OpsRunner执行,不经过 ACLNN。这在源码中有清晰印证:
- GmmDeqSwigluQuantGmmDeqOperation::CreateRunner 通过
RunnerTypeRegister::GetRunnerTypeIdx("GmmDeqSwigluQuantGmmDeqOpsRunner")从 RunnerPool 获取 runner 类型,并MallocRunner<GmmDeqSwigluQuantGmmDeqOpsRunner>创建实例,失败时降级为直接make_shared创建; - GmmDeqSwigluQuantGmmDeqOpsRunner 继承自
OpsRunner,通过REG_RUNNER_TYPE注册,其核心工作是重写SetupKernelGraph; - SetupKernelGraph 构建一个仅含1 个节点的 KernelGraph:把 7 个输入张量与 1 个输出张量绑定到名为
GmmDeqSwigluQuantGmmDeqOperation的节点上,并把用户传入的 ATB 参数(infer::GmmDeqSwigluQuantGmmDeqParam)转换为 MKI 层参数AtbOps::OpParam::GmmDeqSwigluQuantGmmDeq(通过三个GetAtbOps*转换函数完成枚举映射)。
整个调用链为:
ATB Operation(GmmDeqSwigluQuantGmmDeqOperation) → RunnerPool 获取 GmmDeqSwigluQuantGmmDeqOpsRunner → SetupKernelGraph 构建单节点 KernelGraph → MKI 层 OperationBase 派生类(GmmDeqSwigluQuantGmmDeqOperation) → GetBestKernel 按 weightUpPermuteType 选择 N256/N128 kernel → tiling 计算并下发执行MKI 层 operation 的GetBestKernel是 kernel 选择的入口:weightUpPermuteType == PERMUTE_N256时选择GmmDeqSwigluQuantGmmDeqN256Kernel,PERMUTE_N128时选择GmmDeqSwigluQuantGmmDeqN128Kernel。这与 op_kernel 目录 下的gmm_deq_swiglu_quant_gmm_deq_n256.cpp、gmm_deq_swiglu_quant_gmm_deq_n128.cpp两个 kernel 变体一一对应。Tiling 则由 gmm_deq_swiglu_quant_gmm_deq_tiling.cpp 中的GmmDeqSwigluQuantGmmDeqTiling完成(声明见 tiling.h)。
参数更新方面,SetParam(ops_runner.cpp)使用operator==比较新旧参数(逐字段比较outputType、groupListType、weightUpPermuteType、transposeWeightUp、transposeWeightDown,定义于 ops_runner.h),仅在参数真正变化时才置位isParamUpdated_触发重建,避免无谓的 kernel 重编译。
七、算子族谱:与相邻算子的关系
知识条目给出了该算子在 ATB 算子族中的两个相关算子:
mm_deq_swiglu_quant_mm_deq(MM 变体):位于 src/ops/ops_infer/mm_deq_swiglu_quant_mm_deq/,对应参数结构 MmDeqSwigluQuantMmDeqParam。二者 Pipeline 结构一致(Matmul1 + Dequant1 + Swiglu + Quant + Matmul2 + Dequant2),区别在于gmm_deq_swiglu_quant_gmm_deq使用grouped matmul(分组矩阵乘)语义并引入groupList输入,而 MM 变体为普通矩阵乘,因此前者可支持多组权重的分组计算(groupCount ≤ 32),适用于 MoE 路由场景,后者面向单组权重;swiglu_quant(单 SwiGLU + Quant):位于 src/ops/ops_infer/swiglu_quant/,可视为本算子 Pipeline 中"阶段 2 + 阶段 3"的独立切片,用于不需要两端量化矩阵乘、仅需完成 SwiGLU 激活与再量化的场景。
三者的关系可以概括为:swiglu_quant是本算子中间两阶段的独立版本,mm_deq_swiglu_quant_mm_deq是本算子的非分组(非 MoE)变体,而gmm_deq_swiglu_quant_gmm_deq是面向分组量化推理的完整融合形态。
八、使用建议与平台限制
综合源码约束,使用该算子需要注意以下几点:
- 平台前提:仅在Atlas 800I A2 推理产品上可用(ParamCheck),其他硬件平台会直接报错;
- 形状强约束:M ≤ 131072、分组数 ≤ 32,第一段矩阵乘固定为
K=7168、N=4096,第二段固定为K=2048、N=7168,输入 x1 的 K 维固定 7168,输出固定[m, 7168]fp16。该算子是面向特定模型规格深度定制,接入前务必核对模型形状; - 参数取值:保持默认值即可覆盖绝大多数场景——
outputType=OUTPUT_FLOAT16、groupListType=GROUP_LIST_CUMSUM、transposeWeightUp=false、transposeWeightDown=true;weightUpPermuteType需与预处理好的权重重排方式一致(N256 或 N128),这会决定实际选择哪个 kernel 变体; - 数据准备:x1、weight1、weight2 必须为 int8,weight 使用 FRACTAL_NZ 格式;scale1、scale2、perTokenScale1 为 fp32 ND 格式;groupList 为 int64、以 cumsum 前缀和形式编码分组边界;
- 运行时:走单一
GmmDeqSwigluQuantGmmDeqOpsRunner(OpsRunner 路径),不经过 ACLNN;参数未变化时不会触发 kernel 重建,可放心复用 runner 实例。
参考与进一步阅读
- 算子知识条目:.agent/knowledge/ops/activation/gmm_deq_swiglu_quant_gmm_deq/index.md
- ATB Operation 层:src/ops/ops_infer/gmm_deq_swiglu_quant_gmm_deq/
- Kernel 层(op_kernel + tiling):src/kernels/mixkernels/gmm_deq_swiglu_quant_gmm_deq/
- 参数结构定义:include/atb/infer_op_params.h
- 相关算子:MM 变体 src/ops/ops_infer/mm_deq_swiglu_quant_mm_deq/、单 SwiGLU+Quant src/ops/ops_infer/swiglu_quant/
【免费下载链接】ascend-transformer-boost本项目是CANN提供的是一款高效、可靠的Transformer加速库,基于华为Ascend AI处理器,提供Transformer定制化场景的高性能融合算子。项目地址: https://gitcode.com/cann/ascend-transformer-boost
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考