已开启
[RFC]: [Feature] KVAllGather CP支持THD格式下的EOD负载均衡 #232
m0_50947149创建于  8月1日
m0_50947149
m0_50947149
8月1日 创建

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 目标

目标:

  • 支持 KVAllGather CP、THD 格式和 EOD Reset 的组合使用。
  • 对每个 EOD 子序列执行对称负载均衡切分,使各 CP rank 的理论因果 Attention 工作量相等。
  • 保证前向输出及 Q/K/V 梯度与非 CP TND Attention 参考结果一致。
  • 支持 FP16、BF16 以及 MHA、GQA、展开后的 MLA。
  • 避免完整 K/V reorder 和输出 restore,将两个本地 Q 分块合并为一次 TND 融合 Attention。
  • 缓存可复用的 cu_seqlens 元数据,减少 Transformer 层间的重复计算。

非目标:

  • 不支持非 causal Attention、Cross Attention 或 Q/K/V token 数不同的场景。
  • 不新增用户命令行参数,不改变 SBHD 路径和其他 CP 算法。
  • 本提案不保证所有模型和序列长度下均获得相同的性能收益;实际收益受计算通信比、CP 规模和 NPU 型号影响。

2. 用例分析

主要用例是在 Ascend NPU 上使用 Transformer Engine 训练 THD packed sequence,并同时开启 KVAllGather CP、causal mask 和 EOD Reset。每个 packed sequence 中包含多个相互隔离的子序列,子序列边界由累计序列长度描述。

功能和 DFX 要求如下:

  • 数据切分:每个子序列切为 2 * CP size 块,rank r 持有块 r 和块 2 * CP size - r - 1
  • 负载均衡:以 sum(q_len * kv_len) 表示理论计算量,各 rank 的该值应相等。
  • 正确性:前向输出及 Q/K/V 梯度与 CP=1 的 TND Attention 结果在对应精度容差内一致。
  • Attention 结构:支持 MHA、GQA 和展开后的 MLA。
  • 兼容性:沿用 kvallgather_cp_algo,SBHD 和其他 CP 策略行为不变。
  • 可测试性:元数据构造可在 CPU 上独立测试;前反向正确性通过双 NPU 分布式用例测试。
  • 可靠性:对 CP rank、序列长度、head 数及 head dimension 进行显式校验,错误输入应尽早失败。

使用限制:

  • 仅支持 causal self-attention 和 THD/TND 布局。
  • 每个 EOD 子序列长度必须为正,并在在线 padding 后可被 2 * context_parallel_size 整除。
  • cu_seqlens_qcu_seqlens_kv 必须表示相同的子序列。
  • Q head 数必须可被 KV head 数整除;Q/K head dimension 必须相等;K head dimension 不小于 V head dimension。

3. 方案设计

3.1 总体方案

数据加载阶段复用 _get_batch_on_this_cp_rank_in_megatron_cp_eod_padding。每个 EOD 子序列经过 padding 后被等分为 2 * CP size 块,每个 rank 获取一前一后的对称块。

Attention 阶段执行以下流程:

EOD packed sequence
        |
        v
对称负载均衡切分 Q/K/V
        |
        v
各 rank 对本地 K/V 执行 AllGather,得到 rank-major K/V
        |
        v
根据 causal 前缀生成 rank-major 直接索引
        |
        v
一次 TND Fusion Attention 计算两个本地 Q 块
        |
        v
反向使用同一索引执行 index_add_,再 ReduceScatter dK/dV

get_thd_load_balanced_cp_metadata 根据 EOD 边界、CP size 和 rank 生成:

  • 本地及全局 token 数。
  • 融合调用需要的 actual_seq_qlenactual_seq_kvlen
  • 从 rank-major AllGather 结果选择 causal K/V 前缀的索引。

前向保存索引及序列元数据供反向复用。反向将融合算子产生的 dK/dV 通过 index_add_ 累积至 rank-major 梯度缓冲区,再通过 ReduceScatter 返回各 rank 的本地梯度。

3.2 技术选型

方案 优点 缺点 结论
完整 K/V reorder,计算后 restore 实现直观 增加全序列重排、显存和带宽开销 不采用
两个本地 Q 块分别调用 TND Attention 逻辑简单 前反向各增加一次算子启动,并重复处理元数据 不采用
在 GPT forward 生成并逐层透传元数据 Attention 内开销低 改动链路长,增加 MindSpeed 与 Megatron 接口耦合 不采用
Attention 内每层重新生成元数据 无缓存状态 重复进行 cu_seqlens 转换和索引构造 不采用
rank-major 直接索引、单次 TND 调用和 LRU 缓存 改动集中,减少重排、启动及重复计算 需要维护索引和有界缓存 采用

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 校验及前反向实现。
  • 单元测试:新增数据切分、元数据、缓存以及 MHA/GQA/MLA 前反向测试。

