news 2026/9/19 3:44:26

CANN ops-transformer 算子实战:aclnnFlashAttentionUnpaddingScoreGradV2 可变长 FlashAttention 反向接口详解

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
CANN ops-transformer 算子实战:aclnnFlashAttentionUnpaddingScoreGradV2 可变长 FlashAttention 反向接口详解

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 可以看到,算子注册了ascend950ascend350(对应 A3 系列)与ascend910bascend910_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 反向(借助正向的softmaxMaxsoftmaxSumattentionIn等中间量)与 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、FLOAT32ND[TND]
keyIn输入公式中的 K数据类型与 query/value 一致FLOAT16、BFLOAT16、FLOAT32ND[TND]
value输入公式中的 V数据类型与 query/keyIn 一致FLOAT16、BFLOAT16、FLOAT32ND[TND]
dy输入公式中的 dY-FLOAT16、BFLOAT16、FLOAT32ND[TND]
pseShiftOptional可选输入公式中的 pse数据类型与 query 一致,需与 pseType 配套使用FLOAT16、BFLOAT16、FLOAT32ND[B,N,1024,Skv]、[1,N,1024,Skv]、[B,N]、[N]
dropMaskOptional可选输入Dropout 掩码-UINT8ND0、1
paddingMaskOptional可选输入预留参数,暂未使用调用时需传空----
qStartIdxOptional可选输入外切场景下,当前分块 query 的 sequence 在全局中的起始索引-INT64ND0、1-
kvStartIdxOptional可选输入外切场景下,当前分块 key/value 的 sequence 在全局中的起始索引-INT64ND0、1-
attenMaskOptional可选输入公式中的 atten_mask取值为 1 代表该位不参与计算,为 0 代表该位参与计算BOOL、UINT8ND[B,N,Sq,Skv]、[B,1,Sq,Skv]、[1,1,Sq,Skv]、[Sq,Skv]
softmaxMaxOptional可选输入正向 softmax 的中间输出-FLOATND[N,T,8]
softmaxSumOptional可选输入正向 softmax 的中间输出-FLOATND[N,T,8]
softmaxInOptional可选输入正向 softmax 的中间输出预留参数,暂未使用----
attentionInOptional可选输入正向注意力输出与 query 数据类型、shape 一致FLOAT16、BFLOAT16、FLOAT32ND[TND]
prefixOptional可选输入prefix 稀疏场景每个 Batch 的 N-INT64ND0、1-
actualSeqQLenOptional可选输入实际 Query 序列长度-INT64ND1-
actualSeqKvLenOptional可选输入实际 Key/Value 序列长度-INT64ND1-
scaleValue可选输入scale 缩放系数-DOUBLE---
keepProb可选输入dropMaskOptional 中 1 的比例-DOUBLE---
preTokens可选输入稀疏计算时滑窗左边界-INT64---
nextTokens可选输入稀疏计算时滑窗右边界-INT64---
headNum输入单卡 head 数量,即 Query 的 N 轴长度-INT64---
inputLayout输入输入 Q/K/V 数据排布支持 TNDString---
innerPrecise可选输入内部计算精度控制保留参数,暂未使用INT64---
sparseMode可选输入稀疏模式支持配置 0~8,不支持 5INT64---
pseType可选输入数据类型支持 INT64支持配置值为 0、1、2、3INT64---
dqOut输出公式中的 dQ,Query 梯度-FLOAT16、BFLOAT16、FLOAT32ND[TND]
dkOut输出公式中的 dK,Key 梯度-FLOAT16、BFLOAT16、FLOAT32ND[TND]
dvOut输出公式中的 dV,Value 梯度-FLOAT16、BFLOAT16、FLOAT32ND[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)通过actualSeqQLenOptionalactualSeqKvLenOptional传入各 sequence 的累积长度来区分。
  • pseShiftOptional:必须与pseType配套使用。不开启 alibi 位置编码压缩时需传入nullptrpseType=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_NULLPTR161001传入参数是必选输入、输出或必选属性,且是空指针
ACLNN_ERR_PARAM_INVALID161002query、keyIn、value、dy、pseShiftOptional、dropMaskOptional、paddingMaskOptional、attenMaskOptional、softmaxMaxOptional、softmaxSumOptional、softmaxInOptional、attentionInOptional、dqOut、dkOut、dvOut 的数据类型不在支持的范围内
ACLNN_ERR_PARAM_INVALID161002上述参数的数据格式不在支持的范围内

