CANN ops-transformer FlashAttnGrad 算子详解:Flash Attention 反向梯度计算的原理、构建与 Torch 接口实战

发布时间:2026/9/20 20:17:03
CANN ops-transformer FlashAttnGrad 算子详解:Flash Attention 反向梯度计算的原理、构建与 Torch 接口实战
CANN ops-transformer FlashAttnGrad 算子详解Flash Attention 反向梯度计算的原理、构建与 Torch 接口实战【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer本篇文章聚焦 CANN ops-transformer 项目attention/flash_attn_grad 目录中的flash_attn_grad反向注意力算子系统讲解其数学原理、输入输出与属性约定、custom 包与 torch 扩展的完整构建安装流程、PyTorch 接口调用规范并结合仓库源码剖析 tiling 与 kernel 的底层实现。读者读完可掌握在 Ascend 950 上独立编译、安装并正确调用flash_attn_grad完成 Flash Attention 训练反向传播的能力。一、算子功能与数学原理flash_attn_grad是 Flash Attention 的反向梯度计算算子。它根据前向注意力计算保存下来的中间结果softmax_lse、attn_out和上游梯度do/dout计算 Q/K/V 的梯度dq/dk/dv从而避免在反向阶段重新计算完整注意力矩阵是 Flash Attention 训练链路中降低显存占用、提升训练吞吐的关键一环。前向计算可表示为$$ Sscale\cdot QK^T $$$$ P_{ij}\exp(S_{ij}-softmax_lse_i) $$$$ YPV $$其中 $scale$ 的取值规则为softmax_scale属性非 0 时取softmax_scale为 0 时取 $scale1/\sqrt{D}$D 为 Query/Key 的 head dim。反向计算依据链式法则展开README.md 与接口文档 torchapi_flash_attn_grad.md 给出的完整公式为$$ dVP^TdY $$$$ dPdYV^T $$$$ sfmgrowsum(dY\odot Y) $$$$ dSP\odot(dP-sfmg) $$$$ dQscale\cdot(dS\cdot K) $$$$ dKscale\cdot(dS^T\cdot Q) $$其中 $Q$、$K$、$V$ 分别对应输入q、k、v$Y$ 对应attn_out$dY$ 对应上游梯度dout。注意反向计算复用了前向的softmax_lse重建概率矩阵 $P$这正是算子可以在不保存完整 $P$ 矩阵的情况下完成反向传播的数学基础——通过softmax_lse每行 log-sum-exp就能在反向时低成本恢复归一化后的 $P$。符号约定B 表示 batch 大小S1 表示 Query 序列长度S2 表示 Key/Value 序列长度N1 表示 Query head 数N2 表示 Key/Value head 数D 表示 Query/Key 的 head dimDv 表示 Value/输出 的 head dim。T1 表示所有 batch 中 Query 序列长度的累加和T2 表示所有 batch 中 Key/Value 序列长度的累加和。GQA 场景满足 N1 是 N2 的整数倍。二、输入、输出与属性2.1 输入13 个名称类型必选说明qBF16/FP16是Query tensorkBF16/FP16是Key tensorvBF16/FP16是Value tensordoBF16/FP16是上游梯度attn_outBF16/FP16是前向注意力输出softmax_lseFP32是前向 softmax LSEcu_seqlens_qINT32否TND layout 累积序列长度cu_seqlens_kvINT32否TND layout 累积序列长度seqused_qINT32否实际使用的 Q 序列长度seqused_kvINT32否实际使用的 KV 序列长度sinksFP32否Sink tensorattn_maskINT8否Attention maskmask_mode3/4 时使用metadataINT32否FAG metadata tensor实际调用时必须传入这些声明与源码 op_host/flash_attn_grad_def.cpp 中的Input(...)注册完全一一对应6 个必选输入q/k/v/dout/attn_out/softmax_lse 7 个可选输入cu_seqlens_q/cu_seqlens_kv/seqused_q/seqused_kv/sinks/attn_mask/metadata全部注册了 ND 格式与AutoContiguous传入非连续 Tensor 时由框架自动转为连续 Tensor。softmax_lse与sinks为 FP32attn_mask为 INT8其余辅助索引张量均为 INT32。2.2 输出3 个名称类型说明dqBF16/FP16Query 梯度shape 与q相同dkBF16/FP16Key 梯度shape 与k相同dvBF16/FP16Value 梯度shape 与v相同shape/dtype 推导逻辑在 op_host/flash_attn_grad_infershape.cpp 中实现dq/dk/dv分别直接继承q/k/v的 shapedtype 统一取q的 dtype。PyTorch 侧的 Meta 实现torch_extension/flash_attn_grad.py 的register_meta也遵循同一规则dq torch.empty(q.size(), ...)。2.3 属性9 个名称类型默认值说明softmax_scaleFloat0.0softmax 缩放因子0.0 表示 1/sqrt(d)mask_modeInt00: 全计算, 3: causal, 4: windowwin_leftInt-1window mask 左窗口-1 表示正无穷win_rightInt-1window mask 右窗口-1 表示正无穷max_seqlen_qInt-1Q 最大序列长度-1 表示由输入 shape 推导max_seqlen_kvInt-1KV 最大序列长度-1 表示由输入 shape 推导layout_qStringBSNDQ 布局BSND/BNSD/TNDlayout_kvStringBSNDKV 布局必须与 layout_q 相同layout_outStringBSND输出布局必须与 layout_q 相同从 op_host/flash_attn_grad_def.cpp 可以看到所有属性均注册为 OPTIONAL 并带默认值算子以DynamicCompileStaticFlag(true)、DynamicFormatFlag(true)、DynamicRankSupportFlag(true)、DynamicShapeSupportFlag(true)注册到ascend950的 AICore 配置中且ExtendCfgInfo(opFile.value, flash_attn_grad)将 opFile 指向内核实现。三、构建与安装Quick Start3.1 custom 包编译与安装完整脚本整合编译与安装步骤CANN_DIR为 CANN 安装根目录按实际环境调整# 前置加载 CANN 环境 source ${CANN_DIR}/cann/set_env.sh # 清理历史构建产物避免残留影响增量编译 rm -rf ./build ./build_out rm -rf ${CANN_DIR}/vendors # 编译flash_attn_grad 与 flash_attn_metadata 需同时编译metadata 生成反向分核信息 bash build.sh --pkg --socascend950 --opsflash_attn_grad,flash_attn_metadata # 安装 cd build_out ./cann-ops-transformer-*.run --install-path${CANN_DIR}编译产物为build_out/cann-ops-transformer-custom_linux-x86_64.run。可选参数说明-j限制并行线程数默认按机器核数并行。当机器内存不足、或 cgroup 实际限制核数小于/proc/cpuinfo报告值导致编译 OOM 或失败时需显式指定较小值bash build.sh --pkg --socascend950 --opsflash_attn_grad,flash_attn_metadata -j163.2 torch 扩展包构建与安装使用 torch 接口前必做。在仓库根目录构建 torch 扩展 whl 并安装# 前置加载 CANN 环境 source ${CANN_DIR}/cann/set_env.sh # 清理 torch 扩展缓存~ 为当前用户 home需与安装/运行 torch 的用户一致避免加载过期编译产物 rm -rf ~/.cache/torch_extensions/* # 构建 torch 扩展 whl全量包包名 cann_ops_transformer 保持不变whl 输出到 build_out/ bash build.sh --torch_extension --socascend950 # 安装 python3 -m pip install build_out/*.whl --force-reinstall --no-deps安装后验证python3 -c from cann_ops_transformer.ops import flash_attn_grad; print(ok)3.3 接口调用调用分两步先用flash_attn_metadata生成反向分核 metadata必须设置is_grad_enabledTrue并保证与主算子的 shape、layout、mask 和序列长度参数一致再调用flash_attn_grad主算子。导入路径与安装包名一致按上述步骤构建的全量包from cann_ops_transformer.ops import flash_attn_grad四、PyTorch 接口详解完整接口文档见 torchapi_flash_attn_grad.md下面摘录核心内容并补充源码佐证。4.1 产品支持情况flash_attn_grad当前仅支持Ascend 950PR/Ascend 950DTAtlas A2/A3 训练与推理系列、Atlas 200I/500 A2 推理产品、Atlas 推理系列、Atlas 训练系列产品均不支持。这与算子定义中AddConfig(ascend950, ...)的注册范围一致。4.2 函数原型cann_ops_transformer.flash_attn_grad( q, k, v, dout, attn_out, softmax_lse, cu_seqlens_qNone, cu_seqlens_kvNone, seqused_qNone, seqused_kvNone, sinksNone, attn_maskNone, metadataNone, softmax_scale0.0, mask_mode0, win_left-1, win_right-1, max_seqlen_q-1, max_seqlen_kv-1, layout_qBSND, layout_kvBSND, layout_outBSND ) - (Tensor, Tensor, Tensor)当前注册的 PyTorch schema 未使用*分隔符因此cu_seqlens_q及之后的参数既可以按位置传入也可以按关键字传入。建议使用关键字传入可选参数。4.3 mask_mode 枚举mask_mode在 Python 接口中支持传入IntEnum枚举或对应 int 值枚举定义于cann_ops_transformer.ops.attention.flash_attn_grad枚举名值含义NO_MASK0全计算模式默认值CAUSAL3Causal 模式SLIDING_WINDOW4Sliding Window 模式枚举为IntEnum可直接作为 int 传入底层算子接口同时兼容传入枚举名对应的字符串不区分大小写与 int 值。值得留意的是host 侧 tiling 的mask_mode位域按 0/1/2 存储、kernel 侧按[0,3,4]解释两者之间的映射在 tilingkey 编码时完成见 op_host/plan/flash_attn_grad_tiling_key.h 的注释bit[4:3] mask_mode host 存 0/1/2kernel values[0,3,4]。4.4 参数说明参数名参数类型可选/必选描述数据类型数据格式维度qTensor必选公式中的Qbfloat16/float16NDBSND(B, S1, N1, D)BNSD(B, N1, S1, D)TND(T1, N1, D)kTensor必选公式中的Kbfloat16/float16NDBSND(B, S2, N2, D)BNSD(B, N2, S2, D)TND(T2, N2, D)vTensor必选公式中的Vbfloat16/float16NDBSND(B, S2, N2, Dv)BNSD(B, N2, S2, Dv)TND(T2, N2, Dv)doutTensor必选公式中的dY前向输出的上游梯度bfloat16/float16NDBSND(B, S1, N1, Dv)BNSD(B, N1, S1, Dv)TND(T1, N1, Dv)attn_outTensor必选公式中的Y即前向接口返回的注意力输出bfloat16/float16NDshape 与dout相同softmax_lseTensor必选前向接口在return_softmax_lseTrue时返回的 log-sum-exp 结果float32NDBSND/BNSD(B, N1, S1)TND(N1, T1)cu_seqlens_qTensor可选TND 布局下 Q 的累积序列长度第一个元素必须为 0最后一个元素等于 T1。layout_q为 TND 时必须传入非 TND 时不支持传入。默认 Noneint32ND(B1,)cu_seqlens_kvTensor可选TND 布局下 KV 的累积序列长度第一个元素必须为 0最后一个元素等于 T2。layout_kv为 TND 时必须传入非 TND 时不支持传入。默认 Noneint32ND(B1,)seqused_qTensor可选每个 batch 实际使用的 Q 序列长度。默认 Noneint32ND(B,)seqused_kvTensor可选每个 batch 实际使用的 KV 序列长度。默认 Noneint32ND(B,)sinksTensor可选Sink 参数用于改善自注意力计算的数值稳定性。默认 Nonefloat32ND(N1,)attn_maskTensor可选掩码矩阵。默认 Noneint8ND(2048, 2048)metadataTensor可选flash_attn_metadata生成的 FAG 任务切分数据。schema 中为可选参数但实际调用时必须传入int32NDshape 根据 batch 大小和 N2 动态计算softmax_scalefloat可选Softmax 缩放系数。默认 0.0表示使用 $1/\sqrt{D}$float32--mask_modeint/MaskMode可选掩码模式支持传入枚举或对应 int 值。默认 0int32--win_leftint可选Window mask 左窗口值需大于等于 -1-1 表示正无穷。默认 -1int32--win_rightint可选Window mask 右窗口值需大于等于 -1-1 表示正无穷。默认 -1int32--max_seqlen_qint可选Q 最大序列长度必须大于等于 -1。BSND/BNSD 场景可保持默认 -1由输入 shape 推导 S1int32--max_seqlen_kvint可选KV 最大序列长度必须大于等于 -1。BSND/BNSD 场景可保持默认 -1由输入 shape 推导 S2int32--layout_qstring可选q 的布局支持 BSND、BNSD、TND。默认 BSNDstring--layout_kvstring可选k 和 v 的布局必须与layout_q相同。默认 BSNDstring--layout_outstring可选dout和attn_out的布局必须与layout_q相同。默认 BSNDstring--q、k、v、dout、attn_out及输出dq、dk、dv支持 float16 和 bfloat16数据类型必须一致。4.5 返回值说明参数名参数类型描述数据类型数据格式维度dqTensor公式中的 dQQuery 的梯度bfloat16/float16NDshape 与q相同dkTensor公式中的 dKKey 的梯度bfloat16/float16NDshape 与k相同dvTensor公式中的 dVValue 的梯度bfloat16/float16NDshape 与v相同4.6 约束说明接口支持以下组合数据类型为 float16 或 bfloat16layout_q、layout_kv、layout_out支持 BSND、BNSD、TND且必须相同Dense Attentionmask_mode0、attn_maskNone、win_left-1、win_right-1TND 布局下cu_seqlens_q、cu_seqlens_kv必须传入非 TND 布局下不支持传入。q、k、v、dout、attn_out的数据类型必须一致。B、S1、S2、N1、N2、D 和 Dv 必须为正数其中 B 的取值范围为 (0, 65536)。N1 必须能被 N2 整除支持 MHAN1N2和 GQAN1N2。Query 和 Key 的 head dim 均为 DValue、dout和attn_out的 head dim 均为 Dv并满足0 Dv D 192。与当前仓库中的flash_attn联合调用时前向接口还要求 DDv 且 D 取 64、128 或 256结合本接口 D 不超过 192 的约束联合调用当前支持 DDv64 或 128。softmax_lse必须为 float32shape 为 (B, N1, S1)BSND/BNSD 场景或 (N1, T1)TND 场景。所有输入的数据格式均为 ND。算子注册了AutoContiguous传入非连续 Tensor 时由框架转换为连续 Tensor。metadata必须由flash_attn_metadata生成并设置is_grad_enabledTrue。生成 metadata 和调用本接口时N1、N2、D、B、S1、S2、layout 及 mask 相关参数必须一致否则行为未定义。is_grad_enabledTrue生成的 metadata 同时包含正向和反向任务切分数据前向flash_attn和反向flash_attn_grad均可使用同一份 metadata无需分别生成。softmax_lse、attn_out必须来自与本次反向计算配置一致的前向调用尤其是softmax_scale和 layout 必须一致。当前仅支持单算子模式。4.7 配套接口 flash_attn_metadata调用flash_attn_grad之前需要通过flash_attn_metadata生成反向任务切分数据cann_ops_transformer.flash_attn_metadata( num_heads_q, num_heads_kv, head_dim, *, cu_seqlens_qNone, cu_seqlens_kvNone, seqused_qNone, seqused_kvNone, batch_sizeNone, max_seqlen_qNone, max_seqlen_kvNone, mask_modeNone, win_leftNone, win_rightNone, layout_qNone, layout_kvNone, layout_outNone, is_grad_enabledFalse ) - Tensor与反向接口直接相关的参数如下参数名参数类型可选/必选描述数据类型数据格式维度num_heads_qint必选Query head 数即 N1int32--num_heads_kvint必选Key/Value head 数即 N2int32--head_dimint必选Query/Key 的 head dim即 Dint32--cu_seqlens_qTensor可选TND 布局下 Q 的累积序列长度。layout_q为 TND 时必须传入非 TND 时不支持传入。默认 Noneint32ND(B1,)cu_seqlens_kvTensor可选TND 布局下 KV 的累积序列长度。layout_kv为 TND 时必须传入非 TND 时不支持传入。默认 Noneint32ND(B1,)seqused_qTensor可选每个 batch 实际使用的 Q 序列长度。默认 Noneint32ND(B,)seqused_kvTensor可选每个 batch 实际使用的 KV 序列长度。默认 Noneint32ND(B,)batch_sizeint可选batch 大小。BSND/BNSD 场景必须传入实际 B取值范围 (0, 65536)int32--max_seqlen_qint可选Q 最大序列长度。BSND/BNSD 场景必须传入实际 S1且大于 0int32--max_seqlen_kvint可选KV 最大序列长度。BSND/BNSD 场景必须传入实际 S2且大于 0int32--mask_modeint/MaskMode可选必须与flash_attn_grad一致int32--win_leftint可选必须与flash_attn_grad一致int32--win_rightint可选必须与flash_attn_grad一致int32--layout_qstring可选必须与flash_attn_grad一致string--layout_kvstring可选必须与flash_attn_grad一致string--layout_outstring可选必须与flash_attn_grad一致string--is_grad_enabledbool可选是否生成反向算子所需的 metadata。调用flash_attn_grad前必须设置为 True。默认 Falsebool--返回的metadata为 int32、ND 格式的一维 Tensor长度根据 batch 大小和 N2 动态计算。4.8 完整调用示例BSNDflash_attn_metadata、flash_attn和flash_attn_grad联合调用示例。is_grad_enabledTrue的flash_attn_metadata会同时生成正向和反向的任务切分数据前向和反向均可使用该 metadata无需分别生成import math import torch import torch_npu import cann_ops_transformer torch_npu.npu.set_device(0) dtype torch.float16 B 2 S1 128 S2 128 N1 8 N2 2 D 128 Dv 128 scale 1.0 / math.sqrt(D) q torch.randn(B, S1, N1, D, dtypedtype, devicenpu) k torch.randn(B, S2, N2, D, dtypedtype, devicenpu) v torch.randn(B, S2, N2, Dv, dtypedtype, devicenpu) metadata cann_ops_transformer.flash_attn_metadata( N1, N2, D, batch_sizeB, max_seqlen_qS1, max_seqlen_kvS2, mask_mode0, win_left-1, win_right-1, layout_qBSND, layout_kvBSND, layout_outBSND, is_grad_enabledTrue, ) attn_out, softmax_lse cann_ops_transformer.flash_attn( q, k, v, metadatametadata, softmax_scalescale, mask_mode0, win_left-1, win_right-1, max_seqlen_qS1, max_seqlen_kvS2, layout_qBSND, layout_kvBSND, layout_outBSND, return_softmax_lseTrue, ) dout torch.randn_like(attn_out) dq, dk, dv cann_ops_transformer.flash_attn_grad( q, k, v, dout, attn_out, softmax_lse, metadatametadata, softmax_scalescale, mask_mode0, win_left-1, win_right-1, max_seqlen_qS1, max_seqlen_kvS2, layout_qBSND, layout_kvBSND, layout_outBSND, ) torch_npu.npu.synchronize() assert dq.shape q.shape assert dk.shape k.shape assert dv.shape v.shape该示例展示的是 GQA 场景N18、N22符合N1 必须能被 N2 整除的约束将 N2 设为 8 即为 MHA 场景。TND 变长序列场景则需额外传入cu_seqlens_q/cu_seqlens_kv。五、算子目录结构与源码级实现README.md 给出的目录结构如下文件说明op_kernel/flash_attn_grad.pypypto-pro kernel 实现BN2GS1S2 模板op_host/flash_attn_grad_tiling.cpptiling 实现tilingkey 编码、workspace 布局op_host/flash_attn_grad_def.cpp算子定义op_proto/输入输出/属性op_host/flash_attn_grad_infershape.cppshape/dtype 推导op_host/config/ascend950/flash_attn_grad_binary.json二进制算子配置torch_extension/flash_attn_grad.pyPyTorch 算子 schema 与 python 绑定torch_extension/csrc/flash_attn_grad.cppPyTorch C 扩展aclnn 调用层5.1 host 侧tiling 流水线op_host/flash_attn_grad_tiling.cpp 是 L0 入口注释明确其职责为解析 - 校验 - 规划 - 写回ParsePlatform解析平台信息AIV/AIC 核数、L2 缓存大小、libapi workspace 大小FlashAttnGradCheck::CheckParams对输入参数做合法性校验ParseFlashAttnGradInfo解析算子信息构造FlashAttnGradTilingRegbase并调用DoTiling完成分核规划。平台信息还会在编译期通过TilingParseForFlashAttnGrad缓存进FlashAttnGradCompileInfo供图模式下 Tiling 阶段拿不到PlatformInfo时回退使用。分核策略进一步拆分为 planop_host/plan/下的tiling_key、tiling_plan、tiling_route、tiling_schedule、tiling_swizzle、tiling_workspace等模块、checkers属性/输入/shape/feature 校验、infotiling 信息解析三层职责边界清晰。TilingKey 的位域编码在 op_host/plan/flash_attn_grad_tiling_key.h 中定义与 kernel 侧FlashAttnGradTilingKey保持一致bit[1:0] template 0BN2GS1S2, 1BN2, 2未用, 3BN2S2(保留勿复用) bit2 layout 0非TND, 1TND bit[4:3] mask_mode host 存 0/1/2kernel values[0,3,4] bit5 swizzle 仅 template0 有效 bit[7:6] d_align 064, 1128, 2192 bit8 dv_align 0128, 1192 bit9 is_bn2_multiblk 仅 template1 有效 bit10 bn2_need_zero 仅 MultiBlk mask 3/4 的无效行/列该文件同时强调合法组合的真源在 pypto 侧的FlashAttnGradTilingKey / is_valid()host 只能产出is_valid为真的组合否则运行期找不到对应二进制。5.2 kernel 侧PyPTO Pro 双模板实现op_kernel/flash_attn_grad.py 是 kernel 入口只做编译期分发根据 tilingkey 的template位选择两个主循环模板flash_attn_grad_bn2gs1s2template0通用切分模板cube 领先 vector 两个 taskPRELOAD_TIMES3dQ 走 fp32 workspace atomicAdddK/dV 在 L0C 跨 s1 累加同时接收cu_seqlens_q/kv、seqused_q/kv、sinks等全部可选输入支持 GQA 与 D192。flash_attn_grad_bn2template1一个核独占一个 (b, n2) head1-ahead ping-pong无 pre/post不接 GQA / D192支持 mask禁止 BN2swizzle 组合。同目录的fag_*模块按职责拆分见 op_kernel/flash_attn_grad.py 的模块清单注释模块职责fag_common基本块尺寸、同步 flag、ConstInfo/RunInfofag_memL1/L0/UB 地址与 mutex idfag_tilinghost 契约TilingData / TilingKeyfag_schedule块有效性、swizzle、RunInfofag_buffers片上 tile 声明fag_block_cubeC1..C5 五个 matmulfag_block_vecV1..V6 与 pre/post 编排fag_vector_api寄存器级 VF 微内核attenmaskmask 搬入与 softmax VFfag_kernel两个模板的主循环从 op_kernel/fag_kernel.py 的头部注释可以还原出完整的片上数据流与数学公式一一对应MM2(QK^T)-UB - V2(softmax-P) - V4(P castND2NZ-L1) MM1(dOV^T)-UB - V3(dS(dP-sfmg)*P castND2NZ-L1) - MM5(P^TdO-dV) - MM3(dSK-dQ) - MM4(dS^TQ-dK) V1 计算 softmaxGradFront sfmg rowsum(dy * y)取自前向输出 y即Cube 侧完成 5 个矩阵乘C1..C5Vector 侧完成 6 个向量/元素级阶段V1..V6CV 之间通过手动 cross-core set/wait 同步并区分正向/反向信号对。实现中还处理了一个典型的布局陷阱tensor_k因 MM2 的转置读被推导为 DN 布局而 MM3 需要非转置的 K因此 MM3 使用独立的tensor_k_nt视图指向同一 buffer由于平台不支持 ND2ZNMM4/MM5 左矩阵的转置dS^T / P^T显式走 transpose(UB) - ND2NZ - L1 的路径。5.3 torch 扩展绑定链路torch_extension/flash_attn_grad.py 定义 PyTorch 算子 schematorch.library注册与 Meta 实现并以PrivateUse1后端注册在torch.ops.cann_ops_transformer.flash_attn_grad名下其FlashAttnGradOpBuilder指定 C 源文件为csrc/attention/flash_attn_grad.cpp。torch_extension/csrc/flash_attn_grad.cpp 是 C 扩展层先TORCH_CHECK校验 6 个必选 Tensor再在 NPU 设备上为dq/dk/dv分配输出最终通过ACLNN_CMD(aclnnInnerFlashAttnGrad, ...)调用底层 aclnn 接口并返回三元组。5.4 二进制算子配置op_host/config/ascend950/flash_attn_grad_binary.json 按 dtype 注册了两份二进制FlashAttnGrad_bf16与FlashAttnGrad_fp16。两份配置的 13 个输入、3 个输出、9 个属性签名完全一致仅在 dtype 上区分bf16/fp16所有张量均为 ND 格式、FormatAgnostic匹配模式、shape 为-2动态 shape。六、aclnn 接口说明本算子 aclnn 接口不对外开放aclnnInnerFlashAttnGradGetWorkspaceSize/aclnnInnerFlashAttnGrad以 inner 符号编入自定义包libcust_opapi.so仅供cann_ops_transformertorch 扩展内部调用不在op_api/include/aclnnop安装头文件中导出也不提供 aclnn 调用示例。对外统一使用 torch 接口cann_ops_transformer.flash_attn_grad。因此在实际项目中集成该算子时应始终走「custom 包 torch 扩展 whl」的组合安装路径见第三节并以flash_attn_metadata - flash_attn - flash_attn_grad的联合调用模式接入训练反向传播链路。【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考