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
185 changes: 185 additions & 0 deletions tests/v1/core/test_prefix_caching.py
Original file line number Diff line number Diff line change
Expand Up @@ -2267,6 +2267,191 @@ def test_null_parent_block_hash():
assert blocks[num_full_blocks - 1].block_hash is not None


def test_block_stored_event_splits_around_null_blocks():
block_size = 4
pool = BlockPool(
num_gpu_blocks=4,
enable_caching=True,
hash_block_size=block_size,
enable_kv_cache_events=True,
)
req = make_request(
"req_null_event_split",
prompt_token_ids=list(range(2 * block_size)),
block_size=block_size,
hash_fn=sha256,
)
blocks = [pool.null_block, *pool.get_new_blocks(1)]

pool.cache_full_blocks(
request=req,
blocks=blocks,
num_cached_blocks=0,
num_full_blocks=2,
block_size=block_size,
kv_cache_group_id=0,
)

events = pool.take_events()
assert len(events) == 1
event = events[0]
assert isinstance(event, BlockStored)
assert event.parent_block_hash == kv_cache_utils.maybe_convert_block_hash(
req.block_hashes[0]
)
assert event.block_hashes == [
kv_cache_utils.maybe_convert_block_hash(req.block_hashes[1])
]
assert event.extra_keys == [None]
assert event.token_ids == list(req.all_token_ids[block_size:])
assert event.skipped_parent_block_hash is None
assert event.skipped_token_ids == list(req.all_token_ids[:block_size])
assert event.skipped_extra_keys == [None]
assert pool.null_block.block_hash is None
assert blocks[1].block_hash is not None


def test_block_stored_event_splits_around_masked_blocks():
block_size = 4
pool = BlockPool(
num_gpu_blocks=4,
enable_caching=True,
hash_block_size=block_size,
enable_kv_cache_events=True,
)
req = make_request(
"req_masked_event_split",
prompt_token_ids=list(range(2 * block_size)),
block_size=block_size,
hash_fn=sha256,
)
blocks = pool.get_new_blocks(2)

pool.cache_full_blocks(
request=req,
blocks=blocks,
num_cached_blocks=0,
num_full_blocks=2,
block_size=block_size,
kv_cache_group_id=0,
block_mask=[False, True],
)

events = pool.take_events()
assert len(events) == 1
event = events[0]
assert isinstance(event, BlockStored)
assert event.parent_block_hash == kv_cache_utils.maybe_convert_block_hash(
req.block_hashes[0]
)
assert event.block_hashes == [
kv_cache_utils.maybe_convert_block_hash(req.block_hashes[1])
]
assert event.extra_keys == [None]
assert event.token_ids == list(req.all_token_ids[block_size:])
assert event.skipped_parent_block_hash is None
assert event.skipped_token_ids == list(req.all_token_ids[:block_size])
assert event.skipped_extra_keys == [None]
assert blocks[0].block_hash is None
assert blocks[1].block_hash is not None


def test_block_stored_event_emits_dense_runs_around_masked_block():
block_size = 4
pool = BlockPool(
num_gpu_blocks=4,
enable_caching=True,
hash_block_size=block_size,
enable_kv_cache_events=True,
)
req = make_request(
"req_contiguous_prefix_event",
prompt_token_ids=list(range(3 * block_size)),
block_size=block_size,
hash_fn=sha256,
)
blocks = pool.get_new_blocks(3)

pool.cache_full_blocks(
request=req,
blocks=blocks,
num_cached_blocks=0,
num_full_blocks=3,
block_size=block_size,
kv_cache_group_id=0,
block_mask=[True, False, True],
)

events = pool.take_events()
assert len(events) == 2
first_event, second_event = events
assert isinstance(first_event, BlockStored)
assert first_event.parent_block_hash is None
assert first_event.block_hashes == [
kv_cache_utils.maybe_convert_block_hash(req.block_hashes[0])
]
assert first_event.extra_keys == [None]
assert first_event.token_ids == list(req.all_token_ids[:block_size])
assert first_event.skipped_parent_block_hash is None
assert first_event.skipped_token_ids is None
assert first_event.skipped_extra_keys is None
assert isinstance(second_event, BlockStored)
expected_parent = kv_cache_utils.maybe_convert_block_hash(req.block_hashes[1])
assert second_event.parent_block_hash == expected_parent
assert second_event.block_hashes == [
kv_cache_utils.maybe_convert_block_hash(req.block_hashes[2])
]
assert second_event.extra_keys == [None]
assert second_event.token_ids == list(req.all_token_ids[2 * block_size :])
expected_skipped_parent = kv_cache_utils.maybe_convert_block_hash(
req.block_hashes[0]
)
assert second_event.skipped_parent_block_hash == expected_skipped_parent
assert second_event.skipped_token_ids == list(
req.all_token_ids[block_size : 2 * block_size]
)
assert second_event.skipped_extra_keys == [None]
assert all(block.block_hash is not None for block in blocks[::2])
assert blocks[1].block_hash is None


