已合并
linear attention cp #1035
linear attention cp #1035
已合并
xu-xianliang创建于 7月13日
xu-xianliang
xu-xianliang
7月13日

What type of PR is this?

/kind feature


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

Qwen3.5 采用 3:1 的 Linear Attention(Gated DeltaNet)与 Full Attention
混合层结构。原有 Qwen3.5 Context Parallel 仅能通过整段序列 gather/slice
处理 Linear Attention,无法利用其按 head 独立计算和递归状态较小的特点。

本 PR 新增 LinearAttentionContextParallel,并接入 Qwen3.5 dense 模型的
CP 并行化流程,支持以下三种执行模式:

  1. ulysses

    • 在本地 sequence shard 上完成 Q/K/V/B/A 投影。
    • 将 Q/K/V/B/A 按 rank 打包,通过一次可微 all-to-all 完成
      Shard(sequence) -> Shard(head)
    • 在完整序列、局部 head 上执行 depthwise Conv1D 和 Gated DeltaNet。
    • 通过反向 all-to-all 将输出恢复为本地 sequence shard,再执行
      RMSNorm、门控和输出投影。
  2. p2p

    • 保持完整 head 和本地 sequence shard。
    • Conv1D 仅通过可微 all-to-all-v 交换相邻 rank 所需的左侧 halo,
      不复制完整序列。
    • 每个 rank 并行计算本地 GDN 的仿射状态摘要
      H_out = M @ H_in + S
    • rank 间仅按序传递较小的 recurrent state;自定义 autograd 边界在
      backward 中按相反方向传递 state gradient。
    • 收到初始状态后,使用已准备的 chunk 中间量计算本地 token output。
  3. all_gather

    • 保持完整 head 和本地 sequence shard。
    • 每个 rank 计算本地 (S, M) 状态摘要并打包为一次可微 all-gather。
    • 本地合并当前 rank 之前的摘要,得到 initial state,再计算本地
      token output。

同时完成以下集成:

  • 在 CP 边界兼容 local tensor 和 DTensor local shard 输入。
  • 使用 platform 层可微 all-to-all/all-gather,使 collective 具备正确反向。
  • P2P/AllGather 的 GDN summary 使用 non-reentrant checkpoint:forward 不保留
    按 chunk 构造 (S, M) 的完整 autograd graph,backward 按需重算 summary,
    降低 eager 小算子实现的训练峰值显存。
  • 对 Ulysses 的 K/V head 与 V head divisibility 做显式校验。
  • Qwen3.5 混合模型中,Full Attention 继续使用现有 ContextParallel
    Linear Attention 使用新增 CP executor。
  • 对当前尚未支持的 Linear Attention TP+CP 组合显式报错,避免静默错误。
  • 保持默认模式为 ulysses,不改变未配置用户的执行行为。

Which issue(s) this PR fixes:

Fixes #<填写对应 Issue/SR/AR 编号>


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

1. UT 与分布式 ST

新增 UT:

  • 三种公开 mode 的构造与非法 mode 校验。
  • fused Q/K/V channel 按投影边界切分。
  • GDN (S, M) 摘要 pack/unpack、仿射前缀合并及 autograd 梯度。
  • checkpoint summary 与普通 eager summary 的 (S, M) 输出及
    K/V/g/beta 梯度一致性。

结果:

7 passed

新增 2 卡 Ascend 910B BF16 分布式 ST:

  • 模型:单层 Qwen3.5 Gated DeltaNet。
  • A2A 边界:覆盖 batch=1, local sequence=1
    Shard(sequence) -> Shard(head) -> Shard(sequence) round-trip
    及其 backward。
  • shape:batch=1,global sequence=128,hidden=128,
    K/V heads=4/8,K/V head dim=16/16。
  • 对比完整序列 reference 与 ulyssesp2pall_gather
  • 覆盖 forward、input gradient、全部 parameter gradient 和 global grad norm。

结果:

mode output max abs output rel-L2 input-grad max abs input-grad rel-L2 param-grad max abs worst per-param rel-L2 grad-norm relative diff
ulysses 0 0 0 0 2.1875e-1 2.4899e-3 7.58e-6
p2p 0 0 7.8125e-3 5.4051e-4 2.1875e-1 4.0091e-3 1.80e-5
all_gather 0 0 7.8125e-3 5.4051e-4 2.1875e-1 4.0091e-3 1.80e-5

