From c66aa48e993b74c46f83261654827b7349b2208c Mon Sep 17 00:00:00 2001 From: Woosuk Kwon Date: Thu, 26 Feb 2026 11:20:35 -0800 Subject: [PATCH 1/2] [Model Runner V2] Add model states [1/N] (#35350) Signed-off-by: Woosuk Kwon --- vllm/v1/worker/gpu/cudagraph_utils.py | 59 ++++++++----------- vllm/v1/worker/gpu/input_batch.py | 3 - vllm/v1/worker/gpu/model_runner.py | 84 +++++++-------------------- vllm/v1/worker/gpu/model_states.py | 74 +++++++++++++++++++++++ 4 files changed, 119 insertions(+), 101 deletions(-) create mode 100644 vllm/v1/worker/gpu/model_states.py diff --git a/vllm/v1/worker/gpu/cudagraph_utils.py b/vllm/v1/worker/gpu/cudagraph_utils.py index d70a4c7ab18..95369005d96 100644 --- a/vllm/v1/worker/gpu/cudagraph_utils.py +++ b/vllm/v1/worker/gpu/cudagraph_utils.py @@ -22,6 +22,7 @@ from vllm.v1.worker.gpu.attn_utils import ( from vllm.v1.worker.gpu.block_table import BlockTables from vllm.v1.worker.gpu.dp_utils import make_num_tokens_across_dp from vllm.v1.worker.gpu.input_batch import InputBuffers +from vllm.v1.worker.gpu.model_states import ModelState from vllm.v1.worker.utils import AttentionGroup @@ -29,13 +30,11 @@ class CudaGraphManager: def __init__( self, vllm_config: VllmConfig, - uses_mrope: bool, use_aux_hidden_state_outputs: bool, device: torch.device, ): self.vllm_config = vllm_config self.scheduler_config = vllm_config.scheduler_config - self.uses_mrope = uses_mrope self.use_aux_hidden_state_outputs = use_aux_hidden_state_outputs self.device = device @@ -88,8 +87,8 @@ class CudaGraphManager: num_tokens: int, capture_cg_mode: CUDAGraphMode, model: nn.Module, + model_state: ModelState, input_buffers: InputBuffers, - mrope_positions: torch.Tensor | None, inputs_embeds: torch.Tensor | None, block_tables: BlockTables, attn_groups: list[list[AttentionGroup]], @@ -113,13 +112,18 @@ class CudaGraphManager: ) else: num_reqs = min(num_tokens, self.max_num_reqs) - input_ids = input_buffers.input_ids[:num_tokens] - positions = input_buffers.positions[:num_tokens] - if self.uses_mrope: - assert mrope_positions is not None - positions = mrope_positions[:, :num_tokens] - if inputs_embeds is not None: - inputs_embeds = inputs_embeds[:num_tokens] + + model_inputs = { + "input_ids": input_buffers.input_ids[:num_tokens], + "positions": input_buffers.positions[:num_tokens], + "inputs_embeds": ( + inputs_embeds[:num_tokens] if inputs_embeds is not None else None + ), + # NOTE: Values returned by `prepare_dummy_inputs` will override the + # default values above. + **model_state.prepare_dummy_inputs(num_reqs, num_tokens), + } + attn_metadata, slot_mappings = prepare_inputs_to_capture( num_reqs, num_tokens, @@ -143,11 +147,7 @@ class CudaGraphManager: num_tokens_across_dp=num_tokens_across_dp, slot_mapping=slot_mappings, ): - model_output = model( - input_ids=input_ids, - positions=positions, - inputs_embeds=inputs_embeds, - ) + model_output = model(**model_inputs) if self.use_aux_hidden_state_outputs: hidden_states, aux_hidden_states = model_output else: @@ -164,9 +164,7 @@ class CudaGraphManager: num_tokens=num_tokens, num_reqs=num_reqs, model=model, - input_ids=input_ids, - positions=positions, - inputs_embeds=inputs_embeds, + model_inputs=model_inputs, num_tokens_across_dp=num_tokens_across_dp, attn_metadata=attn_metadata, slot_mappings=slot_mappings, @@ -178,9 +176,7 @@ class CudaGraphManager: num_tokens: int, num_reqs: int, model: nn.Module, - input_ids: torch.Tensor, - positions: torch.Tensor, - inputs_embeds: torch.Tensor | None, + model_inputs: dict[str, torch.Tensor | None], num_tokens_across_dp: torch.Tensor, attn_metadata: dict[str, Any] | None, slot_mappings: dict[str, torch.Tensor] | None, @@ -206,11 +202,8 @@ class CudaGraphManager: ), torch.cuda.graph(graph, self.pool), ): - model_output = model( - input_ids=input_ids, - positions=positions, - inputs_embeds=inputs_embeds, - ) + model_output = model(**model_inputs) + # Join offloader's copy stream after forward to avoid unjoined # stream error. The last layer's start_prefetch forks copy_stream, # but wait_prefetch only happens in the next forward pass. @@ -235,9 +228,7 @@ class CudaGraphManager: num_tokens: int, num_reqs: int, model: nn.Module, - input_ids: torch.Tensor, - positions: torch.Tensor, - inputs_embeds: torch.Tensor | None, + model_inputs: dict[str, torch.Tensor | None], num_tokens_across_dp: torch.Tensor, attn_metadata: dict[str, Any] | None, slot_mappings: dict[str, torch.Tensor] | None, @@ -256,18 +247,14 @@ class CudaGraphManager: batch_descriptor=batch_descriptor, slot_mapping=slot_mappings, ): - model( - input_ids=input_ids, - positions=positions, - inputs_embeds=inputs_embeds, - ) + model(**model_inputs) @torch.inference_mode() def capture( self, model: nn.Module, + model_state: ModelState, input_buffers: InputBuffers, - mrope_positions: torch.Tensor | None, inputs_embeds: torch.Tensor | None, block_tables: BlockTables, attn_groups: list[list[AttentionGroup]], @@ -278,8 +265,8 @@ class CudaGraphManager: device=self.device, capture_fn=self.capture_graph, model=model, + model_state=model_state, input_buffers=input_buffers, - mrope_positions=mrope_positions, inputs_embeds=inputs_embeds, block_tables=block_tables, attn_groups=attn_groups, diff --git a/vllm/v1/worker/gpu/input_batch.py b/vllm/v1/worker/gpu/input_batch.py index a15da926da4..87b8bbf18b8 100644 --- a/vllm/v1/worker/gpu/input_batch.py +++ b/vllm/v1/worker/gpu/input_batch.py @@ -65,8 +65,6 @@ class InputBatch: input_ids: torch.Tensor # [num_tokens_after_padding] positions: torch.Tensor - # [3, num_tokens_after_padding] - mrope_positions: torch.Tensor | None # [num_tokens_after_padding, hidden_size] inputs_embeds: torch.Tensor | None @@ -143,7 +141,6 @@ class InputBatch: seq_lens=seq_lens, input_ids=input_ids, positions=positions, - mrope_positions=None, inputs_embeds=None, attn_metadata=None, # type: ignore slot_mappings=None, # type: ignore diff --git a/vllm/v1/worker/gpu/model_runner.py b/vllm/v1/worker/gpu/model_runner.py index 26eb3ecf72a..949f09f54ba 100644 --- a/vllm/v1/worker/gpu/model_runner.py +++ b/vllm/v1/worker/gpu/model_runner.py @@ -77,7 +77,7 @@ from vllm.v1.worker.gpu.kv_connector import ( ) from vllm.v1.worker.gpu.lora_utils import LoraState from vllm.v1.worker.gpu.mm.encoder_runner import EncoderRunner -from vllm.v1.worker.gpu.mm.mrope_utils import MRopeState +from vllm.v1.worker.gpu.model_states import ModelState from vllm.v1.worker.gpu.pp_utils import pp_broadcast, pp_receive from vllm.v1.worker.gpu.sample.output import SamplerOutput from vllm.v1.worker.gpu.sample.prompt_logprob import PromptLogprobsWorker @@ -140,14 +140,6 @@ class GPUModelRunner(LoRAModelRunnerMixin): dtype=self.dtype, device=self.device, ) - self.uses_mrope = self.model_config.uses_mrope - if self.uses_mrope: - self.mrope_states = MRopeState( - max_num_reqs=self.max_num_reqs, - max_num_tokens=self.max_num_tokens, - max_model_len=self.max_model_len, - device=self.device, - ) self.use_async_scheduling = self.scheduler_config.async_scheduling self.output_copy_stream = torch.cuda.Stream(self.device) @@ -212,7 +204,6 @@ class GPUModelRunner(LoRAModelRunnerMixin): # CUDA graphs. self.cudagraph_manager = CudaGraphManager( self.vllm_config, - self.uses_mrope, self.use_aux_hidden_state_outputs, self.device, ) @@ -271,6 +262,9 @@ class GPUModelRunner(LoRAModelRunnerMixin): if self.speculator is not None: prepare_communication_buffer_for_model(self.speculator) + # Initialize the components that require the model. + self.model_state = ModelState(self.vllm_config, self.model, self.device) + def get_model(self) -> nn.Module: return self.model @@ -481,16 +475,13 @@ class GPUModelRunner(LoRAModelRunnerMixin): start_free_gpu_memory = torch.cuda.mem_get_info()[0] with self.maybe_setup_dummy_loras(self.lora_config): - mrope_positions = None - if self.uses_mrope: - mrope_positions = self.mrope_states.mrope_positions inputs_embeds = None if self.supports_mm_inputs: inputs_embeds = self.encoder_runner.inputs_embeds self.cudagraph_manager.capture( model=self.model, + model_state=self.model_state, input_buffers=self.input_buffers, - mrope_positions=mrope_positions, inputs_embeds=inputs_embeds, block_tables=self.block_tables, attn_groups=self.attn_groups, @@ -554,14 +545,7 @@ class GPUModelRunner(LoRAModelRunnerMixin): if self.supports_mm_inputs: self.encoder_runner.add_request(req_id, new_req_data.mm_features) - # Pre-compute M-RoPE positions for prefill. - if self.uses_mrope: - self.mrope_states.init_prefill_mrope_positions( - req_index, - self.model, # type: ignore - new_req_data.prefill_token_ids, - mm_features=new_req_data.mm_features, - ) + self.model_state.add_request(req_index, new_req_data) self.block_tables.append_block_ids( req_index, new_req_data.block_ids, overwrite=True @@ -577,8 +561,7 @@ class GPUModelRunner(LoRAModelRunnerMixin): if scheduler_output.scheduled_new_reqs: self.req_states.apply_staged_writes() self.sampler.apply_staged_writes() - if self.uses_mrope: - self.mrope_states.apply_staged_writes() + self.model_state.apply_staged_writes() def update_requests(self, scheduler_output: SchedulerOutput) -> None: # Add new blocks for the existing requests. @@ -692,15 +675,6 @@ class GPUModelRunner(LoRAModelRunnerMixin): ) dcp_local_seq_lens = self.input_buffers.dcp_local_seq_lens[:num_reqs] - # Prepare M-RoPE positions. - if self.uses_mrope: - self.mrope_states.prepare_mrope_positions( - idx_mapping, - query_start_loc, - self.req_states.prefill_len.gpu, - self.req_states.num_computed_tokens.gpu, - ) - # Some input token ids are directly read from the last sampled tokens # and draft tokens. Also, get the logits indices to sample tokens from. logits_indices = combine_sampled_and_draft_tokens( @@ -744,10 +718,6 @@ class GPUModelRunner(LoRAModelRunnerMixin): input_ids = self.input_buffers.input_ids[:num_tokens_after_padding] positions = self.input_buffers.positions[:num_tokens_after_padding] - mrope_positions = None - if self.uses_mrope: - mrope_positions = self.mrope_states.mrope_positions - mrope_positions = mrope_positions[:, :num_tokens_after_padding] return InputBatch( req_ids=req_ids, num_reqs=num_reqs, @@ -764,7 +734,6 @@ class GPUModelRunner(LoRAModelRunnerMixin): seq_lens=seq_lens, input_ids=input_ids, positions=positions, - mrope_positions=mrope_positions, inputs_embeds=None, attn_metadata=attn_metadata, slot_mappings=slot_mappings_by_layer, @@ -959,14 +928,24 @@ class GPUModelRunner(LoRAModelRunnerMixin): input_buffers=self.input_buffers, device=self.device, ) - if self.uses_mrope: - input_batch.mrope_positions = self.mrope_states.mrope_positions[ - :, :num_tokens_after_padding - ] if not skip_attn_for_dummy_run: self.prepare_dummy_attn_metadata(input_batch) # FIXME(woosuk): Fix warmup for LoRA. + model_inputs = { + "input_ids": input_batch.input_ids, + "positions": input_batch.positions, + "inputs_embeds": input_batch.inputs_embeds, + # NOTE: Values returned by `prepare_inputs` will override the default + # values above. + **self.model_state.prepare_inputs(input_batch, self.req_states), + } + if not self.is_first_pp_rank: + # Update for non-first PP ranks. + model_inputs["input_ids"] = None + model_inputs["inputs_embeds"] = None + model_inputs["intermediate_tensors"] = intermediate_tensors + # Run model. if cudagraph_runtime_mode == CUDAGraphMode.FULL: # Use explicit cudagraph replay for FULL mode. @@ -983,20 +962,6 @@ class GPUModelRunner(LoRAModelRunnerMixin): aux_hidden_states = None else: # For piecewise and eager mode, just call model(). - positions = input_batch.positions - if self.uses_mrope: - assert input_batch.mrope_positions is not None - positions = input_batch.mrope_positions - - if self.is_first_pp_rank: - input_ids = input_batch.input_ids - inputs_embeds = input_batch.inputs_embeds - assert intermediate_tensors is None - else: - input_ids = None - inputs_embeds = None - assert intermediate_tensors is not None - batch_descriptor = BatchDescriptor( num_tokens=input_batch.num_tokens_after_padding, has_lora=self.lora_config is not None, @@ -1012,12 +977,7 @@ class GPUModelRunner(LoRAModelRunnerMixin): slot_mapping=input_batch.slot_mappings, ): self.kv_connector.pre_forward(scheduler_output) - model_output = self.model( - input_ids=input_ids, - positions=positions, - inputs_embeds=inputs_embeds, - intermediate_tensors=intermediate_tensors, - ) + model_output = self.model(**model_inputs) if self.use_aux_hidden_state_outputs: hidden_states, aux_hidden_states = model_output else: diff --git a/vllm/v1/worker/gpu/model_states.py b/vllm/v1/worker/gpu/model_states.py new file mode 100644 index 00000000000..03574b2ad55 --- /dev/null +++ b/vllm/v1/worker/gpu/model_states.py @@ -0,0 +1,74 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +import torch +import torch.nn as nn + +from vllm.config import VllmConfig +from vllm.v1.core.sched.output import NewRequestData +from vllm.v1.worker.gpu.input_batch import InputBatch +from vllm.v1.worker.gpu.mm.mrope_utils import MRopeState +from vllm.v1.worker.gpu.states import RequestState + + +class ModelState: + def __init__(self, vllm_config: VllmConfig, model: nn.Module, device: torch.device): + self.vllm_config = vllm_config + self.model_config = vllm_config.model_config + self.scheduler_config = vllm_config.scheduler_config + self.model = model + self.device = device + + self.max_model_len = self.model_config.max_model_len + self.max_num_reqs = self.scheduler_config.max_num_seqs + self.max_num_tokens = self.scheduler_config.max_num_batched_tokens + + self.uses_mrope = self.model_config.uses_mrope + if self.uses_mrope: + self.mrope_state = MRopeState( + max_num_reqs=self.max_num_reqs, + max_num_tokens=self.max_num_tokens, + max_model_len=self.max_model_len, + device=self.device, + ) + + def add_request(self, req_index: int, new_req_data: NewRequestData) -> None: + if self.uses_mrope: + # Pre-compute M-RoPE positions for prefill. + assert new_req_data.prefill_token_ids is not None + self.mrope_state.init_prefill_mrope_positions( + req_index, + self.model, # type: ignore + new_req_data.prefill_token_ids, + mm_features=new_req_data.mm_features, + ) + + def apply_staged_writes(self) -> None: + if self.uses_mrope: + self.mrope_state.apply_staged_writes() + + def prepare_inputs( + self, input_batch: InputBatch, req_states: RequestState + ) -> dict[str, torch.Tensor | None]: + if not self.uses_mrope: + # Common case (1D positions). + return {} + + # Prepare M-RoPE positions. + self.mrope_state.prepare_mrope_positions( + input_batch.idx_mapping, + input_batch.query_start_loc, + req_states.prefill_len.gpu, + req_states.num_computed_tokens.gpu, + ) + mrope_positions = self.mrope_state.mrope_positions[ + :, : input_batch.num_tokens_after_padding + ] + return {"positions": mrope_positions} + + def prepare_dummy_inputs( + self, num_reqs: int, num_tokens: int + ) -> dict[str, torch.Tensor | None]: + if not self.uses_mrope: + return {} + mrope_positions = self.mrope_state.mrope_positions[:, :num_tokens] + return {"positions": mrope_positions} From 3d66502e1bf48d4ca92ca0d54f7c9bba39a8556c Mon Sep 17 00:00:00 2001 From: Woosuk Kwon Date: Thu, 26 Feb 2026 11:47:02 -0800 Subject: [PATCH 2/2] [Model Runner V2] Prepare attn metadata in ModelState [2/N] (#35383) Signed-off-by: Woosuk Kwon --- vllm/v1/worker/gpu/input_batch.py | 11 +- vllm/v1/worker/gpu/model_runner.py | 154 ++++++++---------- vllm/v1/worker/gpu/model_states.py | 31 ++++ .../gpu/spec_decode/eagle/speculator.py | 6 +- 4 files changed, 110 insertions(+), 92 deletions(-) diff --git a/vllm/v1/worker/gpu/input_batch.py b/vllm/v1/worker/gpu/input_batch.py index 87b8bbf18b8..75655258c18 100644 --- a/vllm/v1/worker/gpu/input_batch.py +++ b/vllm/v1/worker/gpu/input_batch.py @@ -1,7 +1,6 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project from dataclasses import dataclass -from typing import Any import numpy as np import torch @@ -60,6 +59,8 @@ class InputBatch: query_start_loc_np: np.ndarray # [num_reqs] seq_lens: torch.Tensor + # [num_reqs] + dcp_local_seq_lens: torch.Tensor | None # [num_tokens_after_padding] input_ids: torch.Tensor @@ -68,11 +69,6 @@ class InputBatch: # [num_tokens_after_padding, hidden_size] inputs_embeds: torch.Tensor | None - # layer_name -> Metadata - attn_metadata: dict[str, Any] - # layer_name -> slot_mapping - slot_mappings: dict[str, torch.Tensor] - # [total_num_logits] logits_indices: torch.Tensor # [num_reqs + 1] @@ -139,11 +135,10 @@ class InputBatch: query_start_loc=query_start_loc, query_start_loc_np=query_start_loc_np, seq_lens=seq_lens, + dcp_local_seq_lens=None, input_ids=input_ids, positions=positions, inputs_embeds=None, - attn_metadata=None, # type: ignore - slot_mappings=None, # type: ignore logits_indices=logits_indices, cu_num_logits=cu_num_logits, cu_num_logits_np=cu_num_logits_np, diff --git a/vllm/v1/worker/gpu/model_runner.py b/vllm/v1/worker/gpu/model_runner.py index 949f09f54ba..7dcdaf1d202 100644 --- a/vllm/v1/worker/gpu/model_runner.py +++ b/vllm/v1/worker/gpu/model_runner.py @@ -46,7 +46,6 @@ from vllm.v1.outputs import DraftTokenIds, ModelRunnerOutput from vllm.v1.worker.cp_utils import check_attention_cp_compatibility from vllm.v1.worker.gpu.async_utils import AsyncOutput from vllm.v1.worker.gpu.attn_utils import ( - build_attn_metadata, build_slot_mappings_by_layer, get_kv_cache_spec, init_attn_backend, @@ -317,31 +316,6 @@ class GPUModelRunner(LoRAModelRunnerMixin): ) self.kv_connector = get_kv_connector(self.vllm_config, kv_caches_dict) - def prepare_dummy_attn_metadata(self, input_batch: InputBatch) -> None: - block_tables = self.block_tables.get_dummy_block_tables(input_batch.num_reqs) - slot_mappings = self.block_tables.get_dummy_slot_mappings( - input_batch.num_tokens - ) - slot_mappings_by_layer = build_slot_mappings_by_layer( - slot_mappings, self.kv_cache_config - ) - attn_metadata = build_attn_metadata( - attn_groups=self.attn_groups, - num_reqs=input_batch.num_reqs, - num_tokens=input_batch.num_tokens, - query_start_loc_gpu=input_batch.query_start_loc, - query_start_loc_cpu=torch.from_numpy(input_batch.query_start_loc_np), - max_query_len=input_batch.num_scheduled_tokens.max().item(), - seq_lens=input_batch.seq_lens, - max_seq_len=self.max_model_len, - block_tables=block_tables, - slot_mappings=slot_mappings, - kv_cache_config=self.kv_cache_config, - dcp_local_seq_lens=self.input_buffers.dcp_local_seq_lens, - ) - input_batch.attn_metadata = attn_metadata - input_batch.slot_mappings = slot_mappings_by_layer - @torch.inference_mode() def _dummy_run( self, num_tokens: int, *args, skip_attn: bool = True, **kwargs @@ -384,7 +358,7 @@ class GPUModelRunner(LoRAModelRunnerMixin): return None, None assert self.execute_model_state is not None - hidden_states, _, input_batch, _ = self.execute_model_state + input_batch, _, _, _, hidden_states, _, _ = self.execute_model_state self.execute_model_state = None assert hidden_states is not None # Last PP rank always has hidden_states sample_hidden_states = hidden_states[input_batch.logits_indices] @@ -546,7 +520,6 @@ class GPUModelRunner(LoRAModelRunnerMixin): self.encoder_runner.add_request(req_id, new_req_data.mm_features) self.model_state.add_request(req_index, new_req_data) - self.block_tables.append_block_ids( req_index, new_req_data.block_ids, overwrite=True ) @@ -624,9 +597,6 @@ class GPUModelRunner(LoRAModelRunnerMixin): idx_mapping, total_num_logits, cu_num_logits, max_expand_len ) - # Block tables: num_kv_cache_groups x [num_reqs, max_num_blocks] - block_tables = self.block_tables.gather_block_tables(idx_mapping) - # Get query_start_loc. query_start_loc_np = np.empty(self.max_num_reqs + 1, dtype=np.int32) query_start_loc_np[0] = 0 @@ -635,11 +605,8 @@ class GPUModelRunner(LoRAModelRunnerMixin): # Some attention backends like FA3 require query_start_loc to be non-decreasing. query_start_loc_np[num_reqs + 1 :] = num_tokens async_copy_to_gpu(query_start_loc_np, out=self.input_buffers.query_start_loc) - query_start_loc_np = query_start_loc_np[: num_reqs + 1] - query_start_loc_cpu = torch.from_numpy(query_start_loc_np) query_start_loc = self.input_buffers.query_start_loc[: num_reqs + 1] - max_query_len = num_scheduled_tokens.max().item() # Get prefill tokens if any. if self.req_states.any_prefills(idx_mapping_np): @@ -663,6 +630,7 @@ class GPUModelRunner(LoRAModelRunnerMixin): ) seq_lens = self.input_buffers.seq_lens[:num_reqs] + dcp_local_seq_lens = None if self.use_dcp: # Prepare dcp local seq_lens. prepare_dcp_local_seq_lens( @@ -673,7 +641,7 @@ class GPUModelRunner(LoRAModelRunnerMixin): self.dcp_rank, self.cp_interleave, ) - dcp_local_seq_lens = self.input_buffers.dcp_local_seq_lens[:num_reqs] + dcp_local_seq_lens = self.input_buffers.dcp_local_seq_lens[:num_reqs] # Some input token ids are directly read from the last sampled tokens # and draft tokens. Also, get the logits indices to sample tokens from. @@ -689,35 +657,6 @@ class GPUModelRunner(LoRAModelRunnerMixin): total_num_logits, ) - # Compute slot mappings: [num_kv_cache_groups, num_tokens] - slot_mappings = self.block_tables.compute_slot_mappings( - idx_mapping, - query_start_loc, - self.input_buffers.positions[:num_tokens], - ) - # Layer name -> slot mapping. - slot_mappings_by_layer = build_slot_mappings_by_layer( - slot_mappings, self.kv_cache_config - ) - - # Layer name -> attention metadata. - attn_metadata = build_attn_metadata( - attn_groups=self.attn_groups, - num_reqs=num_reqs, - num_tokens=num_tokens, - query_start_loc_gpu=query_start_loc, - query_start_loc_cpu=query_start_loc_cpu, - max_query_len=max_query_len, - seq_lens=self.input_buffers.seq_lens, - max_seq_len=self.max_model_len, - block_tables=block_tables, - slot_mappings=slot_mappings, - kv_cache_config=self.kv_cache_config, - dcp_local_seq_lens=dcp_local_seq_lens, - ) - - input_ids = self.input_buffers.input_ids[:num_tokens_after_padding] - positions = self.input_buffers.positions[:num_tokens_after_padding] return InputBatch( req_ids=req_ids, num_reqs=num_reqs, @@ -732,17 +671,38 @@ class GPUModelRunner(LoRAModelRunnerMixin): query_start_loc=query_start_loc, query_start_loc_np=query_start_loc_np, seq_lens=seq_lens, - input_ids=input_ids, - positions=positions, + dcp_local_seq_lens=dcp_local_seq_lens, + input_ids=self.input_buffers.input_ids[:num_tokens_after_padding], + positions=self.input_buffers.positions[:num_tokens_after_padding], inputs_embeds=None, - attn_metadata=attn_metadata, - slot_mappings=slot_mappings_by_layer, logits_indices=logits_indices, cu_num_logits=cu_num_logits, cu_num_logits_np=cu_num_logits_np, has_structured_output_reqs=scheduler_output.has_structured_output_requests, ) + def prepare_attn( + self, input_batch: InputBatch + ) -> tuple[tuple[torch.Tensor, ...], torch.Tensor]: + # Block tables: num_kv_cache_groups x [num_reqs, max_num_blocks] + block_tables = self.block_tables.gather_block_tables(input_batch.idx_mapping) + # Compute slot mappings: [num_kv_cache_groups, num_tokens] + slot_mappings = self.block_tables.compute_slot_mappings( + input_batch.idx_mapping, + input_batch.query_start_loc, + input_batch.positions, + ) + return block_tables, slot_mappings + + def prepare_dummy_attn( + self, input_batch: InputBatch + ) -> tuple[tuple[torch.Tensor, ...], torch.Tensor]: + block_tables = self.block_tables.get_dummy_block_tables(input_batch.num_reqs) + slot_mappings = self.block_tables.get_dummy_slot_mappings( + input_batch.num_tokens + ) + return block_tables, slot_mappings + @torch.inference_mode() def get_mm_embeddings( self, @@ -899,6 +859,8 @@ class GPUModelRunner(LoRAModelRunnerMixin): input_batch = self.prepare_inputs( scheduler_output, num_tokens_after_padding ) + block_tables, slot_mappings = self.prepare_attn(input_batch) + if self.lora_config: # Activate LoRA adapters. lora_inputs = self.lora_state.make_lora_inputs( @@ -929,9 +891,28 @@ class GPUModelRunner(LoRAModelRunnerMixin): device=self.device, ) if not skip_attn_for_dummy_run: - self.prepare_dummy_attn_metadata(input_batch) + block_tables, slot_mappings = self.prepare_dummy_attn(input_batch) + else: + block_tables = None + slot_mappings = None # FIXME(woosuk): Fix warmup for LoRA. + attn_metadata = None + slot_mappings_by_layer = None + if not (dummy_run and skip_attn_for_dummy_run): + assert slot_mappings is not None + slot_mappings_by_layer = build_slot_mappings_by_layer( + slot_mappings, self.kv_cache_config + ) + assert block_tables is not None + attn_metadata = self.model_state.prepare_attn( + input_batch, + block_tables, + slot_mappings, + self.attn_groups, + self.kv_cache_config, + ) + model_inputs = { "input_ids": input_batch.input_ids, "positions": input_batch.positions, @@ -968,13 +949,13 @@ class GPUModelRunner(LoRAModelRunnerMixin): ) with set_forward_context( - input_batch.attn_metadata, + attn_metadata, self.vllm_config, num_tokens=input_batch.num_tokens_after_padding, cudagraph_runtime_mode=cudagraph_runtime_mode, num_tokens_across_dp=num_tokens_across_dp, batch_descriptor=batch_descriptor, - slot_mapping=input_batch.slot_mappings, + slot_mapping=slot_mappings_by_layer, ): self.kv_connector.pre_forward(scheduler_output) model_output = self.model(**model_inputs) @@ -985,22 +966,23 @@ class GPUModelRunner(LoRAModelRunnerMixin): aux_hidden_states = None kv_connector_output = self.kv_connector.post_forward(scheduler_output) + self.execute_model_state = ( + input_batch, + model_inputs, + attn_metadata, + slot_mappings_by_layer, + hidden_states, + aux_hidden_states, + kv_connector_output, + ) if not self.is_last_pp_rank: # Non-last PP rank: return IntermediateTensors for sending. assert isinstance(hidden_states, IntermediateTensors) hidden_states.kv_connector_output = kv_connector_output - self.execute_model_state = (None, None, input_batch, kv_connector_output) return hidden_states - # Last rank (or no PP): hidden_states is a tensor for sampling. assert isinstance(hidden_states, torch.Tensor) - self.execute_model_state = ( - hidden_states, - aux_hidden_states, - input_batch, - kv_connector_output, - ) return None @torch.inference_mode() @@ -1010,9 +992,15 @@ class GPUModelRunner(LoRAModelRunnerMixin): if self.execute_model_state is None: # The prior execute_model call must have failed. return None - hidden_states, aux_hidden_states, input_batch, kv_connector_output = ( - self.execute_model_state - ) + ( + input_batch, + model_inputs, + attn_metadata, + slot_mappings_by_layer, + hidden_states, + aux_hidden_states, + kv_connector_output, + ) = self.execute_model_state self.execute_model_state = None if not self.is_last_pp_rank: @@ -1075,6 +1063,8 @@ class GPUModelRunner(LoRAModelRunnerMixin): if self.speculator is not None: draft_tokens = self.speculator.propose( input_batch, + attn_metadata, + slot_mappings_by_layer, hidden_states, aux_hidden_states, num_sampled, diff --git a/vllm/v1/worker/gpu/model_states.py b/vllm/v1/worker/gpu/model_states.py index 03574b2ad55..838f177b39c 100644 --- a/vllm/v1/worker/gpu/model_states.py +++ b/vllm/v1/worker/gpu/model_states.py @@ -1,13 +1,18 @@ # 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.config import VllmConfig from vllm.v1.core.sched.output import NewRequestData +from vllm.v1.kv_cache_interface import KVCacheConfig +from vllm.v1.worker.gpu.attn_utils import build_attn_metadata from vllm.v1.worker.gpu.input_batch import InputBatch from vllm.v1.worker.gpu.mm.mrope_utils import MRopeState from vllm.v1.worker.gpu.states import RequestState +from vllm.v1.worker.utils import AttentionGroup class ModelState: @@ -72,3 +77,29 @@ class ModelState: return {} mrope_positions = self.mrope_state.mrope_positions[:, :num_tokens] return {"positions": mrope_positions} + + def prepare_attn( + self, + input_batch: InputBatch, + block_tables: tuple[torch.Tensor, ...], + slot_mappings: torch.Tensor, + attn_groups: list[list[AttentionGroup]], + kv_cache_config: KVCacheConfig, + ) -> dict[str, Any]: + query_start_loc_cpu = torch.from_numpy(input_batch.query_start_loc_np) + max_query_len = input_batch.num_scheduled_tokens.max().item() + attn_metadata = build_attn_metadata( + attn_groups=attn_groups, + num_reqs=input_batch.num_reqs, + num_tokens=input_batch.num_tokens, + query_start_loc_gpu=input_batch.query_start_loc, + query_start_loc_cpu=query_start_loc_cpu, + max_query_len=max_query_len, + seq_lens=input_batch.seq_lens, + max_seq_len=self.max_model_len, + block_tables=block_tables, + slot_mappings=slot_mappings, + kv_cache_config=kv_cache_config, + dcp_local_seq_lens=input_batch.dcp_local_seq_lens, + ) + return attn_metadata diff --git a/vllm/v1/worker/gpu/spec_decode/eagle/speculator.py b/vllm/v1/worker/gpu/spec_decode/eagle/speculator.py index 6cd13cebf99..0c85bf65ee9 100644 --- a/vllm/v1/worker/gpu/spec_decode/eagle/speculator.py +++ b/vllm/v1/worker/gpu/spec_decode/eagle/speculator.py @@ -182,6 +182,8 @@ class EagleSpeculator: def propose( self, input_batch: InputBatch, + attn_metadata: dict[str, Any], + slot_mappings: dict[str, torch.Tensor], # [num_tokens, hidden_size] last_hidden_states: torch.Tensor, # num_layers x [num_tokens, hidden_size] @@ -229,8 +231,8 @@ class EagleSpeculator: # TODO(woosuk): Support CUDA graph for prefill. last_hidden_states, hidden_states = self.run_model( num_tokens, - input_batch.attn_metadata, - input_batch.slot_mappings, + attn_metadata, + slot_mappings, num_tokens_across_dp=None, # FIXME ) sample_hidden_states = last_hidden_states[last_token_indices]