已开启
【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— Part 6 端到端使能模型 #144
changzherui创建于 5月15日
5月15日 修改标题为 “hyper-parallel 支持torchtitan 风格 Module/Config —— M6 端到端使能模型”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M6”
5月15日 修改标题为 “hyper-parallel 支持torchtitan 风格 Module/Config —— M6 端到端使能模型”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M6”
5月15日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M6 端到端使能模型”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M6 端到端使能模型”
5月15日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M6 端到端使能模型”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M6 端到端使能模型”
5月25日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— Part 6 端到端使能模型”,原标题为“【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M6 端到端使能模型”
5月25日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— Part 6 端到端使能模型”,原标题为“【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M6 端到端使能模型”
5月25日 修改了issue 的描述
Part 6 — 首个模型迁移
models/qwen3_5_v2/1. 目标
models/common/组件重写 Qwen3.5 dense 模型,验证 Module 协议在真实模型上可用。set_qwen3_5_sharding_config(),把 sharding 声明灌进 model config 树。parallelize_qwen3_5_v2():CP →model.parallelize(tp_mesh)→ AC → FSDP。qwen3_5同种子 8-card 训练 loss 误差 ≤ 1e-4。2. 任务边界
新增目录
hyper_parallel/models/qwen3_5_v2/:model.pymodels/qwen3_5/model.pyQwen3_5Model(Decoder)+Qwen3_5TransformerBlock(TransformerBlock),全部基于 M4models/common/组件;Qwen3_5Model.Config(BaseModel.Config)字段对齐Qwen3_5Config(model.py:82)sharding.pyset_qwen3_5_sharding_config(config, *, loss_parallel, enable_sp):调 M4 助手;处理attn_output_gate时q_proj列数翻倍的特殊情况parallelize.pymodels/qwen3_5/parallelize.py:102parallelize_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.pymodels/qwen3_5/state_dict.py:Qwen3_5StateDictAdapterBaseStateDictAdapter,复用旧 HF↔hyper 键名映射__init__.pymodels/qwen3_5/__init__.py:75 register_specregister_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)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 + FSDPparallelize_qwen3_5_v2(model, parallel_dims, training, ...)含model.parallelize(tp_mesh)自递归tp>1raise NotImplementedError(仅 full_attention 层支持,parallelize.py:105)register_spec("qwen3_5", ...)(__init__.py:75)register_spec("qwen3_5_v2", ...)(同一 registry)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 model4.3
set_qwen3_5_sharding_config处理attn_output_gatedef 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_configclass 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 接口差异说明
CrossEntropyLoss组件计算;与 torchtitan 一致GatedDeltaNet算子复杂,TP 切分需 M9local_mapapply_ac / apply_fsdp直接搬运旧parallelize.py:34-100算法attn_output_gate=True时q_proj列数翻倍仍走Shard(0)Shard(1)特殊路径6. 开发步骤
model.py:重写Qwen3_5TextModel / Qwen3_5Decoder / Qwen3_5ForCausalLM为 Module 协议;forward返回 logitssharding.py:复用 M4 助手;处理attn_output_gate / partial_rotary_factor特殊情况parallelize.py:搬运_apply_ac+_apply_fsdp;插入model.parallelize(tp_mesh)调用state_dict.py:继承BaseStateDictAdapter,逻辑搬models/qwen3_5/state_dict.py__init__.py+qwen3_5_v2_spec_factory("0.8B" / "4B")7. 验证标准
新建
tests/torch/st/qwen3_5_v2/:test_unit_forward.pyQwen3_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-exacttest_st_8card_loss.pytest_st_1card_loss.pytp_mesh.size()==1时model.parallelize不坏;loss 与旧路径 bit-exacttest_state_dict_roundtrip.py通过门槛:
qwen3_5训练 0 受影响。8. 工期 & 依赖