所有测试通过。BF16 parameter gradient 的最大单点误差来自不同分布式
计算顺序下的舍入;最差单个 parameter 的 relative-L2 保持在
4.1e-3 以内,整体 grad norm 相对误差保持在 2e-5 以内。

ST 不把当前 eager 小算子的单点舍入特征固化为公共接口契约。为兼容后续
融合后端采用不同 reduction order,max-abs 阈值设置为:

output:        5e-2
input grad:    5e-1

Parameter gradient 的原始 max-abs 受 loss、shape 和随机梯度量级影响,
只打印用于诊断,不作为通过条件。精度通过条件使用尺度化校验:

  • output/input gradient:
    max_abs / reference_abs_max <= 1e-1relative-L2 <= 5e-2
  • 每一个 parameter gradient:
    max_abs / reference_abs_max <= 1e-1relative-L2 <= 1e-1
  • 每个 parameter gradient 必须存在且全部 finite。

global grad-norm relative difference 阈值为 1e-2。宽松 max-abs 只作为
output/input 的异常保险;逐参数 relative-max 和 relative-L2 避免大权重
梯度掩盖 A_logdt_bias 等较小参数的错误。

为评估阈值稳定性,额外在相同 2 卡 NPU ST shape 上扫描了 100 组模型初始化、
输入和 grad-output 随机种子。三种 mode 共 300 组前反向全部通过,跨全部
样本观测到的最大值为:

metric 100-seed maximum threshold margin
output max abs 1.9531e-3 5e-2 25.6x
input-grad max abs 3.1250e-2 5e-1 16.0x
max relative to reference scale 9.7087e-3 1e-1 10.3x
output/input relative-L2 2.6903e-3 5e-2 18.6x
worst per-param grad relative-L2 7.7551e-3 1e-1 12.9x
grad-norm relative diff 4.5545e-4 1e-2 22.0x

100 组中,所有参数参考梯度的最小 max-abs6.6406e-2,最小
L2 norm 为 1.0297e-1,均来自 dt_bias,未出现接近零的相对误差
分母。原始 parameter-gradient max-abs 最高为 9.0625e-1。旧的
param_grad_max_abs <= 3e-1 会失败,说明该指标过度依赖固定随机样本,
因此已从 assert 中移除。提交的 ST 仍使用固定 seed 保证 CI 可复现;
100-seed sweep 只作为阈值制定实验,不增加日常 CI 时间。

2. CP8 100 步整网精度

配置:

  • Ascend 910B,BF16,CP=8。
  • Qwen3.5 dense 4 层,layer types 为
    [linear_attention, linear_attention, linear_attention, full_attention]
  • hidden=2048,intermediate=6144,K/V heads=16/32,
    K/V head dim=128/128。
  • global sequence=1024,batch=1。
  • rank0 运行完整序列单卡 reference;CP 模型使用 FSDP+CP。
  • 初始化前逐参数确认权重完全一致。
  • 每一步使用不同的 input/labels,但 reference 与 CP 使用相同数据。
  • AdamW,learning rate=1e-4,weight decay=0.01,共 100 步。
  • 比较每一步的 global loss 和 global grad norm。

结果:

mode max loss abs max loss relative diff max grad-norm abs max grad-norm relative diff
ulysses 3.0422e-4 3.5852e-5 2.0742e-5 1.2499e-5
p2p 2.3651e-4 2.7910e-5 2.4796e-5 1.4909e-5
all_gather 2.8992e-4 3.4228e-5 2.0266e-5 1.2210e-5

100 步内未观察到误差随 optimizer step 持续放大,三种 CP 路径均与单卡
reference 保持一致的 loss 与 grad-norm 轨迹。

3. 性能验证

统一设置:

  • Ascend 910B,BF16,batch=1。
  • hidden=2048,K/V heads=16/32,K/V head dim=128/128。
  • warmup=3,repeat=10。
  • 每个 rank 同步后计时,并对 elapsed 执行 rank-max。
  • 报告 10 次测量的 median;当前使用 eager GDN 小算子后端。

单层 Linear Attention,CP4/global sequence=32K:

mode forward median forward+backward median peak allocated
ulysses 340.560 ms 2932.831 ms 5014.9 MiB
p2p 130.122 ms 1295.501 ms 5288.8 MiB
all_gather 133.388 ms 1295.456 ms 5309.8 MiB

summary checkpoint 的单独收益:

mode implementation forward+backward median peak allocated
p2p 保留 summary graph 1259.897 ms 6242.0 MiB
p2p checkpoint summary 1295.501 ms 5288.8 MiB
all_gather 保留 summary graph 1252.530 ms 6259.0 MiB
all_gather checkpoint summary 1295.456 ms 5309.8 MiB

