gpu_input_batch.py 35.9 KB
Newer Older
1
# SPDX-License-Identifier: Apache-2.0
2
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
3
# Datastructures defining a GPU input batch
4
5

from dataclasses import dataclass
6
from typing import Optional, cast
7
8
9
10

import numpy as np
import torch

11
from vllm.lora.request import LoRARequest
12
from vllm.multimodal.inputs import MultiModalFeatureSpec
13
from vllm.pooling_params import PoolingParams
14
from vllm.sampling_params import SamplingParams, SamplingType
15
from vllm.utils import length_from_prompt_token_ids_or_embeds, swap_dict_values
16
from vllm.v1.outputs import LogprobsTensors
17
from vllm.v1.pool.metadata import PoolingMetadata
18
19
20
21
22
from vllm.v1.sample.logits_processor import (
    BatchUpdateBuilder,
    LogitsProcessors,
    MoveDirectionality,
)
23
from vllm.v1.sample.metadata import SamplingMetadata
24
from vllm.v1.spec_decode.utils import is_spec_decode_unsupported
25
from vllm.v1.utils import copy_slice
26
from vllm.v1.worker.block_table import MultiGroupBlockTable
27
28
29
30
31


@dataclass
class CachedRequestState:
    req_id: str
32
    prompt_token_ids: Optional[list[int]]
33
    mm_features: list[MultiModalFeatureSpec]
34
35
    sampling_params: Optional[SamplingParams]
    pooling_params: Optional[PoolingParams]
36
37
    generator: Optional[torch.Generator]

38
    block_ids: tuple[list[int], ...]
39
    num_computed_tokens: int
40
    output_token_ids: list[int]
41

42
43
44
    mrope_positions: Optional[torch.Tensor] = None
    mrope_position_delta: Optional[int] = None

45
    lora_request: Optional[LoRARequest] = None
46
    prompt_embeds: Optional[torch.Tensor] = None
47

48
    def __post_init__(self):
49
        self.num_prompt_tokens = length_from_prompt_token_ids_or_embeds(
50
51
            self.prompt_token_ids, self.prompt_embeds
        )
52

53
54
    @property
    def num_tokens(self) -> int:
55
56
57
58
        return self.num_prompt_tokens + len(self.output_token_ids)

    def get_token_id(self, idx: int) -> int:
        if idx < self.num_prompt_tokens:
59
60
61
            if self.prompt_token_ids is None:
                raise ValueError(
                    f"Tried to access token index {idx}, but that token was "
62
63
                    "provided via prompt_embeds, and its ID is unknown."
                )
64
            return self.prompt_token_ids[idx]
65
66
67
68
        elif idx - self.num_prompt_tokens < len(self.output_token_ids):
            return self.output_token_ids[idx - self.num_prompt_tokens]
        else:
            return -1
69
70
71
72


class InputBatch:
    def __init__(
73
74
75
76
77
78
79
80
        self,
        max_num_reqs: int,
        max_model_len: int,
        max_num_batched_tokens: int,
        device: torch.device,
        pin_memory: bool,
        vocab_size: int,
        block_sizes: list[int],  # The block_size of each kv cache group
81
        logitsprocs: Optional[LogitsProcessors] = None,
82
        is_spec_decode: bool = False,
83
        is_pooling_model: bool = False,
84
        num_speculative_tokens: int = 0,
85
    ):
86
        self.is_pooling_model = is_pooling_model
87
        self.is_spec_decode = is_spec_decode
88
89
        self.max_num_reqs = max_num_reqs
        self.max_model_len = max_model_len
90
        self.max_num_batched_tokens = max_num_batched_tokens
91
92
        self.device = device
        self.pin_memory = pin_memory
93
        self.vocab_size = vocab_size
94

95
96
        self._req_ids: list[Optional[str]] = []
        self.req_id_to_index: dict[str, int] = {}
97

98
99
        # TODO(woosuk): This buffer could be too large if max_model_len is big.
        # Find a way to reduce the CPU memory usage.
100
101
        # This buffer is not directly transferred to the GPU, so it does not
        # need to be pinned.
102
103
104
105
        self.token_ids_cpu_tensor = torch.zeros(
            (max_num_reqs, max_model_len),
            device="cpu",
            dtype=torch.int32,
106
            pin_memory=False,
107
108
        )
        self.token_ids_cpu = self.token_ids_cpu_tensor.numpy()
109
110
111
        self.is_token_ids = torch.zeros(
            (max_num_reqs, max_model_len), device="cpu", dtype=bool, pin_memory=False
        )
112
113
114
115
        # Store prompt embeddings per request to avoid OOM from large upfront
        # allocation if max_model_len is big.
        # Maps req_index -> tensor of shape (num_prompt_tokens, hidden_size)
        self.req_prompt_embeds: dict[int, torch.Tensor] = {}
116
        self.num_tokens = np.zeros(max_num_reqs, dtype=np.int32)
117
        self.num_tokens_no_spec = np.zeros(max_num_reqs, dtype=np.int32)
