当前源码实际目录名是 hyper_parallel.auto_models,不是 auto_module。7 个场景共用以下结构,只替换 cp_size 和 inner_wrapper。
model:
attn_implementation: sdpa
force_hf: true
accelerator:
cp_size: 2
sequence_parallel: true
plan_overrides:
- match: "*.self_attn"
when: cp
region_dispatch: false
inner_target: self
inner_wrapper:
_target_: <下面对应的 wrapper>
- 同步 Colossal
inner_wrapper:
_target_: hyper_parallel.auto_models.components.distributed.cp_wrappers.sdpa_hf_cp_wrapper
- 同步 Pure Ulysses
inner_wrapper:
_target_: hyper_parallel.auto_models.components.distributed.cp_wrappers.sdpa_hf_ulysses_cp_wrapper
要求 attention head 数能被 cp_size 整除。
- 同步 Hybrid
accelerator:
cp_size: 4
plan_overrides:
- match: "*.self_attn"
when: cp
region_dispatch: false
inner_target: self
inner_wrapper:
_target_: hyper_parallel.auto_models.components.distributed.cp_wrappers.sdpa_hf_hybrid_cp_wrapper
ulysses_degree: 2
约束:
1 < ulysses_degree < cp_size
cp_size % ulysses_degree == 0
- 异步 Colossal
inner_wrapper:
_target_: hyper_parallel.auto_models.components.distributed.cp_wrappers.qwen3_moe_async_colossal_cp_wrapper
- 异步 Pure Ulysses
inner_wrapper:
_target_: hyper_parallel.auto_models.components.distributed.cp_wrappers.qwen3_moe_async_ulysses_cp_wrapper
Q head 和 KV head 都必须能被 cp_size 整除。
- 异步 Hybrid
accelerator:
cp_size: 4
plan_overrides:
- match: "*.self_attn"
when: cp
region_dispatch: false
inner_target: self
inner_wrapper:
_target_: hyper_parallel.auto_models.components.distributed.cp_wrappers.qwen3_moe_async_hybrid_cp_wrapper
ulysses_degree: 2
- 同步 Colossal Head-Tail 负载均衡
inner_wrapper:
_target_: hyper_parallel.auto_models.components.distributed.cp_wrappers.sdpa_hf_load_balance_cp_wrapper
要求本地序列长度为偶数,也就是全局序列长度最好满足:
global_seq_len % (2 * cp_size) == 0
需要注意:
- 前四个同步 HF wrapper 通过拦截
F.scaled_dot_product_attention工作,不能同时把 attention 替换成不调用 SDPA 的 fused forward。 - 三个异步 wrapper 只支持
Qwen3MoeAttention,会直接替换整个 attention forward,并使用 NPU fused attention。 match: "*.self_attn"和inner_target: self对 Qwen3/Qwen3-MoE 当前模型结构是正确的。- Hybrid 才需要显式配置
ulysses_degree;Pure Ulysses 默认整个 CP group 都是 Ulysses group。
实现位置见 cp_wrappers.py。