从 aclnn_flash_attention_score_grad.cpp 源码可以看到,第一段接口内部会执行入参空指针检查、shape 校验(如 D 维是否需要 pad/transpose 预处理),随后通过INFER_SHAPEADD_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_qlenactual_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需传入nullptrpseType需传入 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.0883881/sqrt(128)(D=128 的标准缩放),keepProb=1表示不启用 dropout 丢弃;preTokens/nextTokens=65536为示例中的大滑窗值(配合 sparseMode=0 时实际可视为全量注意力范围)。

八、源码级实现佐证

8.1 算子定义层(op_host)

flash_attention_score_grad_def.cpp 通过OpDef注册了完整的输入/输出/属性集合:

  • querykeyvaluedy为必选输入,支持 FLOAT16/BFLOAT16/FLOAT32(Ascend 950 及 A3 的 AiCore 配置还登记了 FLOAT8/HIFLOAT8 量化输入);
  • pse_shiftdrop_maskatten_masksoftmax_maxsoftmax_sumattention_in等为可选输入,其中prefixactual_seq_qlenactual_seq_kvlenq_start_idxkv_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);
  • prefixactual_seq_**_start_idxaclIntArray转换为 INT64 的 aclTensor 供底层 shape 推导使用;
  • 通过INFER_SHAPEADD_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 训练的关键反向算子接口,核心价值在于:

  1. 可变长支持:通过actualSeqQLenOptional/actualSeqKvLenOptional一次处理长度不等的多个 sequence,天然适配 LLM 训练中 padding 消除后的 unpadding 数据流;
  2. pseType 扩展:0/1/2/3 四种取值覆盖了外部/内部生成位置编码与 mul/add/sqrt 多种融合语义,兼容不同位置编码实现;
  3. 稀疏与长序列外切:sparseMode 0~8(不含 5)与 qStartIdx/kvStartIdx 组合支撑 band、causal、prefix 及多卡 sequence 外切等训练优化手段;
  4. 两段式编程模型:GetWorkspaceSize 负责校验与资源估算,执行段按 workspaceSize 申请内存后即可异步下发。

实际使用时建议:不开启 pse 时显式传pseShiftOptional=nullptrpseType=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),仅供参考

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

.NET 8 vs Java:产线物联网后端的资源确定性选型指南

1. 为什么一个物联网产线项目&#xff0c;会因为选 .NET 还是 Java 被客户当场质疑&#xff1f;我干工业物联网系统集成快八年了&#xff0c;跑过三十多个产线级项目&#xff0c;从汽车焊装车间到食品灌装线&#xff0c;从半导体晶圆厂到中药提取车间。最常被客户运维团队堵在控…

作者头像 李华
网站建设 2026/9/19 3:41:38

人脸识别只是入口:从检测对齐到边缘部署的AI全链路拆解

/* 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 3:39:48

一文讲透信令流程讲义:附着、切换与定时器排障要点

简介&#xff1a;以GSM信令流程为主线的完整讲义PPT&#xff0c;面向通信工程专业学生、网络运维与优化人员&#xff0c;也适合作为企业内训或高校教学的辅助材料&#xff0c;帮助读者系统建立移动核心网信令分析框架。内容从GSM网络拓扑入手&#xff0c;阐释MSC、BSC、BTS、HL…

作者头像 李华
网站建设 2026/9/19 3:39:11

Agent测试框架Harbor:从路由识别到工具调用的全链路实践

Agent项目跑起来容易&#xff0c;测起来是真的烦。你千辛万苦把Agent接上大模型&#xff0c;调通工具调用&#xff0c;结果一改prompt&#xff0c;原来的路由识别直接跑偏&#xff0c;你还没法像传统接口那样用断言一把梭。我最近在做一个内部Agent平台&#xff0c;把大模型接入…

作者头像 李华
网站建设 2026/9/19 3:39:01

2026年前端进阶路线:从基础三件套到微前端与AI应用

2026年再看软件行业&#xff0c;前端早就不是当年那个“改改页面切切图”的岗位了。这一年企业招聘里出现一个挺明显的信号&#xff1a;前端岗位的需求量虽然没爆炸式增长&#xff0c;但对候选人的要求已经截然不同。纯靠Vue或React写几个页面就能拿offer的时代彻底过去了&…

作者头像 李华