From 5e3f8f59cd83441beabf7284df4ed001496fc55f Mon Sep 17 00:00:00 2001 From: yaoyu-33 Date: Mon, 20 Jul 2026 00:37:36 -0700 Subject: [PATCH] [model] fix: Preserve GLM-5 MTP layers during conversion Signed-off-by: Yu Yao --- src/megatron/bridge/models/glm_moe_dsa/glm5_bridge.py | 4 ---- .../unit_tests/models/glm_moe_dsa/test_glm5_bridge.py | 10 +++++++++- 2 files changed, 9 insertions(+), 5 deletions(-) diff --git a/src/megatron/bridge/models/glm_moe_dsa/glm5_bridge.py b/src/megatron/bridge/models/glm_moe_dsa/glm5_bridge.py index 3b776e3e97..e27069dd79 100644 --- a/src/megatron/bridge/models/glm_moe_dsa/glm5_bridge.py +++ b/src/megatron/bridge/models/glm_moe_dsa/glm5_bridge.py @@ -74,10 +74,6 @@ def provider_bridge(self, hf_pretrained: PreTrainedCausalLM) -> MLAModelProvider provider.qk_layernorm = True provider.multi_latent_attention = True - # Disable MTP (Multi-Token Prediction) by default - # HF config has num_nextn_predict_layers=1 - provider.mtp_num_layers = None - provider.moe_grouped_gemm = True provider.moe_router_pre_softmax = True provider.moe_token_dispatcher_type = "alltoall" diff --git a/tests/unit_tests/models/glm_moe_dsa/test_glm5_bridge.py b/tests/unit_tests/models/glm_moe_dsa/test_glm5_bridge.py index 3ea8ceba72..7e42a530c1 100644 --- a/tests/unit_tests/models/glm_moe_dsa/test_glm5_bridge.py +++ b/tests/unit_tests/models/glm_moe_dsa/test_glm5_bridge.py @@ -42,6 +42,7 @@ def _provider_from_hf_config(monkeypatch: pytest.MonkeyPatch, **config_overrides config = { "first_k_dense_replace": 3, "num_hidden_layers": 78, + "num_nextn_predict_layers": 1, "moe_intermediate_size": 2048, "n_shared_experts": 1, "rope_parameters": {"rope_theta": 1_000_000}, @@ -55,11 +56,18 @@ def _provider_from_hf_config(monkeypatch: pytest.MonkeyPatch, **config_overrides monkeypatch.setattr( MegatronModelBridge, "provider_bridge", - lambda _self, _hf_pretrained: SimpleNamespace(), + lambda _self, hf_pretrained: SimpleNamespace(mtp_num_layers=hf_pretrained.config.num_nextn_predict_layers), ) return GLM5Bridge().provider_bridge(SimpleNamespace(config=SimpleNamespace(**config))) +def test_provider_bridge_preserves_mtp_architecture_from_hf_config(monkeypatch: pytest.MonkeyPatch) -> None: + """GLM-5 conversion preserves the MTP layers declared by the checkpoint config.""" + provider = _provider_from_hf_config(monkeypatch) + + assert provider.mtp_num_layers == 1 + + @pytest.mark.parametrize( ("config_overrides", "expected_topk_freq", "expected_skip_topk_offset"), [