已开启
【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— Part 6 端到端使能模型 #144
changzherui创建于  5月15日
changzherui
changzherui成员
5月15日 创建

Part 6 — 首个模型迁移 models/qwen3_5_v2/

把现有 models/qwen3_5/ 重写为 Module 协议版,端到端打通"新 Config 路径"。models/qwen3_5/ 完全保留


1. 目标

  1. 用 M4 的 models/common/ 组件重写 Qwen3.5 dense 模型,验证 Module 协议在真实模型上可用。
  2. 实现 set_qwen3_5_sharding_config(),把 sharding 声明灌进 model config 树。
  3. 实现 parallelize_qwen3_5_v2():CP → model.parallelize(tp_mesh) → AC → FSDP。
  4. 数值与旧 qwen3_5 同种子 8-card 训练 loss 误差 ≤ 1e-4

2. 任务边界

新增目录 hyper_parallel/models/qwen3_5_v2/

新文件 旧对应文件(仅参考) 内容
model.py models/qwen3_5/model.py Qwen3_5Model(Decoder) + Qwen3_5TransformerBlock(TransformerBlock),全部基于 M4 models/common/ 组件;Qwen3_5Model.Config(BaseModel.Config) 字段对齐 Qwen3_5Configmodel.py:82
sharding.py set_qwen3_5_sharding_config(config, *, loss_parallel, enable_sp):调 M4 助手;处理 attn_output_gateq_proj 列数翻倍的特殊情况
parallelize.py models/qwen3_5/parallelize.py:102 parallelize_qwen3_5_v2(model, parallel_dims, training, parallelism, activation_checkpoint, ...):CP → model.parallelize(tp_mesh) → AC → FSDP。apply_ac / apply_fsdp 复用旧版算法(_apply_ac / _apply_fsdp
state_dict.py models/qwen3_5/state_dict.py:Qwen3_5StateDictAdapter 继承 BaseStateDictAdapter,复用旧 HF↔hyper 键名映射
__init__.py models/qwen3_5/__init__.py:75 register_spec register_spec("qwen3_5_v2", ModelSpec(name="qwen3_5_v2", model=None, parallelize_fn=parallelize_qwen3_5_v2, state_dict_adapter=Qwen3_5_v2StateDictAdapter)) + 工厂函数 qwen3_5_v2_spec_factory(flavor: str)

3. 与旧 spec 的关键差别

维度 qwen3_5(旧) qwen3_5_v2(新)
模型类 Qwen3_5ForCausalLM(nn.Module)model.py:268 Qwen3_5Model(Decoder)(继承 BaseModel
Config Qwen3_5Config(@dataclass)model.py:82 Qwen3_5Model.Config(BaseModel.Config)
并行声明 类属性 _tp_plan = {"*.q_proj": "colwise", ...}model.py:284 set_qwen3_5_sharding_config(cfg, ...)ShardingConfig 到 cfg 树
并行入口 parallelize_qwen3_5(model, mesh, cfg)parallelize.py:102)仅 AC + FSDP parallelize_qwen3_5_v2(model, parallel_dims, training, ...)model.parallelize(tp_mesh) 自递归
TP 支持 tp>1 raise NotImplementedError(仅 full_attention 层支持,parallelize.py:105 相同限制(M6 不解决 linear_attn 的 TP 问题,留 M9)
注册 register_spec("qwen3_5", ...)__init__.py:75 register_spec("qwen3_5_v2", ...)(同一 registry)
loss 计算 内置在 forward 末尾(model.py:361 由外部 CrossEntropyLoss.Config().build() 计算(M3)

4. 核心实现

4.1 Qwen3_5Model.Config 字段对齐旧 Qwen3_5Config

直接对照 models/qwen3_5/model.py:82 的所有字段(vocab_size / hidden_size / intermediate_size / num_hidden_layers / num_attention_heads / num_key_value_heads / head_dim / max_position_embeddings / rms_norm_eps / attention_bias / tie_word_embeddings / attn_output_gate / rope_theta / partial_rotary_factor / mrope_section / full_attention_interval / linear_num_value_heads / ...)。__post_init__ 计算 layer_types(与 model.py:126 完全一致)。

4.2 parallelize_qwen3_5_v2 流程

def parallelize_qwen3_5_v2(
    model, parallel_dims, *,
    training, parallelism, activation_checkpoint, compile, dump_folder,
):
    # 0. 校验:TP 暂不支持 linear_attn 层
    if parallelism.tensor_parallel_degree > 1:
        raise NotImplementedError(
            "Qwen3_5 v2 TP for linear-attention layers is not yet implemented. "
            "Set parallelism.tensor_parallel_degree=1."
        )
    if parallelism.expert_parallel_degree > 1:
        raise NotImplementedError("Qwen3_5 v2 dense has no experts.")

    # 1. CP: 包装 inner attention(仅当 cp > 1 时)
    if parallel_dims.cp_enabled:
        apply_cp_to_forward(model, parallel_dims.world_mesh["cp"])

    # 2. TP: 声明式递归(仅当 tp > 1)
    if parallel_dims.tp_enabled:
        model.parallelize(parallel_dims.world_mesh["tp"])

    # 3. AC(复用旧 _apply_ac 算法)
    _apply_ac(model, activation_checkpoint)

    # 4. FSDP(复用旧 _apply_fsdp 算法)
    _apply_fsdp(model, parallel_dims.world_mesh, training)

    return model

4.3 set_qwen3_5_sharding_config 处理 attn_output_gate

def set_qwen3_5_sharding_config(cfg: Qwen3_5Model.Config, *, loss_parallel, enable_sp):
    set_decoder_sharding_config(
        cfg, loss_parallel=loss_parallel, enable_sp=enable_sp,
    )
    # Qwen3.5 特殊:attn_output_gate=True 时 q_proj 列数翻倍
    # 仍然按 Shard(0) 切,head 维度沿 TP 切分自然成立
    if cfg.attn_output_gate:
        for block_cfg in cfg.layers:
            if block_cfg.layer_type == "full_attention":
                # q_proj.out_features = num_heads * head_dim * 2
                # colwise_config() 已经返回 Shard(0),与 gate 拆分兼容
                pass    # 不需要额外处理,注释说明语义

4.4 update_from_config

class Qwen3_5Model(Decoder):
    def update_from_config(self, trainer_config):
        # 把 trainer_config 的运行时参数同步到 model config 树
        self.rope.max_seq_len = trainer_config.training.seq_len
        # 调声明式 sharding 助手
        set_qwen3_5_sharding_config(
            self.config,
            loss_parallel=trainer_config.parallelism.loss_parallel,
            enable_sp=trainer_config.parallelism.enable_sp,
        )

5. 与 torchtitan 接口差异说明

# 差异点 原因
1 模型 forward 返回 logits(不算 loss) loss 由 M3 CrossEntropyLoss 组件计算;与 torchtitan 一致
2 TP 限制:linear_attn 层暂不支持(与旧版一致) GatedDeltaNet 算子复杂,TP 切分需 M9 local_map
3 apply_ac / apply_fsdp 直接搬运旧 parallelize.py:34-100 算法 数值兼容、不重复造轮子
4 attn_output_gate=Trueq_proj 列数翻倍仍走 Shard(0) 与 head 维度自然兼容,不需要 Shard(1) 特殊路径

6. 开发步骤

内容 工期
1 model.py:重写 Qwen3_5TextModel / Qwen3_5Decoder / Qwen3_5ForCausalLM 为 Module 协议;forward 返回 logits 1 d
2 sharding.py:复用 M4 助手;处理 attn_output_gate / partial_rotary_factor 特殊情况 0.5 d
3 parallelize.py:搬运 _apply_ac + _apply_fsdp;插入 model.parallelize(tp_mesh) 调用 1 d
4 state_dict.py:继承 BaseStateDictAdapter,逻辑搬 models/qwen3_5/state_dict.py 0.5 d
5 __init__.py + qwen3_5_v2_spec_factory("0.8B" / "4B") 0.5 d

7. 验证标准

新建 tests/torch/st/qwen3_5_v2/

测试 断言要点
test_unit_forward.py 1-card:构造 Qwen3_5Model.Config(num_hidden_layers=4, hidden_size=256, ...)init_states + forward,与旧 Qwen3_5ForCausalLM(Qwen3_5Config(num_hidden_layers=4, hidden_size=256)) 同种子 logits bit-exact
test_st_8card_loss.py 8-card:TP=2 / FSDP=4 训练 10 步,对比旧 qwen3_5 同 cfg 同种子的 loss / grad_norm,误差容忍 1e-4
test_st_1card_loss.py 1-card:tp_mesh.size()==1model.parallelize 不坏;loss 与旧路径 bit-exact
test_state_dict_roundtrip.py 加载同一 HF 权重到旧 / 新模型,state_dict 张量逐元素一致

通过门槛

  • 4 个 ST 全绿。
  • qwen3_5 训练 0 受影响

8. 工期 & 依赖

工期 3.5 天
依赖 M1 + M4 + M5
软依赖 M2 / M3(可以先用 stub Config 跑通;最终验收需 M2 / M3 完成)
下游 M7 / M8
likedislike
changzheruichangzherui成员
5月15日 修改标题为 “hyper-parallel 支持torchtitan 风格 Module/Config —— M6 端到端使能模型”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M6”
changzheruichangzherui成员
5月15日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M6 端到端使能模型”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M6 端到端使能模型”
changzheruichangzherui成员
5月25日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— Part 6 端到端使能模型”,原标题为“【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M6 端到端使能模型”
changzheruichangzherui成员
5月25日 修改了issue 的描述