diff --git a/tests/v1/attention/test_flashinfer_plan_from_bounds.py b/tests/v1/attention/test_flashinfer_plan_from_bounds.py new file mode 100644 index 000000000000..c1a6e61e505f --- /dev/null +++ b/tests/v1/attention/test_flashinfer_plan_from_bounds.py @@ -0,0 +1,477 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""FlashInfer plans from the CPU bounds on seq_lens instead of copying seq_lens +from the device.""" + +import itertools +import types +import unittest.mock +from typing import Any + +import pytest + +from vllm.platforms import current_platform + +if not current_platform.is_cuda(): + pytest.skip("FlashInfer backend requires a CUDA platform.", allow_module_level=True) + +import flashinfer +import torch + +from tests.v1.attention.utils import ( + BatchSpec, + create_common_attn_metadata, + create_standard_kv_cache_spec, + create_vllm_config, +) +from vllm.config import set_current_vllm_config +from vllm.utils.math_utils import cdiv +from vllm.v1.attention.backends.flashinfer import ( + FlashInferMetadataBuilder, + _PinnedPlanWorkspaces, +) +from vllm.v1.attention.backends.utils import CommonAttentionMetadata, PerLayerParameters + +NUM_QO_HEADS, NUM_KV_HEADS, HEAD_SIZE = 8, 2, 128 +NUM_SPEC = 3 +BLOCK_SIZE = 16 +# Six decodes, two draft verifications and a prefill chunk. The upper bounds +# count up to NUM_SPEC drafts that may still be rejected; rows 3 and 5 get a +# page too many. Lower bounds are NUM_SPEC below, except on the prefill. +QUERY_LENS = [1] * 6 + [1 + NUM_SPEC] * 2 + [17] +EXACT = [1328, 18, 463, 64, 65, 127, 300, 1000, 2000] +UPPER = [1328, 21, 466, 65, 67, 129, 303, 1003, 2000] +# GPU clock cycles to spin so that an event recorded afterwards is still pending. +PENDING_CYCLES = 200_000_000 + + +@pytest.mark.parametrize( + "split", ["default", "disabled", "split_every_block_size_pages"] +) +@pytest.mark.parametrize("block_size", [16, 64]) +@pytest.mark.parametrize("kind", ["prefill", "decode"]) +def test_fa2_plan_from_upper_bound_matches_exact_plan( + kind: str, block_size: int, split: str +) -> None: + """fa2 planned from the upper bound and then given the exact last-page + lengths attends like fa2 planned from the exact lengths.""" + torch.manual_seed(0) + exact = [1328, 18, 463, 64, 65, 127] + if kind == "prefill": + # Several query tokens: the upper bound stays on the same page. + qo_lens = [1, 4, 1, 3, 2, 1] + upper = [n + min(NUM_SPEC, -n % block_size) for n in exact] + else: + # One query token: the upper bound may add one or two pages. + qo_lens = [1] * len(exact) + upper = [n + s for n, s in zip(exact, [0, 3, 6, 1, 5, 2 * block_size])] + extra: dict[str, Any] = {} + if split == "disabled": + extra["disable_split_kv"] = True + elif split == "split_every_block_size_pages": + # fixed_split_size counts pages: chunks of block_size pages each. + extra["fixed_split_size"] = block_size + pages = [cdiv(n, block_size) for n in upper] + block_ids = torch.randperm(4096, dtype=torch.int32, device="cuda") + rows = block_ids[: sum(pages)].split(pages) + kv_cache = torch.randn( + 4096, + 2, + block_size, + NUM_KV_HEADS, + HEAD_SIZE, + dtype=torch.bfloat16, + device="cuda", + ) + query = torch.randn( + sum(qo_lens), NUM_QO_HEADS, HEAD_SIZE, dtype=torch.bfloat16, device="cuda" + ) + qo_indptr = torch.tensor([0, *itertools.accumulate(qo_lens)], dtype=torch.int32) + + def run(lens: list[int]) -> torch.Tensor: + num_pages = [cdiv(n, block_size) for n in lens] + indptr = torch.tensor([0, *itertools.accumulate(num_pages)], dtype=torch.int32) + indices = torch.cat([row[:p] for row, p in zip(rows, num_pages)]) + last_page_len = torch.tensor( + [n - (p - 1) * block_size for n, p in zip(lens, num_pages)], + dtype=torch.int32, + ) + workspace = torch.empty(128 * 1024 * 1024, dtype=torch.uint8, device="cuda") + args = (NUM_QO_HEADS, NUM_KV_HEADS, HEAD_SIZE, block_size) + dtypes = dict(q_data_type=torch.bfloat16, kv_data_type=torch.bfloat16) + if kind == "prefill": + wrapper = flashinfer.BatchPrefillWithPagedKVCacheWrapper( + workspace, "NHD", backend="fa2" + ) + wrapper.plan( + qo_indptr, + indptr, + indices, + last_page_len, + *args, + causal=True, + **dtypes, + **extra, + ) + else: + wrapper = flashinfer.BatchDecodeWithPagedKVCacheWrapper( + workspace, "NHD", use_tensor_cores=True, backend="fa2" + ) + wrapper.plan( + indptr, + indices, + last_page_len, + *args, + pos_encoding_mode="NONE", + **dtypes, + **extra, + ) + # What the builder writes after planning from the upper bound. + exact_last_page_len = [n - (p - 1) * block_size for n, p in zip(exact, pages)] + if lens is upper: + wrapper._paged_kv_last_page_len_buf[: len(exact)].copy_( + torch.tensor(exact_last_page_len, dtype=torch.int32) + ) + return wrapper.run(query, kv_cache) + + output, expected = run(upper), run(exact) + if split == "disabled": + assert torch.equal(output, expected) + else: + # Surplus pages can form extra empty chunks, changing the summation order. + torch.testing.assert_close(output, expected, atol=1e-2, rtol=1e-2) + + +class _Fa2PrefillWrapper(flashinfer.BatchPrefillWithPagedKVCacheWrapper): + def __init__(self, *args, **kwargs): + super().__init__(*args, **{**kwargs, "backend": "fa2"}) + + +class _Fa2DecodeWrapper(flashinfer.BatchDecodeWithPagedKVCacheWrapper): + def __init__(self, *args, **kwargs): + super().__init__(*args, **{**kwargs, "backend": "fa2"}) + + +@pytest.fixture(autouse=True) +def _fa2_wrappers(monkeypatch): + """GPUs differ in the kernels FlashInfer picks; the builder tests need fa2.""" + module = "vllm.v1.attention.backends.flashinfer." + monkeypatch.setattr( + module + "BatchPrefillWithPagedKVCacheWrapper", _Fa2PrefillWrapper + ) + monkeypatch.setattr( + module + "BatchDecodeWithPagedKVCacheWrapper", _Fa2DecodeWrapper + ) + + +def _make_builder(vllm_config) -> FlashInferMetadataBuilder: + # Keep every row on the FlashInfer wrappers, also where TRTLLM is the default. + vllm_config.attention_config.use_trtllm_attention = False + head_size = vllm_config.model_config.get_head_size() + + def per_layer_parameters(vllm_config, layer_names, impl_cls): + params = PerLayerParameters( + window_left=-1, + logits_soft_cap=0.0, + sm_scale=head_size**-0.5, + has_sinks=False, + ) + return {name: params for name in layer_names} + + with ( + set_current_vllm_config(vllm_config), + unittest.mock.patch( + "vllm.v1.attention.backends.flashinfer.get_per_layer_parameters", + per_layer_parameters, + ), + ): + builder = FlashInferMetadataBuilder( + create_standard_kv_cache_spec(vllm_config), + ["model.layers.0.self_attn.attn"], + vllm_config, + torch.device("cuda"), + ) + # The bounds in these tests count drafts, as with speculative decoding. + # Without it the builder takes the upper bound as exact. + builder._num_speculative_tokens = NUM_SPEC + return builder + + +@pytest.mark.parametrize( + "case, copies", + [ + ("fa2", False), + ("straddling_verification_row", True), + ("not_fa2", True), + ("not_fa2_exact_bounds", False), + ("no_lower_bound", True), + ], +) +def test_build_copies_seq_lens_only_when_the_bounds_do_not_suffice( + case: str, copies: bool, monkeypatch +) -> None: + vllm_config = create_vllm_config( + model_name="Qwen/Qwen3-0.6B", block_size=BLOCK_SIZE, max_model_len=4096 + ) + builder = _make_builder(vllm_config) + exact, upper = list(EXACT), list(UPPER) + if case == "straddling_verification_row": + exact[7], upper[7] = 1008, 1009 + lower = [n - NUM_SPEC for n in upper[:-1]] + upper[-1:] + if case == "not_fa2_exact_bounds": + upper = lower = exact + cam = _mixed_metadata( + exact, upper, None if case == "no_lower_bound" else lower, vllm_config + ) + + def build(): + with set_current_vllm_config(vllm_config): + return builder.build(common_prefix_len=0, common_attn_metadata=cam) + + # Warm up the kernels, both wrappers and their pinned buffer rings. + for _ in range(6): + build() + if case.startswith("not_fa2"): + # As if the wrappers ran kernels that take the lengths from the plan. + monkeypatch.setattr( + "vllm.v1.attention.backends.flashinfer._reads_kv_lens_from_device", + lambda wrapper: False, + ) + torch.accelerator.synchronize() + # Any read of seq_lens back from the device now raises. + torch.cuda.set_sync_debug_mode("error") + try: + if copies: + with pytest.raises(RuntimeError, match="synchroniz"): + build() + return + attn_metadata = build() + finally: + torch.cuda.set_sync_debug_mode("default") + torch.accelerator.synchronize() + + exact_gpu = torch.tensor(exact, dtype=torch.int32, device="cuda") + num_decodes = QUERY_LENS.count(1) + for wrapper, start, stop in ( + (attn_metadata.decode.wrapper, 0, num_decodes), + (attn_metadata.prefill.wrapper, num_decodes, len(exact)), + ): + indptr = wrapper._paged_kv_indptr_buf[: stop - start + 1] + last_page_len = wrapper._paged_kv_last_page_len_buf[: stop - start] + lens = (indptr[1:] - indptr[:-1] - 1) * BLOCK_SIZE + last_page_len + assert torch.equal(lens, exact_gpu[start:stop]) + if case == "fa2": + indptr = builder.paged_kv_indptr.gpu[: len(exact) + 1] + assert torch.any(indptr[1:] - indptr[:-1] > cdiv(exact_gpu, BLOCK_SIZE)) + + +def _mixed_metadata( + exact: list[int], upper: list[int], lower: list[int] | None, vllm_config +) -> CommonAttentionMetadata: + """Decodes, verifications and a prefill (QUERY_LENS) with exact seq_lens on + the device and the CPU bounds.""" + cam = create_common_attn_metadata( + BatchSpec(seq_lens=upper, query_lens=QUERY_LENS), + BLOCK_SIZE, + torch.device("cuda"), + max_block_idx=vllm_config.cache_config.num_gpu_blocks, + ) + return cam.replace( + seq_lens=torch.tensor(exact, dtype=torch.int32, device="cuda"), + seq_lens_cpu_lower_bound=( + None if lower is None else torch.tensor(lower, dtype=torch.int32) + ), + ) + + +def test_seq_lens_bounds_check_does_not_synchronize() -> None: + """With VLLM_DEBUG_SEQ_LENS_BOUNDS the builder asserts on the device that + the exact seq_lens lie within the bounds it planned from, without reading + anything back.""" + vllm_config = create_vllm_config( + model_name="Qwen/Qwen3-0.6B", block_size=BLOCK_SIZE, max_model_len=4096 + ) + builder = _make_builder(vllm_config) + builder._check_seq_lens_bounds = True + lower = [n - NUM_SPEC for n in UPPER[:-1]] + UPPER[-1:] + cam = _mixed_metadata(list(EXACT), list(UPPER), lower, vllm_config) + + def build(): + with set_current_vllm_config(vllm_config): + return builder.build(common_prefix_len=0, common_attn_metadata=cam) + + for _ in range(6): + build() + torch.accelerator.synchronize() + torch.cuda.set_sync_debug_mode("error") + try: + build() + finally: + torch.cuda.set_sync_debug_mode("default") + # The bounds hold: the device-side assertion stays quiet. + torch.accelerator.synchronize() + + +def _decode_metadata( + upper: list[int], lower: list[int], vllm_config +) -> CommonAttentionMetadata: + cam = create_common_attn_metadata( + BatchSpec(seq_lens=upper, query_lens=[1] * len(upper)), + BLOCK_SIZE, + torch.device("cuda"), + max_block_idx=vllm_config.cache_config.num_gpu_blocks, + ) + return cam.replace(seq_lens_cpu_lower_bound=torch.tensor(lower, dtype=torch.int32)) + + +def _fake_wrapper(backend: str = "fa2", **missing: bool) -> types.SimpleNamespace: + wrapper = types.SimpleNamespace(_backend=backend) + for name in ("_pin_memory_int_workspace_buffer", "_paged_kv_last_page_len_buf"): + if not missing.get(name): + setattr(wrapper, name, torch.empty(0)) + return wrapper + + +@pytest.mark.parametrize( + "case, expected", + [ + ("fa2", "upper"), + ("wrapper_not_created_yet", None), + ("backend_not_picked_yet", None), + ("fa3", None), + ("fa3_exact_bounds", "exact"), + ("no_pinned_buffer", None), + ("no_pinned_buffer_exact_bounds", None), + ("no_last_page_len_buffer", None), + ("lower_above_upper", None), + ], +) +def test_bounds_are_used_only_for_the_wrapper_that_can_take_them( + case: str, expected: str | None +) -> None: + vllm_config = create_vllm_config( + model_name="Qwen/Qwen3-0.6B", block_size=BLOCK_SIZE, max_model_len=4096 + ) + builder = _make_builder(vllm_config) + upper = [100, 203, 300] + lower = [n - NUM_SPEC for n in upper] + if case.endswith("exact_bounds"): + lower = upper + elif case == "lower_above_upper": + lower[1] = upper[1] + 1 + wrapper = { + "wrapper_not_created_yet": None, + "backend_not_picked_yet": _fake_wrapper("auto"), + "fa3": _fake_wrapper("fa3"), + "fa3_exact_bounds": _fake_wrapper("fa3"), + "no_pinned_buffer": _fake_wrapper(_pin_memory_int_workspace_buffer=True), + "no_pinned_buffer_exact_bounds": _fake_wrapper( + _pin_memory_int_workspace_buffer=True + ), + "no_last_page_len_buffer": _fake_wrapper(_paged_kv_last_page_len_buf=True), + }.get(case, _fake_wrapper()) + + from_bounds = builder._seq_lens_cpu_from_bounds( + _decode_metadata(upper, lower, vllm_config), + num_decodes=len(upper), + decode_uses_trtllm=False, + prefill_uses_trtllm=False, + decode_wrapper=wrapper, + prefill_wrapper=None, + ) + + if expected is None: + assert from_bounds is None + else: + seq_lens_cpu, seq_lens_exact = from_bounds + assert seq_lens_cpu.tolist() == upper + assert seq_lens_exact == (expected == "exact") + + +def test_bounds_wait_for_the_first_plan_of_an_auto_wrapper() -> None: + """A wrapper created with backend="auto" picks its kernels on its first + plan(). bf16 tensor-core decode picks fa2 on every GPU.""" + vllm_config = create_vllm_config( + model_name="Qwen/Qwen3-0.6B", block_size=BLOCK_SIZE, max_model_len=4096 + ) + builder = _make_builder(vllm_config) + upper = [100, 203, 300] + cam = _decode_metadata(upper, [n - NUM_SPEC for n in upper], vllm_config) + wrapper = flashinfer.BatchDecodeWithPagedKVCacheWrapper( + torch.empty(128 * 1024 * 1024, dtype=torch.uint8, device="cuda"), + "NHD", + use_tensor_cores=True, + ) + + def from_bounds(): + return builder._seq_lens_cpu_from_bounds( + cam, len(upper), False, False, decode_wrapper=wrapper, prefill_wrapper=None + ) + + assert from_bounds() is None + pages = [cdiv(n, BLOCK_SIZE) for n in upper] + wrapper.plan( + torch.tensor([0, *itertools.accumulate(pages)], dtype=torch.int32), + torch.arange(sum(pages), dtype=torch.int32, device="cuda"), + torch.tensor([n - (p - 1) * BLOCK_SIZE for n, p in zip(upper, pages)]).int(), + NUM_QO_HEADS, + NUM_KV_HEADS, + HEAD_SIZE, + BLOCK_SIZE, + q_data_type=torch.bfloat16, + kv_data_type=torch.bfloat16, + ) + assert from_bounds() is not None + + +def test_plan_without_a_pinned_buffer_is_called_directly() -> None: + vllm_config = create_vllm_config( + model_name="Qwen/Qwen3-0.6B", block_size=BLOCK_SIZE, max_model_len=4096 + ) + builder = _make_builder(vllm_config) + calls: list[dict[str, int]] = [] + builder._plan(object(), lambda **kwargs: calls.append(kwargs), batch_size=3) + assert calls == [{"batch_size": 3}] + + +class _Wrapper: + pass + + +def test_pinned_plan_workspace_is_not_reused_while_its_copy_is_pending() -> None: + workspaces = _PinnedPlanWorkspaces() + wrapper = _Wrapper() + own = torch.zeros(1024, dtype=torch.uint8, pin_memory=True) + buf, event = workspaces.acquire(wrapper, own) + assert buf is own + torch.cuda._sleep(PENDING_CYCLES) + event.record() + + second, second_event = workspaces.acquire(wrapper, own) + assert second is not own + assert second.is_pinned() and not second.any() + second_event.record() + torch.accelerator.synchronize() + # Once the copies have run, both buffers are reused in turn. + assert all(workspaces.acquire(wrapper, own)[0] is b for b in (own, second, own)) + + +def test_pinned_plan_workspaces_add_at_most_max_extra_buffers(monkeypatch) -> None: + monkeypatch.setattr(_PinnedPlanWorkspaces, "MAX_EXTRA_BUFFERS", 1) + workspaces = _PinnedPlanWorkspaces() + first, second = _Wrapper(), _Wrapper() + own = [torch.zeros(1024, dtype=torch.uint8, pin_memory=True) for _ in range(2)] + for wrapper, pinned in zip((first, second), own): + _, event = workspaces.acquire(wrapper, pinned) + torch.cuda._sleep(PENDING_CYCLES) + event.record() + + # The first wrapper gets the one extra buffer. + extra, event = workspaces.acquire(first, own[0]) + assert extra is not own[0] + event.record() + # The second waits for its copy instead of adding a buffer. + buf, event = workspaces.acquire(second, own[1]) + assert buf is own[1] + event.record() + torch.accelerator.synchronize() diff --git a/tests/v1/worker/test_gpu_ubatch_slicing.py b/tests/v1/worker/test_gpu_ubatch_slicing.py index ebf334e45d4f..2dea5a7fbbb6 100644 --- a/tests/v1/worker/test_gpu_ubatch_slicing.py +++ b/tests/v1/worker/test_gpu_ubatch_slicing.py @@ -121,6 +121,9 @@ def _make_input_batch( query_start_loc_np=query_start_loc_np, seq_lens=buffers.seq_lens[:num_reqs_padded], seq_lens_cpu_upper_bound=torch.from_numpy(seq_lens_upper_bound), + seq_lens_cpu_lower_bound=torch.from_numpy( + np.maximum(seq_lens_upper_bound - 2, 0) + ), num_computed_tokens_np=np.array(seq_lens, dtype=np.int32) - np.array(query_lens, dtype=np.int32), prefill_len_np=np.zeros(num_reqs, dtype=np.int32), @@ -187,6 +190,11 @@ def test_slicing_matches_v1_split_attn_metadata(batch_name: str): torch.testing.assert_close( v2_ubatch.seq_lens_cpu_upper_bound, v1_ubatch.seq_lens_cpu_upper_bound ) + # The lower bound is sliced and truncated like the upper bound. + torch.testing.assert_close( + v2_ubatch.seq_lens_cpu_lower_bound, + (v2_ubatch.seq_lens_cpu_upper_bound - 2).clamp(min=0), + ) def test_microbatches_do_not_share_buffers(): diff --git a/tests/v1/worker/test_gpu_warmup_blocks.py b/tests/v1/worker/test_gpu_warmup_blocks.py index e0a99d5e8f8e..4aab30767e4d 100644 --- a/tests/v1/worker/test_gpu_warmup_blocks.py +++ b/tests/v1/worker/test_gpu_warmup_blocks.py @@ -89,6 +89,7 @@ def _make_runner( return SimpleNamespace( num_speculative_steps=num_spec_steps, adaptive_verification=None, + emit_seq_lens_cpu_lower_bound=True, rejection_sampler=None, decode_query_len=num_spec_steps + 1, is_pooling_model=False, diff --git a/tests/v1/worker/test_mamba_hybrid_model_state.py b/tests/v1/worker/test_mamba_hybrid_model_state.py index fa54ccfcedc5..7b0402049aee 100644 --- a/tests/v1/worker/test_mamba_hybrid_model_state.py +++ b/tests/v1/worker/test_mamba_hybrid_model_state.py @@ -43,6 +43,7 @@ def test_prepare_attn_forwards_positions(monkeypatch: pytest.MonkeyPatch) -> Non num_scheduled_tokens=torch.tensor([1], dtype=torch.int32), max_query_len=None, seq_lens_cpu_upper_bound=torch.tensor([1537], dtype=torch.int32), + seq_lens_cpu_lower_bound=None, seq_lens=torch.tensor([1537], dtype=torch.int32), is_prefilling_np=torch.tensor([False]).numpy(), prefill_runs_as_decode_np=None, @@ -117,6 +118,7 @@ def test_padded_prompt_tail_builds_as_spec_decode( num_draft_tokens_per_req=np.array([k, k, 0], dtype=np.int32), max_query_len=None, seq_lens_cpu_upper_bound=torch.tensor(seq_lens, dtype=torch.int32), + seq_lens_cpu_lower_bound=None, seq_lens=torch.tensor(seq_lens, dtype=torch.int32), is_prefilling_np=np.array(is_prefilling), prefill_runs_as_decode_np=np.array([False, True, False]), @@ -190,6 +192,7 @@ def _input_batch( num_scheduled_tokens=np.array(num_scheduled_tokens, dtype=np.int32), max_query_len=None, seq_lens_cpu_upper_bound=seq_lens, + seq_lens_cpu_lower_bound=None, is_prefilling_np=np.array(is_prefilling), prefill_runs_as_decode_np=None, num_draft_tokens_per_req=None diff --git a/tests/v1/worker/test_mixed_warmup_gate.py b/tests/v1/worker/test_mixed_warmup_gate.py index ba532c2c06b4..7cd98190a875 100644 --- a/tests/v1/worker/test_mixed_warmup_gate.py +++ b/tests/v1/worker/test_mixed_warmup_gate.py @@ -74,11 +74,13 @@ def test_kernel_warmup_restores_uncalibrated_adaptive_manager(monkeypatch, fail_ runner = SimpleNamespace( adaptive_verification=manager, rejection_sampler=rejection_sampler, + emit_seq_lens_cpu_lower_bound=True, ) def run_steps(model_runner, execute, sample): assert model_runner.adaptive_verification is None assert not model_runner.rejection_sampler.enable_adaptive_verification + assert not model_runner.emit_seq_lens_cpu_lower_bound if fail_warmup: raise RuntimeError("warmup failed") @@ -91,3 +93,4 @@ def run_steps(model_runner, execute, sample): assert runner.adaptive_verification is manager assert manager.cost_tables is None assert rejection_sampler.enable_adaptive_verification + assert runner.emit_seq_lens_cpu_lower_bound diff --git a/tests/v1/worker/test_seq_len_bounds.py b/tests/v1/worker/test_seq_len_bounds.py new file mode 100644 index 000000000000..f373202fba3d --- /dev/null +++ b/tests/v1/worker/test_seq_len_bounds.py @@ -0,0 +1,96 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""The CPU bounds on seq_lens that FlashInfer plans from, and the check of the +exact seq_lens against them. No GPU needed.""" + +import numpy as np +import pytest +import torch + +from vllm import envs +from vllm.v1.attention.backends.utils import check_seq_lens_bounds +from vllm.v1.worker.gpu.model_runner import compute_seq_lens_cpu_lower_bound +from vllm.v1.worker.gpu.spec_decode.speculator import ( + compute_draft_seq_lens_cpu_lower_bound, +) + +NUM_SPEC = 3 +NUM_REQS = 6 +# Six requests padded to eight. Row 3 is prefilling; rows 4 and 5 are the +# first tokens of a request, where the drafts in flight would go below zero. +UPPER = [1328, 21, 466, 65, 2, 3, 0, 0] +IS_PREFILLING = [False, False, False, True, False, False] +LOWER = [1325, 18, 463, 65, 0, 0, 0, 0] +EXACT = [1327, 20, 463, 65, 1, 3] + + +def test_runner_lower_bound_subtracts_the_drafts_in_flight() -> None: + upper_np = np.array(UPPER, dtype=np.int32) + lower_np = compute_seq_lens_cpu_lower_bound( + upper_np, np.array(IS_PREFILLING), NUM_SPEC, NUM_REQS + ) + assert lower_np.tolist() == LOWER + assert lower_np.dtype == np.int32 + # A new array; the upper bound is left as it was. + assert upper_np.tolist() == UPPER + + +def test_runner_lower_bound_without_drafts_is_the_upper_bound() -> None: + upper_np = np.array(UPPER, dtype=np.int32) + lower_np = compute_seq_lens_cpu_lower_bound( + upper_np, np.array(IS_PREFILLING), 0, NUM_REQS + ) + assert lower_np.tolist() == UPPER + assert lower_np is not upper_np + + +@pytest.mark.parametrize( + "step, expected", + [ + (1, [1323, 16, 461, 63, 0, 0]), + (2, [1324, 17, 462, 64, 0, 0]), + (3, [1325, 18, 463, 65, 0, 0]), + ], +) +def test_draft_lower_bound_per_step(step: int, expected: list[int]) -> None: + """Draft step ``step`` has ``step`` more tokens than the target lower bound + and up to NUM_SPEC fewer: the verification may have rejected that many.""" + target_lower = torch.tensor(LOWER, dtype=torch.int32) + draft_lower = compute_draft_seq_lens_cpu_lower_bound( + target_lower, step, NUM_SPEC, NUM_REQS, num_reqs_padded=8 + ) + assert draft_lower.tolist() == expected + [0, 0] + assert draft_lower.dtype == torch.int32 + assert target_lower.tolist() == LOWER + + +def _bounds() -> tuple[torch.Tensor, torch.Tensor]: + return ( + torch.tensor(LOWER, dtype=torch.int32), + torch.tensor(UPPER, dtype=torch.int32), + ) + + +def test_check_passes_within_the_bounds() -> None: + seq_lens = torch.tensor(EXACT, dtype=torch.int32) + # Longer bounds are cut to the requests; the bounds are inclusive. + check_seq_lens_bounds(seq_lens, *_bounds()) + check_seq_lens_bounds(seq_lens, seq_lens, seq_lens) + + +@pytest.mark.parametrize( + "row, value", + [(0, 1329), (1, 17), (3, 64), (3, 66), (4, 3), (5, -1)], +) +def test_check_raises_outside_the_bounds(row: int, value: int) -> None: + seq_lens = torch.tensor(EXACT, dtype=torch.int32) + seq_lens[row] = value + with pytest.raises(RuntimeError): + check_seq_lens_bounds(seq_lens, *_bounds()) + + +def test_check_is_off_by_default(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("VLLM_DEBUG_SEQ_LENS_BOUNDS", raising=False) + assert envs.VLLM_DEBUG_SEQ_LENS_BOUNDS is False + monkeypatch.setenv("VLLM_DEBUG_SEQ_LENS_BOUNDS", "1") + assert envs.VLLM_DEBUG_SEQ_LENS_BOUNDS is True diff --git a/vllm/envs.py b/vllm/envs.py index 851741d5253b..6ff82770f625 100755 --- a/vllm/envs.py +++ b/vllm/envs.py @@ -298,6 +298,7 @@ VLLM_NCCL_INCLUDE_PATH: str | None = None VLLM_GC_DEBUG: str = "" VLLM_DEBUG_WORKSPACE: bool = False + VLLM_DEBUG_SEQ_LENS_BOUNDS: bool = False VLLM_DISABLE_SHARED_EXPERTS_STREAM: bool = False VLLM_DISABLE_DSV4_MEGAMOE_SHARED_EXPERT_FUSION: bool = False VLLM_SHARED_EXPERTS_STREAM_TOKEN_THRESHOLD: int = 256 @@ -2078,6 +2079,13 @@ def _resolve_rust_cli_path() -> str | None: # Debug workspace allocations. # logging of workspace resize operations. "VLLM_DEBUG_WORKSPACE": lambda: bool(int(os.getenv("VLLM_DEBUG_WORKSPACE", "0"))), + # Check on the device that the exact seq_lens lie within the CPU bounds + # FlashInfer plans from. Does not synchronize. Debug aid only: on CUDA a + # violation is a device-side assertion that leaves the CUDA context + # unusable. The FlashInfer builder reads this once at construction. + "VLLM_DEBUG_SEQ_LENS_BOUNDS": lambda: bool( + int(os.getenv("VLLM_DEBUG_SEQ_LENS_BOUNDS", "0")) + ), # Disables parallel execution of shared_experts via separate cuda stream "VLLM_DISABLE_SHARED_EXPERTS_STREAM": lambda: bool( int(os.getenv("VLLM_DISABLE_SHARED_EXPERTS_STREAM", "0")) diff --git a/vllm/v1/attention/backend.py b/vllm/v1/attention/backend.py index 3387a31794c9..ea4c6b9d0187 100644 --- a/vllm/v1/attention/backend.py +++ b/vllm/v1/attention/backend.py @@ -454,6 +454,10 @@ class CommonAttentionMetadata: table. Rows of one request are adjacent, so equal neighbours are PCP chunks sharing one KV context.""" + seq_lens_cpu_lower_bound: torch.Tensor | None = None + """(batch_size,) CPU lower bound on seq_lens. It differs from the upper + bound only while drafts of a step still in flight may be rejected.""" + mm_req_doc_ranges: dict[int, list[tuple[int, int]]] | None = None """PrefixLM bidirectional ranges for multimodal tokens. Maps request index to list of (start, end) token position ranges @@ -483,6 +487,9 @@ def naive_query_lens(self) -> torch.Tensor: return self.query_start_loc[1:] - self.query_start_loc[:-1] def replace(self, **kwargs) -> "CommonAttentionMetadata": + if "seq_lens_cpu_upper_bound" in kwargs or "seq_lens" in kwargs: + # A lower bound not updated along with seq_lens is stale. + kwargs.setdefault("seq_lens_cpu_lower_bound", None) return replace(self, **kwargs) def compute_num_computed_tokens(self) -> torch.Tensor: @@ -543,6 +550,7 @@ def unpadded( self.dcp_local_seq_lens_cpu_upper_bound ), seq_lens_cpu_upper_bound=maybe_slice_reqs(self.seq_lens_cpu_upper_bound), + seq_lens_cpu_lower_bound=maybe_slice_reqs(self.seq_lens_cpu_lower_bound), is_prefilling=maybe_slice_reqs(self.is_prefilling), req_idx=maybe_slice_reqs(self.req_idx), rswa_prefix_lens=maybe_slice_reqs(self.rswa_prefix_lens), diff --git a/vllm/v1/attention/backends/flashinfer.py b/vllm/v1/attention/backends/flashinfer.py index a21720ab26e6..baa42b42e3e1 100755 --- a/vllm/v1/attention/backends/flashinfer.py +++ b/vllm/v1/attention/backends/flashinfer.py @@ -2,6 +2,8 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project """Attention layer with FlashInfer.""" +import weakref +from collections.abc import Callable from dataclasses import dataclass, replace from enum import Enum from functools import partial @@ -74,6 +76,7 @@ max_decode_query_len, ) from vllm.v1.attention.backends.utils import ( + check_seq_lens_bounds, get_dcp_local_seq_lens, get_flashinfer_layout_string, get_num_attention_heads_from_layers, @@ -745,6 +748,64 @@ class FlashInferMetadata: cascade_wrapper: MultiLevelCascadeAttentionWrapper | None +class _PinnedPlanWorkspaces: + """Per-wrapper rings of pinned host buffers for plan() to stage through. + + plan() copies its pinned buffer to the GPU asynchronously. When build() + does not wait for the GPU, the next plan() could overwrite that buffer + before the copy has run, so a new zeroed buffer is used instead of + waiting. Buffers are never shared between wrappers: under CUDA graphs the + kernels read entries past the planned batch, which must be zeros or left + over from plans of the same shape. At most MAX_EXTRA_BUFFERS buffers are + added over all rings; after that a ring waits for its oldest copy. + """ + + MAX_EXTRA_BUFFERS = 8 + + def __init__(self): + self._rings: weakref.WeakKeyDictionary[ + object, list[tuple[torch.Tensor, torch.cuda.Event]] + ] = weakref.WeakKeyDictionary() + self._next: weakref.WeakKeyDictionary[object, int] = weakref.WeakKeyDictionary() + self._num_extra = 0 + + def acquire( + self, wrapper, pinned: torch.Tensor + ) -> tuple[torch.Tensor, torch.cuda.Event]: + ring = self._rings.get(wrapper) + if ring is None: + ring = [(pinned, torch.cuda.Event())] + self._rings[wrapper] = ring + i = self._next.get(wrapper, 0) + if ring[i][1].query(): + pass + elif self._num_extra >= self.MAX_EXTRA_BUFFERS: + with gpu_sync_allowed(): + ring[i][1].synchronize() + else: + ring.insert( + i, + ( + torch.zeros( + pinned.shape, dtype=pinned.dtype, device="cpu", pin_memory=True + ), + torch.cuda.Event(), + ), + ) + self._num_extra += 1 + self._next[wrapper] = (i + 1) % len(ring) + return ring[i] + + +def _reads_kv_lens_from_device(wrapper) -> bool: + """Whether wrapper runs FlashInfer's fa2 kernels, which take each request's + KV length from its last-page length on the device. A wrapper created with + backend="auto" picks its kernels on its first plan() and keeps them.""" + return getattr(wrapper, "_backend", None) == "fa2" and hasattr( + wrapper, "_paged_kv_last_page_len_buf" + ) + + class FlashInferMetadataBuilder(AttentionMetadataBuilder[FlashInferMetadata]): kv_cache_spec: AttentionSpec reorder_batch_threshold: int = 1 @@ -999,6 +1060,13 @@ def __init__( self.paged_kv_last_page_len = CpuGpuBuffer( max_num_reqs, dtype=torch.int32, device=self.device, pin_memory=False ) + # Exact last-page lengths, computed on the device when planning from the + # CPU upper bound on seq_lens. + self.paged_kv_last_page_len_exact = torch.zeros( + max_num_reqs, dtype=torch.int32, device=self.device + ) + self._plan_workspaces = _PinnedPlanWorkspaces() + self._check_seq_lens_bounds = envs.VLLM_DEBUG_SEQ_LENS_BOUNDS @property def kv_cache_layout(self) -> KVCacheLayout: @@ -1377,12 +1445,20 @@ def _compute_flashinfer_kv_metadata( block_table_tensor: torch.Tensor, num_reqs: int, page_size: int, + device_seq_lens: torch.Tensor, ) -> torch.Tensor: """Compute paged_kv_indptr, paged_kv_indices and paged_kv_last_page_len. Results are stored in self.paged_kv_indptr, self.paged_kv_indices, self.paged_kv_last_page_len buffers. + ``seq_lens_np`` are the CPU lengths to plan from, exact or the upper + bound. ``device_seq_lens`` are the exact lengths on the device; the + last-page lengths they imply for this page table are written to + self.paged_kv_last_page_len_exact. That buffer is read only when + ``seq_lens_np`` is the upper bound, never with DCP or cascade + attention, where ``seq_lens_np`` counts other tokens. + Returns paged_kv_indices, a GPU tensor with shape [num_actual_pages]. """ # write self.paged_kv_indptr_cpu inplace (0-index is always 0) @@ -1399,6 +1475,9 @@ def _compute_flashinfer_kv_metadata( block_table_tensor, block_table_tensor.stride(0), paged_kv_indptr, + device_seq_lens, + self.paged_kv_last_page_len_exact, + page_size, BLOCK_SIZE=1024, ) @@ -1412,6 +1491,80 @@ def _compute_flashinfer_kv_metadata( self.paged_kv_last_page_len.copy_to_gpu(num_reqs) return paged_kv_indices + def _seq_lens_cpu_from_bounds( + self, + common_attn_metadata: CommonAttentionMetadata, + num_decodes: int, + decode_uses_trtllm: bool, + prefill_uses_trtllm: bool, + decode_wrapper, + prefill_wrapper, + ) -> tuple[torch.Tensor, bool] | None: + """CPU seq_lens to plan from without reading them back from the device, + and whether they are exact. None if the CPU bounds do not allow it. + + decode_wrapper and prefill_wrapper are the wrappers this build plans, or + None if they do not exist yet. The plan uses the upper bound. fa2 + wrappers get the exact lengths on the device after plan(), so their + rows may count a page too many, except rows with several query tokens: + fa2 lays out split-KV from the planned lengths. Other FlashInfer kernels + need exact bounds. TRTLLM reads the lengths from the device. + """ + upper = common_attn_metadata.seq_lens_cpu_upper_bound + lower = common_attn_metadata.seq_lens_cpu_lower_bound + if upper is None or lower is None: + return None + num_reqs = common_attn_metadata.num_reqs + upper_np, lower_np = upper.numpy()[:num_reqs], lower.numpy()[:num_reqs] + if min(len(upper_np), len(lower_np)) < num_reqs or (lower_np > upper_np).any(): + return None + exact = upper_np == lower_np + page_size = self.page_size + same_page = (lower_np + page_size - 1) // page_size == ( + upper_np + page_size - 1 + ) // page_size + query_start_loc = common_attn_metadata.query_start_loc_cpu.numpy() + single_query = np.diff(query_start_loc[: num_reqs + 1]) == 1 + fa2_ok = same_page | single_query + all_exact = True + for start, stop, uses_trtllm, wrapper in ( + (0, num_decodes, decode_uses_trtllm, decode_wrapper), + (num_decodes, num_reqs, prefill_uses_trtllm, prefill_wrapper), + ): + if uses_trtllm or start == stop: + continue + # Without the wait for the GPU, plan() must stage through the ring. + if not hasattr(wrapper, "_pin_memory_int_workspace_buffer"): + return None + if exact[start:stop].all(): + continue + if not (_reads_kv_lens_from_device(wrapper) and fa2_ok[start:stop].all()): + return None + all_exact = False + return upper, all_exact + + def _plan(self, wrapper, plan: Callable[..., None], **kwargs) -> None: + """Run ``plan(**kwargs)`` for ``wrapper`` with a pinned buffer from its + ring (see _PinnedPlanWorkspaces).""" + pinned = getattr(wrapper, "_pin_memory_int_workspace_buffer", None) + if pinned is None: + plan(**kwargs) + return + buf, event = self._plan_workspaces.acquire(wrapper, pinned) + wrapper._pin_memory_int_workspace_buffer = buf + try: + plan(**kwargs) + finally: + event.record() + + def _write_exact_last_page_len(self, wrapper, start: int, num: int) -> None: + """Hand a wrapper planned from the CPU upper bound the exact lengths: + fa2 derives each request's KV length from its last-page length. It + does not read the KV lengths plan() uploads; only trtllm-gen does.""" + wrapper._paged_kv_last_page_len_buf[:num].copy_( + self.paged_kv_last_page_len_exact[start : start + num] + ) + def build( self, common_prefix_len: int, @@ -1534,13 +1687,46 @@ def build( # seq_lens_cpu is not needed since TRTLLM paths use GPU tensors # (block_tables, seq_lens) directly. needs_seq_lens_cpu = self.use_dcp or use_cascade or not all_uses_trtllm + decode_uses_cudagraph = ( + self.enable_cuda_graph + and num_prefills == 0 + and num_decode_tokens <= self._decode_cudagraph_max_bs + ) + seq_lens_exact = True if needs_seq_lens_cpu: + from_bounds = None if ( self._num_speculative_tokens == 0 and common_attn_metadata.seq_lens_cpu_upper_bound is not None ): # No speculative token accounting, so the upper bound is exact. - seq_lens_cpu = common_attn_metadata.seq_lens_cpu_upper_bound + from_bounds = (common_attn_metadata.seq_lens_cpu_upper_bound, True) + elif not (self.use_dcp or use_cascade or self.has_sinks): + from_bounds = self._seq_lens_cpu_from_bounds( + common_attn_metadata, + num_decodes, + decode_with_flashinfer_trtllm_api, + prefill_use_trtllm, + decode_wrapper=( + self._decode_wrappers_cudagraph.get(num_decode_tokens) + if decode_uses_cudagraph + else self._decode_wrapper + ), + prefill_wrapper=( + self._prefill_wrapper + if causal + else self._noncausal_prefill_wrapper + ), + ) + if from_bounds is not None: + seq_lens_cpu, seq_lens_exact = from_bounds + if self._check_seq_lens_bounds: + assert common_attn_metadata.seq_lens_cpu_lower_bound is not None + check_seq_lens_bounds( + seq_lens[:num_reqs], + common_attn_metadata.seq_lens_cpu_lower_bound, + seq_lens_cpu, + ) else: with gpu_sync_allowed(): seq_lens_cpu = common_attn_metadata.seq_lens.cpu() @@ -1605,6 +1791,7 @@ def build( block_table_tensor, num_reqs, page_size, + device_seq_lens=seq_lens, ) else: paged_kv_indices = None @@ -1793,7 +1980,9 @@ def build( o_dtype = ( FP8_DTYPE if self.nvfp4_trtllm else self.model_config.dtype ) - prefill_wrapper.plan( + self._plan( + prefill_wrapper, + prefill_wrapper.plan, qo_indptr=qo_indptr_prefill_cpu, paged_kv_indptr=paged_kv_indptr_prefill_cpu, paged_kv_indices=paged_kv_indices, @@ -1813,6 +2002,10 @@ def build( fixed_split_size=self.prefill_fixed_split_size, disable_split_kv=self.disable_split_kv, ) + if not seq_lens_exact: + self._write_exact_last_page_len( + prefill_wrapper, prefill_start, num_prefills + ) attn_metadata.prefill = FIPrefill(wrapper=prefill_wrapper) ## DECODE PATHWAY @@ -1861,16 +2054,10 @@ def build( ) else: assert seq_lens_cpu is not None - pure_decode = num_prefills == 0 - use_cudagraph = ( - self.enable_cuda_graph - and pure_decode - and num_decode_tokens <= self._decode_cudagraph_max_bs - ) num_input_tokens = num_decode_tokens decode_wrapper = self._get_decode_wrapper( - num_input_tokens, use_cudagraph + num_input_tokens, decode_uses_cudagraph ) # Use the persistent buffer with padding length, # instead of the same address but chunked version @@ -1899,8 +2086,9 @@ def build( ) if PIN_MEMORY: kv_lens_decode_cpu = kv_lens_decode_cpu.pin_memory() - fast_plan_decode( + self._plan( decode_wrapper, + partial(fast_plan_decode, decode_wrapper), indptr_cpu=paged_kv_indptr_cpu, indices=paged_kv_indices, last_page_len_cpu=paged_kv_last_page_len_cpu, @@ -1920,6 +2108,8 @@ def build( fixed_split_size=self.decode_fixed_split_size, disable_split_kv=self.disable_split_kv, ) + if not seq_lens_exact: + self._write_exact_last_page_len(decode_wrapper, 0, num_decodes) attn_metadata.decode = FIDecode(wrapper=decode_wrapper) return attn_metadata @@ -2889,6 +3079,9 @@ def _copy_page_indices_kernel( block_table, block_table_stride, cu_num_blocks, + seq_lens, + last_page_len, + page_size, BLOCK_SIZE: tl.constexpr, ): req_idx = tl.program_id(0) @@ -2905,3 +3098,7 @@ def _copy_page_indices_kernel( block_ids, mask=i + offset < num_blocks, ) + + # Zero or negative when num_blocks counts a page past the exact length. + seq_len = tl.load(seq_lens + req_idx) + tl.store(last_page_len + req_idx, seq_len - (num_blocks - 1) * page_size) diff --git a/vllm/v1/attention/backends/utils.py b/vllm/v1/attention/backends/utils.py index f93156d53fe5..6b5b917653d8 100644 --- a/vllm/v1/attention/backends/utils.py +++ b/vllm/v1/attention/backends/utils.py @@ -1112,6 +1112,36 @@ def compute_causal_conv1d_metadata( return nums_dict, batch_ptr, token_chunk_offset_ptr +def check_seq_lens_bounds( + seq_lens: torch.Tensor, lower: torch.Tensor, upper: torch.Tensor +) -> None: + """Assert lower <= seq_lens <= upper element-wise without synchronizing. + + ``seq_lens`` are the exact lengths, on any device. ``lower`` and ``upper`` + may be longer and may live on the CPU; they are copied to the device of + ``seq_lens`` asynchronously. On the CPU a violation raises at once. On + CUDA it surfaces as a device-side assertion at the next synchronization, + which leaves the CUDA context unusable: a debug aid, not a recoverable + check. Enabled with VLLM_DEBUG_SEQ_LENS_BOUNDS, which the FlashInfer + builder reads once at construction. + """ + num_reqs = seq_lens.shape[0] + + def on_device(bound: torch.Tensor) -> torch.Tensor: + bound = bound[:num_reqs] + if bound.device == seq_lens.device: + return bound + if PIN_MEMORY and not bound.is_pinned(): + bound = bound.pin_memory() + return bound.to(seq_lens.device, non_blocking=True) + + lower, upper = on_device(lower), on_device(upper) + torch._assert_async( + ((seq_lens >= lower) & (seq_lens <= upper)).all(), + "seq_lens outside the CPU bounds FlashInfer was planned from", + ) + + def get_dcp_local_seq_lens( seq_lens: torch.Tensor, dcp_size: int = 1, diff --git a/vllm/v1/worker/gpu/attn_utils.py b/vllm/v1/worker/gpu/attn_utils.py index 4fa308bcf471..a4fa52c67fe6 100644 --- a/vllm/v1/worker/gpu/attn_utils.py +++ b/vllm/v1/worker/gpu/attn_utils.py @@ -440,6 +440,7 @@ def build_attn_metadata( dcp_local_seq_lens_cpu_upper_bound: torch.Tensor | None = None, positions: torch.Tensor | None = None, is_prefilling: torch.Tensor | None = None, + seq_lens_cpu_lower_bound: torch.Tensor | None = None, mm_req_doc_ranges: dict[int, list[tuple[int, int]]] | None = None, model_specific_attn_metadata: ModelSpecificAttnMetadata | None = None, for_cudagraph_capture: bool = False, @@ -458,6 +459,8 @@ def build_attn_metadata( ] if seq_lens_cpu_upper_bound is not None: seq_lens_cpu_upper_bound = seq_lens_cpu_upper_bound[:num_reqs] + if seq_lens_cpu_lower_bound is not None: + seq_lens_cpu_lower_bound = seq_lens_cpu_lower_bound[:num_reqs] attn_metadata: dict[str, Any] = {} token_to_req_indices: torch.Tensor | None = None @@ -497,6 +500,7 @@ def build_attn_metadata( query_start_loc_cpu=query_start_loc_cpu, seq_lens=seq_lens, seq_lens_cpu_upper_bound=seq_lens_cpu_upper_bound, + seq_lens_cpu_lower_bound=seq_lens_cpu_lower_bound, max_seq_len=max_seq_len, num_reqs=num_reqs, num_actual_tokens=num_tokens, diff --git a/vllm/v1/worker/gpu/input_batch.py b/vllm/v1/worker/gpu/input_batch.py index e2f68ad59865..f6225175e36e 100644 --- a/vllm/v1/worker/gpu/input_batch.py +++ b/vllm/v1/worker/gpu/input_batch.py @@ -123,6 +123,9 @@ class InputBatch: # None if there are no prefills. prefill_runs_as_decode_np: np.ndarray | None = None + # [num_reqs] CPU lower bound on seq_lens (see CommonAttentionMetadata). + seq_lens_cpu_lower_bound: torch.Tensor | None = None + @classmethod def make_dummy( cls, diff --git a/vllm/v1/worker/gpu/model_runner.py b/vllm/v1/worker/gpu/model_runner.py index 373042a41502..50347cb08f98 100644 --- a/vllm/v1/worker/gpu/model_runner.py +++ b/vllm/v1/worker/gpu/model_runner.py @@ -185,6 +185,31 @@ logger = init_logger(__name__) +def compute_seq_lens_cpu_lower_bound( + seq_lens_cpu_upper_bound_np: np.ndarray, + is_prefilling_np: np.ndarray, + num_speculative_steps: int, + num_reqs: int, +) -> np.ndarray: + """CPU lower bound on seq_lens, as a new array. + + The upper bound counts the drafts of the step in flight as accepted; up to + num_speculative_steps of them may be rejected, none on rows still + prefilling. Entries past num_reqs are copied as they are. + """ + lower_bound_np = seq_lens_cpu_upper_bound_np.copy() + if num_speculative_steps > 0: + lower = lower_bound_np[:num_reqs] + np.subtract( + lower, + num_speculative_steps, + out=lower, + where=~is_prefilling_np[:num_reqs], + ) + np.maximum(lower, 0, out=lower) + return lower_bound_np + + class GPUModelRunner(LoRAModelRunnerMixin): def __init__(self, vllm_config: VllmConfig, device: torch.device): self.vllm_config = vllm_config @@ -286,6 +311,8 @@ def __init__(self, vllm_config: VllmConfig, device: torch.device): # Speculative decoding. self.speculator = None self.use_aux_hidden_state_outputs = False + # Off during kernel warm-up, whose steps never subtract rejected drafts. + self.emit_seq_lens_cpu_lower_bound = True if self.speculative_config is not None: if self.is_last_pp_rank: self.speculator = init_speculator( @@ -1489,6 +1516,25 @@ def prepare_inputs( ) seq_lens_cpu_upper_bound = torch.from_numpy(seq_lens_cpu_upper_bound_np) + # CPU lower bound on seq_lens (see compute_seq_lens_cpu_lower_bound). + # There is no such bound with adaptive verification, with more steps + # in flight or with pipeline parallelism. + seq_lens_cpu_lower_bound = None + if ( + self.emit_seq_lens_cpu_lower_bound + and adaptive_verification is None + and self.vllm_config.max_concurrent_batches <= 2 + and not self.use_pp + ): + seq_lens_cpu_lower_bound = torch.from_numpy( + compute_seq_lens_cpu_lower_bound( + seq_lens_cpu_upper_bound_np, + batch_req_state.is_prefilling_np, + self.num_speculative_steps, + num_reqs, + ) + ) + prompt_lens = None if self.model_config.rswa_window is not None: # prompt_lens is only used in R-SWA case. @@ -1511,6 +1557,7 @@ def prepare_inputs( query_start_loc_np=query_start_loc_np, seq_lens=seq_lens, seq_lens_cpu_upper_bound=seq_lens_cpu_upper_bound, + seq_lens_cpu_lower_bound=seq_lens_cpu_lower_bound, dcp_local_seq_lens=None, num_computed_tokens_np=num_computed_tokens_np, prefill_len_np=batch_req_state.prefill_len_np, diff --git a/vllm/v1/worker/gpu/model_states/default.py b/vllm/v1/worker/gpu/model_states/default.py index 35ca9ccba102..6766fff4aa38 100644 --- a/vllm/v1/worker/gpu/model_states/default.py +++ b/vllm/v1/worker/gpu/model_states/default.py @@ -223,6 +223,7 @@ def prepare_attn( slot_mappings=slot_mappings, kv_cache_config=kv_cache_config, seq_lens_cpu_upper_bound=seq_lens_cpu_upper_bound, + seq_lens_cpu_lower_bound=input_batch.seq_lens_cpu_lower_bound, dcp_local_seq_lens=input_batch.dcp_local_seq_lens, dcp_local_seq_lens_cpu_upper_bound=input_batch.dcp_local_seq_lens_cpu_upper_bound, positions=input_batch.positions, diff --git a/vllm/v1/worker/gpu/model_states/mamba_hybrid.py b/vllm/v1/worker/gpu/model_states/mamba_hybrid.py index c45304efa512..fb15d6b16a5d 100644 --- a/vllm/v1/worker/gpu/model_states/mamba_hybrid.py +++ b/vllm/v1/worker/gpu/model_states/mamba_hybrid.py @@ -328,6 +328,7 @@ def prepare_attn( slot_mappings=slot_mappings, kv_cache_config=kv_cache_config, seq_lens_cpu_upper_bound=seq_lens_cpu_upper_bound, + seq_lens_cpu_lower_bound=input_batch.seq_lens_cpu_lower_bound, dcp_local_seq_lens=input_batch.dcp_local_seq_lens, positions=input_batch.positions, model_specific_attn_metadata=mamba_attn_metadata, diff --git a/vllm/v1/worker/gpu/pcp_manager.py b/vllm/v1/worker/gpu/pcp_manager.py index 6aec96bfe14b..12683efc8ae1 100644 --- a/vllm/v1/worker/gpu/pcp_manager.py +++ b/vllm/v1/worker/gpu/pcp_manager.py @@ -665,6 +665,7 @@ def partition_batch( query_start_loc_np=local_query_start_loc_np[: num_reqs_after_padding + 1], seq_lens=seq_lens, seq_lens_cpu_upper_bound=torch.from_numpy(seq_lens_cpu_upper_bound_np), + seq_lens_cpu_lower_bound=None, dcp_local_seq_lens=None, dcp_local_seq_lens_cpu_upper_bound=dcp_local_seq_lens_cpu_upper_bound, num_computed_tokens_np=local_start_pos_np, diff --git a/vllm/v1/worker/gpu/spec_decode/speculator.py b/vllm/v1/worker/gpu/spec_decode/speculator.py index d07c043de915..85074bd9b443 100644 --- a/vllm/v1/worker/gpu/spec_decode/speculator.py +++ b/vllm/v1/worker/gpu/spec_decode/speculator.py @@ -44,6 +44,29 @@ logger = init_logger(__name__) +def compute_draft_seq_lens_cpu_lower_bound( + seq_lens_cpu_lower_bound: torch.Tensor, + step: int, + num_speculative_steps: int, + num_reqs: int, + num_reqs_padded: int, +) -> torch.Tensor: + """CPU lower bound on seq_lens for draft step ``step``, as a new tensor. + + The target lower bound already allows for the drafts of the step in + flight. The verification just run may reject up to num_speculative_steps + more, which the target bounds still count. Padded entries are zero. + """ + draft_lower_bound = torch.zeros(num_reqs_padded, dtype=torch.int32, device="cpu") + torch.add( + seq_lens_cpu_lower_bound[:num_reqs], + step - num_speculative_steps, + out=draft_lower_bound[:num_reqs], + ) + draft_lower_bound[:num_reqs].clamp_(min=0) + return draft_lower_bound + + def _target_feeds_hc_residual(vllm_config: VllmConfig) -> bool: """Whether the target replaces the drafter's input with its HC residual. @@ -310,6 +333,7 @@ def _build_attn_metadata( causal: bool | Mapping[int, bool] = True, dcp_local_seq_lens: torch.Tensor | None = None, slot_mappings: torch.Tensor | None = None, + seq_lens_cpu_lower_bound: torch.Tensor | None = None, ) -> dict[str, Any] | None: num_reqs_padded = batch_desc.num_reqs or num_reqs # A FULL graph replays a captured shape whose padded requests each hold @@ -342,6 +366,15 @@ def _build_attn_metadata( out=draft_seq_lens_cpu_upper_bound[:num_reqs], ) draft_seq_lens_cpu_upper_bound[:num_reqs].clamp_(max=self.max_model_len) + draft_seq_lens_cpu_lower_bound = None + if seq_lens_cpu_lower_bound is not None: + draft_seq_lens_cpu_lower_bound = compute_draft_seq_lens_cpu_lower_bound( + seq_lens_cpu_lower_bound, + step, + self.num_speculative_steps, + num_reqs, + num_reqs_padded, + ) if dcp_local_seq_lens is None and self.dcp_size > 1: # Draft steps advance and rewind their own global sequence lengths, # so the target model's DCP-local lengths may already be stale. @@ -376,6 +409,7 @@ def _build_attn_metadata( seq_lens_cpu_upper_bound=draft_seq_lens_cpu_upper_bound, positions=self.input_buffers.positions[:num_tokens], is_prefilling=self.draft_is_prefilling[:num_reqs_padded], + seq_lens_cpu_lower_bound=draft_seq_lens_cpu_lower_bound, ) return attn_metadata @@ -551,6 +585,7 @@ def _build_uniform_attn_metadata( step: int, causal: bool | Mapping[int, bool] = True, dcp_local_seq_lens: torch.Tensor | None = None, + seq_lens_cpu_lower_bound: torch.Tensor | None = None, ) -> dict[str, Any] | None: query_start_loc_np = self.arange_np[: num_reqs + 1] * num_query_per_req return self._build_attn_metadata( @@ -561,6 +596,7 @@ def _build_uniform_attn_metadata( step=step, causal=causal, dcp_local_seq_lens=dcp_local_seq_lens, + seq_lens_cpu_lower_bound=seq_lens_cpu_lower_bound, ) def _update_draft_decode_metadata( diff --git a/vllm/v1/worker/gpu/spec_decode/target_dependent_ar/speculator.py b/vllm/v1/worker/gpu/spec_decode/target_dependent_ar/speculator.py index 9a8ed7b8b4ca..814d8bf6fad9 100644 --- a/vllm/v1/worker/gpu/spec_decode/target_dependent_ar/speculator.py +++ b/vllm/v1/worker/gpu/spec_decode/target_dependent_ar/speculator.py @@ -392,6 +392,7 @@ def propose( num_tokens_across_dp, input_batch.seq_lens_cpu_upper_bound, num_speculative_tokens, + input_batch.seq_lens_cpu_lower_bound, ) self.on_multi_step_decode_end(num_reqs) @@ -519,6 +520,7 @@ def _multi_step_decode( num_tokens_across_dp: torch.Tensor | None, seq_lens_cpu_upper_bound: torch.Tensor, num_speculative_steps: int, + seq_lens_cpu_lower_bound: torch.Tensor | None = None, ) -> None: positions = self.input_buffers.positions[:num_reqs] query_start_loc = self.input_buffers.query_start_loc[: num_reqs + 1] @@ -545,6 +547,7 @@ def _multi_step_decode( num_query_per_req=1, seq_lens_cpu_upper_bound=seq_lens_cpu_upper_bound, step=step, + seq_lens_cpu_lower_bound=seq_lens_cpu_lower_bound, ) self.current_draft_step.fill_(step) @@ -571,6 +574,7 @@ def _fused_multi_step_decode( num_tokens_across_dp: torch.Tensor | None, seq_lens_cpu_upper_bound: torch.Tensor, num_speculative_steps: int, + seq_lens_cpu_lower_bound: torch.Tensor | None = None, ) -> None: positions = self.input_buffers.positions[:num_reqs] query_start_loc = self.input_buffers.query_start_loc[: num_reqs + 1] @@ -599,6 +603,7 @@ def _fused_multi_step_decode( num_query_per_req=1, seq_lens_cpu_upper_bound=seq_lens_cpu_upper_bound, step=1, + seq_lens_cpu_lower_bound=seq_lens_cpu_lower_bound, ) self._generate_fused_drafts( diff --git a/vllm/v1/worker/gpu/ubatch_utils.py b/vllm/v1/worker/gpu/ubatch_utils.py index 3fa9783bf50c..560cd48a3db7 100644 --- a/vllm/v1/worker/gpu/ubatch_utils.py +++ b/vllm/v1/worker/gpu/ubatch_utils.py @@ -142,6 +142,10 @@ def _slice_input_batch( tokens_truncated = max(0, int(input_batch.query_start_loc_np[req_stop]) - tok_stop) if tokens_truncated: seq_lens_cpu_upper_bound[-1] -= tokens_truncated + seq_lens_cpu_lower_bound = input_batch.seq_lens_cpu_lower_bound + if seq_lens_cpu_lower_bound is not None: + seq_lens_cpu_lower_bound = seq_lens_cpu_lower_bound[req_start:req_stop].clone() + seq_lens_cpu_lower_bound[-1:].sub_(tokens_truncated).clamp_(min=0) # Query lengths of the truncated requests, so consumers that derive # max_query_len from this array see the microbatch's own lengths. @@ -174,6 +178,7 @@ def _slice_input_batch( query_start_loc_np=query_start_loc_np, seq_lens=seq_lens, seq_lens_cpu_upper_bound=seq_lens_cpu_upper_bound, + seq_lens_cpu_lower_bound=seq_lens_cpu_lower_bound, dcp_local_seq_lens=dcp_local_seq_lens, num_computed_tokens_np=input_batch.num_computed_tokens_np[req_start:req_stop], prefill_len_np=input_batch.prefill_len_np[req_start:req_stop], diff --git a/vllm/v1/worker/gpu/warmup.py b/vllm/v1/worker/gpu/warmup.py index 894e2e5b6bd9..0dbcf5263947 100644 --- a/vllm/v1/worker/gpu/warmup.py +++ b/vllm/v1/worker/gpu/warmup.py @@ -239,6 +239,9 @@ def warmup_kernels( if adaptive_sampling: assert rejection_sampler is not None rejection_sampler.enable_adaptive_verification = False + # The synthetic steps below never subtract rejected drafts. + emit_lower_bound = model_runner.emit_seq_lens_cpu_lower_bound + model_runner.emit_seq_lens_cpu_lower_bound = False try: _warmup_kernels(model_runner, worker_execute_model, worker_sample_tokens) finally: @@ -246,6 +249,7 @@ def warmup_kernels( if adaptive_sampling: assert rejection_sampler is not None rejection_sampler.enable_adaptive_verification = True + model_runner.emit_seq_lens_cpu_lower_bound = emit_lower_bound def _warmup_kernels(