diff --git a/vllm/v1/worker/gpu/buffer_utils.py b/vllm/v1/worker/gpu/buffer_utils.py index ad910933aa2..1dd85efca4f 100644 --- a/vllm/v1/worker/gpu/buffer_utils.py +++ b/vllm/v1/worker/gpu/buffer_utils.py @@ -1,6 +1,7 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project from collections.abc import Iterable, Sequence +from dataclasses import dataclass from functools import partial import numpy as np @@ -35,6 +36,61 @@ def async_copy_to_gpu( return out.copy_(tmp, non_blocking=True) +@dataclass +class PrepareInputsRingSlot: + idx_mapping_cpu: torch.Tensor + query_start_loc_cpu: torch.Tensor + cu_num_logits_cpu: torch.Tensor + event: torch.cuda.Event + in_use: bool = False + + +class PrepareInputsBuffers: + def __init__(self, ring_size: int, max_num_reqs: int, device: torch.device): + self.device = device + self.ring_slots = [ + PrepareInputsRingSlot( + idx_mapping_cpu=torch.empty( + max_num_reqs, dtype=torch.int32, pin_memory=True + ), + query_start_loc_cpu=torch.empty( + max_num_reqs + 1, dtype=torch.int32, pin_memory=True + ), + cu_num_logits_cpu=torch.empty( + max_num_reqs + 1, dtype=torch.int32, pin_memory=True + ), + event=torch.cuda.Event(), + ) + for _ in range(ring_size) + ] + self.ring_slot_idx = -1 + self.idx_mapping_gpu = torch.empty( + max_num_reqs, dtype=torch.int32, device=device + ) + self.cu_num_logits_gpu = torch.empty( + max_num_reqs + 1, dtype=torch.int32, device=device + ) + self.arange_reqs_np = np.arange(max_num_reqs + 1, dtype=np.int32) + self.arange_reqs_gpu = torch.arange( + max_num_reqs + 1, dtype=torch.int32, device=device + ) + self.zero_local_pos_gpu = torch.zeros( + max_num_reqs, dtype=torch.int32, device=device + ) + + def acquire_ring_slot(self) -> PrepareInputsRingSlot: + slot_idx = (self.ring_slot_idx + 1) % len(self.ring_slots) + self.ring_slot_idx = slot_idx + slot = self.ring_slots[slot_idx] + if slot.in_use and not slot.event.query(): + slot.event.synchronize() + return slot + + def mark_ring_slot_inflight(self, slot: PrepareInputsRingSlot) -> None: + slot.event.record(torch.cuda.current_stream(self.device)) + slot.in_use = True + + class UvaBuffer: def __init__(self, size: int | Sequence[int], dtype: torch.dtype): if not is_uva_available(): 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..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,21 +59,16 @@ 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 # [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 - # 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] @@ -141,12 +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, - mrope_positions=None, 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 ecd199cc938..0ecf442ce87 100644 --- a/vllm/v1/worker/gpu/model_runner.py +++ b/vllm/v1/worker/gpu/model_runner.py @@ -46,13 +46,13 @@ 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, init_kv_cache, ) from vllm.v1.worker.gpu.block_table import BlockTables +from vllm.v1.worker.gpu.buffer_utils import PrepareInputsBuffers from vllm.v1.worker.gpu.cp_utils import prepare_dcp_local_seq_lens from vllm.v1.worker.gpu.cudagraph_utils import CudaGraphManager from vllm.v1.worker.gpu.dp_utils import ( @@ -76,7 +76,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 @@ -139,14 +139,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) @@ -199,40 +191,10 @@ class GPUModelRunner(LoRAModelRunnerMixin): device=self.device, ) # ring-buffered staging for prepare_inputs - self._prepare_inputs_ring_size = max(self.pp_size, 2) - self._prepare_inputs_ring_slot = -1 - self._prepare_inputs_slot_in_use = [False] * self._prepare_inputs_ring_size - self._prepare_inputs_slot_events = [ - torch.cuda.Event() for _ in range(self._prepare_inputs_ring_size) - ] - - self._idx_mapping_cpu = torch.empty( - (self._prepare_inputs_ring_size, self.max_num_reqs), - dtype=torch.int32, - pin_memory=True, - ) - self._idx_mapping_gpu = torch.empty( - self.max_num_reqs, dtype=torch.int32, device=self.device - ) - self._query_start_loc_cpu = torch.empty( - (self._prepare_inputs_ring_size, self.max_num_reqs + 1), - dtype=torch.int32, - pin_memory=True, - ) - self._cu_num_logits_cpu = torch.empty( - (self._prepare_inputs_ring_size, self.max_num_reqs + 1), - dtype=torch.int32, - pin_memory=True, - ) - self._cu_num_logits_gpu = torch.empty( - self.max_num_reqs + 1, dtype=torch.int32, device=self.device - ) - self._arange_reqs_np = np.arange(self.max_num_reqs + 1, dtype=np.int32) - self._arange_reqs_gpu = torch.arange( - self.max_num_reqs + 1, dtype=torch.int32, device=self.device - ) - self._zero_local_pos_gpu = torch.zeros( - self.max_num_reqs, dtype=torch.int32, device=self.device + self._prepare_inputs_buffers = PrepareInputsBuffers( + ring_size=max(self.pp_size, 2), + max_num_reqs=self.max_num_reqs, + device=self.device, ) self.sampler = Sampler( max_num_reqs=self.max_num_reqs, @@ -247,7 +209,6 @@ class GPUModelRunner(LoRAModelRunnerMixin): # CUDA graphs. self.cudagraph_manager = CudaGraphManager( self.vllm_config, - self.uses_mrope, self.use_aux_hidden_state_outputs, self.device, ) @@ -264,20 +225,6 @@ class GPUModelRunner(LoRAModelRunnerMixin): # For transferring state from execute_model to subsequent sample_tokens call. self.execute_model_state: tuple | None = None - def _get_prepare_inputs_ring_slot(self) -> int: - slot = (self._prepare_inputs_ring_slot + 1) % self._prepare_inputs_ring_size - self._prepare_inputs_ring_slot = slot - if self._prepare_inputs_slot_in_use[slot]: - event = self._prepare_inputs_slot_events[slot] - if not event.query(): - event.synchronize() - return slot - - def _mark_prepare_inputs_ring_slot_inflight(self, slot: int) -> None: - event = self._prepare_inputs_slot_events[slot] - event.record(torch.cuda.current_stream(self.device)) - self._prepare_inputs_slot_in_use[slot] = True - def update_max_model_len(self, max_model_len: int) -> None: self.max_model_len = max_model_len self.req_states.max_model_len = max_model_len @@ -319,6 +266,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 @@ -371,31 +321,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 @@ -438,7 +363,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] @@ -529,16 +454,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, @@ -602,15 +524,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 ) @@ -625,8 +539,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. @@ -645,7 +558,8 @@ class GPUModelRunner(LoRAModelRunnerMixin): assert num_tokens > 0 num_tokens_per_req = scheduler_output.num_scheduled_tokens num_reqs = len(num_tokens_per_req) - slot = self._get_prepare_inputs_ring_slot() + prep_buffers = self._prepare_inputs_buffers + slot = prep_buffers.acquire_ring_slot() # Decode first, then prefill. # batch_idx -> req_id @@ -655,9 +569,9 @@ class GPUModelRunner(LoRAModelRunnerMixin): idx_mapping_iter = map(self.req_states.req_id_to_index.get, req_ids) idx_mapping_np = np.fromiter(idx_mapping_iter, dtype=np.int32, count=num_reqs) - idx_mapping_cpu = self._idx_mapping_cpu[slot, :num_reqs] + idx_mapping_cpu = slot.idx_mapping_cpu[:num_reqs] np.copyto(idx_mapping_cpu.numpy(), idx_mapping_np, casting="no") - idx_mapping = self._idx_mapping_gpu[:num_reqs] + idx_mapping = prep_buffers.idx_mapping_gpu[:num_reqs] idx_mapping.copy_(idx_mapping_cpu, non_blocking=True) # Get the number of draft tokens for each request. @@ -666,10 +580,10 @@ class GPUModelRunner(LoRAModelRunnerMixin): # No draft token scheduled (common case). total_num_draft_tokens = 0 total_num_logits = num_reqs - cu_num_logits_np = self._arange_reqs_np[: num_reqs + 1] - cu_num_logits = self._arange_reqs_gpu[: num_reqs + 1] + cu_num_logits_np = prep_buffers.arange_reqs_np[: num_reqs + 1] + cu_num_logits = prep_buffers.arange_reqs_gpu[: num_reqs + 1] expanded_idx_mapping = idx_mapping - expanded_local_pos = self._zero_local_pos_gpu[:num_reqs] + expanded_local_pos = prep_buffers.zero_local_pos_gpu[:num_reqs] else: num_draft_tokens = np.array( [len(draft_tokens.get(req_id, ())) for req_id in req_ids], @@ -679,24 +593,22 @@ class GPUModelRunner(LoRAModelRunnerMixin): total_num_logits = num_reqs + total_num_draft_tokens num_logits = num_draft_tokens + 1 - cu_num_logits_cpu = self._cu_num_logits_cpu[slot, : num_reqs + 1] + cu_num_logits_cpu = slot.cu_num_logits_cpu[: num_reqs + 1] cu_num_logits_cpu_np = cu_num_logits_cpu.numpy() cu_num_logits_cpu_np[0] = 0 np.cumsum(num_logits, out=cu_num_logits_cpu_np[1:]) - cu_num_logits = self._cu_num_logits_gpu[: num_reqs + 1] + cu_num_logits = prep_buffers.cu_num_logits_gpu[: num_reqs + 1] cu_num_logits.copy_(cu_num_logits_cpu, non_blocking=True) - cu_num_logits_np = cu_num_logits_cpu_np + # keep an independent CPU snapshot because ring slots are reused. + cu_num_logits_np = cu_num_logits_cpu_np.copy() max_expand_len = self.num_speculative_steps + 1 expanded_idx_mapping, expanded_local_pos = expand_idx_mapping( 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_cpu_full = self._query_start_loc_cpu[slot] + query_start_loc_cpu_full = slot.query_start_loc_cpu query_start_loc_np_full = query_start_loc_cpu_full.numpy() query_start_loc_np_full[0] = 0 np.cumsum(num_scheduled_tokens, out=query_start_loc_np_full[1 : num_reqs + 1]) @@ -706,12 +618,11 @@ class GPUModelRunner(LoRAModelRunnerMixin): self.input_buffers.query_start_loc.copy_( query_start_loc_cpu_full, non_blocking=True ) - self._mark_prepare_inputs_ring_slot_inflight(slot) + prep_buffers.mark_ring_slot_inflight(slot) - query_start_loc_cpu = query_start_loc_cpu_full[: num_reqs + 1] - query_start_loc_np = query_start_loc_np_full[: num_reqs + 1] + # keep an independent CPU snapshot because ring slots are reused. + query_start_loc_np = query_start_loc_np_full[: num_reqs + 1].copy() 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): @@ -735,6 +646,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( @@ -745,16 +657,7 @@ class GPUModelRunner(LoRAModelRunnerMixin): self.dcp_rank, self.cp_interleave, ) - 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, - ) + 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. @@ -770,39 +673,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] - 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, @@ -817,18 +687,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, - mrope_positions=mrope_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, @@ -985,6 +875,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( @@ -1014,14 +906,43 @@ 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) + 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, + "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. @@ -1038,41 +959,22 @@ 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, ) 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( - 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: @@ -1080,22 +982,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() @@ -1105,9 +1008,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: @@ -1170,6 +1079,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 new file mode 100644 index 00000000000..838f177b39c --- /dev/null +++ b/vllm/v1/worker/gpu/model_states.py @@ -0,0 +1,105 @@ +# 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: + 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} + + 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]