gpu_input_batch.py 36.5 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
        if idx - self.num_prompt_tokens < len(self.output_token_ids):
66
            return self.output_token_ids[idx - self.num_prompt_tokens]
67
        return -1
68
69
70
71


class InputBatch:
    def __init__(
72
73
74
75
76
77
78
79
        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
80
        kernel_block_sizes: list[int],
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
            kernel_block_sizes=kernel_block_sizes,
136
            num_speculative_tokens=num_speculative_tokens,
137
138
139
        )

        # Sampling-related.
140
141
142
143
144
145
        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
        )
146
        self.temperature_cpu = self.temperature_cpu_tensor.numpy()
147
148
        self.greedy_reqs: set[str] = set()
        self.random_reqs: set[str] = set()
149

150
151
152
153
        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
        )
154
        self.top_p_cpu = self.top_p_cpu_tensor.numpy()
155
        self.top_p_reqs: set[str] = set()
156

157
158
159
160
        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
        )
161
        self.top_k_cpu = self.top_k_cpu_tensor.numpy()
162
        self.top_k_reqs: set[str] = set()
163

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

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

        # Presence penalty related data structures
178
179
        self.presence_penalties = torch.empty(
            (max_num_reqs,), dtype=torch.float, device=device
180
        )
181
182
183
184
        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()
185
        self.presence_penalties_reqs: set[str] = set()
186
187

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

197
        # Speculative decoding
198
199
200
201
        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()
202

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

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

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

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

221
222
223
224
225
226
        # 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
227
        self.has_allowed_token_ids: set[str] = set()
228
229
        # NOTE(lufang): In the mask tensor, if the corresponding token allowed,
        # the value is False. Since we use masked_fill_ to set -inf.
230
231
        self.allowed_token_ids_mask: Optional[torch.Tensor] = None
        self.allowed_token_ids_mask_cpu_tensor: Optional[torch.Tensor] = None
232

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

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

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

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

244
245
246
        # Store last speculative tokens for sampler.
        self.spec_token_ids: list[Optional[list[int]]] = []

247
248
249
        # This is updated each time the batch constituents change.
        self.sampling_metadata = self._make_sampling_metadata()

250
251
        self.pooling_params: dict[str, PoolingParams] = {}

252
253
254
255
        # Cached reference to the GPU tensor of previously sampled tokens
        self.prev_sampled_token_ids: Optional[torch.Tensor] = None
        self.prev_req_id_to_index: Optional[dict[str, int]] = None

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

262
    def _register_add_request(self, request: "CachedRequestState") -> int:
263
264
265
266
267
268
269
270
271
272
        """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
273
274
275
276
277
        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(
278
279
280
281
282
283
284
                (
                    new_req_index,
                    request.sampling_params,
                    request.prompt_token_ids,
                    request.output_token_ids,
                )
            )
285

286
        return new_req_index
287

288
289
290
    def add_request(
        self,
        request: "CachedRequestState",
291
    ) -> int:
292
        req_index = self._register_add_request(request)
293
294

        req_id = request.req_id
295
296
297
        if req_index == len(self._req_ids):
            self._req_ids.append(req_id)
            self.req_output_token_ids.append(request.output_token_ids)
298
            self.spec_token_ids.append([])
299
300
301
        else:
            self._req_ids[req_index] = req_id
            self.req_output_token_ids[req_index] = request.output_token_ids
302
            self.spec_token_ids[req_index] = []
303

304
305
306
        self.req_id_to_index[req_id] = req_index

        # Copy the prompt token ids and output token ids.
307
        num_prompt_tokens = length_from_prompt_token_ids_or_embeds(
308
309
            request.prompt_token_ids, request.prompt_embeds
        )
310
        self.num_prompt_tokens[req_index] = num_prompt_tokens
311
312
        start_idx = num_prompt_tokens
        end_idx = start_idx + len(request.output_token_ids)
313
        if request.prompt_token_ids is not None:
314
            self.token_ids_cpu[req_index, :num_prompt_tokens] = request.prompt_token_ids
315
316
317
318
319
            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
320
        self.token_ids_cpu[req_index, start_idx:end_idx] = request.output_token_ids
321
322
        self.is_token_ids[req_index, start_idx:end_idx] = True
        # Number of token ids in prompt (token_ids_cpu or prompt_embeds).
323
        # NOTE(woosuk): This may include spec decode tokens.
324
        self.num_tokens[req_index] = request.num_tokens
325
326
        # Number of tokens without spec decode tokens.
        self.num_tokens_no_spec[req_index] = request.num_tokens
327
328

        self.num_computed_tokens_cpu[req_index] = request.num_computed_tokens
329
        self.block_table.add_row(request.block_ids, req_index)
330

331
        if sampling_params := request.sampling_params:
332
            if self.is_spec_decode and is_spec_decode_unsupported(sampling_params):
333
                self.spec_decode_unsupported_reqs.add(req_id)
334
            if sampling_params.sampling_type == SamplingType.GREEDY:
335
336
                # Should avoid division by zero later when apply_temperature.
                self.temperature_cpu[req_index] = 0.0
337
338
339
340
341
342
343
344
345
346
347
348
349
350
                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
351
            self.frequency_penalties_cpu[req_index] = sampling_params.frequency_penalty
352
353
            if sampling_params.frequency_penalty != 0.0:
                self.frequency_penalties_reqs.add(req_id)
354
            self.presence_penalties_cpu[req_index] = sampling_params.presence_penalty
355
356
            if sampling_params.presence_penalty != 0.0:
                self.presence_penalties_reqs.add(req_id)
357
358
359
            self.repetition_penalties_cpu[req_index] = (
                sampling_params.repetition_penalty
            )
360
361
362
363
364
365
366
367
368
            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:
369
370
371
372
373
                self.num_logprobs[req_id] = (
                    self.vocab_size
                    if sampling_params.logprobs == -1
                    else sampling_params.logprobs
                )
374
            if sampling_params.prompt_logprobs is not None:
375
                self.num_prompt_logprobs[req_id] = (
376
377
378
379
                    self.vocab_size
                    if sampling_params.prompt_logprobs == -1
                    else sampling_params.prompt_logprobs
                )
380
381
382
383
384
385
386
387
388
389

            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,
390
391
                        device=self.device,
                    )
392
393
394
395
                    self.allowed_token_ids_mask_cpu_tensor = torch.zeros(
                        self.max_num_reqs,
                        self.vocab_size,
                        dtype=torch.bool,
396
397
                        device="cpu",
                    )
398
                self.allowed_token_ids_mask_cpu_tensor[req_index] = True
399
                # False means we don't fill with -inf.
400
                self.allowed_token_ids_mask_cpu_tensor[req_index][
401
402
                    sampling_params.allowed_token_ids
                ] = False
403

404
            if sampling_params.bad_words_token_ids:
405
406
407
                self.bad_words_token_ids[req_index] = (
                    sampling_params.bad_words_token_ids
                )
408
409
410
        elif pooling_params := request.pooling_params:
            self.pooling_params[req_id] = pooling_params
            self.logits_processing_needs_token_ids[req_index] = (
411
412
                pooling_params.requires_token_ids
            )
413
        else:
414
            raise NotImplementedError("Unrecognized request type")
415

416
417
418
        # Speculative decoding: by default 1 token is generated.
        self.num_accepted_tokens_cpu[req_index] = 1

419
420
421
422
423
424
425
426
427
428
429
430
431
        # 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

432
433
        return req_index

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

437
438
439
440
441
442
        Args:
          req_id: request to remove

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

444
445
446
        req_index = self.req_id_to_index.pop(req_id, None)
        if req_index is None:
            return None
447
448

        self.batch_update_builder.removed_append(req_index)
449
450
        self._req_ids[req_index] = None
        self.req_output_token_ids[req_index] = None
451
        self.spec_token_ids[req_index] = None
452

453
454
455
456
457
458
459
460
461
462
463
464
465
466
        # 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

467
468
469
470
        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)
471
        self.spec_decode_unsupported_reqs.discard(req_id)
472
473
474
        self.frequency_penalties_reqs.discard(req_id)
        self.presence_penalties_reqs.discard(req_id)
        self.repetition_penalties_reqs.discard(req_id)
475
476
        self.generators.pop(req_index, None)
        self.num_logprobs.pop(req_id, None)
477
        self.num_prompt_logprobs.pop(req_id, None)
478
        self.in_progress_prompt_logprobs_cpu.pop(req_id, None)
479

480
481
        self.has_allowed_token_ids.discard(req_id)
        if self.allowed_token_ids_mask_cpu_tensor is not None:
482
            # False means we don't fill with -inf.
483
            self.allowed_token_ids_mask_cpu_tensor[req_index].fill_(False)
484
        self.bad_words_token_ids.pop(req_index, None)
485
486
        return req_index

487
488
489
    def swap_states(self, i1: int, i2: int) -> None:
        old_id_i1 = self._req_ids[i1]
        old_id_i2 = self._req_ids[i2]
490
491
492
493
494
        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],
        )
495
496
497
498
        self.spec_token_ids[i1], self.spec_token_ids[i2] = (
            self.spec_token_ids[i2],
            self.spec_token_ids[i1],
        )
499
        assert old_id_i1 is not None and old_id_i2 is not None
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
        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],
        )
520

521
522
523
524
525
526
527
528
529
        # 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

530
531
532
533
534
535
536
537
538
539
540
541
542
543
        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)

544
        self.block_table.swap_row(i1, i2)
545

546
547
548
549
        self.request_lora_mapping[i1], self.request_lora_mapping[i2] = (
            self.request_lora_mapping[i2],
            self.request_lora_mapping[i1],
        )
550

551
552
553
554
555
556
        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.
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
        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],
        )
581
582
583
584

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

585
        if self.allowed_token_ids_mask_cpu_tensor is not None:
586
587
588
589
590
591
592
            (
                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],
            )
593

594
595
596
597
598
599
600
601
602
    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
603
        """
