diff --git a/tests/config/test_virtual_tp.py b/tests/config/test_virtual_tp.py index a1d67ddb1f49..429b9c7d412a 100644 --- a/tests/config/test_virtual_tp.py +++ b/tests/config/test_virtual_tp.py @@ -45,6 +45,7 @@ def get_model_arch_config(self): class FakeGlmDsaModelConfig: def __init__(self): + self.verified_attention_heads = None self.hf_text_config = SimpleNamespace( model_type="glm_moe_dsa", architectures=["GlmMoeDsaForCausalLM"], @@ -67,6 +68,10 @@ def get_model_arch_config(self): total_num_attention_heads=self.hf_text_config.num_attention_heads, ) + def verify_with_parallel_config(self, parallel_config): + self.verified_attention_heads = self.hf_text_config.num_attention_heads + assert self.verified_attention_heads % parallel_config.tensor_parallel_size == 0 + class FakeWrappedGlmDsaModelConfig(FakeGlmDsaModelConfig): def __init__(self): @@ -334,6 +339,23 @@ def test_b12x_virtual_tp_padding_glm_dsa_draft_tp6(): assert draft_model_config.model_arch_config.total_num_attention_heads == 66 +def test_b12x_virtual_tp_padding_glm_dsa_draft_precedes_validation(): + target_model_config = FakeGlmDsaModelConfig() + draft_model_config = FakeGlmDsaModelConfig() + spec_config = SpeculativeConfig( + method="ngram", + num_speculative_tokens=1, + ) + spec_config.method = "mtp" + spec_config.target_model_config = target_model_config + spec_config.draft_model_config = draft_model_config + spec_config.draft_parallel_config = ParallelConfig(tensor_parallel_size=6) + + spec_config._verify_args() + + assert draft_model_config.verified_attention_heads == 66 + + def test_b12x_virtual_tp_padding_minimax_m3_tp3_only(): vllm_config = _fake_vllm_config( model_config=FakeMiniMaxM3ModelConfig(), diff --git a/vllm/config/speculative.py b/vllm/config/speculative.py index da989a3af16a..e641127b7390 100644 --- a/vllm/config/speculative.py +++ b/vllm/config/speculative.py @@ -1184,6 +1184,24 @@ def create_draft_parallel_config( return draft_parallel_config + def _maybe_apply_virtual_tp_to_draft(self) -> None: + if ( + self.method != "mtp" + or self.draft_model_config is None + or self.draft_parallel_config is None + or self.draft_model_config is self.target_model_config + ): + return + + from vllm.config.virtual_tp import ( + apply_b12x_virtual_tp_padding_to_model_config, + ) + + apply_b12x_virtual_tp_padding_to_model_config( + self.draft_model_config, + self.draft_parallel_config, + ) + @field_validator("draft_attention_backend", mode="before") @classmethod def _parse_draft_attention_backend(cls, value: Any) -> Any: @@ -1244,6 +1262,7 @@ def _verify_args(self) -> Self: ) if self.draft_model_config: + self._maybe_apply_virtual_tp_to_draft() self.draft_model_config.verify_with_parallel_config( self.draft_parallel_config )