已合并
feat: add fused GDN backend for linear attention CP #1114
feat: add fused GDN backend for linear attention CP #1114
已合并
xu-xianliang创建于 8月5日
xu-xianliang
xu-xianliang
8月5日

What type of PR is this?

/kind feature


What does this PR do / why do we need it:

本 PR 为已有的 Qwen3.5 Linear Attention Context Parallel 增加显式可选的 Triton-Ascend
Gated Delta Rule(GDN)融合后端,并为 State-P2P 模式增加融合的状态仿射 summary 路径。

主要改动如下:

  1. 增加 GDN Triton-Ascend 前反向算子,覆盖 token output、final state、全部输入梯度以及
    initial/final-state gradient。
  2. 增加 State-P2P 所需的 state summary、summary apply 和反向 gradient-summary 算子,
    将跨 rank 依赖收敛为较小 recurrent state 的传递,避免传输完整序列激活。
  3. Ulysses 和 P2P 均可通过 linear_attention_gdn_backend: triton 显式启用融合后端;默认
    仍为 eager,不改变已有配置的行为。
  4. 对 Triton 版本、设备、dtype、head dimension、chunk size、local sequence 和 CP mode
    进行 capability 检查;不支持的组合直接报错,不进行静默回退。
  5. 保持 AllGather Triton、融合 Conv1D 和自动 backend 选择不在本 PR 范围内;
    all_gather + eager 仍按已有实现工作。
  6. 增加 wheel 打包规则和第三方许可证文件,并保证 CPU/MindSpore import 不会主动加载
    Torch Triton kernel。

该功能用于解决 eager GDN 由大量小算子及长 autograd graph 带来的性能和显存问题,并使
P2P Linear Attention CP 在长序列和较大 CP size 下避免 Ulysses 的大张量 all-to-all 以及
head 分片后 kernel 并行度下降问题。


Which issue(s) this PR fixes:
完整的测试验证报告见关联的issue #320


Test Plan and Test result:What scenarios were tested, and what were the verification results(Function, performance, reliability, etc.)

测试环境:8 张 Ascend 910B3,CANN 9.1.0-beta.3,PyTorch 2.10.0,torch-npu 2.10.0,
triton-ascend 3.2.1,主要计算类型为 BF16,GDN gate 累计及 recurrent state 使用 FP32。

