已开启
[RFC]: [Feature] KVAllGather CP支持THD格式下的EOD负载均衡 #232
m0_50947149创建于 8月1日
8月1日 添加了label:rfc
8月1日 修改了issue 的描述
8月1日 修改了issue 的描述
8月1日 修改了issue 的描述
26 天前 修改了issue 的描述
26 天前 修改了issue 的描述
26 天前 修改了issue 的描述
26 天前 修改了issue 的描述
26 天前 issue类型由 Bug-Report 改变为 RFC
26 天前 修改了issue 的描述
8 天前 关联了pull request:feat: adapt kvallgather THD load balance to core r0.16.0
1. 概述
1.1 简介
本提案为 KVAllGather Context Parallel(CP)补充 THD 数据格式下的 EOD 感知负载均衡能力。方案沿用 Megatron CP 的对称切分规则,使每个 CP rank 获得因果 Attention 计算量相等的两个序列分块,并在 MindSpeed 内完成适配。
方案支持 MHA、GQA 和展开后的 MLA,通过 rank-major K/V 直接索引、单次 TND 融合 Attention 和元数据缓存,减少重排、算子启动及重复元数据计算开销。
1.2 动机
KVAllGather CP 已支持 SBHD 和 THD 格式,但原有 THD EOD Reset 路径采用普通序列切分。因果 Attention 中,序列前部和后部的计算量不同,普通切分会造成各 rank 工作量不均,慢 rank 成为整次迭代的性能瓶颈。
THD 常用于多个变长样本拼接后的无填充训练。缺少负载均衡会限制 KVAllGather CP 在长序列、EOD Reset、GQA/MLA 等场景的可用性。为 THD 增加与 Megatron CP 一致的对称切分,可在保持数学等价的同时改善 rank 间负载均衡,并降低单卡序列相关显存压力。
1.3 目标
目标:
cu_seqlens元数据,减少 Transformer 层间的重复计算。非目标:
2. 用例分析
主要用例是在 Ascend NPU 上使用 Transformer Engine 训练 THD packed sequence,并同时开启 KVAllGather CP、causal mask 和 EOD Reset。每个 packed sequence 中包含多个相互隔离的子序列,子序列边界由累计序列长度描述。
功能和 DFX 要求如下:
2 * CP size块,rankr持有块r和块2 * CP size - r - 1。sum(q_len * kv_len)表示理论计算量,各 rank 的该值应相等。kvallgather_cp_algo,SBHD 和其他 CP 策略行为不变。使用限制:
2 * context_parallel_size整除。cu_seqlens_q与cu_seqlens_kv必须表示相同的子序列。3. 方案设计
3.1 总体方案
数据加载阶段复用
_get_batch_on_this_cp_rank_in_megatron_cp_eod_padding。每个 EOD 子序列经过 padding 后被等分为2 * CP size块,每个 rank 获取一前一后的对称块。Attention 阶段执行以下流程:
get_thd_load_balanced_cp_metadata根据 EOD 边界、CP size 和 rank 生成:actual_seq_qlen和actual_seq_kvlen。前向保存索引及序列元数据供反向复用。反向将融合算子产生的 dK/dV 通过
index_add_累积至 rank-major 梯度缓冲区,再通过 ReduceScatter 返回各 rank 的本地梯度。3.2 技术选型
cu_seqlens转换和索引构造3.3 功能与性能设计
功能影响范围:
get_batch_utils.py:使kvallgather_cp_algo的 causal EOD Reset 场景进入 Megatron CP EOD padding 切分路径。context_parallel.py:THD 格式路由至负载均衡 Autograd Function。kvallgather_context_parallel.py:新增元数据、缓存、shape 校验及前反向实现。性能设计:
cu_seqlens修改后缓存自动失效。本提案的验收重点是正确性、rank 间工作量均衡以及消除额外完整 reorder/restore。端到端耗时和显存需在目标模型上通过 CP=1 与 CP>1 对照测试评估,不设置与硬件无关的固定收益值。
3.4 安全隐私与DFX设计
cu_seqlenstensor;测试结束可显式清理缓存。3.5 编程与调用设计
该特性通过现有训练参数自动启用,不要求普通用户直接调用内部 Python 接口。
3.5.1 编程模型基本设计
开发环境:
开发约束:
kvallgather_cp_algo、THD、causal self-attention 和 EOD Reset 组合下使用新增路径。2 * CP size对齐要求。可验收设计:
atol=5e-3, rtol=5e-3。atol=2.5e-2, rtol=2.5e-2。3.5.2 接口定义与设计
本提案不新增公共 API。以下为 MindSpeed 内部集成接口。
3.5.2.1 get_thd_load_balanced_cp_metadata
get_thd_load_balanced_cp_metadata(cu_seqlens, cp_size, rank, device) -> dict[0, cp_size)2 * cp_size对齐时抛出AssertionError。3.5.2.2 AttnFuncWithCPAndKVAllGatherForTHDLoadBalanced
thdcausal[0, 1)[local_tokens, q_heads, v_dim]AssertionError。KVAllGatherCPStrategy调用,不建议用户直接调用。调用方式沿用现有参数:
3.5.3 编程手册设计
RFC 评审通过后,在现有《KVAllGather长序列并行》文档中补充 THD EOD Reset 负载均衡原理、启用参数、支持的 Attention 结构、对齐约束和验证方法,不单独新增用户操作手册。
4. 测试设计
单元测试:
sum(q_len * kv_len)相等。index_add_结果。集成测试:
端到端测试:
max allocated。测试命令:
5. 缺点和风险
index_select、index_add_和梯度缓冲区会产生额外显存及内存带宽开销,但低于完整 reorder/restore 方案。cu_seqlenstensor 引用;容量限制为 8,并提供清理接口降低长期占用风险。6. 现有技术
KVAllGather CP 的基本思想参考 Llama 3:各 rank 持有局部 Q,并通过 AllGather 获取 K/V。MindSpeed 已有 SBHD 负载均衡实现,本提案将 Megatron CP 的对称因果切分扩展到 THD packed sequence。
与普通 Megatron Ring CP 不同,本方案一次性 AllGather K/V,并使用直接索引构造本地 causal prefix;与原 THD KVAllGather 实现相比,本方案增加 EOD 粒度的负载均衡,并移除完整 reorder/restore 和双 TND step。
7. 未解决问题
附录
欢迎加入社区,感谢您对社区的贡献 🎉!