
ascend-transformer-boost 的 MlaPreprocessOperationMLA 预处理四阶段融合算子深度解析【免费下载链接】ascend-transformer-boost本项目是CANN提供的是一款高效、可靠的Transformer加速库基于华为Ascend AI处理器提供Transformer定制化场景的高性能融合算子。项目地址: https://gitcode.com/cann/ascend-transformer-boost导读MlaPreprocessOperation是 CANN ascend-transformer-boost 加速库中面向 Multi-head Latent AttentionMLA推理场景的高性能融合算子它将 RMS Norm 量化、QKV 投影Matmul、旋转位置编码RoPE与 KV Cache 写入四个阶段融合为单个算子从而减少 kernel 启动次数与中间张量搬运开销。本文以.agent/knowledge/ops/attention/mla_preprocess/index.md知识条目为骨架结合算子源码、参数头文件与示例工程系统讲解其参数约束、计算流水线、Runner 决策树、输入输出规格与运行方式帮助读者在 Atlas 800I A2/A3 推理产品上正确配置并使用该算子。1. 算子定位为 MLA 推理服务的预处理融合MLAMulti-head Latent Attention是 DeepSeek 等大模型采用的低秩注意力机制其核心思想是对 QKV 进行低秩压缩latent 表示在推理时通常需要对输入的 hidden state 依次执行 rmsNormQuant、QKV 投影、RoPE 旋转编码以及 KV Cache 写入。在未融合的实现中这些阶段各自对应独立 kernel中间结果需要反复读写显存造成额外开销。MlaPreprocessOperation将这些阶段融合为一个多阶段融合算子pipeline_type: multi_stage_fusion一次执行即可产出下游 MLA 计算所需的全部中间张量。该算子的知识条目位于 .agent/knowledge/ops/attention/mla_preprocess/index.md源代码位于 src/ops/ops_infer/mla_preprocess/底层 kernel 位于 src/kernels/mixkernels/mla_preprocess。注意参数头文件明确指出该算子仅 Atlas 800I A2 推理产品支持见 include/atb/infer_op_params.h 中MlaPreprocessParam的注释算子创建入口CreateOperation中也会校验平台非 910B 产品会直接报错拒绝创建。2. 参数结构详解MlaPreprocessParam算子的全部配置通过atb::infer::MlaPreprocessParam传入完整定义位于 include/atb/infer_op_params.h。知识条目中的核心字段与约束如下struct MlaPreprocessParam { uint32_t wdqDim 0; // matmul 后拆分 dim uint32_t qRopeDim 0; // Q 进入 RoPE 的 dim uint32_t kRopeDim 0; // K 进入 RoPE 的 dim float epsilon 1e-5; // Norm epsilon int32_t qRotaryCoeff 2; // Q 旋转系数2/4/headDim int32_t kRotaryCoeff 2; // K 旋转系数2/4/headDim bool transposeWdq true; // WDQ 是否转置 bool transposeWuq true; // WUQ 是否转置 bool transposeWuk true; // WUK 是否转置 enum CacheMode { KVCACHE0, KROPE_CTKV, INT8_NZCACHE, NZCACHE }; CacheMode cacheMode KVCACHE; };各字段含义与约束汇总参数约束 / 说明wdqDimmatmul 后 Q 的拆分维度即投影结果中被切分为 rope 部分与 nope 部分的分界qRopeDim/kRopeDimQ / K 送入 RoPE 旋转编码的维度大小epsilonRMS Norm 分母上的小量防止除 0默认1e-5qRotaryCoeff/kRotaryCoeff旋转系数对半旋转取 2支持配置 2、4 或 headDimtransposeWdq/transposeWuq/transposeWuk三组投影权重是否转置默认均为 truecacheModeCache 格式见下2.1 CacheMode 与 QuantMode 枚举除了知识条目列出的CacheMode头文件还定义了QuantMode枚举用于控制 RMS Norm 量化的类型枚举取值说明CacheMode::KVCACHE0标准 KV Cache 输出2 路输出CacheMode::KROPE_CTKV1Split 变体K/Rope 与 CTKV 分别输出4 路输出CacheMode::INT8_NZCACHE2int8 量化 NZ 格式 CacheCacheMode::NZCACHE3NZ 格式 CacheQuantMode::PER_TENSOR_QUANT_ASYMM0Per-tensor 静态非对称量化默认QuantMode::PER_TOKEN_QUANT_SYMM1Per-token 静态对称量化QuantMode::PER_TOKEN_QUANT_ASYMM2Per-token 静态非对称量化不支持QuantMode::UNQUANT3不量化不支持cacheMode的值直接影响输出路数见第 4 节与 runner 选择非法取值会在CreateOperation阶段被拒绝cacheMode NZCACHE报错。quantMode同样有强约束创建算子时若传入PER_TOKEN_QUANT_ASYMM或UNQUANT会直接返回ERROR_INVALID_PARAM也就是说当前实现仅支持PER_TENSOR_QUANT_ASYMM与PER_TOKEN_QUANT_SYMM两种量化模式这一校验逻辑可在 src/ops/ops_infer/mla_preprocess/mla_preprocess_operation.cpp 中确认。2.2 参数的实际生效范围需要特别说明的是头文件注释将wdqDim、qRopeDim、kRopeDim、epsilon、旋转系数与转置开关标注为目前均为未使用的预留参数需支持泛化后启用。从知识条目与 demo 场景看这些字段主要服务于 DeepSeek 场景的 Split/泛化路径例如 example/op_demo/mla_preprocess/README.md 中 mlapo_ds_demo 显式配置了wdqDim1536、qRopeDim64、kRopeDim64等而针对固定 7168 hidden size 的标准路径则依赖源码中固化的一组内部维度常量。3. 计算流水线四阶段融合知识条目给出的融合流水线如下pipeline_type: multi_stage_fusion stages: - stage: QKV 投影 note: x WDQ → Q, x WUK → K, x WUV → V - stage: RoPE note: Q[..., :qRopeDim] K[..., :kRopeDim] 旋转位置编码 - stage: KV Cache note: 根据 cacheMode 将 K,V 写入指定 Cache 格式 - stage: RMS Norm note: 可选 Normepsilon 控制按知识条目及示例工程mlapo_demo.cpp 的注释实际执行顺序可进一步细化为rmsNormQuant对输入input使用gamma0/beta0做 RMS Norm并按quantScale0/quantOffset0做量化int8matmul_0WDQKV 投影量化后的输入与wdqkv权重做矩阵乘输出经deScale0/bias0反量化并加偏置得到中间 2112 维表示rmsNormQuant对 2112 维中间结果再次做 RMS Norm 量化使用gamma1/beta1与quantScale1/quantOffset1matmul_1WUQ 投影与wuq权重相乘输出经deScale1/bias1还原得到 Q1536 维其中 rope 部分 64 维、nope 部分其余维度RMS Norm对 rope 部分使用gamma2做 NormRoPE对 Q 的[..., :qRopeDim]与 K 的[..., :kRopeDim]施加旋转位置编码cos/sin由外部传入reshapeAndCache根据slotmapping将 K/V以及量化场景下的 scale写入 KV Cache当cacheMode为 Split/INT8/NZ 变体时还会将 rope 后的 K 与 CTKV 分别写入kvCacheRope并按 NZ 格式输出。其中可选 Norm与是否走量化由源码中的doRmsNorm_标志动态决定当wdqkv权重 dtype 与输入 dtype 一致且均为 float16/bf16 时跳过输入侧 rmsNormQuant见CheckAclnnKernel逻辑mla_preprocess_operation.cpp。4. 输入输出规格24 输入、2/4 路输出该算子的接口规模较大固定24 个输入张量IN_TENSOR_NUM 24输出路数取决于cacheMode——标准KVCACHE模式下为2 路输出OUT_TENSOR_NUM 2其余 Split/INT8/NZ 变体为4 路输出OUT_TENSOR_NUM_SPLIT 4该规则实现在 mla_preprocess_operation.cpp。4.1 24 路输入索引表按源码DimCheck中的 shape 校验表mla_preprocess_operation.cpp各输入张量的含义与形状约束如下索引张量期望形状说明0input[tokenNum, hiddenSize]输入 hidden state1gamma0[hiddenSize]第 1 次 RMS Norm 的 gamma2beta0[hiddenSize]第 1 次 RMS Norm 的 beta3quantScale0[1]第 1 次量化 scale4quantOffset0[1]第 1 次量化 offset5wdqkv权重WDQKV 投影权重NZ 格式6deScale0[2112]第 1 次 matmul 反量化 scale7bias0[2112]第 1 次 matmul bias8gamma1[1536]第 2 次 RMS Norm 的 gamma9beta1[1536]第 2 次 RMS Norm 的 beta10quantScale1[1]第 2 次量化 scale11quantOffset1[1]第 2 次量化 offset12wuq权重WUQ 投影权重NZ 格式13deScale1[headNum * 192]第 2 次 matmul 反量化 scale14bias1[headNum * 192]第 2 次 matmul bias15gamma2[512]RoPE 前 Norm 的 gamma16cos[tokenNum, 64]RoPE cos 表17sin[tokenNum, 64]RoPE sin 表18wuk权重WUK 投影权重19ctkvkvCacheCache 张量CTKV / KV CacheNZ 格式20kRopekvCacheRopeCache 张量K 的 RoPE 部分 Cache21slotmapping[tokenNum]token 到 cache 槽位的映射22ctkvScale[1]CTKV 量化 scale23qNopeScale[headNum]Q nope 部分量化 scale4.2 输出规格InferShapeImplmla_preprocess_operation.cpp根据cacheMode推导输出形状KVCACHE标准out0 [tokenNum, headNum, 576]Q 输出out1 kvCacheIn 同规格的 Cache 输出KROPE_CTKV / INT8_NZCACHE / NZCACHESplit 变体4 路输出——out0 [tokenNum, headNum, 512]Q 的 nope 部分INT8_NZCACHE 模式下 dtype 为ACL_INT8out1 [blockNum, blockSize, 1, 512]CTKV Cache 输出out2 [tokenNum, headNum, 64]Q 的 rope 部分out3 [blockNum, blockSize, 1, 64]K rope Cache 输出。其中tokenNum取自输入第 0 维headNum取自wuk索引 18的第 0 维。从代码可以推断576 512(nope) 64(rope) 的拆分正是 MLA 低秩注意力中 Q 分离式表示的典型配置。4.3 关键校验约束tokenNum必须 0 且 ≤ 1024MAX_TOKEN_NUMblockSize范围[1, 128] ∪ {256}INT8_NZCACHE / NZCACHE 模式强制为 128hiddenSize非泛化非 ACLNN路径要求输入 hiddenSize 严格等于7168启用 ACLNN kernel 后支持泛化范围[2048, 8192]MIN_HIDDEN_SIZE/MAX_HIDDEN_SIZEgamma0 / beta0必须为 1 维且长度等于 hiddenSize输出维度、Cache 维度如[blockNum, blockSize, 1, 512]均有OutTensorCheck/OutTensorCheckSplit严格校验。5. Runner 执行路径与平台适配知识条目给出了算子内部的执行路径决策树源码CreateRunner与之完全对应MlaPreprocessOperation::CreateRunner() │ ├── [ACLNN 路径] → MlaPreprocessAclnnRunner └── [Ops 路径] ├── MlaPreprocessOpsRunner标准 └── MlaPreprocessOpsRunnerSplitSplit 变体平台支持情况知识条目 源码佐证平台Runner说明910BAtlas 800I A2/A3OpsRunner标准 / Split原生 ATB kernel 路径950ACLNN Runner OpsRunner支持 ACLNN 加速路径CreateRunner的完整决策逻辑位于 mla_preprocess_operation.cpp构造MlaPreprocessOperation时先尝试MlaPreprocessAclnnRunner::LoadMethod()加载 ACLNN 函数符号加载失败则isAclnnFuncLoaded false此时泛化 hiddenSize 与跳过 rmsNormQuant能力不可用并打印 WARN 日志CheckAclnnKernel判断是否需要走 ACLNN 路径仅当 (a) hiddenSize ≠ 7168泛化或 (b) 需要跳过输入 rmsNormQuant权重为 fp16/bf16 且与输入同 dtype时才启用 ACLNN kernel否则使用 ATB 原生 kernel启用 ACLNN 时还强制要求quantMode PER_TENSOR_QUANT_ASYMMCreateRunner中useAclnnKernel_为真则创建MlaPreprocessAclnnRunner否则cacheMode KVCACHE创建MlaPreprocessOpsRunner其余 cacheMode 一律创建MlaPreprocessOpsRunnerSplit。ACLNN Runner 的 Workspace 计算与 ACLNN API 调用封装在 mla_preprocess_aclnn_runner.cppATB 原生路径见 mla_preprocess_ops_runner.cpp 与 mla_preprocess_ops_runner_split.cppACL 侧辅助逻辑位于 atb_acl_mla_preprocess.cpp。6. Kernel 依赖mix kernel 实现该算子依赖的底层 kernel 为mla_preprocess位于 src/kernels/mixkernels/mla_preprocess目录结构包含文件角色mla_preprocess_kernel.cppKernel 入口host 侧调度mla_preprocess_operation.cppKernel 算子实现op_kernel/mla_preprocess_mix.cce主 kernel.cce 设备侧实现op_kernel/mla_preprocess_mix_bf16.ccebf16 变体 kerneltiling/mla_preprocess_tiling.cpp / .hTiling 策略与参数计算tiling/mla_preprocess_tilingdata.hTiling 数据结构定义从文件组成可以推断融合 kernel 同时提供 fp16 与 bf16 两套设备侧实现并通过独立 tiling 逻辑在 host 侧完成分块切分与参数下发这与算子侧多阶段融合、单次下发的设计目标一致——避免各阶段独立 kernel 带来的多次启动开销。7. 示例工程与运行方式仓库在 example/op_demo/mla_preprocess/ 提供了两个 C demo完整演示了 24 路输入的准备、VariantPack组装、Setup计算 workspace、循环Execute的调用范式mlapo_demo.cppint8 量化叠加 rope 切分场景cacheMode INT8_NZCACHEctkv 与 qNope 经 per-head 静态对称量化输出 int8krope 与 ctkv 以 NZ 格式输出mlapo_ds_demo.cppDeepSeek 场景使用 rope 拆分 kvcache 与 queryper-tensor 静态非对称量化并显式配置了wdqDim1536、qRopeDim64、kRopeDim64等参数。7.1 编译与运行按 example/op_demo/mla_preprocess/README.md 的说明加载 CANN 与 nnal 环境源码编译场景则 source 构建产物下的 set_env.shsource /usr/local/Ascend/ascend-toolkit/set_env.sh source [加速库源码路径]/output/atb/set_env.sh # 例如 ./ascend-transformer-boost/output/atb/set_env.sh编译运行 demobash build.sh注意_GLIBCXX_USE_CXX11_ABI需与所链接加速库的 ABI 设置一致cxx_abi0 时为-D_GLIBCXX_USE_CXX11_ABI0运行参数dtypefloat16/bf16、tokenNum、headNum默认值分别为float16、4、128mlapo_demo与float16、32、128mlapo_ds_demo。7.2 核心调用范式摘自 mlapo_demo.cppatb::infer::MlaPreprocessParam param; param.cacheMode atb::infer::MlaPreprocessParam::CacheMode::INT8_NZCACHE; atb::Operation *mlaPreprocessOp nullptr; atb::CreateOperation(param, mlaPreprocessOp); // 创建算子 atb::VariantPack variantPack; // PrepareInTensor1 / PrepareInTensor2 组装 24 路输入... // 4 路输出qOut0[int8, tokenNum,headNum,512]、kvCacheOut0、 // qOut1[dtype, tokenNum,headNum,64]、kvCacheOut1 variantPack.outTensors {qOut0, kvCacheOut0, qOut1, kvCacheOut1}; uint64_t workspaceSize 0; mlaPreprocessOp-Setup(variantPack, workspaceSize, context); // 推导 workspace uint8_t *workspacePtr ...; // 按需 aclrtMalloc for (size_t i 0; i 10; i) { mlaPreprocessOp-Execute(variantPack, workspacePtr, workspaceSize, context); aclrtSynchronizeStream(stream); // 流同步 } atb::DestroyOperation(mlaPreprocessOp);一个值得注意的工程细节mlapo_demo.cpp中输出kvCacheOut0/kvCacheOut1直接复用输入kvCache/kvCacheRope的 deviceData原地更新 Cache释放内存时通过比对deviceData指针避免重复释放这体现了该算子 Cache 原地写入的使用习惯。数据规格方面示例工程 README 给出了每一路输入输出的 dtype、formatnd / nz与完整维度信息例如input为[tokenNum, 7168]的 float16/bf16、wdqkv为 NZ 格式 int8、输出qOut0为[tokenNum, headNum, 512]的 int8 等是自行构造测试数据时最直接的参考。同时示例数据生成也可参考测试用例目录tests/apitest/opstest/python/operations/mla_preprocess/。8. 已知问题与注意事项知识条目记录了当前实现的一个关键注意点#问题状态1融合 4 阶段投影 RoPE Cache NormGolden 需匹配 Pipeline注意即由于算子将多个数学阶段融合在单次执行内进行精度验证时参考实现Golden必须按同样的阶段顺序与量化/反量化规则逐一复现而不能只做端到端的结果对比同时还需注意以下实现约束算子仅在 Atlas 800I A2/A3910B推理产品上可用创建时会做平台校验标准非泛化路径要求 hiddenSize 固定为 7168超出该范围必须依赖 ACLNN kernel 且要求环境已正确加载 ACLNN 函数quantMode不支持PER_TOKEN_QUANT_ASYMM与UNQUANTINT8_NZCACHE / NZCACHE 模式下 blockSize 必须为 128Cache 原地写入语义输出 Cache 张量可与输入 Cache 张量共享存储释放资源时需避免重复 free。9. 相关算子脉络从知识条目与仓库整体结构看mla_preprocess处于 MLA 推理链路的上游算子与 mla_preprocess 的关系multi_latent_attention下游算子直接消费 preprocess 的 Q / CTKV / Cache 输出完成 MLA 注意力计算实现见 src/ops/ops_infer/multi_latent_attention/ring_mla分布式 MLA 场景Ring 通信 MLA同样依赖 preprocess 产出结果实现见 src/ops/ops_infer/ring_mla/rope融合进本算子的 RoPE 阶段独立的 RoPE 算子实现见 src/ops/ops_infer/rope/理解mla_preprocess的输出格式512 64 的 Q 拆分、NZ Cache 布局是正确使用上述下游 MLA 算子的前提若需按知识条目推荐的阅读顺序深入源码建议依次阅读mla_preprocess_operation.h接口与 InferShape 签名、mla_preprocess_operation.cppCreateRunner 决策树、mla_preprocess_aclnn_runner.cppACLNN 路径与 Workspace 计算最后对照mla_preprocess_ops_runner.cpp与mla_preprocess_ops_runner_split.cpp理解原生 ATB 调用链与平台适配细节。【免费下载链接】ascend-transformer-boost本项目是CANN提供的是一款高效、可靠的Transformer加速库基于华为Ascend AI处理器提供Transformer定制化场景的高性能融合算子。项目地址: https://gitcode.com/cann/ascend-transformer-boost创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考