CANN ops-transformer 算子实战:aclnnFlashAttentionUnpaddingScoreGradV2 可变长 FlashAttention 反向接口详解
【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer
本文以 CANN ops-transformer 开源仓库中attention/flash_attention_score_grad模块的 API 参考文档为主体,系统讲解aclnnFlashAttentionUnpaddingScoreGradV2两段式接口的数学原理、完整参数语义、pseType 扩展能力与约束边界,并结合仓库源码(op_def 算子定义、op_api 封装、Tiling 与 Kernel 设计)说明其底层实现路径。读者读完可掌握在 Ascend 950PR/A950DT、Atlas A2/A3 训练与推理系列产品上,为变长序列(TND 排布)FlashAttention 训练编写反向梯度计算调用的完整方法。
一、产品支持情况
该接口是训练场景下注意力反向计算的 aclnn 单算子 API,其硬件适配情况如下:
| 产品 | 支持情况 |
|---|---|
| Ascend 950PR / Ascend 950DT | 支持 |
| Atlas A3 训练系列产品 / Atlas A3 推理系列产品 | 支持 |
| Atlas A2 训练系列产品 / Atlas A2 推理系列产品 | 支持 |
| Atlas 200I/500 A2 推理产品 | 不支持 |
| Atlas 推理系列产品(310P) | 不支持 |
| Atlas 训练系列产品(910) | 不支持 |
从仓库的算子定义文件 flash_attention_score_grad_def.cpp 可以看到,算子注册了ascend950、ascend350(对应 A3 系列)与ascend910b、ascend910_93(对应 A2 训练系列)等 AiCore 配置,与上表支持矩阵一一对应。
二、功能说明与计算公式
2.1 接口定位
aclnnFlashAttentionUnpaddingScoreGradV2是 aclnnFlashAttentionVarLenScoreV2(正向接口)的反向计算。与基础版反向接口 aclnnFlashAttentionUnpaddingScoreGrad 相比,本接口的核心增量是新增了pseType参数:
pseType=1时,实现与aclnnFlashAttentionUnpaddingScoreGrad完全一致(先 add 再 mul);pseType取其他值时,位置编码与 QK 得分的融合顺序变为先 mul 再 add。
2.2 正向计算公式
已知注意力的正向计算(以 pseType≠1 为例):
$$ Y=Dropout(Softmax(Mask(\frac{QK^T}{\sqrt{d}}+pse),atten_mask),keep_prob)V $$
为便于表达,引入中间变量 $S$ 与 $P$:
$$ S=Mask(\frac{QK^T}{\sqrt{d}}+pse,atten_mask) $$
$$ P=Dropout(Softmax(S),keep_prob) $$
$$ Y=PV $$
2.3 反向计算公式
注意力的反向计算公式为:
$$ dV=P^TdY $$
$$ dQ=\frac{((dS)*K)}{\sqrt{d}} $$
$$ dK=\frac{((dS)^T*Q)}{\sqrt{d}} $$
其中 $dS$ 由 softmax 反向(借助正向的softmaxMax、softmaxSum、attentionIn等中间量)与 dropout 掩码共同推导得出。
2.4 pseType 语义对照
| pseType | 含义 | 备注 |
|---|---|---|
| 0 | 外部传入 pse,先 mul 再 add | - |
| 1 | 外部传入 pse,先 add 再 mul | 与 FlashAttentionUnpaddingScoreGrad 实现一致 |
| 2 | 内部生成 pse,先 mul 再 add | - |
| 3 | 内部生成 pse,先 mul 再 add 再 sqrt | - |
可以推断,pseType的存在是为了兼容不同位置编码(如 alibi 的乘法式 bias 与加法式 bias)在 attention 得分中的不同融合习惯,从而在算子内部完成 fused 计算,避免框架层拆分多个 kernel。
三、两段式接口与函数原型
与其他 CANN 单算子 API 一样,本算子遵循两段式接口约定:必须先调用aclnnFlashAttentionUnpaddingScoreGradV2GetWorkspaceSize获取计算所需 workspace 大小以及包含算子计算流程的执行器,再调用aclnnFlashAttentionUnpaddingScoreGradV2执行计算。
第一段接口原型:
aclnnStatus aclnnFlashAttentionUnpaddingScoreGradV2GetWorkspaceSize( const aclTensor *query, const aclTensor *keyIn, const aclTensor *value, const aclTensor *dy, const aclTensor *pseShiftOptional, const aclTensor *dropMaskOptional, const aclTensor *paddingMaskOptional, const aclTensor *attenMaskOptional, const aclTensor *softmaxMaxOptional, const aclTensor *softmaxSumOptional, const aclTensor *softmaxInOptional, const aclTensor *attentionInOptional, 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 *dqOut, const aclTensor *dkOut, const aclTensor *dvOut, const aclTensor *dpseOut, uint64_t *workspaceSize, aclOpExecutor **executor)第二段接口原型:
aclnnStatus aclnnFlashAttentionUnpaddingScoreGradV2( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, const aclrtStream stream)第一段接口完成入参校验、shape 推导并返回 workspace 大小;第二段接口在指定 stream 上真正下发执行。workspace 是算子在 NPU 上完成计算所需的临时内存(不含输入/输出本身),其大小必须由第一段接口计算得出。
四、aclnnFlashAttentionUnpaddingScoreGradV2GetWorkspaceSize 参数详解
4.1 参数表
下表完整列出第一段接口的全部参数(信息取自 aclnnFlashAttentionUnpaddingScoreGradV2.md):
| 参数名 | 输入/输出 | 描述 | 使用说明 | 数据类型 | 数据格式 | 维度(shape) | 非连续Tensor |
|---|---|---|---|---|---|---|---|
| query | 输入 | 公式中的 Q | 数据类型与 keyIn/value 一致 | FLOAT16、BFLOAT16、FLOAT32 | ND | [TND] | √ |
| keyIn | 输入 | 公式中的 K | 数据类型与 query/value 一致 | FLOAT16、BFLOAT16、FLOAT32 | ND | [TND] | √ |
| value | 输入 | 公式中的 V | 数据类型与 query/keyIn 一致 | FLOAT16、BFLOAT16、FLOAT32 | ND | [TND] | √ |
| dy | 输入 | 公式中的 dY | - | FLOAT16、BFLOAT16、FLOAT32 | ND | [TND] | √ |
| pseShiftOptional | 可选输入 | 公式中的 pse | 数据类型与 query 一致,需与 pseType 配套使用 | FLOAT16、BFLOAT16、FLOAT32 | ND | [B,N,1024,Skv]、[1,N,1024,Skv]、[B,N]、[N] | √ |
| dropMaskOptional | 可选输入 | Dropout 掩码 | - | UINT8 | ND | 0、1 | √ |
| paddingMaskOptional | 可选输入 | 预留参数,暂未使用 | 调用时需传空 | - | - | - | - |
| qStartIdxOptional | 可选输入 | 外切场景下,当前分块 query 的 sequence 在全局中的起始索引 | - | INT64 | ND | 0、1 | - |
| kvStartIdxOptional | 可选输入 | 外切场景下,当前分块 key/value 的 sequence 在全局中的起始索引 | - | INT64 | ND | 0、1 | - |
| attenMaskOptional | 可选输入 | 公式中的 atten_mask | 取值为 1 代表该位不参与计算,为 0 代表该位参与计算 | BOOL、UINT8 | ND | [B,N,Sq,Skv]、[B,1,Sq,Skv]、[1,1,Sq,Skv]、[Sq,Skv] | √ |
| softmaxMaxOptional | 可选输入 | 正向 softmax 的中间输出 | - | FLOAT | ND | [N,T,8] | √ |
| softmaxSumOptional | 可选输入 | 正向 softmax 的中间输出 | - | FLOAT | ND | [N,T,8] | √ |
| softmaxInOptional | 可选输入 | 正向 softmax 的中间输出 | 预留参数,暂未使用 | - | - | - | - |
| attentionInOptional | 可选输入 | 正向注意力输出 | 与 query 数据类型、shape 一致 | FLOAT16、BFLOAT16、FLOAT32 | ND | [TND] | √ |
| prefixOptional | 可选输入 | prefix 稀疏场景每个 Batch 的 N | - | INT64 | ND | 0、1 | - |
| actualSeqQLenOptional | 可选输入 | 实际 Query 序列长度 | - | INT64 | ND | 1 | - |
| actualSeqKvLenOptional | 可选输入 | 实际 Key/Value 序列长度 | - | INT64 | ND | 1 | - |
| scaleValue | 可选输入 | scale 缩放系数 | - | DOUBLE | - | - | - |
| keepProb | 可选输入 | dropMaskOptional 中 1 的比例 | - | DOUBLE | - | - | - |
| preTokens | 可选输入 | 稀疏计算时滑窗左边界 | - | INT64 | - | - | - |
| nextTokens | 可选输入 | 稀疏计算时滑窗右边界 | - | INT64 | - | - | - |
| headNum | 输入 | 单卡 head 数量,即 Query 的 N 轴长度 | - | INT64 | - | - | - |
| inputLayout | 输入 | 输入 Q/K/V 数据排布 | 支持 TND | String | - | - | - |
| innerPrecise | 可选输入 | 内部计算精度控制 | 保留参数,暂未使用 | INT64 | - | - | - |
| sparseMode | 可选输入 | 稀疏模式 | 支持配置 0~8,不支持 5 | INT64 | - | - | - |
| pseType | 可选输入 | 数据类型支持 INT64 | 支持配置值为 0、1、2、3 | INT64 | - | - | - |
| dqOut | 输出 | 公式中的 dQ,Query 梯度 | - | FLOAT16、BFLOAT16、FLOAT32 | ND | [TND] | √ |
| dkOut | 输出 | 公式中的 dK,Key 梯度 | - | FLOAT16、BFLOAT16、FLOAT32 | ND | [TND] | √ |
| dvOut | 输出 | 公式中的 dV,Value 梯度 | - | FLOAT16、BFLOAT16、FLOAT32 | ND | [TND] | √ |
| dpseOut | 输出 | d(pse) 梯度 | 预留参数,暂未使用 | - | - | - | - |
| workspaceSize | 输出 | 返回需要在 Device 侧申请的 workspace 大小 | - | - | - | - | - |
| executor | 输出 | 返回 op 执行器,包含算子计算流程 | - | - | - | - | - |
4.2 关键参数解读
- TND 排布:本接口仅支持
inputLayout="TND"。T 是 B 与 S 合轴后的总 token 数(每个 batch 的 SeqLenQ 与 SeqLenKV 紧密排列),N 为多头数,D 为 Head-Dim。这一点与正向接口 aclnnFlashAttentionVarLenScoreV2 一致,可变长序列(一次传入多个长度不等的 sequence)通过actualSeqQLenOptional与actualSeqKvLenOptional传入各 sequence 的累积长度来区分。 - pseShiftOptional:必须与
pseType配套使用。不开启 alibi 位置编码压缩时需传入nullptr且pseType=1(见下文约束)。 - softmaxMax / softmaxSum:由正向 FlashAttention 计算产出的中间量(shape 为 [N,T,8],TND 场景下为 [T,N,8]),反向计算借助它们完成 softmax 梯度推导,无需重算完整 softmax。
- qStartIdx / kvStartIdx:服务于 varlen 长序列外切(sequence parallel 多卡切分)场景,指明当前分块在全局序列中的起始位置。
4.3 返回值与错误码
两段接口均返回aclnnStatus状态码,具体可参见 aclnn返回码。第一段接口完成入参校验,出现以下场景时报错:
| 返回值 | 错误码 | 描述 |
|---|---|---|
| ACLNN_ERR_PARAM_NULLPTR | 161001 | 传入参数是必选输入、输出或必选属性,且是空指针 |
| ACLNN_ERR_PARAM_INVALID | 161002 | query、keyIn、value、dy、pseShiftOptional、dropMaskOptional、paddingMaskOptional、attenMaskOptional、softmaxMaxOptional、softmaxSumOptional、softmaxInOptional、attentionInOptional、dqOut、dkOut、dvOut 的数据类型不在支持的范围内 |
| ACLNN_ERR_PARAM_INVALID | 161002 | 上述参数的数据格式不在支持的范围内 |
从 aclnn_flash_attention_score_grad.cpp 源码可以看到,第一段接口内部会执行入参空指针检查、shape 校验(如 D 维是否需要 pad/transpose 预处理),随后通过INFER_SHAPE与ADD_TO_LAUNCHER_LIST_AICORE完成 shape 推导与 kernel 下发准备。
五、aclnnFlashAttentionUnpaddingScoreGradV2 参数说明
第二段接口参数较少,语义如下:
| 参数名 | 输入/输出 | 描述 |
|---|---|---|
| workspace | 输入 | 在 Device 侧申请的 workspace 内存地址 |
| workspaceSize | 输入 | 在 Device 侧申请的 workspace 大小,由第一段接口获取 |
| executor | 输入 | op 执行器,包含算子计算流程 |
| stream | 输入 | 指定执行任务的 Stream |
注意:第二段接口不能重复调用(同一 executor 只可执行一次),参见两段式接口说明。
六、约束说明
以下约束来自原文档并结合源码补充,是正确使用该接口的关键:
6.1 确定性计算
- 本接口默认非确定性实现,支持通过
aclrtCtxSetSysParamOpt开启确定性计算。关于确定性计算的详细机制可参考 determinism_compute.md。从FAG算子设计介绍可见,确定性计算模板通过特定分核方式避免多核同地址累加来保证结果可复现。
6.2 版本与 dtype 约束
- 与 PyTorch 配合使用时,需保证 CANN 相关包与 PyTorch 相关包版本匹配。
- query、key、value、dy 的 B(batchsize)必须相等,inputLayout 必须一致。
- Head-Dim 需满足
qD == kD && kD >= vD。 - query、key、value、pseShiftOptional 的数据类型必须一致。
- key/value 的 shape 除 D 外必须一致;在 query/key/value 的 D 大小相同的情况下,query/dy 的 shape 必须一致。
- 支持 query 的 N 与 key/value 的 N 不相等,但必须成比例,即
Nq/Nkv必须是非 0 整数,Nq 取值范围 1~256。
6.3 shape 取值范围(TND 场景)
- T:1 ~ 1M
- N:1 ~ 256
- D:1 ~ 768
- KeepProb:(0, 1]
6.4 TND 与 actual_seq 语义
- TND 格式下,支持尾部部分 Batch 不参与计算:此时
actual_seq_qlen和actual_seq_kvlen尾部传入对应个数个 0 即可。假设真实 S 长度为 [2, 3, 4, 5, 6],后两个 Batch 不参与计算,则传入的 actual_seq_qlen 为 [2, 5, 9, 0, 0]。 actualSeqQLenOptional的长度取值范围为 1~2K;当存在prefixOptional输入时,长度最大支持 1K。actualSeqQLenOptional支持某个 Batch 上的 S 长度为 0,此时不支持可选输入 pseShiftOptional。
6.5 pseShiftOptional 与 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];
- 如果不开启该参数,
pseShiftOptional需传入nullptr,pseType需传入 1。
6.6 sparseMode 约束
sparseMode 支持配置 0~8,不支持 5,具体约束如下:
- 当所有
attenMaskOptional的 shape 小于 2048 且相同时,建议使用 default 模式(0),减少内存使用量; - 配置为 1、2、3 时,用户配置的 preTokens、nextTokens 不会生效;
- 配置为 0、4 时,须保证 attenMaskOptional 与 preTokens、nextTokens 的范围一致;
- 用户不特意指定时建议传入 0;
- 配置为 7 时,不支持可选输入 pseShiftOptional;
- 配置为 8 时,当每个 sequence 的 q、kv 等长时支持可选输入 pseShiftOptional(针对全局做 pse 生成);支持 q 方向外切,需要外切前每个 sequence 的 q、kv 等长,外切后需满足
actualSeqQLenOptional[0] - actualSeqKvLenOptional[0] + qStartIdxOptional - kvStartIdxOptional == 0(实验性功能)。
各稀疏模式(defaultMask、allMask、leftUpCausal、rightDownCausal、band、prefix 压缩/非压缩、varlen 外切等)的完整说明见 sparse模式说明。
6.7 其他约束
- prefixOptional 稀疏计算(sparseMode=6)仅支持压缩场景:当 Sq > Skv 时,prefix 的 N 值取值范围 [0, Skv];当 Sq <= Skv 时,取值范围 [Skv-Sq, Skv]。
- softmaxMax 与 softmaxSum输入格式固定为 [B, N, S, 8];TND 场景除外,此时为 [T, N, 8](T = B*S)。
- headNum的取值必须和传入 Query 中的 N 值保持一致。
- 部分场景下计算量过大可能导致算子执行超时(aicore error 类型报错,errorStr 为
timeout or trap error),此时建议做轴切分处理。计算量受 B、S、N、D 等参数影响,值越大计算量越大。
七、调用示例
以下完整示例来自原文档,展示了两段式接口的标准调用流程(资源初始化 → 构造输入输出 → 第一段接口 → 申请 workspace → 第二段接口 → 同步 → 结果回拷 → 资源释放)。仓库中对应的可编译样例可参考 test_aclnn_flash_attention_unpadding_score_grad.cpp,具体编译与执行流程参见编译与运行样例。
#include <iostream> #include <vector> #include "acl/acl.h" #include "aclnnop/aclnn_flash_attention_score_grad.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> dxShape = {256, 1, 128}; std::vector<int64_t> attenmaskShape = {256, 256}; std::vector<int64_t> softmaxMaxShape = {256, 1, 8}; std::vector<int64_t> softmaxSumShape = {256, 1, 8}; std::vector<int64_t> attentionInShape = {256, 1, 128}; std::vector<int64_t> dqShape = {256, 1, 128}; std::vector<int64_t> dkShape = {256, 1, 128}; std::vector<int64_t> dvShape = {256, 1, 128}; void* qDeviceAddr = nullptr; void* kDeviceAddr = nullptr; void* vDeviceAddr = nullptr; void* dxDeviceAddr = nullptr; void* attenmaskDeviceAddr = nullptr; void* softmaxMaxDeviceAddr = nullptr; void* softmaxSumDeviceAddr = nullptr; void* attentionInDeviceAddr = nullptr; void* dqDeviceAddr = nullptr; void* dkDeviceAddr = nullptr; void* dvDeviceAddr = nullptr; aclTensor* q = nullptr; aclTensor* k = nullptr; aclTensor* v = nullptr; aclTensor* dx = nullptr; aclTensor* pse = nullptr; aclTensor* dropMask = nullptr; aclTensor* padding = nullptr; aclTensor* attenmask = nullptr; aclTensor* softmaxMax = nullptr; aclTensor* softmaxSum = nullptr; aclTensor* softmaxIn = nullptr; aclTensor* attentionIn = nullptr; aclTensor* dq = nullptr; aclTensor* dk = nullptr; aclTensor* dv = nullptr; aclTensor* dpse = nullptr; std::vector<float> qHostData(32768, 1); std::vector<float> kHostData(32768, 1); std::vector<float> vHostData(32768, 1); std::vector<float> dxHostData(32768, 1); std::vector<uint8_t> attenmaskHostData(65536, 0); std::vector<float> softmaxMaxHostData(2048, 3.0); std::vector<float> softmaxSumHostData(2048, 3.0); std::vector<float> attentionInHostData(32768, 1); std::vector<float> dqHostData(32768, 0); std::vector<float> dkHostData(32768, 0); std::vector<float> dvHostData(32768, 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(dxHostData, dxShape, &dxDeviceAddr, aclDataType::ACL_FLOAT, &dx); 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(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); ret = CreateAclTensor(attentionInHostData, attentionInShape, &attentionInDeviceAddr, aclDataType::ACL_FLOAT, &attentionIn); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(dqHostData, dqShape, &dqDeviceAddr, aclDataType::ACL_FLOAT, &dq); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(dkHostData, dkShape, &dkDeviceAddr, aclDataType::ACL_FLOAT, &dk); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(dvHostData, dvShape, &dvDeviceAddr, aclDataType::ACL_FLOAT, &dv); CHECK_RET(ret == ACL_SUCCESS, return ret); std::vector<int64_t> prefixOp = {0}; aclIntArray* prefix = aclCreateIntArray(prefixOp.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()); std::vector<int64_t> qStartIdxOp = {0}; std::vector<int64_t> kvStartIdxOp = {0}; aclIntArray *qStartIdx = aclCreateIntArray(qStartIdxOp.data(), 1); aclIntArray *kvStartIdx = aclCreateIntArray(kvStartIdxOp.data(), 1); 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; // 调用aclnnFlashAttentionUnpaddingScoreGradV2第一段接口 ret = aclnnFlashAttentionUnpaddingScoreGradV2GetWorkspaceSize(q, k, v, dx, pse, dropMask, padding, attenmask, softmaxMax, softmaxSum, softmaxIn, attentionIn, prefix, acSeqQLen, acSeqKvLen, qStartIdx, kvStartIdx, scaleValue, keepProb, preTokens, nextTokens, headNum, layOut, innerPrecise, sparseMode, pseType, dq, dk, dv, dpse, &workspaceSize, &executor); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnFlashAttentionUnpaddingScoreGradV2GetWorkspaceSize 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); } // 调用aclnnFlashAttentionUnpaddingScoreGradV2第二段接口 ret = aclnnFlashAttentionUnpaddingScoreGradV2(workspaceAddr, workspaceSize, executor, stream); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnFlashAttentionUnpaddingScoreGradV2 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(dqShape, &dqDeviceAddr); PrintOutResult(dkShape, &dkDeviceAddr); PrintOutResult(dvShape, &dvDeviceAddr); // 6. 释放aclTensor和aclScalar,需要根据具体API的接口定义修改 aclDestroyTensor(q); aclDestroyTensor(k); aclDestroyTensor(v); aclDestroyTensor(dx); aclDestroyTensor(attenmask); aclDestroyTensor(softmaxMax); aclDestroyTensor(softmaxSum); aclDestroyTensor(attentionIn); aclDestroyTensor(dq); aclDestroyTensor(dk); aclDestroyTensor(dv); // 7. 释放device资源 aclrtFree(qDeviceAddr); aclrtFree(kDeviceAddr); aclrtFree(vDeviceAddr); aclrtFree(dxDeviceAddr); aclrtFree(attenmaskDeviceAddr); aclrtFree(softmaxMaxDeviceAddr); aclrtFree(softmaxSumDeviceAddr); aclrtFree(attentionInDeviceAddr); aclrtFree(dqDeviceAddr); aclrtFree(dkDeviceAddr); aclrtFree(dvDeviceAddr); if (workspaceSize > 0) { aclrtFree(workspaceAddr); } aclrtDestroyStream(stream); aclrtResetDevice(deviceId); aclFinalize(); return 0; }示例中scaleValue=0.088388即1/sqrt(128)(D=128 的标准缩放),keepProb=1表示不启用 dropout 丢弃;preTokens/nextTokens=65536为示例中的大滑窗值(配合 sparseMode=0 时实际可视为全量注意力范围)。
八、源码级实现佐证
8.1 算子定义层(op_host)
flash_attention_score_grad_def.cpp 通过OpDef注册了完整的输入/输出/属性集合:
query、key、value、dy为必选输入,支持 FLOAT16/BFLOAT16/FLOAT32(Ascend 950 及 A3 的 AiCore 配置还登记了 FLOAT8/HIFLOAT8 量化输入);pse_shift、drop_mask、atten_mask、softmax_max、softmax_sum、attention_in等为可选输入,其中prefix、actual_seq_qlen、actual_seq_kvlen、q_start_idx、kv_start_idx均标记ValueDepend(OPTIONAL)——这意味着其取值会参与 shape 推导与 tiling 计算,是可变长与稀疏场景的关键信息;- 属性
scale_value(默认 1.0)、keep_prob(默认 1.0)、pre_tockens(默认 INT_MAX)、next_tockens(默认 INT_MAX)、head_num(必填)、input_layout(必填)、inner_precise(默认 0)、sparse_mode(默认 0)、pse_type(默认 1)等与 aclnn 接口参数一一对应。
8.2 API 封装层(op_api)
aclnn_flash_attention_score_grad.cpp 中实现了入参校验与预处理逻辑,例如:
- 对 D 维是否为 192、72、88 等特殊 Head-Dim 判断是否需要 pad 或 transpose 预处理(
CheckIsNeedPad); - 将
prefix、actual_seq_*、*_start_idx等aclIntArray转换为 INT64 的 aclTensor 供底层 shape 推导使用; - 通过
INFER_SHAPE与ADD_TO_LAUNCHER_LIST_AICORE完成推导与 kernel 下发。
l0op 层实现 则负责分配 dqOut/dkOut/dvOut/dpseOut 等输出 tensor,并在输出为 FP8 时按outDType转为 FLOAT16/BF16。
8.3 Tiling 与 Kernel 层
从 FAG算子设计介绍 可了解底层实现脉络:该算子按 FlashAttention 反向流程分为六个计算阶段——重计算 p → 计算 dp → 计算 ds → 计算 dq → 计算 dk → 计算 dv;在 NPU 上通过 Cube(AIC)与 Vector(AIV)分离的架构并行执行,依据 shape 特征路由到 B 模板、N2 模板、SameAB 模板、S1S2 模板、TND 模板(A2 系列)或 BN2、BN2S2、BN2GS1S2、确定性计算模板(950 系列)等不同 tiling 模板,并配套 double buffer 与 CV 流水设计。TND 场景对应arch22/flash_attention_score_grad_tiling_unpadded_attension.cpp等 tiling 文件与 op_kernel 下的 TND 模板 kernel。
8.4 测试与样例
仓库中提供了多种可参考的验证样例,例如 test_aclnn_flash_attention_unpadding_score_grad.cpp,以及 tests 下的 ut/st/pytest 用例,可用于对照验证本接口在不同 shape、sparseMode 与 pseType 组合下的行为。
九、总结与使用建议
aclnnFlashAttentionUnpaddingScoreGradV2是 CANN ops-transformer 面向可变长序列(TND 排布)FlashAttention 训练的关键反向算子接口,核心价值在于:
- 可变长支持:通过
actualSeqQLenOptional/actualSeqKvLenOptional一次处理长度不等的多个 sequence,天然适配 LLM 训练中 padding 消除后的 unpadding 数据流; - pseType 扩展:0/1/2/3 四种取值覆盖了外部/内部生成位置编码与 mul/add/sqrt 多种融合语义,兼容不同位置编码实现;
- 稀疏与长序列外切:sparseMode 0~8(不含 5)与 qStartIdx/kvStartIdx 组合支撑 band、causal、prefix 及多卡 sequence 外切等训练优化手段;
- 两段式编程模型:GetWorkspaceSize 负责校验与资源估算,执行段按 workspaceSize 申请内存后即可异步下发。
实际使用时建议:不开启 pse 时显式传pseShiftOptional=nullptr且pseType=1;稀疏场景优先sparseMode=0并保持 attenMask 与 preTokens/nextTokens 范围一致;开启确定性计算需配合aclrtCtxSetSysParamOpt;遇到 aicore timeout 时优先考虑对 B/S/N/D 轴做切分。
【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考