CANN ops-transformer FlashAttnGrad 算子详解:Flash Attention 反向梯度计算的原理、构建与 Torch 接口实战
2026/9/20 20:16:41 网站建设 项目流程

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_lseattn_out)和上游梯度(do/dout),计算 Q/K/V 的梯度dq/dk/dv,从而避免在反向阶段重新计算完整注意力矩阵,是 Flash Attention 训练链路中降低显存占用、提升训练吞吐的关键一环。

前向计算可表示为:

$$ S=scale\cdot QK^T $$

$$ P_{ij}=\exp(S_{ij}-softmax_lse_i) $$

$$ Y=PV $$

其中 $scale$ 的取值规则为:softmax_scale属性非 0 时取softmax_scale,为 0 时取 $scale=1/\sqrt{D}$(D 为 Query/Key 的 head dim)。

反向计算依据链式法则展开,README.md 与接口文档 torchapi_flash_attn_grad.md 给出的完整公式为:

$$ dV=P^TdY $$

$$ dP=dYV^T $$

$$ sfmg=rowsum(dY\odot Y) $$

$$ dS=P\odot(dP-sfmg) $$

$$ dQ=scale\cdot(dS\cdot K) $$

$$ dK=scale\cdot(dS^T\cdot Q) $$

其中 $Q$、$K$、$V$ 分别对应输入qkv,$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 dim,Dv 表示 Value/输出 的 head dim。T1 表示所有 batch 中 Query 序列长度的累加和,T2 表示所有 batch 中 Key/Value 序列长度的累加和。GQA 场景满足 N1 是 N2 的整数倍。

二、输入、输出与属性

2.1 输入(13 个)

名称类型必选说明
qBF16/FP16Query tensor
kBF16/FP16Key tensor
vBF16/FP16Value tensor
doBF16/FP16上游梯度
attn_outBF16/FP16前向注意力输出
softmax_lseFP32前向 softmax LSE
cu_seqlens_qINT32TND layout 累积序列长度
cu_seqlens_kvINT32TND layout 累积序列长度
seqused_qINT32实际使用的 Q 序列长度
seqused_kvINT32实际使用的 KV 序列长度
sinksFP32Sink tensor
attn_maskINT8Attention mask(mask_mode=3/4 时使用)
metadataINT32FAG 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_lsesinks为 FP32,attn_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的 shape,dtype 统一取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: window
win_leftInt-1window mask 左窗口,-1 表示正无穷
win_rightInt-1window mask 右窗口,-1 表示正无穷
max_seqlen_qInt-1Q 最大序列长度,-1 表示由输入 shape 推导
max_seqlen_kvInt-1KV 最大序列长度,-1 表示由输入 shape 推导
layout_qString"BSND"Q 布局(BSND/BNSD/TND)
layout_kvString"BSND"KV 布局,必须与 layout_q 相同
layout_outString"BSND"输出布局,必须与 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 Start)

3.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 --soc=ascend950 --ops=flash_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 --soc=ascend950 --ops=flash_attn_grad,flash_attn_metadata -j16

3.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 --soc=ascend950 # 安装 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_enabled=True,并保证与主算子的 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 950DT;Atlas 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_q=None, cu_seqlens_kv=None, seqused_q=None, seqused_kv=None, sinks=None, attn_mask=None, metadata=None, softmax_scale=0.0, mask_mode=0, win_left=-1, win_right=-1, max_seqlen_q=-1, max_seqlen_kv=-1, layout_q="BSND", layout_kv="BSND", layout_out="BSND" ) -> (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/2,kernel 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_lse=True时返回的 log-sum-exp 结果float32NDBSND/BNSD:(B, N1, S1);TND:(N1, T1)
cu_seqlens_qTensor可选TND 布局下 Q 的累积序列长度,第一个元素必须为 0,最后一个元素等于 T1。layout_q为 TND 时必须传入,非 TND 时不支持传入。默认 Noneint32ND(B+1,)
cu_seqlens_kvTensor可选TND 布局下 KV 的累积序列长度,第一个元素必须为 0,最后一个元素等于 T2。layout_kv为 TND 时必须传入,非 TND 时不支持传入。默认 Noneint32ND(B+1,)
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。默认 "BSND"string--
layout_kvstring可选k 和 v 的布局,必须与layout_q相同。默认 "BSND"string--
layout_outstring可选doutattn_out的布局,必须与layout_q相同。默认 "BSND"string--

