FEATURED · 精选文章

PyPTO-Gym 算子骨架 SK-16 实战:单趟 Attention 反向传播(FA Grad)的设计与实现

发布时间 / 2026/9/20 2:05:03
来源 / 创域科博编辑部
栏目 / 资讯中心
PyPTO-Gym 算子骨架 SK-16 实战:单趟 Attention 反向传播(FA Grad)的设计与实现 PyPTO-Gym 算子骨架 SK-16 实战单趟 Attention 反向传播FA Grad的设计与实现【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym导读SK-16Single-Pass Attention Backward即 FA Grad是 PyPTO-Gym 仓库中cannbot-skills/ops/pypto-op-design/patterns/skeletons/SK-16-attention-backward.md定义的一套 Attention 反向传播算子骨架它要求输出 dQ/dK/dV 三路梯度并直接消费前向预计算的 softmax 统计量 l/m在单趟single-pass双层 tile 循环内完成全部梯度计算。读完本文你将掌握这套 C1→V1→C2→V2 的 C-V-C-V 循环排布方式、完整可复刻的 kernel 骨架代码、14 条关键编码特征与开箱性能优化配置并通过仓库中的 A3 达标实现与测试用例理解其落地形态。一、为什么需要单趟 Attention 反向传播骨架Attention 反向传播算子需要同时输出 dQ/dK/dV 三路梯度其循环方向存在天然冲突dQ 的最优循环方向是 q-outer / kv-inner而 dK/dV 的最优循环方向是 kv-outer / q-inner。这一方向冲突决定了算子结构选型。在 SK-01Online Flash Attention 的「反向传播结构决策」一节中给出了两条路线结构循环组织梯度写回softmax 统计量QK^T 次数单 Pass默认推荐batch → head → q_tile → kv_tile单重嵌套三路梯度在同一块对内完成pypto.atomic_add直写 GM 输出host 侧torch.zeros预清零输入含 l/m 时直接消费零重算无则块内在线自算每 (q, kv) 块对 1 次双 Pass安全回退Pass1 q-outer 计算 dQPass2 kv-outer 计算 dK/dVUB 累积器 三路分支 assemble尾块写回Pass1 自算后经 GM scratch 传递至 Pass2每块对 2~3 次决策规则的核心是默认选单 Pass用atomic_add消解循环方向冲突从而消除双 Pass 的 QK^T 重算与 GM scratch 往返仅在平台不支持atomic_add、精度语义要求确定性归约顺序、或输出 buffer 不可预清零时才回退双 Pass且必须写明合法理由。SK-16 正是「单 Pass」决策的具体实现骨架二者在骨架匹配阶段联动先做 SK-01 的反向传播结构决策选单 Pass 即落入 SK-16。二、CV 排布C1 → V1 → C2 → V2 循环模式SK-16 的硬件级计算排布为C1QK^T 与 dOV^T 双 MatMul产出 S 与 dPV1D 行和 softmax 归一化 dS 推导C2dV / dQ / dK 三 MatMulV2scale 缩放 atomic_add三路写回。整个 C-V-C-V 序列在 s2KV tile循环内重复执行。注意与 SK-01 前向的 C1-V1-C2 相比反向多出 C2→V2 一段因为三路梯度 MatMul 产出后还需要缩放与原子累加写回因此sg_set_scope的链合分段也从两段变为 V11 / V22 两段。三、骨架结构完整代码与逐段解读以下骨架代码直接继承自 SK-16-attention-backward.md其中s1_tile、s2_tile、c_tile、v_tile_s、v_tile_d由设计阶段确定实现时定义为编译期整数常量及常量列表不作为 kernel 的运行时输入。# s1_tile、s2_tile、c_tile、v_tile_s、v_tile_d 由设计确定。 # 实现时定义为编译期整数常量及常量列表不作为 kernel 的运行时输入。 pypto.frontend.jit( runtime_options{ stitch_function_max_num: 512, device_sched_mode: 3, ready_on_host_tensors: [actual_q, actual_kv], # varlen 控制流必配 max_workspace_kb: platform_value, # 大 workspace (A3 实测 25170624) }, pass_options{ cube_l1_reuse_setting: {-1: 1, 0: 4}, }, ) def attention_backward_kernel(q, k, v, o, do, l_input, m_input, dq, dk, dv, actual_q, actual_kv): num_heads q.shape[1] head_dim q.shape[2] hidden_dim num_heads * head_dim total q.shape[0] scale 1.0 / (head_dim ** 0.5) # Python float, 编译期折叠 pypto.experimental.set_operation_options(combine_axisTrue) # 3D → 2D inplace reshape零拷贝便于按 [seq, hidden_dim] 切片 q_2d pypto.reshape(q, [total, hidden_dim], inplaceTrue) ... # k/v/o/do 同 l_2d pypto.reshape(l_input, [total, num_heads], inplaceTrue) m_2d pypto.reshape(m_input, [total, num_heads], inplaceTrue) for b_idx in pypto.loop(batch_size, nameLOOP_b): # Loop: Batch (pypto.loop) q_start actual_q[b_idx] # per-batch 动态 offset s1 actual_q[b_idx 1] - q_start # per-batch 动态 seqlen kv_start actual_kv[b_idx] s2 actual_kv[b_idx 1] - kv_start s1_loop (s1 s1_tile - 1) // s1_tile # 循环次数动态推导 s2_loop (s2 s2_tile - 1) // s2_tile for n_idx in pypto.loop(num_heads, nameLOOP_n): # Loop: Head (pypto.loop) h_ofs n_idx * head_dim for s1_idx in pypto.loop(s1_loop, nameLOOP_s1): # Loop: Q tile for s2_idx in pypto.loop(s2_loop, nameLOOP_s2): # Loop: KV tile (C-V-C-V) s1_off q_start s1_idx * s1_tile actual_s1 (s1 - s1_idx * s1_tile).min(s1_tile) s2_off kv_start s2_idx * s2_tile actual_s2 (s2 - s2_idx * s2_tile).min(s2_tile) # view valid_shape无手动掩码 q_i pypto.view(q_2d, [s1_tile, head_dim], [s1_off, h_ofs], valid_shape[actual_s1, head_dim]) ... # k_j/v_j/do_i/o_i 同; m_i/l_i view 自 m_2d/l_2d [s1_tile, 1] # C1: 双 MatMulS QK^T, dP dOV^T pypto.set_cube_tile_shapes(c_tile[0], c_tile[1], c_tile[2]) s_ij pypto.matmul(q_i, k_j, pypto.DT_FP32, b_transTrue) dp_ij pypto.matmul(do_i, v_j, pypto.DT_FP32, b_transTrue) # V1: D softmax 归一化 dSsg_set_scope1 链合 pypto.set_pass_options(sg_set_scope1) pypto.set_vec_tile_shapes(v_tile_s[0], v_tile_s[1]) d_i pypto.sum(pypto.mul(cast(o_i, FP32), cast(do_i, FP32)), -1, keepdimTrue) s_ij pypto.mul(s_ij, scale) p_ij pypto.exp(pypto.sub(s_ij, m_i)) # m_i 直用前向输入 p_ij pypto.div(p_ij, l_i, # l_i 直用前向输入 precision_typepypto.PrecisionType.INTRINSIC) ds_ij pypto.mul(p_ij, pypto.sub(dp_ij, d_i)) ds_bf16 pypto.cast(ds_ij, pypto.DT_BF16) p_bf16 pypto.cast(p_ij, pypto.DT_BF16) pypto.set_pass_options(sg_set_scope-1) # C2: 三 MatMuldV/dQ/dK pypto.set_cube_tile_shapes(c_tile[0], c_tile[1], c_tile[2]) dv_tile pypto.matmul(p_bf16, do_i, pypto.DT_FP32, a_transTrue) dq_tile pypto.matmul(ds_bf16, k_j, pypto.DT_FP32) dk_tile pypto.matmul(ds_bf16, q_i, pypto.DT_FP32, a_transTrue) # V2: scale atomic_add 写回sg_set_scope2 链合 pypto.set_pass_options(sg_set_scope2) pypto.set_vec_tile_shapes(v_tile_d[0], v_tile_d[1]) pypto.atomic_add(dv_tile, [s2_off, h_ofs], dv) pypto.atomic_add(pypto.mul(dq_tile, scale), [s1_off, h_ofs], dq) pypto.atomic_add(pypto.mul(dk_tile, scale), [s2_off, h_ofs], dk) pypto.set_pass_options(sg_set_scope-1)3.1 数学语义每个 (s1, s2) 块对内发生了什么以 flash_attention_mha_grad/README.md 中的计算流程为参照每个 q_tile × kv_tile 块对内依次完成1. S_tile Q_tile K_tile^T [sq, s2] BF16 matmul → FP32 2. P_tile exp(S * scale - M) / L [sq, s2] FP32 3. dP_tile dO_tile V_tile^T [sq, s2] BF16 matmul → FP32 4. D sum(O_tile * dO_tile, -1) [sq, 1] BF16 → FP32 5. dS_tile P * (dP - D) [sq, s2] FP32 → cast BF16 6. dK_tile dS^T Q_tile * scale [s2, D] BF16 matmul → FP32 → BF16 7. dV_tile P^T dO_tile [s2, D] BF16 matmul → BF16 8. dQ_partial dS K_tile * scale [sq, D] BF16 matmul → FP32跨 kv_tile 累加其中第 4 步的 D 是标准 Flash Attention 反向传播中的行和项D rowsum(O ⊙ dO)dS 的推导式dS P ⊙ (dP - D)正是由 softmax 前向P exp(S·scale - M) / L求导得到这也是为什么反向必须拿到前向的 l/m——没有它们就无法还原 P 的归一化分母与数值稳定偏移。3.2 与仓库 A3 实现的对应关系仓库中的 A3 达标实现 flash_attention_mha_grad_impl_a3.py 与上述骨架逐行对应可作为骨架的「标准答案」对照阅读runtime_options中stitch_function_max_num512、device_sched_mode3、ready_on_host_tensors[actual_q, actual_kv]、max_workspace_kb25170624全部落地第 44-53 行pass_options的cube_l1_reuse_setting{-1: 1, 0: 4}落地第 51-53 行入口通过pypto.reshape(..., inplaceTrue)将[total, num_heads, head_dim]三维输入零拷贝展平为[total, hidden_dim]l/m 展平为[total, num_heads]第 86-92 行四级循环全部使用带name/idx_name的pypto.loop第 100-113 行每个 view 均携带valid_shape第 120-126 行全程无where手动掩码V1 段用sg_set_scope1包裹、V2 写回段用sg_set_scope2包裹、段尾复位-1第 132/146/153/160 行。此外还存在一个非 A3 版本 flash_attention_mha_grad_impl.py它把 tile 配置细分为c_tile_mm / c_tile_dq / c_tile_dkv三组独立 cube tileC1 双 MatMul 与 C2 三 MatMul 采用不同分块并根据pypto.platform.npuarch DAV_3510在关键段落切换sg_set_ooo_scope展示了对不同 NPU 架构做分段调度的做法——这印证了 SK-16 骨架在具体实现时可按架构进一步细分 tile 与调度作用域。四、关键编码特征14 条硬性规则SK-16 用一张特征表给出了反向算子区别于前向骨架的核心编码纪律必须逐条遵守特征规则单趟结构同一 s2 tile 循环内完成 S/P/dP/dS 与 dQ/dK/dV 三路梯度QK^T 每 (s1,s2) 块对仅 1 次l/m 直接消费前向预计算的 l/m 从输入 view 读取禁止 kernel 内重算重算 每 q_tile 多一趟完整 QK^T 多一倍 KV 内存流量全动态循环batch/head/s1/s2 四级全部pypto.loop带name/idx_name禁止 Pythonfor展开静态展开阻止编译器全图调度任务数爆炸atomic_add 三路写回dQ/dK/dV 均atomic_add直写 GMhost 侧torch.zeros预清零kernel 不做 assemble 清零无手动掩码仅view valid_shape声明尾块有效形状禁止where手动清零 padding 行强制 vec 路径破坏 matmul→atomic_add 融合scale 编译期折叠scale 1.0 / (head_dim ** 0.5)Python float禁止作为运行时 tensor 传入每次 mul 多走一遍 vecTile 配置使用设计确定的编译期常量调整后重新编译并验证不将 TileShape 作为运行时输入div 硬件 intrinsicpypto.div(p, l, precision_typeINTRINSIC)inplace reshapekernel 入口 3D→2D 用pypto.reshape(..., inplaceTrue)零拷贝dS/P cast BF16FP32 → BF16 后再进 C2 matmul与前向 golden 精度语义一致4.1 l/m 禁止重算这是精度与性能的双重红线骨架特别给出注意项l/m 是 kernel 需用的输入golden 可重算不使用kernel 仍直接读取测试输入须提供真实 l/m。也就是说即便测试侧的 golden reference 可以自行重算 softmax 统计量见pypto-golden-generate的 reference-normalization.md 规范也不得以 golden 重算为由反过来让 kernel 重算。原因很直接若 kernel 内重算 l/m每个 q_tile 都要额外执行一趟完整 QK^T 并多读一倍 KV 数据单趟结构退化为事实上的多趟性能无法达标。4.2 Dtype 转换纪律从仓库 README 的 dtype 转换流程表可以更精确地理解骨架中每一处 cast 的用意阶段操作Dtype输入Q/K/V/O/dOBF16D 计算cast(O, FP32) * cast(dO, FP32) → sumBF16 → FP32S/dP matmulQK^T, dOV^TBF16 → FP32 (out_dtypeFP32)softmaxexp(S*scale - M) / LFP32 全程中间 castdS, P → BF16FP32 → BF16dK matmuldS^TQ → FP32 * scaleBF16 → FP32 → BF16dV matmulP^TdOBF16 → BF16 (out_dtypeBF16)dQ matmuldSK → FP32 * scale累加 FP32BF16 → FP32 → BF16即C1 的两个 MatMul 以 BF16 输入、FP32 输出softmax 全程 FP32C2 的三个 MatMul 输入是p_bf16/ds_bf16FP32→BF16 后的中间量其中 dV 的out_dtype可保持 BF16dQ 累加与 dK 均需经 FP32 累积后缩放。五、适用条件何时选用 SK-16SK-16 的适用范围由以下条件界定is_backward true输出 dQ/dK/dV 三路梯度输入含前向 l/m 统计量l_input/m_inputhas_matmul true且matmul_count 5C1 双 C2 三与 SK-01 的关系SK-01 是前向 online softmaxSK-16 是反向单趟梯度。骨架匹配阶段先做 SK-01 的反向传播结构决策单 Pass vs 双 Pass选单 Pass 即落入本骨架。任务数预估公式源自 SK-01 决策规则为任务数 ≈ batch × heads × ⌈sq/Q_TILE⌉ × ⌈skv/KV_TILE⌉ × Pass 趟数单 Pass 的 Pass 趟数为 1。该公式同时是后续调优的抓手结合实际调度开销评估增大 tile 或减少计算趟数是否可行。六、开箱性能优化提示A3 达标配置解读SK-16 的性能优化表以仓库中 flash_attention_mha_grad_impl_a3.pyA3 达标实现为实证来源每一项都标注了「必配」或「禁止」的强制程度维度推荐配置取值经验作用runtime_options.stitch_function_max_num必配512高于 SK-01 的 128单趟梯度子图更大128 会切碎runtime_options.device_sched_mode必配3多 MatMul 并行调度与 SK-06 FFN 一致runtime_options.ready_on_host_tensors必配[actual_q, actual_kv]varlen 控制流 tensor host 预发射消除调度等待气泡runtime_options.max_workspace_kb必配按平台实测A3:25170624memory-driven stitch workspacepass_options.cube_l1_reuse_setting必配{-1: 1, 0: 4}C1/C2 cube L1 复用pypto.set_pass_options(sg_set_scope...)必配V11/ V22/ 段尾-1softmax 链合 写回链合避免 vec 图过小、调度缝隙多combine_axisTrue必配jit 函数体首行尾轴 broadcast 内联 brcb见 F-15手写where掩码禁止仅valid_shape手动掩码强制 vec 路径破坏 matmul→atomic_add 融合且增加 vec 开销Pythonfor展开 batch/head禁止全pypto.loop静态展开阻止全图调度root 数爆炸实测 P0 256 roots → profiling 后处理挂死两趟重算 l/m禁止单趟直用两趟 2× KV 内存流量 额外一趟 QK^T性能无法达标几个关键配置的机理stitch_function_max_num从 128 提到 512SK-01 前向的 C-V-C 子图较小128 足够反向单趟梯度包含 C1 双 MatMul V1 链 C2 三 MatMul V2 链子图显著更大128 会把完整链切碎成过多小函数增加调度开销device_sched_mode3允许多个 MatMul 在设备侧并行调度与 SK-06FFN/SwiGLU的多 MatMul 场景同源ready_on_host_tensors对 varlen 是必配actual_q/actual_kv是控制流读值 tensor若不提前在 host 侧发射循环次数推导会形成调度等待气泡sg_set_scope分段V1softmax 归一化 dS 推导用 scope1 链合、V2scale atomic_add 写回用 scope2 链合、段尾复位 -1防止 vec 图过小导致调度缝隙过多。七、落地验证测试用例与运行方式仓库在 tests/ops/experimental/ops_transformer/flash_attention_mha_grad/ 提供了test_flash_attention_mha_grad_a3.py等测试验证 SK-16 骨架的正确性。测试的语义约定与骨架一致Q 侧s1_size包含 Q/O/dO/L/M/dQKV 侧s2_size包含 K/V/dK/dVS2_TILE将 KV 序列维分块把中间注意力矩阵从[s1_size, s2_size]降到[s1_size, S2_TILE]dK/dV 跨 kv tile 累加。README 中给出的典型测试用例矩阵覆盖了多种形状组合用例batchheadsQ:s1KV:s2dim说明test_018832032064默认配置test_02882432243264大序列长度 (skip)test_03816323232多头小维度test_04816646432多头中等序列test_0588323264短序列test_06846464128少头大维度精度校验使用numpy.testing.assert_allclosertol 0.00781251/128等于 BF16 machine epsilon、atol 0.0001。值得强调的是Golden reference 严格模拟 kernel 内部的 dtype 转换流程包括中间 BF16 cast确保对比基准与硬件行为一致——这与骨架「dS/P cast BF16 后再进 C2 matmul与前向 golden 精度语义一致」的规则互为印证。运行方式在具备 NPU 环境的机器上# 设置设备 ID export TILE_FWK_DEVICE_ID0 # 运行全部测试用例 python flash_attention_mha_grad.pymain()中通过注释/取消注释test_funcs列表控制要执行的用例run_test支持batch_size、num_heads、s1_size、s2_size、dim、tile_config等可选参数便于针对新规格扩展用例。八、设计落地时的性能评估建议SK-16 在给出配置表的同时也提醒评估重算统计量的开销、循环展开产生的任务数以及掩码带来的额外计算采用输入统计量、单趟计算或 atomic_add 时仍需满足参考计算、依赖关系和精度要求。落地时建议按以下顺序自查结构决策已在 DESIGN.md 记录是否选择了单 Pass 路线双 Pass 回退是否写明理由平台不支持atomic_add/ 精度语义要求确定性归约 / 输出 buffer 不可预清零l/m 输入已接真实值测试输入必须提供前向真实 l/mkernel 不做重算四级循环全部pypto.loop无 Pythonfor静态展开 batch/head无手写掩码所有尾块边界仅靠view valid_shape表达scale 为编译期常量不作为运行时 tensor 传入sg_set_scope 分段完整V11 / V22 / 段尾-1。结语SK-16 是 PyPTO-Gym 中 Attention 反向传播算子族FA MHA Grad、FA Score Grad、SparseAttn Grad 等的默认实现骨架它把「单 Pass」结构决策固化为一套可复刻的 C-V-C-V 循环模板全动态pypto.loop四级循环、l/m 直用零重算、三路atomic_add写回、sg_set_scope分段链合配合 A3 达标实现的配置实证可直接作为新反向算子的设计起点。对照 SK-01 理解前向骨架对照 flash_attention_mha_grad_impl_a3.py 理解落地形态即可快速掌握 PyPTO 上 Attention 反向传播的高性能实现范式。【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED — 相关阅读

相关资讯

LATEST — 最新资讯

最新发布

TODAY — 本日精选

新闻

WEEKLY — 本周精选

新闻

MONTHLY — 本月精选

新闻