性能设计:

  • K/V AllGather 后保持 rank-major 布局,只为当前 rank 构造 causal prefix 索引。
  • 两个本地 Q 块合并到一次 TND Fusion Attention,前向和反向各调用一次融合算子。
  • 元数据缓存容量为 8,使用 tensor identity 和 tensor version 区分输入;cu_seqlens 修改后缓存自动失效。
  • 缓存仅保存序列端点、索引及长度信息,不保存 Q/K/V 或模型参数。

本提案的验收重点是正确性、rank 间工作量均衡以及消除额外完整 reorder/restore。端到端耗时和显存需在目标模型上通过 CP=1 与 CP>1 对照测试评估,不设置与硬件无关的固定收益值。

3.4 安全隐私与DFX设计

  • 安全隐私:功能不访问网络、文件或用户明文数据;缓存不保存 token 内容,仅保存长度和索引元数据。
  • 兼容性:不新增 CLI,SBHD、Ulysses、Ring、Hamilton 等路径不变;适配 MindSpeed master 与 Megatron-Core v0.12.1。
  • 可维护性:元数据构造与 Autograd 计算解耦,并提供统一的缓存清理函数。
  • 可测试性:纯元数据逻辑可使用 CPU 测试,算子正确性使用 NPU 分布式测试。
  • 可靠性:在通信和算子调用前校验序列、rank、tensor shape、dtype 和 head 关系,防止无效索引或静默错误。
  • 资源控制:LRU 缓存容量有上限,避免持续保留不同 cu_seqlens tensor;测试结束可显式清理缓存。

3.5 编程与调用设计

该特性通过现有训练参数自动启用,不要求普通用户直接调用内部 Python 接口。

3.5.1 编程模型基本设计

开发环境:

  • Ascend NPU、PyTorch、torch_npu、Transformer Engine。
  • MindSpeed master 与 Megatron-Core v0.12.1。
  • 使用 PyTest 执行单元测试,使用仓库 pre-commit 工具链执行静态检查。

开发约束:

  • 仅在 kvallgather_cp_algo、THD、causal self-attention 和 EOD Reset 组合下使用新增路径。
  • 每个 EOD 子序列需满足 2 * CP size 对齐要求。
  • 当前通信仍采用 K/V AllGather 和 dK/dV ReduceScatter,不在本提案中引入双流通信计算重叠。

可验收设计:

  • CP=2、CP=4 元数据工作量相等。
  • FP16 容差为 atol=5e-3, rtol=5e-3
  • BF16 容差为 atol=2.5e-2, rtol=2.5e-2
  • MHA、GQA、展开后的 MLA 均通过双 NPU 前反向对照测试。
3.5.2 接口定义与设计

本提案不新增公共 API。以下为 MindSpeed 内部集成接口。

3.5.2.1 get_thd_load_balanced_cp_metadata
  • 接口描述:生成当前 CP rank 的 THD 负载均衡索引及融合 Attention 长度元数据。
  • 接口原型:
get_thd_load_balanced_cp_metadata(cu_seqlens, cp_size, rank, device) -> dict
  • 输入参数:
参数名称 输入/输出 类型 描述 取值范围
cu_seqlens 输入 Sequence[int] EOD 子序列累计结束位置 非空、严格递增
cp_size 输入 int CP 通信域大小 大于等于 1
rank 输入 int 当前 CP rank [0, cp_size)
device 输入 torch.device 索引 tensor 所在设备 CPU 或 NPU
  • 返回参数:
参数名称 类型 描述 取值范围
metadata dict token 数、rank-major 索引和 Q/KV 累计长度 由输入序列决定
  • 异常处理:参数非法、子序列长度非正或未按 2 * cp_size 对齐时抛出 AssertionError
  • 约束说明:该接口为内部接口,不保证跨版本稳定。
  • 变更说明:新增接口。
3.5.2.2 AttnFuncWithCPAndKVAllGatherForTHDLoadBalanced
  • 接口描述:执行 EOD 负载均衡 THD KVAllGather Attention 的前向和反向。
  • 接口原型:
AttnFuncWithCPAndKVAllGatherForTHDLoadBalanced.apply(
    q, k, v, n_head, attention_mask, qkv_format, attn_mask_type,
    attention_dropout, softmax_scale, deterministic, cp_group,
    cu_seqlens_q, cu_seqlens_kv,
)
  • 输入/输出参数:
参数名称 输入/输出 类型 描述 取值范围
q/k/v 输入 torch.Tensor 本地 THD Q/K/V 三维、同 dtype、同 token 数
n_head 输入 int Query head 数 大于 0
attention_mask 输入 torch.Tensor causal mask 算子支持的 mask
qkv_format 输入 str QKV 布局 thd
attn_mask_type 输入 str mask 类型 包含 causal
attention_dropout 输入 float Attention dropout [0, 1)
softmax_scale 输入 float/None Softmax 缩放系数 None 或正数
deterministic 输入 bool 确定性配置 True/False
cp_group 输入 ProcessGroup CP 通信组 有效通信组
cu_seqlens_q/kv 输入 Sequence[int] Q/KV EOD 累计结束位置 两者等价
output 输出 torch.Tensor 当前 rank 的 Attention 输出 [local_tokens, q_heads, v_dim]
  • 异常处理:不支持的 mask、shape、head 关系或序列布局抛出 AssertionError
  • 约束说明:由 KVAllGatherCPStrategy 调用,不建议用户直接调用。
  • 变更说明:THD KVAllGather 路径由原实现切换到负载均衡实现,SBHD 不变。

调用方式沿用现有参数:

--transformer-impl transformer_engine \
--context-parallel-size 2 \
--context-parallel-algo kvallgather_cp_algo \
--attention-mask-type causal \
--reset-attention-mask \
--variable-seq-lengths
3.5.3 编程手册设计

RFC 评审通过后,在现有《KVAllGather长序列并行》文档中补充 THD EOD Reset 负载均衡原理、启用参数、支持的 Attention 结构、对齐约束和验证方法,不单独新增用户操作手册。

4. 测试设计

单元测试:

  • 验证 CP=2、CP=4 下的对称切分和元数据正确性。
  • 验证各 rank 的 sum(q_len * kv_len) 相等。
  • 验证 rank-major K/V 索引及梯度 index_add_ 结果。
  • 验证相同 tensor 命中 LRU 缓存,tensor 修改后缓存失效。

集成测试:

  • 双 NPU 验证 EOD batch 切分结果。
  • 双 NPU 分别使用 FP16、BF16 验证 MHA、GQA、展开后的 MLA。
  • 将 CP 前向输出及 Q/K/V 梯度与非 CP TND Fusion Attention 参考结果比较。

端到端测试:

  • 使用相同数据、模型和随机种子运行 CP=1 与 CP=2 训练。
  • 对比 LM loss、grad norm、elapsed time per iteration 和日志中的 max allocated
  • 检查训练无发散、无死锁、无 OOM,loss 和梯度误差符合精度预期。

测试命令:

pytest -v tests_extend/unit_tests/mindspeed/te/test_kvallgather_context_parallel.py

5. 缺点和风险

  • K/V AllGather 在 MHA 或较大 KV head 数下通信量较大,性能可能低于 CP=1;应结合模型结构和序列长度评估。
  • index_selectindex_add_ 和梯度缓冲区会产生额外显存及内存带宽开销,但低于完整 reorder/restore 方案。
  • LRU 缓存会短期持有 cu_seqlens tensor 引用;容量限制为 8,并提供清理接口降低长期占用风险。
  • 仅支持 causal self-attention,错误配置会显式失败,不能自动回退到通用 THD 路径。
  • EOD padding 会改变实际 token 数;数据管线必须开启变长序列支持并保证长度对齐。
  • 本方案与 Megatron-Core v0.12.1 适配,升级 Megatron 后需重新检查调用签名和 THD 元数据语义。

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. 未解决问题

  • 是否需要为非 causal THD Attention 提供通用负载均衡策略。
  • 是否进一步通过双流实现 K/V AllGather 与 TND Attention 的通信计算重叠。
  • 是否将元数据前移至数据加载或 GPT forward,以进一步减少 Attention 内开销。
  • 是否为不同模型和 NPU 型号定义统一的性能验收阈值。
  • 后续 Megatron-Core 版本升级时,内部接口应采用何种兼容策略。

附录

  • 参考资料链接
  • 术语表
    • CP:Context Parallel,上下文并行。
    • THD/TND:以总 token 数、head 数、head dimension 组织的变长序列布局。
    • EOD:End of Document,文档结束边界。
    • MHA/GQA/MLA:Multi-Head、Grouped-Query、Multi-head Latent Attention。
    • rank-major:AllGather 后按 rank 顺序拼接的 tensor 布局。

欢迎加入社区,感谢您对社区的贡献 🎉!

likedislike
ascend-robotascend-robot成员
8月1日 添加了label:rfc
m0_50947149m0_50947149
8月1日 修改了issue 的描述
m0_50947149m0_50947149
8月1日 修改了issue 的描述
m0_50947149m0_50947149
8月1日 修改了issue 的描述
m0_50947149m0_50947149
26 天前 修改了issue 的描述
m0_50947149m0_50947149
26 天前 修改了issue 的描述
m0_50947149m0_50947149
26 天前 修改了issue 的描述
m0_50947149m0_50947149
26 天前 修改了issue 的描述
m0_50947149m0_50947149
26 天前 issue类型由 Bug-Report 改变为 RFC
m0_50947149m0_50947149
26 天前 修改了issue 的描述
m0_50947149m0_50947149
8 天前 关联了pull request:feat: adapt kvallgather THD load balance to core r0.16.0