qkvdoutattn_out及输出dqdkdv支持 float16 和 bfloat16,数据类型必须一致。

4.5 返回值说明

参数名参数类型描述数据类型数据格式维度
dqTensor公式中的 dQ,Query 的梯度bfloat16/float16NDshape 与q相同
dkTensor公式中的 dK,Key 的梯度bfloat16/float16NDshape 与k相同
dvTensor公式中的 dV,Value 的梯度bfloat16/float16NDshape 与v相同

4.6 约束说明

  • 接口支持以下组合:
    • 数据类型为 float16 或 bfloat16;
    • layout_qlayout_kvlayout_out支持 BSND、BNSD、TND,且必须相同;
    • Dense Attention:mask_mode=0attn_mask=Nonewin_left=-1win_right=-1
    • TND 布局下cu_seqlens_qcu_seqlens_kv必须传入;非 TND 布局下不支持传入。
  • qkvdoutattn_out的数据类型必须一致。
  • B、S1、S2、N1、N2、D 和 Dv 必须为正数,其中 B 的取值范围为 (0, 65536)。
  • N1 必须能被 N2 整除,支持 MHA(N1=N2)和 GQA(N1>N2)。
  • Query 和 Key 的 head dim 均为 D;Value、doutattn_out的 head dim 均为 Dv,并满足0 < Dv <= D <= 192
  • 与当前仓库中的flash_attn联合调用时,前向接口还要求 D=Dv 且 D 取 64、128 或 256;结合本接口 D 不超过 192 的约束,联合调用当前支持 D=Dv=64 或 128。
  • softmax_lse必须为 float32,shape 为 (B, N1, S1)(BSND/BNSD 场景)或 (N1, T1)(TND 场景)。
  • 所有输入的数据格式均为 ND。算子注册了AutoContiguous,传入非连续 Tensor 时由框架转换为连续 Tensor。
  • metadata必须由flash_attn_metadata生成,并设置is_grad_enabled=True。生成 metadata 和调用本接口时,N1、N2、D、B、S1、S2、layout 及 mask 相关参数必须一致,否则行为未定义。
  • is_grad_enabled=True生成的 metadata 同时包含正向和反向任务切分数据,前向flash_attn和反向flash_attn_grad均可使用同一份 metadata,无需分别生成。
  • softmax_lseattn_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_q=None, cu_seqlens_kv=None, seqused_q=None, seqused_kv=None, batch_size=None, max_seqlen_q=None, max_seqlen_kv=None, mask_mode=None, win_left=None, win_right=None, layout_q=None, layout_kv=None, layout_out=None, is_grad_enabled=False ) -> 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(B+1,)
cu_seqlens_kvTensor可选TND 布局下 KV 的累积序列长度。layout_kv为 TND 时必须传入,非 TND 时不支持传入。默认 Noneint32ND(B+1,)
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 完整调用示例(BSND)

flash_attn_metadataflash_attnflash_attn_grad联合调用示例。is_grad_enabled=Trueflash_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, dtype=dtype, device="npu") k = torch.randn(B, S2, N2, D, dtype=dtype, device="npu") v = torch.randn(B, S2, N2, Dv, dtype=dtype, device="npu") metadata = cann_ops_transformer.flash_attn_metadata( N1, N2, D, batch_size=B, max_seqlen_q=S1, max_seqlen_kv=S2, mask_mode=0, win_left=-1, win_right=-1, layout_q="BSND", layout_kv="BSND", layout_out="BSND", is_grad_enabled=True, ) attn_out, softmax_lse = cann_ops_transformer.flash_attn( q, k, v, metadata=metadata, softmax_scale=scale, mask_mode=0, win_left=-1, win_right=-1, max_seqlen_q=S1, max_seqlen_kv=S2, layout_q="BSND", layout_kv="BSND", layout_out="BSND", return_softmax_lse=True, ) dout = torch.randn_like(attn_out) dq, dk, dv = cann_ops_transformer.flash_attn_grad( q, k, v, dout, attn_out, softmax_lse, metadata=metadata, softmax_scale=scale, mask_mode=0, win_left=-1, win_right=-1, max_seqlen_q=S1, max_seqlen_kv=S2, layout_q="BSND", layout_kv="BSND", layout_out="BSND", ) torch_npu.npu.synchronize() assert dq.shape == q.shape assert dk.shape == k.shape assert dv.shape == v.shape

该示例展示的是 GQA 场景(N1=8、N2=2),符合"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 入口,注释明确其职责为"解析 -> 校验 -> 规划 -> 写回":

  1. ParsePlatform解析平台信息(AIV/AIC 核数、L2 缓存大小、libapi workspace 大小);
  2. FlashAttnGradCheck::CheckParams对输入参数做合法性校验;
  3. ParseFlashAttnGradInfo解析算子信息;
  4. 构造FlashAttnGradTilingRegbase并调用DoTiling完成分核规划。

平台信息还会在编译期通过TilingParseForFlashAttnGrad缓存进FlashAttnGradCompileInfo,供图模式下 Tiling 阶段拿不到PlatformInfo时回退使用。分核策略进一步拆分为 plan(op_host/plan/下的tiling_keytiling_plantiling_routetiling_scheduletiling_swizzletiling_workspace等模块)、checkers(属性/输入/shape/feature 校验)、info(tiling 信息解析)三层,职责边界清晰。

TilingKey 的位域编码在 op_host/plan/flash_attn_grad_tiling_key.h 中定义,与 kernel 侧FlashAttnGradTilingKey保持一致:

bit[1:0] template 0=BN2GS1S2, 1=BN2, 2=未用, 3=BN2S2(保留,勿复用) bit2 layout 0=非TND, 1=TND bit[4:3] mask_mode host 存 0/1/2,kernel values=[0,3,4] bit5 swizzle 仅 template=0 有效 bit[7:6] d_align 0=64, 1=128, 2=192 bit8 dv_align 0=128, 1=192 bit9 is_bn2_multiblk 仅 template=1 有效 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_bn2gs1s2(template=0):通用切分模板,cube 领先 vector 两个 task(PRELOAD_TIMES=3),dQ 走 fp32 workspace + atomicAdd,dK/dV 在 L0C 跨 s1 累加;同时接收cu_seqlens_q/kvseqused_q/kvsinks等全部可选输入,支持 GQA 与 D=192。
  • flash_attn_grad_bn2(template=1):一个核独占一个 (b, n2) head,1-ahead ping-pong,无 pre/post;不接 GQA / D=192,支持 mask,禁止 BN2+swizzle 组合。

同目录的fag_*模块按职责拆分(见 op_kernel/flash_attn_grad.py 的模块清单注释):

模块职责
fag_common基本块尺寸、同步 flag、ConstInfo/RunInfo
fag_memL1/L0/UB 地址与 mutex id
fag_tilinghost 契约(TilingData / TilingKey)
fag_schedule块有效性、swizzle、RunInfo
fag_buffers片上 tile 声明
fag_block_cubeC1..C5 五个 matmul
fag_block_vecV1..V6 与 pre/post 编排
fag_vector_api寄存器级 VF 微内核
attenmaskmask 搬入与 softmax VF
fag_kernel两个模板的主循环

从 op_kernel/fag_kernel.py 的头部注释可以还原出完整的片上数据流,与数学公式一一对应:

MM2(Q@K^T)->UB -> V2(softmax->P) -> V4(P cast+ND2NZ->L1) MM1(dO@V^T)->UB -> V3(dS=(dP-sfmg)*P cast+ND2NZ->L1) -> MM5(P^T@dO->dV) -> MM3(dS@K->dQ) -> MM4(dS^T@Q->dK) V1 计算 softmaxGradFront sfmg = rowsum(dy * y),取自前向输出 y

即:Cube 侧完成 5 个矩阵乘(C1..C5),Vector 侧完成 6 个向量/元素级阶段(V1..V6);CV 之间通过手动 cross-core set/wait 同步,并区分正向/反向信号对。实现中还处理了一个典型的布局陷阱:tensor_k因 MM2 的转置读被推导为 DN 布局,而 MM3 需要非转置的 K,因此 MM3 使用独立的tensor_k_nt视图指向同一 buffer;由于平台不支持 ND2ZN,MM4/MM5 左矩阵的转置(dS^T / P^T)显式走 transpose(UB) -> ND2NZ -> L1 的路径。

5.3 torch 扩展绑定链路

torch_extension/flash_attn_grad.py 定义 PyTorch 算子 schema(torch.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_bf16FlashAttnGrad_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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询