def test_block_stored_event_skipped_context_includes_extra_keys():
block_size = 4
pool = BlockPool(
num_gpu_blocks=4,
enable_caching=True,
hash_block_size=block_size,
enable_kv_cache_events=True,
)
req = make_request(
"req_skipped_event_context_extra_keys",
prompt_token_ids=list(range(2 * block_size)),
block_size=block_size,
hash_fn=sha256,
cache_salt="salt",
)
blocks = pool.get_new_blocks(2)

pool.cache_full_blocks(
request=req,
blocks=blocks,
num_cached_blocks=0,
num_full_blocks=2,
block_size=block_size,
kv_cache_group_id=0,
block_mask=[False, True],
)

events = pool.take_events()
assert len(events) == 1
event = events[0]
assert isinstance(event, BlockStored)
assert event.extra_keys == [None]
assert event.skipped_parent_block_hash is None
assert event.skipped_token_ids == list(req.all_token_ids[:block_size])
assert event.skipped_extra_keys == [("salt",)]


@pytest.mark.parametrize("blocks_to_cache", [2, 3, 10])
def test_kv_cache_events_with_lora(blocks_to_cache: int):
"""Test BlockStored events contain correct lora_id when using LoRA requests."""
Expand Down
13 changes: 13 additions & 0 deletions vllm/distributed/kv_events.py
Original file line number Diff line number Diff line change
Expand Up @@ -71,6 +71,16 @@ class BlockStored(KVCacheEvent):
"""Store events carry cache-spec metadata so consumers can classify and
filter groups as they are learned. Remove events only need group_idx+hash.
"""

skipped_parent_block_hash: ExternalBlockHash | None = None
"""Parent hash for skipped token context preceding this store event."""

skipped_token_ids: list[int] | None = None
"""Logical token span for skipped blocks preceding this store event."""

skipped_extra_keys: list[tuple[Any, ...] | None] | None = None
"""Extra keys for ``skipped_token_ids``, one entry per skipped block."""

group_idx: int | None = None
kv_cache_spec_kind: str | None = None
kv_cache_spec_sliding_window: int | None = None
Expand All @@ -95,6 +105,9 @@ def __hash__(self) -> int:
self.lora_id,
self.medium,
tuple(self.extra_keys) if self.extra_keys else None,
self.skipped_parent_block_hash,
tuple(self.skipped_token_ids) if self.skipped_token_ids else None,
tuple(self.skipped_extra_keys) if self.skipped_extra_keys else None,
self.group_idx,
self.kv_cache_spec_kind,
self.kv_cache_spec_sliding_window,
Expand Down
115 changes: 90 additions & 25 deletions vllm/v1/core/block_pool.py
Original file line number Diff line number Diff line change
Expand Up @@ -265,14 +265,26 @@ def cache_full_blocks(
)

new_block_hashes = block_hashes[num_cached_blocks:]
new_hashes: list[ExternalBlockHash] | None = (
[] if self.enable_kv_cache_events else None
)
# KV events must stay dense: split around skipped logical blocks so each
# event's parent is the immediate predecessor of its first block.
event_runs: list[tuple[int, int, list[ExternalBlockHash]]] = []
event_start_idx: int | None = None
event_hashes: list[ExternalBlockHash] = []
for i, blk in enumerate(new_full_blocks):
# Some blocks may be null or masked out when enabling sparse attention
# like sliding window attention, or Mamba models with prefix-caching
# in align mode. We skip null blocks here.
if blk.is_null or (block_mask is not None and not block_mask[i]):
if self.enable_kv_cache_events and event_start_idx is not None:
event_runs.append(
(
event_start_idx,
num_cached_blocks + i,
event_hashes,
)
)
event_start_idx = None
event_hashes = []
continue
block_hash = new_block_hashes[i]
num_hash_tokens = (num_cached_blocks + i + 1) * block_size
Expand All @@ -295,51 +307,92 @@ def cache_full_blocks(
blk,
num_tokens=num_hash_tokens,
)
if new_hashes is not None:
new_hashes.append(maybe_convert_block_hash(block_hash))
if self.enable_kv_cache_events:
if event_start_idx is None:
event_start_idx = num_cached_blocks + i
event_hashes.append(maybe_convert_block_hash(block_hash))

