Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
40 changes: 40 additions & 0 deletions tests/kernels/test_engram.py
Original file line number Diff line number Diff line change
Expand Up @@ -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":
Expand Down Expand Up @@ -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",
Expand Down
11 changes: 11 additions & 0 deletions tests/v1/attention/test_attention_splitting.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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()
Expand All @@ -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
Expand Down
141 changes: 141 additions & 0 deletions tests/v1/worker/test_ubatch_inputs.py
Original file line number Diff line number Diff line change
@@ -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()
42 changes: 34 additions & 8 deletions vllm/models/deepseek_v41/common/engram.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)

Expand Down Expand Up @@ -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:
Expand All @@ -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:
Expand Down
15 changes: 10 additions & 5 deletions vllm/models/deepseek_v41/nvidia/engram.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand All @@ -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)
13 changes: 11 additions & 2 deletions vllm/models/deepseek_v41/nvidia/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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[
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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]
Expand All @@ -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,
Expand All @@ -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:
Expand Down
Loading
Loading