604
605
        num_reqs = self.num_reqs

606
607
608
609
        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
610
        if num_reqs == 0:
611
            # The batched states are empty.
612
613
            self._req_ids.clear()
            self.req_output_token_ids.clear()
614
            self.spec_token_ids.clear()
615
616
617
618
            return

        # NOTE(woosuk): This function assumes that the empty_req_indices
        # is sorted in descending order.
619
        last_req_index = num_reqs + len(empty_req_indices) - 1
620
621
622
623
624
625
        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.
626
627
            empty_index = self.batch_update_builder.peek_removed()
            assert empty_index is not None
628
629
630
            if empty_index >= last_req_index:
                break

631
632
633
            # Move active request down into empty request
            # index.
            self.batch_update_builder.pop_removed()
634
635
            req_id = self._req_ids[last_req_index]
            output_token_ids = self.req_output_token_ids[last_req_index]
636
            assert req_id is not None
637
638
639
640
            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
641
642
            self.req_id_to_index[req_id] = empty_index

643
644
645
646
            spec_token_ids = self.spec_token_ids[last_req_index]
            self.spec_token_ids[empty_index] = spec_token_ids
            self.spec_token_ids[last_req_index] = None

647
648
            num_tokens = self.num_tokens[last_req_index]
            self.token_ids_cpu[empty_index, :num_tokens] = self.token_ids_cpu[
649
650
                last_req_index, :num_tokens
            ]
651
            self.is_token_ids[empty_index, :num_tokens] = self.is_token_ids[
652
653
                last_req_index, :num_tokens
            ]
654
            if last_req_index in self.req_prompt_embeds:
655
656
657
                self.req_prompt_embeds[empty_index] = self.req_prompt_embeds.pop(
                    last_req_index
                )
658
            self.num_tokens[empty_index] = num_tokens
659
            self.num_tokens_no_spec[empty_index] = self.num_tokens_no_spec[
660
661
662
663
664
665
                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
            ]
666
            self.block_table.move_row(last_req_index, empty_index)
667
668

            self.request_lora_mapping[empty_index] = self.request_lora_mapping[
669
670
                last_req_index
            ]
671
672
673

            if self.is_pooling_model:
                last_req_index -= 1
co63oc's avatar
co63oc committed
674
                # Sampling state not used by pooling models.
675
676
677
678
679
                continue

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

683
            self.temperature_cpu[empty_index] = self.temperature_cpu[last_req_index]
684
685
            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]
686
687
688
689
690
691
692
693
694
695
696
697
            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
            ]
698
699
700
701
            generator = self.generators.pop(last_req_index, None)
            if generator is not None:
                self.generators[empty_index] = generator

702
            # TODO convert these to LogitsProcessors
703
            if self.allowed_token_ids_mask_cpu_tensor is not None:
704
705
706
                self.allowed_token_ids_mask_cpu_tensor[empty_index] = (
                    self.allowed_token_ids_mask_cpu_tensor[last_req_index]
                )
707

708
            bad_words_token_ids = self.bad_words_token_ids.pop(last_req_index, None)
709
710
            if bad_words_token_ids is not None:
                self.bad_words_token_ids[empty_index] = bad_words_token_ids
711

712
713
714
            # Decrement last_req_index since it is now empty.
            last_req_index -= 1

715
        # Trim lists to the batch size.
716
717
        del self._req_ids[num_reqs:]
        del self.req_output_token_ids[num_reqs:]
718
        del self.spec_token_ids[num_reqs:]
719

720
    def refresh_metadata(self):
721
        """Apply any batch updates to sampling metadata."""
722

723
        if self.is_pooling_model:
724
725
726
            batch_changed = self.batch_update_builder.reset()
            if batch_changed:
                self.sampling_metadata = self._make_sampling_metadata()