118
        self.num_prompt_tokens = np.zeros(max_num_reqs, dtype=np.int32)
119
        self.num_computed_tokens_cpu_tensor = torch.zeros(
120
            (max_num_reqs,),
121
122
123
124
            device="cpu",
            dtype=torch.int32,
            pin_memory=pin_memory,
        )
125
        self.num_computed_tokens_cpu = self.num_computed_tokens_cpu_tensor.numpy()
126

127
        # Block table.
128
        self.block_table = MultiGroupBlockTable(
129
            max_num_reqs=max_num_reqs,
130
            max_model_len=max_model_len,
131
            max_num_batched_tokens=max_num_batched_tokens,
132
            pin_memory=pin_memory,
133
            device=device,
134
            block_sizes=block_sizes,
135
            num_speculative_tokens=num_speculative_tokens,
136
137
138
        )

        # Sampling-related.
139
140
141
142
143
144
        self.temperature = torch.empty(
            (max_num_reqs,), dtype=torch.float32, device=device
        )
        self.temperature_cpu_tensor = torch.empty(
            (max_num_reqs,), dtype=torch.float32, device="cpu", pin_memory=pin_memory
        )
145
        self.temperature_cpu = self.temperature_cpu_tensor.numpy()
146
147
        self.greedy_reqs: set[str] = set()
        self.random_reqs: set[str] = set()
148

149
150
151
152
        self.top_p = torch.empty((max_num_reqs,), dtype=torch.float32, device=device)
        self.top_p_cpu_tensor = torch.empty(
            (max_num_reqs,), dtype=torch.float32, device="cpu", pin_memory=pin_memory
        )
153
        self.top_p_cpu = self.top_p_cpu_tensor.numpy()
154
        self.top_p_reqs: set[str] = set()
155

156
157
158
159
        self.top_k = torch.empty((max_num_reqs,), dtype=torch.int32, device=device)
        self.top_k_cpu_tensor = torch.empty(
            (max_num_reqs,), dtype=torch.int32, device="cpu", pin_memory=pin_memory
        )
160
        self.top_k_cpu = self.top_k_cpu_tensor.numpy()
161
        self.top_k_reqs: set[str] = set()
162

163
164
        # IDs of requests which do not support spec decoding
        self.spec_decode_unsupported_reqs: set[str] = set()
165

166
        # Frequency penalty related data structures
167
168
169
        self.frequency_penalties = torch.empty(
            (max_num_reqs,), dtype=torch.float, device=device
        )
170
        self.frequency_penalties_cpu_tensor = torch.empty(
171
172
173
            (max_num_reqs,), dtype=torch.float, device="cpu", pin_memory=pin_memory
        )
        self.frequency_penalties_cpu = self.frequency_penalties_cpu_tensor.numpy()
174
        self.frequency_penalties_reqs: set[str] = set()
175
176

        # Presence penalty related data structures
177
178
        self.presence_penalties = torch.empty(
            (max_num_reqs,), dtype=torch.float, device=device
179
        )
180
181
182
183
        self.presence_penalties_cpu_tensor = torch.empty(
            (max_num_reqs,), dtype=torch.float, device="cpu", pin_memory=pin_memory
        )
        self.presence_penalties_cpu = self.presence_penalties_cpu_tensor.numpy()
184
        self.presence_penalties_reqs: set[str] = set()
185
186

        # Repetition penalty related data structures
187
188
189
        self.repetition_penalties = torch.empty(
            (max_num_reqs,), dtype=torch.float, device=device
        )
190
        self.repetition_penalties_cpu_tensor = torch.empty(
191
192
193
            (max_num_reqs,), dtype=torch.float, device="cpu", pin_memory=pin_memory
        )
        self.repetition_penalties_cpu = self.repetition_penalties_cpu_tensor.numpy()
194
        self.repetition_penalties_reqs: set[str] = set()
195

196
        # Speculative decoding
197
198
199
200
        self.num_accepted_tokens_cpu_tensor = torch.ones(
            (max_num_reqs,), dtype=torch.int64, device="cpu", pin_memory=pin_memory
        )
        self.num_accepted_tokens_cpu = self.num_accepted_tokens_cpu_tensor.numpy()
201

202
        # lora related
203
        self.request_lora_mapping = np.zeros((self.max_num_reqs,), dtype=np.int32)
204
205
        self.lora_id_to_request_ids: dict[int, set[str]] = {}
        self.lora_id_to_lora_request: dict[int, LoRARequest] = {}
206

207
        # req_index -> generator
208
209
        # NOTE(woosuk): The indices of the requests that do not have their own
        # generator should not be included in the dictionary.
210
        self.generators: dict[int, torch.Generator] = {}
211

212
        self.num_logprobs: dict[str, int] = {}
213
214
        # NOTE(rob): num_prompt_logprobs only includes reqs
        # that are currently in the prefill phase.
215
        self.num_prompt_logprobs: dict[str, int] = {}
