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是 ATBAscend Transformer Boost算子体系中的一个composite复合类型算子归类于activation类别规模等级为tier: XL。它并非一个单一数学运算而是把 MoEMixture of Experts类模型中一段典型的量化推理计算链完整融合进单个算子从而减少中间张量的搬运与 kernel 启动开销。从知识条目与源码可以确认如下事实类别activation类型 compositesrc/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)逐段展开为阶段 1GMM1 Dequant1。第一个分组矩阵乘接收 int8 激活 x1 与 int8 权重 weight1累加得到 fp32 中间结果后使用 scale1per-channel/权重侧尺度与 perTokenScale1per-token 尺度执行反量化输出 fp16阶段 2SwiGLU。对 fp16 结果施加 SwiGLU 门控激活即silu(x) * y形式的 gated linear unit中间仍为 fp16阶段 3Quant。将 fp16 激活再次量化为 int8作为第二个矩阵乘的输入同时生成对应的量化尺度阶段 4GMM2 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、weight1int8kernel 校验x1.dtype TENSOR_DTYPE_INT8、weight1.dtype TENSOR_DTYPE_INT8GMM1 Dequant 之后SwiGLU 输入/输出fp16知识条目定义Quant 之后GMM2 输入int8kernel 校验weight2.dtype TENSOR_DTYPE_INT8最终输出fp16InferShapeoutDesc.dtype ACL_FLOAT16这条链路上存在2 个 dtype 断点断点 1int8 → fp16位于 GMM1 累加与反量化之后。int8 激活与 int8 权重在矩阵乘单元中累加为 fp32再经 scale1 × perTokenScale1 反量化后落到 fp16进入 SwiGLU断点 2fp16 → int8位于 Quant 阶段。SwiGLU 的 fp16 输出被再量化为 int8作为 GMM2 的激活输入。正是这两个断点让算子可以全程使用 int8 矩阵乘的高吞吐计算单元GMM1、GMM2 均为 int8 输入同时在精度敏感的反量化与激活环节保持 fp16兼顾性能与精度。两个尺度张量scale1 与 scale2与 per-token 尺度perTokenScale1均为 fp32kernel 校验保证了反量化计算的精度。四、输入输出契约7 输入 1 输出该算子固定为7 个输入、1 个输出源码常量INPUT_NUM 7、OUTPUT_NUM 1见 operation.cpp。输入按顺序为序号名称语义维度dtypeformat0x1第一个矩阵乘的激活2 维[m, K1]int8ND1weight1第一个矩阵乘的权重3 维[groupCount, K1, N1]int8FRACTAL_NZ2scale1第一个矩阵乘的反量化尺度2 维[groupCount, N1]fp32ND3perTokenScale1per-token 反量化尺度1 维[m]fp32ND4groupList分组边界列表cumsum 形式1 维[groupCount]int64ND5weight2第二个矩阵乘的权重3 维[groupCount, N2, K2]转置语义int8FRACTAL_NZ6scale2第二个矩阵乘的反量化尺度2 维[groupCount, N2]fp32ND输出为 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 trueN 在 dim1、K 在 dim2这与参数transposeWeightUp false、transposeWeightDown true的约束一一对应见下节。五、参数定义与强制约束参数结构体GmmDeqSwigluQuantGmmDeqParam定义于 include/atb/infer_op_params.h包含三个枚举与两个布尔开关枚举定义与默认值参数可选值默认值说明outputTypeOUTPUT_FLOAT16 0、OUTPUT_BFLOAT16、OUTPUT_INVALIDOUTPUT_FLOAT16输出数据类型groupListTypeGROUP_LIST_CUMSUM 0、GROUP_LIST_SINGLE、GROUP_LIST_INVALIDGROUP_LIST_CUMSUMgroupList 的编码形式weightUpPermuteTypePERMUTE_N256 0、PERMUTE_N128、PERMUTE_INVALIDPERMUTE_N256weight1 与 scale1 的重排方式transposeWeightUpboolfalseweight1 是否转置transposeWeightDownbooltrueweight2 是否转置当前版本下的强制约束ParamCheck 实现平台限定仅支持Atlas 800I A2 推理产品Config::Is910B()校验其他平台直接报错拒绝运行outputType仅支持OUTPUT_FLOAT16groupListType仅支持GROUP_LIST_CUMSUMgroupList 以 cumsum 前缀和形式给出各分组边界weightUpPermuteType不能为PERMUTE_INVALID且必须与 kernel 变体匹配N256/N128transposeWeightUp仅支持falsetransposeWeightDown仅支持true。这些约束在 operation 层与 kernel 层被重复校验kernel 侧 CheckGmmDeqSwigluQuantGmmDeq保证两层校验语义一致。形状约束CheckInTensorsShape同样严格约束项取值Mx1 的 batch/token 维≤ 131072MAX_M分组数groupCount≤ 32MAX_GROUP_COUNTx1 的 K 维固定 7168SUPPORTED_K1weight1 的 K/N固定 7168 / 4096SUPPORTED_K1/SUPPORTED_N1weight2 的 N/K固定 7168 / 2048SUPPORTED_N2/SUPPORTED_K2输出 N 维固定 7168SUPPORTED_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 类型并MallocRunnerGmmDeqSwigluQuantGmmDeqOpsRunner创建实例失败时降级为直接make_shared创建GmmDeqSwigluQuantGmmDeqOpsRunner 继承自OpsRunner通过REG_RUNNER_TYPE注册其核心工作是重写SetupKernelGraphSetupKernelGraph 构建一个仅含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.cpp、gmm_deq_swiglu_quant_gmm_deq_n128.cpp两个 kernel 变体一一对应。Tiling 则由 gmm_deq_swiglu_quant_gmm_deq_tiling.cpp 中的GmmDeqSwigluQuantGmmDeqTiling完成声明见 tiling.h。参数更新方面SetParamops_runner.cpp使用operator比较新旧参数逐字段比较outputType、groupListType、weightUpPermuteType、transposeWeightUp、transposeWeightDown定义于 ops_runner.h仅在参数真正变化时才置位isParamUpdated_触发重建避免无谓的 kernel 重编译。七、算子族谱与相邻算子的关系知识条目给出了该算子在 ATB 算子族中的两个相关算子mm_deq_swiglu_quant_mm_deqMM 变体位于 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第一段矩阵乘固定为K7168、N4096第二段固定为K2048、N7168输入 x1 的 K 维固定 7168输出固定[m, 7168]fp16。该算子是面向特定模型规格深度定制接入前务必核对模型形状参数取值保持默认值即可覆盖绝大多数场景——outputTypeOUTPUT_FLOAT16、groupListTypeGROUP_LIST_CUMSUM、transposeWeightUpfalse、transposeWeightDowntrueweightUpPermuteType需与预处理好的权重重排方式一致N256 或 N128这会决定实际选择哪个 kernel 变体数据准备x1、weight1、weight2 必须为 int8weight 使用 FRACTAL_NZ 格式scale1、scale2、perTokenScale1 为 fp32 ND 格式groupList 为 int64、以 cumsum 前缀和形式编码分组边界运行时走单一GmmDeqSwigluQuantGmmDeqOpsRunnerOpsRunner 路径不经过 ACLNN参数未变化时不会触发 kernel 重建可放心复用 runner 实例。参考与进一步阅读算子知识条目.agent/knowledge/ops/activation/gmm_deq_swiglu_quant_gmm_deq/index.mdATB Operation 层src/ops/ops_infer/gmm_deq_swiglu_quant_gmm_deq/Kernel 层op_kernel tilingsrc/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/、单 SwiGLUQuant src/ops/ops_infer/swiglu_quant/【免费下载链接】ascend-transformer-boost本项目是CANN提供的是一款高效、可靠的Transformer加速库基于华为Ascend AI处理器提供Transformer定制化场景的高性能融合算子。项目地址: https://gitcode.com/cann/ascend-transformer-boost创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考