拓冰建站拓冰建站
首页 / 资讯中心 / 正文

PyPTO-Gym 中的 FlashAttentionScoreGrad 反向算子实现:数学原理、PyPTO Kernel 与测试验证全解析

PyPTO-Gym 中的 FlashAttentionScoreGrad 反向算子实现数学原理、PyPTO Kernel 与测试验证全解析【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym本指南围绕 PyPTO-Gym 仓库中 FlashAttentionScoreGrad 算子的 PyPTO 实现展开完整解读 Flash Attention 反向传播的数学推导、基于 Online Softmax 的梯度重算策略、PyPTO JIT Kernel 的分块实现细节以及从环境准备到 golden 精度对比的完整测试流程。阅读后你将掌握如何在 Ascend NPU 上编写与验证一个 Flash Attention 反向算子并能将同样的两趟two-pass分块模式复用到其他自注意力类算子的梯度计算中。产品支持情况该实现已在以下硬件/软件平台上验证支持依据 README.md产品支持情况Ascend 950PR支持Atlas A3 训练系列产品 / Atlas A3 推理系列产品支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品支持需要说明的是同一目录下的前向算子 flash_attention_score 明确标注Ascend 950PR不支持而本反向算子对 950PR 是支持的二者能力边界不同迁移使用时需分别确认。算子概述与数学原理FlashAttentionScoreGrad 计算 Flash Attention 前向传播的反向梯度输入为前向传播保存的中间结果softmax_max、softmax_sum、attention_out与输出梯度dY输出为三个梯度张量dQ、dK、dV。前向公式前向注意力计算可写为Y Softmax(Q K^T / sqrt(D)) V其中sqrt(D)的缩放即scale 1/sqrt(D)也可并入分数计算中。前向算子同时输出用于反向的统计量softmax_max每行最大值与softmax_sum每行归一化因子见 flash_attention_score_impl.py 中flash_attention_score_kernel_npu的pypto.assemble(m_v, [out_row, 0], softmax_max)调用。反向公式Online Softmax 重算反向阶段利用前向保存的统计量重算概率矩阵P避免存储全量注意力矩阵这正是 Flash Attention 的核心内存优化P exp(Q K^T * scale - softmax_max) / softmax_sum D sum(dY * attention_out, dim-1, keepdimTrue) dP dY V^T dS P * (dP - D) dV P^T dY dQ dS K * scale dK dS^T Q * scale其中D是损失对attention_out的梯度中沿 HeadDim 维的聚合项在 Flash Attention 的梯度推导中常称为rowsum(dY ⊙ O)或D_idP是损失对概率矩阵P的梯度dS是损失对注意力分数矩阵S的梯度最后将dS分别与K、Q做矩阵乘并乘以scale得到dQ、dK。上述公式同时完整出现在 golden 参考实现 flash_attention_score_grad_golden.py 的模块 docstring 中其注释明确说明与 PyPTO kernels_tile128计算流完全对齐。核心特性该算子的 PyPTO 实现具备以下关键设计动态轴Dynamic AxisBatchB与 SequenceS均为动态轴由pypto.Tensor([pypto.DYN, ...])声明Kernel 在运行时根据输入 shape 推导分块循环次数Online Softmax 重算利用前向保存的softmax_max/softmax_sum重算P无需存储全量[B, N, S, S]注意力矩阵内存复杂度从 O(S²) 降为 O(S)数据类型BF16 输入/输出FP32 中间计算s_ij、p_ij、dP、dS、累加器dq_acc/dk_acc/dv_acc均为 FP32兼顾精度与昇腾硬件 BF16 算力布局BNSD即[Batch, NumHeads, SeqLen, HeadDim]Kernel 内部将 4D 张量展平为 2D 视图[B*N*S, D]以适配分块矩阵乘。目录结构与仓库真实路径README 中给出了算子目录的结构路径为custom/flash_attention_score_grad/系文档作者的工作目录示意。在 PyPTO-Gym 仓库中该算子的真实文件分布如下src/pypto_gym/ops/pypto_tensor/experimental/ops_transformer/flash_attention_score_grad/ ├── flash_attention_score_grad_impl.py # PyPTO JIT kernel wrapper性能优化版 └── README.md # 本文档对应的说明文件 tests/ops/experimental/ops_transformer/flash_attention_score_grad/ ├── flash_attention_score_grad_golden.py # 纯 PyTorch golden 参考实现 └── test_flash_attention_score_grad.py # 测试入口其中实现代码位于 flash_attention_score_grad_impl.pygolden 与测试位于 tests/ops/experimental/ops_transformer/flash_attention_score_grad/。文档中的运行命令路径需相应替换为上述仓库路径。PyPTO Kernel 实现解析Kernel 签名与动态维度推导Kernel 通过pypto.frontend.jit装饰所有输入张量均声明为pypto.Tensor([pypto.DYN, ...])动态第一维 其余静态softmax_max/softmax_sum为 FP32其余为 BF16见 flash_attention_score_grad_impl.pypypto.frontend.jit( runtime_options{ stitch_function_max_num: 128, device_sched_mode: 1, }, pass_options{ cube_l1_reuse_setting: {0: 8}, cube_nbuffer_setting: {0: 4}, } ) def flash_attention_score_grad_kernel( q: pypto.Tensor([pypto.DYN, ...], pypto.DT_BF16), k: pypto.Tensor([pypto.DYN, ...], pypto.DT_BF16), v: pypto.Tensor([pypto.DYN, ...], pypto.DT_BF16), dy: pypto.Tensor([pypto.DYN, ...], pypto.DT_BF16), softmax_max: pypto.Tensor([pypto.DYN, ...], pypto.DT_FP32), softmax_sum: pypto.Tensor([pypto.DYN, ...], pypto.DT_FP32), attention_out: pypto.Tensor([pypto.DYN, ...], pypto.DT_BF16), dq: pypto.Tensor([pypto.DYN, ...], pypto.DT_BF16), dk: pypto.Tensor([pypto.DYN, ...], pypto.DT_BF16), dv: pypto.Tensor([pypto.DYN, ...], pypto.DT_BF16), batch_size: pypto.Tensor([pypto.DYN], pypto.DT_INT32), scale_value: float, num_heads: int, ):Kernel 内部从 shape 推导几何参数b batch_size.shape[0]、total q.shape[0]、head_dim q.shape[1]、s total // b // num_heads即S total / (B × N)。JIT 装饰器上的pass_options与runtime_options是本次实现的性能优化关键cube_l1_reuse_setting{0: 8}让Q常驻 L1cube_nbuffer_setting{0: 4}配置 Cube 多缓冲stitch_function_max_num128与device_sched_mode1控制算子拼接与设备调度。分块Tiling配置模块级常量S_TILE 128优化前为 64增大后减少循环迭代次数Sequence 维被切分为s_loop (s S_TILE - 1) // S_TILE个块。每个块采用如下 tile 配置c_tile [[S_TILE, S_TILE], [head_dim, 256], [S_TILE, S_TILE]] # cube: [M,N],[K1,K2],[M,N] v_tile_s [S_TILE, S_TILE] # vec: [S,S] v_tile_d [S_TILE, head_dim] # vec: [S,D]c_tile为 Cube 单元矩阵乘的 tile shape用于QK^T、dYV^T、dSK、dS^TQ、P^TdY等矩阵乘v_tile_s/v_tile_d为 Vector 单元逐元素/规约的 tile shape用于乘法、减法、exp、div、sum等操作。单块计算compute_tilecompute_tile函数实现一个(s1_tile, s2_tile)块的P_ij与dS_ij计算完整对应反向公式的前四行见 flash_attention_score_grad_impl.pydef compute_tile(q_i, k_j, v_j, dy_i, smax_i, ssum_i, d_i, ...): # S_ij Q_i K_j^T * scale s_ij pypto.matmul(q_i, k_j, pypto.DT_FP32, b_transTrue) s_ij pypto.view(s_ij, [s_tile_size, s_tile_size], [0, 0], valid_shape[actual_s1, actual_s2]) s_ij pypto.mul(s_ij, scale_value) p_ij pypto.exp(pypto.sub(s_ij, smax_i), precision_typepypto.PrecisionType.HIGH_PRECISION) p_ij pypto.div(p_ij, ssum_i) # dP_ij dY_i V_j^T dp_ij pypto.matmul(dy_i, v_j, pypto.DT_FP32, b_transTrue) dp_ij pypto.view(dp_ij, [s_tile_size, s_tile_size], [0, 0], valid_shape[actual_s1, actual_s2]) # dS_ij P_ij * (dP_ij - D_i) ds_ij pypto.mul(p_ij, pypto.sub(dp_ij, d_i)) return p_ij, ds_ij细节说明pypto.matmul(..., pypto.DT_FP32, b_transTrue)以 FP32 累加执行转置矩阵乘Q K^T、dY V^Tpypto.exp(..., precision_typepypto.PrecisionType.HIGH_PRECISION)显式指定高精度 exp降低 BF16 中间计算带来的数值误差pypto.view(..., valid_shape[actual_s1, actual_s2])原生处理尾块Sequence 长度不是 S_TILE 整数倍时的最后一个不完整块其中actual_s1 (s - s1_idx * S_TILE).min(S_TILE)无需 host 端 paddingd_i pypto.sum(pypto.cast(pypto.mul(dy_i, ao_i), pypto.DT_FP32), -1, keepdimTrue)在调用方完成即D sum(dY * attention_out, dim-1, keepdimTrue)其中attention_out为前向输出ao_i来自attention_out张量的分块视图。两趟循环结构Two-Pass整个反向计算采用先算 dQ再算 dK/dV的两趟two-pass结构趟 1计算 dQ外层LOOP_s1_dq遍历 Q 行块内层LOOP_s2_dq遍历 K/V 列块见 flash_attention_score_grad_impl.py对每个 Q 行块计算d_i初始化 FP32 累加器dq_acc内层遍历所有 K/V 列块调用compute_tile得到ds_ijpypto.cast回 BF16 后做dq_tile pypto.matmul(ds_bf16, k_j, pypto.DT_FP32)利用pypto.is_loop_begin/is_loop_end区分首块初始化与末块收尾最后dq_final cast(mul(dq_acc, scale_value), BF16)并pypto.assemble写回dq_2d。趟 2计算 dK 与 dV外层LOOP_s2_dkv遍历 K/V 列块内层LOOP_s1_dkv遍历 Q 行块见 flash_attention_score_grad_impl.pydk_tile pypto.matmul(ds_bf16, q_i, pypto.DT_FP32, a_transTrue) # dK dS^T Q dv_tile pypto.matmul(p_bf16, dy_i, pypto.DT_FP32, a_transTrue) # dV P^T dYa_transTrue表示对左操作数转置ds_ij与p_ij均先 cast 为 BF16 再做矩阵乘累加完成后dk_final cast(mul(dk_acc, scale_value), BF16)dK 需要乘 scaledV 不需要最后pypto.assemble写回dk_2d/dv_2d。两个内层循环都带有unroll_list[8, 4, 2, 1]的循环展开提示用于提升指令级并行。Wrapper 与形状校验flash_attention_score_grad_wrapper是对外暴露的 Python 接口flash_attention_score_grad_impl.py解析query.shape [B, N, S, D]校验num_heads与head_dim一致性不一致时抛出带详细信息的ValueError将 Q/K/V/dY/attention_out 展平为[-1, D]的 2D 连续张量softmax_max/softmax_sum展平为[-1, 8]最后一维为 8 是前向统计量在 NPU 上的对齐存储形式Kernel 通过pypto.view(sm_i_8, [S_TILE, 1], [0, 0])取第 0 列作为 1D 统计量用torch.empty_like分配dq/dk/dv输出构造batch_tensor形状[B]的 int32 张量作为动态 batch 维度的载体调用 JIT kernel返回 reshape 回[B, N, S, D]的三元组(dq, dk, dv)。输入/输出规格以下规格表完整继承自 README.md其中BBatch、NNumHeads、SSeqLen、DHeadDimTensorShapeDType说明query[B, N, S, D]BF16Query 张量key[B, N, S, D]BF16Key 张量value[B, N, S, D]BF16Value 张量dy[B, N, S, D]BF16输出梯度softmax_max[B, N, S, 8]FP32前向 softmax maxsoftmax_sum[B, N, S, 8]FP32前向 softmax sumattention_out[B, N, S, D]BF16前向输出dQ输出[B, N, S, D]BF16Query 梯度dK输出[B, N, S, D]BF16Key 梯度dV输出[B, N, S, D]BF16Value 梯度注意softmax_max/softmax_sum的最后一维为 8对齐存储其有效数据位于第 0 列scale_value通常取1/sqrt(D)与num_heads以标量参数形式传给 Kernel。从 wrapper 的 shape 校验看实现要求query/key/value/dy/attention_out的N、D与传入参数严格一致。运行方法环境准备在装有 CANN 工具链与 NPU 驱动的主机上执行source /usr/local/Ascend/ascend-toolkit/set_env.sh export TILE_FWK_DEVICE_ID0 export PTO_TILE_LIB_CODE_PATH/path/to/pto-isaTILE_FWK_DEVICE_ID指定使用的 NPU 设备号测试代码通过os.environ.get(TILE_FWK_DEVICE_ID, 0)读取默认 0PTO_TILE_LIB_CODE_PATH指向 PyPTO Tile 指令库pto-isa路径需按实际安装位置替换/path/to/pto-isa。验证 Golden纯 PyTorch无需 NPUpython3 tests/ops/experimental/ops_transformer/flash_attention_score_grad/flash_attention_score_grad_golden.pygolden 脚本包含两部分flash_attention_score_grad_golden.pygenerate_forward_data以固定随机种子torch.manual_seed(42)生成 Q/K/V/dY并用_compute_forward分块在线 softmax 前向_BLOCK_Q32、_BLOCK_KV64复算softmax_max/softmax_sum/attention_out与前向算子BLOCK_Q64, BLOCK_KV64保持一致的统计量语义flash_attention_score_grad_golden按S_TILE128分块、与 PyPTO kernel 相同的两趟结构_grad_pass1_dq与_grad_pass2_dkdv计算 dQ/dK/dV因此不仅输出正确连分块数据流都与 kernel 逐块对齐。运行测试# 运行所有测试级别 python3 tests/ops/experimental/ops_transformer/flash_attention_score_grad/test_flash_attention_score_grad.py # 运行指定级别 python3 tests/ops/experimental/ops_transformer/flash_attention_score_grad/test_flash_attention_score_grad.py 0 # 查看可用级别 python3 tests/ops/experimental/ops_transformer/flash_attention_score_grad/test_flash_attention_score_grad.py --list测试入口 test_flash_attention_score_grad.py 支持三个级别与运行模式切换级别配置B, N, S, D说明0(1, 8, 128, 64)最小功能验证默认运行1(2, 8, 128, 64)典型小规模标记为 large test case默认 skip2(2, 8, 256, 64)中等规模标记为 large test case默认 skip命令行参数level位置参数0/1/2指定测试级别缺省则按 0→1→2 顺序运行全部--list列出可用级别--run_mode {npu,sim}默认npu选择sim时在 CPU 上以 SIM 模式运行此时devicecpu注意 Kernel 本身仍需 NPU 环境。run_test的验证流程为生成前向数据 → 调用flash_attention_score_grad_wrapper得到 PyPTO 结果 → 调用flash_attention_score_grad_golden得到参考结果 → 对 dQ/dK/dV 分别计算 max diff 并打印 → 用numpy.testing.assert_allclose以rtol1e-2, atol2e-2判定精度全部通过则输出[PRECISION_PASS]。由于 kernel 内部 exp 采用高精度计算且累加在 FP32 中进行测试用例seq_len128/256中 Sequence 维恰好为 S_TILE 的整数倍尾块路径由valid_shape机制兜底。简化说明与已知边界README 明确说明初版实现聚焦核心 dQ/dK/dV 计算跳过以下可选功能这些均不影响当前测试用例的正确性但迁移到实际模型前需评估PSE位置偏移编码不处理位置相关分数偏置Dropout不接收 dropout masksoftmax_sum无需按1/keep_prob缩放Attention Mask不支持下三角/因果等掩码模式RoPE旋转位置编码假设 Q/K 已由上游完成位置编码FP8 量化输入输出保持 BF16稀疏注意力模式按稠密分块计算未利用稀疏性。对比前向算子 flash_attention_score 已支持 PSE/Dropout/GQA 的现状可以推断反向算子的这些扩展大概率是后续演进方向。此外实现约束还包括softmax_max/softmax_sum需要前向算子以[B, N, S, 8]有效列在第 0 列的存储格式输出且 Kernel 依赖num_heads、head_dim等编译期几何信息不同 shape 需要重新 JIT 编译。小结FlashAttentionScoreGrad 的 PyPTO 实现完整展示了在 Ascend NPU 上落地一个 Flash Attention 反向算子的标准套路用 Online Softmax 统计量重算概率矩阵以省内存、用两趟分块循环分别累积 dQ 与 dK/dV、用 FP32 中间累加 高精度 exp 保证 BF16 精度、用valid_shape原生处理尾块、用 golden 分块对齐实现做逐块精度验证。这套模式不仅适用于本算子也可直接复用于 MLA 梯度、稀疏注意力梯度等其他自注意力类反向算子仓库中同类实现可见 flash_attention_score_grad_tnd 等目录是 PyPTO-Gym 中值得精读的算子开发范例。【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门