216

217
218
219
        # To accumulate prompt logprobs tensor chunks across prefill steps.
        self.in_progress_prompt_logprobs_cpu: dict[str, LogprobsTensors] = {}

220
221
222
223
224
225
        # Internal representation of per-step batch state changes, used for
        # reordering persistent batch and generating logitsprocs batch state
        # updates. Should reset each step.
        self.batch_update_builder = BatchUpdateBuilder()

        # TODO convert this to LogitsProcessor
226
        self.has_allowed_token_ids: set[str] = set()
227
228
        # NOTE(lufang): In the mask tensor, if the corresponding token allowed,
        # the value is False. Since we use masked_fill_ to set -inf.
229
230
        self.allowed_token_ids_mask: Optional[torch.Tensor] = None
        self.allowed_token_ids_mask_cpu_tensor: Optional[torch.Tensor] = None
231

232
233
234
        # req_index -> bad_words_token_ids
        self.bad_words_token_ids: dict[int, list[list[int]]] = {}

235
        self.logits_processing_needs_token_ids = np.zeros(max_num_reqs, dtype=bool)
236

237
        self.req_output_token_ids: list[Optional[list[int]]] = []
238

239
240
241
242
        # Store provided logitsprocs. If none are provided, initialize empty
        # data structure
        self.logitsprocs = logitsprocs or LogitsProcessors()

243
244
245
        # This is updated each time the batch constituents change.
        self.sampling_metadata = self._make_sampling_metadata()

246
247
        self.pooling_params: dict[str, PoolingParams] = {}

248
249
250
251
252
        # Cached reference to the GPU tensor of previously sampled tokens
        self.prev_sampled_token_ids: Optional[torch.Tensor] = None
        self.prev_sampled_token_ids_invalid_indices: Optional[set[int]] = None
        self.prev_req_id_to_index: Optional[dict[str, int]] = None

253
    @property
254
    def req_ids(self) -> list[str]:
255
256
        # None elements should only be present transiently
        # while performing state updates to the batch.
257
        return cast(list[str], self._req_ids)
258

259
    def _register_add_request(self, request: "CachedRequestState") -> int:
260
261
262
263
264
265
266
267
268
269
        """Track add-request operations for logits processors.
        Not applicable to pooling models.
        """

        # Fill the next empty index if there is one.
        if (new_req_index := self.batch_update_builder.pop_removed()) is None:
            # Append to end otherwise.
            new_req_index = self.num_reqs

        assert new_req_index < self.max_num_reqs
270
271
272
273
274
        self.batch_update_builder.batch_changed = True
        if request.sampling_params:
            # Detailed added request metadata is only required for non-pooling
            # models, to support logitsprocs.
            self.batch_update_builder.added.append(
275
276
277
278
279
280
281
                (
                    new_req_index,
                    request.sampling_params,
                    request.prompt_token_ids,
                    request.output_token_ids,
                )
            )
282

283
        return new_req_index
284

285
286
287
    def add_request(
        self,
        request: "CachedRequestState",
288
    ) -> int:
289
        req_index = self._register_add_request(request)
290
291

        req_id = request.req_id
292
293
294
295
296
297
298
        if req_index == len(self._req_ids):
            self._req_ids.append(req_id)
            self.req_output_token_ids.append(request.output_token_ids)
        else:
            self._req_ids[req_index] = req_id
            self.req_output_token_ids[req_index] = request.output_token_ids

299
300
301
        self.req_id_to_index[req_id] = req_index

        # Copy the prompt token ids and output token ids.
302
        num_prompt_tokens = length_from_prompt_token_ids_or_embeds(
303
304
            request.prompt_token_ids, request.prompt_embeds
        )
305
        self.num_prompt_tokens[req_index] = num_prompt_tokens
306
307
        start_idx = num_prompt_tokens
        end_idx = start_idx + len(request.output_token_ids)
308
        if request.prompt_token_ids is not None:
309
            self.token_ids_cpu[req_index, :num_prompt_tokens] = request.prompt_token_ids
310
311
312
313
314
            self.is_token_ids[req_index, :num_prompt_tokens] = True
        else:
            self.is_token_ids[req_index, :num_prompt_tokens] = False
        if request.prompt_embeds is not None:
            self.req_prompt_embeds[req_index] = request.prompt_embeds
315
        self.token_ids_cpu[req_index, start_idx:end_idx] = request.output_token_ids
316
317
        self.is_token_ids[req_index, start_idx:end_idx] = True
        # Number of token ids in prompt (token_ids_cpu or prompt_embeds).
318
        # NOTE(woosuk): This may include spec decode tokens.
319
        self.num_tokens[req_index] = request.num_tokens
320
321
        # Number of tokens without spec decode tokens.
        self.num_tokens_no_spec[req_index] = request.num_tokens
322
323

        self.num_computed_tokens_cpu[req_index] = request.num_computed_tokens
324
        self.block_table.add_row(request.block_ids, req_index)
325

326
        if sampling_params := request.sampling_params:
