Pull Request已成功合入, 合并人@liuchongming74
(感谢 xu-xianliang 的贡献)变更摘要
此 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参数与后端调度:LinearAttentionContextParallel、LinearAttentionP2PCPWrapper、LinearAttentionUlyssesCPWrapper及chunk_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.py、chunk_o.py、chunk_scaled_dot_kkt.py、cumsum.py、solve_tril.py、wy_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_backend与linear_attention_cp_mode:在AcceleratorConfig中新增两个字段,默认值分别为"eager"和"ulysses",并在parallelize_qwen3_5中读取并传递给 CP 并行化函数,使用户可通过训练配置控制线性注意力层的后端选择和 CP 模式。


代码审查
审查总结
本次审查覆盖了全部 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 |
💬 仅评论


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


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 路径。
主要改动如下:
initial/final-state gradient。
将跨 rank 依赖收敛为较小 recurrent state 的传递,避免传输完整序列激活。
linear_attention_gdn_backend: triton显式启用融合后端;默认仍为
eager,不改变已有配置的行为。进行 capability 检查;不支持的组合直接报错,不进行静默回退。
all_gather + eager仍按已有实现工作。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,并取所有 rankelapsed 的最大值,报告 median/min/p90。显存取所有 rank 的 peak allocated 最大值。
单卡 GDN 算子精度、性能和显存
B=1,H=32,dk=dv=128,序列长度 8K/16K/32K;覆盖有/无 initial state、三个随机种子以及随机
dO/dHT。dq/dk/dv/dg/dbeta/dH0,最差relative-L2 为
5.86e-3,全部张量 finite,误差未随序列长度增长。12.65x、16K22.37x、32K
40.08x、64K74.17x;增量 peak allocated 均降低约3.84x。64K 仅补测性能和显存。该数据只代表 GDN core,不代表完整层或整网收益。
CP4 完整 Linear Attention 层功能和精度
hidden=2048,QK/V heads=16/32,head dim=128,BF16;无 CP eager 层作为reference。
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、跨 rankmask 切片和 GDN partial chunk;三种模式的 padding output 均为 exact zero。
四层 Qwen3.5 多步训练精度
hidden=2048,BF16,AdamW,非零学习率;每一步读取不同的真实 token 序列。
200 个 optimizer steps,并覆盖七跳 P2P 状态链。
3.8e-4,最大 grad-norm relativediff 为
3.8e-3~4.0e-3;全部步骤无 NaN/Inf、无误差持续扩大。8.19e-5/2.87e-4。完整层性能和显存
Qwen3.5 整网性能和显存
完整层收益进入整网后会被 MLP、Full Attention、embedding、lm_head、loss 和 FSDP
等公共开销稀释。CP4/64K 四层整网仍保留约 92% 的三层 Linear Attention 绝对节省;
CP8 下 Ulysses 的 all-to-all 和小 head kernel 开销进一步增加,因此 P2P 整网收益扩大到
约 15.9%。
Profile、兼容性和工程检查
0.12~0.14 ms;中间 rank 存在计算通信重叠,状态张量传输带宽不是主要瓶颈。
backend=auto、all_gather + triton、不支持的 head dimension/chunk size/localsequence,以及 Triton 缺失或版本不满足时,均按预期 fail-fast。
10 passed。git diff --check、pylint/lizard 流水线、wheel 构建及隔离安装检查通过;wheel 中包含GDN kernel 和许可证文件。
Self-checklist:(请自检,在[ ]内打上x,我们将检视你的完成情况,否则会导致pr无法合入)