mindspore-ai / mindspore-ai/hyper-parallel
【RFC】新 Trainer Context Parallel 能力迁移与统一接入
Nobody has claimed this yet.
- Dominant language
- Python
- Stars
- 53
- Forks
- 63
- Avg merge
- 23h 45m
- Merged PRs (30d)
- 63
Description
HyperModels 新 Trainer Context Parallel 接入设计方案
0. 基本信息与术语
| 项目 | 内容 |
|---|---|
| 目标入口 | TextTrainer / BaseTrainer / HyperAutoModel |
| 主 mesh | dp_replicate、dp_shard、cp、tp;EP 为派生逻辑分组 |
| Phase 1 | 核心 CP 功能闭环:Colossal、Pure Ulysses、Hybrid、Async、Head-Tail load balance |
| Phase 2 | fused/varlen backend、模型专用 attention 和组合并行扩展 |
| Phase 3 | 大规模性能优化、通信调度优化和生产化能力完善 |
| 首期后端 | Ascend NPU/HCCL;功能、精度和性能验收均直接使用 NPU/HCCL |
| 术语 | 含义 |
|---|---|
| Colossal CP(AllGather CP) | Q 保持 local sequence,K/V 在 CP group 内 AllGather 为 global sequence |
| Ulysses CP | Q/K/V 从 sequence shard A2A 到 head shard,attention 后再转回 sequence shard |
| Hybrid CP | 在 Ulysses 子组执行 head A2A,在 Colossal 子组执行 K/V AllGather |
| Async CP | 在安全 handoff 点发起 collective,在 attention consumer 前等待,以覆盖 projection/RoPE 后的可并行计算 |
| Head-Tail load balance | 对 Colossal CP 的 Q sequence 做 Head-Tail 重排和双 attention,平衡 causal attention 计算量 |
| CP rank | 持有一个连续 sequence shard 的 rank |
| CP wrapper | 适配模型 attention forward 签名、并调用 CP strategy 的 wrapper |
1. 目标与范围
1.1 目标
新 Trainer 需要提供从 YAML 到训练 step 的完整 CP 能力:
CP 拓扑与方法声明
-> 主 DeviceMesh / CP 子组
-> batch sequence shard
-> attention wrapper + CP collective
-> local logits/loss
-> DP+CP loss 聚合
-> FSDP/TP 梯度同步
-> optimizer step
CP 的标准使用也必须通过 plan_overrides 显式声明 attention 方法。Planner 可以识别
attention boundary 和默认 placement,但不得替用户选择 CP 通信方法;Applier 只应用
用户声明的 wrapper/compute。这样可以避免模型结构被误判后静默使用错误的通信语义。
1.2 Phase 1 目标
交付范围如下:
| 场景 | Phase 1 交付要求 | 关键约束 |
|---|---|---|
| 同步 Colossal | 必须交付 | Q 保持 local sequence,K/V 在完整 cp_mesh 上 AllGather |
| 同步 Pure Ulysses | 必须交付 | ulysses_degree == cp_size,Q/K/V head 满足整除约束 |
| 同步 Hybrid | 必须交付 | 1 < ulysses_degree < cp_size 且 cp_size % ulysses_degree == 0 |
| 异步 Colossal | 必须交付 | K/V 在最后一个安全 handoff 后发起异步 AllGather,attention 前等待 |
| 异步 Pure Ulysses | 必须交付 | Q/K/V 在安全 handoff 后发起异步 A2A,attention 前等待 |
| 异步 Hybrid | 必须交付 | Ulysses 子组 A2A 和 Colossal 子组 K/V gather 建立显式依赖链,不允许伪异步 |
| Head-Tail load balance | 必须交付 | 仅作用于 Colossal 语义;重排、mask、双 attention、输出还原和 backward 必须闭环 |
具体目标:
- YAML 只声明
cp_size和每个 attention 边界使用的 CP 方法;方法通过
plan_overrides.inner_wrapper(或自定义local_compute_fn)注入。 - 同步、异步和 load-balance 不是全局开关,而是不同的显式 wrapper。Hybrid 的
ulysses_degree等方法私有数据可以写在inner_wrapperTarget 内;异步 handoff 由
模型专用 wrapper 在代码中显式处理,当前实现不提供通用AsyncCPHandoffProvider注册协议。
不得重新引入accelerator.cp.algorithm/schedule/ulysses_degree。 - CP 方法的 mesh、layout、head、handoff 和 batch contract 在 apply/preflight 阶段校验;
不提供算法默认值、隐式 wrapper 选择或 runtime fallback。 - 提供 Colossal AllGather、Ulysses A2A、Hybrid 子组通信、异步 launch/wait 和
Head-Tail 重排所需的可微 primitive;底层 collective 统一复用platform。 - 标准 HF SDPA/QKV attention 至少提供同步和异步接入路径;无法安全识别 handoff 的
HF 模型使用模型专用 wrapper,未匹配的 attention 明确报错。 - 固化 global position、causal mask、labels、padding、packed metadata、local loss、
backward 和 optimizer-step contract。 - 对上述七个交付场景完成 CP=1/CP>1 的 logits、loss、gradient 和 optimizer-step parity;
异步场景还必须通过 NPU profiler 证明存在有效通信计算重叠。
Phase 1 不要求任意能力自由组合。支持单元是一个明确命名或 Target 声明的 wrapper;例如
async + load_balance 只有在存在独立 wrapper 和完整 contract 时才算支持,不能把两个开关
拼接后自动生成行为。Ulysses/Hybrid 与 Head-Tail load balance 的组合默认非法并 fail-fast。
1.3 非目标
以下能力不阻塞 Phase 1:
- DSA indexer/sparse attention/indexer loss;
- Qwen3.5-VL/MRoPE、复杂多模态输入;
- PP+CP、EP+CP 生产组合和不同 CP size 的 checkpoint reshard;
- 未注册的 FlashAttention、NPU fused attention 和任意自定义 varlen kernel。
2. 总体方案
2.1 分层架构
Trainer YAML / CLI
|
v
TrainerConfig
accelerator.cp_size + plan_overrides
|
v
DistributedSetup / MeshContext
device_mesh / cp_mesh / dp_cp_mesh
|
+-----------------------+
| |
v v
PlanOverrideResolver ShardingPlanner
match/when 校验 parameter + boundary placement
| |
+-----------+-----------+
v
ShardingPlan
module specs + explicit CP intent
|
v
ShardingApplier
TP placement -> declared CP wrapper -> FSDP/HSDP
|
v
Model attention boundary
|
v
CPStrategy runtime
Colossal / Ulysses / Hybrid / Async / LoadBalance
|
v
local logits/loss/backward
2.2 职责边界
| 模块 | 职责 |
|---|---|
TrainerConfig |
读取 cp_size 和 plan_overrides,不保存通用 CP 算法配置 |
MeshContext |
构造 cp_mesh、dp_cp_mesh,提供 rank/group/topology |
PlanOverrideResolver |
按 FQN 和 when: cp 合并用户声明,缺失或冲突时 fail-fast |
ShardingPlanner |
识别 attention/MLP/MoE 边界,生成 placement;不推断 CP 方法 |
ShardingApplier |
只注入显式声明的 inner_wrapper/local_compute_fn,处理双模式 |
CPStrategy |
作为 wrapper 内部实现 Colossal/Ulysses/Hybrid、异步 launch/wait、Head-Tail 和 autograd |
cp_utils.py |
保存同步/异步 tensor collective、layout transform、Head-Tail 和 batch shard helper |
BaseTrainer/TextTrainer |
在 forward 前准备 local batch,执行 loss/gradient/optimizer 生命周期 |
mean_global_loss |
在 DP+CP 域聚合 token 和 loss,生成 backward/logging 结果 |
CPStrategy 是 wrapper 内部运行时接口,不作为用户 YAML 的通用配置对象:wrapper 负责适配
模型签名并明确选择 strategy,strategy 负责通信算法。每个包含 CP collective 的注入都必须
声明 region_dispatch: false,防止声明式边界再次发起冲突通信。
2.3 一次训练 step
1. DataLoader/collator 生成 global micro batch
2. PlanOverrideResolver 校验 `when: cp` 条目、目标模块和显式 CP 方法
3. CP batch prepare:global labels/positions/mask -> local sequence shard
4. model forward:由 YAML 指定的 wrapper 执行 local Q + strategy(global K/V 或 head A2A)
5. local logits/loss;不在 CP 边界 gather logits
6. mean_global_loss 在 dp_cp_mesh 聚合有效 token 和 loss
7. backward;FSDP/TP 在正确同步域归约参数梯度
8. optimizer.step;同一 CP group 的 replicated parameter 保持一致
3. 对外接口
3.1 Trainer YAML
cp_size 只表示 CP 拓扑 degree。CP 算法、attention forward 约定和通信边界必须在
plan_overrides 中逐模块声明;不提供 accelerator.cp.algorithm、schedule、
ulysses_degree、load_balance 或 async_fallback 等通用 CP 配置字段。
accelerator:
dp_shard_size: 1
dp_replicate_size: 1
tp_size: 1
cp_size: 2
ep_size: 1
pp_size: 1
sequence_parallel: false
loss_parallel: false
plan_overrides:
# HF 风格 forward(hidden_states),内部调用 F.scaled_dot_product_attention。
# sdpa_hf 的通信方法由 wrapper 固定,框架不再根据模型结构自动猜测。
- match: "*.self_attn"
when: cp
region_dispatch: false
inner_target: self
inner_wrapper: sdpa_hf
字段语义和校验:
| 字段 | 约束 |
|---|---|
cp_size |
>=1,必须与主 mesh/world size 匹配 |
plan_overrides.when |
cp_size>1 时必须命中;cp_size=1 可跳过,但应记录 INFO |
inner_target |
指定被包装的 inner attention;无法自动定位时必须显式填写 |
inner_wrapper |
注册表名或 _target_ callable,唯一决定该边界的 CP 方法 |
local_compute_fn |
仅用于用户接管完整 local-region compute;函数内含 collective 时必须伴随 region_dispatch: false |
region_dispatch |
CP wrapper/compute 含通信或自定义 kernel 时必须为 false,不得省略 |
| 未命中/缺少方法 | CP>1 在 apply 前直接报错,不自动选择 wrapper、不回退 AllGather |
同一份 YAML 可以通过 when: cp 在 CP=1 调试时跳过 CP 注入,但这不等于存在默认
CP 方法。CP>1 时每一个实际 attention boundary 都必须有且只有一个明确的方法声明。
内置 wrapper 的选择示例:
| attention 形态 | inner_wrapper |
CP 方法 |
|---|---|---|
HF forward(hidden_states) + SDPA |
sdpa_hf |
K/V AllGather |
分离 QKV forward(q, k, v) + SDPA |
sdpa_qkv |
K/V AllGather |
| HF + FlexAttention | flex_hf |
由 wrapper 约定的 Flex CP |
| 分离 QKV + FlexAttention | flex_qkv |
由 wrapper 约定的 Flex CP |
| Pure Ulysses | 独立 Ulysses wrapper 或 _target_ callable |
固定 sequence/head A2A 语义 |
| Hybrid | _target_ callable |
wrapper-local ulysses_degree 决定 Ulysses/Colossal 子组 |
| Async Colossal/Ulysses/Hybrid | 独立 async _target_ callable |
wrapper 显式声明 handoff、launch、wait 和 backward contract |
| Head-Tail load balance | 独立 load-balance _target_ callable |
固定 Colossal + Head-Tail 重排语义 |
| 自定义 fused attention | 用户注册名或 _target_ callable |
wrapper 自行实现并明确校验 layout/backend |
Ulysses/Hybrid/Async/load balance 等方法不能通过额外的 YAML algorithm 或 schedule
字段覆盖 wrapper 行为;需要使用对应的 wrapper 实现(或用户自定义 callable)。通信方法
由 wrapper 类型固定,ulysses_degree 等方法私有数据通过该 Target 的具名参数提供并由 wrapper
校验。handoff 不是普通配置数据:模型专用 wrapper 必须在代码中解析真实模块,或替换 forward
插入显式 handoff。这样配置入口仍只有一个:plan_overrides.inner_wrapper。
同一 attention 只能选择一个 wrapper。Pure Ulysses 不需要额外 degree 参数,因为其 degree
固定等于 cp_size;Hybrid wrapper 可以在自身 Target 内声明 ulysses_degree,但不能在
accelerator.cp 下另写全局 degree。wrapper 必须自行校验 cp_size、head 整除关系、
sequence/head layout、handoff 和 backend 能力。
3.1.1 各 CP 场景的 plan_overrides 写法
以下示例统一采用 HF SDPA attention。公共字段相同:
plan_overrides:
- match: "*.self_attn"
when: cp
region_dispatch: false
inner_target: self
inner_wrapper: <按场景替换>
sdpa_hf、sdpa_hf_ulysses 以及 sdpa_qkv、sdpa_qkv_ulysses 是当前实现中的内置注册名;
FlexAttention 对应 flex_hf、flex_qkv 及其 Ulysses 变体。通用 CP Target 路径指向
hyper_parallel.distributed.context_parallel.wrappers 中带 @inner_wrapper 的 callable;模型专用
wrapper 则位于对应模型的 adapter 目录。wrapper 可以返回 replacement forward 交给 rewriter 安装;
外部 wrapper 也可以原地替换 target_module.forward 并返回 None。QKV 风格 attention 使用对应的
*_qkv_* wrapper,算法参数和校验规则不变。
同步 Colossal/AllGather
inner_wrapper:
_target_: hyper_parallel.distributed.context_parallel.wrappers.sdpa_hf_cp_wrapper
同步 Pure Ulysses
inner_wrapper:
_target_: hyper_parallel.distributed.context_parallel.wrappers.sdpa_hf_ulysses_cp_wrapper
Pure Ulysses 的 ulysses_degree 固定等于 cp_size,不单独配置。
同步 Hybrid
inner_wrapper:
_target_: hyper_parallel.distributed.context_parallel.wrappers.sdpa_hf_hybrid_cp_wrapper
ulysses_degree: 2
异步 Colossal/AllGather(Qwen3-MoE)
inner_wrapper:
_target_: hyper_parallel.models.qwen3_moe.adapter.distributed.context_parallel_async.qwen3_moe_async_colossal_cp_wrapper
异步 Pure Ulysses(Qwen3-MoE)
inner_wrapper:
_target_: hyper_parallel.models.qwen3_moe.adapter.distributed.context_parallel_async.qwen3_moe_async_ulysses_cp_wrapper
异步 Hybrid(Qwen3-MoE)
inner_wrapper:
_target_: hyper_parallel.models.qwen3_moe.adapter.distributed.context_parallel_async.qwen3_moe_async_hybrid_cp_wrapper
ulysses_degree: 2
其中 _target_ 是 YAML 中的字面量键名。当前异步 wrapper 直接适配 Qwen3-MoE attention
的 forward/projection/fused-attention contract,并在实现内部完成异步 launch、依赖等待和
反向通信;它不是通用 HF SDPA wrapper,也不通过 YAML 接收模型名、handoff 路径或任意字符串。
其他模型必须提供自己的 @inner_wrapper callable,并在代码中明确 handoff、handle 生命周期、
前向等待点和 backward 顺序;不存在可自动套用的 AsyncCPHandoffProvider 注册接口。
Colossal Head-Tail load balance
inner_wrapper:
_target_: hyper_parallel.distributed.context_parallel.wrappers.sdpa_hf_load_balance_cp_wrapper
Head-Tail wrapper 是 Colossal 语义的本地 tensor 实现,负责 Q 的 Head-Tail 重排、K/V
AllGather、双 SDPA、输出还原及 backward 通信。当前实现通过 wrapper 自身识别并改写目标
attention 的 forward;preflight 无法拦截或识别目标 attention 时直接报错。
模型专用异步接入应在 QK-Norm、RoPE 和 layout transform 完成后发起 collective,并尽量保留
可覆盖的 projection/预处理计算;禁止在 SDPA consumer 处才 launch 后立即 wait,并将这种没有
有效 overlap 的路径标记为异步支持。
async + load_balance 不从上述两个 wrapper 自动组合。需要支持时必须另行提供明确命名的
wrapper Target,并声明 Head-Tail 两个 Q 分支各自的 handoff、handle 生命周期和 backward 次序。
3.2 Python API
CP 不新增一个承载算法选择的通用 Python/YAML 配置对象。Planner 只接收拓扑参数和
用户声明的 plan_overrides;每个 CP wrapper 在 apply 时接收通用 mesh 上下文,并由
自身固定通信方法和输入输出 contract。
plan = ShardingPlanner(plan_overrides={
"model.layers.*.self_attn": ModuleShardingSpec(
inner_target="self",
inner_wrapper="sdpa_hf",
region_dispatch=False,
),
}).plan(
model,
device_mesh,
tp_size=1,
cp_size=2,
ep_size=1,
sequence_parallel=False,
loss_parallel=False,
)
model, tp_grad_info = apply_sharding_plan(
model,
plan,
device_mesh,
)
配置传播链:
YAML -> TrainerConfig.accelerator.cp_size + plan_overrides
-> PlanOverrideResolver(match/when/字段校验)
-> ShardingPlanner(placement + 显式 CP intent)
-> ShardingPlan.module_specs
-> ShardingApplier(只应用声明的 wrapper/compute)
-> wrapper 内部 CPStrategy
3.3 CP strategy 接口
class ContextParallelStrategy:
def validate(self, mesh, model) -> None:
...
def shard_batch(self, batch, cp_mesh) -> dict:
...
def attention(self, q, k, v, *, cp_mesh, attention_meta):
...
def launch(self, q, k, v, *, cp_mesh, handoff_meta):
"""异步 strategy 返回按 layer/invocation 隔离的 handle state。"""
...
def wait(self, state, q, k, v, *, attention_meta):
"""在 consumer 前 materialize 通信结果并恢复 layout/autograd。"""
...
def restore_output(self, output, *, cp_mesh, attention_meta):
...
实现要求:
AllGatherStrategy复用cp_utils.flex_cp_allgather;UlyssesStrategy实现 sequence/head A2A,并提供可微 backward;HybridStrategy管理 Ulysses 子组和 AllGather 子组;- Async strategy 管理安全 handoff、异步 launch/wait、反向逆通信和 handle 生命周期;
- Head-Tail strategy 管理 Colossal Q 重排、双 attention、mask 和输出还原;
- strategy 由显式 wrapper 固定,不读取
accelerator.cp或其他隐式全局配置; - wrapper 必须在注入时校验自己的 sequence/head/layout/backend 约束;
- strategy 不依赖 HF 模型类名,只接收 layout、mask 和 mesh metadata。
3.4 模型/Wrapper 扩展接口
Planner 可以通过模板识别 attention 的 placement,但不再设置隐式的 CP wrapper。标准
模型和非标准模型都由用户在 YAML 中指定 inner_target、inner_wrapper 或
local_compute_fn。模型差异只影响 wrapper 的实现,不需要新建一套 Trainer。
标准 HF attention 的显式声明:
plan_overrides:
- match: "*.self_attn"
when: cp
region_dispatch: false
inner_target: self
inner_wrapper:
_target_: hyper_parallel.distributed.context_parallel.wrappers.sdpa_hf_cp_wrapper
非标准模型可注册模型规格,但规格只提供目标 FQN、batch contract 和可用 wrapper,不能
绕过 YAML 的方法选择:
register_cp_model_spec(
architecture="deepseek_v3",
attention_targets=("*.self_attn",),
batch_prepare_fn=prepare_deepseek_v3_cp_batch,
wrappers=("deepseek_v3_mla_all_gather", "deepseek_v3_mla_ulysses"),
)
用户自定义 attention wrapper 的约定:
@inner_wrapper
def my_cp_attention_wrapper(
target_module,
mesh,
tp_mesh,
cp_mesh,
ep_mesh,
):
"""原地替换 target_module.forward;内部明确实现一种 CP 通信方法。"""
del mesh, tp_mesh, ep_mesh
original_forward = target_module.forward
def cp_forward(*args, **kwargs):
# 1. 解包 local/DTensor 输入;2. 按本 wrapper 的方法通信;
# 3. 执行 attention;4. 恢复 boundary contract。
return original_forward(*args, **kwargs)
target_module.forward = cp_forward
然后在 YAML 中显式绑定实现:
plan_overrides:
- match: "*.self_attn"
when: cp
region_dispatch: false
inner_target: core_attention
inner_wrapper:
_target_: my_project.cp.my_cp_attention_wrapper
plan_overrides 是唯一的 CP 注入入口,用于指定每个模块实际使用的 CP 方法:
spec = ModuleShardingSpec(
inner_target="self",
inner_wrapper="my_cp_wrapper",
region_dispatch=False,
)
planner = ShardingPlanner(
plan_overrides={"model.layers.0.self_attn": spec}
)
inner_wrapper 与 local_compute_fn 二选一:前者替换/包装 attention inner forward,
后者替换 local-region 的计算函数。两者都必须由用户显式声明;CP>1 时只声明
inner_target 或只依赖模板推断均不构成完整 CP 接入。
3.5 低层 cp_utils/platform 接口
Colossal 继续复用现有 AllGather primitive:
local_batch = shard_batch_for_cp(global_batch, mesh_context.cp_mesh)
global_k, global_v = flex_cp_allgather(
local_k,
local_v,
cp_dim=2,
cp_mesh=mesh_context.cp_mesh,
)
其余场景需要提供同一层级的能力:
- Ulysses:sequence-to-head/head-to-sequence 可微 A2A;
- Hybrid:稳定
(co, ds)子 mesh、A2A + K/V AllGather layout transform; - Async:platform-owned launch、handle、wait、stream/event 和可微逆通信;
- Head-Tail:Q 重排/还原、双 attention metadata 和 backward 映射。
实现约束:所有 collective 复用 cp_mesh.get_group() 或从 root mesh 缓存得到的稳定子组,禁止在
forward 中 dist.new_group();底层通信不得绕过 platform。AllGather backward 必须提供
reduce-scatter 语义,A2A backward 必须执行逆 layout transform;cp_size=1 返回 identity。
4. 核心实现设计
4.1 Mesh 与通信域
主 mesh:
world_size = pp × dp_replicate × dp_shard × cp × tp
CP 相关子组:
| 子组 | 用途 |
|---|---|
cp_mesh |
attention AllGather 或 Ulysses A2A |
dp_cp_mesh |
global loss/token 聚合 |
| Ulysses subgroup | Hybrid 中 sequence/head A2A |
| Colossal subgroup | Hybrid 中 K/V AllGather |
| FSDP mesh | 参数/梯度同步;至少覆盖需要同步的 DP 与 CP 域 |
EP 不进入主 mesh 乘积;EP mesh 在 apply 阶段从 dense rank domain 派生。PP 首期固定为 1。
4.2 Batch contract
CP 切分前 batch 必须仍是 global sequence。处理顺序:
global collated batch
-> 生成 global position_ids、labels/shift_labels、padding metadata
-> 按 strategy 的 padding policy 对齐
-> 按 cp_rank 切出 local sequence
-> model forward
| 字段 | 规则 |
|---|---|
input_ids |
沿 sequence 维切分 |
labels |
先保证 global next-token 语义,再切分 |
shift_labels |
global shift 后切分,禁止各 rank 独立 shift 丢失边界 target |
position_ids |
global 位置生成后切分,rank 不从 0 重新编号 |
attention_mask |
根据 local-Q/global-K 或 head-shard 语义生成,禁止盲切 |
inputs_embeds |
沿 batch 的 sequence 维切分 |
seq_lens/seq_lens_padded |
按 local chunk 和 pack boundary 重算 |
past_key_values/use_cache |
训练首期拒绝 |
| 未知 tensor | 原样透传或 fail-fast,不默认沿最后一维切分 |
padding policy:
| strategy | 约束 |
|---|---|
all_gather |
S 对齐到 cp_size 的倍数 |
ulysses |
S 对齐到 CP degree,且 head 维满足 A2A 整除 |
hybrid |
同时满足 Ulysses degree 和子组约束 |
load_balance |
另行使用 2*cp_size 或 zigzag 对齐,不影响普通 AllGather |
4.3 三种同步 strategy
Colossal / AllGather
Q: [B, Nq, S/C, D]
K/V: [B, Nk, S/C, D]
-> K/V AllGather on cp_mesh
K/V: [B, Nk, S, D]
-> local-Q/global-KV attention
output: [B, Nq, S/C, D]
不要求 Q head 被 CP 整除,适合 GQA/MQA;causal mask 必须使用 local Q 的 global offset。
Pure Ulysses
local sequence shard
-> sequence/head A2A on cp_mesh
head shard + global sequence
-> local attention
-> head/sequence A2A
local sequence output
要求 ulysses_degree == cp_size,Q head 和参与通信的 KV head 满足整除约束。
Hybrid
CP mesh
-> Ulysses subgroup:sequence -> head A2A
-> Colossal subgroup:K/V AllGather
-> attention
-> inverse A2A
要求 cp_size % ulysses_degree == 0。每个子组必须使用稳定的 rank 顺序和已缓存 process group。
4.4 Async 和 load balance
Async strategy 是独立的双边界注入,不是同步 wrapper 上的布尔开关。它在最后一个安全
projection/RoPE/layout handoff 后提前发起通信,在 attention consumer 前等待:
q_proj/k_proj/v_proj (+ QK norm/RoPE)
-> async communication launch
-> reshape/mask/metadata
-> attention boundary wait
forward/backward 必须形成对称闭环:
forward: handoff post-hook -> launch -> consumer pre-hook wait -> attention
backward: attention autograd -> launch inverse collective
-> handoff/projection backward pre-hook wait -> projection grad GEMM
异步边界由模型专用 wrapper 暴露,不由框架根据模型类名猜测结构。当前 Qwen3-MoE
实现直接在 wrapper 内完成 handoff 与 collective 生命周期;通用 HF 模型若需异步,
必须新增并注册自己的 @inner_wrapper callable:
@inner_wrapper
def my_async_cp_wrapper(target_module, mesh, tp_mesh, cp_mesh, ep_mesh):
"""Replace target_module.forward and own launch/wait/backward contract."""
...
模型侧应在已有 forward 中保留原 attention 数学逻辑,仅在 projection/layout 完成后增加
明确的 launch/wait 边界,不复制整个 attention 实现:
q = project_norm_and_layout_q(hidden_states)
k = project_norm_and_layout_k(hidden_states)
q, k = apply_rotary_pos_emb(q, k, cos, sin)
q = self.cp_q_handoff(q) # wrapper 在 post-hook launch Q
k = self.cp_k_handoff(k) # wrapper 在 post-hook launch K
v = project_and_layout_v(hidden_states) # 与已 launch 的 Q/K 通信重叠
v = self.cp_v_handoff(v) # launch V
q, k, v = self.cp_attn_wait(q, k, v) # pre-hook wait/materialize
output = attention_interface(q, k, v, attention_mask)
通用 async wrapper 从 target_module.get_async_cp_handoffs() 取得真实 Module,注册 launch、wait
和 backward hook。preflight 必须检查四个 Module 唯一、属于当前 attention、调用次数匹配且输出
layout 符合 wrapper contract。没有实现协议的 attention 不能选择异步 wrapper。
各模式的 Phase 1 异步语义:
| 模式 | launch 点 | consumer 前等待内容 |
|---|---|---|
| Async Colossal | K/V 最后安全 handoff | K/V async AllGather |
| Async Pure Ulysses | Q/K/V 最后安全 handoff | Q/K/V sequence-to-head async A2A |
| Async Hybrid | Q/K/V 最后安全 handoff | Ulysses 子组 A2A;K/V 继续进入 Colossal 子组 gather 依赖链 |
Hybrid 的两段 collective 存在数据依赖。实现可以通过 platform 的异步依赖链在 A2A 完成后
继续发起 K/V AllGather,并在 attention 前统一 materialize;如果首版只能异步 A2A、同步
AllGather,必须在 wrapper 名称、能力矩阵和 profiler 结果中标记为 half_async,不能宣称为
完整 Async Hybrid。完整交付要求 forward 和 backward 的两段依赖、stream/event 顺序及
DTensor layout 恢复全部闭环。
handoff 不存在、跨越不安全或 wrapper 无法保证顺序时:
- 异步 wrapper 在 apply/preflight 阶段直接报错;
- 需要同步通信时,用户改为显式选择对应的同步 wrapper;不得由 runtime 自动回退。
HF forward(hidden_states) 的通用 SDPA consumer interception 通常只能在 Q/K/V 已全部生成后
看到张量,适合同步 wrapper,但不足以保证有效异步 overlap。异步 HF wrapper 必须依赖上述
handoff provider,而不是替换整个 attention forward。仅在 SDPA 调用点 launch 并立即 wait
属于伪异步,验收时按同步路径处理。
Head-Tail load balance 只允许 AllGather:batch 先按 zigzag 重排,attention 内交换 Q 片段并执行
双 attention,输出再还原原始 sequence 顺序。该能力必须有独立的 mask、padding 和 backward contract。
它只能通过显式的 load-balance wrapper 接入,不能通过额外的全局 YAML 开关打开。
Phase 1 至少交付同步 Colossal Head-Tail wrapper。async + load_balance 不是自动组合能力;只有
提供独立 wrapper、明确两个 Q 分支的 launch/wait 次序并通过 profiler 验证后才标记支持。
4.5 Attention 与模型适配
| 模型/类型 | 接入方式 |
|---|---|
HF SDPA hidden_states |
sdpa_hf primitive interception |
| QKV 风格 core attention | sdpa_qkv |
| FlexAttention | flex_hf / flex_qkv,block mask 使用 global KV 长度 |
| DeepSeek MLA | 在压缩 K/V expand_kv 前接入专用 wrapper,不能直接套普通 SDPA |
| GLM5 DSA | indexer、sparse attention、indexer loss 分别注册 placement/wrapper |
| Qwen3.5 linear/VL | 模型规格提供 sequence、MRoPE、multimodal mask 和 attention boundary |
所有 wrapper 都必须支持 production/validate 双模式,且未命中实际 attention primitive 时 fail-fast。
4.6 Loss、梯度与 optimizer
统一内部 loss contract:
local_loss_sum
local_valid_tokens
mean_global_loss 在 dp_cp_mesh 上聚合 token 和 loss;backward scale 与 FSDP SUM/AVG 语义一致,
logging loss 与可微 loss 分开生成。
要求:
- CP rank 的 valid token 只统计一次;
- optimizer step 前 replicated parameter gradient 在目标同步域一致;
- optimizer step 后 CP group 参数一致;
- TP/EP 副本不重复计数;
- gradient accumulation、activation checkpoint 和 mixed precision 不改变 CP=1 数值语义。
5. 代码改动边界
| 模块 | 主要改动 |
|---|---|
trainer/config.py |
AcceleratorConfig 保留 cp_size,读取 plan_overrides |
components/distributed/config.py |
校验 when: cp、wrapper/compute 声明和 mesh 约束;不提供通用 CP 算法配置 |
components/distributed/infrastructure.py |
构造 cp_mesh/dp_cp_mesh、Hybrid 子 mesh 和 wrapper 所需稳定通信组 |
cp_utils.py |
提供 Colossal/Ulysses/Hybrid 同步与异步 primitive、Head-Tail、batch/mask/padding 工具 |
sharding_config.py |
保存 CP layout 和显式 wrapper/compute intent |
sharding_planner.py |
合并 plan_overrides,校验 head/layout/backend,禁止隐式 CP 方法 |
sharding_applier.py |
解析并注入用户声明的 wrapper/compute,填充 wrapper Target 上下文并执行 fail-fast |
trainer/base.py |
forward 前 CP batch prepare,统一 global label/position contract |
loss_utils.py / FSDP2 |
DP+CP token/loss 和 replicated gradient 域闭环 |
| 模型实现 | 为无安全 handoff 的 HF attention 和 MLA/DSA/VL 等非标准 attention 增加专用 wrapper/handoff contract |
所有 cp_size>1 的模型都必须在 YAML 中为 attention 声明 plan_overrides。只有
cp_size=1 可以借助 when: cp 跳过该条目;不得因为模型是标准 HF attention 就省略
CP 方法声明。
6. 支持矩阵与整体迁移范围
6.1 支持矩阵
| 能力 | Phase 1 | Phase 2 | Phase 3 |
|---|---|---|---|
cp_size + 显式 plan_overrides 方法 |
必须 | 保持兼容 | 保持兼容 |
| 同步 Colossal/AllGather | 完整支持 | 性能优化 | 持续回归 |
| 同步 Pure Ulysses | 完整支持 | 性能优化 | 持续回归 |
| 同步 Hybrid | 完整支持 | 子组/拓扑优化 | 持续回归 |
| Async Colossal | 完整支持并证明 overlap | handoff 覆盖扩展 | 调度优化 |
| Async Pure Ulysses | 完整支持并证明 overlap | handoff 覆盖扩展 | 调度优化 |
| Async Hybrid | 完整支持;half-async 仅算中间里程碑 | collective 依赖链优化 | 调度优化 |
| Head-Tail load balance | 同步 Colossal 完整支持 | 独立 async 组合按 wrapper 扩展 | 性能优化 |
| SDPA/Flex | 标准 SDPA/QKV 必须;Flex wrapper 明确支持边界 | fused/varlen 扩展 | 持续回归 |
| GQA/MQA | Colossal 支持;Ulysses/Hybrid 按 head 约束 fail-fast | KV 策略扩展 | 完整优化 |
| TP+CP | 组件和 Trainer E2E 闭环 | 稳定支持 | 持续回归 |
| FSDP+CP | loss/gradient/optimizer 正确性闭环 | 稳定支持 | 持续回归 |
| DSA/VL/MLA | 不作为核心场景的统一 wrapper;已有专用路径需回归 | 按模型逐项扩展 | 持续扩展 |
6.2 本次迁移的整体工作面
| 类别 | 迁移内容 |
|---|---|
| 基础闭环 | cp_size、mesh/group、batch shard、显式 wrapper、offset mask、global loss、FSDP gradient domain |
| 同步算法 | Colossal、Pure Ulysses、Hybrid;包括子组构造、layout transform、autograd 和输出恢复 |
| 异步算法 | Colossal/Ulysses/Hybrid 的 handoff、launch/wait、反向逆通信、stream/event 和 profiler 验收 |
| 负载均衡 | Colossal Head-Tail 的输入重排、双 attention、mask、输出还原、backward 和 padding contract |
| 模型扩展 | 仅处理无法套用标准 SDPA 的 MLA、DSA、linear/VL attention;按模型增加 wrapper 和 batch contract |
以下能力不纳入本次 CP 迁移主线:inference cache、PP 生产组合、EP 生产组合、跨 CP size checkpoint
reshard,以及未注册的 fused kernel。它们只在后续组合并行/工程 Issue 中单独跟踪。
7. 验证设计
7.1 测试分层
| 层级 | 验证内容 |
|---|---|
| UT | plan_overrides、wrapper-local 参数、padding、batch field、offset mask、head/layout/handoff 约束 |
| Component | NPU/HCCL 上验证 Colossal/Ulysses/Hybrid 同步与异步 forward/backward、load balance、process group/handle 生命周期 |
| Trainer Level 0 | tiny Causal LM,CP=1/2/Hybrid factor,固定 seed,2~5 step parity |
| Trainer Level 1 | Llama/Qwen 为核心,NPU/HCCL,多 micro-batch/FSDP;专用模型路径按支持矩阵回归 |
| Performance | NPU profiler 检查 async launch/wait、stream 依赖和有效 overlap;禁止以 API 名称替代性能证据 |
7.2 Phase 1 验收项
- YAML 中每个 CP attention 的
when: cp、inner_target、inner_wrapper/
local_compute_fn和region_dispatch能到达 apply;缺失或冲突声明会在启动前报错; - global input、position、labels 在 CP rank 拼接后恢复 reference;
- Colossal AllGather backward 为 reduce-scatter 语义,且不泄漏 process group;
- Pure Ulysses 和 Hybrid 的 forward/backward 与 CP=1/reference parity,Hybrid 子组 rank 顺序稳定;
- local-Q/global-KV causal mask 与 CP=1 一致;
- Async Colossal/Ulysses/Hybrid 与对应同步路径数值一致,forward/backward 均无悬空 handle;
- NPU profiler 证明异步 collective 与有效计算区间重叠;launch 后立即 wait 不通过验收;
- Head-Tail load balance 的重排/还原、mask、padding、loss token 和 backward 全部一致;
- local logits 拼接、global loss、valid token 数与 CP=1 一致;
- 每个 trainable parameter gradient 与 CP=1 一致;
- optimizer step 后 replicated parameter 在 CP group 内一致;
- production/validate 输出和梯度一致;
- 不支持的模型、backend、layout、handoff、cache、packed/TND 和非法能力组合 fail-fast;
cp_size=1与未启用 CP 的原始 Trainer 回归一致。
7.3 Phase 2/3 扩展验收项
- fused/varlen backend 与 Phase 1 同步 reference parity;
- Async collective 依赖链在更多模型 handoff 和组合并行下保持正确 overlap;
- GQA/MQA 的扩展 KV 策略有明确精度、通信量和 head-layout 证据;
- MLA/DSA/VL 模型只在专用 wrapper/batch contract 完成后标记支持;
- 大规模场景验证通信组生命周期、显存峰值、吞吐和长稳可靠性。
8. 交付拆分
- CP-Declaration:
plan_overridesschema、when: cp、wrapper/compute 解析、拓扑摘要和非法组合校验; - CP-Colossal:AllGather primitive、batch contract、显式 SDPA/QKV wrapper、offset mask 和 Trainer parity;
- CP-Gradient:DP+CP loss、FSDP/HSDP gradient domain、TP+CP 基础 E2E;
- CP-Ulysses-Hybrid:Pure Ulysses/Hybrid 同步 strategy、子组、layout、autograd 和输出恢复;
- CP-Async:Colossal/Ulysses/Hybrid 的 handoff、launch/wait、反向通信和 profiler overlap 验收;
- CP-LoadBalance:Colossal Head-Tail 重排、双 attention、mask、padding 和 backward;
- CP-Model:无法复用标准 SDPA/QKV contract 的 MLA/DSA/VL 等模型专用 wrapper。
第 1~6 项共同构成 Phase 1 核心交付,不是按“先只交付 AllGather、其余以后再补”的可选顺序;
允许按 PR 拆分实现,但最终验收必须覆盖完整功能面。只有同时满足显式配置可达、组件测试、
Trainer parity、异步 profiler 证据和明确支持矩阵,才能将某项能力标记为新 Trainer 已支持;
任何“自动猜测 wrapper”“缺失时默认 AllGather”或“launch 后立即 wait”的实现都不算支持。
schema_version: 1
source: gitcode
gitcode_repo: mindspore/hyper-parallel
gitcode_issue: 324
source_url: https://gitcode.com/mindspore/hyper-parallel/issues/324
Contributor guide
No contributing guide indexed for this repository
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Start by reading the named entry points: TrainerConfig, MeshContext, PlanOverrideResolver, ShardingPlanner, ShardingApplier, CPStrategy, and cp_utils.py, plus the wrappers under hyper_parallel.distributed.context_parallel.wrappers. Trace how BaseTrainer/TextTrainer prepares batches and aggregates loss; done requires the RFC’s Phase 1 CP scenarios, contracts, parity checks, and profiler validation to be implemented and verified.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python
- Domain
- distributed-systems, machine-learning
- Issue type
- Feature
- Difficulty
- 5/5
- Estimated time
- Over a week
- Activity status
- Active
- Clarity
- Mostly clear
- Newbie friendliness
- 25/100