CANN ops-transformer 中 aclnnFlashAttentionVarLenScoreV2 接口实战:基于 TND 排布的可变长 FlashAttention 训练加速
【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer
本篇技术指南围绕 CANN ops-transformer 仓库中 FlashAttentionScore 算子族的aclnnFlashAttentionVarLenScoreV2两段式接口展开,面向训练场景下一次性传入多个长度不相等 sequence 的自注意力(self-attention)计算需求。读完本文,你将掌握该接口的产品支持范围、TND 数据排布与累积长度语义、全部 26 个入参与 4 个输出的含义与取值约束、pseType/sparseMode 等关键属性的组合规则,并能基于仓库中的完整示例代码编写出可编译运行的调用程序,同时了解其底层 op_api、op_host 实现链路。
一、接口概述与产品支持情况
aclnnFlashAttentionVarLenScoreV2是 CANN ops-transformer 中 FlashAttentionScore 算子在训练场景下的可变长(VarLen)V2 版本。它与定长接口 aclnnFlashAttentionScoreV2 的核心区别在于:该接口支持可变长 S 的计算,即一次调用可传入多个长度不相等的 sequence。使用此接口时,query、key、value 使用 TND 格式传入数据,其中 T(total number)表示所有 sequence 的 length 总和,同时使用actualSeqQLenOptional与actualSeqKvLenOptional传入每个 sequence 依次的累积长度以区分不同 sequence,每个 sequence 单独计算其注意力结果。
其产品支持情况如下表所示:
| 产品 | 是否支持 |
|---|---|
| Ascend 950PR / Ascend 950DT | 支持 |
| Atlas A3 训练系列产品 / Atlas A3 推理系列产品 | 支持 |
| Atlas A2 训练系列产品 / Atlas A2 推理系列产品 | 支持 |
| Atlas 200I/500 A2 推理产品 | 不支持 |
| Atlas 推理系列产品(310P) | 不支持 |
| Atlas 训练系列产品(910) | 不支持 |
说明:仓库根目录 README.md 指出 ops-transformer 定位为 CANN 提供的 transformer 类大模型算子库,FlashAttentionScore 系列算子(README.md)即其注意力计算的核心成员,本文接口即隶属于该算子族。
二、功能说明与计算公式
该接口在训练场景下使用 FlashAttention 算法实现 self-attention 计算,其正向计算公式根据pseType的取值分为两种情况:
- pseType = 1 时:与 aclnnFlashAttentionVarLenScore 的计算公式相同,即先 add 再 mul:
$$ attention_out=Dropout(Softmax(Mask(scale*(pse+query*key^T),atten_mask)),keep_prob)*value $$
- pseType = 其他取值时(0、2、3),先 mul 再 add:
$$ attention_out=Dropout(Softmax(Mask(scale*(query*key^T) + pse),atten_mask),keep_prob)*value $$
从仓库源码结构看,op_api 层的 flash_attention_score.cpp 将pse_type作为算子属性透传给底层(OP_ATTR(...sparseMode, pseType, ...)),而算子原型定义 flash_attention_score_def.cpp 中Attr("pse_type").AttrType(OPTIONAL).Int(1)表明其默认值为 1(先 add 再 mul),这与 README 中“pseType=1 时需要先 add 再 mul,pseType≠1 时需要先 mul 再 add”的描述一致。
三、两段式接口与函数原型
与其他 aclnn 单算子 API 一致,该算子采用两段式接口调用方式:必须先调用第一段接口aclnnFlashAttentionVarLenScoreV2GetWorkspaceSize获取计算所需 workspace 大小以及包含了算子计算流程的执行器(executor),再调用第二段接口aclnnFlashAttentionVarLenScoreV2执行计算。
两段接口的函数原型如下:
aclnnStatus aclnnFlashAttentionVarLenScoreV2GetWorkspaceSize( const aclTensor *query, const aclTensor *key, const aclTensor *value, const aclTensor *realShiftOptional, const aclTensor *dropMaskOptional, const aclTensor *paddingMaskOptional, const aclTensor *attenMaskOptional, const aclIntArray *prefixOptional, const aclIntArray *actualSeqQLenOptional, const aclIntArray *actualSeqKvLenOptional, const aclIntArray *qStartIdxOptional, const aclIntArray *kvStartIdxOptional, double scaleValue, double keepProb, int64_t preTokens, int64_t nextTokens, int64_t headNum, char *inputLayout, int64_t innerPrecise, int64_t sparseMode, int64_t pseType, const aclTensor *softmaxMaxOut, const aclTensor *softmaxSumOut, const aclTensor *softmaxOutOut, const aclTensor *attentionOutOut, uint64_t *workspaceSize, aclOpExecutor **executor)aclnnStatus aclnnFlashAttentionVarLenScoreV2( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, const aclrtStream stream)第一段接口完成入参校验、shape 推断与 workspace 大小计算;第二段接口在指定的 stream 上真正下发计算任务。需要注意的是,第二段接口不能重复调用。
四、第一段接口参数详解
aclnnFlashAttentionVarLenScoreV2GetWorkspaceSize的参数可分为输入 Tensor、可选输入 Tensor、标量属性与输出四类,完整参数说明如下:
| 参数名 | 输入/输出 | 描述 | 使用说明 | 数据类型 | 数据格式 | 维度(shape) | 非连续Tensor |
|---|---|---|---|---|---|---|---|
| query | 输入 | 公式中的 query | 数据类型与 key/value 一致 | FLOAT16、BFLOAT16、FLOAT32 | ND | [TND] | √ |
| key | 输入 | 公式中的 key | 数据类型与 query/value 一致 | FLOAT16、BFLOAT16、FLOAT32 | ND | [TND] | √ |
| value | 输入 | 公式中的 value | 数据类型与 query/key 一致 | FLOAT16、BFLOAT16、FLOAT32 | ND | [TND] | √ |
| realShiftOptional | 可选输入 | 公式中的 pse | 数据类型与 query 一致,需与 pseType 配套使用 | FLOAT16、BFLOAT16、FLOAT32 | ND | [B,N,1024,Skv]、[1,N,1024,Skv]、[pseTotalLen]、[B,N]、[N] | √ |
| dropMaskOptional | 可选输入 | 公式中的 Dropout | - | UINT8 | ND | 0、1 | √ |
| paddingMaskOptional | 可选输入 | 预留参数,暂未使用 | - | - | - | - | - |
| attenMaskOptional | 可选输入 | 公式中的 atten_mask | 取值为 1 代表该位不参与计算,为 0 代表该位参与计算 | BOOL、UINT8 | ND | [B,N,Sq,Skv]、[B,1,Sq,Skv]、[1,1,Sq,Skv]、[Sq,Skv] | √ |
| prefixOptional | 可选输入 | 代表 prefix 稀疏计算场景每个 Batch 的 N 值 | - | INT64 | ND | 0、1 | - |
| actualSeqQLenOptional | 可选输入 | 描述每个 Batch 对应的 query 的 sequence length | - | INT64 | ND | 0、1 | - |
| actualSeqKvLenOptional | 可选输入 | 描述每个 Batch 对应的 key/value 的 sequence length | - | INT64 | ND | 0、1 | - |
| qStartIdxOptional | 可选输入 | 代表外切场景,当前分块的 query 的 sequence 在全局中的起始索引 | - | INT64 | ND | 0、1 | - |
| kvStartIdxOptional | 可选输入 | 代表外切场景,当前分块的 key 和 value 的 sequence 在全局中的起始索引 | - | INT64 | ND | 0、1 | - |
| scaleValue | 可选输入 | 公式中的 scale,代表缩放系数 | - | DOUBLE | - | - | - |
| keepProb | 可选输入 | 代表 dropMaskOptional 中 1 的比例 | - | DOUBLE | - | - | - |
| preTokens | 可选输入 | 用于稀疏计算,表示 sliding window 的左边界 | - | INT64 | - | - | - |
| nextTokens | 可选输入 | 用于稀疏计算,表示 sliding window 的右边界 | - | INT64 | - | - | - |
| headNum | 输入 | 代表单卡的 head 个数,即输入 query 的 N 轴长度 | - | INT64 | - | - | - |
| inputLayout | 输入 | 代表输入 query、key、value 的数据排布格式 | 支持 TND | String | - | - | - |
| innerPrecise | 可选输入 | 用于提升精度 | 默认配置为 0 即可 | INT64 | - | - | - |
| sparseMode | 可选输入 | 表示 sparse 的模式 | 支持配置值为 0、1、2、3、4、6、7、8 | INT64 | - | - | - |
| pseType | 可选输入 | 控制 mul 与 add 计算顺序 | 支持配置值为 0、1、2、3 | INT64 | - | - | - |
| softmaxMaxOut | 输出 | Softmax 计算的 Max 中间结果,用于反向计算 | - | FLOAT | ND | [N,T,8] | √ |
| softmaxSumOut | 输出 | Softmax 计算的 Sum 中间结果,用于反向计算 | - | FLOAT | ND | [N,T,8] | √ |
| softmaxOutOut | 输出 | 预留参数,暂未使用 | - | - | - | - | - |
| attentionOutOut | 输出 | 计算公式的最终输出 | 数据类型和 shape 类型与 query 保持一致 | FLOAT16、BFLOAT16、FLOAT32 | ND | [TND] | √ |
| workspaceSize | 输出 | 返回需要在 Device 侧申请的 workspace 大小 | - | - | - | - | - |
| executor | 输出 | 返回 op 执行器,包含了算子计算流程 | - | - | - | - | - |
从仓库源码看,infershape 实现 flash_attention_score_infershape.cpp 对 TND 场景做了专门校验:TND 排布下要求headNum必须等于 query 的 N 轴长度(For TND layout, headNum must equal the N dim of query),并校验 query 与 key 的 D 轴相等、key 的 D 轴不小于 value 的 D 轴(qD == kD && kD >= vD),这些校验与文档“约束说明”中的 B/D/inputLayout 约束相互印证。
五、返回值与错误码
两个接口均返回aclnnStatus状态码,具体定义参见 aclnn 返回码。第一段接口完成入参校验,出现以下场景时报错:
| 返回值 | 错误码 | 描述 |
|---|---|---|
| ACLNN_ERR_PARAM_NULLPTR | 161001 | 传入参数是必选输入、输出或者必选属性,且是空指针 |
| ACLNN_ERR_PARAM_INVALID | 161002 | query、key、value、realShiftOptional、dropMaskOptional、paddingMaskOptional、attenMaskOptional、softmaxMaxOut、softmaxSumOut、softmaxOutOut、attentionOutOut 的数据类型不在支持的范围内 |
| ACLNN_ERR_PARAM_INVALID | 161002 | query、key、value、realShiftOptional、dropMaskOptional、paddingMaskOptional、attenMaskOptional、softmaxMaxOut、softmaxSumOut、softmaxOutOut、attentionOutOut 的数据格式不在支持的范围内 |
六、第二段接口参数说明
aclnnFlashAttentionVarLenScoreV2的参数较少,执行时仅需传入第一段接口产出的三要素与目标 stream:
| 参数名 | 输入/输出 | 描述 |
|---|---|---|
| workspace | 输入 | 在 Device 侧申请的 workspace 内存地址 |
| workspaceSize | 输入 | 在 Device 侧申请的 workspace 大小,由第一段接口获取 |
| executor | 输入 | op 执行器,包含了算子计算流程 |
| stream | 输入 | 指定执行任务的 Stream |
七、约束说明
7.1 通用约束
- 确定性计算:
aclnnFlashAttentionVarLenScoreV2为默认确定性实现(相关背景可参考 确定性计算)。 - 该接口与 PyTorch 配合使用时,需要保证 CANN 相关包与 PyTorch 相关包的版本匹配。
- 输入 query、key、value 的约束:
- B:batchsize 必须相等;
- D:Head-Dim 必须满足
qD == kD && kD >= vD; - inputLayout 必须一致。
- shape 取值范围:
- T:1~1M;
- N:1~256;
- D:1~768。
7.2 TND 数据排布
query、key、value 数据排布格式仅支持 TND。T 是 B 和 S 合轴紧密排列的数据(每个 batch 的 SeqLenQ 和 SeqLenKV),其中 B(Batch)表示输入样本批量大小、S(Seq-Length)表示输入样本序列长度、H(Head-Size)表示隐藏层的大小、N(Head-Num)表示多头数、D(Head-Dim)表示隐藏层最小的单元尺寸,且满足 D = H / N。
7.3 realShiftOptional(pse)与 pseType
alibi 位置编码压缩(内存优化):如果 Sq 大于 1024,且每个 batch 的 Sq 与 Skv 等长,且是 sparseMode 为 0、2、3 的下三角掩码场景,可开启 alibi 位置编码压缩,此时只需输入原始 PSE 最后 1024 行,即
alibi_compress = ori_pse[:, :, -1024:, :],具体规则如下:- 参数每个 batch 不相同时,shape 为
BNHSkv(H=1024); - 每个 batch 相同时,shape 为
1NHSkv(H=1024); - TND 场景下,每个 batch 段内部仍按
[N, Sq_i, Skv_i]生成,但存储与传参时统一 flatten。若第 i 个 batch 段的真实 query 长度为 Sq_i、真实 key/value 长度为 Skv_i,则该段 PSE 元素个数为N * Sq_i * Skv_i,整段 PSE 总长度pseTotalLen为sum_i(N * Sq_i * Skv_i); - 如果 pseType 为 2 或 3 时,数据类型需为 FLOAT32,对应 shape 支持范围是
[B,N]或[N]; - 如果不开启该参数,realShiftOptional 需要传入 nullptr,pseType 需要传入 1。
- 参数每个 batch 不相同时,shape 为
pseType 各取值含义:
| pseType | 含义 | 备注 |
|---|---|---|
| 0 | 外部传入 pse 先 mul 再 add | - |
| 1 | 外部传入 pse 先 add 再 mul | 跟 aclnnFlashAttentionUnpaddingScoreGrad 实现一致 |
| 2 | 内部生成 pse 先 mul 再 add | - |
| 3 | 内部生成 pse 先 mul 再 add 再 sqrt | - |
- pseType 为 2 或 3 时,当前只支持 Sq 和 Skv 等长。
从源码看,infershape 中的 flash_attention_score_infershape.cpp 对 pseType 2/3 场景做了强校验:内部生成 alibi pse 时realShiftOptional不允许为空指针,且其数据类型必须是 FLOAT(否则返回参数校验错误);tiling 实现 flash_attention_score_tiling_varlen.cpp 中也可看到 pse alibi 场景对 sparse_type 的限制(必须为 CAUSAL 或 RIGHT_DOWN_CAUSAL)。
7.4 innerPrecise
当前 0、1 为保留配置值,2 为开启无效行计算,其功能是避免在计算过程中存在整行 mask 进而导致精度有损失,但该配置会导致性能下降。如果算子可判断出存在无效行场景,会自动开启无效行计算,例如 sparseMode 为 3、Sq > Skv 场景。
7.5 sparseMode 约束
- 当所有的 attenMaskOptional 的 shape 小于 2048 且相同的时候,建议使用 default 模式(0),以减少内存使用量;
- 配置为 1、2、3 时,用户配置的 preTokens、nextTokens 不会生效;
- 配置为 0、4 时,须保证 attenMaskOptional 与 preTokens、nextTokens 的范围一致;
- 用户不特意指定时建议传入 0;
- sparse 不同模式的详细说明请参见 sparseMode 介绍(其中 varlen 场景仅支持 0/1/2/3/4/6/7/8,非压缩 prefix 模式 5 与 treeMask 模式 9 不属于本接口支持范围);
- 配置为 3 时,不支持无效行计算,需要满足每个 batch 的 Sq <= Skv;
- 配置为 7 时,不支持可选输入 realShiftOptional;
- 配置为 8 时,当每个 sequence 的 q、kv 等长时支持可选输入 realShiftOptional,针对全局做 pse 生成;支持 q 方向进行外切,需要外切前每个 sequence 的 q、kv 等长,外切后传入的
actualSeqQLenOptional[0] - actualSeqKvLenOptional[0] + qStartIdxOptional - kvStartIdxOptional == 0(本功能属实验性功能)。
tiling 侧 flash_attention_score_tiling_varlen.cpp 中也存在对 sparseMode 与 preTokens/nextTokens 匹配性的校验(preTokens and nextTokens not match sparseMode),与文档约束一致。
7.6 其他约束
- 部分场景下,如果计算量过大可能会导致算子执行超时(aicore error 类型报错,errorStr 为:
timeout or trap error),此时建议做轴切分处理。这里的计算量受 B、S、N、D 等参数影响,值越大计算量越大。 - band 场景,preTokens 和 nextTokens 之间必须要有交集。
- prefixOptional 稀疏计算场景即 sparseMode=6:当 Sq > Skv 时,prefix 的 N 值取值范围
[0, Skv];当 Sq <= Skv 时,prefix 的 N 值取值范围[Skv-Sq, Skv]。 - actualSeqQLenOptional 输入支持某个 Batch 上的 S 长度为 0,此时不支持可选输入 realShiftOptional;actualSeqQLenOptional 的长度取值范围为 1~2K,当存在 prefixOptional 输入时,其长度最大支持 1K。例如真实的 S 长度为
[2,2,0,2,2],则传入的 actualSeqQLenOptional 为[2,4,4,6,8](累积长度语义,0 长度 batch 使用与前一项相同的累积值)。 - attenMaskOptional 输入不支持补 pad,即 attenMaskOptional 中不能存在某一行全 1 的场景。
八、完整调用示例
以下调用示例代码来自仓库文档,仅供参考,具体编译和执行过程请参考 编译与运行样例。仓库 examples/test_aclnn_flash_attention_score.cpp 亦提供了可参考的样例工程。
#include <iostream> #include <vector> #include "acl/acl.h" #include "aclnnop/aclnn_flash_attention_score.h" #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; } void PrintOutResult(std::vector<int64_t> &shape, void** deviceAddr) { auto size = GetShapeSize(shape); std::vector<float> resultData(size, 0); auto ret = aclrtMemcpy(resultData.data(), resultData.size() * sizeof(resultData[0]), *deviceAddr, 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); for (int64_t i = 0; i < size; i++) { LOG_PRINT("mean result[%ld] is: %f\n", i, resultData[i]); } } 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的接口自定义构造 std::vector<int64_t> qShape = {256, 1, 128}; std::vector<int64_t> kShape = {256, 1, 128}; std::vector<int64_t> vShape = {256, 1, 128}; std::vector<int64_t> attenmaskShape = {256, 256}; std::vector<int64_t> attentionOutShape = {256, 1, 128}; std::vector<int64_t> softmaxMaxShape = {256, 1, 8}; std::vector<int64_t> softmaxSumShape = {256, 1, 8}; void* qDeviceAddr = nullptr; void* kDeviceAddr = nullptr; void* vDeviceAddr = nullptr; void* attenmaskDeviceAddr = nullptr; void* attentionOutDeviceAddr = nullptr; void* softmaxMaxDeviceAddr = nullptr; void* softmaxSumDeviceAddr = nullptr; aclTensor* q = nullptr; aclTensor* k = nullptr; aclTensor* v = nullptr; aclTensor* pse = nullptr; aclTensor* dropMask = nullptr; aclTensor* padding = nullptr; aclTensor* attenmask = nullptr; aclTensor* attentionOut = nullptr; aclTensor* softmaxMax = nullptr; aclTensor* softmaxSum = nullptr; aclTensor* softmaxOut = nullptr; std::vector<float> qHostData(32768, 1); std::vector<float> kHostData(32768, 1); std::vector<float> vHostData(32768, 1); std::vector<uint8_t> attenmaskHostData(65536, 0); std::vector<float> attentionOutHostData(32768, 0); std::vector<float> softmaxMaxHostData(2048, 3.0); std::vector<float> softmaxSumHostData(2048, 3.0); ret = CreateAclTensor(qHostData, qShape, &qDeviceAddr, aclDataType::ACL_FLOAT, &q); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(kHostData, kShape, &kDeviceAddr, aclDataType::ACL_FLOAT, &k); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(vHostData, vShape, &vDeviceAddr, aclDataType::ACL_FLOAT, &v); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(attenmaskHostData, attenmaskShape, &attenmaskDeviceAddr, aclDataType::ACL_UINT8, &attenmask); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(attentionOutHostData, attentionOutShape, &attentionOutDeviceAddr, aclDataType::ACL_FLOAT, &attentionOut); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(softmaxMaxHostData, softmaxMaxShape, &softmaxMaxDeviceAddr, aclDataType::ACL_FLOAT, &softmaxMax); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(softmaxSumHostData, softmaxSumShape, &softmaxSumDeviceAddr, aclDataType::ACL_FLOAT, &softmaxSum); CHECK_RET(ret == ACL_SUCCESS, return ret); std::vector<int64_t> prefixOp = {0}; aclIntArray *prefix = aclCreateIntArray(prefixOp.data(), 1); std::vector<int64_t> qStartIdxOp = {0}; std::vector<int64_t> kvStartIdxOp = {0}; aclIntArray *qStartIdx = aclCreateIntArray(qStartIdxOp.data(), 1); aclIntArray *kvStartIdx = aclCreateIntArray(kvStartIdxOp.data(), 1); std::vector<int64_t> acSeqQLenOp = {256}; std::vector<int64_t> acSeqKvLenOp = {256}; aclIntArray* acSeqQLen = aclCreateIntArray(acSeqQLenOp.data(), acSeqQLenOp.size()); aclIntArray* acSeqKvLen = aclCreateIntArray(acSeqKvLenOp.data(), acSeqKvLenOp.size()); double scaleValue = 0.088388; double keepProb = 1; int64_t preTokens = 65536; int64_t nextTokens = 65536; int64_t headNum = 1; int64_t innerPrecise = 0; int64_t sparseMode = 0; int64_t pseType = 1; char layOut[5] = {'T', 'N', 'D', 0}; // 3. 调用CANN算子库API,需要修改为具体的Api名称 uint64_t workspaceSize = 0; aclOpExecutor* executor; // 调用aclnnFlashAttentionVarLenScoreV2第一段接口 ret = aclnnFlashAttentionVarLenScoreV2GetWorkspaceSize( q, k, v, pse, dropMask, padding, attenmask, prefix, acSeqQLen, acSeqKvLen, qStartIdx, kvStartIdx, scaleValue, keepProb, preTokens, nextTokens, headNum, layOut, innerPrecise, sparseMode, pseType, softmaxMax, softmaxSum, softmaxOut, attentionOut, &workspaceSize, &executor); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnFlashAttentionVarLenScoreV2GetWorkspaceSize 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); } // 调用aclnnFlashAttentionVarLenScoreV2第二段接口 ret = aclnnFlashAttentionVarLenScoreV2(workspaceAddr, workspaceSize, executor, stream); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnFlashAttentionVarLenScoreV2 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的接口定义修改 PrintOutResult(attentionOutShape, &attentionOutDeviceAddr); PrintOutResult(softmaxMaxShape, &softmaxMaxDeviceAddr); PrintOutResult(softmaxSumShape, &softmaxSumDeviceAddr); // 6. 释放aclTensor和aclScalar,需要根据具体API的接口定义修改 aclDestroyTensor(q); aclDestroyTensor(k); aclDestroyTensor(v); aclDestroyTensor(attenmask); aclDestroyTensor(attentionOut); aclDestroyTensor(softmaxMax); aclDestroyTensor(softmaxSum); // 7. 释放device资源 aclrtFree(qDeviceAddr); aclrtFree(kDeviceAddr); aclrtFree(vDeviceAddr); aclrtFree(attenmaskDeviceAddr); aclrtFree(attentionOutDeviceAddr); aclrtFree(softmaxMaxDeviceAddr); aclrtFree(softmaxSumDeviceAddr); if (workspaceSize > 0) { aclrtFree(workspaceAddr); } aclrtDestroyStream(stream); aclrtResetDevice(deviceId); aclFinalize(); return 0; }示例中要点说明:
- 构造 Tensor 时需自行通过
aclrtMalloc/aclrtMemcpy完成 Device 侧内存申请与数据搬运,并用aclCreateTensor按 ND 格式创建 aclTensor; actualSeqQLenOptional/actualSeqKvLenOptional传入的是累积长度(cumulative),示例中单 batch 场景为{256};scaleValue = 0.088388即常见的1/sqrt(D)(D=128 时1/√128≈0.088388)缩放系数;preTokens/nextTokens = 65536配合sparseMode = 0表示不做窗口限制(取较大值覆盖全序列);- 未使用 pse 时
pse传nullptr,对应pseType = 1(不开 alibi 压缩时的标准配置)。
九、源码实现佐证
9.1 接口声明与 op_api 层
接口声明位于 aclnn_flash_attention_score.h,其中aclnnFlashAttentionVarLenScoreV2GetWorkspaceSize与aclnnFlashAttentionVarLenScoreV2的形参顺序与文档原型完全一致,并在注释中标明@domain aclnn_ops_train,即训练域算子。
在 flash_attention_score.cpp 的 L0 实现中,prefixOptional、actualSeqQLenOptional、actualSeqKvLenOptional、qStartIdxOptional、kvStartIdxOptional等aclIntArray可选入参会被转换为 INT64 的 ND Tensor 并统一置为 ND 格式后进入INFER_SHAPE与ADD_TO_LAUNCHER_LIST_AICORE流程;softmaxMaxOut/softmaxSumOut固定按 FLOAT 类型分配,attentionOutOut默认与 query 同 dtype 分配。这一实现细节解释了为什么文档中 softmax 中间输出仅支持 FLOAT 类型。
9.2 算子原型定义
flash_attention_score_def.cpp 定义了 FlashAttentionScore 算子原型,与本文接口直接相关的属性默认值如下:
scale_value:可选,默认 1.0;keep_prob:可选,默认 1.0;pre_tockens/next_tockens:可选,默认 2147483647(int64 最大值,等价于不限窗口);inner_precise:可选,默认 0;sparse_mode:可选,默认 0(defaultMask 模式);pse_type:可选,默认 1(外部 pse 先 add 再 mul)。
该文件还声明了ascend910b/ascend910_93(Atlas A2 训练系列)与ascend950(Ascend 950 系列)的 AICore 配置,与文档“产品支持情况”中的硬件范围一致。
9.3 Host 侧 infershape 与 tiling
flash_attention_score_infershape.cpp 中:
AnalysisAxisForTnd按(T, N, D)解析 TND 排布,并校验headNum == query 的 N 轴;- 校验 query/key/value 三者的数据类型必须一致,pseType 为 2/3(内部 alibi)时 realShift 必须为非空 FLOAT Tensor;
- 对
actualSeqQLen/actualSeqKvLen的长度设置了告警上限(20000 以内),对应文档中“长度取值范围 1~2K”的约束语义。
tiling 侧 flash_attention_score_tiling_varlen.cpp 负责可变长场景下的切分与稀疏模式映射(SparseMode枚举映射),并针对内部 pse alibi 场景限制了 sparseType 必须为 causal 类模式,从底层印证了文档 7.3、7.5 节的约束条款。
9.4 测试与样例资源
仓库为 FlashAttentionScore 算子族提供了系统性的验证资源,可帮助读者进一步理解该接口的正确用法:
- 样例工程:examples/test_aclnn_flash_attention_score.cpp;
- CPU/GPU 参考实现:tests/pytest/cpu_impl.py、tests/pytest/npu_impl.py;
- 单测用例:tests/ut/op_api/test_aclnn_flash_attention_score.cpp 与 host 侧 infershape/inferdatatype 用例(tests/ut/op_host)。
十、与其他 VarLen 版本接口的关系
在 docs 目录下,本接口还对应一系列演进版本,实际选型时可结合需求参考:
| 接口 | 相对 V2 的差异(从接口签名与文档看) |
|---|---|
| aclnnFlashAttentionVarLenScore | V1:无 qStartIdx/kvStartIdx(不支持外切)、无 pseType 属性,pse 固定先 add 再 mul |
| aclnnFlashAttentionVarLenScoreV2 | 本文主题:新增 qStartIdx/kvStartIdx 外切支持与 pseType 属性(0/1/2/3) |
| aclnnFlashAttentionVarLenScoreV3 | 新增 queryRope/keyRope 输入,支持旋转位置编码在算子内完成 |
| aclnnFlashAttentionVarLenScoreV4 | 新增 softmaxOutLayout 参数 |
| aclnnFlashAttentionVarLenScoreV5 | 在 V3 基础上叠加 sink 输入与 softmaxOutLayout,并新增 GetMaxWorkspaceSize 变体 |
十一、常见问题与调优建议
- 可变长 sequence 如何组织输入?TND 排布要求 query/key/value 的第 0 维为所有 sequence 长度总和 T,各 batch 段按顺序紧密排列;
actualSeqQLenOptional/actualSeqKvLenOptional传入的是累积长度而非逐段长度(如两段各 2 与 3,应传[2,5]),0 长度 batch 的累积值沿用前一项。 - 是否需要传 attenMask?当 attenMaskOptional 为 None 时,sparseMode、preTokens、nextTokens 参数不生效,固定为全计算;若需要 causal 等掩码,须按 sparseMode 对应模式传入正确 shape 的掩码矩阵,且不支持补 pad(不能出现某一行全 1)。
- pse 与 pseType 如何搭配?不使用位置编码时 pse 传 nullptr 且 pseType=1;需要 alibi 且 Sq>1024 时可开启压缩(只传最后 1024 行);pseType=2/3 为内部生成 pse,要求 Sq==Skv 且 pse 数据类型为 FLOAT32。
- 精度与性能取舍:
innerPrecise=2开启无效行计算可避免整行 mask 导致的精度损失,但会带来性能下降;默认 0 即可。 - 执行超时处理:B、S、N、D 过大会导致 aicore error(
timeout or trap error),应做轴切分(例如对 S 或 N 进行切块后多次调用)。 - 与 PyTorch 混用:需保证 CANN 相关包与 PyTorch 相关包版本匹配,避免底层接口签名不一致导致的异常。
通过本文对接口原型、参数语义、约束体系、示例代码与源码链路的系统梳理,读者应能独立完成aclnnFlashAttentionVarLenScoreV2的接入与调参,并基于仓库中的 设计介绍 进一步深入 FlashAttention 在 NPU 上的分块(tiling)与稀疏计算原理。
【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考