已完成复现,仓库代码没有修改,当前分支仍为 f66ab371。
实验配置
- 网络:
/tmp/tiny_qwen3_head_dim64hidden_size=512num_hidden_layers=4num_attention_heads=16head_dim=64vocab_size=256seq_length=16
- 数据:HP 自带
build_indexed_text_dataset,mock_data=true - 后端:Torch + HCCL,8 卡
- 优化器:AdamW,
lr=1e-5 GBS=8,micro_batch_size=1- CP 算法:新 Trainer 流程的
sdpa_hf_ulysses - FSDP 混合精度:
param_dtype: bfloat16
reduce_dtype: float32
output_dtype: bfloat16
拓扑为:
CP=1: DP=8,1 个 micro-step
CP=2: DP=4,2 个 micro-step
CP=2 时,每个 CP rank 只处理序列长度 16 / 2 = 8 的局部片段,但一个完整 Trainer step 仍然包含 2 个 micro-step。
复现结果
CP=1 与 CP=2 均使用相同初始权重、随机种子和 Indexed 数据。
| 指标 | CP=1 | CP=2 | 绝对差 |
|---|---|---|---|
| 第 1 步 loss | 5.78125 | 5.71875 | 0.0625 |
| 第 1 步 grad norm | 31.1074237823 | 31.1070785522 | 3.45e-4 |
20 步整体对比:
loss max abs = 6.25e-2
loss mean abs = 1.9140625e-2
grad norm max abs = 9.963989e-3
grad norm mean abs = 3.350353e-3
gradient max abs = 9.277344e-3
gradient mean abs = 6.820986e-5
gradient rel L2 max= 1.5206347e-2
gradient min cosine= 0.9998844
最差梯度出现在:
model.embed_tokens.weight:
step 6,max abs = 9.277344e-3
model.layers.3.self_attn.k_norm.weight:
step 18,relative L2 = 1.5206e-2
cosine = 0.9998844
完整结果保存在:
定位结论
CP 的数据切分和 attention 结果没有发现错误。
CPBatchSharder 会沿序列维度切分 input_ids 和 labels:batch_parallel.py。在第 1 步初始权重下:
- CP=2 重建后的全局 logits 与 CP=1 逐元素一致,
max abs = 0 - labels 完全一致
- CP=2 的局部
outputs.loss与 CP=1 不同是正常现象,因为 CP=2 每个 rank 只在自己的半段序列上计算 local loss
真正的问题出现在 loss 的 dtype:
-
fsdp2.py将配置直接转换为 FSDP2 的MixedPrecisionPolicy:fsdp2.py -
Transformers 的
ForCausalLMLoss会先将 logits 转成 FP32 计算交叉熵。 -
但是 FSDP2 设置了
output_dtype=bfloat16后,PyTorch FSDP 的 post-forward 会递归转换所有浮点输出,包括 HuggingFaceModelOutput.loss。因此原本计算好的 FP32 loss 在返回 Trainer 前已经被转换为 BF16。 -
ModelOutputLoss只是直接取出model_output.loss,没有恢复 FP32:model_output.py -
随后
mean_global_loss在 BF16 loss 上执行 token weighting、除法和全局聚合,并且global_mean也通过cur_loss.new_tensor(...)继承 BF16:loss_utils.py -
CP=1 每个 step 基本只有一份完整序列 loss;CP=2 则有多个 CP-local loss,并且当前配置下还会有 2 个 micro-step。每个局部 loss 先被 BF16 舍入,再参与全局加权和 Trainer 累加,所以 CP=1 与 CP=2 的舍入路径不同。
reduce_dtype=float32 已经生效,捕获到的梯度 dtype 均为 torch.float32,但它只能保证梯度通信/累加的 reduce dtype,无法恢复 loss 在 forward 输出阶段已经丢失的 FP32 精度。
隔离实验
保持:
param_dtype: bfloat16
reduce_dtype: float32
仅改为:
output_dtype: null
单步结果:
| 指标 | CP=1 | CP=2 | 绝对差 |
|---|---|---|---|
| loss | 5.751282691955566 | 5.7512829303741455 | 2.38e-7 |
| grad norm | 31.1074237823 | 31.1070785522 | 3.45e-4 |
此时:
初始全局 logits max abs = 0
labels 完全一致
gradient relative L2 max = 2.119e-3
gradient min cosine = 0.9999915
这说明当前明显的 loss 差异和后续 20 步梯度/参数轨迹差异,主要由 output_dtype=bfloat16 对 loss 的递归转换引起,而不是 CP wrapper 或 Indexed 数据错位。
当前建议的临时规避方式是:
fsdp_config:
mix_precision:
param_dtype: bfloat16
reduce_dtype: float32
output_dtype: null
如果必须保留 output_dtype=bfloat16,后续需要在 FSDP root 的输出转换路径中跳过 loss 字段,或者让 loss 使用独立的 FP32 输出路径。仅在 mean_global_loss 入口再执行 .float() 已经无法恢复此前丢失的精度。


