Skip to content
54 changes: 49 additions & 5 deletions tests/model_executor/model_loader/test_weight_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,10 +30,19 @@
class WeightCacheDaemon:
"""Context manager running the real weight cache daemon as a subprocess."""

def __init__(self, model: str, tp_size: int, extra_args: list[str] | None = None):
def __init__(
self,
model: str,
tp_size: int,
extra_args: list[str] | None = None,
num_groups: int = 1,
):
# Short base path: Unix socket paths are limited to ~107 characters.
self.socket_dir = tempfile.mkdtemp(prefix="vllm_ipc_")
self.tp_size = tp_size
# Each daemon group (target plus cached drafts) binds one socket per
# local rank.
self.num_sockets = tp_size * num_groups
self._cmd = [
sys.executable,
"-m",
Expand Down Expand Up @@ -87,7 +96,7 @@ def _wait_ready(self, timeout_s: float) -> None:
f"Weight cache daemon exited with {self._proc.returncode}:\n"
f"{self._logs()}"
)
if len(glob.glob(pattern)) >= self.tp_size:
if len(glob.glob(pattern)) >= self.num_sockets:
return
time.sleep(1.0)
raise TimeoutError(
Expand All @@ -112,6 +121,9 @@ class ModelCase:
images: list | None = None
llm_kwargs: dict[str, Any] = field(default_factory=dict)
daemon_args: list[str] = field(default_factory=list)
# Daemon groups the launcher starts; 2 when a cached draft group joins the
# target group (MTP/EAGLE speculative decoding).
daemon_groups: int = 1


def generate(
Expand Down Expand Up @@ -172,14 +184,41 @@ def generate(
daemon_args=["--trust-remote-code"],
)

# Qwen3.5-0.8B ships one MTP layer in the target checkpoint, so method="mtp"
# loads the draft from the same model. The daemon must cache it in its draft
# group for the warm runs (fallback=False) to succeed.
QWEN_MTP_CASE = ModelCase(
model="Qwen/Qwen3.5-0.8B",
prompts=[
"Hello, my name is",
"The capital of France is",
],
llm_kwargs=dict(
gpu_memory_utilization=0.3,
enforce_eager=True,
enable_chunked_prefill=True,
speculative_config={"method": "mtp", "num_speculative_tokens": 1},
),
daemon_args=[
"--speculative-config",
'{"method": "mtp", "num_speculative_tokens": 1}',
],
daemon_groups=2,
)

@pytest.mark.parametrize("case", [QWEN_CASE, K3_CASE], ids=["qwen3.5", "kimi-k3"])

@pytest.mark.parametrize(
"case",
[QWEN_CASE, K3_CASE, QWEN_MTP_CASE],
ids=["qwen3.5", "kimi-k3", "qwen3.5-mtp"],
)
def test_ipc_cache_cold_start_and_warm_restart(vllm_runner, case: ModelCase):
"""Cold start falls back to disk; warm restarts load weights via CUDA IPC.

All runs must produce outputs identical to a default-loader baseline. The
warm runs disable the disk fallback, so they only pass if the weights
really came from the daemon.
really came from the daemon — for the MTP case, both the target's and the
draft's daemon groups.
"""
if not current_platform.is_cuda_alike():
pytest.skip("Weight cache IPC sharing requires CUDA or ROCm")
Expand All @@ -199,7 +238,12 @@ def test_ipc_cache_cold_start_and_warm_restart(vllm_runner, case: ModelCase):
with tempfile.TemporaryDirectory(prefix="vllm_ipc_empty_") as empty_socket_dir:
cold_outputs = generate(vllm_runner, case, empty_socket_dir, fallback=True)

with WeightCacheDaemon(case.model, tp_size=1, extra_args=case.daemon_args) as d:
with WeightCacheDaemon(
case.model,
tp_size=1,
extra_args=case.daemon_args,
num_groups=case.daemon_groups,
) as d:
warm_outputs = generate(vllm_runner, case, d.socket_dir, fallback=False)
# Warm restart: a second engine lifetime against the same daemon.
restart_outputs = generate(vllm_runner, case, d.socket_dir, fallback=False)
Expand Down
3 changes: 2 additions & 1 deletion tests/model_executor/test_qwen3_omni.py
Original file line number Diff line number Diff line change
Expand Up @@ -369,11 +369,12 @@ def test_dspark_shares_target_embedding_with_smaller_draft_vocabulary():
draft_parallel_config=SimpleNamespace(tensor_parallel_size=1),
attention_backend=None,
kv_cache_dtype=None,
draft_load_config=None,
),
parallel_config=ParallelConfig(),
attention_config=SimpleNamespace(backend=None),
cache_config=SimpleNamespace(),
load_config=SimpleNamespace(),
load_config=SimpleNamespace(load_format="auto"),
model_config=SimpleNamespace(get_vocab_size=Mock(return_value=100)),
)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@

import pytest

from vllm.config import LoadConfig
from vllm.config import LoadConfig, SpeculativeConfig
from vllm.v1.worker.gpu.spec_decode.eagle.utils import load_eagle_model


Expand All @@ -37,6 +37,9 @@ class _SpeculativeConfig:
moe_backend: str | None = None
kv_cache_dtype: str | None = None
draft_model_config: object = None
draft_load_config: object = None

apply_draft_overrides = SpeculativeConfig.apply_draft_overrides


@dataclass
Expand Down
5 changes: 4 additions & 1 deletion tests/v1/spec_decode/test_draft_moe_backend_override.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@

import pytest