327
            if self.is_spec_decode and is_spec_decode_unsupported(sampling_params):
328
                self.spec_decode_unsupported_reqs.add(req_id)
329
            if sampling_params.sampling_type == SamplingType.GREEDY:
330
331
                # Should avoid division by zero later when apply_temperature.
                self.temperature_cpu[req_index] = 0.0
332
333
334
335
336
337
338
339
340
341
342
343
344
345
                self.greedy_reqs.add(req_id)
            else:
                self.temperature_cpu[req_index] = sampling_params.temperature
                self.random_reqs.add(req_id)

            self.top_p_cpu[req_index] = sampling_params.top_p
            if sampling_params.top_p < 1:
                self.top_p_reqs.add(req_id)
            top_k = sampling_params.top_k
            if 0 < top_k < self.vocab_size:
                self.top_k_reqs.add(req_id)
            else:
                top_k = self.vocab_size
            self.top_k_cpu[req_index] = top_k
346
            self.frequency_penalties_cpu[req_index] = sampling_params.frequency_penalty
347
348
            if sampling_params.frequency_penalty != 0.0:
                self.frequency_penalties_reqs.add(req_id)
349
            self.presence_penalties_cpu[req_index] = sampling_params.presence_penalty
350
351
            if sampling_params.presence_penalty != 0.0:
                self.presence_penalties_reqs.add(req_id)
352
353
354
            self.repetition_penalties_cpu[req_index] = (
                sampling_params.repetition_penalty
            )
355
356
357
358
359
360
361
362
363
            if sampling_params.repetition_penalty != 1.0:
                self.repetition_penalties_reqs.add(req_id)

            # NOTE(woosuk): self.generators should not include the requests that
            # do not have their own generator.
            if request.generator is not None:
                self.generators[req_index] = request.generator

            if sampling_params.logprobs is not None:
364
365
366
367
368
                self.num_logprobs[req_id] = (
                    self.vocab_size
                    if sampling_params.logprobs == -1
                    else sampling_params.logprobs
                )
369
            if sampling_params.prompt_logprobs is not None:
370
                self.num_prompt_logprobs[req_id] = (
371
372
373
374
                    self.vocab_size
                    if sampling_params.prompt_logprobs == -1
                    else sampling_params.prompt_logprobs
                )
375
376
377
378
379
380
381
382
383
384

            if sampling_params.allowed_token_ids:
                self.has_allowed_token_ids.add(req_id)
                if self.allowed_token_ids_mask_cpu_tensor is None:
                    # Lazy allocation for this tensor, which can be large.
                    # False means we don't fill with -inf.
                    self.allowed_token_ids_mask = torch.zeros(
                        self.max_num_reqs,
                        self.vocab_size,
                        dtype=torch.bool,
385
386
                        device=self.device,
                    )
387
388
389
390
                    self.allowed_token_ids_mask_cpu_tensor = torch.zeros(
                        self.max_num_reqs,
                        self.vocab_size,
                        dtype=torch.bool,
391
392
                        device="cpu",
                    )
393
                self.allowed_token_ids_mask_cpu_tensor[req_index] = True
394
                # False means we don't fill with -inf.
395
                self.allowed_token_ids_mask_cpu_tensor[req_index][
396
397
                    sampling_params.allowed_token_ids
                ] = False
398

399
            if sampling_params.bad_words_token_ids:
400
401
402
                self.bad_words_token_ids[req_index] = (
                    sampling_params.bad_words_token_ids
                )
403
404
405
        elif pooling_params := request.pooling_params:
            self.pooling_params[req_id] = pooling_params
            self.logits_processing_needs_token_ids[req_index] = (
406
407
                pooling_params.requires_token_ids
            )
408
        else:
409
            raise NotImplementedError("Unrecognized request type")
410

411
412
413
        # Speculative decoding: by default 1 token is generated.
        self.num_accepted_tokens_cpu[req_index] = 1

414
415
416
417
418
419
420
421
422
423
424
425
426
        # Add request lora ID
        if request.lora_request:
            lora_id = request.lora_request.lora_int_id
            if lora_id not in self.lora_id_to_request_ids:
                self.lora_id_to_request_ids[lora_id] = set()

            self.request_lora_mapping[req_index] = lora_id
            self.lora_id_to_request_ids[lora_id].add(request.req_id)
            self.lora_id_to_lora_request[lora_id] = request.lora_request
        else:
            # No LoRA
            self.request_lora_mapping[req_index] = 0

427
428
        return req_index

429
    def remove_request(self, req_id: str) -> Optional[int]:
430
        """This method must always be followed by a call to condense().
431

432
433
434
435
436
437
        Args:
          req_id: request to remove

        Returns:
          Removed request index, or `None` if `req_id` not recognized
        """
438

439
440
441
        req_index = self.req_id_to_index.pop(req_id, None)
        if req_index is None:
            return None
442
443

        self.batch_update_builder.removed_append(req_index)
