CANN ops-transformer IncreFlashAttention 增量推理融合算子设计解析

发布时间:2026/9/19 13:38:43
CANN ops-transformer IncreFlashAttention 增量推理融合算子设计解析 CANN ops-transformer IncreFlashAttention 增量推理融合算子设计解析【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer本篇技术指南围绕 CANN ops-transformer 仓库中的 IncreFlashAttention 算子设计介绍 展开系统讲解面向增量推理Incremental Inference场景的融合注意力算子 IFA 的实现原理、模板体系、计算切分与 Tiling 设计。读者读完本文后将理解 IFA 如何复用 FlashAttention/FlashSoftmax 算法支撑自回归逐 token 生成、掌握 CV、All-Vector、matmul 基础 API 等模板的适用场景以及 TilingKey 如何编码模板与特性组合并能在 incre_flash_attention 源码目录中定位到每一处设计对应的实现文件。增量推理与 IFA 算子的提出背景为不断优化提升增量推理性能CANN ops-transformer 提出了支持增量推理的 IncreFlashAttention简称 IFA融合算子需求。对于自回归Auto-regressive语言模型随着新词的生成推理输入长度不断增大相比全量推理增量推理主要有如下差异输入数据特点是 query 的 S 轴固定为 1Key 和 Value 是经过 kv cache 后将之前推理过的 state 信息叠加在一起每个 Batch 对应的 S 轴的实际长度可能不一样输入的数据是经过 padding 后的固定长度数据。kvCache 是大模型推理性能优化的常用技术采样时模型先以 prompt/context 并行处理随后逐一生成额外 token在采样过程中 Transformer 需要为序列中的每个 token 提取键值向量并存储在 kvCache 矩阵中供后续增量计算复用。增量推理的流程与正常全量推理并不完全等价但增量推理的精度并无明显劣化。全量场景可参考同仓库的 PromptFlashAttention。自注意力的标准计算公式为$$ Attention(Q,K,V)Softmax(\frac{QK^T}{\sqrt{d}})V $$其中 $Q$、$K$、$V$ 为输入经过空间变换得到的特征$\frac{1}{\sqrt{d}}$ 即缩放系数 scaleValue避免注意力分值过大。实现原理沿 FlashAttention 正向流程融合IFA 按照 FlashAttention 正向计算流程实现整体计算流程如下对应仓库图片 IncreFlashAttention.pngQK 矩阵乘与遮蔽query 与转置后的 key 做 matmul 计算后得到最初步的 attention_score然后与位置编码 pse 相加再乘以缩放系数 scale_value。此时的结果通过 atten_mask 进行 select 操作将 atten_mask 中为 true 的位置进行遮蔽true 位置在 select 后结果为负的极小值经过 softmax 计算之后变成 0从而达到遮蔽效果。FlashSoftmax 与输出刷新为了实现 FlashAttention 加速算法使用 FlashSoftmax 操作对 masked_attention_score 进行运算用以代替原公式中的 softmax 运算而后将结果与 value 做 matmul 运算。由于 FlashSoftmax 操作对 masked_attention_score 的 Skv输入 key、value 的 sequence length方向进行了切分故实现过程中存在一个刷新update流程具体如下每次 FlashSoftmax 计算只对切分后的一个 SkvSplit针对 Skv 轴切分之后的序列长度的简称进行操作并从第二次循环开始记录 exp其中 i 表示 Skv 切分后的循环变量针对 exp 的 i 从 1 开始exp 的计算公式为 $$ exp[i] e^{max_{i - 1} - max_{i}} $$从 i 1 开始需要增加 Mul 和 Add 操作即将上一次的 MM[PV] 的结果和当前 exp 相乘相乘的结果和本次 MM[PV] 的结果相加得到的结果保存到 GM 中。以此类推遍历 Skv 计算完成。由于 FlashSoftmax 计算中的除 sum 被后移到输出 attention_out 之前因此最后需要将 UB 中的 attention_out 按行除以 softmax_sum并将最终完整的结果保存到输出内存 attention_out(Final) 上。模板设计按输入特征拆分复用为了使不同的输入可以复用相同的 tiling 和流水IFA 采用模板的方式实现融合算子但不同输入全部使用同一套模板又无法达到性能最优和功能泛化因此需要根据输入 shape 的特征区分不同的模板实现。以下模板文件均可在 op_kernel 目录 中对应找到。模板类型与适用场景模板对应源文件定位与关键差异CV 模板incre_flash_attention_split_Bbn2s2_Us2.hIFA 基础模板支持绝大多数输入场景同时开启 VectorCore 和 CubeCorematmul 计算放在 CubeCore 执行且 matmul 调用 AscendC 提供的高阶 APIAll-Vector 模板incre_flash_attention_allvec_new.h对 CV 模板的补充主流程基本一致但 matmul 由 vector 实现降低 Cube 启动和 CV 通信开销对部分输入类型有更好的性能表现matmul 基础 API 模板incre_flash_attention_preload.h基于 CV 模板使用 AscendC matmul 基础 API 对 matmul 部分重写优化性能伪量化 MSD DD 模板incre_flash_attention_preload_dd.h基于 preload 模板开发用于伪量化 MTP 场景优化了 MSD 算法当前仅支持 FIA 算子调用IFA 算子不会调用MLA 全量化模板incre_flash_attention_preload_mla.h适用于 MLA 场景 query、key、value 为 INT8 且 query_rope、key_rope 为 BF16 的 attention 计算基于 preload 模板开发并将 matmul 相关计算抽取到 ifa_service_matmul_full_quant.h当前仅支持 FIA 算子调用各模板支持场景如下All-Vector 模板Atlas 推理系列产品 全部使用该模板。Atlas A2 训练系列产品/Atlas A2 推理系列产品 非 PA、非 GQA且 Q、KV、Output 类型全部为 FP16。matmul 基础 API 模板在 CV 模板基础上主要做了三处改动切换编程视角CV 模板使用基于 VEC 的编程视角1 个 VEC 需要处理 1 次 matmul 的全部计算结果本模板使用基于 CUBE 的编程视角1 次 matmul 的计算结果会被切分为 2 份2 个 VEC 分别处理 1 份。优化 CUBE 和 VEC 之间的核间流水CV 模板使用顺序流水本模板使用 N-Buffer 流水连续执行 N 次某个计算阶段之后再连续 N 次下一个计算阶段。优化 CUBE 核内流水将 CUBE 核内的 Buffer 资源在 FA 的两个 matmul 计算之间统一调度使得 CUBE 核内的搬运和计算流水更加紧凑从而提升性能。matmul 基础 API 模板的支持范围参考 Tiling 中的 EnableCubeViewMM 函数。伪量化 MSD DD 模板基于 incre_flash_attention_preload.h 开发当前仅支持 FIA 算子调用。MLA 全量化模板基于 incre_flash_attention_preload.h 开发matmul 相关计算抽取到 ifa_service_matmul_full_quant.h当前仅支持 FIA 算子调用。从源码结构看模板的分发由 incre_flash_attention_arch22.h 完成该文件按 Q/KV/Output 的原始数据类型组合如ORIG_DTYPE_QUERY DT_FLOAT16 ORIG_DTYPE_KEY DT_INT8分组通过TILING_KEY_IS(...)匹配 TilingKey再经INVOKE_IFA_GENERAL_OP_IMPL、INVOKE_IFA_ALL_VEC_OP_IMPL、INVOKE_IFA_NO_KFC_OP_IMPL等宏实例化对应模板类并调用Init/InitQuant/Process。其中 C1V1CV 配比 1:1场景通过KERNEL_TASK_TYPE(..., KERNEL_TYPE_MIX_AIC_1_1)指定任务类型。伪量化 MSD DD 与 MLA 全量化模板分别通过INVOKE_IFA_NO_KFC_DD_OP_IMPL与INVOKE_IFA_NO_KFC_MLA_OP_IMPL仅在FIA_ENABLE_MLA宏下编译调用。计算过程详解数据切分由于硬件 buffer 大小有限而计算数据量巨大无法一次计算完需要进行 tiling 切分shape 不同会导致算子的切分轴不同而切分轴会影响模板的功能及性能。简单的 element-wise 类算子往往将所有的轴 fuse 成一根轴切分逻辑简单、模板单一而融合算子融合了 element-wise、broadcast、reduce 及 matmul 等多类场景功能复杂为达到较高性能往往需要根据切分轴进行模板拆分。模板拆分时需考虑如下几个点a. 将核心的数量用满防止部分核闲置b. 每一个核心被分配的计算量相对均匀避免出现某些核计算的数据量过大、其余核空闲的情况c. AIC 和 AIV 之间处理的数据量要符合其对应的算力避免 AIC 或 AIV 出现长时间的空闲。IFA 算子包含 B、N2key 和 value 的 N、Gquery_N/kv_N、S1query 的 S、S2key 和 value 的 S共 5 个轴S1 轴固定为 1不参与切分。G 轴只在 Vector 计算时切块BN2S2 切分逻辑如下核间数据外切是为了最大限度的利用多个 Core 并行工作通常先按照 BN2 分核即将 BN2 个 SD 块分配到多个核上每个核计算一定数量的 SD 块当 BN2 小于阈值0.4 × 总核数时需再对 S2 轴进行外切SplitKV 份总块数为 BN2 × SplitKv每个核分配一定数量的子块当所有子块计算完成后再进行规约即 FlashDecode 流程。核内由于单 core 缓存有限需根据设定的缓存大小对 S2 轴或 KV 子块的 S2 轴进行切分此即 FlashAttention 过程。主流程IFA 单核计算主流程可用如下伪代码概括与原文档一致// 单核计算伪代码 void compute() { loops blocks_to_compute_of_this_core(); // 当前核需要计算几个数据块 for (i 0; i loops; i) { block get_curr_block(i); bidx, nidx, sidx dims_of_this_block(block); innerloops get_inner_loops_of_this_block_by_actual_seq_len(bidx, nidx, sidx); // 数据块实际内切份数 q_offset get_offset_of_query(bidx, nidx); softmax_sum {0}; softmax_exp {0}; softmax_max {min_float}; for (j 0; j innerloops; j) { // flash attention循环 kv_offset get_offset_of_kv_block(j); qk_res matmul(q q_offset, k kv_offset); qk_res elementwise(qk_res); // pse, atten-mask qk_res, softmax_max, softmax_sum, softmax_exp softmaxflash(qk_res, softmax_max, softmax_sum); res matmul(qk_res, v kv_offset); prev_res load_prev_res(); res prev_res * softmax_exp; // flash attention update store(res); if (j innerloops - 1) { res div(res, softmax_sum); output(res); } } } }结合源码实际内核入口在 incre_flash_attention.cpp模板类如IncreFlashAttentionAttenSplitBbn2s2Us2通过Init完成 GM 指针、tiling 与流水TPipe的绑定InitQuant传入反量化/伪量化相关缩放参数Process驱动上述循环。FlashDecode 规约当 S2 轴外切分配到不同核上完成 Attention 计算后需要对结果进行 Reduce 操作总共 BN2 个 SD 大块每个 core 对一个大块中的所有子块进行合并即 FlashDecode 流程。// FlashDecode规约单核流程 void combine() { SyncAll(); // 核间同步确保所有子块计算完成 splits get_real_splits_of_this_block_by_actual_seq_len(); lse load_lse_of_this_block(); scale[0:splits] exp(lse[i]) / Sum(exp(lse[i])); // i [0, splits) res {0}; split_res load_split_res(); for (j 0; j splits; j) { res split_res[j] * scale[j]; } output(res); }AntiQuant MSD 算法IFA AntiQuant 场景矩阵计算公式为$$ C A * (B offset)\times scale $$其中 A 矩阵为 FP16/BF16 类型B 矩阵为 int8 类型。经典的反量化方案需要对整个 B 矩阵进行反量化操作将矩阵 B 搬入 Vector 处理而 B 矩阵数据量较大严重影响计算性能。IFA 场景下 A 矩阵较小可以通过变换 A 矩阵来适配 B 矩阵基本流程矩阵 A 进入 Vector 展开成多行每行 An 均用 int8 格式存储将这些 An 打包成新的矩阵 AA计算 CC AA * B按 int8 × int8 int32 计算对 MatMul 结果 CC 进行 Reduce 操作得到 C。PageAttentionKV block 内存不连续时MatMul 针对这种场景提供了回调函数进行 B 矩阵的拷贝GM→L1。IFA 中实现相应的拷贝函数回调函数在 Cube 中执行参数通过 GM 传递Vector 设置相应的参数后到 GM确保 DCCI后再通知 MatMul 工作。GQA 支持G queryHeadNum / KvHeadNum。Vector 上 G 轴的切分由当前操作所涉及的输入输出 UB 大小决定当 G 过大UB 缓存不足以一次加载全部数据进行计算时需要在 G 轴上进行切分// GQA vector切G伪码 void process() { g target_ub_size() / column_size; if (g G) { g G; } process_sub_block(g, column); // sub_block: g * column }Tiling 设计分核设计Tiling 操作的目的是找到一种更高效的 NPU 执行方式。原始数据量一般非常大无法通过一次指令调用完成所有计算因此需要将数据量分到多个核上并行计算且每个核上也需要考虑如何循环计算性能最优不同输入可能有不同的最优执行方式所以需要通过 Tiling 策略决定如何将数据分配到各个核上进行计算。如前所述总块数为 BN2 或 BN2 × SplitKv输入核数 块数 块负载通常为每个分块的 S 轴实际长度处理根据负载值对连续的块进行组合重排达到核间负载差值最小输出blockid 数组每个元素对应一个核的起始 blockid最后附加一个元素等于总块数前后元素差值为该核处理的块数。Tiling 的宿主侧实现位于 incre_flash_attention_tiling.cpp、incre_flash_attention_tiling_impl.h 及 incre_flash_attention_tiling.h并针对 arch38 有独立的 incre_flash_attention_tiling_arch38.cpp。TilingKey 规划TilingKey 为 uint64 类型通常每个模板参数对应 TilingKey 中的一个十进制位部分 BOOL 类型的模板参数采用组合方式在一个十进制位中表示。具体实现参考 Tiling 中的 GenTilingKey 函数声明见 incre_flash_attention_tiling_impl.h。constexpr uint64_t RecursiveSum() { return 0; } template typename T, typename... Args constexpr uint64_t RecursiveSum(T templateId, Args... templateIds) { return static_castuint64_t(templateId) 10U * RecursiveSum(templateIds...); } constexpr uint64_t IFA_TILINGKEYOFFSET uint64_t(10000000000000000UL); // 10^16 constexpr uint64_t IFA_PERF_MODE_TILINGKEYOFFSET uint64_t(1000000000000000UL); // 10^15 template typename... Args constexpr uint64_t IFA_GET_TILINGKEY(Args... templateIds) { return RecursiveSum(templateIds...); } GenTilingKey() { ... uint64_t baseOffset modeVal * IFA_TILINGKEYOFFSET (static_castuint64_t(perfMode_)) * IFA_PERF_MODE_TILINGKEYOFFSET; if (antiquantMode_ PER_TOKEN_MODE || antiquantMode_ PER_CHANNEL_MODE){ context_-tilingKey baseOffset IFA_GET_TILINGKEY(layoutVal, inputQVal, inputKvVal, outputVal, originVal, (paVal splitKvVal antiquantModeVal), 0, kvLayoutInfo.kvLayoutVal, kvLayoutInfo.amlaMode, balanceMode); } else { context_-tilingKey baseOffset IFA_GET_TILINGKEY(layoutVal, inputQVal, inputKvVal, outputVal, originVal, (paVal splitKvVal), antiquantMode_, kvLayoutInfo.kvLayoutVal, kvLayoutInfo.amlaMode, balanceMode); } ... }TilingKey 各十进制位字段说明如下十进制位变量说明0layoutValQ 的 Shape 格式0: BNSD1: BSH/BSND2: TND1inputQValquery 数据类型0: FP162: BF163: INT82inputKvValKV 数据类型0: FP162: BF163: INT84: INT43outputValoutput 数据类型0: FP162: BF163: INT84originVal同 inputQval5 [bit0]splitKvVal开启 FlashDecode 标志1: enable0: disable5 [bit1]paVal开启 PageAttention 标志1: enable0: disable5 [bit2]antiquantModeVal开启 PerToken 伪量化标记1: enable0: disable6antiquantMode_量化模式0: 无效值2: K-perChannel-V-perToken7kvLayoutInfo.kvLayoutValKV 的 shape 格式仅伪量化 MSD DD 模板和 MLA 全量化模板该字段有效其余模板该字段值为 00: BNSD1: BSH/BSND2: NZ8kvLayoutInfo.amlaMode该字段废弃取值只能为 09balanceMode开启新的负载均衡算法的标志1: enable0: disable仅 MLA 全量化模板可开启10...14预留字段值为 015perfMode_模板编号0: C1_V2CV 配比 1:21: 全V2: C1_V1CV 配比 1:13: matmul 基础 API 模板5: MLA 全量化模板6: 伪量化 MSD DD 模板16modeVal1: IFA TilingKey Base2: IFA 启用 SysPrefix 功能在仓库中这些 TilingKey 的组合被固化为宏常量集中定义在 incre_flash_attention_tilingkey.h。例如QF16_KVF16_OUTF16_ANTIPERCHANNEL_C1V2_TILING值10000000000000000表示 Q/KV/Out 均为 FP16、per-channel 反量化、perfMode 为 C1V2 的场景QF16_KVF16_OUTF16_ANTIPERCHANNEL_PAGEDCACHE_C1V2_TILING值10000000000200000则叠加了位 5 中的 PAPageAttention标志以150000...开头的宏对应 MLA 模板perfMode5以160000...开头的宏对应伪量化 MSD DD 模板perfMode6以200000...开头的宏对应启用 SysPrefix 功能modeVal2。宏值中的1/2/3等数字即对应 layoutVal、inputKvVal 等字段的取值内核侧正是依据这些宏与TILING_KEY_IS完成模板分发。接口、调用与约束IFA 通过 aclnn 接口对外提供仓库 README 中给出了产品支持情况 Atlas A3 训练/推理系列产品 、 Atlas A2 训练/推理系列产品 、 Atlas 推理系列产品 支持 Ascend 950PR/Ascend 950DT 、 Atlas 200I/500 A2 推理产品 、 Atlas 训练系列产品 不支持。核心输入输出参数包括参数名输入/输出/属性描述数据类型数据格式query输入公式中的输入 QFLOAT、FLOAT16NDkey输入公式中的输入 KFLOAT、INT8、FLOAT16NDvalue输入公式中的输入 VFLOAT、INT8、FLOAT16NDscaleValue属性公式中 d 开根号的倒数DOUBLE-attentionOut输出公式中的输出FLOAT、INT8、FLOAT16ND同时 IFA 支持量化、位置编码、PageAttention、kvCache 反量化和 KV 左 Padding 等特性主要约束包括query 与 attentionOut 的 shape 需完全一致key、value 对应 tensor 的 shape 需完全一致。query 的 N 与 numHeads 相等key、value 的 N 与 numKeyValueHeads 相等且 numHeads 是 numKeyValueHeads 的倍数关系。Atlas A2 训练/推理系列产品、Ascend 950PR/Ascend 950DT支持 B 轴 ≤ 65536N 轴 ≤ 256D 轴 ≤ 512。Atlas 推理系列产品支持 B 轴 ≤ 256N 轴 ≤ 256D 轴 ≤ 512key、value 的 S 轴 ≤ 65536数据类型仅支持 FLOAT16BNSD 排布下 numHeads 与 numKeyValueHeads 比值不大于 8。PageAttention 场景blockTable 存在且有效blockSize 需要传入非 0 值且不超过 512key、value 为 FLOAT16/BFLOAT16 时需 16 对齐INT8 时需 32 对齐推荐 128kvCache 排布支持blocknum, blocksize, H与blocknum, KV_N, blocksize, D两种格式。kv 左 padding 场景搬运起点为 Smax - kvPaddingSize - actualSeqLengths终点为 Smax - kvPaddingSizekvPaddingSize 小于 0 时置为 0。调用示例见 examples/test_aclnn_incre_flash_attention.cpp它展示了通过aclnnIncreFlashAttentionV4接口完成 device 初始化、aclCreateTensor构造输入示例中 query shape 为{1, 2, 1, 16}、key/value shape 为{1, 2, 2, 16}即 S11、S22 的增量场景、aclrtMemcpy搬入数据并执行推理的完整流程算子更详细的接口与版本差异可参考 aclnnIncreFlashAttention.md 及 V2/V3/V4 各版本文档。小结IFA 是 ops-transformer 中面向增量推理的融合注意力算子通过 FlashAttention/FlashSoftmax 算法在 NPU 上完成 QK 矩阵乘、pse 位置编码、atten-mask 遮蔽、softmax 归一化与 PV 矩阵乘的融合计算并以 CV、All-Vector、matmul 基础 API、伪量化 MSD DD、MLA 全量化五类模板覆盖不同硬件与输入特征的性能诉求。Tiling 层通过 BN2S2 分核、FlashDecode 规约以及按十进制位编码的 TilingKey 机制将数据切分、模板选择与量化/PA/FlashDecode 等特性统一纳入运行时决策相关实现均可在 op_kernel 与 op_host 目录中对照查阅。【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考