CANN ops-transformer 中 aclnnFlashAttentionVarLenScoreV2 接口实战:基于 TND 排布的可变长 FlashAttention 训练加速
2026/9/19 0:53:45 网站建设 项目流程

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 总和,同时使用actualSeqQLenOptionalactualSeqKvLenOptional传入每个 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、FLOAT32ND[TND]
key输入公式中的 key数据类型与 query/value 一致FLOAT16、BFLOAT16、FLOAT32ND[TND]
value输入公式中的 value数据类型与 query/key 一致FLOAT16、BFLOAT16、FLOAT32ND[TND]
realShiftOptional可选输入公式中的 pse数据类型与 query 一致,需与 pseType 配套使用FLOAT16、BFLOAT16、FLOAT32ND[B,N,1024,Skv]、[1,N,1024,Skv]、[pseTotalLen]、[B,N]、[N]
dropMaskOptional可选输入公式中的 Dropout-UINT8ND0、1
paddingMaskOptional可选输入预留参数,暂未使用-----
attenMaskOptional可选输入公式中的 atten_mask取值为 1 代表该位不参与计算,为 0 代表该位参与计算BOOL、UINT8ND[B,N,Sq,Skv]、[B,1,Sq,Skv]、[1,1,Sq,Skv]、[Sq,Skv]
prefixOptional可选输入代表 prefix 稀疏计算场景每个 Batch 的 N 值-INT64ND0、1-
actualSeqQLenOptional可选输入描述每个 Batch 对应的 query 的 sequence length-INT64ND0、1-
actualSeqKvLenOptional可选输入描述每个 Batch 对应的 key/value 的 sequence length-INT64ND0、1-
qStartIdxOptional可选输入代表外切场景,当前分块的 query 的 sequence 在全局中的起始索引-INT64ND0、1-
kvStartIdxOptional可选输入代表外切场景,当前分块的 key 和 value 的 sequence 在全局中的起始索引-INT64ND0、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 的数据排布格式支持 TNDString---
innerPrecise可选输入用于提升精度默认配置为 0 即可INT64---
sparseMode可选输入表示 sparse 的模式支持配置值为 0、1、2、3、4、6、7、8INT64---
pseType可选输入控制 mul 与 add 计算顺序支持配置值为 0、1、2、3INT64---
softmaxMaxOut输出Softmax 计算的 Max 中间结果,用于反向计算-FLOATND[N,T,8]
softmaxSumOut输出Softmax 计算的 Sum 中间结果,用于反向计算-FLOATND[N,T,8]
softmaxOutOut输出预留参数,暂未使用-----
attentionOutOut输出计算公式的最终输出数据类型和 shape 类型与 query 保持一致FLOAT16、BFLOAT16、FLOAT32ND[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_NULLPTR161001传入参数是必选输入、输出或者必选属性,且是空指针
ACLNN_ERR_PARAM_INVALID161002query、key、value、realShiftOptional、dropMaskOptional、paddingMaskOptional、attenMaskOptional、softmaxMaxOut、softmaxSumOut、softmaxOutOut、attentionOutOut 的数据类型不在支持的范围内
ACLNN_ERR_PARAM_INVALID161002query、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 总长度pseTotalLensum_i(N * Sq_i * Skv_i)
    • 如果 pseType 为 2 或 3 时,数据类型需为 FLOAT32,对应 shape 支持范围是[B,N][N]
    • 如果不开启该参数,realShiftOptional 需要传入 nullptr,pseType 需要传入 1。
  • 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 时psenullptr,对应pseType = 1(不开 alibi 压缩时的标准配置)。

九、源码实现佐证

9.1 接口声明与 op_api 层

接口声明位于 aclnn_flash_attention_score.h,其中aclnnFlashAttentionVarLenScoreV2GetWorkspaceSizeaclnnFlashAttentionVarLenScoreV2的形参顺序与文档原型完全一致,并在注释中标明@domain aclnn_ops_train,即训练域算子。

在 flash_attention_score.cpp 的 L0 实现中,prefixOptionalactualSeqQLenOptionalactualSeqKvLenOptionalqStartIdxOptionalkvStartIdxOptionalaclIntArray可选入参会被转换为 INT64 的 ND Tensor 并统一置为 ND 格式后进入INFER_SHAPEADD_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 的差异(从接口签名与文档看)
aclnnFlashAttentionVarLenScoreV1:无 qStartIdx/kvStartIdx(不支持外切)、无 pseType 属性,pse 固定先 add 再 mul
aclnnFlashAttentionVarLenScoreV2本文主题:新增 qStartIdx/kvStartIdx 外切支持与 pseType 属性(0/1/2/3)
aclnnFlashAttentionVarLenScoreV3新增 queryRope/keyRope 输入,支持旋转位置编码在算子内完成
aclnnFlashAttentionVarLenScoreV4新增 softmaxOutLayout 参数
aclnnFlashAttentionVarLenScoreV5在 V3 基础上叠加 sink 输入与 softmaxOutLayout,并新增 GetMaxWorkspaceSize 变体

十一、常见问题与调优建议

  1. 可变长 sequence 如何组织输入?TND 排布要求 query/key/value 的第 0 维为所有 sequence 长度总和 T,各 batch 段按顺序紧密排列;actualSeqQLenOptional/actualSeqKvLenOptional传入的是累积长度而非逐段长度(如两段各 2 与 3,应传[2,5]),0 长度 batch 的累积值沿用前一项。
  2. 是否需要传 attenMask?当 attenMaskOptional 为 None 时,sparseMode、preTokens、nextTokens 参数不生效,固定为全计算;若需要 causal 等掩码,须按 sparseMode 对应模式传入正确 shape 的掩码矩阵,且不支持补 pad(不能出现某一行全 1)。
  3. pse 与 pseType 如何搭配?不使用位置编码时 pse 传 nullptr 且 pseType=1;需要 alibi 且 Sq>1024 时可开启压缩(只传最后 1024 行);pseType=2/3 为内部生成 pse,要求 Sq==Skv 且 pse 数据类型为 FLOAT32。
  4. 精度与性能取舍innerPrecise=2开启无效行计算可避免整行 mask 导致的精度损失,但会带来性能下降;默认 0 即可。
  5. 执行超时处理:B、S、N、D 过大会导致 aicore error(timeout or trap error),应做轴切分(例如对 S 或 N 进行切块后多次调用)。
  6. 与 PyTorch 混用:需保证 CANN 相关包与 PyTorch 相关包版本匹配,避免底层接口签名不一致导致的异常。

通过本文对接口原型、参数语义、约束体系、示例代码与源码链路的系统梳理,读者应能独立完成aclnnFlashAttentionVarLenScoreV2的接入与调参,并基于仓库中的 设计介绍 进一步深入 FlashAttention 在 NPU 上的分块(tiling)与稀疏计算原理。

【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询