727
728
729
730
731
            return

        # For non-pooling models - generate and apply logitsprocs update;
        # reset batch update tracking.
        # Update sampling metadata if batch state is changed.
732
733
734
735
736
        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()
737
738
739

    def _make_sampling_metadata(self) -> SamplingMetadata:
        num_reqs = self.num_reqs
740
        if not self.all_greedy:
741
742
743
            temperature = copy_slice(
                self.temperature_cpu_tensor, self.temperature, num_reqs
            )
744
745
        else:
            temperature = None
746
747
748
749
750
751
752
753
754
        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.
755
756
757
758
759
760
761
762
763
764
765
            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,
            )
766

767
768
        needs_prompt_token_ids = (
            not self.no_penalties
769
770
            or self.logits_processing_needs_token_ids[:num_reqs].any()
        )
771
772
773
774
775
776
777
        # 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.
        prompt_token_ids = (
            self._make_prompt_token_ids_tensor() if needs_prompt_token_ids else None
        )
778

779
780
781
        allowed_token_ids_mask: Optional[torch.Tensor] = None
        if not self.no_allowed_token_ids:
            assert self.allowed_token_ids_mask is not None
782
783
784
785
786
            copy_slice(
                self.allowed_token_ids_mask_cpu_tensor,
                self.allowed_token_ids_mask,
                num_reqs,
            )
787
788
            allowed_token_ids_mask = self.allowed_token_ids_mask[:num_reqs]

789
        return SamplingMetadata(
790
            temperature=temperature,
791
792
            all_greedy=self.all_greedy,
            all_random=self.all_random,
793
794
            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],
795
796
            generators=self.generators,
            max_num_logprobs=self.max_num_logprobs,
797
798
799
800
            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],
801
            output_token_ids=cast(list[list[int]], self.req_output_token_ids),
802
            spec_token_ids=cast(list[list[int]], self.spec_token_ids),
803
            no_penalties=self.no_penalties,
804
            allowed_token_ids_mask=allowed_token_ids_mask,
805
            bad_words_token_ids=self.bad_words_token_ids,
806
            logitsprocs=self.logitsprocs,
807
808
        )

809
810
811
812
813
814
    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()
815
816

        return PoolingMetadata(
817
            prompt_lens=torch.from_numpy(self.num_prompt_tokens[: self.num_reqs]),
818
819
820
821
            prompt_token_ids=self.sampling_metadata.prompt_token_ids,
            pooling_params=pooling_params,
        )

822
    def _make_prompt_token_ids_tensor(self) -> torch.Tensor:
823
824
        num_reqs = self.num_reqs
        max_prompt_len = self.num_prompt_tokens[:num_reqs].max()
825
826
827
828
        prompt_token_ids_cpu_tensor = torch.empty(
            (self.num_reqs, max_prompt_len),
            device="cpu",
            dtype=torch.int64,
829
830
            pin_memory=self.pin_memory,
        )
831
        prompt_token_ids = prompt_token_ids_cpu_tensor.numpy()
832
        prompt_token_ids[:] = self.token_ids_cpu[:num_reqs, :max_prompt_len]
833
834
        # Use the value of vocab_size as a pad since we don't have a
        # token_id of this value.
835
        for i in range(num_reqs):
836
837
            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)
838

839
840
    def make_lora_inputs(
        self, num_scheduled_tokens: np.ndarray
841
    ) -> tuple[tuple[int, ...], tuple[int, ...], set[LoRARequest]]:
842
843
844
845
846
847
848
849
850
851
852
        """
        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.
        """

853
        req_lora_mapping = self.request_lora_mapping[: self.num_reqs]
854
        prompt_lora_mapping = tuple(req_lora_mapping)
855
        token_lora_mapping = tuple(req_lora_mapping.repeat(num_scheduled_tokens))
856
        active_lora_requests: set[LoRARequest] = set(
857
858
            self.lora_id_to_lora_request.values()
        )
859
860
861

        return prompt_lora_mapping, token_lora_mapping, active_lora_requests

862
863
864
865
866
867
868
869
870
871
872
873
874
875
876
877
878
879
880
881
    @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

882
883
    @property
    def no_penalties(self) -> bool:
884
885
886
887
888
        return (
            len(self.presence_penalties_reqs) == 0
            and len(self.frequency_penalties_reqs) == 0
            and len(self.repetition_penalties_reqs) == 0
        )
889

890
    @property
891
892
    def max_num_logprobs(self) -> Optional[int]:
        return max(self.num_logprobs.values()) if self.num_logprobs else None
893
894
895

    @property
    def no_prompt_logprob(self) -> bool:
896
        return not self.num_prompt_logprobs
897
898
899
900

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