From 13a3b1077cd40659f252b18d8fec7603a1ced916 Mon Sep 17 00:00:00 2001 From: liusy58 Date: Thu, 17 Sep 2026 07:09:58 +0000 Subject: [PATCH 01/12] add Signed-off-by: liusy58 --- .../model_loader/test_weight_cache.py | 68 +++++++ vllm/config/load.py | 5 + vllm/model_executor/model_loader/utils.py | 37 +++- .../model_loader/weight_cache/daemon.py | 168 +++++++++++++++--- .../model_loader/weight_cache/ipc_loader.py | 26 ++- .../model_loader/weight_cache/protocol.py | 57 +++++- vllm/v1/spec_decode/llm_base_proposer.py | 3 +- .../v1/worker/gpu/spec_decode/dflash/utils.py | 3 +- .../v1/worker/gpu/spec_decode/dspark/utils.py | 3 +- vllm/v1/worker/gpu/spec_decode/eagle/utils.py | 5 +- .../gpu/spec_decode/gemma4/speculator.py | 3 +- 11 files changed, 336 insertions(+), 42 deletions(-) diff --git a/tests/model_executor/model_loader/test_weight_cache.py b/tests/model_executor/model_loader/test_weight_cache.py index 1195b8b09d1c..204347d24ea0 100644 --- a/tests/model_executor/model_loader/test_weight_cache.py +++ b/tests/model_executor/model_loader/test_weight_cache.py @@ -207,3 +207,71 @@ def test_ipc_cache_cold_start_and_warm_restart(vllm_runner, case: ModelCase): assert cold_outputs == baseline_outputs assert warm_outputs == baseline_outputs assert restart_outputs == baseline_outputs + + +def test_weight_cache_caches_only_mtp_drafts(): + from types import SimpleNamespace + + from vllm.model_executor.model_loader.weight_cache.protocol import ( + caches_draft_model, + ) + + draft = object() + assert caches_draft_model(SimpleNamespace(method="mtp", draft_model_config=draft)) + assert not caches_draft_model( + SimpleNamespace(method="eagle3", draft_model_config=draft) + ) + assert not caches_draft_model( + SimpleNamespace(method="mtp", draft_model_config=None) + ) + assert not caches_draft_model(None) + + +def test_weight_cache_target_and_draft_use_distinct_sockets(tmp_path): + from vllm.model_executor.model_loader.weight_cache.protocol import get_socket_path + + target_path = get_socket_path("GPU-abc", str(tmp_path)) + draft_path = get_socket_path("GPU-abc", str(tmp_path), is_draft_model=True) + + assert target_path != draft_path + assert draft_path.endswith("GPU-abc_draft0.sock") + + +def test_draft_load_config_under_ipc_cache(): + """An MTP draft is routed to the daemon's draft group; any other draft + falls back to disk instead of hitting the target daemon; an explicit + draft_load_config always wins.""" + from types import SimpleNamespace + + from vllm.config import LoadConfig + from vllm.model_executor.model_loader.utils import get_draft_load_config + + ipc = LoadConfig( + load_format="ipc_cache", model_loader_extra_config={"fallback": False} + ) + explicit = LoadConfig(load_format="fastsafetensors") + + def cfg(method, draft_load_config=None, load_config=ipc): + return SimpleNamespace( + load_config=load_config, + speculative_config=SimpleNamespace( + method=method, + draft_model_config=object(), + draft_load_config=draft_load_config, + ), + ) + + mtp = get_draft_load_config(cfg("mtp")) + assert mtp.load_format == "ipc_cache" + assert mtp.weight_cache_is_draft_model + assert mtp.weight_cache_draft_model_idx == 0 + assert mtp.model_loader_extra_config == {"fallback": False} + + eagle = get_draft_load_config(cfg("eagle3")) + assert eagle.load_format == "auto" + assert not eagle.weight_cache_is_draft_model + assert eagle.model_loader_extra_config == {} + + assert get_draft_load_config(cfg("mtp", explicit)) is explicit + disk = LoadConfig(load_format="fastsafetensors") + assert get_draft_load_config(cfg("mtp", load_config=disk)) is disk diff --git a/vllm/config/load.py b/vllm/config/load.py index ee13e2fccd7c..c6ced1c84224 100644 --- a/vllm/config/load.py +++ b/vllm/config/load.py @@ -97,6 +97,11 @@ class LoadConfig: model_loader_extra_config: dict | TensorizerConfig = Field(default_factory=dict) """Extra config for model loader. This will be passed to the model loader corresponding to the chosen load_format.""" + weight_cache_is_draft_model: bool = False + """Whether an ``ipc_cache`` load targets the weight cache daemon group that + serves the speculative draft model instead of the target model.""" + weight_cache_draft_model_idx: int | None = None + """Index of the speculative draft model within the draft cache role.""" device: str | None = None """Device to which model weights will be loaded, default to device_config.device""" diff --git a/vllm/model_executor/model_loader/utils.py b/vllm/model_executor/model_loader/utils.py index d80b7d17027d..bd1c76754015 100644 --- a/vllm/model_executor/model_loader/utils.py +++ b/vllm/model_executor/model_loader/utils.py @@ -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 ( @@ -34,6 +40,35 @@ logger = init_logger(__name__) +def get_draft_load_config(vllm_config: VllmConfig) -> LoadConfig: + """Load config for the speculative draft model. + + An explicit ``draft_load_config`` always wins. Otherwise the draft inherits + the target's load config, except under ``ipc_cache``: an MTP draft is + routed to the daemon's draft group, and any other draft (which the daemon + does not cache) falls back to disk loading instead of being sent to the + target daemon with a mismatching fingerprint. + """ + from vllm.model_executor.model_loader.weight_cache.protocol import ( + caches_draft_model, + ) + + speculative_config = vllm_config.speculative_config + assert speculative_config is not None + load_config = vllm_config.load_config + if speculative_config.draft_load_config is not None: + return speculative_config.draft_load_config + if load_config.load_format != "ipc_cache": + return load_config + if caches_draft_model(speculative_config): + return replace( + load_config, + weight_cache_is_draft_model=True, + weight_cache_draft_model_idx=0, + ) + return replace(load_config, load_format="auto", model_loader_extra_config={}) + + @instrument(span_name="Initialize model") def initialize_model( vllm_config: VllmConfig, diff --git a/vllm/model_executor/model_loader/weight_cache/daemon.py b/vllm/model_executor/model_loader/weight_cache/daemon.py index c009cc6ffa29..033d8f8cd10f 100644 --- a/vllm/model_executor/model_loader/weight_cache/daemon.py +++ b/vllm/model_executor/model_loader/weight_cache/daemon.py @@ -41,6 +41,14 @@ ``r * (tp_size // nnodes) + i``, matching vLLM's contiguous per-node rank assignment, so each engine worker maps its shard from the daemon on its own node. + +With MTP speculative decoding (``--speculative-config`` with ``method=mtp``) +the launcher additionally starts a draft daemon group that caches the MTP +draft model. It uses its own cache key, Unix sockets (``*_draft0.sock``) and +rendezvous port (``--weight-cache-draft-master-port``, default +``--weight-cache-master-port + 1``), so each process serves exactly one model +role. Other draft types (e.g. EAGLE3 heads) are not cached and keep loading +from disk in the engine. """ import contextlib @@ -55,7 +63,12 @@ import torch -from vllm.config import ParallelConfig, VllmConfig, set_current_vllm_config +from vllm.config import ( + ParallelConfig, + VllmConfig, + replace, + set_current_vllm_config, +) from vllm.distributed import ( ensure_model_parallel_initialized, init_distributed_environment, @@ -68,9 +81,11 @@ TensorEntry, WeightCacheKey, WeightCacheUnavailableError, + caches_draft_model, check_ipc_platform_support, check_ipc_quant_support, ensure_private_socket_dir, + format_daemon_role, get_current_device_uuid, get_socket_path, recv_msg, @@ -164,12 +179,17 @@ def __init__( local_rank: int, distributed_init_method: str, socket_dir: str | None = None, + is_draft_model: bool = False, + draft_model_idx: int | None = None, ): self.vllm_config = vllm_config self.tp_rank = tp_rank self.local_rank = local_rank self.distributed_init_method = distributed_init_method self.socket_dir = socket_dir + self.is_draft_model = is_draft_model + self.draft_model_idx = draft_model_idx + self.role = "draft" if is_draft_model else "target" self.model: torch.nn.Module | None = None # Fingerprint before loading: process_weights_after_loading may # mutate hf_config.quantization_config. @@ -177,6 +197,8 @@ def __init__( vllm_config.model_config, tp_size=vllm_config.parallel_config.tensor_parallel_size, tp_rank=tp_rank, + is_draft_model=is_draft_model, + draft_model_idx=draft_model_idx, ) def load_model(self) -> None: @@ -193,7 +215,8 @@ def load_model(self) -> None: ensure_model_parallel_initialized(tp_size, 1) self.model = get_daemon_model(self.vllm_config) logger.info( - "Weight cache daemon rank %d loaded model", + "Weight cache %s daemon rank %d loaded model", + self.role, self.tp_rank, ) @@ -223,7 +246,10 @@ def serve_forever(self, ready_callback: Callable[[], None] | None = None) -> Non os.chmod(socket_path, 0o600) server.listen() logger.info( - "Weight cache daemon rank %d serving on %s", self.tp_rank, socket_path + "Weight cache %s daemon rank %d serving on %s", + self.role, + self.tp_rank, + socket_path, ) if ready_callback is not None: ready_callback() @@ -270,7 +296,12 @@ def _acquire_gpu_lock(self, socket_path: str) -> int: @property def _socket_path(self) -> str: - return get_socket_path(get_current_device_uuid(), self.socket_dir) + return get_socket_path( + get_current_device_uuid(), + self.socket_dir, + is_draft_model=self.is_draft_model, + draft_model_idx=self.draft_model_idx, + ) def _handle_connection(self, conn: socket.socket) -> None: request = recv_msg(conn) @@ -307,7 +338,8 @@ def _handle_get_state(self, conn: socket.socket, request: dict) -> None: }, ) logger.info_once( - "Weight cache daemon rank %d sent %d tensors (+%d aliases) to engine", + "Weight cache %s daemon rank %d sent %d tensors (+%d aliases) to engine", + self.role, self.tp_rank, len(entries), len(aliases), @@ -316,7 +348,11 @@ def _handle_get_state(self, conn: socket.socket, request: dict) -> None: def _handle_release(self, conn: socket.socket) -> None: self.model = None torch.accelerator.empty_cache() - logger.info("Weight cache daemon rank %d released cached weights", self.tp_rank) + logger.info( + "Weight cache %s daemon rank %d released cached weights", + self.role, + self.tp_rank, + ) send_msg(conn, {"status": "ok"}) @@ -326,13 +362,47 @@ def _run_daemon( vllm_config: VllmConfig, distributed_init_method: str, socket_dir: str | None, - ready_queue: "multiprocessing.Queue[int]", + ready_queue: "multiprocessing.Queue[tuple[str, int]]", + is_draft_model: bool = False, + draft_model_idx: int | None = None, ) -> None: daemon = WeightCacheDaemon( - vllm_config, tp_rank, local_rank, distributed_init_method, socket_dir + vllm_config, + tp_rank, + local_rank, + distributed_init_method, + socket_dir, + is_draft_model, + draft_model_idx, ) daemon.load_model() - daemon.serve_forever(ready_callback=lambda: ready_queue.put(tp_rank)) + daemon.serve_forever(ready_callback=lambda: ready_queue.put((daemon.role, tp_rank))) + + +def get_cached_draft_vllm_config(vllm_config: VllmConfig) -> VllmConfig | None: + """Config for the draft daemon group, or None when the draft is not cached. + + Only MTP drafts are cached (see ``caches_draft_model``). The draft group + reuses the target's parallel topology, so the draft must be sharded with + the same tensor parallel size. + """ + speculative_config = vllm_config.speculative_config + if speculative_config is None or not caches_draft_model(speculative_config): + return None + draft_parallel_config = speculative_config.draft_parallel_config + target_tp = vllm_config.parallel_config.tensor_parallel_size + if draft_parallel_config.tensor_parallel_size != target_tp: + raise ValueError( + "The weight cache daemon requires the MTP draft and the target to " + f"use the same tensor parallel size, got " + f"{draft_parallel_config.tensor_parallel_size} != {target_tp}" + ) + _reject_unsupported_parallelism(draft_parallel_config) + return replace( + vllm_config, + model_config=speculative_config.draft_model_config, + parallel_config=draft_parallel_config, + ) def _reject_unsupported_parallelism(parallel_config: ParallelConfig) -> None: @@ -369,6 +439,14 @@ def main() -> None: "serving) and match across nodes. Required when --nnodes > 1; defaults " "to a free port for single-node.", ) + parser.add_argument( + "--weight-cache-draft-master-port", + type=int, + default=None, + help="Rendezvous port for the MTP draft daemon group. Defaults to " + "--weight-cache-master-port + 1 for multi-node, or a free port for " + "single-node.", + ) args = parser.parse_args() engine_args = EngineArgs.from_cli_args(args) vllm_config = engine_args.create_engine_config() @@ -398,28 +476,59 @@ def main() -> None: master_port = args.weight_cache_master_port or get_open_port() distributed_init_method = get_distributed_init_method(master_addr, master_port) + # (is_draft_model, draft_model_idx, config, rendezvous) per daemon group. + groups: list[tuple[bool, int | None, VllmConfig, str]] = [ + (False, None, vllm_config, distributed_init_method) + ] + draft_vllm_config = get_cached_draft_vllm_config(vllm_config) + if draft_vllm_config is not None: + draft_master_port = args.weight_cache_draft_master_port + if draft_master_port is None: + draft_master_port = master_port + 1 if nnodes > 1 else get_open_port() + if draft_master_port == master_port: + raise ValueError( + "--weight-cache-draft-master-port must differ from " + "--weight-cache-master-port" + ) + groups.append( + ( + True, + 0, + draft_vllm_config, + get_distributed_init_method(master_addr, draft_master_port), + ) + ) + ctx = multiprocessing.get_context("spawn") - ready_queue: multiprocessing.Queue[int] = ctx.Queue() + ready_queue: multiprocessing.Queue[tuple[str, int]] = ctx.Queue() # Global TP rank of local GPU i on this node; local index == device index. global_ranks = [ node_rank * local_world_size + local_rank for local_rank in range(local_world_size) ] - procs = [ - ctx.Process( - target=_run_daemon, - args=( - global_rank, - local_rank, - vllm_config, - distributed_init_method, - args.weight_cache_socket_dir, - ready_queue, - ), - name=f"vllm-weight-cache-daemon-{global_rank}", - ) - for local_rank, global_rank in enumerate(global_ranks) - ] + procs = [] + expected_ready: set[tuple[str, int]] = set() + for is_draft_model, draft_model_idx, config, init_method in groups: + role = "draft" if is_draft_model else "target" + suffix = format_daemon_role(is_draft_model, draft_model_idx) + for local_rank, global_rank in enumerate(global_ranks): + expected_ready.add((role, global_rank)) + procs.append( + ctx.Process( + target=_run_daemon, + args=( + global_rank, + local_rank, + config, + init_method, + args.weight_cache_socket_dir, + ready_queue, + is_draft_model, + draft_model_idx, + ), + name=f"vllm-weight-cache-{role}{suffix}-{global_rank}", + ) + ) for proc in procs: proc.start() @@ -430,10 +539,10 @@ def _shutdown(signum, frame): signal.signal(signal.SIGINT, _shutdown) signal.signal(signal.SIGTERM, _shutdown) - ready_ranks: set[int] = set() - while len(ready_ranks) < local_world_size: + ready: set[tuple[str, int]] = set() + while len(ready) < len(expected_ready): try: - ready_ranks.add(ready_queue.get(timeout=1.0)) + ready.add(ready_queue.get(timeout=1.0)) except queue.Empty: dead = [p for p in procs if p.exitcode is not None] if dead: @@ -450,10 +559,11 @@ def _shutdown(signum, frame): socket_dir_msg = args.weight_cache_socket_dir or "the default socket dir" logger.info_once( "===== Weight cache daemon READY: node %d/%d serving %d local rank(s) " - "in %s =====", + "x %d role(s) in %s =====", node_rank, nnodes, local_world_size, + len(groups), socket_dir_msg, ) diff --git a/vllm/model_executor/model_loader/weight_cache/ipc_loader.py b/vllm/model_executor/model_loader/weight_cache/ipc_loader.py index 6c3fb742cdf7..87d58cc7aa43 100644 --- a/vllm/model_executor/model_loader/weight_cache/ipc_loader.py +++ b/vllm/model_executor/model_loader/weight_cache/ipc_loader.py @@ -60,7 +60,8 @@ class IpcModelLoader(BaseModelLoader): Extra config keys (via --model-loader-extra-config): - socket_path: explicit daemon socket path. Defaults to a per-GPU path - derived from the physical GPU uuid. + derived from the physical GPU uuid and the cache role + (``load_config.weight_cache_is_draft_model``). - socket_dir: directory containing the daemon sockets. - mode: "zero_copy" (default) or "copy". - fallback: fall back to disk loading when the daemon is unavailable or @@ -78,6 +79,14 @@ def __init__(self, load_config: LoadConfig): extra_config = copy(load_config.model_loader_extra_config or {}) self.socket_path: str | None = extra_config.pop("socket_path", None) self.socket_dir: str | None = extra_config.pop("socket_dir", None) + self.is_draft_model = load_config.weight_cache_is_draft_model + self.draft_model_idx = load_config.weight_cache_draft_model_idx + if self.is_draft_model and self.socket_path is not None: + raise ValueError( + "socket_path cannot be combined with the draft weight cache role; " + "use socket_dir so the target and draft sockets are derived " + "independently" + ) self.mode: str = extra_config.pop("mode", "zero_copy") self.fallback: bool = extra_config.pop("fallback", True) self.connect_timeout_s: float = float( @@ -269,6 +278,8 @@ def _fetch_entries( model_config, tp_size=get_tensor_model_parallel_world_size(), tp_rank=get_tensor_model_parallel_rank(), + is_draft_model=self.is_draft_model, + draft_model_idx=self.draft_model_idx, ) return self._request_state(cache_config) @@ -316,7 +327,12 @@ def _connect(self, timeout: float) -> socket.socket: def _resolve_socket_path(self) -> str: if self.socket_path is not None: return self.socket_path - return get_socket_path(get_current_device_uuid(), self.socket_dir) + return get_socket_path( + get_current_device_uuid(), + self.socket_dir, + is_draft_model=self.is_draft_model, + draft_model_idx=self.draft_model_idx, + ) def _check_gpu_uuid(self, daemon_uuid: str | None) -> None: if daemon_uuid is None: @@ -340,7 +356,11 @@ def _fallback_load_config(self) -> LoadConfig: # DefaultModelLoader must not see load_format="ipc_cache" or the ipc # extra config keys. return dataclasses.replace( - self.load_config, load_format="auto", model_loader_extra_config={} + self.load_config, + load_format="auto", + model_loader_extra_config={}, + weight_cache_is_draft_model=False, + weight_cache_draft_model_idx=None, ) def _fallback_load( diff --git a/vllm/model_executor/model_loader/weight_cache/protocol.py b/vllm/model_executor/model_loader/weight_cache/protocol.py index 852856b139f1..1a34fce58b87 100644 --- a/vllm/model_executor/model_loader/weight_cache/protocol.py +++ b/vllm/model_executor/model_loader/weight_cache/protocol.py @@ -35,7 +35,7 @@ from vllm.platforms import current_platform from vllm.utils.hashing import safe_hash -SOCKET_NAME_TEMPLATE = "vllm_weight_cache_{gpu_uuid}.sock" +SOCKET_NAME_TEMPLATE = "vllm_weight_cache_{gpu_uuid}{role}.sock" SOCKET_DIR_TEMPLATE = "vllm_weight_cache_{uid}" _LEN_STRUCT = struct.Struct("!Q") @@ -128,9 +128,48 @@ def get_socket_dir(socket_dir: str | None = None) -> str: ) -def get_socket_path(gpu_uuid: str, socket_dir: str | None = None) -> str: +# Speculative methods whose draft model the daemon caches in its own group. +# Other drafts (e.g. EAGLE3 heads) keep loading from disk in the engine. +WEIGHT_CACHE_DRAFT_METHODS = frozenset({"mtp"}) + + +def caches_draft_model(speculative_config: Any) -> bool: + """Whether the daemon serves the speculative draft as a separate role.""" + return ( + speculative_config is not None + and speculative_config.method in WEIGHT_CACHE_DRAFT_METHODS + and speculative_config.draft_model_config is not None + ) + + +def normalize_draft_model_idx(draft_model_idx: int | None) -> int: + return -1 if draft_model_idx is None else draft_model_idx + + +def format_daemon_role( + is_draft_model: bool = False, draft_model_idx: int | None = None +) -> str: + """Socket-name suffix distinguishing the draft daemon group from the target.""" + if not is_draft_model: + return "" + return f"_draft{draft_model_idx if draft_model_idx is not None else 0}" + + +def get_socket_path( + gpu_uuid: str, + socket_dir: str | None = None, + *, + is_draft_model: bool = False, + draft_model_idx: int | None = None, +) -> str: directory = get_socket_dir(socket_dir) - return os.path.join(directory, SOCKET_NAME_TEMPLATE.format(gpu_uuid=gpu_uuid)) + return os.path.join( + directory, + SOCKET_NAME_TEMPLATE.format( + gpu_uuid=gpu_uuid, + role=format_daemon_role(is_draft_model, draft_model_idx), + ), + ) def ensure_private_socket_dir(directory: str, strict_perms: bool = True) -> None: @@ -272,10 +311,18 @@ class WeightCacheKey: quant_config_hash: str revision: str | None vllm_version: str + is_draft_model: bool = False + draft_model_idx: int = -1 @classmethod def from_model_config( - cls, model_config: ModelConfig, tp_size: int, tp_rank: int + cls, + model_config: ModelConfig, + tp_size: int, + tp_rank: int, + *, + is_draft_model: bool = False, + draft_model_idx: int | None = None, ) -> "WeightCacheKey": """Build the fingerprint for a model configuration. @@ -302,6 +349,8 @@ def from_model_config( quant_config_hash=_hash_quant_config(quant_config), revision=model_config.revision, vllm_version=vllm.version.__version__, + is_draft_model=is_draft_model, + draft_model_idx=normalize_draft_model_idx(draft_model_idx), ) def mismatched_fields(self, other: "WeightCacheKey") -> list[str]: diff --git a/vllm/v1/spec_decode/llm_base_proposer.py b/vllm/v1/spec_decode/llm_base_proposer.py index 9f7ad68a88a6..e0a95f2ce5fb 100644 --- a/vllm/v1/spec_decode/llm_base_proposer.py +++ b/vllm/v1/spec_decode/llm_base_proposer.py @@ -25,6 +25,7 @@ from vllm.logger import init_logger from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase from vllm.model_executor.model_loader import get_model +from vllm.model_executor.model_loader.utils import get_draft_load_config from vllm.model_executor.models import ( supports_multimodal, supports_multimodal_embeddings, @@ -1315,7 +1316,7 @@ def _get_model(self) -> nn.Module: model = get_model( vllm_config=draft_vllm_config, model_config=self.speculative_config.draft_model_config, - load_config=self.speculative_config.draft_load_config, + load_config=get_draft_load_config(draft_vllm_config), ) return model diff --git a/vllm/v1/worker/gpu/spec_decode/dflash/utils.py b/vllm/v1/worker/gpu/spec_decode/dflash/utils.py index 7323b29312c1..2e5826362346 100644 --- a/vllm/v1/worker/gpu/spec_decode/dflash/utils.py +++ b/vllm/v1/worker/gpu/spec_decode/dflash/utils.py @@ -4,6 +4,7 @@ from vllm.config import VllmConfig, replace from vllm.model_executor.model_loader import get_model +from vllm.model_executor.model_loader.utils import get_draft_load_config from vllm.v1.worker.gpu.spec_decode.eagle.utils import ( _should_share, get_target_lm_head, @@ -38,7 +39,7 @@ def load_dflash_model(target_model: nn.Module, vllm_config: VllmConfig) -> nn.Mo if speculative_config.kv_cache_dtype is not None else vllm_config.cache_config ), - load_config=get_pp_safe_draft_load_config(vllm_config.load_config), + load_config=get_pp_safe_draft_load_config(get_draft_load_config(vllm_config)), ) with set_model_tag("dflash_head"): dflash_model = get_model( diff --git a/vllm/v1/worker/gpu/spec_decode/dspark/utils.py b/vllm/v1/worker/gpu/spec_decode/dspark/utils.py index bf1a5761b2ef..489d0da554a2 100644 --- a/vllm/v1/worker/gpu/spec_decode/dspark/utils.py +++ b/vllm/v1/worker/gpu/spec_decode/dspark/utils.py @@ -5,6 +5,7 @@ from vllm.config import ModelConfig, ParallelConfig, VllmConfig, replace from vllm.logger import init_logger +from vllm.model_executor.model_loader.utils import get_draft_load_config from vllm.v1.attention.backends.registry import AttentionBackendEnum from vllm.v1.worker.gpu.spec_decode.utils import get_pp_safe_draft_load_config @@ -94,7 +95,7 @@ def load_dspark_model(target_model: nn.Module, vllm_config: VllmConfig) -> nn.Mo if speculative_config.kv_cache_dtype is not None else vllm_config.cache_config ), - load_config=get_pp_safe_draft_load_config(vllm_config.load_config), + load_config=get_pp_safe_draft_load_config(get_draft_load_config(vllm_config)), ) # VllmConfig post-init restores the target's quant config because the target # config is retained for DSpark's target-layer metadata, so we must override it. diff --git a/vllm/v1/worker/gpu/spec_decode/eagle/utils.py b/vllm/v1/worker/gpu/spec_decode/eagle/utils.py index eb0f4793bf1f..54b0b7776a98 100644 --- a/vllm/v1/worker/gpu/spec_decode/eagle/utils.py +++ b/vllm/v1/worker/gpu/spec_decode/eagle/utils.py @@ -7,6 +7,7 @@ from vllm.distributed.parallel_state import get_pp_group from vllm.lora.layers.base import BaseLayerWithLoRA from vllm.model_executor.model_loader import get_model +from vllm.model_executor.model_loader.utils import get_draft_load_config from vllm.model_executor.models.utils import PPMissingLayer from vllm.v1.worker.gpu.spec_decode.utils import get_pp_safe_draft_load_config @@ -103,7 +104,9 @@ def load_eagle_model(target_model: nn.Module, vllm_config: VllmConfig) -> nn.Mod backend=speculative_config.attention_backend, ), ) - draft_load_config = get_pp_safe_draft_load_config(vllm_config.load_config) + draft_load_config = get_pp_safe_draft_load_config( + get_draft_load_config(vllm_config) + ) if draft_load_config is not vllm_config.load_config: vllm_config = replace(vllm_config, load_config=draft_load_config) with set_model_tag("eagle_head"): diff --git a/vllm/v1/worker/gpu/spec_decode/gemma4/speculator.py b/vllm/v1/worker/gpu/spec_decode/gemma4/speculator.py index dfa2c680109d..9a489bcc4844 100644 --- a/vllm/v1/worker/gpu/spec_decode/gemma4/speculator.py +++ b/vllm/v1/worker/gpu/spec_decode/gemma4/speculator.py @@ -16,6 +16,7 @@ from vllm.distributed.parallel_state import get_pp_group from vllm.logger import init_logger from vllm.model_executor.model_loader import get_model +from vllm.model_executor.model_loader.utils import get_draft_load_config from vllm.v1.worker.gpu.spec_decode.autoregressive.speculator import ( AutoRegressiveSpeculator, ) @@ -40,7 +41,7 @@ def load_draft_model( draft_model = get_model( vllm_config=draft_vllm_config, model_config=self.speculative_config.draft_model_config, - load_config=self.speculative_config.draft_load_config, + load_config=get_draft_load_config(draft_vllm_config), ) self._setup_gemma4_kv_sharing(draft_model, target_attn_layer_names) self._share_embeddings(draft_model, target_model) From e5d0102f88f019a921888ec22bbe6a36736e88af Mon Sep 17 00:00:00 2001 From: liusy58 Date: Thu, 17 Sep 2026 13:32:39 +0000 Subject: [PATCH 02/12] add Signed-off-by: liusy58 --- .../model_loader/test_weight_cache.py | 36 +++++-- vllm/model_executor/model_loader/utils.py | 8 +- .../model_loader/weight_cache/daemon.py | 98 ++++++++++++------- .../model_loader/weight_cache/ipc_loader.py | 21 ++-- .../model_loader/weight_cache/protocol.py | 24 ++++- 5 files changed, 130 insertions(+), 57 deletions(-) diff --git a/tests/model_executor/model_loader/test_weight_cache.py b/tests/model_executor/model_loader/test_weight_cache.py index 204347d24ea0..c36a673db07b 100644 --- a/tests/model_executor/model_loader/test_weight_cache.py +++ b/tests/model_executor/model_loader/test_weight_cache.py @@ -209,7 +209,7 @@ def test_ipc_cache_cold_start_and_warm_restart(vllm_runner, case: ModelCase): assert restart_outputs == baseline_outputs -def test_weight_cache_caches_only_mtp_drafts(): +def test_weight_cache_caches_mtp_and_eagle_drafts(): from types import SimpleNamespace from vllm.model_executor.model_loader.weight_cache.protocol import ( @@ -217,9 +217,12 @@ def test_weight_cache_caches_only_mtp_drafts(): ) draft = object() - assert caches_draft_model(SimpleNamespace(method="mtp", draft_model_config=draft)) + for method in ("mtp", "eagle", "eagle3"): + assert caches_draft_model( + SimpleNamespace(method=method, draft_model_config=draft) + ) assert not caches_draft_model( - SimpleNamespace(method="eagle3", draft_model_config=draft) + SimpleNamespace(method="dflash", draft_model_config=draft) ) assert not caches_draft_model( SimpleNamespace(method="mtp", draft_model_config=None) @@ -227,6 +230,21 @@ def test_weight_cache_caches_only_mtp_drafts(): assert not caches_draft_model(None) +def test_weight_cache_exports_eagle_ownership_flags(): + """The engine never runs load_weights for cached drafts, so the flags that + decide whether to share the target's embed/lm_head travel with the state.""" + import torch + + from vllm.model_executor.model_loader.weight_cache.protocol import ( + export_model_attrs, + ) + + model = torch.nn.Module() + assert export_model_attrs(model) == {} + model.has_own_lm_head = True + assert export_model_attrs(model) == {"has_own_lm_head": True} + + def test_weight_cache_target_and_draft_use_distinct_sockets(tmp_path): from vllm.model_executor.model_loader.weight_cache.protocol import get_socket_path @@ -238,7 +256,7 @@ def test_weight_cache_target_and_draft_use_distinct_sockets(tmp_path): def test_draft_load_config_under_ipc_cache(): - """An MTP draft is routed to the daemon's draft group; any other draft + """A cached draft is routed to the daemon's draft group; any other draft falls back to disk instead of hitting the target daemon; an explicit draft_load_config always wins.""" from types import SimpleNamespace @@ -268,9 +286,13 @@ def cfg(method, draft_load_config=None, load_config=ipc): assert mtp.model_loader_extra_config == {"fallback": False} eagle = get_draft_load_config(cfg("eagle3")) - assert eagle.load_format == "auto" - assert not eagle.weight_cache_is_draft_model - assert eagle.model_loader_extra_config == {} + assert eagle.load_format == "ipc_cache" + assert eagle.weight_cache_is_draft_model + + dflash = get_draft_load_config(cfg("dflash")) + assert dflash.load_format == "auto" + assert not dflash.weight_cache_is_draft_model + assert dflash.model_loader_extra_config == {} assert get_draft_load_config(cfg("mtp", explicit)) is explicit disk = LoadConfig(load_format="fastsafetensors") diff --git a/vllm/model_executor/model_loader/utils.py b/vllm/model_executor/model_loader/utils.py index bd1c76754015..ab950be92eb3 100644 --- a/vllm/model_executor/model_loader/utils.py +++ b/vllm/model_executor/model_loader/utils.py @@ -44,10 +44,10 @@ def get_draft_load_config(vllm_config: VllmConfig) -> LoadConfig: """Load config for the speculative draft model. An explicit ``draft_load_config`` always wins. Otherwise the draft inherits - the target's load config, except under ``ipc_cache``: an MTP draft is - routed to the daemon's draft group, and any other draft (which the daemon - does not cache) falls back to disk loading instead of being sent to the - target daemon with a mismatching fingerprint. + the target's load config, except under ``ipc_cache``: a cached draft + (MTP, EAGLE, EAGLE3) is routed to the daemon's draft group, and any other + draft falls back to disk loading instead of being sent to the target + daemon with a mismatching fingerprint. """ from vllm.model_executor.model_loader.weight_cache.protocol import ( caches_draft_model, diff --git a/vllm/model_executor/model_loader/weight_cache/daemon.py b/vllm/model_executor/model_loader/weight_cache/daemon.py index 033d8f8cd10f..d6f41b4c3512 100644 --- a/vllm/model_executor/model_loader/weight_cache/daemon.py +++ b/vllm/model_executor/model_loader/weight_cache/daemon.py @@ -42,13 +42,12 @@ assignment, so each engine worker maps its shard from the daemon on its own node. -With MTP speculative decoding (``--speculative-config`` with ``method=mtp``) -the launcher additionally starts a draft daemon group that caches the MTP -draft model. It uses its own cache key, Unix sockets (``*_draft0.sock``) and -rendezvous port (``--weight-cache-draft-master-port``, default -``--weight-cache-master-port + 1``), so each process serves exactly one model -role. Other draft types (e.g. EAGLE3 heads) are not cached and keep loading -from disk in the engine. +With MTP, EAGLE or EAGLE3 speculative decoding the launcher additionally +starts a draft daemon group that caches the draft model. It uses its own cache +key, Unix sockets (``*_draft0.sock``) and rendezvous port +(``--weight-cache-draft-master-port``, default ``--weight-cache-master-port + +1``), so each process serves exactly one model role. Other draft types are not +cached and keep loading from disk in the engine. """ import contextlib @@ -64,6 +63,7 @@ import torch from vllm.config import ( + ModelConfig, ParallelConfig, VllmConfig, replace, @@ -85,6 +85,7 @@ check_ipc_platform_support, check_ipc_quant_support, ensure_private_socket_dir, + export_model_attrs, format_daemon_role, get_current_device_uuid, get_socket_path, @@ -145,7 +146,9 @@ def _add(name: str, tensor: torch.Tensor, kind: str) -> None: return entries, aliases -def get_daemon_model(vllm_config: VllmConfig) -> torch.nn.Module: +def get_daemon_model( + vllm_config: VllmConfig, model_config: ModelConfig | None = None +) -> torch.nn.Module: """Load the daemon's model, composed from the configured loader. Runs the quantization check after model creation but before the slow @@ -153,7 +156,8 @@ def get_daemon_model(vllm_config: VllmConfig) -> torch.nn.Module: always fails the check, so load_model's finalize step for it is unnecessary here. """ - model_config = vllm_config.model_config + if model_config is None: + model_config = vllm_config.model_config load_config = vllm_config.load_config loader = get_model_loader(load_config) device_config = vllm_config.device_config @@ -181,8 +185,12 @@ def __init__( socket_dir: str | None = None, is_draft_model: bool = False, draft_model_idx: int | None = None, + model_config: ModelConfig | None = None, ): self.vllm_config = vllm_config + # A draft daemon builds the draft with the target's VllmConfig, like + # the engine does; only the ModelConfig differs. + self.model_config = model_config or vllm_config.model_config self.tp_rank = tp_rank self.local_rank = local_rank self.distributed_init_method = distributed_init_method @@ -194,7 +202,7 @@ def __init__( # Fingerprint before loading: process_weights_after_loading may # mutate hf_config.quantization_config. self.cache_config = WeightCacheKey.from_model_config( - vllm_config.model_config, + self.model_config, tp_size=vllm_config.parallel_config.tensor_parallel_size, tp_rank=tp_rank, is_draft_model=is_draft_model, @@ -213,7 +221,7 @@ def load_model(self) -> None: ) with set_current_vllm_config(self.vllm_config): ensure_model_parallel_initialized(tp_size, 1) - self.model = get_daemon_model(self.vllm_config) + self.model = get_daemon_model(self.vllm_config, self.model_config) logger.info( "Weight cache %s daemon rank %d loaded model", self.role, @@ -334,6 +342,7 @@ def _handle_get_state(self, conn: socket.socket, request: dict) -> None: "status": "ok", "entries": entries, "aliases": aliases, + "attrs": export_model_attrs(self.model), "gpu_uuid": gpu_uuid, }, ) @@ -365,6 +374,7 @@ def _run_daemon( ready_queue: "multiprocessing.Queue[tuple[str, int]]", is_draft_model: bool = False, draft_model_idx: int | None = None, + model_config: ModelConfig | None = None, ) -> None: daemon = WeightCacheDaemon( vllm_config, @@ -374,35 +384,48 @@ def _run_daemon( socket_dir, is_draft_model, draft_model_idx, + model_config, ) daemon.load_model() daemon.serve_forever(ready_callback=lambda: ready_queue.put((daemon.role, tp_rank))) -def get_cached_draft_vllm_config(vllm_config: VllmConfig) -> VllmConfig | None: - """Config for the draft daemon group, or None when the draft is not cached. +def get_draft_daemon_config( + vllm_config: VllmConfig, +) -> tuple[VllmConfig, ModelConfig] | None: + """Configs for the draft daemon group, or None when the draft is not cached. - Only MTP drafts are cached (see ``caches_draft_model``). The draft group - reuses the target's parallel topology, so the draft must be sharded with - the same tensor parallel size. + Mirrors how the engine loads a draft: the target's VllmConfig with the + speculative kernel overrides, plus the draft's ModelConfig passed + separately, because draft classes read the target from + ``vllm_config.model_config``. """ speculative_config = vllm_config.speculative_config - if speculative_config is None or not caches_draft_model(speculative_config): + if not caches_draft_model(speculative_config): return None - draft_parallel_config = speculative_config.draft_parallel_config - target_tp = vllm_config.parallel_config.tensor_parallel_size - if draft_parallel_config.tensor_parallel_size != target_tp: - raise ValueError( - "The weight cache daemon requires the MTP draft and the target to " - f"use the same tensor parallel size, got " - f"{draft_parallel_config.tensor_parallel_size} != {target_tp}" + if speculative_config.moe_backend is not None: + vllm_config = replace( + vllm_config, + kernel_config=replace( + vllm_config.kernel_config, moe_backend=speculative_config.moe_backend + ), ) - _reject_unsupported_parallelism(draft_parallel_config) - return replace( - vllm_config, - model_config=speculative_config.draft_model_config, - parallel_config=draft_parallel_config, - ) + if speculative_config.attention_backend is not None: + vllm_config = replace( + vllm_config, + attention_config=replace( + vllm_config.attention_config, + backend=speculative_config.attention_backend, + ), + ) + if speculative_config.kv_cache_dtype is not None: + vllm_config = replace( + vllm_config, + cache_config=replace( + vllm_config.cache_config, cache_dtype=speculative_config.kv_cache_dtype + ), + ) + return vllm_config, speculative_config.draft_model_config def _reject_unsupported_parallelism(parallel_config: ParallelConfig) -> None: @@ -477,11 +500,14 @@ def main() -> None: distributed_init_method = get_distributed_init_method(master_addr, master_port) # (is_draft_model, draft_model_idx, config, rendezvous) per daemon group. - groups: list[tuple[bool, int | None, VllmConfig, str]] = [ - (False, None, vllm_config, distributed_init_method) + # (is_draft_model, draft_model_idx, vllm_config, model_config, rendezvous) + # per daemon group; model_config is None for the target. + groups: list[tuple[bool, int | None, VllmConfig, ModelConfig | None, str]] = [ + (False, None, vllm_config, None, distributed_init_method) ] - draft_vllm_config = get_cached_draft_vllm_config(vllm_config) - if draft_vllm_config is not None: + draft = get_draft_daemon_config(vllm_config) + if draft is not None: + draft_vllm_config, draft_model_config = draft draft_master_port = args.weight_cache_draft_master_port if draft_master_port is None: draft_master_port = master_port + 1 if nnodes > 1 else get_open_port() @@ -495,6 +521,7 @@ def main() -> None: True, 0, draft_vllm_config, + draft_model_config, get_distributed_init_method(master_addr, draft_master_port), ) ) @@ -508,7 +535,7 @@ def main() -> None: ] procs = [] expected_ready: set[tuple[str, int]] = set() - for is_draft_model, draft_model_idx, config, init_method in groups: + for is_draft_model, draft_model_idx, config, model_config, init_method in groups: role = "draft" if is_draft_model else "target" suffix = format_daemon_role(is_draft_model, draft_model_idx) for local_rank, global_rank in enumerate(global_ranks): @@ -525,6 +552,7 @@ def main() -> None: ready_queue, is_draft_model, draft_model_idx, + model_config, ), name=f"vllm-weight-cache-{role}{suffix}-{global_rank}", ) diff --git a/vllm/model_executor/model_loader/weight_cache/ipc_loader.py b/vllm/model_executor/model_loader/weight_cache/ipc_loader.py index 87d58cc7aa43..4daea1fd8f95 100644 --- a/vllm/model_executor/model_loader/weight_cache/ipc_loader.py +++ b/vllm/model_executor/model_loader/weight_cache/ipc_loader.py @@ -117,7 +117,7 @@ def load_weights(self, model: nn.Module, model_config: ModelConfig) -> None: loaded through this loader). """ device_index = torch.accelerator.current_device_index() - entries, _ = self._fetch_entries(model_config) + entries, _, _ = self._fetch_entries(model_config) params = dict(model.named_parameters()) buffers = dict(model.named_buffers()) for name, entry in entries.items(): @@ -137,10 +137,10 @@ def load_model( check_ipc_platform_support() state_fetched = False try: - entries, aliases = self._fetch_entries(model_config) + entries, aliases, attrs = self._fetch_entries(model_config) state_fetched = True return self._build_model( - vllm_config, model_config, prefix, entries, aliases + vllm_config, model_config, prefix, entries, aliases, attrs ) except (WeightCacheUnavailableError, CacheConfigMismatchError) as e: if not self.fallback: @@ -174,6 +174,7 @@ def _build_model( prefix: str, entries: dict[str, TensorEntry], aliases: dict[str, str], + attrs: dict[str, bool], ) -> nn.Module: device_config = vllm_config.device_config load_device = ( @@ -196,6 +197,10 @@ def _build_model( ) check_ipc_quant_support(model) self._apply_entries(model, entries, aliases, device_index) + # Flags that load_weights would have set (e.g. EAGLE ownership of + # embed_tokens / lm_head); the daemon ran it, this process did not. + for name, value in attrs.items(): + setattr(model, name, value) # The daemon exports tensors that already went through # process_weights_after_loading; re-run it in pre-processed mode # so quant methods only rebuild Python-side state (e.g. the MoE @@ -273,7 +278,7 @@ def _register(name: str, tensor: torch.Tensor, is_param: bool) -> None: def _fetch_entries( self, model_config: ModelConfig - ) -> tuple[dict[str, TensorEntry], dict[str, str]]: + ) -> tuple[dict[str, TensorEntry], dict[str, str], dict[str, bool]]: cache_config = WeightCacheKey.from_model_config( model_config, tp_size=get_tensor_model_parallel_world_size(), @@ -285,7 +290,7 @@ def _fetch_entries( def _request_state( self, cache_config: WeightCacheKey - ) -> tuple[dict[str, TensorEntry], dict[str, str]]: + ) -> tuple[dict[str, TensorEntry], dict[str, str], dict[str, bool]]: with self._connect(self.state_timeout_s) as conn: send_msg(conn, {"cmd": "get_state", "cache_config": cache_config}) response = recv_msg(conn) @@ -299,7 +304,11 @@ def _request_state( f"Weight cache daemon error: {response.get('message')}" ) self._check_gpu_uuid(response.get("gpu_uuid")) - return response["entries"], response.get("aliases", {}) + return ( + response["entries"], + response.get("aliases", {}), + response.get("attrs", {}), + ) def _connect(self, timeout: float) -> socket.socket: socket_path = self._resolve_socket_path() diff --git a/vllm/model_executor/model_loader/weight_cache/protocol.py b/vllm/model_executor/model_loader/weight_cache/protocol.py index 1a34fce58b87..7e8733d12499 100644 --- a/vllm/model_executor/model_loader/weight_cache/protocol.py +++ b/vllm/model_executor/model_loader/weight_cache/protocol.py @@ -20,14 +20,14 @@ import struct import tempfile from dataclasses import dataclass, fields -from typing import Any +from typing import Any, TypeGuard import torch from torch.multiprocessing.reductions import rebuild_cuda_tensor, reduce_tensor from transformers.utils import SAFE_WEIGHTS_INDEX_NAME import vllm.version -from vllm.config import ModelConfig +from vllm.config import ModelConfig, SpeculativeConfig from vllm.model_executor.layers.quantization.base_config import QuantizeMethodBase from vllm.model_executor.model_loader.weight_utils import ( filter_duplicate_safetensors_files, @@ -129,11 +129,25 @@ def get_socket_dir(socket_dir: str | None = None) -> str: # Speculative methods whose draft model the daemon caches in its own group. -# Other drafts (e.g. EAGLE3 heads) keep loading from disk in the engine. -WEIGHT_CACHE_DRAFT_METHODS = frozenset({"mtp"}) +# Other drafts keep loading from disk in the engine. +WEIGHT_CACHE_DRAFT_METHODS = frozenset({"mtp", "eagle", "eagle3"}) +# Python-side flags that weight loading sets on EAGLE-style drafts; the engine +# never runs load_weights for cached models, so the daemon ships them. +EXPORTED_MODEL_ATTRS = ("has_own_embed_tokens", "has_own_lm_head") -def caches_draft_model(speculative_config: Any) -> bool: + +def export_model_attrs(model: Any) -> dict[str, bool]: + return { + name: bool(getattr(model, name)) + for name in EXPORTED_MODEL_ATTRS + if hasattr(model, name) + } + + +def caches_draft_model( + speculative_config: SpeculativeConfig | None, +) -> TypeGuard[SpeculativeConfig]: """Whether the daemon serves the speculative draft as a separate role.""" return ( speculative_config is not None From 22dd56fb0ae1678f78107209bc12ea8f2641e550 Mon Sep 17 00:00:00 2001 From: liusy58 Date: Sun, 20 Sep 2026 08:20:02 +0000 Subject: [PATCH 03/12] fix Signed-off-by: liusy58 --- .../model_loader/test_weight_cache.py | 11 ++-- vllm/config/load.py | 6 +- vllm/model_executor/model_loader/utils.py | 8 +-- .../model_loader/weight_cache/daemon.py | 34 +++++------ .../model_loader/weight_cache/ipc_loader.py | 8 +-- .../model_loader/weight_cache/protocol.py | 57 +++---------------- .../model_loader/weight_cache/utils.py | 56 ++++++++++++++++++ 7 files changed, 89 insertions(+), 91 deletions(-) create mode 100644 vllm/model_executor/model_loader/weight_cache/utils.py diff --git a/tests/model_executor/model_loader/test_weight_cache.py b/tests/model_executor/model_loader/test_weight_cache.py index c36a673db07b..c2979cd16153 100644 --- a/tests/model_executor/model_loader/test_weight_cache.py +++ b/tests/model_executor/model_loader/test_weight_cache.py @@ -212,7 +212,7 @@ def test_ipc_cache_cold_start_and_warm_restart(vllm_runner, case: ModelCase): def test_weight_cache_caches_mtp_and_eagle_drafts(): from types import SimpleNamespace - from vllm.model_executor.model_loader.weight_cache.protocol import ( + from vllm.model_executor.model_loader.weight_cache.utils import ( caches_draft_model, ) @@ -235,7 +235,7 @@ def test_weight_cache_exports_eagle_ownership_flags(): decide whether to share the target's embed/lm_head travel with the state.""" import torch - from vllm.model_executor.model_loader.weight_cache.protocol import ( + from vllm.model_executor.model_loader.weight_cache.utils import ( export_model_attrs, ) @@ -249,7 +249,7 @@ def test_weight_cache_target_and_draft_use_distinct_sockets(tmp_path): from vllm.model_executor.model_loader.weight_cache.protocol import get_socket_path target_path = get_socket_path("GPU-abc", str(tmp_path)) - draft_path = get_socket_path("GPU-abc", str(tmp_path), is_draft_model=True) + draft_path = get_socket_path("GPU-abc", str(tmp_path), draft_model_idx=0) assert target_path != draft_path assert draft_path.endswith("GPU-abc_draft0.sock") @@ -281,17 +281,16 @@ def cfg(method, draft_load_config=None, load_config=ipc): mtp = get_draft_load_config(cfg("mtp")) assert mtp.load_format == "ipc_cache" - assert mtp.weight_cache_is_draft_model assert mtp.weight_cache_draft_model_idx == 0 assert mtp.model_loader_extra_config == {"fallback": False} eagle = get_draft_load_config(cfg("eagle3")) assert eagle.load_format == "ipc_cache" - assert eagle.weight_cache_is_draft_model + assert eagle.weight_cache_draft_model_idx == 0 dflash = get_draft_load_config(cfg("dflash")) assert dflash.load_format == "auto" - assert not dflash.weight_cache_is_draft_model + assert dflash.weight_cache_draft_model_idx is None assert dflash.model_loader_extra_config == {} assert get_draft_load_config(cfg("mtp", explicit)) is explicit diff --git a/vllm/config/load.py b/vllm/config/load.py index c6ced1c84224..29d5f87d9ba4 100644 --- a/vllm/config/load.py +++ b/vllm/config/load.py @@ -97,11 +97,9 @@ class LoadConfig: model_loader_extra_config: dict | TensorizerConfig = Field(default_factory=dict) """Extra config for model loader. This will be passed to the model loader corresponding to the chosen load_format.""" - weight_cache_is_draft_model: bool = False - """Whether an ``ipc_cache`` load targets the weight cache daemon group that - serves the speculative draft model instead of the target model.""" weight_cache_draft_model_idx: int | None = None - """Index of the speculative draft model within the draft cache role.""" + """Index of the speculative draft model whose weight cache daemon group an + ``ipc_cache`` load reads from. ``None`` selects the target model's group.""" device: str | None = None """Device to which model weights will be loaded, default to device_config.device""" diff --git a/vllm/model_executor/model_loader/utils.py b/vllm/model_executor/model_loader/utils.py index ab950be92eb3..849f5e126795 100644 --- a/vllm/model_executor/model_loader/utils.py +++ b/vllm/model_executor/model_loader/utils.py @@ -49,7 +49,7 @@ def get_draft_load_config(vllm_config: VllmConfig) -> LoadConfig: draft falls back to disk loading instead of being sent to the target daemon with a mismatching fingerprint. """ - from vllm.model_executor.model_loader.weight_cache.protocol import ( + from vllm.model_executor.model_loader.weight_cache.utils import ( caches_draft_model, ) @@ -61,11 +61,7 @@ def get_draft_load_config(vllm_config: VllmConfig) -> LoadConfig: if load_config.load_format != "ipc_cache": return load_config if caches_draft_model(speculative_config): - return replace( - load_config, - weight_cache_is_draft_model=True, - weight_cache_draft_model_idx=0, - ) + return replace(load_config, weight_cache_draft_model_idx=0) return replace(load_config, load_format="auto", model_loader_extra_config={}) diff --git a/vllm/model_executor/model_loader/weight_cache/daemon.py b/vllm/model_executor/model_loader/weight_cache/daemon.py index d6f41b4c3512..44d92d531977 100644 --- a/vllm/model_executor/model_loader/weight_cache/daemon.py +++ b/vllm/model_executor/model_loader/weight_cache/daemon.py @@ -81,18 +81,20 @@ TensorEntry, WeightCacheKey, WeightCacheUnavailableError, - caches_draft_model, check_ipc_platform_support, check_ipc_quant_support, ensure_private_socket_dir, - export_model_attrs, - format_daemon_role, get_current_device_uuid, get_socket_path, recv_msg, send_msg, verify_peer_is_owner, ) +from vllm.model_executor.model_loader.weight_cache.utils import ( + caches_draft_model, + export_model_attrs, + format_daemon_role, +) from vllm.platforms import current_platform from vllm.utils.argparse_utils import FlexibleArgumentParser from vllm.utils.network_utils import get_distributed_init_method, get_open_port @@ -183,7 +185,6 @@ def __init__( local_rank: int, distributed_init_method: str, socket_dir: str | None = None, - is_draft_model: bool = False, draft_model_idx: int | None = None, model_config: ModelConfig | None = None, ): @@ -195,9 +196,8 @@ def __init__( self.local_rank = local_rank self.distributed_init_method = distributed_init_method self.socket_dir = socket_dir - self.is_draft_model = is_draft_model self.draft_model_idx = draft_model_idx - self.role = "draft" if is_draft_model else "target" + self.role = format_daemon_role(draft_model_idx) self.model: torch.nn.Module | None = None # Fingerprint before loading: process_weights_after_loading may # mutate hf_config.quantization_config. @@ -205,7 +205,6 @@ def __init__( self.model_config, tp_size=vllm_config.parallel_config.tensor_parallel_size, tp_rank=tp_rank, - is_draft_model=is_draft_model, draft_model_idx=draft_model_idx, ) @@ -307,7 +306,6 @@ def _socket_path(self) -> str: return get_socket_path( get_current_device_uuid(), self.socket_dir, - is_draft_model=self.is_draft_model, draft_model_idx=self.draft_model_idx, ) @@ -372,7 +370,6 @@ def _run_daemon( distributed_init_method: str, socket_dir: str | None, ready_queue: "multiprocessing.Queue[tuple[str, int]]", - is_draft_model: bool = False, draft_model_idx: int | None = None, model_config: ModelConfig | None = None, ) -> None: @@ -382,7 +379,6 @@ def _run_daemon( local_rank, distributed_init_method, socket_dir, - is_draft_model, draft_model_idx, model_config, ) @@ -499,11 +495,10 @@ def main() -> None: master_port = args.weight_cache_master_port or get_open_port() distributed_init_method = get_distributed_init_method(master_addr, master_port) - # (is_draft_model, draft_model_idx, config, rendezvous) per daemon group. - # (is_draft_model, draft_model_idx, vllm_config, model_config, rendezvous) - # per daemon group; model_config is None for the target. - groups: list[tuple[bool, int | None, VllmConfig, ModelConfig | None, str]] = [ - (False, None, vllm_config, None, distributed_init_method) + # (draft_model_idx, vllm_config, model_config, rendezvous) per daemon + # group; the target group has no draft index and no separate model_config. + groups: list[tuple[int | None, VllmConfig, ModelConfig | None, str]] = [ + (None, vllm_config, None, distributed_init_method) ] draft = get_draft_daemon_config(vllm_config) if draft is not None: @@ -518,7 +513,6 @@ def main() -> None: ) groups.append( ( - True, 0, draft_vllm_config, draft_model_config, @@ -535,9 +529,8 @@ def main() -> None: ] procs = [] expected_ready: set[tuple[str, int]] = set() - for is_draft_model, draft_model_idx, config, model_config, init_method in groups: - role = "draft" if is_draft_model else "target" - suffix = format_daemon_role(is_draft_model, draft_model_idx) + for draft_model_idx, config, model_config, init_method in groups: + role = format_daemon_role(draft_model_idx) for local_rank, global_rank in enumerate(global_ranks): expected_ready.add((role, global_rank)) procs.append( @@ -550,11 +543,10 @@ def main() -> None: init_method, args.weight_cache_socket_dir, ready_queue, - is_draft_model, draft_model_idx, model_config, ), - name=f"vllm-weight-cache-{role}{suffix}-{global_rank}", + name=f"vllm-weight-cache-{role}-{global_rank}", ) ) for proc in procs: diff --git a/vllm/model_executor/model_loader/weight_cache/ipc_loader.py b/vllm/model_executor/model_loader/weight_cache/ipc_loader.py index 4daea1fd8f95..e7080227e5c6 100644 --- a/vllm/model_executor/model_loader/weight_cache/ipc_loader.py +++ b/vllm/model_executor/model_loader/weight_cache/ipc_loader.py @@ -61,7 +61,7 @@ class IpcModelLoader(BaseModelLoader): - socket_path: explicit daemon socket path. Defaults to a per-GPU path derived from the physical GPU uuid and the cache role - (``load_config.weight_cache_is_draft_model``). + (``load_config.weight_cache_draft_model_idx``). - socket_dir: directory containing the daemon sockets. - mode: "zero_copy" (default) or "copy". - fallback: fall back to disk loading when the daemon is unavailable or @@ -79,9 +79,8 @@ def __init__(self, load_config: LoadConfig): extra_config = copy(load_config.model_loader_extra_config or {}) self.socket_path: str | None = extra_config.pop("socket_path", None) self.socket_dir: str | None = extra_config.pop("socket_dir", None) - self.is_draft_model = load_config.weight_cache_is_draft_model self.draft_model_idx = load_config.weight_cache_draft_model_idx - if self.is_draft_model and self.socket_path is not None: + if self.draft_model_idx is not None and self.socket_path is not None: raise ValueError( "socket_path cannot be combined with the draft weight cache role; " "use socket_dir so the target and draft sockets are derived " @@ -283,7 +282,6 @@ def _fetch_entries( model_config, tp_size=get_tensor_model_parallel_world_size(), tp_rank=get_tensor_model_parallel_rank(), - is_draft_model=self.is_draft_model, draft_model_idx=self.draft_model_idx, ) return self._request_state(cache_config) @@ -339,7 +337,6 @@ def _resolve_socket_path(self) -> str: return get_socket_path( get_current_device_uuid(), self.socket_dir, - is_draft_model=self.is_draft_model, draft_model_idx=self.draft_model_idx, ) @@ -368,7 +365,6 @@ def _fallback_load_config(self) -> LoadConfig: self.load_config, load_format="auto", model_loader_extra_config={}, - weight_cache_is_draft_model=False, weight_cache_draft_model_idx=None, ) diff --git a/vllm/model_executor/model_loader/weight_cache/protocol.py b/vllm/model_executor/model_loader/weight_cache/protocol.py index 7e8733d12499..9a1fa412a2eb 100644 --- a/vllm/model_executor/model_loader/weight_cache/protocol.py +++ b/vllm/model_executor/model_loader/weight_cache/protocol.py @@ -20,15 +20,19 @@ import struct import tempfile from dataclasses import dataclass, fields -from typing import Any, TypeGuard +from typing import Any import torch from torch.multiprocessing.reductions import rebuild_cuda_tensor, reduce_tensor from transformers.utils import SAFE_WEIGHTS_INDEX_NAME import vllm.version -from vllm.config import ModelConfig, SpeculativeConfig +from vllm.config import ModelConfig from vllm.model_executor.layers.quantization.base_config import QuantizeMethodBase +from vllm.model_executor.model_loader.weight_cache.utils import ( + format_socket_role_suffix, + normalize_draft_model_idx, +) from vllm.model_executor.model_loader.weight_utils import ( filter_duplicate_safetensors_files, ) @@ -128,60 +132,19 @@ def get_socket_dir(socket_dir: str | None = None) -> str: ) -# Speculative methods whose draft model the daemon caches in its own group. -# Other drafts keep loading from disk in the engine. -WEIGHT_CACHE_DRAFT_METHODS = frozenset({"mtp", "eagle", "eagle3"}) - -# Python-side flags that weight loading sets on EAGLE-style drafts; the engine -# never runs load_weights for cached models, so the daemon ships them. -EXPORTED_MODEL_ATTRS = ("has_own_embed_tokens", "has_own_lm_head") - - -def export_model_attrs(model: Any) -> dict[str, bool]: - return { - name: bool(getattr(model, name)) - for name in EXPORTED_MODEL_ATTRS - if hasattr(model, name) - } - - -def caches_draft_model( - speculative_config: SpeculativeConfig | None, -) -> TypeGuard[SpeculativeConfig]: - """Whether the daemon serves the speculative draft as a separate role.""" - return ( - speculative_config is not None - and speculative_config.method in WEIGHT_CACHE_DRAFT_METHODS - and speculative_config.draft_model_config is not None - ) - - -def normalize_draft_model_idx(draft_model_idx: int | None) -> int: - return -1 if draft_model_idx is None else draft_model_idx - - -def format_daemon_role( - is_draft_model: bool = False, draft_model_idx: int | None = None -) -> str: - """Socket-name suffix distinguishing the draft daemon group from the target.""" - if not is_draft_model: - return "" - return f"_draft{draft_model_idx if draft_model_idx is not None else 0}" - - def get_socket_path( gpu_uuid: str, socket_dir: str | None = None, *, - is_draft_model: bool = False, draft_model_idx: int | None = None, ) -> str: + """Socket path of a daemon group; ``draft_model_idx`` None is the target.""" directory = get_socket_dir(socket_dir) return os.path.join( directory, SOCKET_NAME_TEMPLATE.format( gpu_uuid=gpu_uuid, - role=format_daemon_role(is_draft_model, draft_model_idx), + role=format_socket_role_suffix(draft_model_idx), ), ) @@ -325,8 +288,8 @@ class WeightCacheKey: quant_config_hash: str revision: str | None vllm_version: str - is_draft_model: bool = False draft_model_idx: int = -1 + """Daemon group the weights come from; -1 is the target model.""" @classmethod def from_model_config( @@ -335,7 +298,6 @@ def from_model_config( tp_size: int, tp_rank: int, *, - is_draft_model: bool = False, draft_model_idx: int | None = None, ) -> "WeightCacheKey": """Build the fingerprint for a model configuration. @@ -363,7 +325,6 @@ def from_model_config( quant_config_hash=_hash_quant_config(quant_config), revision=model_config.revision, vllm_version=vllm.version.__version__, - is_draft_model=is_draft_model, draft_model_idx=normalize_draft_model_idx(draft_model_idx), ) diff --git a/vllm/model_executor/model_loader/weight_cache/utils.py b/vllm/model_executor/model_loader/weight_cache/utils.py new file mode 100644 index 000000000000..7b944494952b --- /dev/null +++ b/vllm/model_executor/model_loader/weight_cache/utils.py @@ -0,0 +1,56 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Helpers shared by the weight cache daemon, the IPC loader and the engine. + +These live outside ``protocol`` so callers that only need to know whether a +draft is cached, or how a daemon group is named, do not have to import the +wire format. +""" + +from typing import Any, TypeGuard + +from vllm.config import SpeculativeConfig + +# Speculative methods whose draft model the daemon caches in its own group. +# Other drafts keep loading from disk in the engine. +WEIGHT_CACHE_DRAFT_METHODS = frozenset({"mtp", "eagle", "eagle3"}) + +# Python-side flags that weight loading sets on EAGLE-style drafts; the engine +# never runs load_weights for cached models, so the daemon ships them. +EXPORTED_MODEL_ATTRS = ("has_own_embed_tokens", "has_own_lm_head") + + +def export_model_attrs(model: Any) -> dict[str, bool]: + return { + name: bool(getattr(model, name)) + for name in EXPORTED_MODEL_ATTRS + if hasattr(model, name) + } + + +def caches_draft_model( + speculative_config: SpeculativeConfig | None, +) -> TypeGuard[SpeculativeConfig]: + """Whether the daemon serves the speculative draft as a separate role.""" + return ( + speculative_config is not None + and speculative_config.method in WEIGHT_CACHE_DRAFT_METHODS + and speculative_config.draft_model_config is not None + ) + + +def normalize_draft_model_idx(draft_model_idx: int | None) -> int: + """Encode a daemon group as a cache key field; -1 is the target model.""" + return -1 if draft_model_idx is None else draft_model_idx + + +def format_daemon_role(draft_model_idx: int | None) -> str: + """Name of a daemon group: the target model, or a draft by index.""" + return "target" if draft_model_idx is None else f"draft{draft_model_idx}" + + +def format_socket_role_suffix(draft_model_idx: int | None) -> str: + """Socket-name suffix keeping each draft group distinct from the target.""" + if draft_model_idx is None: + return "" + return f"_{format_daemon_role(draft_model_idx)}" From 1b1fbbc3ac0d6b4c2432881817e1d1a496301c5f Mon Sep 17 00:00:00 2001 From: Isotr0py Date: Mon, 21 Sep 2026 07:50:58 +0000 Subject: [PATCH 04/12] use draft_model_idx=None for target model Signed-off-by: Isotr0py --- vllm/model_executor/model_loader/weight_cache/protocol.py | 7 +++---- vllm/model_executor/model_loader/weight_cache/utils.py | 5 ----- 2 files changed, 3 insertions(+), 9 deletions(-) diff --git a/vllm/model_executor/model_loader/weight_cache/protocol.py b/vllm/model_executor/model_loader/weight_cache/protocol.py index 9a1fa412a2eb..75611ee67802 100644 --- a/vllm/model_executor/model_loader/weight_cache/protocol.py +++ b/vllm/model_executor/model_loader/weight_cache/protocol.py @@ -31,7 +31,6 @@ from vllm.model_executor.layers.quantization.base_config import QuantizeMethodBase from vllm.model_executor.model_loader.weight_cache.utils import ( format_socket_role_suffix, - normalize_draft_model_idx, ) from vllm.model_executor.model_loader.weight_utils import ( filter_duplicate_safetensors_files, @@ -288,8 +287,8 @@ class WeightCacheKey: quant_config_hash: str revision: str | None vllm_version: str - draft_model_idx: int = -1 - """Daemon group the weights come from; -1 is the target model.""" + draft_model_idx: int | None = None + """Daemon group the weights come from; None is the target model.""" @classmethod def from_model_config( @@ -325,7 +324,7 @@ def from_model_config( quant_config_hash=_hash_quant_config(quant_config), revision=model_config.revision, vllm_version=vllm.version.__version__, - draft_model_idx=normalize_draft_model_idx(draft_model_idx), + draft_model_idx=draft_model_idx, ) def mismatched_fields(self, other: "WeightCacheKey") -> list[str]: diff --git a/vllm/model_executor/model_loader/weight_cache/utils.py b/vllm/model_executor/model_loader/weight_cache/utils.py index 7b944494952b..33d68c1a4fa0 100644 --- a/vllm/model_executor/model_loader/weight_cache/utils.py +++ b/vllm/model_executor/model_loader/weight_cache/utils.py @@ -39,11 +39,6 @@ def caches_draft_model( ) -def normalize_draft_model_idx(draft_model_idx: int | None) -> int: - """Encode a daemon group as a cache key field; -1 is the target model.""" - return -1 if draft_model_idx is None else draft_model_idx - - def format_daemon_role(draft_model_idx: int | None) -> str: """Name of a daemon group: the target model, or a draft by index.""" return "target" if draft_model_idx is None else f"draft{draft_model_idx}" From 42a33abaa13c649dc8bcd4e1bffa7db6d98727c4 Mon Sep 17 00:00:00 2001 From: Isotr0py Date: Mon, 21 Sep 2026 12:00:11 +0000 Subject: [PATCH 05/12] cleanup and test with mtp e2e Signed-off-by: Isotr0py --- .../model_loader/test_weight_cache.py | 140 +++++++----------- vllm/config/speculative.py | 27 +++- .../model_loader/weight_cache/daemon.py | 31 +--- .../model_loader/weight_cache/utils.py | 3 +- vllm/v1/worker/gpu/spec_decode/eagle/utils.py | 29 +--- 5 files changed, 90 insertions(+), 140 deletions(-) diff --git a/tests/model_executor/model_loader/test_weight_cache.py b/tests/model_executor/model_loader/test_weight_cache.py index c2979cd16153..6b8c1a2a93bd 100644 --- a/tests/model_executor/model_loader/test_weight_cache.py +++ b/tests/model_executor/model_loader/test_weight_cache.py @@ -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", @@ -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( @@ -172,6 +181,27 @@ 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}', + ], +) + @pytest.mark.parametrize("case", [QWEN_CASE, K3_CASE], ids=["qwen3.5", "kimi-k3"]) def test_ipc_cache_cold_start_and_warm_restart(vllm_runner, case: ModelCase): @@ -209,90 +239,32 @@ def test_ipc_cache_cold_start_and_warm_restart(vllm_runner, case: ModelCase): assert restart_outputs == baseline_outputs -def test_weight_cache_caches_mtp_and_eagle_drafts(): - from types import SimpleNamespace - - from vllm.model_executor.model_loader.weight_cache.utils import ( - caches_draft_model, - ) - - draft = object() - for method in ("mtp", "eagle", "eagle3"): - assert caches_draft_model( - SimpleNamespace(method=method, draft_model_config=draft) - ) - assert not caches_draft_model( - SimpleNamespace(method="dflash", draft_model_config=draft) - ) - assert not caches_draft_model( - SimpleNamespace(method="mtp", draft_model_config=None) - ) - assert not caches_draft_model(None) - - -def test_weight_cache_exports_eagle_ownership_flags(): - """The engine never runs load_weights for cached drafts, so the flags that - decide whether to share the target's embed/lm_head travel with the state.""" - import torch - - from vllm.model_executor.model_loader.weight_cache.utils import ( - export_model_attrs, - ) +def test_ipc_cache_caches_mtp_draft(vllm_runner): + """The daemon caches the MTP draft in its draft group alongside the target. - model = torch.nn.Module() - assert export_model_attrs(model) == {} - model.has_own_lm_head = True - assert export_model_attrs(model) == {"has_own_lm_head": True} - - -def test_weight_cache_target_and_draft_use_distinct_sockets(tmp_path): - from vllm.model_executor.model_loader.weight_cache.protocol import get_socket_path - - target_path = get_socket_path("GPU-abc", str(tmp_path)) - draft_path = get_socket_path("GPU-abc", str(tmp_path), draft_model_idx=0) - - assert target_path != draft_path - assert draft_path.endswith("GPU-abc_draft0.sock") - - -def test_draft_load_config_under_ipc_cache(): - """A cached draft is routed to the daemon's draft group; any other draft - falls back to disk instead of hitting the target daemon; an explicit - draft_load_config always wins.""" - from types import SimpleNamespace - - from vllm.config import LoadConfig - from vllm.model_executor.model_loader.utils import get_draft_load_config - - ipc = LoadConfig( - load_format="ipc_cache", model_loader_extra_config={"fallback": False} - ) - explicit = LoadConfig(load_format="fastsafetensors") - - def cfg(method, draft_load_config=None, load_config=ipc): - return SimpleNamespace( - load_config=load_config, - speculative_config=SimpleNamespace( - method=method, - draft_model_config=object(), - draft_load_config=draft_load_config, - ), - ) + The warm runs use fallback=False, so both the target and the draft fail + hard unless their daemon groups served the weights; matching the + disk-loaded baseline proves the cached draft produces identical drafts. + """ + if not current_platform.is_cuda_alike(): + pytest.skip("Weight cache IPC sharing requires CUDA or ROCm") - mtp = get_draft_load_config(cfg("mtp")) - assert mtp.load_format == "ipc_cache" - assert mtp.weight_cache_draft_model_idx == 0 - assert mtp.model_loader_extra_config == {"fallback": False} + case = QWEN_MTP_CASE + # Baseline: target and MTP draft both loaded from disk. + baseline_outputs = generate(vllm_runner, case, None, fallback=True) + assert all(text for _, texts in baseline_outputs for text in texts) - eagle = get_draft_load_config(cfg("eagle3")) - assert eagle.load_format == "ipc_cache" - assert eagle.weight_cache_draft_model_idx == 0 + # Cold start: no daemon is serving, so both models fall back to disk. + with tempfile.TemporaryDirectory(prefix="vllm_ipc_empty_") as empty_socket_dir: + cold_outputs = generate(vllm_runner, case, empty_socket_dir, fallback=True) - dflash = get_draft_load_config(cfg("dflash")) - assert dflash.load_format == "auto" - assert dflash.weight_cache_draft_model_idx is None - assert dflash.model_loader_extra_config == {} + with WeightCacheDaemon( + case.model, tp_size=1, extra_args=case.daemon_args, num_groups=2 + ) 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) - assert get_draft_load_config(cfg("mtp", explicit)) is explicit - disk = LoadConfig(load_format="fastsafetensors") - assert get_draft_load_config(cfg("mtp", load_config=disk)) is disk + assert cold_outputs == baseline_outputs + assert warm_outputs == baseline_outputs + assert restart_outputs == baseline_outputs diff --git a/vllm/config/speculative.py b/vllm/config/speculative.py index ef9a16f903a5..eb9d54813703 100644 --- a/vllm/config/speculative.py +++ b/vllm/config/speculative.py @@ -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 @@ -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" @@ -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.""" @@ -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: diff --git a/vllm/model_executor/model_loader/weight_cache/daemon.py b/vllm/model_executor/model_loader/weight_cache/daemon.py index 44d92d531977..2639be100dbc 100644 --- a/vllm/model_executor/model_loader/weight_cache/daemon.py +++ b/vllm/model_executor/model_loader/weight_cache/daemon.py @@ -66,7 +66,6 @@ ModelConfig, ParallelConfig, VllmConfig, - replace, set_current_vllm_config, ) from vllm.distributed import ( @@ -158,8 +157,7 @@ def get_daemon_model( always fails the check, so load_model's finalize step for it is unnecessary here. """ - if model_config is None: - model_config = vllm_config.model_config + model_config = model_config or vllm_config.model_config load_config = vllm_config.load_config loader = get_model_loader(load_config) device_config = vllm_config.device_config @@ -399,29 +397,10 @@ def get_draft_daemon_config( speculative_config = vllm_config.speculative_config if not caches_draft_model(speculative_config): return None - if speculative_config.moe_backend is not None: - vllm_config = replace( - vllm_config, - kernel_config=replace( - vllm_config.kernel_config, moe_backend=speculative_config.moe_backend - ), - ) - if speculative_config.attention_backend is not None: - vllm_config = replace( - vllm_config, - attention_config=replace( - vllm_config.attention_config, - backend=speculative_config.attention_backend, - ), - ) - if speculative_config.kv_cache_dtype is not None: - vllm_config = replace( - vllm_config, - cache_config=replace( - vllm_config.cache_config, cache_dtype=speculative_config.kv_cache_dtype - ), - ) - return vllm_config, speculative_config.draft_model_config + return ( + speculative_config.apply_draft_overrides(vllm_config), + speculative_config.draft_model_config, + ) def _reject_unsupported_parallelism(parallel_config: ParallelConfig) -> None: diff --git a/vllm/model_executor/model_loader/weight_cache/utils.py b/vllm/model_executor/model_loader/weight_cache/utils.py index 33d68c1a4fa0..e9cb9a7ce2c7 100644 --- a/vllm/model_executor/model_loader/weight_cache/utils.py +++ b/vllm/model_executor/model_loader/weight_cache/utils.py @@ -10,6 +10,7 @@ from typing import Any, TypeGuard from vllm.config import SpeculativeConfig +from vllm.model_executor.models.interfaces import SupportsEagleBase # Speculative methods whose draft model the daemon caches in its own group. # Other drafts keep loading from disk in the engine. @@ -17,7 +18,7 @@ # Python-side flags that weight loading sets on EAGLE-style drafts; the engine # never runs load_weights for cached models, so the daemon ships them. -EXPORTED_MODEL_ATTRS = ("has_own_embed_tokens", "has_own_lm_head") +EXPORTED_MODEL_ATTRS = tuple(SupportsEagleBase.__annotations__) def export_model_attrs(model: Any) -> dict[str, bool]: diff --git a/vllm/v1/worker/gpu/spec_decode/eagle/utils.py b/vllm/v1/worker/gpu/spec_decode/eagle/utils.py index 54b0b7776a98..06bed9391e03 100644 --- a/vllm/v1/worker/gpu/spec_decode/eagle/utils.py +++ b/vllm/v1/worker/gpu/spec_decode/eagle/utils.py @@ -76,34 +76,7 @@ def load_eagle_model(target_model: nn.Module, vllm_config: VllmConfig) -> nn.Mod speculative_config = vllm_config.speculative_config assert speculative_config is not None draft_model_config = speculative_config.draft_model_config - if speculative_config.moe_backend is not None: - # Otherwise the draft inherits the target's --moe-backend, which - # fails when the draft is unquantized and that backend is not. - vllm_config = replace( - vllm_config, - kernel_config=replace( - vllm_config.kernel_config, - moe_backend=speculative_config.moe_backend, - ), - ) - if speculative_config.kv_cache_dtype is not None: - vllm_config = replace( - vllm_config, - cache_config=replace( - vllm_config.cache_config, - cache_dtype=speculative_config.kv_cache_dtype, - ), - ) - if speculative_config.attention_backend is not None: - # Before get_model(): the backend is read off the constructed layers. - # Only when set, so the draft keeps a KV cache layout the target shares. - vllm_config = replace( - vllm_config, - attention_config=replace( - vllm_config.attention_config, - backend=speculative_config.attention_backend, - ), - ) + vllm_config = speculative_config.apply_draft_overrides(vllm_config) draft_load_config = get_pp_safe_draft_load_config( get_draft_load_config(vllm_config) ) From 88424f4fd46c9ccb8d24af5e684dba6f3171dbc2 Mon Sep 17 00:00:00 2001 From: Isotr0py Date: Mon, 21 Sep 2026 13:40:39 +0000 Subject: [PATCH 06/12] further clean Signed-off-by: Isotr0py --- vllm/config/load.py | 3 -- vllm/model_executor/model_loader/utils.py | 43 +++++++++++-------- .../model_loader/weight_cache/__init__.py | 23 ---------- .../model_loader/weight_cache/daemon.py | 37 ++++++++-------- .../model_loader/weight_cache/ipc_loader.py | 27 +++++++++--- .../model_loader/weight_cache/protocol.py | 14 +++--- .../model_loader/weight_cache/utils.py | 20 ++++----- 7 files changed, 78 insertions(+), 89 deletions(-) diff --git a/vllm/config/load.py b/vllm/config/load.py index 29d5f87d9ba4..ee13e2fccd7c 100644 --- a/vllm/config/load.py +++ b/vllm/config/load.py @@ -97,9 +97,6 @@ class LoadConfig: model_loader_extra_config: dict | TensorizerConfig = Field(default_factory=dict) """Extra config for model loader. This will be passed to the model loader corresponding to the chosen load_format.""" - weight_cache_draft_model_idx: int | None = None - """Index of the speculative draft model whose weight cache daemon group an - ``ipc_cache`` load reads from. ``None`` selects the target model's group.""" device: str | None = None """Device to which model weights will be loaded, default to device_config.device""" diff --git a/vllm/model_executor/model_loader/utils.py b/vllm/model_executor/model_loader/utils.py index 849f5e126795..465703f75983 100644 --- a/vllm/model_executor/model_loader/utils.py +++ b/vllm/model_executor/model_loader/utils.py @@ -29,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 @@ -41,28 +44,30 @@ def get_draft_load_config(vllm_config: VllmConfig) -> LoadConfig: - """Load config for the speculative draft model. - - An explicit ``draft_load_config`` always wins. Otherwise the draft inherits - the target's load config, except under ``ipc_cache``: a cached draft - (MTP, EAGLE, EAGLE3) is routed to the daemon's draft group, and any other - draft falls back to disk loading instead of being sent to the target - daemon with a mismatching fingerprint. - """ - from vllm.model_executor.model_loader.weight_cache.utils import ( - caches_draft_model, - ) - + """Get load config for the speculative draft model.""" speculative_config = vllm_config.speculative_config - assert speculative_config is not None - load_config = vllm_config.load_config - if speculative_config.draft_load_config is not None: + if ( + speculative_config is not None + and speculative_config.draft_load_config is not None + ): return speculative_config.draft_load_config - if load_config.load_format != "ipc_cache": + load_config = vllm_config.load_config + if load_config is not None and load_config.load_format != "ipc_cache": return load_config - if caches_draft_model(speculative_config): - return replace(load_config, weight_cache_draft_model_idx=0) - return replace(load_config, load_format="auto", model_loader_extra_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") diff --git a/vllm/model_executor/model_loader/weight_cache/__init__.py b/vllm/model_executor/model_loader/weight_cache/__init__.py index f9dbe9bb32a1..6655f8913623 100644 --- a/vllm/model_executor/model_loader/weight_cache/__init__.py +++ b/vllm/model_executor/model_loader/weight_cache/__init__.py @@ -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", -] diff --git a/vllm/model_executor/model_loader/weight_cache/daemon.py b/vllm/model_executor/model_loader/weight_cache/daemon.py index 2639be100dbc..56e950c69f18 100644 --- a/vllm/model_executor/model_loader/weight_cache/daemon.py +++ b/vllm/model_executor/model_loader/weight_cache/daemon.py @@ -44,7 +44,7 @@ With MTP, EAGLE or EAGLE3 speculative decoding the launcher additionally starts a draft daemon group that caches the draft model. It uses its own cache -key, Unix sockets (``*_draft0.sock``) and rendezvous port +key, Unix sockets (``*_draft.sock``) and rendezvous port (``--weight-cache-draft-master-port``, default ``--weight-cache-master-port + 1``), so each process serves exactly one model role. Other draft types are not cached and keep loading from disk in the engine. @@ -90,9 +90,9 @@ verify_peer_is_owner, ) from vllm.model_executor.model_loader.weight_cache.utils import ( - caches_draft_model, export_model_attrs, format_daemon_role, + is_draft_model_cacheable, ) from vllm.platforms import current_platform from vllm.utils.argparse_utils import FlexibleArgumentParser @@ -183,7 +183,7 @@ def __init__( local_rank: int, distributed_init_method: str, socket_dir: str | None = None, - draft_model_idx: int | None = None, + is_draft: bool = False, model_config: ModelConfig | None = None, ): self.vllm_config = vllm_config @@ -194,8 +194,8 @@ def __init__( self.local_rank = local_rank self.distributed_init_method = distributed_init_method self.socket_dir = socket_dir - self.draft_model_idx = draft_model_idx - self.role = format_daemon_role(draft_model_idx) + self.is_draft = is_draft + self.role = format_daemon_role(is_draft) self.model: torch.nn.Module | None = None # Fingerprint before loading: process_weights_after_loading may # mutate hf_config.quantization_config. @@ -203,7 +203,7 @@ def __init__( self.model_config, tp_size=vllm_config.parallel_config.tensor_parallel_size, tp_rank=tp_rank, - draft_model_idx=draft_model_idx, + is_draft=is_draft, ) def load_model(self) -> None: @@ -304,7 +304,7 @@ def _socket_path(self) -> str: return get_socket_path( get_current_device_uuid(), self.socket_dir, - draft_model_idx=self.draft_model_idx, + is_draft=self.is_draft, ) def _handle_connection(self, conn: socket.socket) -> None: @@ -368,7 +368,7 @@ def _run_daemon( distributed_init_method: str, socket_dir: str | None, ready_queue: "multiprocessing.Queue[tuple[str, int]]", - draft_model_idx: int | None = None, + is_draft: bool = False, model_config: ModelConfig | None = None, ) -> None: daemon = WeightCacheDaemon( @@ -377,7 +377,7 @@ def _run_daemon( local_rank, distributed_init_method, socket_dir, - draft_model_idx, + is_draft, model_config, ) daemon.load_model() @@ -395,8 +395,9 @@ def get_draft_daemon_config( ``vllm_config.model_config``. """ speculative_config = vllm_config.speculative_config - if not caches_draft_model(speculative_config): + if not is_draft_model_cacheable(speculative_config): return None + assert speculative_config is not None return ( speculative_config.apply_draft_overrides(vllm_config), speculative_config.draft_model_config, @@ -474,10 +475,10 @@ def main() -> None: master_port = args.weight_cache_master_port or get_open_port() distributed_init_method = get_distributed_init_method(master_addr, master_port) - # (draft_model_idx, vllm_config, model_config, rendezvous) per daemon - # group; the target group has no draft index and no separate model_config. - groups: list[tuple[int | None, VllmConfig, ModelConfig | None, str]] = [ - (None, vllm_config, None, distributed_init_method) + # (is_draft, vllm_config, model_config, rendezvous) per daemon group; the + # target group is not a draft and has no separate model_config. + groups: list[tuple[bool, VllmConfig, ModelConfig | None, str]] = [ + (False, vllm_config, None, distributed_init_method) ] draft = get_draft_daemon_config(vllm_config) if draft is not None: @@ -492,7 +493,7 @@ def main() -> None: ) groups.append( ( - 0, + True, draft_vllm_config, draft_model_config, get_distributed_init_method(master_addr, draft_master_port), @@ -508,8 +509,8 @@ def main() -> None: ] procs = [] expected_ready: set[tuple[str, int]] = set() - for draft_model_idx, config, model_config, init_method in groups: - role = format_daemon_role(draft_model_idx) + for is_draft, config, model_config, init_method in groups: + role = format_daemon_role(is_draft) for local_rank, global_rank in enumerate(global_ranks): expected_ready.add((role, global_rank)) procs.append( @@ -522,7 +523,7 @@ def main() -> None: init_method, args.weight_cache_socket_dir, ready_queue, - draft_model_idx, + is_draft, model_config, ), name=f"vllm-weight-cache-{role}-{global_rank}", diff --git a/vllm/model_executor/model_loader/weight_cache/ipc_loader.py b/vllm/model_executor/model_loader/weight_cache/ipc_loader.py index e7080227e5c6..ec028794913b 100644 --- a/vllm/model_executor/model_loader/weight_cache/ipc_loader.py +++ b/vllm/model_executor/model_loader/weight_cache/ipc_loader.py @@ -60,8 +60,7 @@ class IpcModelLoader(BaseModelLoader): Extra config keys (via --model-loader-extra-config): - socket_path: explicit daemon socket path. Defaults to a per-GPU path - derived from the physical GPU uuid and the cache role - (``load_config.weight_cache_draft_model_idx``). + derived from the physical GPU uuid and the cache role (target/draft). - socket_dir: directory containing the daemon sockets. - mode: "zero_copy" (default) or "copy". - fallback: fall back to disk loading when the daemon is unavailable or @@ -79,8 +78,10 @@ def __init__(self, load_config: LoadConfig): extra_config = copy(load_config.model_loader_extra_config or {}) self.socket_path: str | None = extra_config.pop("socket_path", None) self.socket_dir: str | None = extra_config.pop("socket_dir", None) - self.draft_model_idx = load_config.weight_cache_draft_model_idx - if self.draft_model_idx is not None and self.socket_path is not None: + # Internal: set by the engine when routing a speculative draft to the + # daemon's draft group. + self.is_draft = bool(extra_config.pop("is_draft", False)) + if self.is_draft and self.socket_path is not None: raise ValueError( "socket_path cannot be combined with the draft weight cache role; " "use socket_dir so the target and draft sockets are derived " @@ -136,6 +137,19 @@ def load_model( check_ipc_platform_support() state_fetched = False try: + # Cross-check the routing flag against the identity of the model + # being loaded: a draft load that lost its flag (or a target load + # that got one) would hit the wrong daemon group and + # fingerprint-mismatch. + spec = vllm_config.speculative_config + inferred = spec is not None and model_config is spec.draft_model_config + if inferred != self.is_draft: + raise CacheConfigMismatchError( + f"Weight cache role mismatch: loading " + f"{'draft' if inferred else 'target'} model but the loader " + f"was configured for the " + f"{'draft' if self.is_draft else 'target'} group" + ) entries, aliases, attrs = self._fetch_entries(model_config) state_fetched = True return self._build_model( @@ -282,7 +296,7 @@ def _fetch_entries( model_config, tp_size=get_tensor_model_parallel_world_size(), tp_rank=get_tensor_model_parallel_rank(), - draft_model_idx=self.draft_model_idx, + is_draft=self.is_draft, ) return self._request_state(cache_config) @@ -337,7 +351,7 @@ def _resolve_socket_path(self) -> str: return get_socket_path( get_current_device_uuid(), self.socket_dir, - draft_model_idx=self.draft_model_idx, + is_draft=self.is_draft, ) def _check_gpu_uuid(self, daemon_uuid: str | None) -> None: @@ -365,7 +379,6 @@ def _fallback_load_config(self) -> LoadConfig: self.load_config, load_format="auto", model_loader_extra_config={}, - weight_cache_draft_model_idx=None, ) def _fallback_load( diff --git a/vllm/model_executor/model_loader/weight_cache/protocol.py b/vllm/model_executor/model_loader/weight_cache/protocol.py index 75611ee67802..7911d0478f61 100644 --- a/vllm/model_executor/model_loader/weight_cache/protocol.py +++ b/vllm/model_executor/model_loader/weight_cache/protocol.py @@ -135,15 +135,15 @@ def get_socket_path( gpu_uuid: str, socket_dir: str | None = None, *, - draft_model_idx: int | None = None, + is_draft: bool = False, ) -> str: - """Socket path of a daemon group; ``draft_model_idx`` None is the target.""" + """Socket path of a daemon group; ``is_draft=False`` is the target.""" directory = get_socket_dir(socket_dir) return os.path.join( directory, SOCKET_NAME_TEMPLATE.format( gpu_uuid=gpu_uuid, - role=format_socket_role_suffix(draft_model_idx), + role=format_socket_role_suffix(is_draft), ), ) @@ -287,8 +287,8 @@ class WeightCacheKey: quant_config_hash: str revision: str | None vllm_version: str - draft_model_idx: int | None = None - """Daemon group the weights come from; None is the target model.""" + is_draft: bool = False + """Daemon group the weights come from; False is the target model.""" @classmethod def from_model_config( @@ -297,7 +297,7 @@ def from_model_config( tp_size: int, tp_rank: int, *, - draft_model_idx: int | None = None, + is_draft: bool = False, ) -> "WeightCacheKey": """Build the fingerprint for a model configuration. @@ -324,7 +324,7 @@ def from_model_config( quant_config_hash=_hash_quant_config(quant_config), revision=model_config.revision, vllm_version=vllm.version.__version__, - draft_model_idx=draft_model_idx, + is_draft=is_draft, ) def mismatched_fields(self, other: "WeightCacheKey") -> list[str]: diff --git a/vllm/model_executor/model_loader/weight_cache/utils.py b/vllm/model_executor/model_loader/weight_cache/utils.py index e9cb9a7ce2c7..1abbf572f847 100644 --- a/vllm/model_executor/model_loader/weight_cache/utils.py +++ b/vllm/model_executor/model_loader/weight_cache/utils.py @@ -7,7 +7,7 @@ wire format. """ -from typing import Any, TypeGuard +from typing import Any from vllm.config import SpeculativeConfig from vllm.model_executor.models.interfaces import SupportsEagleBase @@ -29,9 +29,7 @@ def export_model_attrs(model: Any) -> dict[str, bool]: } -def caches_draft_model( - speculative_config: SpeculativeConfig | None, -) -> TypeGuard[SpeculativeConfig]: +def is_draft_model_cacheable(speculative_config: SpeculativeConfig | None) -> bool: """Whether the daemon serves the speculative draft as a separate role.""" return ( speculative_config is not None @@ -40,13 +38,11 @@ def caches_draft_model( ) -def format_daemon_role(draft_model_idx: int | None) -> str: - """Name of a daemon group: the target model, or a draft by index.""" - return "target" if draft_model_idx is None else f"draft{draft_model_idx}" +def format_daemon_role(is_draft: bool) -> str: + """Name of a daemon group: the target model or the draft.""" + return "draft" if is_draft else "target" -def format_socket_role_suffix(draft_model_idx: int | None) -> str: - """Socket-name suffix keeping each draft group distinct from the target.""" - if draft_model_idx is None: - return "" - return f"_{format_daemon_role(draft_model_idx)}" +def format_socket_role_suffix(is_draft: bool) -> str: + """Socket-name suffix keeping the draft group distinct from the target.""" + return "_draft" if is_draft else "" From 6e03b0842b4f202ac9023be62a030990d19f6d07 Mon Sep 17 00:00:00 2001 From: Isotr0py Date: Mon, 21 Sep 2026 13:45:08 +0000 Subject: [PATCH 07/12] don't pass vllm_config and model_config together Signed-off-by: Isotr0py --- .../model_loader/weight_cache/daemon.py | 52 ++++++++----------- 1 file changed, 22 insertions(+), 30 deletions(-) diff --git a/vllm/model_executor/model_loader/weight_cache/daemon.py b/vllm/model_executor/model_loader/weight_cache/daemon.py index 56e950c69f18..d006a06e18fa 100644 --- a/vllm/model_executor/model_loader/weight_cache/daemon.py +++ b/vllm/model_executor/model_loader/weight_cache/daemon.py @@ -148,7 +148,7 @@ def _add(name: str, tensor: torch.Tensor, kind: str) -> None: def get_daemon_model( - vllm_config: VllmConfig, model_config: ModelConfig | None = None + vllm_config: VllmConfig, model_config: ModelConfig ) -> torch.nn.Module: """Load the daemon's model, composed from the configured loader. @@ -157,7 +157,6 @@ def get_daemon_model( always fails the check, so load_model's finalize step for it is unnecessary here. """ - model_config = model_config or vllm_config.model_config load_config = vllm_config.load_config loader = get_model_loader(load_config) device_config = vllm_config.device_config @@ -184,12 +183,16 @@ def __init__( distributed_init_method: str, socket_dir: str | None = None, is_draft: bool = False, - model_config: ModelConfig | None = None, ): self.vllm_config = vllm_config # A draft daemon builds the draft with the target's VllmConfig, like # the engine does; only the ModelConfig differs. - self.model_config = model_config or vllm_config.model_config + if is_draft: + spec = vllm_config.speculative_config + assert spec is not None and spec.draft_model_config is not None + self.model_config = spec.draft_model_config + else: + self.model_config = vllm_config.model_config self.tp_rank = tp_rank self.local_rank = local_rank self.distributed_init_method = distributed_init_method @@ -369,7 +372,6 @@ def _run_daemon( socket_dir: str | None, ready_queue: "multiprocessing.Queue[tuple[str, int]]", is_draft: bool = False, - model_config: ModelConfig | None = None, ) -> None: daemon = WeightCacheDaemon( vllm_config, @@ -378,30 +380,24 @@ def _run_daemon( distributed_init_method, socket_dir, is_draft, - model_config, ) daemon.load_model() daemon.serve_forever(ready_callback=lambda: ready_queue.put((daemon.role, tp_rank))) -def get_draft_daemon_config( - vllm_config: VllmConfig, -) -> tuple[VllmConfig, ModelConfig] | None: - """Configs for the draft daemon group, or None when the draft is not cached. +def get_draft_daemon_config(vllm_config: VllmConfig) -> VllmConfig | None: + """VllmConfig for the draft daemon group, or None when the draft is not + cached. Mirrors how the engine loads a draft: the target's VllmConfig with the - speculative kernel overrides, plus the draft's ModelConfig passed - separately, because draft classes read the target from - ``vllm_config.model_config``. + speculative kernel overrides; the draft's ModelConfig stays on + ``speculative_config.draft_model_config`` for the daemon to derive, + because draft classes read the target from ``vllm_config.model_config``. """ - speculative_config = vllm_config.speculative_config - if not is_draft_model_cacheable(speculative_config): + if not is_draft_model_cacheable(vllm_config.speculative_config): return None - assert speculative_config is not None - return ( - speculative_config.apply_draft_overrides(vllm_config), - speculative_config.draft_model_config, - ) + assert vllm_config.speculative_config is not None + return vllm_config.speculative_config.apply_draft_overrides(vllm_config) def _reject_unsupported_parallelism(parallel_config: ParallelConfig) -> None: @@ -475,14 +471,12 @@ def main() -> None: master_port = args.weight_cache_master_port or get_open_port() distributed_init_method = get_distributed_init_method(master_addr, master_port) - # (is_draft, vllm_config, model_config, rendezvous) per daemon group; the - # target group is not a draft and has no separate model_config. - groups: list[tuple[bool, VllmConfig, ModelConfig | None, str]] = [ - (False, vllm_config, None, distributed_init_method) + # (is_draft, vllm_config, rendezvous) per daemon group. + groups: list[tuple[bool, VllmConfig, str]] = [ + (False, vllm_config, distributed_init_method) ] - draft = get_draft_daemon_config(vllm_config) - if draft is not None: - draft_vllm_config, draft_model_config = draft + draft_vllm_config = get_draft_daemon_config(vllm_config) + if draft_vllm_config is not None: draft_master_port = args.weight_cache_draft_master_port if draft_master_port is None: draft_master_port = master_port + 1 if nnodes > 1 else get_open_port() @@ -495,7 +489,6 @@ def main() -> None: ( True, draft_vllm_config, - draft_model_config, get_distributed_init_method(master_addr, draft_master_port), ) ) @@ -509,7 +502,7 @@ def main() -> None: ] procs = [] expected_ready: set[tuple[str, int]] = set() - for is_draft, config, model_config, init_method in groups: + for is_draft, config, init_method in groups: role = format_daemon_role(is_draft) for local_rank, global_rank in enumerate(global_ranks): expected_ready.add((role, global_rank)) @@ -524,7 +517,6 @@ def main() -> None: args.weight_cache_socket_dir, ready_queue, is_draft, - model_config, ), name=f"vllm-weight-cache-{role}-{global_rank}", ) From b6b1ff87305c525425f80fb7ceb61ddd7a367012 Mon Sep 17 00:00:00 2001 From: Isotr0py Date: Mon, 21 Sep 2026 14:07:59 +0000 Subject: [PATCH 08/12] minor refactor Signed-off-by: Isotr0py --- .../model_loader/weight_cache/daemon.py | 61 ++++++++----------- .../model_loader/weight_cache/ipc_loader.py | 43 ++++++------- .../model_loader/weight_cache/protocol.py | 13 +++- 3 files changed, 55 insertions(+), 62 deletions(-) diff --git a/vllm/model_executor/model_loader/weight_cache/daemon.py b/vllm/model_executor/model_loader/weight_cache/daemon.py index d006a06e18fa..7179a360f213 100644 --- a/vllm/model_executor/model_loader/weight_cache/daemon.py +++ b/vllm/model_executor/model_loader/weight_cache/daemon.py @@ -63,7 +63,6 @@ import torch from vllm.config import ( - ModelConfig, ParallelConfig, VllmConfig, set_current_vllm_config, @@ -147,31 +146,6 @@ def _add(name: str, tensor: torch.Tensor, kind: str) -> None: return entries, aliases -def get_daemon_model( - vllm_config: VllmConfig, model_config: ModelConfig -) -> torch.nn.Module: - """Load the daemon's model, composed from the configured loader. - - Runs the quantization check after model creation but before the slow - weight load, so an unsupported method fails fast. Online quantization - always fails the check, so load_model's finalize step for it is - unnecessary here. - """ - load_config = vllm_config.load_config - loader = get_model_loader(load_config) - device_config = vllm_config.device_config - target_device = torch.device( - device_config.device if load_config.device is None else load_config.device - ) - with set_default_torch_dtype(model_config.dtype): - with target_device: - model = loader.create_model(vllm_config, model_config) - check_ipc_quant_support(model) - loader.load_weights(model, model_config) - process_weights_after_loading(model, model_config, target_device) - return model.eval() - - class WeightCacheDaemon: """Per-GPU process that loads one TP shard and serves CUDA IPC handles.""" @@ -221,13 +195,37 @@ def load_model(self) -> None: ) with set_current_vllm_config(self.vllm_config): ensure_model_parallel_initialized(tp_size, 1) - self.model = get_daemon_model(self.vllm_config, self.model_config) + self.model = self.get_model() logger.info( "Weight cache %s daemon rank %d loaded model", self.role, self.tp_rank, ) + def get_model(self) -> torch.nn.Module: + """Load the daemon's model, composed from the configured loader. + + Runs the quantization check after model creation but before the slow + weight load, so an unsupported method fails fast. Online quantization + always fails the check, so load_model's finalize step for it is + unnecessary here. + """ + vllm_config = self.vllm_config + model_config = self.model_config + load_config = vllm_config.load_config + loader = get_model_loader(load_config) + device_config = vllm_config.device_config + target_device = torch.device( + device_config.device if load_config.device is None else load_config.device + ) + with set_default_torch_dtype(model_config.dtype): + with target_device: + model = loader.create_model(vllm_config, model_config) + check_ipc_quant_support(model) + loader.load_weights(model, model_config) + process_weights_after_loading(model, model_config, target_device) + return model.eval() + def serve_forever(self, ready_callback: Callable[[], None] | None = None) -> None: """Serve requests until terminated. @@ -386,14 +384,7 @@ def _run_daemon( def get_draft_daemon_config(vllm_config: VllmConfig) -> VllmConfig | None: - """VllmConfig for the draft daemon group, or None when the draft is not - cached. - - Mirrors how the engine loads a draft: the target's VllmConfig with the - speculative kernel overrides; the draft's ModelConfig stays on - ``speculative_config.draft_model_config`` for the daemon to derive, - because draft classes read the target from ``vllm_config.model_config``. - """ + """VllmConfig for the draft daemon group""" if not is_draft_model_cacheable(vllm_config.speculative_config): return None assert vllm_config.speculative_config is not None diff --git a/vllm/model_executor/model_loader/weight_cache/ipc_loader.py b/vllm/model_executor/model_loader/weight_cache/ipc_loader.py index ec028794913b..e22fc2b0296f 100644 --- a/vllm/model_executor/model_loader/weight_cache/ipc_loader.py +++ b/vllm/model_executor/model_loader/weight_cache/ipc_loader.py @@ -25,9 +25,9 @@ ) from vllm.model_executor.model_loader.weight_cache.protocol import ( CacheConfigMismatchError, - TensorEntry, UnsupportedQuantForIPCError, WeightCacheKey, + WeightCacheState, WeightCacheUnavailableError, check_ipc_platform_support, check_ipc_quant_support, @@ -117,7 +117,7 @@ def load_weights(self, model: nn.Module, model_config: ModelConfig) -> None: loaded through this loader). """ device_index = torch.accelerator.current_device_index() - entries, _, _ = self._fetch_entries(model_config) + entries = self._fetch_entries(model_config).entries params = dict(model.named_parameters()) buffers = dict(model.named_buffers()) for name, entry in entries.items(): @@ -150,11 +150,9 @@ def load_model( f"was configured for the " f"{'draft' if self.is_draft else 'target'} group" ) - entries, aliases, attrs = self._fetch_entries(model_config) + state = self._fetch_entries(model_config) state_fetched = True - return self._build_model( - vllm_config, model_config, prefix, entries, aliases, attrs - ) + return self._build_model(vllm_config, model_config, prefix, state) except (WeightCacheUnavailableError, CacheConfigMismatchError) as e: if not self.fallback: raise @@ -185,9 +183,7 @@ def _build_model( vllm_config: VllmConfig, model_config: ModelConfig, prefix: str, - entries: dict[str, TensorEntry], - aliases: dict[str, str], - attrs: dict[str, bool], + state: WeightCacheState, ) -> nn.Module: device_config = vllm_config.device_config load_device = ( @@ -209,10 +205,10 @@ def _build_model( prefix=prefix, ) check_ipc_quant_support(model) - self._apply_entries(model, entries, aliases, device_index) + self._apply_entries(model, state, device_index) # Flags that load_weights would have set (e.g. EAGLE ownership of # embed_tokens / lm_head); the daemon ran it, this process did not. - for name, value in attrs.items(): + for name, value in state.attrs.items(): setattr(model, name, value) # The daemon exports tensors that already went through # process_weights_after_loading; re-run it in pre-processed mode @@ -229,7 +225,7 @@ def _build_model( self._send_release() logger.info( "Mapped %d tensors from the weight cache daemon (%s mode)", - len(entries), + len(state.entries), self.mode, ) return model.eval() @@ -237,8 +233,7 @@ def _build_model( def _apply_entries( self, model: nn.Module, - entries: dict[str, TensorEntry], - aliases: dict[str, str], + state: WeightCacheState, device_index: int, ) -> None: # remove_duplicate=False keeps tied module aliases reachable by name: @@ -269,7 +264,7 @@ def _register(name: str, tensor: torch.Tensor, is_param: bool) -> None: module.register_buffer(leaf, obj) registered[name] = obj - for name, entry in entries.items(): + for name, entry in state.entries.items(): tensor = entry.rebuild(device_index) if self.mode == "copy": tensor = tensor.clone() @@ -278,7 +273,7 @@ def _register(name: str, tensor: torch.Tensor, is_param: bool) -> None: # Re-establish tied-weight aliases by registering the *same* object the # canonical name resolved to, so parameter identity (and the tie) is # preserved instead of allocating uninitialized memory. - for alias_name, canonical_name in aliases.items(): + for alias_name, canonical_name in state.aliases.items(): obj = registered.get(canonical_name) if obj is None: logger.warning( @@ -289,9 +284,7 @@ def _register(name: str, tensor: torch.Tensor, is_param: bool) -> None: continue _register(alias_name, obj, isinstance(obj, nn.Parameter)) - def _fetch_entries( - self, model_config: ModelConfig - ) -> tuple[dict[str, TensorEntry], dict[str, str], dict[str, bool]]: + def _fetch_entries(self, model_config: ModelConfig) -> WeightCacheState: cache_config = WeightCacheKey.from_model_config( model_config, tp_size=get_tensor_model_parallel_world_size(), @@ -300,9 +293,7 @@ def _fetch_entries( ) return self._request_state(cache_config) - def _request_state( - self, cache_config: WeightCacheKey - ) -> tuple[dict[str, TensorEntry], dict[str, str], dict[str, bool]]: + def _request_state(self, cache_config: WeightCacheKey) -> WeightCacheState: with self._connect(self.state_timeout_s) as conn: send_msg(conn, {"cmd": "get_state", "cache_config": cache_config}) response = recv_msg(conn) @@ -316,10 +307,10 @@ def _request_state( f"Weight cache daemon error: {response.get('message')}" ) self._check_gpu_uuid(response.get("gpu_uuid")) - return ( - response["entries"], - response.get("aliases", {}), - response.get("attrs", {}), + return WeightCacheState( + entries=response["entries"], + aliases=response.get("aliases", {}), + attrs=response.get("attrs", {}), ) def _connect(self, timeout: float) -> socket.socket: diff --git a/vllm/model_executor/model_loader/weight_cache/protocol.py b/vllm/model_executor/model_loader/weight_cache/protocol.py index 7911d0478f61..567103eb4856 100644 --- a/vllm/model_executor/model_loader/weight_cache/protocol.py +++ b/vllm/model_executor/model_loader/weight_cache/protocol.py @@ -20,7 +20,7 @@ import struct import tempfile from dataclasses import dataclass, fields -from typing import Any +from typing import Any, NamedTuple import torch from torch.multiprocessing.reductions import rebuild_cuda_tensor, reduce_tensor @@ -368,6 +368,17 @@ def rebuild(self, device_index: int) -> torch.Tensor: return rebuild_cuda_tensor(*args) +class WeightCacheState(NamedTuple): + """Client-side decode of a daemon's get_state response payload.""" + + entries: dict[str, TensorEntry] + """Model tensors, exported as CUDA IPC handles or shipped by value.""" + aliases: dict[str, str] + """Duplicate (tied) weight names aliased to their canonical entry.""" + attrs: dict[str, bool] + """Python-side flags set by load_weights, e.g. EAGLE ownership flags.""" + + def send_msg(sock: socket.socket, obj: Any) -> None: payload = pickle.dumps(obj, protocol=pickle.HIGHEST_PROTOCOL) sock.sendall(_LEN_STRUCT.pack(len(payload))) From 5683fe49cd4516987d97f04e8c3bb4184c842bbe Mon Sep 17 00:00:00 2001 From: Isotr0py Date: Mon, 21 Sep 2026 14:12:41 +0000 Subject: [PATCH 09/12] avoid duplicate test impl Signed-off-by: Isotr0py --- .../model_loader/test_weight_cache.py | 47 ++++++------------- 1 file changed, 15 insertions(+), 32 deletions(-) diff --git a/tests/model_executor/model_loader/test_weight_cache.py b/tests/model_executor/model_loader/test_weight_cache.py index 6b8c1a2a93bd..eb672bc35b2c 100644 --- a/tests/model_executor/model_loader/test_weight_cache.py +++ b/tests/model_executor/model_loader/test_weight_cache.py @@ -121,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( @@ -200,16 +203,22 @@ def generate( "--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") @@ -229,37 +238,11 @@ 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: - 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) - - assert cold_outputs == baseline_outputs - assert warm_outputs == baseline_outputs - assert restart_outputs == baseline_outputs - - -def test_ipc_cache_caches_mtp_draft(vllm_runner): - """The daemon caches the MTP draft in its draft group alongside the target. - - The warm runs use fallback=False, so both the target and the draft fail - hard unless their daemon groups served the weights; matching the - disk-loaded baseline proves the cached draft produces identical drafts. - """ - if not current_platform.is_cuda_alike(): - pytest.skip("Weight cache IPC sharing requires CUDA or ROCm") - - case = QWEN_MTP_CASE - # Baseline: target and MTP draft both loaded from disk. - baseline_outputs = generate(vllm_runner, case, None, fallback=True) - assert all(text for _, texts in baseline_outputs for text in texts) - - # Cold start: no daemon is serving, so both models fall back to disk. - 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, num_groups=2 + 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. From a33ed326b6b05a6df1160b6bac65d802e69783e0 Mon Sep 17 00:00:00 2001 From: Isotr0py Date: Mon, 21 Sep 2026 14:27:21 +0000 Subject: [PATCH 10/12] clean Signed-off-by: Isotr0py --- .../model_loader/weight_cache/daemon.py | 50 ++++++++++--------- 1 file changed, 26 insertions(+), 24 deletions(-) diff --git a/vllm/model_executor/model_loader/weight_cache/daemon.py b/vllm/model_executor/model_loader/weight_cache/daemon.py index 7179a360f213..1eeff81f65dc 100644 --- a/vllm/model_executor/model_loader/weight_cache/daemon.py +++ b/vllm/model_executor/model_loader/weight_cache/daemon.py @@ -59,6 +59,7 @@ import socket import sys from collections.abc import Callable +from itertools import product import torch @@ -468,9 +469,9 @@ def main() -> None: ] draft_vllm_config = get_draft_daemon_config(vllm_config) if draft_vllm_config is not None: - draft_master_port = args.weight_cache_draft_master_port - if draft_master_port is None: - draft_master_port = master_port + 1 if nnodes > 1 else get_open_port() + draft_master_port = args.weight_cache_draft_master_port or ( + master_port + 1 if nnodes > 1 else get_open_port() + ) if draft_master_port == master_port: raise ValueError( "--weight-cache-draft-master-port must differ from " @@ -491,27 +492,28 @@ def main() -> None: node_rank * local_world_size + local_rank for local_rank in range(local_world_size) ] - procs = [] - expected_ready: set[tuple[str, int]] = set() - for is_draft, config, init_method in groups: - role = format_daemon_role(is_draft) - for local_rank, global_rank in enumerate(global_ranks): - expected_ready.add((role, global_rank)) - procs.append( - ctx.Process( - target=_run_daemon, - args=( - global_rank, - local_rank, - config, - init_method, - args.weight_cache_socket_dir, - ready_queue, - is_draft, - ), - name=f"vllm-weight-cache-{role}-{global_rank}", - ) - ) + expected_ready = { + (format_daemon_role(is_draft), global_rank) + for (is_draft, _, _), global_rank in product(groups, global_ranks) + } + procs = [ + ctx.Process( + target=_run_daemon, + args=( + global_rank, + local_rank, + config, + init_method, + args.weight_cache_socket_dir, + ready_queue, + is_draft, + ), + name=f"vllm-weight-cache-{format_daemon_role(is_draft)}-{global_rank}", + ) + for (is_draft, config, init_method), (local_rank, global_rank) in product( + groups, enumerate(global_ranks) + ) + ] for proc in procs: proc.start() From cec5d084f0aa69c55169b19086d0fefa85bd1135 Mon Sep 17 00:00:00 2001 From: Isotr0py Date: Mon, 21 Sep 2026 16:31:52 +0000 Subject: [PATCH 11/12] fix tests Signed-off-by: Isotr0py --- tests/model_executor/test_qwen3_omni.py | 3 ++- .../v1/spec_decode/test_draft_attention_backend_override.py | 5 ++++- tests/v1/spec_decode/test_draft_moe_backend_override.py | 5 ++++- 3 files changed, 10 insertions(+), 3 deletions(-) diff --git a/tests/model_executor/test_qwen3_omni.py b/tests/model_executor/test_qwen3_omni.py index cf3df85caa58..71c74863e996 100644 --- a/tests/model_executor/test_qwen3_omni.py +++ b/tests/model_executor/test_qwen3_omni.py @@ -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)), ) diff --git a/tests/v1/spec_decode/test_draft_attention_backend_override.py b/tests/v1/spec_decode/test_draft_attention_backend_override.py index 410afe8d61d3..26bcd68fea8f 100644 --- a/tests/v1/spec_decode/test_draft_attention_backend_override.py +++ b/tests/v1/spec_decode/test_draft_attention_backend_override.py @@ -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 @@ -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 diff --git a/tests/v1/spec_decode/test_draft_moe_backend_override.py b/tests/v1/spec_decode/test_draft_moe_backend_override.py index 5a4522e5b301..a0cc26c16df4 100644 --- a/tests/v1/spec_decode/test_draft_moe_backend_override.py +++ b/tests/v1/spec_decode/test_draft_moe_backend_override.py @@ -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 @@ -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 From 9d0ad0c0a68ecbeddf025150e0d991c796d29b40 Mon Sep 17 00:00:00 2001 From: Isotr0py Date: Tue, 22 Sep 2026 01:38:51 +0000 Subject: [PATCH 12/12] shorten socket path Signed-off-by: Isotr0py --- .../model_loader/weight_cache/protocol.py | 18 ++++++++++-------- 1 file changed, 10 insertions(+), 8 deletions(-) diff --git a/vllm/model_executor/model_loader/weight_cache/protocol.py b/vllm/model_executor/model_loader/weight_cache/protocol.py index 567103eb4856..3803030cd440 100644 --- a/vllm/model_executor/model_loader/weight_cache/protocol.py +++ b/vllm/model_executor/model_loader/weight_cache/protocol.py @@ -137,15 +137,17 @@ def get_socket_path( *, is_draft: bool = False, ) -> str: - """Socket path of a daemon group; ``is_draft=False`` is the target.""" - directory = get_socket_dir(socket_dir) - return os.path.join( - directory, - SOCKET_NAME_TEMPLATE.format( - gpu_uuid=gpu_uuid, - role=format_socket_role_suffix(is_draft), - ), + """Socket path of a daemon group; ``is_draft=False`` is the target. + + The GPU uuid is hashed to keep the name well under the AF_UNIX path + limit (~108 bytes) even with the draft role suffix. + """ + gpu_id = safe_hash(gpu_uuid.encode()).hexdigest() + name = SOCKET_NAME_TEMPLATE.format( + gpu_uuid=gpu_id, + role=format_socket_role_suffix(is_draft), ) + return os.path.join(get_socket_dir(socket_dir), name) def ensure_private_socket_dir(directory: str, strict_perms: bool = True) -> None: