PyPTO-Gym 实战:基于 PyPTO 框架的 Fused SwiGLU 反向传播算子(Ascend NPU 三 Kernel 融合实现)
PyPTO-Gym 实战基于 PyPTO 框架的 Fused SwiGLU 反向传播算子Ascend NPU 三 Kernel 融合实现【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym导读本文以 PyPTO-Gym 仓库中的 Fused SwiGLU 反向传播算子为实例完整讲解基于 PyPTO 框架在 Ascend NPU 上实现深度学习算子反向传播的工程方法从数学公式推导、三 Kernel 任务划分、参数规格与动态 batch 维度设计到分块配置、Pass 优化选项、精度校验与测试运行。读完本文你将掌握如何用pypto.frontend.jit编写融合反向 kernel理解 BF16 精度下 matmul/逐元素运算的 dtype 转换流程并能直接运行本仓库提供的测试用例复现算子行为。算子背景SwiGLU 激活与反向传播的融合动机SwiGLUSwish-Gated Linear Unit是当前大模型 FFN前馈网络中广泛采用的激活函数其前向计算由 Gate 分支与 FC 分支组成Gate x W_g b_g FC x W_fc b_fc y SiLU(Gate) × FC Gate × σ(Gate) × FC其中 SiLUSwish x·σ(x)σ 为 Sigmoid。仓库中对应的前向 kernel 位于 fused_swiglu_impl.py其计算公式与本文反向算子一一对应是理解反向输入来源g、fc、x、w_g、w_fc的关键参照。反向传播若按朴素方式逐个计算中间梯度、权重梯度与输入梯度会产生多次 kernel 启动和重复的中间张量访存。PyPTO-Gym 的 fused_swiglu_grad 示例将反向过程融合为三个独立 kernel 顺序调用的方案覆盖了 FFN 反向的全部计算需求。产品支持情况依据 fused_swiglu_grad README本算子支持以下硬件平台Ascend 950PR支持Atlas A3 训练系列产品 / Atlas A3 推理系列产品支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品支持这意味着算子实现中不依赖特定架构私有指令且代码中还会针对具体 NPU 架构如DAV_3510做分块参数的动态适配详见下文「NPU 架构适配」一节。文件说明文件说明fused_swiglu_grad_impl.pyKernel 实现含 3 个子 kerneltest_fused_swiglu_grad.py测试用例 Golden 参考实现两个文件同属ops_transformer目录该目录集中存放 Transformer 结构相关算子样例FlashAttention 系列、MLA、Lightning Indexer、Sparse Attention 等SwiGLU 的前向与反向在此并列放置便于对照阅读。算法概述与数学公式本算子实现 SwiGLU 的反向传播将三个独立的梯度计算 kernel 融合分别计算b_kerneldg、dfc中间梯度 db_g、db_fc偏置梯度w_kerneldw_g、dw_fc权重梯度x_kerneldx输入梯度给定上游梯度 dy各梯度的数学定义如下。1. dg 和 dfc中间梯度$$ dg dy \cdot FC \cdot \sigma(Gate) $$$$ \text{SiLU}(g) \sigma(g) \cdot (1 g \cdot (1 - \sigma(g))) $$$$ dfc dy \cdot \text{SiLU}(Gate) dy \cdot g \cdot \sigma(g) $$对照实现代码fused_swiglu_grad_impl.py先算exp_g exp(g)与sigmoid_g exp_g / (1 exp_g)再组合出silu_bwd与silu_g最终dg dy_mul_fc * silu_bwd、dfc dy_fp32 * silu_g。注意这里的 sigmoid 实现为exp(g) / (1 exp(g))的等价形式并以precision_typepypto.PrecisionType.INTRINSIC声明精度模式。2. db_g 和 db_fc偏置梯度$$ db_g \sum_{i} dg_i \quad \text{(沿 batch 维度求和)} $$$$ db_fc \sum_{i} dfc_i \quad \text{(沿 batch 维度求和)} $$实现上通过pypto.sum(dg_fp32, dim0, keepdimTrue)沿 batch 维规约再与输出张量做db_g[:] db_g cast(dg_sum, BF16)的累加——这正是 README 强调「db_g、db_fc 需初始化为 0」的原因多 tile 循环时靠累加合并各分块的部分和。3. dw_g 和 dw_fc权重梯度$$ dw_g x^T \cdot dg $$$$ dw_fc x^T \cdot dfc $$4. dx输入梯度$$ dx dg \cdot W_g^T dfc \cdot W_{fc}^T $$权重梯度与输入梯度均由 matmul 完成前者以a_transTrue, b_transFalse计算 x^T dg后者以a_transFalse, b_transTrue计算 dg w_g^T与公式一一对应。Kernel 概览与调用链反向传播拆分为三个独立的 kernel按顺序调用Kernel功能输入输出fused_swiglu_bwd_b_kernel计算中间梯度 偏置梯度dy, g, fcdg, dfc, db_g, db_fcfused_swiglu_bwd_w_kernel计算权重梯度x, dg, dfcdw_g, dw_fcfused_swiglu_bwd_x_kernel计算输入梯度dg, dfc, w_g, w_fcdx三者存在严格的数据依赖w_kernel 与 x_kernel 都消费 b_kernel 产出的 dg、dfc因此必须顺序执行。这一调用链在测试中体现得最直观test_fused_swiglu_grad.pyfused_swiglu_bwd_b_kernel(dy, g, fc, dg_out, dfc_out, db_g_out, db_fc_out) fused_swiglu_bwd_w_kernel(x, dg_out, dfc_out, dw_g_out, dw_fc_out) fused_swiglu_bwd_x_kernel(dg_out, dfc_out, w_g, w_fc, dx_out)Kernel 签名与参数规格Kernel 1: fused_swiglu_bwd_b_kernelfused_swiglu_bwd_b_kernel( dy, # [M, N] BF16 — 上游梯度动态 batch 维度 g, # [M, N] BF16 — Gate 中间值前向传播输出 fc, # [M, N] BF16 — FC 中间值前向传播输出 dg, # [M, N] BF16 — Gate 梯度输出 dfc, # [M, N] BF16 — FC 梯度输出 db_g, # [1, N] BF16 — Gate 偏置梯度输出需初始化为 0 db_fc # [1, N] BF16 — FC 偏置梯度输出需初始化为 0 )参数类型Shape数据类型说明dy输入[M, N]BF16上游梯度M 为动态维度g输入[M, N]BF16Gate 中间值前向传播保存fc输入[M, N]BF16FC 中间值前向传播保存dg输出[M, N]BF16Gate 梯度dfc输出[M, N]BF16FC 梯度db_g输出[1, N]BF16Gate 偏置梯度需初始化为 0db_fc输出[1, N]BF16FC 偏置梯度需初始化为 0Kernel 2: fused_swiglu_bwd_w_kernelfused_swiglu_bwd_w_kernel( x, # [M, K] BF16 — 输入张量前向传播输入 dg, # [M, N] BF16 — Gate 梯度 dfc, # [M, N] BF16 — FC 梯度 dw_g, # [K, N] BF16 — Gate 权重梯度输出需初始化为 0 dw_fc # [K, N] BF16 — FC 权重梯度输出需初始化为 0 )参数类型Shape数据类型说明x输入[M, K]BF16输入张量前向传播保存dg输入[M, N]BF16Gate 梯度dfc输入[M, N]BF16FC 梯度dw_g输出[K, N]BF16Gate 权重梯度需初始化为 0dw_fc输出[K, N]BF16FC 权重梯度需初始化为 0Kernel 3: fused_swiglu_bwd_x_kernelfused_swiglu_bwd_x_kernel( dg, # [M, N] BF16 — Gate 梯度 dfc, # [M, N] BF16 — FC 梯度 w_g, # [K, N] BF16 — Gate 权重 w_fc, # [K, N] BF16 — FC 权重 dx # [M, K] BF16 — 输入梯度输出 )参数类型Shape数据类型说明dg输入[M, N]BF16Gate 梯度dfc输入[M, N]BF16FC 梯度w_g输入[K, N]BF16Gate 权重w_fc输入[K, N]BF16FC 权重dx输出[M, K]BF16输入梯度动态轴MBatch dimension动态维度运行时可变。在 PyPTO 中通过pypto.Tensor([pypto.DYNAMIC, ...], pypto.DT_BF16)声明见实现代码允许同一份编译产物处理不同 batch 大小的输入这在实际训练中避免了 batch 变化导致的重编译。分块Tiling配置三个 kernel 均沿 batch 维M 维以固定 tile 循环切分并用valid_shape处理末尾不满一块的边界Kerneltile_mvec_tilecube_tile (MNK)b_kernel1024[128, 128]-w_kernel2048[128, 128][128, 128], [128, 256], [128, 256]x_kernel1024[128, 128/256][128, 128], [64, 256], [256, 256]代码中三个 kernel 统一采用tile_m分块 pypto.loop循环的写法例如 b_kernel 的核心循环fused_swiglu_grad_impl.pyfor idx in pypto.loop(loop_count, nameLOOP_BWD_DG, idx_nameidx): tile_offset idx * tile_m valid_m (m - tile_offset).min(tile_m) dy_tile pypto.view(dy, [tile_m, n], [tile_offset, 0], valid_shape[valid_m, n])设计要点b_kernel 只做 element-wise 与规约无 matmul因此不配置 cube_tilew_kernel 的 tile_m 取 2048大于其他两个 kernel 的 1024因为其 matmul 形态为x^T dg输出 shape 为 [K, N]对 M 维分块的粒度要求更宽松更大的 tile 可减少循环次数cube_tile 按 MNK 语义配置w_kernel 的 matmul 是[K,M] [M,N]形态x_kernel 的 matmul 是[M,N] [N,K]形态二者分别用pypto.set_cube_tile_shapes([128, 128], [128, 256], [128, 256])与([128, 128], [64, 256], [256, 256])指定。运行测试测试入口与运行方式如下# 设置设备 ID export TILE_FWK_DEVICE_ID0 # 运行测试 python test_fused_swiglu_grad.py说明设备 ID 通过环境变量TILE_FWK_DEVICE_ID指定测试代码中get_device_id()会读取该变量缺省为 0test_fused_swiglu_grad.py测试前会先把src与src/pypto_gym/ops/pypto_tensor加入sys.path再以from experimental.ops_transformer.fused_swiglu_grad.fused_swiglu_grad_impl import ...导入三个 kernel若以 pytest 方式运行整个测试目录仓库根目录的 conftest.py 与 tests/ops/experimental/conftest.py 会预先导入torch_npu保证pypto.frontend.jit装饰器在模块收集阶段能正常调用torch.npu.is_available()。测试用例用例MKN说明test_bwd2200005121024大 batch 标准 FFN 维度用例以m220000, k512, n1024复现大模型 FFN 的典型维度hidden512、中间层1024同时用torch.randn(...) / sqrt(m)、/ sqrt(k)对输入、权重做缩放使梯度量级保持稳定避免 BF16 下溢出或精度崩溃。随机种子固定为 0np.random.seed(0)与torch.manual_seed(0)保证结果可复现。精度校验使用numpy.testing.assert_allclose进行精度验证rtol 0.0078125 # 1/128等于 BF16 machine epsilon atol 0.0001校验输出包括dx输入梯度dw_gGate 权重梯度dw_fcFC 权重梯度db_gGate 偏置梯度db_fcFC 偏置梯度Golden referencegolden_fused_swiglu_bwd见 test_fused_swiglu_grad.py严格模拟 kernel 内部的 dtype 转换流程——全部先float()提升到 FP32 计算再在规约与 matmul 输出处to(dy.dtype)截断回 BF16从而保证对比基准与硬件行为一致而不是用一个理想化 FP32 结果去苛求 BF16 kernel。值得注意的实现细节Golden 中db_g、db_fc是单次全局 sum而 kernel 是按 1024 行的 tile 分块求和后再累加二者在浮点舍入路径上不同但仍落在 rtol/atol 容差内这也验证了分块累加方案的数值稳定性。Dtype 转换流程阶段操作Dtype输入dy, g, fc, x, w_g, w_fcBF16sigmoidexp(g) / (1 exp(g))BF16element-wisedg, dfc 计算BF16sumdb_g, db_fcBF16matmul (w_kernel)x^T dg, x^T dfcBF16 → BF16matmul (x_kernel)dg w_g^T, dfc w_fc^TBF16 → BF16输出dg, dfc, dx, dw_g, dw_fc, db_g, db_fcBF16对照源码可以看到更精确的中间细节所有算术运算实际在 FP32 中进行。每个 kernel 的第一步都是把输入 tilepypto.cast到pypto.DT_FP32计算完成后仅在写回前 cast 回 BF16。例如 b_kernel 中dy_fp32 pypto.cast(dy_tile, pypto.DT_FP32)matmul 均显式指定pypto.DT_FP32作为累加精度pypto.matmul(x_fp32, dg_fp32, pypto.DT_FP32, ...)。这种「BF16 存取、FP32 计算」的策略既享受 BF16 的存储/带宽优势又避免了 BF16 逐元素累加带来的精度损失——这也是为什么最终 rtol 取 BF16 machine epsilon1/128仍能通过校验。实现细节Pass 配置三个 kernel 均使用pypto.frontend.jit装饰器通过pass_options与runtime_options控制编译与运行行为b_kernel:pypto.frontend.jit( pass_options{ pg_upper_bound: 5000000, vec_nbuffer_setting: {-1: 2, 0: 4} }, runtime_options{ run_mode: global_run_mode, stitch_function_max_num: 128, device_sched_mode: 3 } )w_kernel:pypto.frontend.jit( pass_options{ pg_upper_bound: 5000000, vec_nbuffer_setting: {-1: 2, 0: 8}, cube_l1_reuse_setting: {-1: 2} }, runtime_options{ run_mode: global_run_mode, stitch_function_max_num: 128, device_sched_mode: 3 } )x_kernel:pypto.frontend.jit( pass_options{ pg_upper_bound: 5000000, vec_nbuffer_setting: {-1: 2, 0: 4}, cube_l1_reuse_setting: {-1: 2} }, runtime_options{ run_mode: global_run_mode, stitch_function_max_num: 128, device_sched_mode: 3 } )各选项的作用与差异选项含义三个 kernel 的取值pg_upper_boundPass group 上界约束编译分组规模均为 5000000vec_nbuffer_setting向量计算 buffer 深度配置-1为默认 buffer0为向量 bufferb/x: {-1:2, 0:4}w: {-1:2, 0:8}cube_l1_reuse_settingCube 单元 L1 复用配置开启后可提升 matmul 访存效率仅 w/x kernel 开启{-1:2}b kernel 无 matmul 故不配置run_mode运行模式NPU 或 SIM 仿真global_run_modestitch_function_max_num最大 stitch指令拼接函数数128device_sched_mode设备调度模式3w_kernel 的vec_nbuffer_setting取0: 8比 b/x kernel 的0: 4多一倍向量 buffer与其更大的 tile_m2048和更重的输出回写量相匹配。NPU 架构适配x_kernel 根据 NPU 架构动态调整 vec_tile 配置if pypto.platform.npuarch DAV_3510: pypto.set_vec_tile_shapes(128, 256) else: pypto.set_vec_tile_shapes(128, 128)DAV_3510对应的架构上向量 tile 的列宽可放宽到 256其余架构保持 128。这解释了 README 分块配置表中 x_kernel 的 vec_tile 写作[128, 128/256]的原因——它是运行时按平台决定的。结合「产品支持情况」一节可知本算子通过这种条件分支在同一份代码上覆盖 Ascend 950PR / Atlas A2 / Atlas A3 多代产品。性能优化分块策略Batch 维度分块tile_m 1024/2048优化内存访问valid_shape处理动态边界避免对非对齐尺寸做 paddingCube 复用w_kernel 和 x_kernel 启用 L1 复用cube_l1_reuse_setting提升 matmul 性能向量化使用 vector buffervec_nbuffer_setting优化 element-wise 操作配合pypto.set_vec_tile_shapes控制向量指令的分块粒度梯度累加db_g、db_fc、dw_g、dw_fc 支持跨 tile 累加——输出张量由调用方预先置零kernel 内以out[:] out cast(partial, BF16)形式原地累加规避了跨 tile 的二次归约开销精度与性能的平衡sigmoid 采用precision_typepypto.PrecisionType.INTRINSIC指令精度模式允许编译器使用硬件近似指令换取吞吐同时依赖 FP32 中间计算兜底精度。依赖Python 3.xPyTorch torch_npuPyPTOpypto包NumPy测试脚本还使用 pytest 组织用例仓库 requirements.txt 中声明了相关依赖。运行环境需具备 Ascend NPU 驱动与 CANN 工具链TILE_FWK_DEVICE_ID指定的设备应为空闲可用状态。小结Fused SwiGLU 反向传播算子是 PyPTO-Gym 中一个典型的「动态 batch 多 kernel 融合 分块累加」算子样例用三个职责单一的 kernel 覆盖 FFN 反向下游的全部梯度需求通过 FP32 中间计算保证 BF16 数值精度以pypto.looppypto.view(valid_shape...)支撑动态 M 维用 Pass 选项与架构分支完成跨平台适配。对照前向实现fused_swiglu_impl.py与测试用例test_fused_swiglu.py、test_fused_swiglu_grad.py阅读可以快速掌握 PyPTO 编写生产级算子反向 kernel 的完整套路并作为自研算子融合与精度调试的参考基线。【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考