mindspore-ai / mindspore-ai/hyper-parallel

【RFC】新 Trainer Context Parallel 能力迁移与统一接入

Open
#593 0 comments 0 reactions 0 assignees View on GitHub

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_replicatedp_shardcptp;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_sizecp_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 必须闭环

具体目标:

  1. YAML 只声明 cp_size 和每个 attention 边界使用的 CP 方法;方法通过
    plan_overrides.inner_wrapper(或自定义 local_compute_fn)注入。
  2. 同步、异步和 load-balance 不是全局开关,而是不同的显式 wrapper。Hybrid 的
    ulysses_degree 等方法私有数据可以写在 inner_wrapper Target 内;异步 handoff 由
    模型专用 wrapper 在代码中显式处理,当前实现不提供通用 AsyncCPHandoffProvider 注册协议。
    不得重新引入 accelerator.cp.algorithm/schedule/ulysses_degree
  3. CP 方法的 mesh、layout、head、handoff 和 batch contract 在 apply/preflight 阶段校验;
    不提供算法默认值、隐式 wrapper 选择或 runtime fallback。
  4. 提供 Colossal AllGather、Ulysses A2A、Hybrid 子组通信、异步 launch/wait 和
    Head-Tail 重排所需的可微 primitive;底层 collective 统一复用 platform
  5. 标准 HF SDPA/QKV attention 至少提供同步和异步接入路径;无法安全识别 handoff 的
    HF 模型使用模型专用 wrapper,未匹配的 attention 明确报错。
  6. 固化 global position、causal mask、labels、padding、packed metadata、local loss、
    backward 和 optimizer-step contract。
  7. 对上述七个交付场景完成 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_sizeplan_overrides,不保存通用 CP 算法配置
MeshContext 构造 cp_meshdp_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.algorithmschedule
ulysses_degreeload_balanceasync_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 algorithmschedule
字段覆盖 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_hfsdpa_hf_ulysses 以及 sdpa_qkvsdpa_qkv_ulysses 是当前实现中的内置注册名;
FlexAttention 对应 flex_hfflex_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_targetinner_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_wrapperlocal_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_lossdp_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 验收项
  1. YAML 中每个 CP attention 的 when: cpinner_targetinner_wrapper/
    local_compute_fnregion_dispatch 能到达 apply;缺失或冲突声明会在启动前报错;
  2. global input、position、labels 在 CP rank 拼接后恢复 reference;
  3. Colossal AllGather backward 为 reduce-scatter 语义,且不泄漏 process group;
  4. Pure Ulysses 和 Hybrid 的 forward/backward 与 CP=1/reference parity,Hybrid 子组 rank 顺序稳定;
  5. local-Q/global-KV causal mask 与 CP=1 一致;
  6. Async Colossal/Ulysses/Hybrid 与对应同步路径数值一致,forward/backward 均无悬空 handle;
  7. NPU profiler 证明异步 collective 与有效计算区间重叠;launch 后立即 wait 不通过验收;
  8. Head-Tail load balance 的重排/还原、mask、padding、loss token 和 backward 全部一致;
  9. local logits 拼接、global loss、valid token 数与 CP=1 一致;
  10. 每个 trainable parameter gradient 与 CP=1 一致;
  11. optimizer step 后 replicated parameter 在 CP group 内一致;
  12. production/validate 输出和梯度一致;
  13. 不支持的模型、backend、layout、handoff、cache、packed/TND 和非法能力组合 fail-fast;
  14. 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. 交付拆分

  1. CP-Declarationplan_overrides schema、when: cp、wrapper/compute 解析、拓扑摘要和非法组合校验;
  2. CP-Colossal:AllGather primitive、batch contract、显式 SDPA/QKV wrapper、offset mask 和 Trainer parity;
  3. CP-Gradient:DP+CP loss、FSDP/HSDP gradient domain、TP+CP 基础 E2E;
  4. CP-Ulysses-Hybrid:Pure Ulysses/Hybrid 同步 strategy、子组、layout、autograd 和输出恢复;
  5. CP-Async:Colossal/Ulysses/Hybrid 的 handoff、launch/wait、反向通信和 profiler overlap 验收;
  6. CP-LoadBalance:Colossal Head-Tail 重排、双 attention、mask、padding 和 backward;
  7. 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

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. 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

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.