ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

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

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 可以看到算子注册了ascend950、ascend350对应 A3 系列与ascend910b、ascend910_93对应 A2 训练系列等 AiCore 配置与上表支持矩阵一一对应。二、功能说明与计算公式2.1 接口定位aclnnFlashAttentionUnpaddingScoreGradV2是 aclnnFlashAttentionVarLenScoreV2正向接口的反向计算。与基础版反向接口 aclnnFlashAttentionUnpaddingScoreGrad 相比本接口的核心增量是新增了pseType参数pseType1时实现与aclnnFlashAttentionUnpaddingScoreGrad完全一致先 add 再 mulpseType取其他值时位置编码与 QK 得分的融合顺序变为先 mul 再 add。2.2 正向计算公式已知注意力的正向计算以 pseType≠1 为例$$ YDropout(Softmax(Mask(\frac{QK^T}{\sqrt{d}}pse),atten_mask),keep_prob)V $$为便于表达引入中间变量 $S$ 与 $P$$$ SMask(\frac{QK^T}{\sqrt{d}}pse,atten_mask) $$$$ PDropout(Softmax(S),keep_prob) $$$$ YPV $$2.3 反向计算公式注意力的反向计算公式为$$ dVP^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)非连续Tensorquery输入公式中的 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输出公式中的 dQQuery 梯度-FLOAT16、BFLOAT16、FLOAT32ND[TND]√dkOut输出公式中的 dKKey 梯度-FLOAT16、BFLOAT16、FLOAT32ND[TND]√dvOut输出公式中的 dVValue 梯度-FLOAT16、BFLOAT16、FLOAT32ND[TND]√dpseOut输出d(pse) 梯度预留参数暂未使用----workspaceSize输出返回需要在 Device 侧申请的 workspace 大小-----executor输出返回 op 执行器包含算子计算流程-----4.2 关键参数解读TND 排布本接口仅支持inputLayoutTND。T 是 B 与 S 合轴后的总 token 数每个 batch 的 SeqLenQ 与 SeqLenKV 紧密排列N 为多头数D 为 Head-Dim。这一点与正向接口 aclnnFlashAttentionVarLenScoreV2 一致可变长序列一次传入多个长度不等的 sequence通过actualSeqQLenOptional与actualSeqKvLenOptional传入各 sequence 的累积长度来区分。pseShiftOptional必须与pseType配套使用。不开启 alibi 位置编码压缩时需传入nullptr且pseType1见下文约束。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_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 的 Bbatchsize必须相等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 场景T1 ~ 1MN1 ~ 256D1 ~ 768KeepProb(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 为BNHSkvH1024每个 batch 相同时shape 为1NHSkvH1024TND 场景下每个 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 稀疏计算sparseMode6仅支持压缩场景当 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::vectorint64_t shape) { int64_t shapeSize 1; for (auto i : shape) { shapeSize * i; } return shapeSize; } void PrintOutResult(std::vectorint64_t shape, void** deviceAddr) { auto size GetShapeSize(shape); std::vectorfloat 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::vectorT hostData, const std::vectorint64_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::vectorint64_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::vectorint64_t qShape {256, 1, 128}; std::vectorint64_t kShape {256, 1, 128}; std::vectorint64_t vShape {256, 1, 128}; std::vectorint64_t dxShape {256, 1, 128}; std::vectorint64_t attenmaskShape {256, 256}; std::vectorint64_t softmaxMaxShape {256, 1, 8}; std::vectorint64_t softmaxSumShape {256, 1, 8}; std::vectorint64_t attentionInShape {256, 1, 128}; std::vectorint64_t dqShape {256, 1, 128}; std::vectorint64_t dkShape {256, 1, 128}; std::vectorint64_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::vectorfloat qHostData(32768, 1); std::vectorfloat kHostData(32768, 1); std::vectorfloat vHostData(32768, 1); std::vectorfloat dxHostData(32768, 1); std::vectoruint8_t attenmaskHostData(65536, 0); std::vectorfloat softmaxMaxHostData(2048, 3.0); std::vectorfloat softmaxSumHostData(2048, 3.0); std::vectorfloat attentionInHostData(32768, 1); std::vectorfloat dqHostData(32768, 0); std::vectorfloat dkHostData(32768, 0); std::vectorfloat 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::vectorint64_t prefixOp {0}; aclIntArray* prefix aclCreateIntArray(prefixOp.data(), 1); std::vectorint64_t acSeqQLenOp {256}; std::vectorint64_t acSeqKvLenOp {256}; aclIntArray* acSeqQLen aclCreateIntArray(acSeqQLenOp.data(), acSeqQLenOp.size()); aclIntArray* acSeqKvLen aclCreateIntArray(acSeqKvLenOp.data(), acSeqKvLenOp.size()); std::vectorint64_t qStartIdxOp {0}; std::vectorint64_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; }示例中scaleValue0.088388即1/sqrt(128)D128 的标准缩放keepProb1表示不启用 dropout 丢弃preTokens/nextTokens65536为示例中的大滑窗值配合 sparseMode0 时实际可视为全量注意力范围。八、源码级实现佐证8.1 算子定义层op_hostflash_attention_score_grad_def.cpp 通过OpDef注册了完整的输入/输出/属性集合query、key、value、dy为必选输入支持 FLOAT16/BFLOAT16/FLOAT32Ascend 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_apiaclnn_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 上通过 CubeAIC与 VectorAIV分离的架构并行执行依据 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 时显式传pseShiftOptionalnullptr且pseType1稀疏场景优先sparseMode0并保持 attenMask 与 preTokens/nextTokens 范围一致开启确定性计算需配合aclrtCtxSetSysParamOpt遇到 aicore timeout 时优先考虑对 B/S/N/D 轴做切分。【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表