diff --git a/tests/kernels/test_engram.py b/tests/kernels/test_engram.py index c2a7607cbac0..15f4c4f69f3a 100644 --- a/tests/kernels/test_engram.py +++ b/tests/kernels/test_engram.py @@ -745,6 +745,7 @@ def test_engram_constructor_honors_offload(monkeypatch, backend, cpu_offload): if cpu_offload is None else EngramConfig(cpu_offload=cpu_offload), scheduler_config=SimpleNamespace(max_num_batched_tokens=8), + parallel_config=SimpleNamespace(use_ubatching=False), ) offloaded = backend == "nvidia" and vllm_config.engram_config.cpu_offload if backend == "common": @@ -800,6 +801,45 @@ def track_lookup(ids, out, background=False): assert streams == [main] +@pytest.mark.skipif(not current_platform.is_cuda(), reason="CUDA required") +@pytest.mark.parametrize("cpu_offload", [False, True]) +def test_engram_prepared_rows_keep_microbatches_isolated(cpu_offload, monkeypatch): + """Preparing a second microbatch must preserve the first one's lookup rows.""" + layer = _make_embedding(cpu_offload) + module = Engram.__new__(Engram) + torch.nn.Module.__init__(module) + module.embed_tokens = layer + module.use_sequence_parallel = False + module._prefetch_stream = torch.cuda.Stream() if cpu_offload else None + module.staged_rows = torch.empty( + 8, 24, layer.dim, dtype=torch.bfloat16, device="cuda" + ) + module._init_lookup_staging(2) + prepared = [] + for slot, tokens in enumerate((7, 3)): + monkeypatch.setattr(engram_ops, "dbo_current_ubatch_id", lambda slot=slot: slot) + ids = torch.randint( + layer.vocab_start_idx, + layer.vocab_end_idx, + (tokens, 24), + dtype=torch.int32, + device="cuda", + ) + rows = module.prepare_embeddings(ids) + prepared.append((ids, rows)) + assert prepared[0][1].data_ptr() != prepared[1][1].data_ptr() + # Explicit rows must also work while another microbatch is current. + for ids, rows in prepared: + expected = _reference_lookup( + layer.weight.cuda(), + layer.weight_scale_inv.cuda(), + ids, + layer.vocab_start_idx, + layer.vocab_end_idx, + ) + torch.testing.assert_close(module.embed(ids, rows), expected, rtol=0, atol=0) + + @pytest.mark.skipif(not current_platform.is_cuda(), reason="CUDA required") @pytest.mark.parametrize( "cpu_offload,delay", diff --git a/tests/v1/attention/test_attention_splitting.py b/tests/v1/attention/test_attention_splitting.py index acbc22a02793..7d30e680ebff 100644 --- a/tests/v1/attention/test_attention_splitting.py +++ b/tests/v1/attention/test_attention_splitting.py @@ -107,8 +107,10 @@ def test_make_metadata_with_slice_decode_batch(small_decode_metadata): """Test slicing decode batch metadata""" # Split first request only ubatch_slice = UBatchSlice(slice(0, 1), slice(0, 1)) + small_decode_metadata.positions = small_decode_metadata.seq_lens - 1 result = _make_metadata_with_slice(ubatch_slice, small_decode_metadata) + torch.testing.assert_close(result.positions, small_decode_metadata.positions[:1]) # Check sliced results assert result.num_reqs == 1 # slice(0, 1) gives 1 requests @@ -330,6 +332,9 @@ def test_prefill_split_across_ubatches( device = torch.device("cpu") batch_spec = BatchSpec(seq_lens=seq_lens, query_lens=query_lens) common = create_common_attn_metadata(batch_spec, block_size=16, device=device) + common.positions = torch.cat( + [torch.arange(seq - query, seq) for seq, query in zip(seq_lens, query_lens)] + ) num_scheduled_tokens = np.array(query_lens, dtype=np.int32) qsl_np = common.query_start_loc_cpu.numpy() @@ -347,10 +352,16 @@ def test_prefill_split_across_ubatches( first_meta = _make_metadata_with_slice(ubatch_slices[0], common) second_meta = _make_metadata_with_slice(ubatch_slices[1], common) + # Compressor ring slots depend on absolute positions, even inside a request. + torch.testing.assert_close(first_meta.positions, common.positions[:split_point]) + torch.testing.assert_close(second_meta.positions, common.positions[split_point:]) # Token counts match the split assert first_meta.num_actual_tokens == split_point assert second_meta.num_actual_tokens == num_tokens - split_point + # These counts are passed directly to Triton kernels. + assert isinstance(first_meta.num_actual_tokens, int) + assert isinstance(second_meta.num_actual_tokens, int) # Number of requests per ubatch assert first_meta.num_reqs == expected_first_reqs diff --git a/tests/v1/worker/test_ubatch_inputs.py b/tests/v1/worker/test_ubatch_inputs.py new file mode 100644 index 000000000000..4a592759d4d5 --- /dev/null +++ b/tests/v1/worker/test_ubatch_inputs.py @@ -0,0 +1,141 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +import pytest +import torch + +from vllm.v1.worker.ubatch_inputs import ( + slice_lookback_token_ids, + update_captured_lookback, +) + + +@pytest.mark.parametrize( + "request_slice,token_slice,expected", + [ + (slice(0, 1), slice(0, 2), [[9, 8, 7], [-1, -1, -1]]), + (slice(0, 2), slice(2, 5), [[11, 10, 9], [29, 28, 27], [-1, -1, -1]]), + (slice(1, 2), slice(3, 5), [[29, 28, 27], [-1, -1, -1]]), + (slice(0, 1), slice(1, 3), [[10, 9, 8], [-1, -1, -1]]), + ], +) +def test_lookbacks_follow_microbatch_chunk_start(request_slice, token_slice, expected): + """A split request reads previous batch tokens before its older history.""" + actual = slice_lookback_token_ids( + torch.tensor([[9, 8, 7], [29, 28, 27]], dtype=torch.int32), + torch.tensor([10, 11, 12, 30, 31]), + torch.tensor([0, 3, 5]), + request_slice, + token_slice, + ) + assert actual.tolist() == expected + + +def test_lookback_padding_does_not_read_the_last_live_request(): + actual = slice_lookback_token_ids( + torch.tensor([[5, -1, -1]]), + torch.tensor([6, 0, 0]), + torch.tensor([0, 1]), + slice(0, 3), + slice(0, 3), + ) + assert actual.tolist() == [[5, -1, -1], [-1, -1, -1], [-1, -1, -1]] + + +def test_replay_refreshes_history_and_clears_unused_request_rows(): + captured = torch.tensor([[1, 2, 3], [4, 5, 6]]) + address = captured.data_ptr() + current = slice_lookback_token_ids( + torch.tensor([[7, 8, 9]]), + torch.tensor([10, 11]), + torch.tensor([0, 2]), + slice(0, 1), + slice(0, 2), + ) + update_captured_lookback(captured, current) + assert captured.data_ptr() == address + assert captured.tolist() == [[7, 8, 9], [-1, -1, -1]] + + +def test_replay_rejects_a_changed_history_signature(): + with pytest.raises(ValueError, match="graph signature"): + update_captured_lookback(torch.empty(2, 3), torch.empty(3, 3)) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +@pytest.mark.parametrize("capture", [False, True]) +def test_wrapper_preserves_split_request_history_across_replay(capture): + """Actual microbatch threads and graph replay must receive refreshed history.""" + from vllm.config import CUDAGraphMode, ParallelConfig, VllmConfig + from vllm.forward_context import ( + BatchDescriptor, + DPMetadata, + create_forward_context, + override_forward_context, + ) + from vllm.v1.worker.gpu_ubatch_wrapper import UBatchWrapper + from vllm.v1.worker.ubatch_utils import UBatchSlice + + config = VllmConfig( + parallel_config=ParallelConfig(data_parallel_size=2, is_moe_model=True) + ) + # The callable below has no attention or MoE layers to configure. + config.parallel_config.ubatch_size = 2 + config.parallel_config.all2all_backend = "deepep_high_throughput" + + def model( + *, input_ids, positions, intermediate_tensors, inputs_embeds, lookback_token_ids + ): + return lookback_token_ids.clone() + + mode = CUDAGraphMode.FULL if capture else CUDAGraphMode.NONE + wrapper = UBatchWrapper(model, config, mode, torch.device("cuda:0")) + ids = torch.arange(10, 18, device="cuda") + positions = torch.arange(8, device="cuda") + history = torch.tensor([[9, 8, 7], [29, 28, 27]], device="cuda") + starts = torch.tensor([0, 6, 8], device="cuda", dtype=torch.int32) + slices = [ + UBatchSlice(slice(0, 1), slice(0, 4)), + UBatchSlice(slice(0, 2), slice(4, 8)), + ] + compute_stream = torch.cuda.Stream() + + def run(runtime_mode): + context = create_forward_context( + None, + config, + dp_metadata=DPMetadata(torch.tensor([8, 8])), + cudagraph_runtime_mode=runtime_mode, + batch_descriptor=BatchDescriptor(num_tokens=8), + ubatch_slices=slices, + ) + compute_stream.wait_stream(torch.cuda.current_stream()) + with override_forward_context(context), torch.cuda.stream(compute_stream): + output = wrapper( + input_ids=ids, + positions=positions, + intermediate_tensors=None, + inputs_embeds=None, + lookback_token_ids=history, + lookback_query_start_loc=starts, + ) + torch.cuda.current_stream().wait_stream(compute_stream) + return output + + # Initialize thread-local CUDA handles before capture. + run(CUDAGraphMode.NONE) + if capture: + run(mode) + for change in (0, 100, 200): + ids.copy_(torch.arange(10, 18, device="cuda") + change) + history.copy_(torch.tensor([[9, 8, 7], [29, 28, 27]], device="cuda") + change) + if change == 200: + starts.copy_(torch.tensor([0, 8, 8], device="cuda")) + slices[1] = UBatchSlice(slice(0, 1), slice(4, 8)) + actual = run(mode) + expected = torch.full((8, 3), -1, device="cuda", dtype=torch.int64) + expected[0] = history[0] + expected[4] = ids[torch.tensor([3, 2, 1], device="cuda")] + if change != 200: + expected[5] = history[1] + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + wrapper.clear_graphs() diff --git a/vllm/models/deepseek_v41/common/engram.py b/vllm/models/deepseek_v41/common/engram.py index c415bbaa6e42..bc8b473416f9 100644 --- a/vllm/models/deepseek_v41/common/engram.py +++ b/vllm/models/deepseek_v41/common/engram.py @@ -51,6 +51,7 @@ from vllm.model_executor.layers.quantization import QuantizationConfig from vllm.model_executor.utils import set_weight_attrs from vllm.triton_utils import tl, triton +from vllm.v1.worker.ubatching import dbo_current_ubatch_id logger = init_logger(__name__) @@ -960,19 +961,43 @@ def _init_staging(self, max_tokens: int, head_dim: int) -> None: head_dim, dtype=torch.bfloat16, ) + parallel_config = get_current_vllm_config().parallel_config + self._init_lookup_staging( + parallel_config.num_ubatches if parallel_config.use_ubatching else 1, + ) + + def _init_lookup_staging(self, num_slots: int) -> None: + self._lookup_rows = [self.staged_rows] + [ + torch.empty_like(self.staged_rows) for _ in range(num_slots - 1) + ] + + def _lookup_slot(self) -> int: + rows = getattr(self, "_lookup_rows", None) + return dbo_current_ubatch_id() if rows is not None and len(rows) > 1 else 0 - def prepare_embeddings(self, hash_ids: torch.Tensor) -> None: - """Look up rows before the decoder layers consume them.""" - rows = self.staged_rows[: hash_ids.shape[0]] + def _lookup_staging(self) -> torch.Tensor: + buffers = getattr(self, "_lookup_rows", None) + return buffers[self._lookup_slot()] if buffers is not None else self.staged_rows + + def prepare_embeddings(self, hash_ids: torch.Tensor) -> torch.Tensor: + """Stage lookup rows in the current microbatch's own buffer.""" + rows = self._lookup_staging()[: hash_ids.shape[0]] assert rows.shape[0] == hash_ids.shape[0], "engram staging buffer too small" self.embed_tokens.lookup(hash_ids, rows) + return rows - def _ready_rows(self, num_tokens: int) -> torch.Tensor: - return self.staged_rows[:num_tokens] + def _ready_rows( + self, num_tokens: int, prepared_rows: torch.Tensor | None = None + ) -> torch.Tensor: + if prepared_rows is None: + prepared_rows = self._lookup_staging() + return prepared_rows[:num_tokens] - def embed(self, hash_ids: torch.Tensor) -> torch.Tensor: + def embed( + self, hash_ids: torch.Tensor, prepared_rows: torch.Tensor | None = None + ) -> torch.Tensor: """Gather heads, returning only local tokens when SP is enabled.""" - rows = self._ready_rows(hash_ids.shape[0]) + rows = self._ready_rows(hash_ids.shape[0], prepared_rows) if self.embed_tokens.tp_size == 1: return rows[:, : self.embed_tokens.n_hash_cols] if self.use_sequence_parallel: @@ -997,11 +1022,12 @@ def forward( hidden_states: torch.Tensor, hash_ids: torch.Tensor, token_mask: torch.Tensor | None = None, + prepared_rows: torch.Tensor | None = None, ) -> torch.Tensor: """hidden_states: [T, hc_mult, dim]; hash_ids: [T, n_hash_cols] (all tokens, pre sequence-parallel shard); token_mask: [T], False shuts the gate so those positions pass through untouched.""" - kv = self.wkv(self.embed(hash_ids).flatten(-2)) + kv = self.wkv(self.embed(hash_ids, prepared_rows).flatten(-2)) num_kv_tokens = hash_ids.shape[0] assert token_mask is None or token_mask.shape == (num_kv_tokens,) if self.use_sequence_parallel: diff --git a/vllm/models/deepseek_v41/nvidia/engram.py b/vllm/models/deepseek_v41/nvidia/engram.py index 624a36c214fc..3f6b9cf3d4d0 100644 --- a/vllm/models/deepseek_v41/nvidia/engram.py +++ b/vllm/models/deepseek_v41/nvidia/engram.py @@ -370,13 +370,14 @@ def _init_staging(self, max_tokens: int, head_dim: int) -> None: if self.embed_tokens.cpu_offload: self._prefetch_stream = torch.cuda.Stream(device=self.staged_rows.device) - def prepare_embeddings(self, hash_ids: torch.Tensor) -> None: + def prepare_embeddings(self, hash_ids: torch.Tensor) -> torch.Tensor: """Prefetch local shared rows or the DP group's gathered hash IDs.""" if self._prefetch_stream is None: return super().prepare_embeddings(hash_ids) - rows = self.staged_rows[: hash_ids.shape[0]] + rows = self._lookup_staging()[: hash_ids.shape[0]] assert rows.shape[0] == hash_ids.shape[0], "engram staging buffer too small" self._start_prefetch(hash_ids, rows, self._prefetch_stream) + return rows @eager_break_during_capture def _start_prefetch( @@ -393,11 +394,15 @@ def _start_prefetch( def _finish_prefetch(self, stream: torch.cuda.Stream) -> None: torch.cuda.current_stream().wait_stream(stream) - def _ready_rows(self, num_tokens: int) -> torch.Tensor: + def _ready_rows( + self, num_tokens: int, prepared_rows: torch.Tensor | None = None + ) -> torch.Tensor: if self._prefetch_stream is not None: self._finish_prefetch(self._prefetch_stream) if self.embed_tokens.dp_size > 1: slot = engram_gathered_num_tokens() - staged = self.staged_rows[: slot * self.embed_tokens.dp_size] + staged = super()._ready_rows( + slot * self.embed_tokens.dp_size, prepared_rows + ) return _gather_engram_rows(staged, num_tokens) - return super()._ready_rows(num_tokens) + return super()._ready_rows(num_tokens, prepared_rows) diff --git a/vllm/models/deepseek_v41/nvidia/model.py b/vllm/models/deepseek_v41/nvidia/model.py index f3530544da7b..54e8aebaf21c 100644 --- a/vllm/models/deepseek_v41/nvidia/model.py +++ b/vllm/models/deepseek_v41/nvidia/model.py @@ -323,6 +323,7 @@ def forward( residual: torch.Tensor | None = None, engram_hashes: torch.Tensor | None = None, engram_mask: torch.Tensor | None = None, + engram_rows: torch.Tensor | None = None, *, capture_previous_aux: bool = False, ) -> tuple[ @@ -385,6 +386,7 @@ def forward( previous_post, engram_hashes[:, self.engram.layer_hash_index], engram_mask, + prepared_rows=engram_rows, ) post_mix, res_mix, x, attn_pre = mhc_pre_delayed_tilelang( residual, @@ -630,6 +632,7 @@ def forward( # profile runs (KV cache unbound). engram_hashes: torch.Tensor | None = None engram_mask: torch.Tensor | None = None + engram_rows: dict[int, torch.Tensor] = {} if ( self.engram_hash is not None and input_ids is not None @@ -681,8 +684,10 @@ def forward( for layer in islice(self.layers, self.start_layer, self.end_layer): engram = getattr(layer, "engram", None) if engram is not None: - engram.prepare_embeddings( - gathered_hashes[:, engram.layer_hash_index] + engram_rows[engram.layer_hash_index] = ( + engram.prepare_embeddings( + gathered_hashes[:, engram.layer_hash_index] + ) ) full_num_tokens = positions.shape[0] @@ -707,6 +712,9 @@ def forward( islice(self.layers, self.start_layer, self.end_layer), start=self.start_layer, ): + prepared_rows = None + if layer.engram is not None and engram_hashes is not None: + prepared_rows = engram_rows[layer.engram.layer_hash_index] hidden_states, residual, post_mix, res_mix, pre_mix, previous_aux = layer( hidden_states, positions, @@ -717,6 +725,7 @@ def forward( residual, engram_hashes, engram_mask, + prepared_rows, capture_previous_aux=idx in self.aux_hidden_state_layers, ) if previous_aux is not None: diff --git a/vllm/v1/worker/gpu_model_runner.py b/vllm/v1/worker/gpu_model_runner.py index 97a59bd273c3..af1abc138830 100644 --- a/vllm/v1/worker/gpu_model_runner.py +++ b/vllm/v1/worker/gpu_model_runner.py @@ -1096,6 +1096,8 @@ def _init_model_kwargs(self, num_reqs: int | None = None): model_kwargs["lookback_token_ids"] = self._prepare_lookback_token_ids( num_reqs ) + if self.parallel_config.use_ubatching: + model_kwargs["lookback_query_start_loc"] = self.query_start_loc.gpu if not self.is_pooling_model: return model_kwargs diff --git a/vllm/v1/worker/gpu_ubatch_wrapper.py b/vllm/v1/worker/gpu_ubatch_wrapper.py index 0735b39f81c2..d68b94112fd0 100644 --- a/vllm/v1/worker/gpu_ubatch_wrapper.py +++ b/vllm/v1/worker/gpu_ubatch_wrapper.py @@ -22,6 +22,10 @@ from vllm.model_executor.offloader.base import get_offloader from vllm.platforms import current_platform from vllm.sequence import IntermediateTensors +from vllm.v1.worker.ubatch_inputs import ( + slice_lookback_token_ids, + update_captured_lookback, +) from vllm.v1.worker.ubatch_utils import create_sm_control_context from vllm.v1.worker.ubatching import UBatchContext, make_ubatch_contexts @@ -53,12 +57,18 @@ class UbatchMetadata: inputs_embeds: torch.Tensor | None intermediate_tensors: IntermediateTensors | None num_tokens: int + lookback_token_ids: torch.Tensor | None = None + + def model_kwargs(self) -> dict[str, torch.Tensor]: + if self.lookback_token_ids is None: + return {} + return {"lookback_token_ids": self.lookback_token_ids} @dataclass class CUDAGraphMetaData: cudagraph: torch.cuda.CUDAGraph - ubatch_metadata: UbatchMetadata + ubatch_metadata: list[UbatchMetadata] outputs: Any | None = None @@ -155,6 +165,7 @@ def _capture_ubatch_thread(results, ubatch_metadata): positions=ubatch_metadata.positions, intermediate_tensors=ubatch_metadata.intermediate_tensors, inputs_embeds=ubatch_metadata.inputs_embeds, + **ubatch_metadata.model_kwargs(), ) results.append((ubatch_metadata.context.id, model_output)) @@ -220,6 +231,7 @@ def _ubatch_thread(results, model, ubatch_metadata): positions=ubatch_metadata.positions, intermediate_tensors=ubatch_metadata.intermediate_tensors, inputs_embeds=ubatch_metadata.inputs_embeds, + **ubatch_metadata.model_kwargs(), ) results.append((ubatch_metadata.context.id, model_output)) @@ -262,6 +274,7 @@ def _make_ubatch_metadata( dp_metadata, batch_descriptor, cudagraph_runtime_mode, + lookback_inputs: list[torch.Tensor | None] | None = None, ) -> list[UbatchMetadata]: # Create one forward context per ubatch forward_contexts = [] @@ -311,6 +324,9 @@ def _make_ubatch_metadata( intermediate_tensors=sliced_intermediate_tensors, num_tokens=ubatch_slice.token_slice.stop - ubatch_slice.token_slice.start, + lookback_token_ids=( + lookback_inputs[i] if lookback_inputs is not None else None + ), ) ) @@ -348,6 +364,7 @@ def _slice_model_inputs( ) def __call__(self, *args, **kwargs): + lookback_query_start_loc = kwargs.pop("lookback_query_start_loc", None) forward_context = get_forward_context() batch_descriptor = forward_context.batch_descriptor ubatch_slices = forward_context.ubatch_slices @@ -380,6 +397,24 @@ def __call__(self, *args, **kwargs): intermediate_tensors = kwargs["intermediate_tensors"] inputs_embeds = kwargs["inputs_embeds"] compute_stream = torch.cuda.current_stream() + history = kwargs.get("lookback_token_ids") + lookback_inputs: list[torch.Tensor | None] = [None] * len(ubatch_slices) + if history is not None: + if lookback_query_start_loc is None or input_ids is None: + raise ValueError( + "Microbatch lookbacks require input_ids and " + "lookback_query_start_loc" + ) + lookback_inputs = [ + slice_lookback_token_ids( + history, + input_ids, + lookback_query_start_loc, + ubatch_slice.request_slice, + ubatch_slice.token_slice, + ) + for ubatch_slice in ubatch_slices + ] dp_metadata = forward_context.dp_metadata @@ -415,6 +450,7 @@ def __call__(self, *args, **kwargs): dp_metadata=ubatch_dp_metadata, batch_descriptor=batch_descriptor, cudagraph_runtime_mode=CUDAGraphMode.NONE, + lookback_inputs=lookback_inputs, ) with self.sm_control: return self._capture_ubatches(ubatch_metadata, self.runnable) @@ -423,6 +459,10 @@ def __call__(self, *args, **kwargs): and cudagraph_runtime_mode is CUDAGraphMode.FULL ): cudagraph_metadata = self.cudagraphs[num_tokens] + for metadata, lookback in zip( + cudagraph_metadata.ubatch_metadata, lookback_inputs, strict=True + ): + update_captured_lookback(metadata.lookback_token_ids, lookback) # Sync offloader before replay - ensures any external dependencies # from pre-capture prefetches are satisfied. get_offloader().sync_prev_onload() @@ -441,6 +481,7 @@ def __call__(self, *args, **kwargs): dp_metadata=ubatch_dp_metadata, batch_descriptor=batch_descriptor, cudagraph_runtime_mode=CUDAGraphMode.NONE, + lookback_inputs=lookback_inputs, ) with self.sm_control: return self._run_ubatches(ubatch_metadata, self.runnable) diff --git a/vllm/v1/worker/ubatch_inputs.py b/vllm/v1/worker/ubatch_inputs.py new file mode 100644 index 000000000000..34e748bc8a0d --- /dev/null +++ b/vllm/v1/worker/ubatch_inputs.py @@ -0,0 +1,63 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Request history inputs at microbatch boundaries.""" + +import torch + + +def slice_lookback_token_ids( + history: torch.Tensor, + input_ids: torch.Tensor, + query_start_loc: torch.Tensor, + request_slice: slice, + token_slice: slice, +) -> torch.Tensor: + """Rebase request lookbacks to the first token executed by this microbatch. + + Rows are padded to the token count so graph replay can change request + lengths without changing the captured history buffer's address or shape. + """ + num_tokens = token_slice.stop - token_slice.start + num_reqs = request_slice.stop - request_slice.start + depth = history.shape[1] + output = history.new_full((num_tokens, depth), -1) + if not num_tokens or not num_reqs: + return output + if num_reqs > num_tokens: + raise ValueError("A microbatch cannot have more request rows than token rows") + if history.shape[0] == 0 or query_start_loc.numel() < 2: + return output + req_ids = torch.arange( + request_slice.start, request_slice.stop, device=history.device + ) + num_rows = min(history.shape[0], query_start_loc.numel() - 1) + safe_req_ids = req_ids.clamp(0, num_rows - 1) + starts = query_start_loc[safe_req_ids].to(torch.int64) + ends = query_start_loc[safe_req_ids + 1] + active = ( + (req_ids < num_rows) & (starts < token_slice.stop) & (ends > token_slice.start) + ) + chunk_starts = starts.clamp_min(token_slice.start) + positions = chunk_starts[:, None] - 1 - torch.arange(depth, device=history.device) + from_batch = positions >= starts[:, None] + batch_ids = input_ids[positions.clamp(0, input_ids.numel() - 1)].to(history.dtype) + history_cols = starts[:, None] - 1 - positions + previous_ids = history[safe_req_ids[:, None], history_cols.clamp(0, depth - 1)] + valid = active[:, None] & (from_batch | (history_cols < depth)) + output[:num_reqs] = torch.where( + valid, torch.where(from_batch, batch_ids, previous_ids), -1 + ) + return output + + +def update_captured_lookback( + captured: torch.Tensor | None, current: torch.Tensor | None +) -> None: + """Refresh graph inputs without replacing their captured storage.""" + if captured is None and current is None: + return + if captured is None or current is None or captured.shape != current.shape: + raise ValueError( + "Microbatch lookback inputs changed their CUDA graph signature" + ) + captured.copy_(current) diff --git a/vllm/v1/worker/ubatch_utils.py b/vllm/v1/worker/ubatch_utils.py index 9743322c4165..93e8755e29b8 100644 --- a/vllm/v1/worker/ubatch_utils.py +++ b/vllm/v1/worker/ubatch_utils.py @@ -169,7 +169,7 @@ def maybe_create_ubatch_slices( num_tokens_padded: int, num_reqs_padded: int, num_ubatches: int, - split_point: list[int] | int | None = None, + split_point: int | None = None, ) -> tuple[UBatchSlices | None, UBatchSlices | None]: if not should_ubatch: return None, None @@ -188,7 +188,7 @@ def maybe_create_ubatch_slices( start_token = 0 # Add the end point to the split points to make iteration easier - all_points = token_split_points + [cu_num_tokens[-1]] + all_points = token_split_points + [int(cu_num_tokens[-1])] for end_token in all_points: token_slice = slice(start_token, end_token) @@ -330,6 +330,11 @@ def _make_metadata_with_slice( max_seq_len=max_seq_len, block_table_tensor=block_table_tensor, slot_mapping=slot_mapping, + positions=( + attn_metadata.positions[token_slice] + if attn_metadata.positions is not None + else None + ), seq_lens_cpu_upper_bound=seq_lens_cpu_upper_bound, )