diff --git a/docs/source/user_guide/feature_guide/speculative_decoding.md b/docs/source/user_guide/feature_guide/speculative_decoding.md index c8bb7ae18511..410f4be4ffb4 100644 --- a/docs/source/user_guide/feature_guide/speculative_decoding.md +++ b/docs/source/user_guide/feature_guide/speculative_decoding.md @@ -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 diff --git a/tests/e2e/pull_request/one_card/spec_decode/test_extract_hidden_states.py b/tests/e2e/pull_request/one_card/spec_decode/test_extract_hidden_states.py index 6a50dcdeff1c..eaeff65c2b73 100644 --- a/tests/e2e/pull_request/one_card/spec_decode/test_extract_hidden_states.py +++ b/tests/e2e/pull_request/one_card/spec_decode/test_extract_hidden_states.py @@ -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 @@ -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" @@ -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 @@ -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 = [ @@ -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", + ), ] @@ -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, @@ -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, diff --git a/tests/ut/worker/test_attn_utils_v2.py b/tests/ut/worker/test_attn_utils_v2.py index 1aead6bb0f42..a0c918a2ccc7 100644 --- a/tests/ut/worker/test_attn_utils_v2.py +++ b/tests/ut/worker/test_attn_utils_v2.py @@ -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, @@ -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 diff --git a/tests/ut/worker/test_extract_hidden_states_speculator_v2.py b/tests/ut/worker/test_extract_hidden_states_speculator_v2.py new file mode 100644 index 000000000000..4fb80ca6759e --- /dev/null +++ b/tests/ut/worker/test_extract_hidden_states_speculator_v2.py @@ -0,0 +1,140 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved. +"""Unit tests for MRV2 extract_hidden_states dispatch (upstream speculator). + +Covers Ascend init_speculator returning upstream ExtractHiddenStatesSpeculator +and upstream propose() behavior (vLLM PR #49811, dp_sync from #53694). +""" + +from __future__ import annotations + +from contextlib import nullcontext +from types import SimpleNamespace +from typing import Any, cast + +import pytest +import torch +import vllm.v1.worker.gpu.spec_decode.extract_hidden_states as upstream_spec_module +from vllm.v1.worker.gpu.spec_decode.extract_hidden_states import ( + ExtractHiddenStatesSpeculator, +) + +from vllm_ascend.worker.v2.spec_decode import init_speculator + + +class _RecordingModel(torch.nn.Module): + 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, + SimpleNamespace(speculative_config=SimpleNamespace(method="extract_hidden_states")), + ) + device = torch.device("cpu") + + def fake_speculator(config, target_device): + return config, target_device + + monkeypatch.setattr( + "vllm.v1.worker.gpu.spec_decode.extract_hidden_states.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(upstream_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 = ExtractHiddenStatesSpeculator.propose( + speculator, + 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 kwargs["num_tokens_across_dp"] is None + 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"): + ExtractHiddenStatesSpeculator.propose( + speculator, + 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_ascend/worker/v2/README.md b/vllm_ascend/worker/v2/README.md index 079231fabb93..093ec92163e2 100644 --- a/vllm_ascend/worker/v2/README.md +++ b/vllm_ascend/worker/v2/README.md @@ -37,3 +37,17 @@ to get specific plans. __ of AutoRegressiveAclGraphManager. Location: `speculator.AscendEagleSpeculator.init_cudagraph_manager`. + +- [x] `extract_hidden_states` (MRV2) + + Why: Upstream vLLM added Model Runner V2 support for + `extract_hidden_states`. Ascend reuses upstream + `ExtractHiddenStatesSpeculator` via `init_speculator` dispatch and keeps + `HiddenStateCacheSpec` on a private single-tensor allocate/reshape path. + `use_aux_hidden_state_outputs` is enabled by upstream + `GPUModelRunner.__init__`. + + After vLLM #51718 (0828 pin), every KV descriptor is a view into one + shared backing with `[B, H, N, C]` pages. Hidden-state dumps stay on + per-layer private buffers so they cannot overlay hybrid Attention/Mamba + storage, and reshape matches `CacheOnlyAttentionLayer.basic_cache`. diff --git a/vllm_ascend/worker/v2/attn_utils.py b/vllm_ascend/worker/v2/attn_utils.py index 6f901ddb67e4..3d987afd2261 100644 --- a/vllm_ascend/worker/v2/attn_utils.py +++ b/vllm_ascend/worker/v2/attn_utils.py @@ -64,6 +64,7 @@ enable_sfa, enable_sfa_dcp_replicated_indexer, get_kv_cache_tensor_layers, + is_hidden_state_cache_spec, vllm_version_is, ) @@ -691,6 +692,28 @@ def _allocate_kv_cache( example_layer_name = shared_names[0] example_spec = layer_kv_cache_spec[example_layer_name] + # extract_hidden_states dumps are live at the same time as the target + # model's Attention/Mamba caches. Keep HiddenStateCacheSpec off the + # #51718 hybrid backing so float32 SSM writes cannot overlay bfloat16 + # hidden states, and size each dump from its own page. + if any(is_hidden_state_cache_spec(layer_kv_cache_spec[ln]) for ln in shared_names): + for layer_idx, layer_name in enumerate(shared_names): + layer_spec = layer_kv_cache_spec[layer_name] + if is_hidden_state_cache_spec(layer_spec) or hybrid_backing is None: + kv_cache_raw_tensors[layer_name] = _allocate_int8_cache_tensor( + kv_cache_config.num_blocks * layer_spec.page_size_bytes, + alignment, + device, + ) + continue + layer_size = kv_cache_config.num_blocks * layer_spec.page_size_bytes + start = kv_cache_tensor.offset + layer_idx * kv_cache_tensor.layer_stride + end = start + layer_size + if end > hybrid_backing.numel(): + raise ValueError(f"Hybrid KV cache view for {layer_name} exceeds the backing allocation.") + kv_cache_raw_tensors[layer_name] = hybrid_backing[start:end] + continue + if hybrid_backing is not None: for layer_idx, layer_name in enumerate(shared_names): layer_spec = layer_kv_cache_spec[layer_name] @@ -1007,6 +1030,44 @@ def _reshape_kv_cache_v2( continue raw_cache = kv_cache_raw_tensors[layer_name] + if is_hidden_state_cache_spec(kv_cache_spec): + # Single tensor for extract_hidden_states (no K/V split). + # HiddenStateCacheSpec subclasses MLAAttentionSpec, so this + # must run before the generic MLA reshape path. + if not isinstance(raw_cache, torch.Tensor): + raise ValueError(f"Hidden-state cache for {layer_name} must use one raw tensor.") + if raw_cache.numel() % kv_cache_spec.page_size_bytes: + raise ValueError(f"KV cache for {layer_name} is not a whole number of pages.") + num_blocks = raw_cache.numel() // kv_cache_spec.page_size_bytes + if num_blocks < kv_cache_config.num_blocks: + raise ValueError(f"Hidden-state cache for {layer_name} has fewer blocks than KVCacheManager.") + # CacheOnlyAttentionBackend dropped get_kv_cache_shape in #51718. + # Spec properties already give the [B, H, N, C] layout that + # basic_cache writes as kv_cache[block, :, offset]. + kv_cache_shape = ( + num_blocks, + kv_cache_spec.num_heads, + kv_cache_spec.num_states, + kv_cache_spec.state_content_size_bytes // get_dtype_size(kv_cache_spec.dtype), + ) + typed_cache = raw_cache.view(kv_cache_spec.dtype) + page_size_padded = getattr(kv_cache_spec, "page_size_padded", None) + if page_size_padded is not None: + dtype_size = get_dtype_size(kv_cache_spec.dtype) + page_stride = page_size_padded // dtype_size + strides = [1] * len(kv_cache_shape) + for dim_idx in range(len(kv_cache_shape) - 2, -1, -1): + strides[dim_idx] = strides[dim_idx + 1] * kv_cache_shape[dim_idx + 1] + strides[0] = page_stride + kv_caches[layer_name] = torch.as_strided( + typed_cache, + size=kv_cache_shape, + stride=tuple(strides), + ) + else: + kv_caches[layer_name] = typed_cache.view(kv_cache_shape) + continue + if is_dsv4_model and isinstance( kv_cache_spec, (AscendMLAAttentionSpec, AscendSlidingWindowMLASpec), diff --git a/vllm_ascend/worker/v2/spec_decode/__init__.py b/vllm_ascend/worker/v2/spec_decode/__init__.py index 36869a4c8c7d..341ee753b67b 100644 --- a/vllm_ascend/worker/v2/spec_decode/__init__.py +++ b/vllm_ascend/worker/v2/spec_decode/__init__.py @@ -29,6 +29,14 @@ def init_speculator( """ speculative_config = vllm_config.speculative_config assert speculative_config is not None + if speculative_config.method == "extract_hidden_states": + # No Ascend-specific behavior beyond update_stream assignment in + # NPUModelRunner; reuse upstream ExtractHiddenStatesSpeculator as-is. + from vllm.v1.worker.gpu.spec_decode.extract_hidden_states import ( + ExtractHiddenStatesSpeculator, + ) + + return ExtractHiddenStatesSpeculator(vllm_config, device) if speculative_config.use_dspark(): from vllm_ascend.worker.v2.spec_decode.dspark.speculator import ( AscendDSparkSpeculator,