From a3cbf97a3ec6b40392c7508e950b39198a1ce0f4 Mon Sep 17 00:00:00 2001 From: Zupeng Wang <71580390+zupengwang@users.noreply.github.com> Date: Sat, 25 Jul 2026 23:39:13 +0800 Subject: [PATCH 1/2] [Feature][Model Runner V2] Support extract_hidden_states speculation Co-authored-by: OpenAI Signed-off-by: Zupeng Wang <71580390+zupengwang@users.noreply.github.com> --- tests/test_config.py | 15 ++ ...st_gpu_extract_hidden_states_speculator.py | 120 ++++++++++++++ vllm/config/vllm.py | 1 + vllm/v1/worker/gpu/model_runner.py | 7 +- vllm/v1/worker/gpu/spec_decode/__init__.py | 8 +- .../gpu/spec_decode/extract_hidden_states.py | 148 ++++++++++++++++++ 6 files changed, 297 insertions(+), 2 deletions(-) create mode 100644 tests/v1/worker/test_gpu_extract_hidden_states_speculator.py create mode 100644 vllm/v1/worker/gpu/spec_decode/extract_hidden_states.py diff --git a/tests/test_config.py b/tests/test_config.py index 70c8728d25d4..6c247de73b19 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -5,6 +5,7 @@ import os from dataclasses import MISSING, Field, asdict, dataclass, field from types import SimpleNamespace +from typing import cast from unittest.mock import patch import pydantic @@ -80,6 +81,20 @@ def test_rocm_defaults_deepseek_v4_to_mrv1(monkeypatch): default_v2_model_runner_architectures.cache_clear() +def test_v2_model_runner_supports_extract_hidden_states(): + config = VllmConfig() + config.speculative_config = cast( + SpeculativeConfig, + SimpleNamespace( + method="extract_hidden_states", + parallel_drafting=False, + enable_adaptive_verification=False, + ), + ) + + assert config._get_v2_model_runner_unsupported_features() == [] + + @pytest.mark.parametrize( ("use_v2_model_runner", "expected_capture_sizes"), [ diff --git a/tests/v1/worker/test_gpu_extract_hidden_states_speculator.py b/tests/v1/worker/test_gpu_extract_hidden_states_speculator.py new file mode 100644 index 000000000000..a5272ce78253 --- /dev/null +++ b/tests/v1/worker/test_gpu_extract_hidden_states_speculator.py @@ -0,0 +1,120 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from contextlib import nullcontext +from types import SimpleNamespace +from typing import Any, cast + +import pytest +import torch + +from vllm.v1.worker.gpu.spec_decode import extract_hidden_states as spec_module +from vllm.v1.worker.gpu.spec_decode import init_speculator +from vllm.v1.worker.gpu.spec_decode.extract_hidden_states import ( + ExtractHiddenStatesSpeculator, +) + + +class _RecordingModel(torch.nn.Module): + def forward(self, *, hidden_states: torch.Tensor) -> None: + self.hidden_states = hidden_states.clone() + + +def test_init_speculator_dispatches_extract_hidden_states(monkeypatch): + vllm_config = cast( + Any, + SimpleNamespace( + speculative_config=SimpleNamespace(method="extract_hidden_states") + ), + ) + device = torch.device("cpu") + + def fake_speculator(config, target_device): + return config, target_device + + monkeypatch.setattr(spec_module, "ExtractHiddenStatesSpeculator", fake_speculator) + + assert init_speculator(vllm_config, device) == (vllm_config, device) + + +def test_propose_caches_hidden_states_and_returns_sampled_tokens(monkeypatch): + contexts = [] + + def fake_set_forward_context(*args, **kwargs): + contexts.append((args, kwargs)) + return nullcontext() + + monkeypatch.setattr(spec_module, "set_forward_context", fake_set_forward_context) + + layer_name = "cache_only_layers.2" + speculator = object.__new__(ExtractHiddenStatesSpeculator) + speculator.vllm_config = cast(Any, SimpleNamespace()) + speculator.num_hidden_states = 2 + speculator.hidden_states = torch.zeros(4, 2, 3) + speculator.draft_attn_layer_names = {layer_name} + speculator.model = _RecordingModel() + + input_batch = cast( + Any, + SimpleNamespace( + idx_mapping=torch.tensor([2, 0], dtype=torch.int32), + is_padding=torch.zeros(4, dtype=torch.bool), + ), + ) + aux_hidden_states = [ + torch.full((4, 3), 1.0), + torch.full((4, 3), 2.0), + ] + attn_metadata = {layer_name: object(), "target_layer": object()} + slot_mappings = { + layer_name: torch.arange(4), + "target_layer": torch.arange(4), + } + last_sampled = torch.tensor([[10], [11], [12]], dtype=torch.int64) + + draft_tokens = speculator.propose( + input_batch=input_batch, + attn_metadata=attn_metadata, + slot_mappings=slot_mappings, + last_hidden_states=torch.empty(0), + aux_hidden_states=aux_hidden_states, + num_sampled=torch.empty(0), + num_rejected=torch.empty(0), + last_sampled=last_sampled, + next_prefill_tokens=torch.empty(0), + temperature=torch.empty(0), + seeds=torch.empty(0), + ) + + expected_hidden_states = torch.stack(aux_hidden_states, dim=1) + assert torch.equal(speculator.model.hidden_states, expected_hidden_states) + assert torch.equal(draft_tokens, torch.tensor([[12], [10]])) + + assert len(contexts) == 1 + args, kwargs = contexts[0] + assert args[0] == {layer_name: attn_metadata[layer_name]} + assert kwargs["num_tokens"] == 4 + assert set(kwargs["slot_mapping"]) == {layer_name} + assert torch.equal(kwargs["slot_mapping"][layer_name], slot_mappings[layer_name]) + + +def test_propose_requires_aux_hidden_states(): + speculator = object.__new__(ExtractHiddenStatesSpeculator) + speculator.num_hidden_states = 2 + input_batch = cast( + Any, SimpleNamespace(idx_mapping=torch.tensor([0], dtype=torch.int32)) + ) + + with pytest.raises(ValueError, match="aux_hidden_states are required"): + speculator.propose( + input_batch=input_batch, + attn_metadata={}, + slot_mappings={}, + last_hidden_states=torch.empty(0), + aux_hidden_states=None, + num_sampled=torch.empty(0), + num_rejected=torch.empty(0), + last_sampled=torch.tensor([[10]]), + next_prefill_tokens=torch.empty(0), + temperature=torch.empty(0), + seeds=torch.empty(0), + ) diff --git a/vllm/config/vllm.py b/vllm/config/vllm.py index 0a4c5cad8838..a3c87776b2a0 100644 --- a/vllm/config/vllm.py +++ b/vllm/config/vllm.py @@ -2386,6 +2386,7 @@ def _get_v2_model_runner_unsupported_features(self) -> list[str]: "mtp", "dflash", "dspark", + "extract_hidden_states", ): unsupported.append(f"speculative method '{speculative_config.method}'") diff --git a/vllm/v1/worker/gpu/model_runner.py b/vllm/v1/worker/gpu/model_runner.py index abbe6d5d3de6..396ea7b1d188 100644 --- a/vllm/v1/worker/gpu/model_runner.py +++ b/vllm/v1/worker/gpu/model_runner.py @@ -242,7 +242,12 @@ def __init__(self, vllm_config: VllmConfig, device: torch.device): if self.is_last_pp_rank: self.speculator = init_speculator(self.vllm_config, self.device) - if self.speculative_config.method in ("eagle3", "dflash", "dspark"): + if self.speculative_config.method in ( + "eagle3", + "dflash", + "dspark", + "extract_hidden_states", + ): # Drafting may require auxiliary hidden states from target model outputs self.use_aux_hidden_state_outputs = True if self.use_pp: diff --git a/vllm/v1/worker/gpu/spec_decode/__init__.py b/vllm/v1/worker/gpu/spec_decode/__init__.py index 4229696f255c..2f75109893d0 100644 --- a/vllm/v1/worker/gpu/spec_decode/__init__.py +++ b/vllm/v1/worker/gpu/spec_decode/__init__.py @@ -8,7 +8,13 @@ def init_speculator(vllm_config: VllmConfig, device: torch.device): speculative_config = vllm_config.speculative_config assert speculative_config is not None - if speculative_config.method == "dflash": + if speculative_config.method == "extract_hidden_states": + from vllm.v1.worker.gpu.spec_decode.extract_hidden_states import ( + ExtractHiddenStatesSpeculator, + ) + + return ExtractHiddenStatesSpeculator(vllm_config, device) + elif speculative_config.method == "dflash": from vllm.v1.worker.gpu.spec_decode.dflash.speculator import ( DFlashSpeculator, ) diff --git a/vllm/v1/worker/gpu/spec_decode/extract_hidden_states.py b/vllm/v1/worker/gpu/spec_decode/extract_hidden_states.py new file mode 100644 index 000000000000..6aea4fd9d676 --- /dev/null +++ b/vllm/v1/worker/gpu/spec_decode/extract_hidden_states.py @@ -0,0 +1,148 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from typing import Any + +import torch +import torch.nn as nn + +from vllm.compilation.backends import set_model_tag +from vllm.config import VllmConfig +from vllm.config.compilation import CUDAGraphMode +from vllm.forward_context import set_forward_context +from vllm.model_executor.model_loader import get_model +from vllm.v1.worker.gpu.input_batch import InputBatch +from vllm.v1.worker.gpu.spec_decode.speculator import DraftModelSpeculator + + +class ExtractHiddenStatesSpeculator(DraftModelSpeculator): + """Cache target hidden states while returning always-accepted draft tokens.""" + + def __init__(self, vllm_config: VllmConfig, device: torch.device): + super().__init__(vllm_config, device) + + if self.num_speculative_steps != 1: + raise ValueError( + "extract_hidden_states requires num_speculative_tokens to be 1" + ) + if self.speculative_config.disable_padded_drafter_batch: + raise ValueError( + "disable_padded_drafter_batch is not supported with " + "extract_hidden_states method" + ) + + self.supports_mm_inputs = False + layer_ids = getattr( + self.draft_model_config.hf_config, + "eagle_aux_hidden_state_layer_ids", + None, + ) + if not layer_ids: + raise ValueError( + "eagle_aux_hidden_state_layer_ids must be set in the draft " + "model config for extract_hidden_states method" + ) + + self.num_hidden_states = len(layer_ids) + assert isinstance(self.dtype, torch.dtype) + self.hidden_states = torch.zeros( + self.max_num_tokens, + self.num_hidden_states, + self.vllm_config.model_config.get_hidden_size(), + dtype=self.dtype, + device=device, + ) + + def load_draft_model( + self, + target_model: nn.Module, + target_attn_layer_names: set[str], + ) -> nn.Module: + del target_model, target_attn_layer_names + with set_model_tag("extract_hidden_states"): + return get_model( + vllm_config=self.vllm_config, + model_config=self.draft_model_config, + ) + + def load_model(self, target_model: nn.Module) -> None: + super().load_model(target_model) + if len(self.draft_attn_layer_names) != 1: + raise ValueError( + "ExtractHiddenStatesModel should have exactly one attention " + f"layer, found {len(self.draft_attn_layer_names)}" + ) + + def init_cudagraph_manager(self, cudagraph_mode: CUDAGraphMode) -> None: + del cudagraph_mode + + def capture(self) -> None: + return None + + @torch.inference_mode() + def propose( + self, + input_batch: InputBatch, + attn_metadata: dict[str, Any], + slot_mappings: dict[str, torch.Tensor], + last_hidden_states: torch.Tensor, + aux_hidden_states: list[torch.Tensor] | None, + num_sampled: torch.Tensor, + num_rejected: torch.Tensor, + last_sampled: torch.Tensor, + next_prefill_tokens: torch.Tensor, + temperature: torch.Tensor, + seeds: torch.Tensor, + num_tokens_across_dp: torch.Tensor | None = None, + dummy_run: bool = False, + skip_attn_for_dummy_run: bool = False, + mm_inputs: tuple[list[torch.Tensor], torch.Tensor] | None = None, + is_profile: bool = False, + ) -> torch.Tensor: + del ( + last_hidden_states, + num_sampled, + num_rejected, + next_prefill_tokens, + temperature, + seeds, + dummy_run, + mm_inputs, + is_profile, + ) + + draft_tokens = last_sampled[input_batch.idx_mapping, :1] + if skip_attn_for_dummy_run: + return draft_tokens + if aux_hidden_states is None: + raise ValueError( + "aux_hidden_states are required when using extract_hidden_states" + ) + if len(aux_hidden_states) != self.num_hidden_states: + raise ValueError( + f"Expected {self.num_hidden_states} auxiliary hidden states, " + f"got {len(aux_hidden_states)}" + ) + + stacked_hidden_states = torch.stack(aux_hidden_states, dim=1) + num_tokens = stacked_hidden_states.shape[0] + self.hidden_states[:num_tokens].copy_(stacked_hidden_states) + + draft_attn_metadata = { + name: attn_metadata[name] for name in self.draft_attn_layer_names + } + draft_slot_mappings = { + name: slot_mappings[name][:num_tokens] + for name in self.draft_attn_layer_names + } + with set_forward_context( + draft_attn_metadata, + self.vllm_config, + num_tokens=num_tokens, + num_tokens_across_dp=num_tokens_across_dp, + cudagraph_runtime_mode=CUDAGraphMode.NONE, + slot_mapping=draft_slot_mappings, + is_padding=input_batch.is_padding[:num_tokens], + ): + self.model(hidden_states=self.hidden_states[:num_tokens]) + + return draft_tokens From 422e6dd696d4048231a1c2237c596d6afafcbee6 Mon Sep 17 00:00:00 2001 From: Misha Goin Date: Wed, 19 Aug 2026 19:47:49 +0000 Subject: [PATCH 2/2] Restrict hidden-state extraction to greedy sampling Co-authored-by: OpenAI Codex Signed-off-by: Misha Goin --- .../test_gpu_extract_hidden_states_speculator.py | 12 ++++++++++++ .../worker/gpu/spec_decode/extract_hidden_states.py | 5 +++++ 2 files changed, 17 insertions(+) diff --git a/tests/v1/worker/test_gpu_extract_hidden_states_speculator.py b/tests/v1/worker/test_gpu_extract_hidden_states_speculator.py index a5272ce78253..d1aa0798feb2 100644 --- a/tests/v1/worker/test_gpu_extract_hidden_states_speculator.py +++ b/tests/v1/worker/test_gpu_extract_hidden_states_speculator.py @@ -19,6 +19,18 @@ def forward(self, *, hidden_states: torch.Tensor) -> None: self.hidden_states = hidden_states.clone() +def test_init_requires_greedy_draft_sampling(): + vllm_config = cast( + Any, + SimpleNamespace( + speculative_config=SimpleNamespace(draft_sample_method="probabilistic") + ), + ) + + with pytest.raises(ValueError, match="only supports draft_sample_method='greedy'"): + ExtractHiddenStatesSpeculator(vllm_config, torch.device("cpu")) + + def test_init_speculator_dispatches_extract_hidden_states(monkeypatch): vllm_config = cast( Any, diff --git a/vllm/v1/worker/gpu/spec_decode/extract_hidden_states.py b/vllm/v1/worker/gpu/spec_decode/extract_hidden_states.py index 6aea4fd9d676..efcdec80182e 100644 --- a/vllm/v1/worker/gpu/spec_decode/extract_hidden_states.py +++ b/vllm/v1/worker/gpu/spec_decode/extract_hidden_states.py @@ -18,6 +18,11 @@ class ExtractHiddenStatesSpeculator(DraftModelSpeculator): """Cache target hidden states while returning always-accepted draft tokens.""" def __init__(self, vllm_config: VllmConfig, device: torch.device): + assert vllm_config.speculative_config is not None + if vllm_config.speculative_config.draft_sample_method != "greedy": + raise ValueError( + "extract_hidden_states only supports draft_sample_method='greedy'" + ) super().__init__(vllm_config, device) if self.num_speculative_steps != 1: