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 训练链路中降低显存占用、提升训练吞吐的关键一环。
前向计算可表示为:
$$ 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$ 分别对应输入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 dim,Dv 表示 Value/输出 的 head dim。T1 表示所有 batch 中 Query 序列长度的累加和,T2 表示所有 batch 中 Key/Value 序列长度的累加和。GQA 场景满足 N1 是 N2 的整数倍。
二、输入、输出与属性
2.1 输入(13 个)
| 名称 | 类型 | 必选 | 说明 |
|---|---|---|---|
| q | BF16/FP16 | 是 | Query tensor |
| k | BF16/FP16 | 是 | Key tensor |
| v | BF16/FP16 | 是 | Value tensor |
| do | BF16/FP16 | 是 | 上游梯度 |
| attn_out | BF16/FP16 | 是 | 前向注意力输出 |
| softmax_lse | FP32 | 是 | 前向 softmax LSE |
| cu_seqlens_q | INT32 | 否 | TND layout 累积序列长度 |
| cu_seqlens_kv | INT32 | 否 | TND layout 累积序列长度 |
| seqused_q | INT32 | 否 | 实际使用的 Q 序列长度 |
| seqused_kv | INT32 | 否 | 实际使用的 KV 序列长度 |
| sinks | FP32 | 否 | Sink tensor |
| attn_mask | INT8 | 否 | Attention mask(mask_mode=3/4 时使用) |
| metadata | INT32 | 否 | 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为 FP32,attn_mask为 INT8,其余辅助索引张量均为 INT32。
2.2 输出(3 个)
| 名称 | 类型 | 说明 |
|---|---|---|
| dq | BF16/FP16 | Query 梯度,shape 与q相同 |
| dk | BF16/FP16 | Key 梯度,shape 与k相同 |
| dv | BF16/FP16 | Value 梯度,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_scale | Float | 0.0 | softmax 缩放因子,0.0 表示 1/sqrt(d) |
| mask_mode | Int | 0 | 0: 全计算, 3: causal, 4: window |
| win_left | Int | -1 | window mask 左窗口,-1 表示正无穷 |
| win_right | Int | -1 | window mask 右窗口,-1 表示正无穷 |
| max_seqlen_q | Int | -1 | Q 最大序列长度,-1 表示由输入 shape 推导 |
| max_seqlen_kv | Int | -1 | KV 最大序列长度,-1 表示由输入 shape 推导 |
| layout_q | String | "BSND" | Q 布局(BSND/BNSD/TND) |
| layout_kv | String | "BSND" | KV 布局,必须与 layout_q 相同 |
| layout_out | String | "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_MASK | 0 | 全计算模式(默认值) |
CAUSAL | 3 | Causal 模式 |
SLIDING_WINDOW | 4 | Sliding 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 参数说明
| 参数名 | 参数类型 | 可选/必选 | 描述 | 数据类型 | 数据格式 | 维度 |
|---|---|---|---|---|---|---|
| q | Tensor | 必选 | 公式中的Q | bfloat16/float16 | ND | BSND:(B, S1, N1, D);BNSD:(B, N1, S1, D);TND:(T1, N1, D) |
| k | Tensor | 必选 | 公式中的K | bfloat16/float16 | ND | BSND:(B, S2, N2, D);BNSD:(B, N2, S2, D);TND:(T2, N2, D) |
| v | Tensor | 必选 | 公式中的V | bfloat16/float16 | ND | BSND:(B, S2, N2, Dv);BNSD:(B, N2, S2, Dv);TND:(T2, N2, Dv) |
| dout | Tensor | 必选 | 公式中的dY,前向输出的上游梯度 | bfloat16/float16 | ND | BSND:(B, S1, N1, Dv);BNSD:(B, N1, S1, Dv);TND:(T1, N1, Dv) |
| attn_out | Tensor | 必选 | 公式中的Y,即前向接口返回的注意力输出 | bfloat16/float16 | ND | shape 与dout相同 |
| softmax_lse | Tensor | 必选 | 前向接口在return_softmax_lse=True时返回的 log-sum-exp 结果 | float32 | ND | BSND/BNSD:(B, N1, S1);TND:(N1, T1) |
| cu_seqlens_q | Tensor | 可选 | TND 布局下 Q 的累积序列长度,第一个元素必须为 0,最后一个元素等于 T1。layout_q为 TND 时必须传入,非 TND 时不支持传入。默认 None | int32 | ND | (B+1,) |
| cu_seqlens_kv | Tensor | 可选 | TND 布局下 KV 的累积序列长度,第一个元素必须为 0,最后一个元素等于 T2。layout_kv为 TND 时必须传入,非 TND 时不支持传入。默认 None | int32 | ND | (B+1,) |
| seqused_q | Tensor | 可选 | 每个 batch 实际使用的 Q 序列长度。默认 None | int32 | ND | (B,) |
| seqused_kv | Tensor | 可选 | 每个 batch 实际使用的 KV 序列长度。默认 None | int32 | ND | (B,) |
| sinks | Tensor | 可选 | Sink 参数,用于改善自注意力计算的数值稳定性。默认 None | float32 | ND | (N1,) |
| attn_mask | Tensor | 可选 | 掩码矩阵。默认 None | int8 | ND | (2048, 2048) |
| metadata | Tensor | 可选 | flash_attn_metadata生成的 FAG 任务切分数据。schema 中为可选参数,但实际调用时必须传入 | int32 | ND | shape 根据 batch 大小和 N2 动态计算 |
| softmax_scale | float | 可选 | Softmax 缩放系数。默认 0.0,表示使用 $1/\sqrt{D}$ | float32 | - | - |
| mask_mode | int/MaskMode | 可选 | 掩码模式,支持传入枚举或对应 int 值。默认 0 | int32 | - | - |
| win_left | int | 可选 | Window mask 左窗口,值需大于等于 -1,-1 表示正无穷。默认 -1 | int32 | - | - |
| win_right | int | 可选 | Window mask 右窗口,值需大于等于 -1,-1 表示正无穷。默认 -1 | int32 | - | - |
| max_seqlen_q | int | 可选 | Q 最大序列长度,必须大于等于 -1。BSND/BNSD 场景可保持默认 -1,由输入 shape 推导 S1 | int32 | - | - |
| max_seqlen_kv | int | 可选 | KV 最大序列长度,必须大于等于 -1。BSND/BNSD 场景可保持默认 -1,由输入 shape 推导 S2 | int32 | - | - |
| layout_q | string | 可选 | q 的布局,支持 BSND、BNSD、TND。默认 "BSND" | string | - | - |
| layout_kv | string | 可选 | k 和 v 的布局,必须与layout_q相同。默认 "BSND" | string | - | - |
| layout_out | string | 可选 | dout和attn_out的布局,必须与layout_q相同。默认 "BSND" | string | - | - |
q、k、v、dout、attn_out及输出dq、dk、dv支持 float16 和 bfloat16,数据类型必须一致。
4.5 返回值说明
| 参数名 | 参数类型 | 描述 | 数据类型 | 数据格式 | 维度 |
|---|---|---|---|---|---|
| dq | Tensor | 公式中的 dQ,Query 的梯度 | bfloat16/float16 | ND | shape 与q相同 |
| dk | Tensor | 公式中的 dK,Key 的梯度 | bfloat16/float16 | ND | shape 与k相同 |
| dv | Tensor | 公式中的 dV,Value 的梯度 | bfloat16/float16 | ND | shape 与v相同 |
4.6 约束说明
- 接口支持以下组合:
- 数据类型为 float16 或 bfloat16;
layout_q、layout_kv、layout_out支持 BSND、BNSD、TND,且必须相同;- Dense Attention:
mask_mode=0、attn_mask=None、win_left=-1、win_right=-1; - TND 布局下
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 整除,支持 MHA(N1=N2)和 GQA(N1>N2)。
- Query 和 Key 的 head dim 均为 D;Value、
dout和attn_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_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_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_q | int | 必选 | Query head 数,即 N1 | int32 | - | - |
| num_heads_kv | int | 必选 | Key/Value head 数,即 N2 | int32 | - | - |
| head_dim | int | 必选 | Query/Key 的 head dim,即 D | int32 | - | - |
| cu_seqlens_q | Tensor | 可选 | TND 布局下 Q 的累积序列长度。layout_q为 TND 时必须传入,非 TND 时不支持传入。默认 None | int32 | ND | (B+1,) |
| cu_seqlens_kv | Tensor | 可选 | TND 布局下 KV 的累积序列长度。layout_kv为 TND 时必须传入,非 TND 时不支持传入。默认 None | int32 | ND | (B+1,) |
| seqused_q | Tensor | 可选 | 每个 batch 实际使用的 Q 序列长度。默认 None | int32 | ND | (B,) |
| seqused_kv | Tensor | 可选 | 每个 batch 实际使用的 KV 序列长度。默认 None | int32 | ND | (B,) |
| batch_size | int | 可选 | batch 大小。BSND/BNSD 场景必须传入实际 B,取值范围 (0, 65536) | int32 | - | - |
| max_seqlen_q | int | 可选 | Q 最大序列长度。BSND/BNSD 场景必须传入实际 S1,且大于 0 | int32 | - | - |
| max_seqlen_kv | int | 可选 | KV 最大序列长度。BSND/BNSD 场景必须传入实际 S2,且大于 0 | int32 | - | - |
| mask_mode | int/MaskMode | 可选 | 必须与flash_attn_grad一致 | int32 | - | - |
| win_left | int | 可选 | 必须与flash_attn_grad一致 | int32 | - | - |
| win_right | int | 可选 | 必须与flash_attn_grad一致 | int32 | - | - |
| layout_q | string | 可选 | 必须与flash_attn_grad一致 | string | - | - |
| layout_kv | string | 可选 | 必须与flash_attn_grad一致 | string | - | - |
| layout_out | string | 可选 | 必须与flash_attn_grad一致 | string | - | - |
| is_grad_enabled | bool | 可选 | 是否生成反向算子所需的 metadata。调用flash_attn_grad前必须设置为 True。默认 False | bool | - | - |
返回的metadata为 int32、ND 格式的一维 Tensor,长度根据 batch 大小和 N2 动态计算。
4.8 完整调用示例(BSND)
flash_attn_metadata、flash_attn和flash_attn_grad联合调用示例。is_grad_enabled=True的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, 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.py | pypto-pro kernel 实现(BN2GS1S2 模板) |
op_host/flash_attn_grad_tiling.cpp | tiling 实现(tilingkey 编码、workspace 布局) |
op_host/flash_attn_grad_def.cpp | 算子定义(op_proto/输入输出/属性) |
op_host/flash_attn_grad_infershape.cpp | shape/dtype 推导 |
op_host/config/ascend950/flash_attn_grad_binary.json | 二进制算子配置 |
torch_extension/flash_attn_grad.py | PyTorch 算子 schema 与 python 绑定 |
torch_extension/csrc/flash_attn_grad.cpp | PyTorch 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时回退使用。分核策略进一步拆分为 plan(op_host/plan/下的tiling_key、tiling_plan、tiling_route、tiling_schedule、tiling_swizzle、tiling_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/kv、seqused_q/kv、sinks等全部可选输入,支持 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_mem | L1/L0/UB 地址与 mutex id |
fag_tiling | host 契约(TilingData / TilingKey) |
fag_schedule | 块有效性、swizzle、RunInfo |
fag_buffers | 片上 tile 声明 |
fag_block_cube | C1..C5 五个 matmul |
fag_block_vec | V1..V6 与 pre/post 编排 |
fag_vector_api | 寄存器级 VF 微内核 |
attenmask | mask 搬入与 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_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),仅供参考