整段 summary checkpoint 分别为 P2P/AllGather 减少约 953/949 MiB 峰值
allocated,代价为约 2.8%/3.4% 的 forward+backward 时间。分段 summary
checkpoint 和 token-scan checkpoint 均在实测中增加了运行时间,且没有
进一步降低训练峰值,因此没有进入提交代码。

CP8 长序列扩展:

global sequence mode forward median forward+backward median peak allocated
32K ulysses 331.511 ms 1789.866 ms 2552.8 MiB
32K p2p 90.001 ms 501.591 ms 2701.1 MiB
32K all_gather 75.676 ms 516.403 ms 2744.8 MiB
64K ulysses 642.818 ms 5788.480 ms 5017.4 MiB
64K p2p 139.325 ms 1301.905 ms 5288.8 MiB
64K all_gather 140.787 ms 1300.613 ms 5330.8 MiB
128K ulysses 1901.831 ms 23496.140 ms 9958.9 MiB
128K p2p 268.767 ms 4925.994 ms 10475.2 MiB
128K all_gather 263.271 ms 4916.733 ms 10521.2 MiB

固定 local sequence=8K 时,从 CP4/global 32K 扩展到 CP8/global 64K,
P2P 和 AllGather 的 forward+backward 分别只增长约 0.5% 和 0.4%,
peak allocated 基本保持不变;当前 eager Ulysses 因完整序列上的 chunk
循环数翻倍而增长约 97.4%。

说明:以上结果用于比较本 PR 的 eager 数据路径,不代表未来融合 GDN
后端的最终性能。Ulysses 在 eager 后端上以完整序列、较少 local heads
执行大量小算子,kernel launch 和小 shape 效率较差,因此该结果不能直接
外推到融合算子实现。checkpoint 后 P2P/AllGather 与 Ulysses 的单层峰值
差距已缩小到约 5%~6%;后续专用 summary forward/backward 算子仍可继续
减少重算开销和临时 workspace。


Self-checklist:

Special notes for your reviewers:

  • 保持 ulysses 为默认模式,未配置用户行为不变。
  • 当前不支持 Linear Attention 的 TP+CP 组合,会显式抛出
    NotImplementedError
  • 不新增第三方依赖,不引入 Triton/MindSpeed 运行时依赖。
  • 当前 P2P/AllGather 使用 eager GDN summary,后续融合后端可在保持 CP
    通信接口不变的情况下替换本地算子。
likedislike
Pull Request已成功合入, 合并人@liuchongming74
(感谢 xu-xianliang 的贡献)
xu-xianliangxu-xianliang
7月13日 创建了 pull request,commit f7bc2dc1
atomgit-bot
atomgit-bot
7月13日 评论:

变更摘要

该 PR 为 Qwen3.5 风格的 Gated DeltaNet(线性注意力)层新增了上下文并行(Context Parallel)支持。与之前简单的序列 gather/slice 钩子方案不同,本次提供了两种 CP 执行策略:纯 Ulysses 模式(通过 differentiable all-to-all 在序列维度和头维度之间重分布 Q/K/V/B/A 张量)和 P2P 模式(每个 rank 持有一个序列分片,通过点对点通信传递循环状态,并支持 recompute 和 graph 两种反向传播实现)。同时将 CP 执行逻辑封装为 LinearAttentionContextParallel(一个 ParallelStyle 子类),并集成到 Qwen3.5 模型的并行化流程中。

