已合并
[master] 优化 DSA 长序列上下文并行显存 #1282
[master] 优化 DSA 长序列上下文并行显存 #1282
已合并
lzy0920232创建于 26 天前
lzy0920232
26 天前

修改背景

DSA2 同步 Colossal CP 中,k_nopevalue 都来自同一份 compressed_kv_norm,原实现却分别执行 SequenceShard -> Replicate AllGather,产生两份内容相同的全局 KV storage;Indexer loss 还会再次构造可复用的全局 key-rope。

本 PR 只处理 DSA2 全局 K/V 通信结果复用,不包含 DSA1 Dense Teacher CP 实现。Dense Teacher 的 DSA 布局语义由 MindFormers PR #8759 处理,并直接复用 HP stock 通用 ContextParallel

修改方案

新增 DSASequenceReplicateCache,在同步 DSA2 CP 的同一层调用周期中:

  • k_nopevalue 共享 main_kv 槽位,只执行一次可微分 AllGather;
  • Indexer loss 复用已得到的全局 key_rope
  • key_index 保持独立 AllGather。主 Indexer forward 的自定义反向会阻断主干 K 梯度,而 Indexer loss 的 K 需要承载预计算的 Indexer 梯度,二者不能合并;
  • 缓存签名校验 shape、dtype、rank list 和 sequence dim;每层调用开始时重置,Indexer loss 参数转换完成后清理映射;
  • Async DSA CP 继续使用既有 prelaunch/wait 生命周期,不启用共享缓存;
  • 默认参数关闭共享,既有 HP 调用方行为保持不变。

为什么共享缓存放在 Hyper-Parallel

该缓存管理的不是普通业务 Tensor,而是 CP placement 转换产生的、带反向通信语义的 AllGather 结果。以下职责属于并行通信层:

  • SequenceShard -> Replicate 的执行和反向梯度转换;
  • CP mesh、rank list、sequence dim 与 Tensor 签名匹配;
  • attention style 和 indexer-loss style 之间的通信结果生命周期;
  • 同步与异步 CP 的不同通信时序。

MindFormers 只负责根据 DSA 阶段和 async 配置决定是否创建并传入共享实例。如果把缓存实现放在 MindFormers,需要复制或依赖 Hyper-Parallel 内部 placement/AllGather 细节,也难以正确控制两个 style hook 的生命周期。因此由 Hyper-Parallel 提供通信缓存能力、MindFormers 负责启用策略。

这与 Dense Teacher CP 的归属不同:Dense Teacher 只需在 MF 中规范化 attention output/LSE 的业务布局,再调用 HP stock ContextParallel,不需要新增 HP DSA 专用通信类。

功能与精度验证

128K、CP4、DSA2、10 step

指标 首步绝对误差 MAE P95 相关系数
loss 0 1.005e-4 2.862e-4 1.000000
indexer_loss 0 2.0e-7 1.0e-6 0.999999
grad_norm - 0.001088 0.002375 1.000000
  • 峰值 allocator 显存:23.8948 GiB -> 23.7698 GiB,降低 128 MiB;
  • 稳定 step 3-10:3610.500 ms/step -> 3606.625 ms/step,无性能劣化;
  • 128 MiB 与 128K * kv_lora_rank 512 * BF16 2 bytes 完全一致,收益随序列长度线性增长。

TP2+CP2、4K、50-step 高学习率压力 A/B

指标 首步绝对误差 MAE P95 相关系数
loss 0 0.004787 0.016029 0.999970
indexer_loss 0 4.054e-5 1.8085e-4 0.999714
grad_norm 1.95e-4 0.227074 1.105011 0.992013

loss 与 indexer_loss 满足既定精度标准。

测试

  • Python compile check、diff check:通过;
  • 新增 UT 覆盖 K/V 共享、key-rope 复用、Indexer K 独立以及共享后父 latent 梯度累加语义;
  • MindSpore NPU 128K CP4、TP2+CP2 功能和精度 A/B:通过。