from vllm.config import LoadConfig
from vllm.config import LoadConfig, SpeculativeConfig
from vllm.v1.worker.gpu.spec_decode.eagle.utils import load_eagle_model


Expand All @@ -34,6 +34,9 @@ class _SpeculativeConfig:
moe_backend: str | None = None
kv_cache_dtype: str | None = None
draft_model_config: object = None
draft_load_config: object = None

apply_draft_overrides = SpeculativeConfig.apply_draft_overrides


@dataclass
Expand Down
27 changes: 26 additions & 1 deletion vllm/config/speculative.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@
from vllm.config.kernel import MoEBackend
from vllm.config.model import HfOverrides, ModelConfig
from vllm.config.parallel import ParallelConfig
from vllm.config.utils import config
from vllm.config.utils import config, replace
from vllm.logger import init_logger
from vllm.transformers_utils.config import get_hf_text_config
from vllm.utils.hashing import safe_hash
Expand All @@ -25,8 +25,10 @@
from transformers import PretrainedConfig

import vllm.model_executor.layers.quantization as me_quant
from vllm.config.vllm import VllmConfig
else:
PretrainedConfig = Any
VllmConfig = Any

me_quant = LazyLoader(
"model_executor", globals(), "vllm.model_executor.layers.quantization"
Expand Down Expand Up @@ -371,6 +373,17 @@ def _validate_qwen3_omni_dspark(
)


# (SpeculativeConfig field, VllmConfig sub-config, overridden field)
_DRAFT_VLLM_CONFIG_OVERRIDES = (
# Otherwise the draft inherits the target's --moe-backend, which fails
# when the draft is unquantized and that backend is not.
("moe_backend", "kernel_config", "moe_backend"),
# Only when set, so the draft keeps a KV cache layout the target shares.
("attention_backend", "attention_config", "backend"),
("kv_cache_dtype", "cache_config", "cache_dtype"),
)


@config
class SpeculativeConfig:
"""Configuration for speculative decoding."""
Expand Down Expand Up @@ -1771,6 +1784,18 @@ def create_draft_parallel_config(

return draft_parallel_config

def apply_draft_overrides(self, vllm_config: VllmConfig) -> VllmConfig:
"""Overlay this config's kernel overrides onto a target VllmConfig.

Only non-None fields override, so an unset field keeps whatever the
target resolved.
"""
for src, config_name, dst in _DRAFT_VLLM_CONFIG_OVERRIDES:
if (value := getattr(self, src)) is not None:
sub_config = replace(getattr(vllm_config, config_name), **{dst: value})
vllm_config = replace(vllm_config, **{config_name: sub_config})
return vllm_config

@field_validator("attention_backend", mode="before")
@classmethod
def _parse_attention_backend(cls, value: Any) -> Any:
Expand Down
38 changes: 37 additions & 1 deletion vllm/model_executor/model_loader/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,13 @@
from typing_extensions import assert_never

import vllm.envs as envs
from vllm.config import ModelConfig, VllmConfig, set_current_vllm_config
from vllm.config import (
LoadConfig,
ModelConfig,
VllmConfig,
replace,
set_current_vllm_config,
)
from vllm.logger import init_logger
from vllm.model_executor.layers.attention import is_deferred_attention_layer
from vllm.model_executor.layers.quantization.base_config import (
Expand All @@ -23,6 +29,9 @@
record_metadata_for_reloading,
set_torchao_reload_attrs,
)
from vllm.model_executor.model_loader.weight_cache.utils import (
is_draft_model_cacheable,
)
from vllm.model_executor.model_loader.weight_tying import maybe_retie_word_embeddings
from vllm.model_executor.models.interfaces import SupportsQuant
from vllm.model_executor.utils import is_weights_pre_processed
Expand All @@ -34,6 +43,33 @@
logger = init_logger(__name__)


def get_draft_load_config(vllm_config: VllmConfig) -> LoadConfig:
"""Get load config for the speculative draft model."""
speculative_config = vllm_config.speculative_config
if (
speculative_config is not None
and speculative_config.draft_load_config is not None
):
return speculative_config.draft_load_config
load_config = vllm_config.load_config
if load_config is not None and load_config.load_format != "ipc_cache":
return load_config
kwargs = (
# Route the draft to the daemon's draft group.
{
"model_loader_extra_config": {
**load_config.model_loader_extra_config,
"is_draft": True,
}
}
if is_draft_model_cacheable(speculative_config)
# No daemon draft group for this method; load from disk instead of
# hitting the target daemon with a mismatching fingerprint.
else {"load_format": "auto", "model_loader_extra_config": {}}
)
return replace(load_config, **kwargs)


@instrument(span_name="Initialize model")
def initialize_model(
vllm_config: VllmConfig,
Expand Down
23 changes: 0 additions & 23 deletions vllm/model_executor/model_loader/weight_cache/__init__.py
Original file line number Diff line number Diff line change
@@ -1,26 +1,3 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

from vllm.model_executor.model_loader.weight_cache.ipc_loader import IpcModelLoader
from vllm.model_executor.model_loader.weight_cache.protocol import (
CacheConfigMismatchError,
TensorEntry,
UnsupportedPlatformForIPCError,
UnsupportedQuantForIPCError,
WeightCacheKey,
WeightCacheUnavailableError,
check_ipc_platform_support,
check_ipc_quant_support,
)

__all__ = [
"CacheConfigMismatchError",
"IpcModelLoader",
"TensorEntry",
"UnsupportedPlatformForIPCError",
"UnsupportedQuantForIPCError",
"WeightCacheKey",
"WeightCacheUnavailableError",
"check_ipc_platform_support",
"check_ipc_quant_support",
]
Loading
Loading