HyperModels 新 Trainer Context Parallel 接入设计方案
0. 基本信息与术语
TextTrainer/BaseTrainer/HyperAutoModeldp_replicate、dp_shard、cp、tp;EP 为派生逻辑分组1. 目标与范围
1.1 目标
新 Trainer 需要提供从 YAML 到训练 step 的完整 CP 能力:
CP 的标准使用也必须通过
plan_overrides显式声明 attention 方法。Planner 可以识别attention boundary 和默认 placement,但不得替用户选择 CP 通信方法;Applier 只应用
用户声明的 wrapper/compute。这样可以避免模型结构被误判后静默使用错误的通信语义。
1.2 Phase 1 目标
交付范围如下:
cp_mesh上 AllGatherulysses_degree == cp_size,Q/K/V head 满足整除约束1 < ulysses_degree < cp_size且cp_size % ulysses_degree == 0具体目标:
cp_size和每个 attention 边界使用的 CP 方法;方法通过plan_overrides.inner_wrapper(或自定义local_compute_fn)注入。ulysses_degree等方法私有数据可以写在inner_wrapperTarget 内;异步 handoff 由模型专用 wrapper 在代码中显式处理,当前实现不提供通用
AsyncCPHandoffProvider注册协议。不得重新引入
accelerator.cp.algorithm/schedule/ulysses_degree。不提供算法默认值、隐式 wrapper 选择或 runtime fallback。
Head-Tail 重排所需的可微 primitive;底层 collective 统一复用
platform。HF 模型使用模型专用 wrapper,未匹配的 attention 明确报错。
backward 和 optimizer-step contract。
异步场景还必须通过 NPU profiler 证明存在有效通信计算重叠。
Phase 1 不要求任意能力自由组合。支持单元是一个明确命名或 Target 声明的 wrapper;例如
async + load_balance只有在存在独立 wrapper 和完整 contract 时才算支持,不能把两个开关拼接后自动生成行为。Ulysses/Hybrid 与 Head-Tail load balance 的组合默认非法并 fail-fast。
1.3 非目标
以下能力不阻塞 Phase 1:
2. 总体方案
2.1 分层架构
2.2 职责边界
TrainerConfigcp_size和plan_overrides,不保存通用 CP 算法配置MeshContextcp_mesh、dp_cp_mesh,提供 rank/group/topologyPlanOverrideResolverwhen: cp合并用户声明,缺失或冲突时 fail-fastShardingPlannerShardingApplierinner_wrapper/local_compute_fn,处理双模式CPStrategycp_utils.pyBaseTrainer/TextTrainermean_global_lossCPStrategy是 wrapper 内部运行时接口,不作为用户 YAML 的通用配置对象:wrapper 负责适配模型签名并明确选择 strategy,strategy 负责通信算法。每个包含 CP collective 的注入都必须
声明
region_dispatch: false,防止声明式边界再次发起冲突通信。2.3 一次训练 step
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.whencp_size>1时必须命中;cp_size=1可跳过,但应记录 INFOinner_targetinner_wrapper_target_callable,唯一决定该边界的 CP 方法local_compute_fnregion_dispatch: falseregion_dispatchfalse,不得省略同一份 YAML 可以通过
when: cp在 CP=1 调试时跳过 CP 注入,但这不等于存在默认CP 方法。CP>1 时每一个实际 attention boundary 都必须有且只有一个明确的方法声明。
内置 wrapper 的选择示例:
inner_wrapperforward(hidden_states)+ SDPAsdpa_hfforward(q, k, v)+ SDPAsdpa_qkvflex_hfflex_qkv_target_callable_target_callableulysses_degree决定 Ulysses/Colossal 子组_target_callable_target_callable_target_callableUlysses/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_wrapperPure 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_wrappercallable,并在代码中明确 handoff、handle 生命周期、前向等待点和 backward 顺序;不存在可自动套用的
AsyncCPHandoffProvider注册接口。Colossal Head-Tail load balance
inner_wrapper: _target_: hyper_parallel.distributed.context_parallel.wrappers.sdpa_hf_load_balance_cp_wrapperHead-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, )配置传播链:
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 子组;accelerator.cp或其他隐式全局配置;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_wrapperplan_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, )其余场景需要提供同一层级的能力:
(co, ds)子 mesh、A2A + K/V AllGather layout transform;实现约束:所有 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:
CP 相关子组:
cp_meshdp_cp_meshEP 不进入主 mesh 乘积;EP mesh 在 apply 阶段从 dense rank domain 派生。PP 首期固定为 1。
4.2 Batch contract
CP 切分前 batch 必须仍是 global sequence。处理顺序:
input_idslabelsshift_labelsposition_idsattention_maskinputs_embedsseq_lens/seq_lens_paddedpast_key_values/use_cachepadding policy:
all_gatherS对齐到cp_size的倍数ulyssesS对齐到 CP degree,且 head 维满足 A2A 整除hybridload_balance2*cp_size或 zigzag 对齐,不影响普通 AllGather4.3 三种同步 strategy
Colossal / AllGather
不要求 Q head 被 CP 整除,适合 GQA/MQA;causal mask 必须使用 local Q 的 global offset。
Pure Ulysses
要求
ulysses_degree == cp_size,Q head 和参与通信的 KV head 满足整除约束。Hybrid
要求
cp_size % ulysses_degree == 0。每个子组必须使用稳定的 rank 顺序和已缓存 process group。4.4 Async 和 load balance
Async strategy 是独立的双边界注入,不是同步 wrapper 上的布尔开关。它在最后一个安全
projection/RoPE/layout handoff 后提前发起通信,在 attention consumer 前等待:
forward/backward 必须形成对称闭环:
异步边界由模型专用 wrapper 暴露,不由框架根据模型类名猜测结构。当前 Qwen3-MoE
实现直接在 wrapper 内完成 handoff 与 collective 生命周期;通用 HF 模型若需异步,
必须新增并注册自己的
@inner_wrappercallable:@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 异步语义:
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 无法保证顺序时:
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 与模型适配
hidden_statessdpa_hfprimitive interceptionsdpa_qkvflex_hf/flex_qkv,block mask 使用 global KV 长度expand_kv前接入专用 wrapper,不能直接套普通 SDPA所有 wrapper 都必须支持 production/validate 双模式,且未命中实际 attention primitive 时 fail-fast。
4.6 Loss、梯度与 optimizer
统一内部 loss contract:
mean_global_loss在dp_cp_mesh上聚合 token 和 loss;backward scale 与 FSDP SUM/AVG 语义一致,logging loss 与可微 loss 分开生成。
要求:
5. 代码改动边界
trainer/config.pyAcceleratorConfig保留cp_size,读取plan_overridescomponents/distributed/config.pywhen: cp、wrapper/compute 声明和 mesh 约束;不提供通用 CP 算法配置components/distributed/infrastructure.pycp_mesh/dp_cp_mesh、Hybrid 子 mesh 和 wrapper 所需稳定通信组cp_utils.pysharding_config.pysharding_planner.pyplan_overrides,校验 head/layout/backend,禁止隐式 CP 方法sharding_applier.pytrainer/base.pyloss_utils.py/ FSDP2所有
cp_size>1的模型都必须在 YAML 中为 attention 声明plan_overrides。只有cp_size=1可以借助when: cp跳过该条目;不得因为模型是标准 HF attention 就省略CP 方法声明。
6. 支持矩阵与整体迁移范围
6.1 支持矩阵
cp_size+ 显式plan_overrides方法6.2 本次迁移的整体工作面
cp_size、mesh/group、batch shard、显式 wrapper、offset mask、global loss、FSDP gradient domain以下能力不纳入本次 CP 迁移主线:inference cache、PP 生产组合、EP 生产组合、跨 CP size checkpoint
reshard,以及未注册的 fused kernel。它们只在后续组合并行/工程 Issue 中单独跟踪。
7. 验证设计
7.1 测试分层
plan_overrides、wrapper-local 参数、padding、batch field、offset mask、head/layout/handoff 约束7.2 Phase 1 验收项
when: cp、inner_target、inner_wrapper/local_compute_fn和region_dispatch能到达 apply;缺失或冲突声明会在启动前报错;cp_size=1与未启用 CP 的原始 Trainer 回归一致。7.3 Phase 2/3 扩展验收项
8. 交付拆分
plan_overridesschema、when: cp、wrapper/compute 解析、拓扑摘要和非法组合校验;第 1~6 项共同构成 Phase 1 核心交付,不是按“先只交付 AllGather、其余以后再补”的可选顺序;
允许按 PR 拆分实现,但最终验收必须覆盖完整功能面。只有同时满足显式配置可达、组件测试、
Trainer parity、异步 profiler 证据和明确支持矩阵,才能将某项能力标记为新 Trainer 已支持;
任何“自动猜测 wrapper”“缺失时默认 AllGather”或“launch 后立即 wait”的实现都不算支持。