已合并
[master] 优化 DSA 长序列上下文并行显存 #1282
lzy0920232创建于 26 天前
[master] 优化 DSA 长序列上下文并行显存 #1282
已合并
Pull Request已成功合入, 合并人@MindSpore-Bot
(感谢 lzy0920232 的贡献)26 天前 创建了 pull request,commit 159db1ac
atomgit-bot
26 天前 评论:
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_nope 与 value 共享 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_cache与share_key_value参数,当开启共享时将key与value都映射到main_kv槽位(key_rope仍独立),Async CP 路径不启用缓存、沿用既有 prelaunch/wait 生命周期。DSAIndexerLossContextParallel复用共享缓存:key走main_kv槽位、key_rope走key_rope槽位复用全局结果,key_indexer保持独立 AllGather 以承载预计算 Indexer 梯度;参数转换完成后通过completion_fn触发缓存清理。- 导出与单元测试:在
hyper_parallel/__init__.py与hyper_parallel/core/context_parallel/__init__.py的__all__与 import 中公开DSADenseAttentionContextParallel、DSASequenceReplicateCache;test_dsa_context_parallel.py新增 dense teacher 各 CP 模式反向轴序、TND 契约校验、K/V 共享与key_rope复用、Indexer K 独立以及共享后父 latent 梯度累加等 UT 用例。


不准确?
atomgit-bot
26 天前 评论:
26 天前 评论:
26 天前 添加了label:mindspore-cla/yes
26 天前 添加了label:pr-check-pass
此处折叠了117条消息 查看更多
司小南(机器人)
22 天前 评论:
22 天前 评论:
| Project Name | Build_Stage | Build Result | Details |
|---|---|---|---|
| Hyper-parallel_Atomgit_Gate | - | ✅ SUCCESS | 8612 |


22 天前 删除了label:ci-pipeline-running
22 天前 添加了label:ci-pipeline-passed
22 天前 通过审查
22 天前 合入了pull request,合并节点 SHA:e8bc77eec2bd3a6ff4cdca829db003eca247640b
修改背景
DSA2 同步 Colossal CP 中,
k_nope与value都来自同一份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_nope与value共享main_kv槽位,只执行一次可微分 AllGather;key_rope;key_index保持独立 AllGather。主 Indexer forward 的自定义反向会阻断主干 K 梯度,而 Indexer loss 的 K 需要承载预计算的 Indexer 梯度,二者不能合并;为什么共享缓存放在 Hyper-Parallel
该缓存管理的不是普通业务 Tensor,而是 CP placement 转换产生的、带反向通信语义的 AllGather 结果。以下职责属于并行通信层:
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
128K * kv_lora_rank 512 * BF16 2 bytes完全一致,收益随序列长度线性增长。TP2+CP2、4K、50-step 高学习率压力 A/B
loss 与 indexer_loss 满足既定精度标准。
测试