news 2026/9/19 20:55:02

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
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 计算,并理解其底层加速机制与接口演进脉络。

一、功能概述:面向自回归推理的增量 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 配置(ascend910bascend910_93mc62ascend310p)与文档声明保持一致,产品支持情况如下:

产品是否支持
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=0innerPrecise=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输入公式中的输入 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。用户不特意指定时建议传入 "BSH"STRING---
numKeyValueHeads输入key、value 中 head 个数用于支持 GQA(Grouped-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.0num_key_value_heads默认0(表示与 query 头数相等)、block_size默认0inner_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 需要小于等于 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 正向计算流程实现:

  1. query 与转置后的 key 做 matmul 得到初始 attention_score,与位置编码 pse 相加后乘以缩放系数 scale_value;随后通过 atten_mask 进行 select 操作,将 mask 中为 true 的位置遮蔽为负的极小值,经 softmax 后变为 0 从而达成遮蔽效果。
  2. 为实现 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.hIFA 基础模板,matmul 在 CubeCore 执行,调用 AscendC 高阶 API
All-Vector 模板incre_flash_attention_allvec_new.hmatmul 由 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.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 与 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 > Gg = 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: BNSD;1: BSH/BSND;2: TND
1inputQValquery 数据类型:0: FP16;2: BF16;3: INT8
2inputKvValKV 数据类型:0: FP16;2: BF16;3: INT8;4: INT4
3outputValoutput 数据类型:0: FP16;2: BF16;3: INT8
4originVal同 inputQVal
5[bit0]splitKvVal开启 FlashDecode 标志
5[bit1]paVal开启 PageAttention 标志
5[bit2]antiquantModeVal开启 PerToken 伪量化标记
6antiquantMode_量化模式:0: 无效值;2: K-perChannel-V-perToken
7kvLayoutValKV 的 shape 格式(仅伪量化 MSD DD 与 MLA 全量化模板有效)
8amlaMode该字段废弃,取值只能为 0
9balanceMode开启新负载均衡算法标志(仅 MLA 全量化模板)
10...14-预留字段,值为 0
15perfMode_模板编号:0: C1_V2;1: 全V;2: C1_V1;3: matmul 基础 API 模板;5: MLA 全量化模板;6: 伪量化 MSD DD 模板
16modeVal1: 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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/19 20:54:38

N_m3u8DL-RE mux failed 深度排查实录

N_m3u8DL-RE mux failed 深度排查实录 【免费下载链接】N_m3u8DL-RE Cross-Platform, modern and powerful stream downloader for MPD/M3U8/ISM. English/简体中文/繁體中文. 项目地址: https://gitcode.com/GitHub_Trending/nm3/N_m3u8DL-RE N_m3u8DL-RE 是一款跨平台…

作者头像 李华
网站建设 2026/9/19 20:53:36

Spark 增量处理:基于 Checkpoint 的状态恢复与增量数据摄取技术详解

Spark 增量处理&#xff1a;基于 Checkpoint 的状态恢复与增量数据摄取技术详解本文深入探讨Spark增量处理方案&#xff0c;重点介绍基于Checkpoint的状态恢复机制与增量数据摄取实现方法&#xff0c;通过示例代码和架构图帮助读者掌握Spark增量处理的核心技术和最佳实践。1. S…

作者头像 李华
网站建设 2026/9/19 20:52:13

JSVMP 逆向 testab 插装日志读不懂?TaoToken 这样给 Codex 配通道

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/19 20:50:10

Vuex 4 快速入门:从零构建你的第一个集中式状态管理 Store

Vuex 4 快速入门&#xff1a;从零构建你的第一个集中式状态管理 Store 【免费下载链接】vuex &#x1f5c3;️ Centralized State Management for Vue.js. 项目地址: https://gitcode.com/gh_mirrors/vu/vuex Vuex 是 Vue.js 官方的集中式状态管理模式与库&#xff0c;而…

作者头像 李华
网站建设 2026/9/19 20:46:26

Codex 下载与本地部署:命令行 AI 编码助手安装与模型接入避坑

上周有位同事在群里发了一张终端截图&#xff0c;满屏红字&#xff0c;最扎眼的是接口返回 404&#xff0c;说找不到/responses这个路径。他为了把 Codex 跑起来折腾了整整两天&#xff0c;中间重装过 Node&#xff0c;换过三个模型&#xff0c;最后发现只是配置文件里少写了一…

作者头像 李华