444
445
        self._req_ids[req_index] = None
        self.req_output_token_ids[req_index] = None
446

447
448
449
450
451
452
453
454
455
456
457
458
459
460
        # LoRA
        lora_id = self.request_lora_mapping[req_index]
        if lora_id != 0:
            lora_req_ids = self.lora_id_to_request_ids[lora_id]
            lora_req_ids.discard(req_id)
            if not lora_req_ids:
                del self.lora_id_to_request_ids[lora_id]
                del self.lora_id_to_lora_request[lora_id]
            self.request_lora_mapping[req_index] = 0

        if self.is_pooling_model:
            self.pooling_params.pop(req_id, None)
            return req_index

461
462
463
464
        self.greedy_reqs.discard(req_id)
        self.random_reqs.discard(req_id)
        self.top_p_reqs.discard(req_id)
        self.top_k_reqs.discard(req_id)
465
        self.spec_decode_unsupported_reqs.discard(req_id)
466
467
468
        self.frequency_penalties_reqs.discard(req_id)
        self.presence_penalties_reqs.discard(req_id)
        self.repetition_penalties_reqs.discard(req_id)
469
470
        self.generators.pop(req_index, None)
        self.num_logprobs.pop(req_id, None)
471
        self.num_prompt_logprobs.pop(req_id, None)
472
        self.in_progress_prompt_logprobs_cpu.pop(req_id, None)
473

474
475
        self.has_allowed_token_ids.discard(req_id)
        if self.allowed_token_ids_mask_cpu_tensor is not None:
476
            # False means we don't fill with -inf.
477
            self.allowed_token_ids_mask_cpu_tensor[req_index].fill_(False)
478
        self.bad_words_token_ids.pop(req_index, None)
479
480
        return req_index

481
482
483
    def swap_states(self, i1: int, i2: int) -> None:
        old_id_i1 = self._req_ids[i1]
        old_id_i2 = self._req_ids[i2]
484
485
486
487
488
        self._req_ids[i1], self._req_ids[i2] = self._req_ids[i2], self._req_ids[i1]  # noqa
        self.req_output_token_ids[i1], self.req_output_token_ids[i2] = (
            self.req_output_token_ids[i2],
            self.req_output_token_ids[i1],
        )
489
        assert old_id_i1 is not None and old_id_i2 is not None
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
        self.req_id_to_index[old_id_i1], self.req_id_to_index[old_id_i2] = (
            self.req_id_to_index[old_id_i2],
            self.req_id_to_index[old_id_i1],
        )
        self.num_tokens[i1], self.num_tokens[i2] = (
            self.num_tokens[i2],
            self.num_tokens[i1],
        )
        self.num_tokens_no_spec[i1], self.num_tokens_no_spec[i2] = (
            self.num_tokens_no_spec[i2],
            self.num_tokens_no_spec[i1],
        )
        self.num_prompt_tokens[i1], self.num_prompt_tokens[i2] = (
            self.num_prompt_tokens[i2],
            self.num_prompt_tokens[i1],
        )
        self.num_computed_tokens_cpu[i1], self.num_computed_tokens_cpu[i2] = (
            self.num_computed_tokens_cpu[i2],
            self.num_computed_tokens_cpu[i1],
        )
510

511
512
513
514
515
516
517
518
519
        # NOTE: the following is unsafe
        # self.token_ids_cpu[i1, ...], self.token_ids_cpu[i2, ...], =\
        #     self.token_ids_cpu[i2, ...], self.token_ids_cpu[i1, ...]
        # instead, we need to temporiarily copy the data for one of the indices
        # TODO(lucas): optimize this by only copying valid indices
        tmp = self.token_ids_cpu[i1, ...].copy()
        self.token_ids_cpu[i1, ...] = self.token_ids_cpu[i2, ...]
        self.token_ids_cpu[i2, ...] = tmp

520
521
522
523
524
525
526
527
528
529
530
531
532
533
        self.is_token_ids[[i1, i2], ...] = self.is_token_ids[[i2, i1], ...]

        # Swap prompt embeddings if they exist
        embeds_i1 = self.req_prompt_embeds.get(i1)
        embeds_i2 = self.req_prompt_embeds.get(i2)
        if embeds_i1 is not None:
            self.req_prompt_embeds[i2] = embeds_i1
        else:
            self.req_prompt_embeds.pop(i2, None)
        if embeds_i2 is not None:
            self.req_prompt_embeds[i1] = embeds_i2
        else:
            self.req_prompt_embeds.pop(i1, None)

534
        self.block_table.swap_row(i1, i2)
535

536
537
538
539
        self.request_lora_mapping[i1], self.request_lora_mapping[i2] = (
            self.request_lora_mapping[i2],
            self.request_lora_mapping[i1],
        )
540

541
542
543
544
545
546
        if self.is_pooling_model:
            # Sampling and logits parameters don't apply to pooling models.
            return

        # For autoregressive models, track detailed request reordering info
        # to support logitsprocs.
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
        self.batch_update_builder.moved.append((i1, i2, MoveDirectionality.SWAP))

        self.temperature_cpu[i1], self.temperature_cpu[i2] = (
            self.temperature_cpu[i2],
            self.temperature_cpu[i1],
        )
        self.top_p_cpu[i1], self.top_p_cpu[i2] = self.top_p_cpu[i2], self.top_p_cpu[i1]
        self.top_k_cpu[i1], self.top_k_cpu[i2] = self.top_k_cpu[i2], self.top_k_cpu[i1]
        self.frequency_penalties_cpu[i1], self.frequency_penalties_cpu[i2] = (
            self.frequency_penalties_cpu[i2],
            self.frequency_penalties_cpu[i1],
        )
        self.presence_penalties_cpu[i1], self.presence_penalties_cpu[i2] = (
            self.presence_penalties_cpu[i2],
            self.presence_penalties_cpu[i1],
        )
        self.repetition_penalties_cpu[i1], self.repetition_penalties_cpu[i2] = (
            self.repetition_penalties_cpu[i2],
            self.repetition_penalties_cpu[i1],
        )
        self.num_accepted_tokens_cpu[i1], self.num_accepted_tokens_cpu[i2] = (
            self.num_accepted_tokens_cpu[i2],
            self.num_accepted_tokens_cpu[i1],
        )
571
572
573
574

        swap_dict_values(self.generators, i1, i2)
        swap_dict_values(self.bad_words_token_ids, i1, i2)

575
        if self.allowed_token_ids_mask_cpu_tensor is not None:
576
577
578
579
580
581
582
            (
                self.allowed_token_ids_mask_cpu_tensor[i1],
                self.allowed_token_ids_mask_cpu_tensor[i2],
            ) = (
                self.allowed_token_ids_mask_cpu_tensor[i2],
                self.allowed_token_ids_mask_cpu_tensor[i1],
            )
583

584
585
586
587
588
589
590
591
592
    def condense(self) -> None:
        """Slide non-empty requests down into lower, empty indices.

        Any consecutive empty indices at the very end of the list are not
        filled.

        Returns:
          swaps: list of (from,to) swap tuples for moved requests
          empty_req_indices: indices not filled by condensation
593
        """
594
595
        num_reqs = self.num_reqs

596
597
598
599
        if not (empty_req_indices := self.batch_update_builder.removed):
            # All removed requests were replaced by added requests, or else no
            # requests were removed at all. No condense() needed
            return
600
        if num_reqs == 0:
601
            # The batched states are empty.
602
603
            self._req_ids.clear()
            self.req_output_token_ids.clear()
604
605
606
607
            return

        # NOTE(woosuk): This function assumes that the empty_req_indices
        # is sorted in descending order.
608
        last_req_index = num_reqs + len(empty_req_indices) - 1
609
610
611
612
613
614
        while empty_req_indices:
            # Find the largest non-empty index.
            while last_req_index in empty_req_indices:
                last_req_index -= 1

            # Find the smallest empty index.
615
616
            empty_index = self.batch_update_builder.peek_removed()
            assert empty_index is not None
617
618
619
            if empty_index >= last_req_index:
                break

620
621
622
            # Move active request down into empty request
            # index.
            self.batch_update_builder.pop_removed()
623
624
            req_id = self._req_ids[last_req_index]
            output_token_ids = self.req_output_token_ids[last_req_index]
625
            assert req_id is not None
626
627
628
629
            self._req_ids[empty_index] = req_id
            self._req_ids[last_req_index] = None
            self.req_output_token_ids[empty_index] = output_token_ids
            self.req_output_token_ids[last_req_index] = None
630
631
            self.req_id_to_index[req_id] = empty_index

632
633
            num_tokens = self.num_tokens[last_req_index]
            self.token_ids_cpu[empty_index, :num_tokens] = self.token_ids_cpu[
634
635
                last_req_index, :num_tokens
            ]
636
            self.is_token_ids[empty_index, :num_tokens] = self.is_token_ids[
637
638
                last_req_index, :num_tokens
            ]
639
            if last_req_index in self.req_prompt_embeds:
640
641
642
                self.req_prompt_embeds[empty_index] = self.req_prompt_embeds.pop(
                    last_req_index
                )
643
            self.num_tokens[empty_index] = num_tokens
644
            self.num_tokens_no_spec[empty_index] = self.num_tokens_no_spec[
645
646
647
648
649
650
                last_req_index
            ]
            self.num_prompt_tokens[empty_index] = self.num_prompt_tokens[last_req_index]
            self.num_computed_tokens_cpu[empty_index] = self.num_computed_tokens_cpu[
                last_req_index
            ]
651
            self.block_table.move_row(last_req_index, empty_index)
652
653

            self.request_lora_mapping[empty_index] = self.request_lora_mapping[
654
655
                last_req_index
            ]
656
657
658

            if self.is_pooling_model:
                last_req_index -= 1
co63oc's avatar
co63oc committed
659
                # Sampling state not used by pooling models.
660
661
662
663
664
                continue

            # Autoregressive models require detailed tracking of condense
            # operations to support logitsprocs
            self.batch_update_builder.moved.append(
665
666
                (last_req_index, empty_index, MoveDirectionality.UNIDIRECTIONAL)
            )
667

668
            self.temperature_cpu[empty_index] = self.temperature_cpu[last_req_index]
669
670
            self.top_p_cpu[empty_index] = self.top_p_cpu[last_req_index]
            self.top_k_cpu[empty_index] = self.top_k_cpu[last_req_index]
671
672
673
674
675
676
677
678
679
680
681
682
            self.frequency_penalties_cpu[empty_index] = self.frequency_penalties_cpu[
                last_req_index
            ]
            self.presence_penalties_cpu[empty_index] = self.presence_penalties_cpu[
                last_req_index
            ]
            self.repetition_penalties_cpu[empty_index] = self.repetition_penalties_cpu[
                last_req_index
            ]
            self.num_accepted_tokens_cpu[empty_index] = self.num_accepted_tokens_cpu[
                last_req_index
            ]
683
684
685
686
            generator = self.generators.pop(last_req_index, None)
            if generator is not None:
                self.generators[empty_index] = generator

687
            # TODO convert these to LogitsProcessors
688
            if self.allowed_token_ids_mask_cpu_tensor is not None:
689
690
691
                self.allowed_token_ids_mask_cpu_tensor[empty_index] = (
                    self.allowed_token_ids_mask_cpu_tensor[last_req_index]
                )
692

693
            bad_words_token_ids = self.bad_words_token_ids.pop(last_req_index, None)
694
695
            if bad_words_token_ids is not None:
                self.bad_words_token_ids[empty_index] = bad_words_token_ids
696

697
698
699
            # Decrement last_req_index since it is now empty.
            last_req_index -= 1

700
        # Trim lists to the batch size.
701
702
        del self._req_ids[num_reqs:]
        del self.req_output_token_ids[num_reqs:]
703

704
    def refresh_metadata(self):
705
        """Apply any batch updates to sampling metadata."""
706

707
        if self.is_pooling_model:
708
709
710
            batch_changed = self.batch_update_builder.reset()
            if batch_changed:
                self.sampling_metadata = self._make_sampling_metadata()
711
712
713
714
715
            return

        # For non-pooling models - generate and apply logitsprocs update;
        # reset batch update tracking.
        # Update sampling metadata if batch state is changed.
716
717
718
719
720
        batch_update = self.batch_update_builder.get_and_reset(self.num_reqs)
        for logit_proc in self.logitsprocs.all:
            logit_proc.update_state(batch_update)
        if batch_update:
            self.sampling_metadata = self._make_sampling_metadata()
721
722
723

    def _make_sampling_metadata(self) -> SamplingMetadata:
        num_reqs = self.num_reqs
724
        if not self.all_greedy:
725
726
727
            temperature = copy_slice(
                self.temperature_cpu_tensor, self.temperature, num_reqs
            )
728
729
        else:
            temperature = None
730
731
732
733
734
735
736
737
738
        if not self.no_top_p:
            copy_slice(self.top_p_cpu_tensor, self.top_p, num_reqs)
        if not self.no_top_k:
            copy_slice(self.top_k_cpu_tensor, self.top_k, num_reqs)

        if not self.no_penalties:
            # Since syncing these tensors is expensive only copy them
            # if necessary i.e. if there are requests which require
            # penalties to be applied during sampling.
739
740
741
742
743
744
745
746
747
748
749
            copy_slice(
                self.frequency_penalties_cpu_tensor, self.frequency_penalties, num_reqs
            )
            copy_slice(
                self.presence_penalties_cpu_tensor, self.presence_penalties, num_reqs
            )
            copy_slice(
                self.repetition_penalties_cpu_tensor,
                self.repetition_penalties,
                num_reqs,
            )
750

751
752
        needs_prompt_token_ids = (
            not self.no_penalties
753
754
            or self.logits_processing_needs_token_ids[:num_reqs].any()
        )
755
756
757
758
759
        if needs_prompt_token_ids:
            # The prompt tokens are used only for applying penalties or
            # step pooling during the sampling/pooling process.
            # Hence copy these tensors only when there are requests which
            # need penalties/step_pooler to be applied.
760
761
762
            prompt_token_ids = self._make_prompt_token_ids_tensor()
        else:
            prompt_token_ids = None
763

764
765
766
        allowed_token_ids_mask: Optional[torch.Tensor] = None
        if not self.no_allowed_token_ids:
            assert self.allowed_token_ids_mask is not None
767
768
769
770
771
            copy_slice(
                self.allowed_token_ids_mask_cpu_tensor,
                self.allowed_token_ids_mask,
                num_reqs,
            )
772
773
            allowed_token_ids_mask = self.allowed_token_ids_mask[:num_reqs]

