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 计算,并理解其底层加速机制与接口演进脉络。
一、功能概述:面向自回归推理的增量 Attention
1.1 为什么需要增量推理
对于自回归(Auto-regressive)的语言模型,随着新词的逐个生成,推理输入长度不断增大。若每次生成新词都做一次全量计算,计算量会随序列长度线性膨胀,推理时延不可接受。IncreFlashAttention(IFA)算子在原来全量推理的基础上实现增量推理:
- query 的 S 轴固定为 1,即每一轮只计算当前待生成 token 的注意力;
- key 和 value 是经过 KV Cache 缓存后,将之前推理过的 state 信息叠加在一起的结果;
- 每个 Batch 对应的 S 轴实际长度可能不一样,输入的数据是经过 padding 后的固定长度数据。
相比全量场景的 FlashAttention 算子(PromptFlashAttention),增量推理的流程与正常全量推理并不完全等价,不过增量推理的精度并无明显劣化。
关于 KV Cache:KV 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置空、blockSize=0、innerPrecise=1,并统一走内层 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) | 非连续Tensor |
|---|---|---|---|---|---|---|---|
| query | 输入 | 公式中的输入 Q | query 和 attentionOut 的 shape 需要完全一致 | FLOAT16、BFLOAT16 | ND | (B, N, S, D) / (B, S, N, D) / (B, S, H) | × |
| key | 输入 | 公式中的输入 K | key、value 中对应 tensor 的 shape 需要完全一致 | FLOAT16、BFLOAT16 | ND | (B, N, S, D) / (B, S, N, D) / (B, S, H) | × |
| value | 输入 | 公式中的输入 V | key、value 中对应 tensor 的 shape 需要完全一致 | FLOAT16、BFLOAT16 | ND | (B, N, S, D) / (B, S, N, D) / (B, S, H) | × |
| pseShift | 输入 | 位置编码 | 预留参数,暂未使用 | FLOAT16、BFLOAT16 | ND | - | - |
| attenMask | 输入 | attention 掩码矩阵 | 支持空 Tensor;当 attenMask 数据类型取 INT8、UINT8 时,其 tensor 中的值需要为 0 或 1 | BOOL、INT8、UINT8 | ND | (B, N, 1, S) / (1, N, 1, S) | × |
| actualSeqLengths | 输入 | key 和 value 的 S 轴实际长度 | 综合约束见约束说明 | INT64 | ND | (B) | - |
| numHeads | 输入 | query 的 head 个数 | numHeads 是 numKeyValueHeads 的倍数关系 | INT64 | - | - | - |
| scaleValue | 输入 | 公式中 d 开根号的倒数 | - | DOUBLE | - | - | - |
| inputLayout | 输入 | 标识输入 query、key、value 的数据排布格式 | 当前支持 BSH、BNSD、BSND。用户不特意指定时建议传入 "BSH" | STRING | - | - | - |
| numKeyValueHeads | 输入 | key、value 中 head 个数 | 用于支持 GQA(Grouped-Query Attention)场景,传入 0 表示和 query 的 head 个数相等 | INT64 | - | - | - |
| attentionOut | 输出 | 公式中的输出 | - | FLOAT16、BFLOAT16 | ND | (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_NULLPTR | 161001 | 传入参数是必选输入、输出或必选属性,且是空指针 |
| ACLNN_ERR_PARAM_INVALID | 161002 | query、key、value、pseShift、attenMask、attentionOut 的数据类型和数据格式不在支持的范围内 |
| ACLNN_ERR_RUNTIME_ERROR | 361001 | API 内存调用 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 需要小于等于 256),shape 除 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 轴小于等于 65536,N 轴小于等于 256,D 轴小于等于 512;
- query 数据类型支持 FLOAT16、BFLOAT16;attentionOut、key 和 value 数据类型支持 FLOAT16 和 BFLOAT16;
- numKeyValueHeads 数据类型支持 INT64。
Atlas 推理系列产品:
- 支持 B 轴小于等于 256,N 轴小于等于 256,D 轴小于等于 512;
- 支持 key、value 的 S 轴小于等于 65536;
- query、key、value 和 attentionOut 数据类型仅支持 FLOAT16;
- numKeyValueHeads 仅支持取值 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::vector<int64_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::vector<T>& hostData, const std::vector<int64_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::vector<int64_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::vector<int64_t> queryShape = {batchSize, numHeads, 1, headDims}; // BNSD std::vector<int64_t> keyShape = {batchSize, keyNumHeads, sequenceLengthKV, headDims}; // BNSD std::vector<int64_t> valueShape = {batchSize, keyNumHeads, sequenceLengthKV, headDims}; // BNSD std::vector<int64_t> attenShape = {batchSize, 1, 1, sequenceLengthKV}; // B11S std::vector<int64_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::vector<float> queryHostData(batchSize * numHeads * headDims, 1.0f); std::vector<float> keyHostData(batchSize * keyNumHeads * sequenceLengthKV * headDims, 1.0f); std::vector<float> valueHostData(batchSize * keyNumHeads * sequenceLengthKV * headDims, 1.0f); std::vector<int8_t> attenHostData(batchSize * sequenceLengthKV, 0); std::vector<float> 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::vector<int64_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::vector<op::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_cast<float>(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 轴固定为 1(
queryShape = {1, 2, 1, 16}),key/value 的 S 轴为 KV 序列长度(16);attenMask 使用(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 操作替代原公式中的 softmax:FlashSoftmax 对 masked_attention_score 的 Skv(key、value 的 sequence length)方向进行切分,因而存在一个刷新流程:
- 每次 FlashSoftmax 只处理切分后的一个 SkvSplit(Skv 轴切分后的序列长度),从第二次循环开始记录 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、N2(key/value 的 N)、G(query_N/kv_N)、S1(query 的 S)、S2(key/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 目录,包括:
| 模板 | 对应文件 | 说明 |
|---|---|---|
| C+V 模板 | incre_flash_attention_split_Bbn2s2_Us2.h | IFA 基础模板,matmul 在 CubeCore 执行,调用 AscendC 高阶 API |
| All-Vector 模板 | incre_flash_attention_allvec_new.h | matmul 由 vector 实现,降低 Cube 启动与 CV 通信开销 |
| matmul 基础 API 模板 | incre_flash_attention_preload.h | 基于 C+V 模板,用 Cube 编程视角重写 matmul,优化 CUBE/VEC 流水(N-Buffer) |
| 伪量化 MSD DD 模板 | incre_flash_attention_preload_dd.h | 用于伪量化 MTP 场景,当前仅 FIA 算子调用 |
| MLA 全量化模板 | incre_flash_attention_preload_mla.h | MLA 场景 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 与 GQA
- AntiQuant MSD 算法:IFA AntiQuant 场景矩阵计算为 $C = A \times (B + offset) \times scale$,A 为 FP16/BF16,B 为 INT8。经典反量化需将较大的 B 矩阵搬入 Vector,性能差;IFA 场景 A 矩阵较小,通过变换 A 来适配 B:A 展开为 int8 存储的多行并打包成新矩阵 AA,计算
CC = AA * B(int8×int8=int32),再对 CC 做 Reduce 得到 C。 - PageAttention:KV block 内存不连续,MatMul 提供回调函数做 B 矩阵的 GM→L1 拷贝,IFA 中实现相应拷贝函数;回调在 Cube 中执行,参数通过 GM 传递,Vector 设置参数到 GM(确保 DCCI)后再通知 MatMul 工作。
- GQA:G = queryHeadNum / KvHeadNum,Vector 上 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 函数)。核心字段(摘自设计文档):
| 十进制位 | 变量 | 说明 |
|---|---|---|
| 0 | layoutVal | Q 的 shape 格式:0: BNSD;1: BSH/BSND;2: TND |
| 1 | inputQVal | query 数据类型:0: FP16;2: BF16;3: INT8 |
| 2 | inputKvVal | KV 数据类型:0: FP16;2: BF16;3: INT8;4: INT4 |
| 3 | outputVal | output 数据类型:0: FP16;2: BF16;3: INT8 |
| 4 | originVal | 同 inputQVal |
| 5[bit0] | splitKvVal | 开启 FlashDecode 标志 |
| 5[bit1] | paVal | 开启 PageAttention 标志 |
| 5[bit2] | antiquantModeVal | 开启 PerToken 伪量化标记 |
| 6 | antiquantMode_ | 量化模式:0: 无效值;2: K-perChannel-V-perToken |
| 7 | kvLayoutVal | KV 的 shape 格式(仅伪量化 MSD DD 与 MLA 全量化模板有效) |
| 8 | amlaMode | 该字段废弃,取值只能为 0 |
| 9 | balanceMode | 开启新负载均衡算法标志(仅 MLA 全量化模板) |
| 10...14 | - | 预留字段,值为 0 |
| 15 | perfMode_ | 模板编号:0: C1_V2;1: 全V;2: C1_V1;3: matmul 基础 API 模板;5: MLA 全量化模板;6: 伪量化 MSD DD 模板 |
| 16 | modeVal | 1: IFA TilingKey Base;2: 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 轴固定为 1,key/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),仅供参考