CANN ops-transformer 中 FlashAttentionScore 算子完全解析:功能公式、参数约束与 NPU 训练加速原理
发布时间:2026/9/18 23:22:00来源:尧图网络
CANN ops-transformer 中 FlashAttentionScore 算子完全解析功能公式、参数约束与 NPU 训练加速原理【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformerFlashAttentionScore 是 CANN ops-transformer 算子库中用于训练场景的 FlashAttention 融合算子它把 QK^T 矩阵乘、位置编码叠加、注意力掩码、FlashSoftmax、Dropout 与 PV 矩阵乘等步骤融合进一次 NPU 计算并输出 Softmax 的 Max/Sum 中间结果供反向算子复用。本文以 attention/flash_attention_score/README.md 为骨架结合算子定义源码 flash_attention_score_def.cpp、Shape 推导实现 flash_attention_score_infershape.cpp、接口文档 aclnnFlashAttentionScoreV2.md 与设计文档 FA算子设计介绍.md完整讲解其功能、全部参数、约束条件、aclnn 调用方式及底层 Tiling 与流水实现帮助你在 NPU 上正确、高效地使用该算子。一、产品支持情况与适用场景FlashAttentionScore 面向训练场景的 self-attention自注意力计算其产品支持情况如下产品是否支持Ascend 950PR / Ascend 950DT√Atlas A3 训练系列产品√Atlas A3 推理系列产品×Atlas A2 训练系列产品√Atlas A2 推理系列产品×需要特别注意的是数据类型差异Atlas A2 训练系列产品不支持 FLOAT8_E5M2、FLOAT8_E4M3FN、HIFLOAT8 三种数据类型。Atlas A3 训练系列产品同样不支持 FLOAT8_E5M2、FLOAT8_E4M3FN、HIFLOAT8 三种数据类型。也就是说FP8 低比特数据类型仅在 Ascend 950PR / Ascend 950DT 上可用。该差异在算子注册源码中也有体现flash_attention_score_def.cpp 中ascend950的 AICore 配置aicore_config_95完整声明了包括ge::DT_FLOAT8_E5M2、ge::DT_FLOAT8_E4M3FN、ge::DT_HIFLOAT8在内的 22 种类型组合而ascend910b/ascend910_93即 A2/A3 训练系列的配置只保留 FLOAT16、BFLOAT16、FLOAT 三类。二、算子功能与计算公式算子功能在训练场景下使用 FlashAttention 算法实现 self-attention自注意力的计算。核心逻辑中位置编码 pse即realShiftOptional与缩放系数scale的执行次序由pseType属性控制pseType 1 时需要先 add 再 mulpseType ≠ 1 时需要先 mul 再 add。对应的正向计算公式如下pseType 1 时$$ attention_outDropout(Softmax(Mask(scale*(pse(queryd_scale_q)(keyd_scale_k)^T), atten_mask)), keep_prob)(value*d_scale_v) $$pseType ≠ 1 时$$ attention_outDropout(Softmax(Mask(scale*((queryd_scale_q)(keyd_scale_k)^T) pse),atten_mask),keep_prob)(value*d_scale_v) $$从公式可以看出该算子将以下计算链路全部融合query、key经d_scale_q、d_scale_kFP8 全量化参数非 FP8 场景默认为 1缩放后做矩阵乘得到注意力分数与位置编码pserealShiftOptional按pseType决定先加后乘或先乘后加乘以缩放系数scalescaleValue通过atten_mask对结果进行遮蔽经 Softmax 归一化内部使用 FlashSoftmax 实现按keep_prob与dropMaskOptional做 Dropout与缩放后的value做矩阵乘得到最终attention_out。d_scale_q、d_scale_k、d_scale_v三个参数对应 FP8 场景下 query / key / value 的量化参数使算子能够在 FP8 输入下先反量化计算、再以 BFLOAT16 输出这一逻辑在 flash_attention_score_infershape.cpp 的InferDataTypeFlashAttentionScore中有直接印证当输入为 FLOAT8_E5M2 / FLOAT8_E4M3FN / HIFLOAT8 时softmax_out与attention_out的输出数据类型统一被推导为DT_BF16。三、参数说明完整对照表下表完整列出算子的全部输入、输出与属性参数。其中“数据格式”一列的 ND 表示非降维的通用格式“-”表示属性参数不涉及数据格式。参数名输入/输出/属性描述数据类型数据格式query输入公式中的输入 queryBFLOAT16、FLOAT16、FLOAT、FLOAT8_E5M2、FLOAT8_E4M3FN、HIFLOAT8NDkey输入公式中的输入 keyBFLOAT16、FLOAT16、FLOAT、FLOAT8_E5M2、FLOAT8_E4M3FN、HIFLOAT8NDvalue输入公式中的输入 valueBFLOAT16、FLOAT16、FLOAT、FLOAT8_E5M2、FLOAT8_E4M3FN、HIFLOAT8NDrealShiftOptional可选输入公式中的 pse表示位置编码BFLOAT16、FLOAT16、FLOATNDdropMaskOptional可选输入公式中的 Dropout表示数据丢弃掩码。取值为 1 代表保留该数据为 0 代表丢弃该数据UINT8NDpaddingMaskOptional可选输入预留参数暂未使用BFLOAT16、FLOAT16、FLOATNDattenMaskOptional可选输入公式中的 atten_mask表示注意力掩码取值为 1 代表该位不参与计算不生效为 0 代表该位参与计算BOOL、UINT8NDprefixOptional可选输入prefix 稀疏计算场景中每个 Batch 的 N 值INT64NDactualSeqQlenOptional可选输入实际 Q 序列长度INT64NDactualSeqKvlenOptional可选输入实际 KV 序列长度INT64NDqStartIdxOptional可选输入Q 起始索引INT64NDkvStartIdxOptional可选输入KV 起始索引INT64NDdScaleQOptional可选输入公式中的 d_scale_qFP8 场景下 query 的全量化参数FLOATNDdScaleKOptional可选输入公式中的 d_scale_kFP8 场景下 key 的全量化参数FLOATNDdScaleVOptional可选输入公式中的 d_scale_vFP8 场景下 value 的全量化参数FLOATNDqueryRopeOptional可选输入Q 的 RoPE 旋转位置编码输入FLOAT8_E5M2、FLOAT8_E4M3FN、FLOAT16、BFLOAT16、FLOATNDkeyRopeOptional可选输入K 的 RoPE 旋转位置编码输入FLOAT8_E5M2、FLOAT8_E4M3FN、FLOAT16、BFLOAT16、FLOATNDsinkOptional可选输入Sinking 参数FLOATNDpScaleOptional可选输入P 的缩放因子FLOATNDscaleValue可选属性公式中的 scale表示缩放系数作为计算流中 Muls 的 scalar 值。默认值为 1.0DOUBLE-keepProb可选属性公式中的 keep_prob表示数据需要保留的概率。默认值为 1.0DOUBLE-preTokens可选属性用于稀疏计算表示 sliding window 的左边界。默认值为 2147483647INT64-nextTokens可选属性用于稀疏计算表示 sliding window 的右边界。默认值为 2147483647INT64-headNum必要属性代表单卡的 head 个数即输入 query 的 N 轴长度INT64-inputLayout必要属性代表输入 query、key、value 的数据排布格式。支持 BSH、SBH、BSND、BNSDSTRING-innerPrecise可选属性用于提升精度。默认值为 0INT64-sparseMode可选属性表示 sparse 的模式。支持配置值为 0、1、2、3、4、5、6。默认值为 0INT64-pseType可选属性控制 add 与 mul 的执行次序支持配置值为 0、1、2、3。默认值为 1INT64-seed可选属性Dropout 随机种子。默认值为 0INT64-offset可选属性Dropout 偏移量。默认值为 0INT64-outDtype可选属性输出精度控制。默认值为 0INT64-softmaxOutLayout可选属性Softmax 输出数据排布格式。默认值为空STRING-softmaxMaxOut输出Softmax 计算的 Max 中间结果用于反向计算FLOATNDsoftmaxSumOut输出Softmax 计算的 Sum 中间结果用于反向计算FLOATNDsoftmaxOut输出预留参数暂未使用BFLOAT16、FLOAT16、FLOATNDattentionOut输出公式中的 attention_outBFLOAT16、FLOAT16、FLOATND参数默认值与算子注册源码的对应关系上表所列默认值均能在算子注册文件 flash_attention_score_def.cpp 的Attr(...)声明中逐一找到this-Attr(scale_value).AttrType(OPTIONAL).Float(1.0); // scaleValue 默认 1.0 this-Attr(keep_prob).AttrType(OPTIONAL).Float(1.0); // keepProb 默认 1.0 this-Attr(pre_tockens).AttrType(OPTIONAL).Int(2147483647); // preTokens 默认 2147483647 this-Attr(next_tockens).AttrType(OPTIONAL).Int(2147483647); // nextTokens 默认 2147483647 this-Attr(head_num).AttrType(REQUIRED).Int(); // headNum 必填 this-Attr(input_layout).AttrType(REQUIRED).String(); // inputLayout 必填 this-Attr(inner_precise).AttrType(OPTIONAL).Int(0); // innerPrecise 默认 0 this-Attr(sparse_mode).AttrType(OPTIONAL).Int(0); // sparseMode 默认 0 this-Attr(pse_type).AttrType(OPTIONAL).Int(1); // pseType 默认 1 this-Attr(seed).AttrType(OPTIONAL).Int(0); // seed 默认 0 this-Attr(offset).AttrType(OPTIONAL).Int(0); // offset 默认 0 this-Attr(out_dtype).AttrType(OPTIONAL).Int(0); // outDtype 默认 0 this-Attr(softmax_out_layout).AttrType(OPTIONAL).String(); // softmaxOutLayout 默认空串其中preTokens/nextTokens的默认值 2147483647即 INT32_MAX表示“不限制窗口”等效于全量注意力pseType默认 1 表示与aclnnFlashAttentionScore基础接口行为一致先 add 再 mul。softmaxMaxOut / softmaxSumOut 的输出 Shape 推导softmaxMaxOut与softmaxSumOut是反向计算必需的中间结果其 Shape 由 flash_attention_score_infershape.cpp 的InferShapeFlashAttentionScore推导非 FP8 场景输出为(B, N, S, 8)的四维张量最后一维固定为 8FlashAttention 按 8 个元素一组维护 Softmax 在线归一化统计量HIFLOAT8 场景最后一维为 1TNDvarlen场景输出为(T, N, 8)三维张量两者数据类型恒为 FLOAT与输入类型无关由InferDataTypeFlashAttentionScore强制设置。同时attention_out的 Shape 与 query 保持一致BSND / BNSD 下第 3 维取 value 的 D 维BSH / SBH 下根据headNum、key/value 的 H 维动态换算D H / N后重算输出 H。四、约束说明使用前必读使用 FlashAttentionScore 前必须满足以下约束数据类型一致性输入 query、key、value、realShiftOptional 的数据类型必须一致。数据排布一致性输入 query、key、value 的 inputLayout 必须一致。Shape 范围约束以 inputLayout 的 BSND、BNSD 为例BSH、SBH 下 H N*DB取值范围 1 ~ 2M。当使用 prefixOptional 时 B 最大支持 2K。N取值范围 1 ~ 256。S取值范围 1 ~ 1M。D取值范围 1 ~ 768。输入 query、key、value 类型为 FLOAT8_E5M2、FLOAT8_E4M3FN、HIFLOAT8 时D 取值范围为 1 ~ 128。FP8 场景功能裁剪输入 query、key、value 类型为 FLOAT8_E5M2、FLOAT8_E4M3FN、HIFLOAT8 时不支持 queryRopeOptional、keyRopeOptional、realShiftOptional、attenMaskOptional、dropMaskOptional、keepProb、pseType 等相关可选参数。keepProb 取值范围(0, 1]。计算量与超时部分场景下如果计算量过大可能会导致算子执行超时aicore error 类型报错errorStr 为timeout or trap error此时建议做轴切分处理。计算量会受 B、S、N、D 等参数影响值越大计算量越大。pseType 2 或 3 时当前只支持 Sq 和 Skv 等长。GQA / MQA 支持在 aclnnFlashAttentionScore.md 接口文档中进一步说明支持输入 query 的 N 和 key/value 的 N 不相等但必须成比例关系即Nq/Nkv必须是非 0 整数Nq 取值范围 1 ~ 256。当Nq/Nkv 1时即为 GQAgrouped-query attention当Nkv 1时即为 MQAmulti-query attention。pseType 的四种取值含义pseType含义备注0外部传入 pse先 mul 再 add-1外部传入 pse先 add 再 mul与aclnnFlashAttentionScore基础接口实现一致2内部生成 pse先 mul 再 add仅支持 Sq 与 Skv 等长3内部生成 pse先 mul 再 add 再 sqrt仅支持 Sq 与 Skv 等长innerPrecise 与精度控制innerPrecise当前 0、1 为保留配置值2 表示开启无效行计算其功能是避免在计算过程中存在整行 mask 进而导致精度损失但该配置会导致性能下降。如果算子可判断出存在无效行场景会自动开启无效行计算例如 sparseMode 为 3 且 Sq Skv 的场景。alibi 位置编码压缩realShiftOptional 的特殊用法如果 Sq 大于 1024 且每个 batch 的 Sq 与 Skv 等长且处于 sparseMode 为 0、2、3 的下三角掩码场景可开启 alibi 位置编码压缩此时只需要输入原始 PSE 最后 1024 行实现内存优化即alibi_compress ori_pse[:, :, -1024:, :]具体为参数每个 batch 不相同时shape 为BNHSkvH1024每个 batch 相同时shape 为1NHSkvH1024如果 pseType 为 2 或 3数据类型需为 FLOAT32对应 shape 支持范围是[B,N]或[N]如果不开启该参数realShiftOptional 需要传入 nullptrpseType 需要传入 1。sparseMode 稀疏模式约束当所有 attenMaskOptional 的 shape 小于 2048 且相同的时候建议使用 default 模式sparseMode0以减少内存使用量。配置为 1、2、3、5 时用户配置的 preTokens、nextTokens 不会生效。配置为 0、4 时须保证 attenMaskOptional 与 preTokens、nextTokens 的范围一致。用户不特意指定时建议传入 0。band 场景下preTokens 和 nextTokens 之间必须要有交集。prefixOptional 稀疏计算场景即 sparseMode5 或 6当 Sq Skv 时prefix 的 N 值取值范围 [0, Skv]当 Sq Skv 时prefix 的 N 值取值范围 [Skv-Sq, Skv]。五、aclnn 调用说明与两段式接口FlashAttentionScore 算子的标准调用方式是aclnn 接口通过 examples/test_aclnn_flash_attention_score.cpp 可以查看完整可运行的调用样例非 TND 场景。该样例通过aclnnFlashAttentionScoreV2接口方式调用接口完整定义见 op_api/aclnn_flash_attention_score.h。两段式 API 模式每个算子分为两段式接口必须先调用xxxGetWorkspaceSize接口获取计算所需 workspace 大小以及包含了算子计算流程的执行器executor再调用xxx接口执行计算。以 V2 接口为例// 第一段获取 workspace 大小与执行器 aclnnStatus aclnnFlashAttentionScoreV2GetWorkspaceSize( 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 *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 aclnnFlashAttentionScoreV2( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, const aclrtStream stream);关键调用步骤结合 examples/test_aclnn_flash_attention_score.cpp 的实现完整调用流程分为以下几步初始化 AscendCLaclInit→aclrtSetDevice→aclrtCreateStream构造输入与输出 Tensor通过aclrtMalloc申请 Device 内存、aclrtMemcpy拷贝 Host 数据、aclCreateTensor创建aclTensorND 格式需正确计算连续 strides样例中默认配置为B1, N11, N21, S1256, S2256, D128inputLayoutSBHscaleValue 1.0/sqrt(128)pseType1调用第一段接口aclnnFlashAttentionScoreV2GetWorkspaceSize获取workspaceSize与executor按需aclrtMalloc申请 workspace调用第二段接口aclnnFlashAttentionScoreV2(workspaceAddr, workspaceSize, executor, stream)执行计算同步等待aclrtSynchronizeStream(stream)等待任务结束回拷结果将attentionOut、softmaxMax、softmaxSum从 Device 拷贝回 Host 并打印释放资源aclDestroyTensor释放 tensor、aclDestroyIntArray释放 int arrayaclrtFree释放 Device 内存、销毁 Stream、Reset Device 并aclFinalize。返回码与入参校验第一段接口完成入参校验常见的错误返回码如下返回值错误码描述ACLNN_ERR_PARAM_NULLPTR161001传入参数是必选输入、输出或必选属性且是空指针ACLNN_ERR_PARAM_INVALID161002query、key、value、realShiftOptional、dropMaskOptional、paddingMaskOptional、attenMaskOptional、softmaxMaxOut、softmaxSumOut、softmaxOutOut、attentionOutOut 的数据类型或数据格式不在支持的范围内接口家族一览该算子模块还提供了多个演进版本接口定义见 op_api/aclnn_flash_attention_score.h方便按能力演进选用接口新增能力aclnnFlashAttentionScore基础版见 aclnnFlashAttentionScore.mdaclnnFlashAttentionScoreV2新增 pseType、qStartIdx、kvStartIdx 参数aclnnFlashAttentionScoreV3新增 sinkOptional 参数aclnnFlashAttentionScoreV4新增 queryRope、keyRope、dScaleQ/K/V、seed、offset、outDtype、softmaxOutLayout 等aclnnFlashAttentionVarLenScore 系列TND / varlen 变长序列场景V2~V5 逐步演进六、底层实现原理从计算流程到多模板设计6.1 训练正向计算流程按照 FlashAttention 正向计算流程实现详见 FA算子设计介绍.md整体计算流程如下Bmm1 与 PSE 融合query 与转置后的 key 做 matmul 得到初步的 attention_score与位置编码 pse 相加按 pseType 调整次序后再乘以缩放系数 scale_value然后通过 atten_mask 进行 select 操作将 atten_mask 中为 true 的位置遮蔽为负的极小值经过 softmax 计算之后变成 0 从而达到遮蔽效果FlashSoftmax 在线归一化使用 FlashSoftmax 操作代替原公式中的 softmax 运算并对 Skvkey/value 的 sequence length方向进行切分因此存在刷新流程每次 FlashSoftmax 只对切分后的一个 SkvSplit 操作从第二次循环开始记录 exp其中exp[i] e^(max_{i-1} - max_i)当 i 0 时MM[PV] 的结果直接保存到ub_attention_out[0]从 i 1 开始将上一次的 MM[PV] 结果与当前 exp 相乘再与本次 MM[PV] 结果相加保存到ub_attention_out[1]以此类推遍历 Skv 完成计算由于除 sum 被后移到输出之前最后需要将 UB 中的结果按行除以 softmax_sum并将最终完整结果写回 GM 中的 attention_out。6.2 每个计算阶段的入口以 S1 模板为例得到每个核处理的任务块数后调用函数ProcessCube1 阶段QK^T计算入口IterateBmm1Vector1 阶段掩码/Softmax/PSE 等计算入口ProcessVec1Cube2 阶段PV计算入口IterateBmm2Vector2 阶段输出归一化等计算入口ProcessVec26.3 Tiling 设计AIC/AIV 分离与 CV 配比Atlas A2 训练系列产品采用 AICCube与 AIVVector分离架构两者拥有独立的 Scalar 计算单元通过 L2 和 GM 交互。本着减少 CV 通信次数、发挥最大算力的原则FA 采用 CV tiling 分离策略Vector 侧数据类型为 float32基本块为8 * 1024最大 buffer 分配 32KBCube 侧输入 float16、输出 float32基本块为128 * 128通过 nRatio8 配比出128 * 1024的数据量实现 C:V 1:16 的算力配比伪代码如下// C-Tiling: (S1_c_i,D)x(D,S2_c_i) (S1_c_i, S2_c_i):(128,1024) // V-Tiling: (S1_v_i, S2_v_i) (8,1024) // C侧matmul计算 Bmm((S1_c_i,D)x(D,S2_c_i)) 128*1024 // 输出结果128*1024放到workspace上 // V侧Vector计算 for S1_c_i/S1_v_i128/8: copy_gm_to_ub(S1_v_i*S2_v_i) // 从bmm的workspace上拷入bmm结果数据 Vector(S1_v_i,S2_v_i) // 进行Vector计算 copy_ub_to_gm(S1_v_i*S2_v_i) // Vector计算结束得到最终输出数据拷贝到GM上也可以在 S1/S2 方向均开启配比这样 Cube 一次可以发射大块数据避免小块数据不断发射带来的通信开销同时最大程度使用 Cube 单元的 buffer。Ascend 950PR / Ascend 950DT既支持 AICAIV 分离架构又支持 AICAIV 融合架构且 UBUnified Buffer是 AIC 与 AIV 之间较 GM 更高效的交互通路——AIC 可直接输出到 UBAIV 利用 UB 上的数据做 vec 运算极大降低数据搬运时间。由于 UB 容量有限D 128 时 mm2 的输出仍需要先保存到 GMmm1 的输出始终保存在 UB。6.4 流水设计为了充分利用硬件资源FA 算子做了精细的流水设计宗旨是尽量使某一条 pipeline 达成 bound 效果V 侧流水主要优化手段是 double bufferping-pong。同一片数据的搬运DataCopy与计算Clc之间串行处理不同数据切片在同一时间点可多个任务并行由此达到任务并行、提升性能的目的。由于 FA 类融合算子 V 侧计算过程较多通常简单的 double buffer 无法覆盖所有情况因此存在多种计算流水排布不同流水适用于不同 shape 特征。CV 流水Atlas A2 训练系列产品FA 的 Cube 双发机制CV 间 preload 流水实现流水优化。在 C:V1:2 的情况下Cube 的搬运时长足以覆盖 Vector 的计算时长只需关注 Cube 的 MTE2 耗时即可最终达成 MTE2 boundCube 双发机制下提前发射两块 Cube 计算使 Cube1、Cube2 计算衔接达成 Cube bound。CV 流水Ascend 950PR / Ascend 950DT设计思路与 A2 基本一致差异点在于 cube 的 preload 次数为 3 次——完成 3 次 mm1 的计算后才开启 mm2 的计算目的是优化启动阶段的 CV 流水使其更紧密。6.5 多模板设计为使不同输入复用相同 tiling 和流水FA 融合算子根据输入 shape 特征区分模板核心思路包括三类根据核内及核间切分进行模板拆分FA 算子包含 B、N2key/value 的 N、Gquery_N/kv_N、S1query 的 S、S2key/value 的 S共 5 个轴切分顺序是先核内再核间。拆分时需考虑将核心数量用满、每个核分配的计算量相对均匀、AIC 与 AIV 处理的数据量符合对应算力。根据特殊场景及特定优化进行模板特化基础模板覆盖功能及大部分 shape 输入的性能特化模板针对特定场景做极致性能优化例如空 tensor 特化的 empty_input 模板、基于 RCM角色化 Cube 核管理的确定性计算模板等。根据不同计算流水进行模板特化不同流水设计对代码架构影响较大为提升可维可测可读性按计算流水进行模板特化。Atlas A2 训练系列产品从 Vector 视角按优先级划分为以下几类模板序号越小优先级越高越先匹配模板切分轴进入条件适用范围TNDSameAB 模板UB 切 S1(B 4 and accumS1 8192 and accumS2 8192 and maxS1 512 and maxS2 512) or (maxS2 5120 and maxS1 5120)TND 场景TND 模板UB 切 S1S2TND 场景但不满足 TNDSameAB 模板条件TND 场景SameAB 模板UB 切 S1非 FP32S2 512 and D % 16 ! 0 or D 96or S2 1024 and 128 D 196普通场景S1S2 模板UB 切 S1S2不满足上述条件且 S2 1024普通场景S1 模板UB 切 S1D中间结果超过 UB 容量如 N2*G*((alignedS1alignedS2)*alignedDalignedS2)*dtypeSize 256*1024 等普通场景B 模板核间切 B不满足以上条件普通场景Ascend 950PR / Ascend 950DT由于核内切分的基本块为 128*128在各种 shape 下性能都能达到最优水平因此只设计一套模板支持全量 shape且不区分 layout 是否为 TND。6.6 编程视角高阶 API 与低阶 APIAscendC 高阶 API 两种模式一是以 Vector 为主核、Cube 为从核Vector0/Vector1 独立发起 Matmul 任务二是以 Cube 为主核、Vector 为从核V0 统一发起 Matmul结果由 V0/V1 共同处理一般各处理一半。模板名以_sab结尾即表示以 Cube 为主核的模板如 flash_attention_score_s1s2_bn2gs1_sab.h。以 Cube 为主核对 FlashAttention 而言V0、V1 的 Matmul 任务可复用左矩阵且部分结果可在 L0C 累加降低带宽依赖大部分场景性能更优。AscendC 低阶 API部分模板更彻底地使用以 Cube 为主核、Vector 为从核的模式Matmul 任务完全从 Cube 侧发起并通过同步通知 Vector 侧。Ascend 950PR / Ascend 950DT 上的 FA 通过 AscendC 低阶 API 实现。七、功能验证pytest 测试框架仓库为算子提供了基于 pytest 的功能验证框架见 tests/pytest/README.md文件结构如下test_case.py测试用例集test_flash_attn.py执行主程序test_utils.py工具方法cpu_impl.pyCPU 实现用于生成 golden 数据npu_impl.pyNPU 实现通过 TorchNPU 算子直调获取实际数据验证思路为CPU 侧复现算子功能生成 golden 数据 → NPU 侧通过 TorchNPU 直调获取实际结果 → 进行 CPU 与 NPU 结果的精度对比以验证算子功能正确性。前置要求为确认 TorchNPU 为最新版本并 source CANN 包环境变量在 pytest 文件夹路径下执行pytest -s即可运行全部用例。此外tests/ut/op_host 下的单元测试如test_flash_attention_score_tiling.cpp、test_flash_attention_score_infershape.cpp及对应 CSV 用例覆盖了 tiling 计算与 infershape 的 C 层验证tests/ut/op_api 则覆盖了 aclnn 接口层的调用验证。八、总结FlashAttentionScore 是 CANN ops-transformer 在 NPU 训练场景下实现 self-attention 的核心融合算子通过 QK^T、FlashSoftmax、Dropout、PV 的端到端融合以及多模板、多流水、CV 配比等底层优化将大模型注意力计算高效落地到 Ascend 硬件。使用时的关键决策点包括根据产品型号确认 FP8 数据类型的可用性仅 Ascend 950PR/DT 支持、按pseType正确配置位置编码的计算次序、按sparseMode与preTokens/nextTokens选择稀疏注意力模式、通过softmaxMaxOut/softmaxSumOut输出中间结果供反向复用以及严格遵循 B/N/S/D 的 Shape 与数据类型一致性约束。更多细节可继续阅读仓库内的 aclnnFlashAttentionScoreV2.md 接口文档与 FA算子设计介绍.md 设计文档。【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
网站建设高端定制企业官网