# SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project from dataclasses import dataclass import numpy as np import torch from vllm.pooling_params import PoolingParams from vllm.tasks import PoolingTask from vllm.utils.torch_utils import PIN_MEMORY @dataclass class PoolingCursor: first_token_indices_gpu: torch.Tensor last_token_indices_gpu: torch.Tensor prompt_lens_cpu: torch.Tensor seq_lens_cpu: torch.Tensor num_scheduled_tokens_cpu: torch.Tensor def __getitem__(self, indices: slice) -> "PoolingCursor": return PoolingCursor( first_token_indices_gpu=self.first_token_indices_gpu[indices], last_token_indices_gpu=self.last_token_indices_gpu[indices], prompt_lens_cpu=self.prompt_lens_cpu[indices], seq_lens_cpu=self.seq_lens_cpu[indices], num_scheduled_tokens_cpu=self.num_scheduled_tokens_cpu[indices], ) def is_partial_prefill(self) -> bool: return not torch.all(self.prompt_lens_cpu == self.num_scheduled_tokens_cpu) def is_finished(self) -> torch.Tensor: return self.prompt_lens_cpu == self.seq_lens_cpu class PoolingStates: def __init__(self) -> None: # for chunked prefill with ALL pooling self.hidden_states_cache: list[torch.Tensor] = [] def clean(self) -> None: self.hidden_states_cache.clear() @dataclass class PoolingMetadata: """Tensors for pooling.""" prompt_lens: torch.Tensor # CPU Tensor prompt_token_ids: torch.Tensor | None # Model-device tensor prompt_token_ids_cpu: torch.Tensor | None # CPU tensor pooling_params: list[PoolingParams] pooling_states: list[PoolingStates] pooling_cursor: PoolingCursor | None = None def __post_init__(self) -> None: pooling_params = self.pooling_params tasks: list[PoolingTask] = [ task for pooling_param in pooling_params if (task := pooling_param.task) is not None ] if len(pooling_params) != len(tasks): raise ValueError( "Every pooling param must have a task set, but got " f"{len(tasks)} tasks for {len(pooling_params)} pooling params" ) self.tasks = tasks def __getitem__(self, indices: slice) -> "PoolingMetadata": return PoolingMetadata( prompt_lens=self.prompt_lens[indices], prompt_token_ids=None if self.prompt_token_ids is None else self.prompt_token_ids[indices], prompt_token_ids_cpu=None if self.prompt_token_ids_cpu is None else self.prompt_token_ids_cpu[indices], pooling_params=self.pooling_params[indices], pooling_states=self.pooling_states[indices], pooling_cursor=None if self.pooling_cursor is None else self.pooling_cursor[indices], ) def _get_prompt_token_ids( self, prompt_token_ids: torch.Tensor | None, ) -> list[torch.Tensor]: if prompt_token_ids is None: raise ValueError( "prompt_token_ids is required but was not set. " "Please set `requires_token_ids=True` in `get_pooling_updates`" ) return [prompt_token_ids[i, :num] for i, num in enumerate(self.prompt_lens)] def get_prompt_token_ids(self) -> list[torch.Tensor]: return self._get_prompt_token_ids(self.prompt_token_ids) def get_prompt_token_ids_cpu(self) -> list[torch.Tensor]: return self._get_prompt_token_ids(self.prompt_token_ids_cpu) def get_pooling_cursor(self) -> PoolingCursor: pooling_cursor = self.pooling_cursor if pooling_cursor is None: raise RuntimeError( "pooling_cursor has not been initialized. " "Call `build_pooling_cursor` before accessing it" ) return pooling_cursor def build_pooling_cursor( self, num_scheduled_tokens_np: np.ndarray, seq_lens_cpu: torch.Tensor, device: torch.device, query_start_loc_gpu: torch.Tensor | None = None, ) -> None: n_seq = len(num_scheduled_tokens_np) prompt_lens = self.prompt_lens if len(prompt_lens) != n_seq: raise ValueError( f"prompt_lens length ({len(prompt_lens)}) does not match " f"the number of sequences ({n_seq})" ) num_scheduled_tokens_cpu = torch.from_numpy(num_scheduled_tokens_np) if query_start_loc_gpu is None: cumsum = torch.zeros( n_seq + 1, dtype=torch.int64, pin_memory=PIN_MEMORY, device="cpu" ) torch.cumsum(num_scheduled_tokens_cpu, dim=0, out=cumsum[1:]) cumsum = cumsum.to(device, non_blocking=True) else: if query_start_loc_gpu.shape[0] != n_seq + 1: raise ValueError( "query_start_loc_gpu length does not match " f"the number of sequences: {query_start_loc_gpu.shape[0]} " f"!= {n_seq + 1}." ) if query_start_loc_gpu.device != device: raise ValueError( "query_start_loc_gpu must be on the same device as the " f"hidden states: {query_start_loc_gpu.device} != {device}." ) cumsum = query_start_loc_gpu self.pooling_cursor = PoolingCursor( first_token_indices_gpu=cumsum[:n_seq], last_token_indices_gpu=cumsum[1:] - 1, prompt_lens_cpu=prompt_lens, seq_lens_cpu=seq_lens_cpu, num_scheduled_tokens_cpu=num_scheduled_tokens_cpu, )