Skip to content

Commit 5e3f8f5

Browse files
committed
[model] fix: Preserve GLM-5 MTP layers during conversion
Signed-off-by: Yu Yao <yaoyu.094@gmail.com>
1 parent e22cb80 commit 5e3f8f5

2 files changed

Lines changed: 9 additions & 5 deletions

File tree

src/megatron/bridge/models/glm_moe_dsa/glm5_bridge.py

Lines changed: 0 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -74,10 +74,6 @@ def provider_bridge(self, hf_pretrained: PreTrainedCausalLM) -> MLAModelProvider
7474
provider.qk_layernorm = True
7575
provider.multi_latent_attention = True
7676

77-
# Disable MTP (Multi-Token Prediction) by default
78-
# HF config has num_nextn_predict_layers=1
79-
provider.mtp_num_layers = None
80-
8177
provider.moe_grouped_gemm = True
8278
provider.moe_router_pre_softmax = True
8379
provider.moe_token_dispatcher_type = "alltoall"

tests/unit_tests/models/glm_moe_dsa/test_glm5_bridge.py

Lines changed: 9 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -42,6 +42,7 @@ def _provider_from_hf_config(monkeypatch: pytest.MonkeyPatch, **config_overrides
4242
config = {
4343
"first_k_dense_replace": 3,
4444
"num_hidden_layers": 78,
45+
"num_nextn_predict_layers": 1,
4546
"moe_intermediate_size": 2048,
4647
"n_shared_experts": 1,
4748
"rope_parameters": {"rope_theta": 1_000_000},
@@ -55,11 +56,18 @@ def _provider_from_hf_config(monkeypatch: pytest.MonkeyPatch, **config_overrides
5556
monkeypatch.setattr(
5657
MegatronModelBridge,
5758
"provider_bridge",
58-
lambda _self, _hf_pretrained: SimpleNamespace(),
59+
lambda _self, hf_pretrained: SimpleNamespace(mtp_num_layers=hf_pretrained.config.num_nextn_predict_layers),
5960
)
6061
return GLM5Bridge().provider_bridge(SimpleNamespace(config=SimpleNamespace(**config)))
6162

6263

64+
def test_provider_bridge_preserves_mtp_architecture_from_hf_config(monkeypatch: pytest.MonkeyPatch) -> None:
65+
"""GLM-5 conversion preserves the MTP layers declared by the checkpoint config."""
66+
provider = _provider_from_hf_config(monkeypatch)
67+
68+
assert provider.mtp_num_layers == 1
69+
70+
6371
@pytest.mark.parametrize(
6472
("config_overrides", "expected_topk_freq", "expected_skip_topk_offset"),
6573
[

0 commit comments

Comments
 (0)