CANN ops-transformer 增量FlashAttention算子 aclnnIncreFlashAttention 使用指南与实现原理解析

CANN ops-transformer 增量FlashAttention算子 aclnnIncreFlashAttention 使用指南与实现原理解析 CANN ops-transformer 增量FlashAttention算子 aclnnIncreFlashAttention 使用指南与实现原理解析【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer本指南以 CANN ops-transformer 仓库中 aclnnIncreFlashAttention 接口文档为核心系统讲解该增量自注意力算子的功能定位、两段式接口原型、全部入参语义与约束、可复制的完整调用示例并结合同目录的算子设计文档与 op_api、op_host 源码深入剖析其 FlashAttention 计算流程、模板拆分与 Tiling 分核原理。读完本文你将掌握如何在 NPU 上正确调用 IFA 接口完成自回归增量推理的 attention 计算并理解其底层加速机制与接口演进脉络。一、功能概述面向自回归推理的增量 Attention1.1 为什么需要增量推理对于自回归Auto-regressive的语言模型随着新词的逐个生成推理输入长度不断增大。若每次生成新词都做一次全量计算计算量会随序列长度线性膨胀推理时延不可接受。IncreFlashAttentionIFA算子在原来全量推理的基础上实现增量推理query 的 S 轴固定为 1即每一轮只计算当前待生成 token 的注意力key 和 value 是经过 KV Cache 缓存后将之前推理过的 state 信息叠加在一起的结果每个 Batch 对应的 S 轴实际长度可能不一样输入的数据是经过 padding 后的固定长度数据。相比全量场景的 FlashAttention 算子PromptFlashAttention增量推理的流程与正常全量推理并不完全等价不过增量推理的精度并无明显劣化。关于 KV CacheKV Cache 是大模型推理性能优化的常用技术。采样时Transformer 模型以给定的 prompt/context 作为初始输入进行推理可并行处理随后逐一生成额外的 token 来完善序列体现自回归性质。采样过程中Transformer 执行自注意力操作需要为当前序列中的每个项目prompt/context 或生成的 token提取键值KV向量这些向量存储在矩阵中即 KV Cache。1.2 计算公式self-attention 利用输入样本自身的关系构建注意力模型假设长度为 $n$ 的输入样本序列 $x$每个元素是 $d$ 维向量可视为 token embedding该序列经 3 个权重矩阵变换得到 3 个 $n \times d$ 矩阵。self-attention 一般定义为$$ Attention(Q,K,V)Score(Q,K)V $$本算子中 Score 函数采用 Softmax计算公式为$$ Attention(Q,K,V)Softmax(\frac{QK^T}{\sqrt{d}})V $$其中 $Q$ 与 $K^T$ 的乘积代表输入 $x$ 的注意力为避免该值过大除以 $d$ 的开根号进行缩放对每行做 softmax 归一化后再与 $V$ 相乘得到 $n \times d$ 的输出矩阵。二、产品支持情况该接口在仓库中通过 算子定义文件 中的 AICore 配置ascend910b、ascend910_93、mc62、ascend310p与文档声明保持一致产品支持情况如下产品是否支持Ascend 950PR/Ascend 950DT不支持Atlas A3 训练系列产品/Atlas A3 推理系列产品不支持Atlas A2 训练系列产品/Atlas A2 推理系列产品支持Atlas 200I/500 A2 推理产品不支持Atlas 推理系列产品支持Atlas 训练系列产品不支持三、函数原型与两段式接口每个算子分为两段式接口必须先调用aclnnIncreFlashAttentionGetWorkspaceSize获取计算所需 workspace 大小以及包含算子计算流程的执行器executor再调用aclnnIncreFlashAttention执行计算。aclnnStatus aclnnIncreFlashAttentionGetWorkspaceSize( const aclTensor *query, const aclTensorList *key, const aclTensorList *value, const aclTensor *pseShift, const aclTensor *attenMask, const aclIntArray *actualSeqLengths, int64_t numHeads, double scaleValue, char *inputLayout, int64_t numKeyValueHeads, const aclTensor *attentionOut, uint64_t *workspaceSize, aclOpExecutor **executor)aclnnStatus aclnnIncreFlashAttention( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, const aclrtStream stream)从 op_api 层源码aclnn_incre_flash_attention.cpp可以看到V1 接口实际上是内部aclnnInnerIncreFlashAttentionGetWorkspaceSize的一个薄封装它固定将pseShift置空、blockSize0、innerPrecise1并统一走内层 V4 接口的入参通道aclnnStatus ret aclnnInnerIncreFlashAttentionGetWorkspaceSize( query, key, value, nullptr, attenMask, actualSeqLengths, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, numHeads, scaleValue, inputLayout, numKeyValueHeads, 0, 1, attentionOut, workspaceSize, executor);因此理解 V1 与 V4 的参数对应关系对后续迁移大有帮助。四、aclnnIncreFlashAttentionGetWorkspaceSize 参数说明第一段接口完成入参校验与 workspace 大小计算参数语义如下参数名输入/输出描述使用说明数据类型数据格式维度(shape)非连续Tensorquery输入公式中的输入 Qquery 和 attentionOut 的 shape 需要完全一致FLOAT16、BFLOAT16ND(B, N, S, D) / (B, S, N, D) / (B, S, H)×key输入公式中的输入 Kkey、value 中对应 tensor 的 shape 需要完全一致FLOAT16、BFLOAT16ND(B, N, S, D) / (B, S, N, D) / (B, S, H)×value输入公式中的输入 Vkey、value 中对应 tensor 的 shape 需要完全一致FLOAT16、BFLOAT16ND(B, N, S, D) / (B, S, N, D) / (B, S, H)×pseShift输入位置编码预留参数暂未使用FLOAT16、BFLOAT16ND--attenMask输入attention 掩码矩阵支持空 Tensor当 attenMask 数据类型取 INT8、UINT8 时其 tensor 中的值需要为 0 或 1BOOL、INT8、UINT8ND(B, N, 1, S) / (1, N, 1, S)×actualSeqLengths输入key 和 value 的 S 轴实际长度综合约束见约束说明INT64ND(B)-numHeads输入query 的 head 个数numHeads 是 numKeyValueHeads 的倍数关系INT64---scaleValue输入公式中 d 开根号的倒数-DOUBLE---inputLayout输入标识输入 query、key、value 的数据排布格式当前支持 BSH、BNSD、BSND。用户不特意指定时建议传入 BSHSTRING---numKeyValueHeads输入key、value 中 head 个数用于支持 GQAGrouped-Query Attention场景传入 0 表示和 query 的 head 个数相等INT64---attentionOut输出公式中的输出-FLOAT16、BFLOAT16ND(B, N, S, D) / (B, S, N, D) / (B, S, H)-workspaceSize输出返回用户需要在 Device 侧申请的 workspace 大小-----executor输出返回 op 执行器包含了算子计算流程-----从 算子定义文件 可以印证各属性的默认值与声明input_layout默认BSH、scale_value默认1.0、num_key_value_heads默认0表示与 query 头数相等、block_size默认0、inner_precise默认1。其中 key、value 在算子定义中被声明为DYNAMIC动态输入因此接口层使用aclTensorList承载。返回值与错误码aclnnStatus返回状态码具体参见 aclnn 返回码。第一段接口完成入参校验以下场景报错返回值错误码描述ACLNN_ERR_PARAM_NULLPTR161001传入参数是必选输入、输出或必选属性且是空指针ACLNN_ERR_PARAM_INVALID161002query、key、value、pseShift、attenMask、attentionOut 的数据类型和数据格式不在支持的范围内ACLNN_ERR_RUNTIME_ERROR361001API 内存调用 npu runtime 的接口异常五、aclnnIncreFlashAttention 参数说明第二段接口执行计算参数名输入/输出描述workspace输入在 Device 侧申请的 workspace 内存地址workspaceSize输入在 Device 侧申请的 workspace 大小由第一段接口获取executor输入op 执行器包含算子计算流程stream输入指定执行任务的 Stream返回值为aclnnStatus状态码。六、约束说明6.1 通用约束确定性计算aclnnIncreFlashAttention 默认确定性实现相关概念可参考确定性计算。非连续场景下参数 key、value 的 tensorlist 中 tensor 的个数等于 query 的 B由于 tensorlist 限制非连续场景下 B 需要小于等于 256shape 除 S 外需要完全一致且 batch 只能为 1。参数 query 中的 N 和 numHeads 值相等key、value 的 N 和 numKeyValueHeads 值相等并且 numHeads 是 numKeyValueHeads 的倍数关系。仅支持 query 的 S 轴等于 1。当 attenMask 数据类型取 INT8、UINT8 时其 tensor 中的值需要为 0 或 1。6.2 分平台约束Atlas A2 训练系列产品/Atlas A2 推理系列产品支持 B 轴小于等于 65536N 轴小于等于 256D 轴小于等于 512query 数据类型支持 FLOAT16、BFLOAT16attentionOut、key 和 value 数据类型支持 FLOAT16 和 BFLOAT16numKeyValueHeads 数据类型支持 INT64。Atlas 推理系列产品支持 B 轴小于等于 256N 轴小于等于 256D 轴小于等于 512支持 key、value 的 S 轴小于等于 65536query、key、value 和 attentionOut 数据类型仅支持 FLOAT16numKeyValueHeads 仅支持取值 0。七、完整调用示例以下示例完整继承自接口文档具体编译和执行过程请参考编译与运行样例。仓库中另有可直接阅读的工程化样例 test_aclnn_incre_flash_attention.cpp对应 V4 接口可供对照。#include iostream #include vector #include math.h #include cstring #include acl/acl.h #include aclnn/opdev/fp16_t.h #include aclnnop/aclnn_incre_flash_attention.h using namespace std; #define CHECK_RET(cond, return_expr) \ do { \ if (!(cond)) { \ return_expr; \ } \ } while (0) #define LOG_PRINT(message, ...) \ do { \ printf(message, ##__VA_ARGS__); \ } while (0) int64_t GetShapeSize(const std::vectorint64_t shape) { int64_t shapeSize 1; for (auto i : shape) { shapeSize * i; } return shapeSize; } int Init(int32_t deviceId, aclrtStream* stream) { // 固定写法资源初始化 auto ret aclInit(nullptr); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclInit failed. ERROR: %d\n, ret); return ret); ret aclrtSetDevice(deviceId); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtSetDevice failed. ERROR: %d\n, ret); return ret); ret aclrtCreateStream(stream); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtCreateStream failed. ERROR: %d\n, ret); return ret); return 0; } template typename T int CreateAclTensor(const std::vectorT hostData, const std::vectorint64_t shape, void** deviceAddr, aclDataType dataType, aclTensor** tensor) { auto size GetShapeSize(shape) * sizeof(T); // 调用aclrtMalloc申请device侧内存 auto ret aclrtMalloc(deviceAddr, size, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtMalloc failed. ERROR: %d\n, ret); return ret); // 调用aclrtMemcpy将host侧数据拷贝到device侧内存上 ret aclrtMemcpy(*deviceAddr, size, hostData.data(), size, ACL_MEMCPY_HOST_TO_DEVICE); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtMemcpy failed. ERROR: %d\n, ret); return ret); // 计算连续tensor的strides std::vectorint64_t strides(shape.size(), 1); for (int64_t i shape.size() - 2; i 0; i--) { strides[i] shape[i 1] * strides[i 1]; } // 调用aclCreateTensor接口创建aclTensor *tensor aclCreateTensor(shape.data(), shape.size(), dataType, strides.data(), 0, aclFormat::ACL_FORMAT_ND, shape.data(), shape.size(), *deviceAddr); return 0; } int main() { // 1.固定写法device/stream初始化参考acl API手册 // 根据自己的实际device填写deviceId int32_t deviceId 0; aclrtStream stream; auto ret Init(deviceId, stream); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(Init acl failed. ERROR: %d\n, ret); return ret); // 2. 构造输入与输出需要根据API的接口自定义构造 int32_t batchSize 1; int32_t numHeads 2; int32_t headDims 16; int32_t keyNumHeads 2; int32_t sequenceLengthKV 16; std::vectorint64_t queryShape {batchSize, numHeads, 1, headDims}; // BNSD std::vectorint64_t keyShape {batchSize, keyNumHeads, sequenceLengthKV, headDims}; // BNSD std::vectorint64_t valueShape {batchSize, keyNumHeads, sequenceLengthKV, headDims}; // BNSD std::vectorint64_t attenShape {batchSize, 1, 1, sequenceLengthKV}; // B11S std::vectorint64_t outShape {batchSize, numHeads, 1, headDims}; // BNSD void *queryDeviceAddr nullptr; void *keyDeviceAddr nullptr; void *valueDeviceAddr nullptr; void *attenDeviceAddr nullptr; void *outDeviceAddr nullptr; aclTensor *queryTensor nullptr; aclTensor *keyTensor nullptr; aclTensor *valueTensor nullptr; aclTensor *attenTensor nullptr; aclTensor *outTensor nullptr; std::vectorfloat queryHostData(batchSize * numHeads * headDims, 1.0f); std::vectorfloat keyHostData(batchSize * keyNumHeads * sequenceLengthKV * headDims, 1.0f); std::vectorfloat valueHostData(batchSize * keyNumHeads * sequenceLengthKV * headDims, 1.0f); std::vectorint8_t attenHostData(batchSize * sequenceLengthKV, 0); std::vectorfloat outHostData(batchSize * numHeads * headDims, 1.0f); // 创建query aclTensor ret CreateAclTensor(queryHostData, queryShape, queryDeviceAddr, aclDataType::ACL_FLOAT16, queryTensor); CHECK_RET(ret ACL_SUCCESS, return ret); // 创建key aclTensor ret CreateAclTensor(keyHostData, keyShape, keyDeviceAddr, aclDataType::ACL_FLOAT16, keyTensor); CHECK_RET(ret ACL_SUCCESS, return ret); int kvTensorNum 1; aclTensor *tensorsOfKey[kvTensorNum]; tensorsOfKey[0] keyTensor; auto tensorKeyList aclCreateTensorList(tensorsOfKey, kvTensorNum); // 创建value aclTensor ret CreateAclTensor(valueHostData, valueShape, valueDeviceAddr, aclDataType::ACL_FLOAT16, valueTensor); CHECK_RET(ret ACL_SUCCESS, return ret); aclTensor *tensorsOfValue[kvTensorNum]; tensorsOfValue[0] valueTensor; auto tensorValueList aclCreateTensorList(tensorsOfValue, kvTensorNum); // 创建atten aclTensor ret CreateAclTensor(attenHostData, attenShape, attenDeviceAddr, aclDataType::ACL_INT8, attenTensor); CHECK_RET(ret ACL_SUCCESS, return ret); // 创建out aclTensor ret CreateAclTensor(outHostData, outShape, outDeviceAddr, aclDataType::ACL_FLOAT16, outTensor); CHECK_RET(ret ACL_SUCCESS, return ret); std::vectorint64_t actualSeqlenVector {sequenceLengthKV}; auto actualSeqLengths aclCreateIntArray(actualSeqlenVector.data(), actualSeqlenVector.size()); int64_t numKeyValueHeads numHeads; double scaleValue 1 / sqrt(headDims); // 1/sqrt(d) string sLayerOut BNSD; char layerOut[sLayerOut.length()1]; strcpy(layerOut, sLayerOut.c_str()); // 3. 调用CANN算子库API uint64_t workspaceSize 0; aclOpExecutor* executor; // 调用第一段接口 ret aclnnIncreFlashAttentionGetWorkspaceSize(queryTensor, tensorKeyList, tensorValueList, nullptr, attenTensor, actualSeqLengths, numHeads, scaleValue, layerOut, numKeyValueHeads, outTensor, workspaceSize, executor); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclnnIncreFlashAttentionGetWorkspaceSize failed. ERROR: %d\n, ret); return ret); // 根据第一段接口计算出的workspaceSize申请device内存 void* workspaceAddr nullptr; if (workspaceSize 0) { ret aclrtMalloc(workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(allocate workspace failed. ERROR: %d\n, ret); return ret); } // 调用第二段接口 ret aclnnIncreFlashAttention(workspaceAddr, workspaceSize, executor, stream); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclnnIncreFlashAttention failed. ERROR: %d\n, ret); return ret); // 4.固定写法同步等待任务执行结束 ret aclrtSynchronizeStream(stream); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtSynchronizeStream failed. ERROR: %d\n, ret); return ret); // 5. 获取输出的值将device侧内存上的结果拷贝至host侧需要根据具体API的接口定义修改 auto size GetShapeSize(outShape); std::vectorop::fp16_t resultData(size, 0); ret aclrtMemcpy(resultData.data(), resultData.size() * sizeof(resultData[0]), outDeviceAddr, size * sizeof(resultData[0]), ACL_MEMCPY_DEVICE_TO_HOST); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(copy result from device to host failed. ERROR: %d\n, ret); return ret); for (int64_t i 0; i size; i) { std::cout index: i : static_castfloat(resultData[i]) std::endl; } // 6. 释放资源 aclDestroyTensor(queryTensor); aclDestroyTensor(keyTensor); aclDestroyTensor(valueTensor); aclDestroyTensor(attenTensor); aclDestroyTensor(outTensor); aclDestroyIntArray(actualSeqLengths); aclrtFree(queryDeviceAddr); aclrtFree(keyDeviceAddr); aclrtFree(valueDeviceAddr); aclrtFree(attenDeviceAddr); aclrtFree(outDeviceAddr); if (workspaceSize 0) { aclrtFree(workspaceAddr); } aclrtDestroyStream(stream); aclrtResetDevice(deviceId); aclFinalize(); return 0; }示例要点解读shape 约定示例采用 BNSD 排布query 的 S 轴固定为 1queryShape {1, 2, 1, 16}key/value 的 S 轴为 KV 序列长度16attenMask 使用(B, 1, 1, S)的 B11S 形状。tensorlist 构造key、value 必须通过aclCreateTensorList封装为aclTensorList传入个数与 B 相等。workspace 生命周期第一段接口返回workspaceSize后通过aclrtMalloc申请 device 内存第二段接口执行完后需aclrtFree释放。属性取值scaleValue 1/sqrt(headDims)numKeyValueHeads在此例中等于numHeads等价于 MHA 场景inputLayout传BNSD。八、底层实现原理结合算子设计文档与源码8.1 整体计算流程按照 IFA 算子设计介绍 的说明算子按照 FlashAttention 正向计算流程实现query 与转置后的 key 做 matmul 得到初始 attention_score与位置编码 pse 相加后乘以缩放系数 scale_value随后通过 atten_mask 进行 select 操作将 mask 中为 true 的位置遮蔽为负的极小值经 softmax 后变为 0 从而达成遮蔽效果。为实现 FlashAttention 加速使用 FlashSoftmax 操作替代原公式中的 softmaxFlashSoftmax 对 masked_attention_score 的 Skvkey、value 的 sequence length方向进行切分因而存在一个刷新流程每次 FlashSoftmax 只处理切分后的一个 SkvSplitSkv 轴切分后的序列长度从第二次循环开始记录 exp$exp[i] e^{max_{i-1} - max_i}$i 为 Skv 切分后的循环变量从 1 开始从 i 1 开始增加 Mul 和 Add 操作将上一次MM[PV]的结果与当前 exp 相乘再与本次MM[PV]相加结果保存到 GM依此类推遍历完 Skv由于 FlashSoftmax 计算中的除 sum 被后移到输出 attention_out 之前最后需要将 UB 中的 attention_out 按行除以 softmax_sum并将最终完整结果写回输出内存。单核主流程伪代码如下摘自设计文档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); } } } }8.2 模板设计与数据切分由于硬件 buffer 有限而数据量巨大无法一次算完需要 Tiling 切分融合算子融合了 element-wise、broadcast、reduce 及 matmul 多类场景需要按切分轴拆分模板。模板拆分需考虑核数用满、各核负载均匀、AIC 与 AIV 间数据量匹配算力。IFA 算子包含 B、N2key/value 的 N、Gquery_N/kv_N、S1query 的 S、S2key/value 的 S共 5 个轴S1 固定为 1 不参与切分G 轴只在 Vector 计算时切块。BN2S2 切分逻辑核间外切先按 BN2 分核将 BN2 个 SD 块分配到多个核当 BN2 小于阈值0.4 × 总核数时再对 S2 轴外切SplitKV 份总块数为 BN2 × SplitKv各核计算子块后规约即 FlashDecode 流程核内由于单 core 缓存有限按缓存大小对 S2 轴或 KV 子块的 S2 轴继续切分即 FlashAttention 过程。仓库中模板文件位于 op_kernel 目录包括模板对应文件说明CV 模板incre_flash_attention_split_Bbn2s2_Us2.hIFA 基础模板matmul 在 CubeCore 执行调用 AscendC 高阶 APIAll-Vector 模板incre_flash_attention_allvec_new.hmatmul 由 vector 实现降低 Cube 启动与 CV 通信开销matmul 基础 API 模板incre_flash_attention_preload.h基于 CV 模板用 Cube 编程视角重写 matmul优化 CUBE/VEC 流水N-Buffer伪量化 MSD DD 模板incre_flash_attention_preload_dd.h用于伪量化 MTP 场景当前仅 FIA 算子调用MLA 全量化模板incre_flash_attention_preload_mla.hMLA 场景 INT8 QKV BF16 rope 的 attention 计算当前仅 FIA 算子调用其中 All-Vector 模板在 Atlas 推理系列产品上全部使用在 Atlas A2 上用于非 PA、非 GQA 且 Q、KV、Output 全为 FP16 的场景。8.3 FlashDecode 规约S2 轴外切到不同核完成 attention 计算后需要对结果做 Reduce 操作共 BN2 个 SD 大块每个 core 合并一个大块的所有子块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); }8.4 特性扩展AntiQuant、PageAttention 与 GQAAntiQuant MSD 算法IFA AntiQuant 场景矩阵计算为 $C A \times (B offset) \times scale$A 为 FP16/BF16B 为 INT8。经典反量化需将较大的 B 矩阵搬入 Vector性能差IFA 场景 A 矩阵较小通过变换 A 来适配 BA 展开为 int8 存储的多行并打包成新矩阵 AA计算CC AA * Bint8×int8int32再对 CC 做 Reduce 得到 C。PageAttentionKV block 内存不连续MatMul 提供回调函数做 B 矩阵的 GM→L1 拷贝IFA 中实现相应拷贝函数回调在 Cube 中执行参数通过 GM 传递Vector 设置参数到 GM确保 DCCI后再通知 MatMul 工作。GQAG queryHeadNum / KvHeadNumVector 上 G 轴切分由当前操作涉及的输入输出 UB 大小决定当 G 过大时在 G 轴切分g target_ub_size() / column_size若g G则g G再按g × column子块处理。8.5 Tiling 分核与 TilingKey 规划Tiling 的目标是找到高效的 NPU 执行方式总块数为 BN2 或 BN2 × SplitKv输入为核数、块数、块负载每个分块的 S 轴实际长度处理上根据负载对连续块组合重排使核间负载差值最小输出为 blockid 数组每个元素对应一个核的起始 blockid末尾追加总块数。TilingKey 为 uint64 类型每个模板参数对应一个十进制位具体实现见 incre_flash_attention_tiling 下的 GenTilingKey 函数。核心字段摘自设计文档十进制位变量说明0layoutValQ 的 shape 格式0: BNSD1: BSH/BSND2: TND1inputQValquery 数据类型0: FP162: BF163: INT82inputKvValKV 数据类型0: FP162: BF163: INT84: INT43outputValoutput 数据类型0: FP162: BF163: INT84originVal同 inputQVal5[bit0]splitKvVal开启 FlashDecode 标志5[bit1]paVal开启 PageAttention 标志5[bit2]antiquantModeVal开启 PerToken 伪量化标记6antiquantMode_量化模式0: 无效值2: K-perChannel-V-perToken7kvLayoutValKV 的 shape 格式仅伪量化 MSD DD 与 MLA 全量化模板有效8amlaMode该字段废弃取值只能为 09balanceMode开启新负载均衡算法标志仅 MLA 全量化模板10...14-预留字段值为 015perfMode_模板编号0: C1_V21: 全V2: C1_V13: matmul 基础 API 模板5: MLA 全量化模板6: 伪量化 MSD DD 模板16modeVal1: IFA TilingKey Base2: IFA 启用 SysPrefix 功能8.6 Infershape 与入参校验在 host 侧incre_flash_attention_infershape.cpp 完成 shape 推导直接将attentionOut的 shape 置为 query 的 shape保证二者一致并根据inputLayout属性校验维度例如 BSH 要求 query 为 3 维、BNSD/BSND 要求 4 维等。这从源码层面印证了接口文档中query 和 attentionOut 的 shape 需要完全一致的约束。九、接口演进与迁移建议该接口文档明确声明aclnnIncreFlashAttention 后续版本会废弃请使用最新接口 aclnnIncreFlashAttentionV4。从 op_api 源码aclnn_incre_flash_attention.cpp可以看到运行时告警OP_LOGW(aclnnIncreFlashAttentionGetWorkspaceSize is scheduled to be deprecated in December 2026, and will be replaced by the aclnnIncreFlashAttentionV4GetWorkspaceSize. ...);接口演进脉络对应 V2、V3、V4 文档V1本文基础增量推理pseShift 预留V2扩展基础能力V3在 V2 基础上新增位置编码pseShift 生效、PageAttention、KV Cache 反量化特性V4兼容 V3 功能新增kv 左 Padding 特性并支持 A3 系列产品。V4 相比 V1 增加了一组量化/反量化因子dequantScale1、quantScale1、dequantScale2、quantScale2、quantOffset2、antiquantScale、antiquantOffset、blocktable、kvPaddingSize、blockSize、innerPrecise 等入参同时 attenMask 支持(B, S)、(B, 1, S)、(B, 1, 1, S)多种形状key/value 支持 INT8 量化输入。若需使用位置编码、page attention、量化等高级特性建议直接迁移到 V4 接口。十、总结与进一步阅读aclnnIncreFlashAttention 是 CANN ops-transformer 中支撑自回归大模型增量推理的核心 attention 算子query 的 S 轴固定为 1key/value 以 KV Cache 形态提供通过 FlashAttention 在线 softmax、FlashDecode 跨核规约、模板化 kernel 拆分与 Tiling 分核等手段在 NPU 上实现高效的逐 token 自注意力计算。调用侧遵循两段式接口规范通过 GetWorkspaceSize 获取 workspace 与执行器后再执行计算。相关资源索引aclnnIncreFlashAttention 接口文档本文核心依据IFA 算子设计介绍计算流程、模板与 Tiling 设计IncreFlashAttention README完整约束与调用说明aclnnIncreFlashAttentionV4 接口文档推荐迁移目标op_api 封装源码op_host 算子定义工程化调用示例配套概念文档两段式接口、aclnn 返回码、编译与运行样例【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考