if self.enable_kv_cache_events:
if num_cached_blocks == 0:
if self.enable_kv_cache_events and event_start_idx is not None:
event_runs.append((event_start_idx, num_full_blocks, event_hashes))

last_event_end_idx = num_cached_blocks
for event_start_idx, event_end_idx, event_hashes in event_runs:
if event_start_idx == 0:
parent_block_hash: ExternalBlockHash | None = None
else:
parent_block_hash = maybe_convert_block_hash(
block_hashes[num_cached_blocks - 1]
block_hashes[event_start_idx - 1]
)

if event_start_idx > last_event_end_idx:
skipped_parent_block_hash = (
None
if last_event_end_idx == 0
else maybe_convert_block_hash(block_hashes[last_event_end_idx - 1])
)
skipped_start_token_idx = last_event_end_idx * block_size
skipped_end_token_idx = event_start_idx * block_size
skipped_extra_keys = self._generate_block_extra_keys(
request,
last_event_end_idx,
event_start_idx,
block_size,
)
else:
skipped_parent_block_hash = None
skipped_start_token_idx = None
skipped_end_token_idx = None
skipped_extra_keys = None

# Calculate token range for the blocks being cached
start_token_idx = num_cached_blocks * block_size
end_token_idx = num_full_blocks * block_size
start_token_idx = event_start_idx * block_size
end_token_idx = event_end_idx * block_size

# Generate extra keys for each block individually.
# Each block may have different extra_keys (e.g., different MM
# features, or cache_salt only for the first block).
# Skip null/masked-out blocks to match the length of new_hashes.
extra_keys_list: list[tuple[Any, ...] | None] = []
curr_mm_idx = 0
for i in range(num_cached_blocks, num_full_blocks):
if blocks[i].is_null:
continue
if block_mask is not None and not block_mask[i - num_cached_blocks]:
continue
block_start = i * block_size
block_end = block_start + block_size
extra_keys, curr_mm_idx = generate_block_hash_extra_keys(
request, block_start, block_end, curr_mm_idx
)
extra_keys_list.append(extra_keys)
extra_keys_list = self._generate_block_extra_keys(
request,
event_start_idx,
event_end_idx,
block_size,
)

self.kv_event_queue.append(
self._build_block_stored_event(
request,
block_hashes=new_hashes,
block_hashes=event_hashes,
parent_block_hash=parent_block_hash,
start_token_idx=start_token_idx,
end_token_idx=end_token_idx,
block_size=block_size,
kv_cache_group_id=kv_cache_group_id,
extra_keys_list=extra_keys_list,
skipped_parent_block_hash=skipped_parent_block_hash,
skipped_start_token_idx=skipped_start_token_idx,
skipped_end_token_idx=skipped_end_token_idx,
skipped_extra_keys=skipped_extra_keys,
)
)
last_event_end_idx = event_end_idx

def _generate_block_extra_keys(
self,
request: Request,
start_block_idx: int,
end_block_idx: int,
block_size: int,
) -> list[tuple[Any, ...] | None]:
extra_keys_list: list[tuple[Any, ...] | None] = []
curr_mm_idx = 0
for i in range(start_block_idx, end_block_idx):
block_start = i * block_size
block_end = block_start + block_size
extra_keys, curr_mm_idx = generate_block_hash_extra_keys(
request, block_start, block_end, curr_mm_idx
)
extra_keys_list.append(extra_keys)
return extra_keys_list

def _build_block_stored_event(
self,
Expand All @@ -351,6 +404,10 @@ def _build_block_stored_event(
block_size: int,
kv_cache_group_id: int,
extra_keys_list: list[tuple[Any, ...] | None],
skipped_parent_block_hash: ExternalBlockHash | None = None,
skipped_start_token_idx: int | None = None,
skipped_end_token_idx: int | None = None,
skipped_extra_keys: list[tuple[Any, ...] | None] | None = None,
) -> BlockStored:
"""Build a ``BlockStored`` KV event for ``request``.

Expand All @@ -367,6 +424,14 @@ def _build_block_stored_event(
medium=MEDIUM_GPU,
lora_name=request.lora_request.name if request.lora_request else None,
extra_keys=extra_keys_list if extra_keys_list else None,
skipped_parent_block_hash=skipped_parent_block_hash,
skipped_token_ids=request.all_token_ids[
skipped_start_token_idx:skipped_end_token_idx
]
if skipped_start_token_idx is not None
and skipped_end_token_idx is not None
else None,
skipped_extra_keys=skipped_extra_keys,
group_idx=kv_cache_group_id,
session_id=request.session_id,
)
Expand Down
Loading