主要改动

  • 新增 LinearAttentionContextParallel 类及两种 CP Wrapper:在 linear_attention_context_parallel.py 中实现 LinearAttentionUlyssesCPWrapper(纯 Ulysses 模式,通过 _differentiable_all_to_all_shard 实现序列/头维度转换)和 LinearAttentionP2PCPWrapper(序列分片模式,通过 _GDNStateP2PFunction_gdn_state_p2p_graph 实现 P2P 状态传递),统一由 LinearAttentionContextParallel.apply() 根据 mode 参数选择并挂载到目标模块。

  • 实现 P2P 循环状态传递的两种反向传播策略_GDNStateP2PFunction 采用 recompute 方式——前向通过 dist.send/dist.recv 传递 hidden state,反向时重新计算局部 GDN 并通过自定义 autograd 计算梯度;_gdn_state_p2p_graph 则将收/发操作拆分为独立的 _RecvInitialStateP2PFunction_SendFinalStateP2PFunction,利用 autograd 图直接进行端到端梯度传递。

  • 重构 Qwen3.5 中线性注意力的 CP 应用方式:删除 parallelize.py 中原有的 _redistribute_first_tensor 和基于 register_forward_pre_hook/register_forward_hook 的 gather/slice 实现,替换为调用 LinearAttentionContextParallel,并新增 linear_attention_cp_mode 配置项支持从外部指定 CP 模式。

  • 新增 TP+CP 兼容性检查与 GQA 展开逻辑修正:在 parallelize_qwen3_5 中添加线性注意力层 TP+CP 不支持的明确报错;修改 _needs_gqa_kv_expand_for_cp 的判断逻辑,增加对 local_q_heads 可被 cp_mesh.size() 整除且能被 local_kv_heads 整除的条件检查。

  • 模块导出更新:在 hyper_parallel/__init__.pyhyper_parallel/core/context_parallel/__init__.py 中注册并导出 LinearAttentionContextParallel

likedislike
不准确?
atomgit-bot
atomgit-bot
7月13日 评论:

代码审查

审查总结

共审查 4 个变更文件,发现 3 个问题(1 个为 P2/P3 重复报告,实际独立问题 2 个),按优先级分布如下:

  • P3: 2 个(均为潜在隐患或 dtype 不一致,当前路径不可达或影响极小)

逐文件审查结果

文件 审查结论
hyper_parallel/__init__.py 无问题(仅新增导入和 __all__ 导出)
hyper_parallel/core/context_parallel/__init__.py 无问题(仅新增导入和 __all__ 导出)
hyper_parallel/core/context_parallel/linear_attention_context_parallel.py 发现 2 个 P3 级问题(recon_concat_dim 负值潜在隐患 + dht dtype 不一致)
hyper_parallel/models/qwen3_5/parallelize.py 无问题(重构为使用 LinearAttentionContextParallel_needs_gqa_kv_expand_for_cp 逻辑改进,新增 TP+CP 守卫正确)

整体风险评估

低风险。新文件 linear_attention_context_parallel.py 整体设计合理,Ulysses CP 和 P2P CP 两条路径的 P2P 通信配对正确,autograd 图构建策略(send_token * 0 注入梯度、_RecvInitialStateP2PFunction / _SendFinalStateP2PFunction 分离前向/反向通信)是正确的。发现的 2 个问题均为潜在隐患或边缘情况,在正常使用场景下不可达或无明显影响。其余 3 个文件的变更为纯导入/导出或重构调用方式,无引入问题。

注:_differentiable_all_to_all_shardrecon_concat_dim 负值问题被报告了两次(第一次误判为 P2,经深入分析确认当前所有调用路径中该代码不可达,修正为 P3)。以 P3 级别报告为准。

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

💬 仅评论

likedislike
不准确?
司小南(机器人)
司小南(机器人)成员7月13日进行代码检视2
hyper_parallel/core/context_parallel/linear_attention_context_parallel.py
已过期
@@ -0,0 +116,2 @@
116+ @staticmethod
117+ def forward( # pylint: disable=arguments-differ
118+ ctx,
司小南(机器人)
司小南(机器人)7月13日评论:

此条代码评论区间+116+420

【openlibing.ci】识别到代码检查告警抑制注释,匹配工具:pylint,请Committer检视其合理性。

likedislike
System
系统消息系统
7月26日 评论:

changed this line on 9aaf9e01 view diff detail

司小南(机器人)
司小南(机器人)成员7月13日进行代码检视1
hyper_parallel/core/context_parallel/linear_attention_context_parallel.py
@@ -0,0 +278,2 @@
278+ @staticmethod
279+ def forward( # pylint: disable=arguments-differ
280+ ctx,
司小南(机器人)
司小南(机器人)7月13日评论:

此条代码评论区间+278+513

【openlibing.ci】识别到代码检查告警抑制注释,匹配工具:pylint,请Committer检视其合理性。

likedislike
此处折叠了154条消息 查看更多
Yyangzhenzhang成员
8月3日 通过审查
阿苏阿苏成员
8月3日 通过审查
Yyao_yf成员
8月3日 通过审查
MindSpore-Bot
MindSpore-Bot成员
8月3日 评论:

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

likedislike
liuchongming74liuchongming74成员
8月3日 合入了pull request,合并节点 SHA:374061bfa944b0c70f8819e023337c3b3e6d4fc1