From 56b11a311ded2278477fc41d7e24cdf2c8889ac6 Mon Sep 17 00:00:00 2001 From: Luca Motz Date: Fri, 4 Sep 2026 20:51:54 +0200 Subject: [PATCH 1/3] [Bugfix][Model][Spec Decode] Defer disposable GLM MTP head GLM MTP checkpoints reuse the target model's lm_head, but model construction still allocates a draft ParallelLMHead that the proposer immediately replaces. Defer only that GLM allocation behind an opt-in, fail-closed placeholder while preserving SharedHead's default behavior. Co-authored-by: OpenAI Codex Signed-off-by: Luca Motz --- tests/v1/spec_decode/test_mtp.py | 53 ++++++++++++++++++++++ vllm/model_executor/models/deepseek_mtp.py | 22 ++++++--- vllm/models/glm5next/nvidia/mtp.py | 5 +- 3 files changed, 72 insertions(+), 8 deletions(-) diff --git a/tests/v1/spec_decode/test_mtp.py b/tests/v1/spec_decode/test_mtp.py index e334371f6d8a..dbcceffbd99d 100644 --- a/tests/v1/spec_decode/test_mtp.py +++ b/tests/v1/spec_decode/test_mtp.py @@ -5,6 +5,7 @@ import pytest import torch +import torch.nn as nn from tests.v1.attention.utils import ( BatchSpec, @@ -22,6 +23,7 @@ VllmConfig, ) from vllm.config.load import LoadConfig +from vllm.model_executor.models.deepseek_mtp import DeferredLMHead, SharedHead from vllm.model_executor.models.llama import LlamaForCausalLM from vllm.platforms import current_platform from vllm.v1.attention.backends.registry import AttentionBackendEnum @@ -31,6 +33,52 @@ DEVICE_TYPE = current_platform.device_type +def test_shared_head_can_defer_lm_head(default_vllm_config): + config = mock.MagicMock(hidden_size=16, rms_norm_eps=1e-5, vocab_size=32) + + with mock.patch( + "vllm.model_executor.models.deepseek_mtp.ParallelLMHead" + ) as parallel_lm_head: + regular = SharedHead(config, "mtp") + deferred = SharedHead(config, "mtp", defer_lm_head=True) + + assert regular.head is parallel_lm_head.return_value + assert isinstance(deferred.head, DeferredLMHead) + assert not tuple(deferred.head.parameters()) + with pytest.raises(RuntimeError, match="was not replaced"): + deferred.head(torch.zeros(1, 16)) + + +def test_glm_mtp_defers_shared_head(): + from vllm.models.glm5next.nvidia import mtp + + config = mock.MagicMock( + hidden_size=16, + rms_norm_eps=1e-5, + index_topk=8, + index_kpool=4, + ) + vllm_config = mock.MagicMock() + vllm_config.speculative_config.draft_model_config.hf_config = config + vllm_config.scheduler_config.max_num_batched_tokens = 4 + + with ( + mock.patch.object( + mtp, "RMSNorm", side_effect=lambda *args, **kwargs: nn.Identity() + ), + mock.patch.object(mtp, "Glm5NextDecoderLayer", return_value=nn.Identity()), + mock.patch.object(mtp, "SharedHead", return_value=nn.Identity()) as shared_head, + mock.patch.object(mtp.current_platform, "device_type", "cpu"), + ): + mtp.Glm5NextMultiTokenPredictorLayer(vllm_config, "model.layers.1") + + shared_head.assert_called_once_with( + config=config, + prefix="model.layers.1", + defer_lm_head=True, + ) + + def _create_mtp_proposer(num_speculative_tokens: int) -> EagleProposer: """Create an MTP proposer with unified model configuration.""" model_config = ModelConfig( @@ -70,6 +118,10 @@ def test_mtp_load_model_unified(mock_get_model, mock_get_layers, mock_get_pp_gro # Setup mocks mock_model = mock.MagicMock() mock_model.model.embed_tokens.weight.shape = (131072, 4096) + draft_head = mock.MagicMock() + mock_model.model.layers = [ + mock.MagicMock(shared_head=mock.MagicMock(head=draft_head)) + ] mock_get_model.return_value = mock_model # MTP does not have its own embed_tokens or lm_head # so it should share them with the target model @@ -111,6 +163,7 @@ class _TargetModelStub(LlamaForCausalLM): mock_get_model.assert_called_once() # MTP shares lm_head with target model assert proposer.model.lm_head == target_model.lm_head + assert proposer.model.model.layers[0].shared_head.head is target_model.lm_head # MTP shares embed_tokens with target model assert proposer.model.model.embed_tokens == target_model.model.embed_tokens diff --git a/vllm/model_executor/models/deepseek_mtp.py b/vllm/model_executor/models/deepseek_mtp.py index 8ac43e0bb2a1..8b3e03444656 100644 --- a/vllm/model_executor/models/deepseek_mtp.py +++ b/vllm/model_executor/models/deepseek_mtp.py @@ -46,21 +46,31 @@ ) +class DeferredLMHead(nn.Module): + def forward(self, *args, **kwargs): + raise RuntimeError("deferred MTP head was not replaced by the target lm_head") + + class SharedHead(nn.Module): def __init__( self, config: PretrainedConfig, prefix: str, quant_config: QuantizationConfig | None = None, + *, + defer_lm_head: bool = False, ) -> None: super().__init__() self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) - self.head = ParallelLMHead( - config.vocab_size, - config.hidden_size, - quant_config=quant_config, - prefix=maybe_prefix(prefix, "head"), - ) + if defer_lm_head: + self.head = DeferredLMHead() + else: + self.head = ParallelLMHead( + config.vocab_size, + config.hidden_size, + quant_config=quant_config, + prefix=maybe_prefix(prefix, "head"), + ) def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: return self.norm(hidden_states) diff --git a/vllm/models/glm5next/nvidia/mtp.py b/vllm/models/glm5next/nvidia/mtp.py index 35b452f44f79..819d38ba01fc 100644 --- a/vllm/models/glm5next/nvidia/mtp.py +++ b/vllm/models/glm5next/nvidia/mtp.py @@ -42,7 +42,6 @@ def __init__(self, vllm_config: VllmConfig, prefix: str) -> None: assert vllm_config.speculative_config is not None config = vllm_config.speculative_config.draft_model_config.hf_config self.config = config - quant_config = vllm_config.quant_config self.enorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) self.hnorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) @@ -66,7 +65,9 @@ def __init__(self, vllm_config: VllmConfig, prefix: str) -> None: device=current_platform.device_type, ) self.shared_head = SharedHead( - config=config, prefix=prefix, quant_config=quant_config + config=config, + prefix=prefix, + defer_lm_head=True, ) # MTP layers sit past the base model's hidden layers; parse the index # from the prefix (e.g. "...layers.32") so the decoder builds an MLA From 4f7151228f5548cf2317ab8af991e218ae11692b Mon Sep 17 00:00:00 2001 From: Luca Motz Date: Sat, 5 Sep 2026 09:54:48 +0200 Subject: [PATCH 2/3] test: assert deferred GLM head skips allocation Co-authored-by: OpenAI Codex Signed-off-by: Luca Motz --- tests/v1/spec_decode/test_mtp.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/v1/spec_decode/test_mtp.py b/tests/v1/spec_decode/test_mtp.py index dbcceffbd99d..5ef1afe7cdb2 100644 --- a/tests/v1/spec_decode/test_mtp.py +++ b/tests/v1/spec_decode/test_mtp.py @@ -42,6 +42,7 @@ def test_shared_head_can_defer_lm_head(default_vllm_config): regular = SharedHead(config, "mtp") deferred = SharedHead(config, "mtp", defer_lm_head=True) + parallel_lm_head.assert_called_once() assert regular.head is parallel_lm_head.return_value assert isinstance(deferred.head, DeferredLMHead) assert not tuple(deferred.head.parameters()) From a3d341f3a472ea35e68055b90e8bdc45264e1146 Mon Sep 17 00:00:00 2001 From: Luca Motz Date: Sat, 5 Sep 2026 19:57:52 +0200 Subject: [PATCH 3/3] test: remove unused vLLM config fixture Signed-off-by: Luca Motz --- tests/v1/spec_decode/test_mtp.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/v1/spec_decode/test_mtp.py b/tests/v1/spec_decode/test_mtp.py index 5ef1afe7cdb2..122aca912f6a 100644 --- a/tests/v1/spec_decode/test_mtp.py +++ b/tests/v1/spec_decode/test_mtp.py @@ -33,7 +33,7 @@ DEVICE_TYPE = current_platform.device_type -def test_shared_head_can_defer_lm_head(default_vllm_config): +def test_shared_head_can_defer_lm_head(): config = mock.MagicMock(hidden_size=16, rms_norm_eps=1e-5, vocab_size=32) with mock.patch(