性能统一使用 warmup=5/repeat=20;计时前后执行 barrier/device synchronize,并取所有 rank
elapsed 的最大值,报告 median/min/p90。显存取所有 rank 的 peak allocated 最大值。

  1. 单卡 GDN 算子精度、性能和显存

    • Shape:B=1,H=32,dk=dv=128,序列长度 8K/16K/32K;覆盖有/无 initial state、
      三个随机种子以及随机 dO/dHT
    • 对比 eager reference 的 output、final state、dq/dk/dv/dg/dbeta/dH0,最差
      relative-L2 为 5.86e-3,全部张量 finite,误差未随序列长度增长。
    • Triton 相对 eager 的 Fwd+Bwd 加速为:8K 12.65x、16K 22.37x
      32K 40.08x、64K 74.17x;增量 peak allocated 均降低约 3.84x
      64K 仅补测性能和显存。该数据只代表 GDN core,不代表完整层或整网收益。
  2. CP4 完整 Linear Attention 层功能和精度

    • 覆盖 projection、Conv1D CP halo、Q/K L2Norm、GDN、RMSNormGated 和 out projection。
    • 配置:hidden=2048,QK/V heads=16/32,head dim=128,BF16;无 CP eager 层作为
      reference。
    • Ulysses Triton、P2P Triton 和 AllGather eager 均完成 8K/32K 前反向对齐;
      P2P Triton 在 32K 下 output/input-grad/parameter-grad relative-L2 分别为
      4.4771e-3/5.8587e-3/5.0757e-3
    • 使用 B=2,global sequence=8192,valid length=7777 验证尾部 padding、跨 rank
      mask 切片和 GDN partial chunk;三种模式的 padding output 均为 exact zero。
    • 连续执行至少 20 次分布式前反向,无 P2P 顺序错配、死锁、NaN/Inf 或显存持续增长。
  3. 四层 Qwen3.5 多步训练精度

    • 模型:3 层 Linear Attention + 1 层 Full Attention,hidden=2048,BF16,AdamW,
      非零学习率;每一步读取不同的真实 token 序列。
    • CP4 + FSDP group 4 完成 100 个 optimizer steps;CP8 + FSDP group 8 完成
      200 个 optimizer steps,并覆盖七跳 P2P 状态链。
    • Ulysses/P2P 的最大 loss relative diff 均低于 3.8e-4,最大 grad-norm relative
      diff 为 3.8e-3~4.0e-3;全部步骤无 NaN/Inf、无误差持续扩大。
    • AllGather eager 额外完成 10 步功能回归,最大 loss/grad-norm relative diff 分别为
      8.19e-5/2.87e-4
  4. 完整层性能和显存

    Topology Global/local seq P2P Fwd+Bwd time reduction vs Ulysses Peak allocated
    CP4 32K/8K 14.4% P2P +12.0 MiB
    CP4 64K/16K 17.0% P2P +12.0 MiB
    CP8 64K/8K 31.6% P2P -52.1 MiB
    CP8 128K/16K 45.8% P2P -116.2 MiB
  5. Qwen3.5 整网性能和显存

    Topology Model Global/local seq P2P Fwd+Bwd time reduction Throughput increase Peak allocated reduction
    CP4 + FSDP4 4 layers 64K/16K 8.83% 9.69% 372 MiB
    CP4 + FSDP4 8 layers 64K/16K 9.31% 10.27% 744 MiB
    CP8 + FSDP8 4 layers 128K/16K 15.94% 18.96% 372 MiB / 2.38%
    CP8 + FSDP8 8 layers 128K/16K 15.84% 18.82% 744 MiB / 2.47%

    完整层收益进入整网后会被 MLP、Full Attention、embedding、lm_head、loss 和 FSDP
    等公共开销稀释。CP4/64K 四层整网仍保留约 92% 的三层 Linear Attention 绝对节省;
    CP8 下 Ulysses 的 all-to-all 和小 head kernel 开销进一步增加,因此 P2P 整网收益扩大到
    约 15.9%。

  6. Profile、兼容性和工程检查

    • CP4/64K Level0 profile 中,单次 P2P send 约 0.12~0.14 ms;中间 rank 存在
      计算通信重叠,状态张量传输带宽不是主要瓶颈。
    • backend=autoall_gather + triton、不支持的 head dimension/chunk size/local
      sequence,以及 Triton 缺失或版本不满足时,均按预期 fail-fast。
    • Linear Attention CP CPU UT:10 passed
    • git diff --check、pylint/lizard 流水线、wheel 构建及隔离安装检查通过;wheel 中包含
      GDN kernel 和许可证文件。