likedislike
Pull Request已成功合入, 合并人@MindSpore-Bot
(感谢 lzy0920232 的贡献)
Llzy0920232
26 天前 创建了 pull request,commit 159db1ac
atomgit-bot
atomgit-bot
26 天前 评论:

变更摘要

本 PR 针对 DSA 长序列上下文并行(CP)做两方面的优化:其一,为 DSA1 Dense Teacher Attention 新增 DSADenseAttentionContextParallel 样式,使其按 YAML 配置支持 Colossal、Ulysses 和 Hybrid CP(Indexer 与 Indexer loss 仍保持 Colossal 调度,DSA2 Sparse Attention 并行算法不变);其二,为 DSA2 同步 Colossal CP 新增 DSASequenceReplicateCache 通信缓存,让同一层调用周期内 k_nopevalue 共享 main_kv 槽位、只执行一次可微分 AllGather,Indexer loss 复用全局 key_rope,从而消除重复的全局 KV storage(128K CP4 场景峰值显存降低 128 MiB,与理论收益一致)。缓存职责放在 Hyper-Parallel 层,MindFormers 仅负责按 DSA 阶段与 async 配置决定是否创建并传入共享实例。

主要改动

  • 新增 DSADenseAttentionContextParallel:在 dsa_context_parallel.py 中新增该类,复用基类 ContextParallel 的 pre-hook,仅重写 Ulysses/Hybrid 反向 all-to-all(_post_hook_ata/_post_hook_hybrid),按 output_dims 分别处理 attention output 与 softmax max/sum 的独立轴序,并保持 local-output 语义;同时支持 BSND/TND 布局及 tnd_lse_canonical 契约校验。
  • 新增 DSASequenceReplicateCache 共享通信缓存:提供 begin/replicate/clear 接口,replicate 按 slot 名对 shape、dtype、rank list 与 sequence dim 做签名校验,首消费者执行可微分 AllGather,后续消费者复用同一 DTensor 使梯度在单次反向 collective 前累加;每层 sparse attention 开始时 begin 重置,Indexer loss 参数转换完成后 clear 释放引用。
  • DSASparseAttentionContextParallel 增加 K/V 共享配置:新增 shared_replicate_cacheshare_key_value 参数,当开启共享时将 keyvalue 都映射到 main_kv 槽位(key_rope 仍独立),Async CP 路径不启用缓存、沿用既有 prelaunch/wait 生命周期。
  • DSAIndexerLossContextParallel 复用共享缓存keymain_kv 槽位、key_ropekey_rope 槽位复用全局结果,key_indexer 保持独立 AllGather 以承载预计算 Indexer 梯度;参数转换完成后通过 completion_fn 触发缓存清理。
  • 导出与单元测试:在 hyper_parallel/__init__.pyhyper_parallel/core/context_parallel/__init__.py__all__ 与 import 中公开 DSADenseAttentionContextParallelDSASequenceReplicateCachetest_dsa_context_parallel.py 新增 dense teacher 各 CP 模式反向轴序、TND 契约校验、K/V 共享与 key_rope 复用、Indexer K 独立以及共享后父 latent 梯度累加等 UT 用例。
likedislike
不准确?
atomgit-bot
atomgit-bot
26 天前 评论:

代码审查

✅ 未发现问题

likedislike
不准确?
MindSpore-BotMindSpore-Bot成员
26 天前 添加了label:mindspore-cla/yes
司小南(机器人)司小南(机器人)成员
26 天前 添加了label:pr-check-pass
此处折叠了117条消息 查看更多
司小南(机器人)
司小南(机器人)成员
22 天前 评论:
Project Name Build_Stage Build Result Details
Hyper-parallel_Atomgit_Gate - SUCCESS 8612
likedislike
司小南(机器人)司小南(机器人)成员
22 天前 删除了label:ci-pipeline-running
司小南(机器人)司小南(机器人)成员
22 天前 添加了label:ci-pipeline-passed
Yyangzhenzhang成员
22 天前 通过审查
MindSpore-BotMindSpore-Bot成员
22 天前 合入了pull request,合并节点 SHA:e8bc77eec2bd3a6ff4cdca829db003eca247640b