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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 16 additions & 1 deletion docs/source/user_guide/feature_guide/speculative_decoding.md
Original file line number Diff line number Diff line change
Expand Up @@ -438,11 +438,26 @@ Suffix Decoding can achieve better performance for tasks with high repetition, s

## Extracting Hidden States

The `extract_hidden_states` method is a special speculative decoding mode that does not perform actual speculation. Instead, it extracts hidden states from specified layers of the target model and saves them to disk. This is primarily used for collecting training data for EAGLE-style draft models.
The `extract_hidden_states` method is a special speculative decoding mode that does not perform actual speculation. Instead, it extracts hidden states from specified layers of the target model and saves them to disk. This is primarily used for collecting training data for EAGLE-style draft models. The dumps are then used to train EAGLE/EAGLE-3 drafts.

> [!NOTE]
> This method produces only 1 output token per request. The primary output is the hidden states saved to disk, not the generated text.

Both Model Runner V1 and Model Runner V2 are supported on Ascend. Enable V2 with:

```shell
export VLLM_USE_V2_MODEL_RUNNER=1
```

> [!NOTE]
> Model Runner V2 support reuses upstream vLLM's `ExtractHiddenStatesSpeculator`
> ([PR #49811](https://github.com/vllm-project/vllm/pull/49811)). Ascend only
> adds `init_speculator` dispatch and NPU KV allocate/reshape for
> `HiddenStateCacheSpec`. After
> [vLLM #51718](https://github.com/vllm-project/vllm/pull/51718) (0828 pin),
> hidden-state layers keep private `[B, H, N, C]` buffers so they cannot overlay
> the standardized hybrid Attention/Mamba backing.

- Offline inference

```python
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,11 @@
* a hybrid attention model (Qwen3.5-0.8B, GatedDeltaNet + full_attention)
loaded with dummy weights as a shape/round-trip smoke test. The hybrid case
mirrors upstream vLLM PR #39949.
* Model Runner V1 (Ascend default) and Model Runner V2 (`VLLM_USE_V2_MODEL_RUNNER=1`),
covering the Ascend adaptation of upstream vLLM PR #49811 on the 0828 pin.
* token-in / token-out via ``skip_tokenizer_init`` + ``TokensPrompt`` on the
text-only dense model (dummy weights). Qwen3.5 is multimodal, so skipping
tokenizer init leaves ``tokenizer=None`` and ``Qwen3VLProcessor`` crashes.
"""

from __future__ import annotations
Expand All @@ -35,6 +40,7 @@
import torch
from vllm import LLM, SamplingParams
from vllm.distributed.kv_transfer.kv_connector.v1 import example_hidden_states_connector
from vllm.inputs import TokensPrompt

os.environ["VLLM_WORKER_MULTIPROC_METHOD"] = "spawn"

Expand All @@ -45,13 +51,20 @@
HYBRID_MODEL = "Qwen/Qwen3.5-0.8B"
HYBRID_AUX_HIDDEN_STATE_LAYER_IDS = [5, 11, 17]

# In-vocab dummy sequences for skip_tokenizer_init (Qwen3 vocab >> 500).
TOKEN_IN_PROMPTS = [
[100, 200, 300, 400, 500],
[7, 8, 9, 10, 11, 12, 13, 14],
]


@dataclass
class ExtractHiddenStatesCase:
model_name: str
aux_hidden_state_layer_ids: list[int]
prompts: list[str]
enforce_eager: bool
prompts: list[str] | None = None
token_prompts: list[list[int]] | None = None
# ``None`` means "do not pass the argument", preserving each model's
# original defaults.
gpu_memory_utilization: float | None = None
Expand All @@ -62,6 +75,10 @@ class ExtractHiddenStatesCase:
verify_nonzero: bool = True
# Hybrid smoke test additionally checks the token_ids round-trip.
verify_token_ids: bool = False
# When True, force Model Runner V2 via VLLM_USE_V2_MODEL_RUNNER.
use_v2_model_runner: bool = False
# Token-in / token-out: skip tokenizer init and pass TokensPrompt.
skip_tokenizer_init: bool = False


CASES = [
Expand Down Expand Up @@ -110,6 +127,72 @@ class ExtractHiddenStatesCase:
),
id="hybrid_dummy_eager",
),
pytest.param(
ExtractHiddenStatesCase(
model_name=DENSE_MODEL,
aux_hidden_state_layer_ids=DENSE_AUX_HIDDEN_STATE_LAYER_IDS,
prompts=[
"Hello, how are you?",
"What is machine learning?",
],
enforce_eager=True,
gpu_memory_utilization=0.8,
max_num_seqs=16,
use_v2_model_runner=True,
),
id="dense_eager_mrv2",
),
pytest.param(
ExtractHiddenStatesCase(
model_name=HYBRID_MODEL,
aux_hidden_state_layer_ids=HYBRID_AUX_HIDDEN_STATE_LAYER_IDS,
prompts=[
"Hello world",
"Test prompt with several tokens",
],
enforce_eager=True,
gpu_memory_utilization=0.4,
max_model_len=256,
load_format="dummy",
verify_nonzero=False,
verify_token_ids=True,
use_v2_model_runner=True,
),
id="hybrid_dummy_eager_mrv2",
),
pytest.param(
ExtractHiddenStatesCase(
model_name=DENSE_MODEL,
aux_hidden_state_layer_ids=DENSE_AUX_HIDDEN_STATE_LAYER_IDS,
token_prompts=TOKEN_IN_PROMPTS,
enforce_eager=True,
gpu_memory_utilization=0.8,
max_num_seqs=16,
max_model_len=256,
load_format="dummy",
verify_nonzero=False,
verify_token_ids=True,
skip_tokenizer_init=True,
),
id="dense_dummy_token_in_token_out",
),
pytest.param(
ExtractHiddenStatesCase(
model_name=DENSE_MODEL,
aux_hidden_state_layer_ids=DENSE_AUX_HIDDEN_STATE_LAYER_IDS,
token_prompts=TOKEN_IN_PROMPTS,
enforce_eager=True,
gpu_memory_utilization=0.8,
max_num_seqs=16,
max_model_len=256,
load_format="dummy",
verify_nonzero=False,
verify_token_ids=True,
use_v2_model_runner=True,
skip_tokenizer_init=True,
),
id="dense_dummy_token_in_token_out_mrv2",
),
]


Expand Down Expand Up @@ -142,9 +225,35 @@ def _verify_output(output, expected_shape, *, verify_nonzero, verify_token_ids):
example_hidden_states_connector.cleanup_hidden_states(hidden_states_path)


def _generate_inputs(case: ExtractHiddenStatesCase):
if case.skip_tokenizer_init:
assert case.token_prompts is not None
return [TokensPrompt(prompt_token_ids=ids) for ids in case.token_prompts]
assert case.prompts is not None
return case.prompts


def _verify_token_in_token_out(output, token_prompt: list[int], *, max_tokens: int):
"""Input token ids round-trip; generated ids are present without detokenizing."""
assert list(output.prompt_token_ids) == token_prompt
assert not output.outputs[0].text
assert len(output.outputs[0].token_ids) == max_tokens


@pytest.mark.parametrize("case", CASES)
def test_extract_hidden_states(case: ExtractHiddenStatesCase, sampling_config):
def test_extract_hidden_states(case: ExtractHiddenStatesCase, sampling_config, monkeypatch):
"""Extract hidden states from the target model and validate the dump."""
if case.use_v2_model_runner:
monkeypatch.setenv("VLLM_USE_V2_MODEL_RUNNER", "1")
else:
monkeypatch.delenv("VLLM_USE_V2_MODEL_RUNNER", raising=False)

generate_inputs = _generate_inputs(case)
if case.skip_tokenizer_init:
sampling = SamplingParams(temperature=0, max_tokens=1, detokenize=False)
else:
sampling = sampling_config

with tempfile.TemporaryDirectory() as tmpdirname:
llm_kwargs = dict(
model=case.model_name,
Expand Down Expand Up @@ -176,18 +285,30 @@ def test_extract_hidden_states(case: ExtractHiddenStatesCase, sampling_config):
llm_kwargs["max_model_len"] = case.max_model_len
if case.load_format is not None:
llm_kwargs["load_format"] = case.load_format
if case.skip_tokenizer_init:
llm_kwargs["skip_tokenizer_init"] = True

llm = LLM(**llm_kwargs)

outputs = llm.generate(case.prompts, sampling_config)
outputs = llm.generate(generate_inputs, sampling)
hidden_size = llm.llm_engine.model_config.get_hidden_size()
num_layers = len(case.aux_hidden_state_layer_ids)
vocab_size = llm.llm_engine.model_config.get_vocab_size()

assert len(outputs) == len(case.prompts)
assert len(outputs) == len(generate_inputs)

for output in outputs:
for idx, output in enumerate(outputs):
num_tokens = len(output.prompt_token_ids)
expected_shape = (num_tokens, num_layers, hidden_size)
if case.skip_tokenizer_init:
assert case.token_prompts is not None
assert sampling.max_tokens is not None
_verify_token_in_token_out(
output,
case.token_prompts[idx],
max_tokens=sampling.max_tokens,
)
assert all(0 <= token_id < vocab_size for token_id in output.outputs[0].token_ids)
_verify_output(
output,
expected_shape,
Expand Down
73 changes: 73 additions & 0 deletions tests/ut/worker/test_attn_utils_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
from vllm.model_executor.models.deepseek_v2 import DeepseekV32IndexerCache
from vllm.v1.kv_cache_interface import (
FullAttentionSpec,
HiddenStateCacheSpec,
KVCacheConfig,
KVCacheGroupSpec,
KVCacheTensor,
Expand Down Expand Up @@ -753,6 +754,78 @@ def test_mrv2_builds_shared_dsa_metadata_for_each_execution_mode(
assert all(call["pcp_cache_group_idx"] is None for call in calls)


def test_mrv2_allocates_and_reshapes_hidden_state_cache(monkeypatch):
"""HiddenStateCacheSpec must stay on a private [B, H, N, C] path after #51718."""
from vllm.model_executor.models.extract_hidden_states import (
CacheOnlyAttentionBackend,
)

layer_name = "draft.cache_only_layers.36"
block_size = 16
num_kv_heads = 3
head_size = 8
num_blocks = 4
dtype = torch.bfloat16
spec = HiddenStateCacheSpec(
block_size=block_size,
num_kv_heads=num_kv_heads,
head_size=head_size,
dtype=dtype,
)
page_bytes = spec.page_size_bytes
tensor_size = num_blocks * page_bytes

kv_cache_config = KVCacheConfig(
num_blocks=num_blocks,
kv_cache_tensors=[_make_kv_cache_tensor(tensor_size, [layer_name], page_bytes)],
kv_cache_groups=[
KVCacheGroupSpec(
layer_names=[layer_name],
kv_cache_spec=spec,
)
],
)

monkeypatch.setattr(
attn_utils,
"get_current_vllm_config",
lambda: SimpleNamespace(
kv_transfer_config=None,
model_config=SimpleNamespace(hf_config=SimpleNamespace(model_type="qwen3")),
quant_config=None,
cache_config=SimpleNamespace(cache_dtype="auto"),
),
)
monkeypatch.setattr(attn_utils, "_is_dsv4_model", lambda _cfg: False)
monkeypatch.setattr(attn_utils, "enable_sfa", lambda _cfg: False)

raw = attn_utils._allocate_kv_cache(kv_cache_config, shared_layers={}, device="cpu")
assert isinstance(raw[layer_name], torch.Tensor)
assert raw[layer_name].numel() == tensor_size

attn_groups = [
AttentionGroup(
backend=CacheOnlyAttentionBackend,
layer_names=[layer_name],
kv_cache_spec=spec,
kv_cache_group_id=0,
)
]
reshaped = attn_utils._reshape_kv_cache_v2(
attn_groups=attn_groups,
kv_cache_raw_tensors=raw,
cache_dtype="auto",
kernel_block_sizes=[block_size],
shared_kv_cache_layers={},
kv_cache_config=kv_cache_config,
)
cache = reshaped[layer_name]
assert isinstance(cache, torch.Tensor)
# vLLM #51718 standardized cache-only writes as kv_cache[block, :, pos].
assert cache.shape == (num_blocks, num_kv_heads, block_size, head_size)
assert cache.dtype == dtype


class _PrefillStateBuilder:
def __init__(self):
self.extra_kwargs = None
Expand Down
Loading
Loading