ATB 算子深度解析:gmm_deq_swiglu_quant_gmm_deq 四阶段量化融合 Pipeline 的实现与使用
2026/9/18 4:37:38 网站建设 项目流程

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. 阶段 1:GMM1 + Dequant1。第一个分组矩阵乘接收 int8 激活 x1 与 int8 权重 weight1,累加得到 fp32 中间结果后,使用 scale1(per-channel/权重侧尺度)与 perTokenScale1(per-token 尺度)执行反量化,输出 fp16;
  2. 阶段 2:SwiGLU。对 fp16 结果施加 SwiGLU 门控激活(即silu(x) * y形式的 gated linear unit),中间仍为 fp16;
  3. 阶段 3:Quant。将 fp16 激活再次量化为 int8,作为第二个矩阵乘的输入,同时生成对应的量化尺度;
  4. 阶段 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)int8kernel 校验:x1.dtype == TENSOR_DTYPE_INT8weight1.dtype == TENSOR_DTYPE_INT8
GMM1 Dequant 之后(SwiGLU 输入/输出)fp16知识条目定义
Quant 之后(GMM2 输入)int8kernel 校验:weight2.dtype == TENSOR_DTYPE_INT8
最终输出fp16InferShape: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 = 7OUTPUT_NUM = 1,见 operation.cpp)。输入按顺序为:

序号名称语义维度dtypeformat
0x1第一个矩阵乘的激活2 维[m, K1]int8ND
1weight1第一个矩阵乘的权重3 维[groupCount, K1, N1]int8FRACTAL_NZ
2scale1第一个矩阵乘的反量化尺度2 维[groupCount, N1]fp32ND
3perTokenScale1per-token 反量化尺度1 维[m]fp32ND
4groupList分组边界列表(cumsum 形式)1 维[groupCount]int64ND
5weight2第二个矩阵乘的权重3 维[groupCount, N2, K2](转置语义)int8FRACTAL_NZ
6scale2第二个矩阵乘的反量化尺度2 维[groupCount, N2]fp32ND

输出为 1 个 fp16、ND 格式、形状[m, 7168]的张量。

上述 dtype 与 format 约束在 kernel 侧校验函数 中逐一强制检查(CheckX1CheckWeight1CheckScale1CheckPerTokenScale1CheckGroupListCheckWeight2CheckScale2),任何一项不满足都会直接返回ERROR_INFERSHAPE_ERROR

值得注意的转置语义:两个权重的维度解读不同。weight1 使用非转置语义(K 在 dim1、N 在 dim2),而 weight2 使用转置语义(IsTrans = true,N 在 dim1、K 在 dim2),这与参数transposeWeightUp = falsetransposeWeightDown = true的约束一一对应(见下节)。

五、参数定义与强制约束

参数结构体GmmDeqSwigluQuantGmmDeqParam定义于 include/atb/infer_op_params.h,包含三个枚举与两个布尔开关:

枚举定义与默认值

参数可选值默认值说明
outputTypeOUTPUT_FLOAT16 = 0OUTPUT_BFLOAT16OUTPUT_INVALIDOUTPUT_FLOAT16输出数据类型
groupListTypeGROUP_LIST_CUMSUM = 0GROUP_LIST_SINGLEGROUP_LIST_INVALIDGROUP_LIST_CUMSUMgroupList 的编码形式
weightUpPermuteTypePERMUTE_N256 = 0PERMUTE_N128PERMUTE_INVALIDPERMUTE_N256weight1 与 scale1 的重排方式
transposeWeightUpboolfalseweight1 是否转置
transposeWeightDownbooltrueweight2 是否转置

当前版本下的强制约束(ParamCheck 实现):

  • 平台限定:仅支持Atlas 800I A2 推理产品Config::Is910B()校验),其他平台直接报错拒绝运行;
  • outputType仅支持OUTPUT_FLOAT16
  • groupListType仅支持GROUP_LIST_CUMSUM(groupList 以 cumsum 前缀和形式给出各分组边界);
  • weightUpPermuteType不能为PERMUTE_INVALID,且必须与 kernel 变体匹配(N256/N128);
  • transposeWeightUp仅支持falsetransposeWeightDown仅支持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

也就是说,这是一个面向特定模型形状深度定制的融合算子,形状不匹配时会在InferShapeCheckImplSetupCheckImpl阶段(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时选择GmmDeqSwigluQuantGmmDeqN256KernelPERMUTE_N128时选择GmmDeqSwigluQuantGmmDeqN128Kernel。这与 op_kernel 目录 下的gmm_deq_swiglu_quant_gmm_deq_n256.cppgmm_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==比较新旧参数(逐字段比较outputTypegroupListTypeweightUpPermuteTypetransposeWeightUptransposeWeightDown,定义于 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是面向分组量化推理的完整融合形态。

八、使用建议与平台限制

综合源码约束,使用该算子需要注意以下几点:

  1. 平台前提:仅在Atlas 800I A2 推理产品上可用(ParamCheck),其他硬件平台会直接报错;
  2. 形状强约束:M ≤ 131072、分组数 ≤ 32,第一段矩阵乘固定为K=7168、N=4096,第二段固定为K=2048、N=7168,输入 x1 的 K 维固定 7168,输出固定[m, 7168]fp16。该算子是面向特定模型规格深度定制,接入前务必核对模型形状;
  3. 参数取值:保持默认值即可覆盖绝大多数场景——outputType=OUTPUT_FLOAT16groupListType=GROUP_LIST_CUMSUMtransposeWeightUp=falsetransposeWeightDown=trueweightUpPermuteType需与预处理好的权重重排方式一致(N256 或 N128),这会决定实际选择哪个 kernel 变体;
  4. 数据准备:x1、weight1、weight2 必须为 int8,weight 使用 FRACTAL_NZ 格式;scale1、scale2、perTokenScale1 为 fp32 ND 格式;groupList 为 int64、以 cumsum 前缀和形式编码分组边界;
  5. 运行时:走单一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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询