Self-checklist:(请自检,在[ ]内打上x,我们将检视你的完成情况,否则会导致pr无法合入

likedislike
Pull Request已成功合入, 合并人@liuchongming74
(感谢 xu-xianliang 的贡献)
xu-xianliangxu-xianliang
8月5日 创建了 pull request,commit 5fbdaf1f
atomgit-bot
atomgit-bot
8月5日 评论:

变更摘要

此 PR 为线性注意力(Linear Attention)上下文并行(Context Parallel)新增了基于 Triton-Ascend 的融合 GDN(Gated DeltaNet)后端。主要通过新增 backend 参数("eager" / "triton"),允许在 LinearAttentionContextParallel 及其相关 wrapper 中显式选择本地 GDN 实现;同时引入了完整的 Triton 算子模块 custom_ops/gdn/,以及配套的 P2P 状态汇总融合路径 _GDNStateP2PTritonFunction,并对 Triton-Ascend 版本、输入 dtype/shape 等进行了严格校验。

主要改动

  • 新增 backend 参数与后端调度LinearAttentionContextParallelLinearAttentionP2PCPWrapperLinearAttentionUlyssesCPWrapperchunk_gated_delta_rule 均新增 backend 参数,支持 "eager""triton" 两种后端;parallelize_qwen3_5_cp 同样透传 linear_attention_gdn_backend 配置项,实现对 Qwen3.5 线性注意力层的后端选择。

  • 新增 _GDNStateP2PTritonFunction 融合 P2P 路径:在 linear_attention_context_parallel.py 中新增自定义 autograd 函数 _GDNStateP2PTritonFunction,利用 state_summary 中的仿射状态汇总算子,将跨 rank 的状态传输与本地 GDN 计算融合,配合异步 send/recv 实现序列分片 CP 下的前向与反向传播。

  • 新增 Triton-Ascend GDN 算子模块:在 hyper_parallel/platform/torch/custom_ops/gdn/ 下新增完整的 Triton 算子集,包括 chunk_gated_delta_rule.py(前向/反向分阶段融合)、state_summary.py(仿射状态汇总与梯度汇总)、以及 chunk_delta_h.pychunk_o.pychunk_scaled_dot_kkt.pycumsum.pysolve_tril.pywy_fast.py 等内部 kernel 文件,适配 Ascend NPU 的 autotune 配置与 block_dim=64 的固定形状约束。

  • 新增 is_triton_gdn_available 运行时校验:在 linear_attention.py 中新增版本解析与输入校验逻辑,要求 triton-ascend 版本在 3.2.x 范围内、输入 q/k/v/beta 为 bf16、g 为 fp32、head_k_dim=head_v_dim=128、chunk_size=64,确保仅在满足合约的 NPU 环境下启用 Triton 后端。

  • 新增配置项 linear_attention_gdn_backendlinear_attention_cp_mode:在 AcceleratorConfig 中新增两个字段,默认值分别为 "eager""ulysses",并在 parallelize_qwen3_5 中读取并传递给 CP 并行化函数,使用户可通过训练配置控制线性注意力层的后端选择和 CP 模式。

likedislike
不准确?
atomgit-bot
atomgit-bot
8月5日 评论:

代码审查

审查总结

本次审查覆盖了全部 19 个变更文件,发现 3 个 P3 级别的低优先级问题,无 P0/P1/P2 级别问题。

各文件审查结果

文件 结果
hyper_parallel/core/context_parallel/linear_attention_context_parallel.py 无问题 — P2P Triton 通信/计算 pipeline 正确,autograd 梯度流正确
hyper_parallel/models/modules/linear_attention.py P3 ×1 — _parse_version 静默回退
hyper_parallel/models/qwen3_5/parallelize.py 无问题 — backend 参数正确透传
hyper_parallel/platform/torch/custom_ops/gdn/LICENSE 无问题 — 标准 MIT 许可证
hyper_parallel/platform/torch/custom_ops/gdn/__init__.py 无问题 — 仅文档字符串
hyper_parallel/platform/torch/custom_ops/gdn/chunk_gated_delta_rule.py 无问题 — saved 函数拆分、autograd 包装正确
hyper_parallel/platform/torch/custom_ops/gdn/state_summary.py 无问题 — affine summary 逻辑正确
hyper_parallel/platform/torch/custom_ops/gdn/triton/__init__.py 无问题 — 仅文档字符串
hyper_parallel/platform/torch/custom_ops/gdn/triton/chunk_delta_h.py P3 ×1 — chunk_gated_delta_rule_fwd_h 返回类型标注错误
hyper_parallel/platform/torch/custom_ops/gdn/triton/chunk_o.py P3 ×1 — grid 函数死代码
hyper_parallel/platform/torch/custom_ops/gdn/triton/chunk_scaled_dot_kkt.py 无问题
hyper_parallel/platform/torch/custom_ops/gdn/triton/cumsum.py 无问题
hyper_parallel/platform/torch/custom_ops/gdn/triton/solve_tril.py 无问题
hyper_parallel/platform/torch/custom_ops/gdn/triton/state_summary.py 无问题
hyper_parallel/platform/torch/custom_ops/gdn/triton/utils.py 无问题
hyper_parallel/platform/torch/custom_ops/gdn/triton/wy_fast.py 无问题
hyper_parallel/trainer/config.py 无问题 — 仅格式化变更 + 新增字段
setup.py 无问题 — LICENSE 正确包含
tests/ut/core/context_parallel/test_linear_attention_context_parallel.py 无问题 — 新增参数校验测试正确

按优先级统计

  • P0: 0
  • P1: 0
  • P2: 0
  • P3: 3(类型标注不匹配、死代码、版本解析静默回退)

整体风险评估

低风险。该 PR 的核心变更(Triton GDN 后端 + P2P state wavefront)实现质量良好:前向/反向通信模式正确、autograd 梯度流完整、边界 rank 处理正确、输入校验充分(模块级 + 运行时双重检查)、延迟导入避免非 NPU 环境崩溃。三个 P3 问题均为次要代码质量问题,不影响运行时正确性。

类型 数量
🔴 阻塞 0
🟡 建议 7

💬 仅评论

likedislike
不准确?
MindSpore-BotMindSpore-Bot成员
8月5日 添加了label:mindspore-cla/yes
司小南(机器人)司小南(机器人)成员
8月5日 添加了label:pr-check-pass
此处折叠了65条消息 查看更多
Yyangzhenzhang成员
8月6日 通过审查
阿苏阿苏成员
8月7日 通过审查
Yyao_yf成员
8月10日 通过审查
MindSpore-Bot
MindSpore-Bot成员
8月10日 评论:

The ci-pipeline-passed label is expired. Please retest again.

likedislike
liuchongming74liuchongming74成员
8月10日 合入了pull request,合并节点 SHA:942b41d2c476dd9a2758e89d483620a9706b5bee