Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 22 additions & 0 deletions tests/config/test_virtual_tp.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"],
Expand All @@ -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):
Expand Down Expand Up @@ -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(),
Expand Down
19 changes: 19 additions & 0 deletions vllm/config/speculative.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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
)
Expand Down