Skip to content
Merged
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
3 changes: 3 additions & 0 deletions tests/model_executor/layers/test_fused_shared_expert.py
Original file line number Diff line number Diff line change
Expand Up @@ -487,6 +487,7 @@ def test_models_fse_init(
quantization_config=None,
runner_type="generate",
is_moe=True,
logits_processors=None,
)
vllm_config.parallel_config.enable_expert_parallel = False
if model_type == "deepseek_v4":
Expand Down Expand Up @@ -545,6 +546,8 @@ def test_models_fse_init(
vllm_config.speculative_config = SimpleNamespace(
draft_model_config=SimpleNamespace(hf_config=config),
method="mtp",
parallel_drafting=False,
enable_adaptive_verification=False,
)
mtp = DeepSeekV4MTP(vllm_config=vllm_config)
assert model.is_fused_shared_expert_enabled is (fse_enabled and not exclude)
Expand Down
85 changes: 54 additions & 31 deletions tests/test_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -132,25 +132,30 @@ def test_v2_model_runner_env_tri_state(monkeypatch, env_value, expected):
def test_rocm_keeps_compiled_deepseek_defaults(monkeypatch):
"""ROCm keeps DeepSeek V3.2 and V4 on their compiled MRV1 paths."""
from vllm.config.vllm import (
ROCM_DEFAULT_MRV1_ARCHITECTURES,
default_breakable_cudagraph_architectures,
default_v2_model_runner_architectures,
)
from vllm.platforms import current_platform

monkeypatch.setattr(current_platform, "is_rocm", lambda: True)
# The lookup is lru_cached against a fixed platform.
default_v2_model_runner_architectures.cache_clear()
default_breakable_cudagraph_architectures.cache_clear()
try:
v2_architectures = default_v2_model_runner_architectures()
breakable_architectures = default_breakable_cudagraph_architectures()
assert "DeepseekV32ForCausalLM" in ROCM_DEFAULT_MRV1_ARCHITECTURES
assert "DeepseekV4ForCausalLM" in ROCM_DEFAULT_MRV1_ARCHITECTURES

assert "DeepseekV32ForCausalLM" not in v2_architectures
assert "DeepseekV4ForCausalLM" not in v2_architectures
breakable_architectures = default_breakable_cudagraph_architectures()
assert "DeepseekV32ForCausalLM" not in breakable_architectures
assert "DeepseekV32MTPModel" not in breakable_architectures

# The carve-out takes effect via the runner-selection property
# (warning_once args must be hashable for its lru_cache).
monkeypatch.delenv("VLLM_USE_V2_MODEL_RUNNER", raising=False)
config = SimpleNamespace(
model_config=SimpleNamespace(architectures=["DeepseekV32ForCausalLM"])
)
assert VllmConfig.use_v2_model_runner.fget(config) is False
finally:
default_v2_model_runner_architectures.cache_clear()
default_breakable_cudagraph_architectures.cache_clear()


Expand All @@ -169,17 +174,13 @@ def test_dsa_models_default_to_mrv2_and_breakable_cudagraph(
from vllm.compilation.breakable_cudagraph import (
is_breakable_cudagraph_enabled,
)
from vllm.config.vllm import (
default_breakable_cudagraph_architectures,
default_v2_model_runner_architectures,
)
from vllm.config.vllm import default_breakable_cudagraph_architectures
from vllm.platforms import current_platform

monkeypatch.delenv("VLLM_USE_BREAKABLE_CUDAGRAPH", raising=False)
monkeypatch.delenv("VLLM_USE_V2_MODEL_RUNNER", raising=False)
monkeypatch.setattr(vllm_config_module, "HAS_TRITON", True)
monkeypatch.setattr(current_platform, "is_rocm", lambda: False)
default_v2_model_runner_architectures.cache_clear()
default_breakable_cudagraph_architectures.cache_clear()

model_config = SimpleNamespace(
Expand All @@ -200,10 +201,6 @@ def test_dsa_models_default_to_mrv2_and_breakable_cudagraph(
),
)
config._dflash_needs_multi_kv_group = lambda: False
config._is_dflash2_draft = lambda: False
config._is_default_v2_model_runner_model = lambda: (
VllmConfig._is_default_v2_model_runner_model(config)
)
config._get_v2_model_runner_unsupported_features = lambda: []
config._uses_breakable_cudagraph_by_default = lambda: (
VllmConfig._uses_breakable_cudagraph_by_default(config)
Expand All @@ -217,7 +214,6 @@ def test_dsa_models_default_to_mrv2_and_breakable_cudagraph(
assert config.compilation_config.cudagraph_mode.has_piecewise_cudagraphs()
finally:
os.environ.pop("VLLM_USE_BREAKABLE_CUDAGRAPH", None)
default_v2_model_runner_architectures.cache_clear()
default_breakable_cudagraph_architectures.cache_clear()


Expand Down Expand Up @@ -387,7 +383,7 @@ def test_resolve_cudagraph_mode_skips_mamba_block_check_while_profiling():
[
(
SimpleNamespace(
model="Qwen/Qwen3-1.7B-Base",
model="Qwen/Qwen3-32B",
architectures=["Qwen3ForCausalLM"],
runner_type="generate",
is_moe=False,
Expand All @@ -397,8 +393,8 @@ def test_resolve_cudagraph_mode_skips_mamba_block_check_while_profiling():
),
(
SimpleNamespace(
model="Qwen/Qwen3-32B",
architectures=["Qwen3ForCausalLM"],
model="Qwen/Qwen2-7B-Instruct",
architectures=["Qwen2ForCausalLM"],
runner_type="generate",
is_moe=False,
is_quantized=False,
Expand Down Expand Up @@ -465,6 +461,16 @@ def test_resolve_cudagraph_mode_skips_mamba_block_check_while_profiling():
),
True,
),
(
SimpleNamespace(
model="deepseek-ai/DeepSeek-V3",
architectures=["DeepseekV3ForCausalLM"],
runner_type="generate",
is_moe=True,
is_quantized=False,
),
True,
),
(
SimpleNamespace(
model="deepseek-ai/DeepSeek-V4-Flash",
Expand Down Expand Up @@ -533,7 +539,7 @@ def test_resolve_cudagraph_mode_skips_mamba_block_check_while_profiling():
is_moe=True,
is_quantized=False,
),
False,
True,
),
(
SimpleNamespace(
Expand All @@ -554,7 +560,7 @@ def test_resolve_cudagraph_mode_skips_mamba_block_check_while_profiling():
is_quantized=False,
is_hybrid=True,
),
False,
True,
),
(
SimpleNamespace(
Expand All @@ -565,7 +571,7 @@ def test_resolve_cudagraph_mode_skips_mamba_block_check_while_profiling():
is_quantized=False,
is_attention_free=True,
),
False,
True,
),
(
SimpleNamespace(
Expand Down Expand Up @@ -602,20 +608,37 @@ def test_resolve_cudagraph_mode_skips_mamba_block_check_while_profiling():
),
],
)
def test_is_default_v2_model_runner_model(model_config, expected, monkeypatch):
from vllm.config.vllm import default_v2_model_runner_architectures
def test_models_default_to_v2_model_runner(model_config, expected, monkeypatch):
from vllm.platforms import current_platform

# The expectations below are the platform-independent defaults; ROCm's
# DeepSeek V4 carve-out is covered by test_rocm_defaults_deepseek_v4_to_mrv1.
# DeepSeek carve-out is covered by test_rocm_keeps_compiled_deepseek_defaults.
monkeypatch.delenv("VLLM_USE_V2_MODEL_RUNNER", raising=False)
monkeypatch.setattr(vllm_config_module, "HAS_TRITON", True)
monkeypatch.setattr(current_platform, "is_rocm", lambda: False)
default_v2_model_runner_architectures.cache_clear()
config = SimpleNamespace(model_config=model_config)
config._get_v2_model_runner_unsupported_features = lambda: []

try:
assert VllmConfig._is_default_v2_model_runner_model(config) is expected
finally:
default_v2_model_runner_architectures.cache_clear()
assert VllmConfig.use_v2_model_runner.fget(config) is expected


def test_v1_model_runner_rejects_v2_only_features():
config = SimpleNamespace(
parallel_config=SimpleNamespace(
prefill_context_parallel_size=2,
enable_batch_sharded_sampling=False,
),
speculative_config=None,
model_config=None,
)
config._dflash_needs_multi_kv_group = lambda: False
config._is_dflash2_draft = lambda: False
config._get_v1_model_runner_unsupported_features = lambda: (
VllmConfig._get_v1_model_runner_unsupported_features(config)
)

with pytest.raises(ValueError, match="prefill context parallel"):
VllmConfig._validate_v1_model_runner(config)


@pytest.mark.skip_global_cleanup
Expand Down
9 changes: 9 additions & 0 deletions vllm/compilation/decorators.py
Original file line number Diff line number Diff line change
Expand Up @@ -591,6 +591,15 @@ def __call__(self: type[_T], *args: Any, **kwargs: Any) -> Any:
**kwargs,
)

# Tensors that enter the graph through the forward context rather than
# model args (e.g. the MoE padding mask) need their batch dim marked
# dynamic too, otherwise AOT-compiled artifacts bake a static guard
# for the compile-time batch size.
if is_forward_context_available():
is_padding = get_forward_context().is_padding
if is_padding is not None:
torch._dynamo.mark_dynamic(is_padding, 0)

original_code_object = self.original_code_object()
logger.debug("Start compiling function %s", original_code_object)

Expand Down
Loading
Loading