774
        return SamplingMetadata(
775
            temperature=temperature,
776
777
            all_greedy=self.all_greedy,
            all_random=self.all_random,
778
779
            top_p=None if self.no_top_p else self.top_p[:num_reqs],
            top_k=None if self.no_top_k else self.top_k[:num_reqs],
780
781
            generators=self.generators,
            max_num_logprobs=self.max_num_logprobs,
782
783
784
785
            prompt_token_ids=prompt_token_ids,
            frequency_penalties=self.frequency_penalties[:num_reqs],
            presence_penalties=self.presence_penalties[:num_reqs],
            repetition_penalties=self.repetition_penalties[:num_reqs],
786
            output_token_ids=cast(list[list[int]], self.req_output_token_ids),
787
            no_penalties=self.no_penalties,
788
            allowed_token_ids_mask=allowed_token_ids_mask,
789
            bad_words_token_ids=self.bad_words_token_ids,
790
            logitsprocs=self.logitsprocs,
791
792
        )

793
794
795
796
797
798
    def get_pooling_params(self) -> list[PoolingParams]:
        assert len(self.req_ids) == len(self.pooling_params)
        return [self.pooling_params[req_id] for req_id in self.req_ids]

    def get_pooling_metadata(self) -> PoolingMetadata:
        pooling_params = self.get_pooling_params()
799
800

        return PoolingMetadata(
801
            prompt_lens=torch.from_numpy(self.num_prompt_tokens[: self.num_reqs]),
802
803
804
805
            prompt_token_ids=self.sampling_metadata.prompt_token_ids,
            pooling_params=pooling_params,
        )

806
    def _make_prompt_token_ids_tensor(self) -> torch.Tensor:
807
808
        num_reqs = self.num_reqs
        max_prompt_len = self.num_prompt_tokens[:num_reqs].max()
809
810
811
812
        prompt_token_ids_cpu_tensor = torch.empty(
            (self.num_reqs, max_prompt_len),
            device="cpu",
            dtype=torch.int64,
813
814
            pin_memory=self.pin_memory,
        )
815
        prompt_token_ids = prompt_token_ids_cpu_tensor.numpy()
816
        prompt_token_ids[:] = self.token_ids_cpu[:num_reqs, :max_prompt_len]
817
818
        # Use the value of vocab_size as a pad since we don't have a
        # token_id of this value.
819
        for i in range(num_reqs):
820
821
            prompt_token_ids[i, self.num_prompt_tokens[i] :] = self.vocab_size
        return prompt_token_ids_cpu_tensor.to(device=self.device, non_blocking=True)
822

823
824
    def make_lora_inputs(
        self, num_scheduled_tokens: np.ndarray
825
    ) -> tuple[tuple[int, ...], tuple[int, ...], set[LoRARequest]]:
826
827
828
829
830
831
832
833
834
835
836
        """
        Given the num_scheduled_tokens for each request in the batch, return
        datastructures used to activate the current LoRAs.
        Returns:
            1. prompt_lora_mapping: A tuple of size self.num_reqs where,
               prompt_lora_mapping[i] is the LoRA id to use for the ith prompt.
            2. token_lora_mapping: A tuple of size np.sum(num_scheduled_tokens)
               where, token_lora_mapping[i] is the LoRA id to use for ith token.
            3. lora_requests: Set of relevant LoRA requests.
        """

837
        req_lora_mapping = self.request_lora_mapping[: self.num_reqs]
838
        prompt_lora_mapping = tuple(req_lora_mapping)
839
        token_lora_mapping = tuple(req_lora_mapping.repeat(num_scheduled_tokens))
840
        active_lora_requests: set[LoRARequest] = set(
841
842
            self.lora_id_to_lora_request.values()
        )
843
844
845

        return prompt_lora_mapping, token_lora_mapping, active_lora_requests

846
847
848
849
850
851
852
853
854
855
856
857
858
859
860
861
862
863
864
865
    @property
    def num_reqs(self) -> int:
        return len(self.req_id_to_index)

    @property
    def all_greedy(self) -> bool:
        return len(self.random_reqs) == 0

    @property
    def all_random(self) -> bool:
        return len(self.greedy_reqs) == 0

    @property
    def no_top_p(self) -> bool:
        return len(self.top_p_reqs) == 0

    @property
    def no_top_k(self) -> bool:
        return len(self.top_k_reqs) == 0

866
867
    @property
    def no_penalties(self) -> bool:
868
869
870
871
872
        return (
            len(self.presence_penalties_reqs) == 0
            and len(self.frequency_penalties_reqs) == 0
            and len(self.repetition_penalties_reqs) == 0
        )
873

874
    @property
875
876
    def max_num_logprobs(self) -> Optional[int]:
        return max(self.num_logprobs.values()) if self.num_logprobs else None
877
878
879

    @property
    def no_prompt_logprob(self) -> bool:
880
        return not self.num_prompt_logprobs
881
882
883
884

    @property
    def no_allowed_token_ids(self) -> bool:
        return len(self.has_allowed_token_ids) == 0