diff --git a/torchtitan/models/qwen3_5/state_dict_adapter.py b/torchtitan/models/qwen3_5/state_dict_adapter.py index c8dad46128..937c2ca9fd 100644 --- a/torchtitan/models/qwen3_5/state_dict_adapter.py +++ b/torchtitan/models/qwen3_5/state_dict_adapter.py @@ -9,12 +9,11 @@ Converts between HuggingFace Qwen3.5 checkpoint format and torchtitan format. -MoE expert weights require two transformations: -- **Transpose**: HF and TT use transposed layouts for grouped 3D expert weights. - E.g. HF down_proj [E, hidden, dim] <-> TT w2 [E, dim, hidden]. -- **Fuse/split gate_up_proj**: HF fuses gate_proj and up_proj into a single - gate_up_proj [E, dim, 2*hidden_dim]. TT stores them separately as +MoE expert weights use the same grouped 3D layouts in HF and TT. HF fuses +gate_proj and up_proj into a single gate_up_proj [E, 2*hidden_dim, dim], while +TT stores them separately as w1 [E, hidden_dim, dim] and w3 [E, hidden_dim, dim]. +HF down_proj and TT w2 both use [E, dim, hidden_dim]. Other notable conversions: - Conv3d patch embedding (HF) <-> Linear (TT) via weight reshape @@ -143,7 +142,7 @@ def to_hf(self, state_dict: dict[str, Any]) -> dict[str, Any]: hf_key = ( f"model.language_model.layers.{layer_num}.mlp.experts.down_proj" ) - hf_state_dict[hf_key] = value.transpose(-2, -1) + hf_state_dict[hf_key] = value continue if tt_abstract_key not in to_hf_map: @@ -214,11 +213,11 @@ def to_hf(self, state_dict: dict[str, Any]) -> dict[str, Any]: # Fuse MoE w1 (gate) + w3 (up) → gate_up_proj for layer_num in moe_w1_by_layer: - w1 = moe_w1_by_layer[layer_num].transpose(-2, -1) - w3 = moe_w3_by_layer[layer_num].transpose(-2, -1) + w1 = moe_w1_by_layer[layer_num] + w3 = moe_w3_by_layer[layer_num] hf_state_dict[ f"model.language_model.layers.{layer_num}.mlp.experts.gate_up_proj" - ] = torch.cat([w1, w3], dim=-1) + ] = torch.cat([w1, w3], dim=-2) # Fuse vision wq/wk/wv → qkv for layer_num, parts in vision_qkv_by_layer.items(): @@ -269,28 +268,28 @@ def from_hf(self, hf_state_dict: dict[str, Any]) -> dict[str, Any]: # pyrefly: ignore [missing-attribute] idx = re.search(r"\d+", hf_key).group(0) - # MoE gate_up_proj → split into w1 + w3 and transpose + # MoE gate_up_proj → split into w1 + w3 if ( hf_abstract_key == "model.language_model.layers.{}.mlp.experts.gate_up_proj" ): - w1_hf, w3_hf = value.chunk(2, dim=-1) + w1_hf, w3_hf = value.chunk(2, dim=-2) tt_state_dict[ f"layers.{idx}.moe.routed_experts.inner_experts.w1_EFD" - ] = w1_hf.transpose(-2, -1) + ] = w1_hf tt_state_dict[ f"layers.{idx}.moe.routed_experts.inner_experts.w3_EFD" - ] = w3_hf.transpose(-2, -1) + ] = w3_hf continue - # MoE down_proj → transpose + # MoE down_proj has the same layout as TT w2 if ( hf_abstract_key == "model.language_model.layers.{}.mlp.experts.down_proj" ): tt_state_dict[ f"layers.{idx}.moe.routed_experts.inner_experts.w2_EDF" - ] = value.transpose(-2, -1) + ] = value continue # GatedDeltaNet fused in_proj_qkv → split into q/k/v