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
37 changes: 37 additions & 0 deletions tests/v1/worker/test_kv_block_zeroer.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@ def test_block_ids_are_not_overwritten_while_copy_is_in_flight():
zeroer._meta = (
torch.tensor([storage.data_ptr()], dtype=torch.uint64, device=device),
torch.tensor([page_size_el], dtype=torch.int64, device=device),
torch.tensor([page_size_el], dtype=torch.int64, device=device),
page_size_el // page_size_el, # max_chunks = 1
page_size_el, # blk_size
1, # n_segs
Expand Down Expand Up @@ -71,6 +72,7 @@ def largest_power_of_2_divisor(n):
device=device,
),
torch.tensor(seg_page_sizes, dtype=torch.int64, device=device),
torch.tensor(seg_page_sizes, dtype=torch.int64, device=device),
max_ps // blk_size,
blk_size,
2,
Expand All @@ -86,3 +88,38 @@ def largest_power_of_2_divisor(n):
assert torch.all(storage[1] == 0)
assert torch.all(storage[2] == 0)
assert torch.all(storage[3] == 1)


@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
def test_packed_segment_zeros_only_its_last_block_page():
"""A packed KV segment steps by block stride but clears only its page."""
device = torch.device("cuda")
num_blocks = 4
block_stride_el = 12
page_size_el = 4
page_offset_el = 3
backing = torch.ones(
(num_blocks, block_stride_el), dtype=torch.int32, device=device
)

zeroer = KVBlockZeroer.__new__(KVBlockZeroer)
zeroer.device = device
zeroer._meta = (
torch.tensor(
[backing.data_ptr() + page_offset_el * backing.element_size()],
dtype=torch.uint64,
device=device,
),
torch.tensor([block_stride_el], dtype=torch.int64, device=device),
torch.tensor([page_size_el], dtype=torch.int64, device=device),
1,
page_size_el,
1,
)

zeroer.zero_block_ids([num_blocks - 1])
torch.accelerator.synchronize()

expected = torch.ones_like(backing)
expected[-1, page_offset_el : page_offset_el + page_size_el] = 0
assert torch.equal(backing, expected)
58 changes: 41 additions & 17 deletions vllm/v1/worker/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@
@triton.jit
def _zero_kv_blocks_kernel(
seg_addrs_ptr,
seg_block_strides_ptr,
seg_page_sizes_ptr,
block_ids_ptr,
n_blocks,
Expand All @@ -57,10 +58,10 @@ def _zero_kv_blocks_kernel(
buffer. For backends where K/V is outermost (block_dim=1) there are
two segments per buffer (one for K, one for V).

Segments may have different page sizes (e.g. models with multiple KV
cache groups like MLA + DSA indexer). Each segment's page size is
read from seg_page_sizes_ptr; programs whose chunk_index falls beyond
their segment's page size early-exit.
Segments may have different block strides and page sizes (e.g. packed
KV views or models with multiple KV cache groups like MLA + DSA
indexer). Each segment's block stride determines where a logical block
begins, while its page size determines how many elements are cleared.

seg_addrs_ptr holds absolute byte addresses (int64) for each segment,
allowing segments to live in different CUDA allocations.
Expand All @@ -75,14 +76,15 @@ def _zero_kv_blocks_kernel(
remainder = pid % work_per_block
seg_index = remainder // MAX_CHUNKS
chunk_index = remainder % MAX_CHUNKS
block_stride_el = tl.load(seg_block_strides_ptr + seg_index)
page_size_el = tl.load(seg_page_sizes_ptr + seg_index)
if chunk_index >= page_size_el // BLOCK_SIZE:
return
block_id = tl.load(block_ids_ptr + block_index)
seg_addr = tl.load(seg_addrs_ptr + seg_index)
ptr = tl.cast(seg_addr, tl.pointer_type(tl.int32))
offset = (
block_id.to(tl.int64) * page_size_el.to(tl.int64)
block_id.to(tl.int64) * block_stride_el.to(tl.int64)
+ chunk_index.to(tl.int64) * BLOCK_SIZE
)
cols = tl.arange(0, BLOCK_SIZE).to(tl.int64)
Expand Down Expand Up @@ -113,18 +115,21 @@ def __init__(

Block IDs from the scheduler reference logical blocks whose size
may differ from the kernel block size (virtual block splitting).
Each segment's page_size_el accounts for this ratio so that
``block_id * page_size_el`` lands at the correct offset.
Each virtual block is represented as an independent segment so its
physical block stride and zeroed page span remain independent.

Only AttentionSpec layers are processed; Mamba layers are skipped.
"""
self.device = device
self._meta: tuple[torch.Tensor, torch.Tensor, int, int, int] | None = None
self._meta: (
tuple[torch.Tensor, torch.Tensor, torch.Tensor, int, int, int] | None
) = None

if runner_only_attn_layers is None:
runner_only_attn_layers = set()
seen_ptrs: set[int] = set()
seg_addrs: list[int] = []
seg_block_strides: list[int] = []
seg_page_sizes: list[int] = []

for group in attn_groups_iter:
Expand All @@ -134,6 +139,7 @@ def __init__(
if group.kv_cache_group_id >= len(kernel_block_sizes):
continue
kernel_bs = kernel_block_sizes[group.kv_cache_group_id]
assert spec.block_size % kernel_bs == 0
ratio = spec.block_size // kernel_bs
block_dim = group.backend.get_kv_cache_block_dim(
kernel_bs,
Expand All @@ -154,22 +160,31 @@ def __init__(
seen_ptrs.add(dp)

el = kv.element_size()
cur_bytes = kv.stride(block_dim) * el
assert cur_bytes % 4 == 0
kernel_block_el = cur_bytes // 4
cur_page_el = kernel_block_el * ratio

block_stride_bytes = cur_bytes
block_stride_bytes = kv.stride(block_dim) * el
assert block_stride_bytes % 4 == 0
assert kv.shape[block_dim] % ratio == 0
outer_dims = [
d
for d in range(block_dim)
if kv.stride(d) * el > block_stride_bytes
]
outer_strides = [kv.stride(d) * el for d in outer_dims]
inner_dims = [
d for d in range(kv.ndim) if d != block_dim and d not in outer_dims
]
kernel_page_bytes = el + sum(
(kv.shape[d] - 1) * kv.stride(d) * el for d in inner_dims
)
assert kernel_page_bytes % 4 == 0
logical_block_stride_bytes = block_stride_bytes * ratio
for outer in iprod(*(range(kv.shape[d]) for d in outer_dims)):
off_bytes = sum(i * s for i, s in zip(outer, outer_strides))
seg_addrs.append(dp + off_bytes)
seg_page_sizes.append(cur_page_el)
for virtual_index in range(ratio):
seg_addrs.append(
dp + off_bytes + virtual_index * block_stride_bytes
)
seg_block_strides.append(logical_block_stride_bytes // 4)
seg_page_sizes.append(kernel_page_bytes // 4)

if not seg_addrs:
self._meta = None
Expand All @@ -182,6 +197,7 @@ def __init__(
)
self._meta = (
torch.tensor(seg_addrs, dtype=torch.uint64, device=self.device),
torch.tensor(seg_block_strides, dtype=torch.int64, device=self.device),
torch.tensor(seg_page_sizes, dtype=torch.int64, device=self.device),
max_page_size_el // blk_size,
blk_size,
Expand All @@ -192,12 +208,20 @@ def zero_block_ids(self, block_ids: list[int]) -> None:
"""Zero the KV cache memory for the given block IDs."""
if not block_ids or self._meta is None:
return
seg_addrs, seg_page_sizes, max_chunks, blk_size, n_segs = self._meta
(
seg_addrs,
seg_block_strides,
seg_page_sizes,
max_chunks,
blk_size,
n_segs,
) = self._meta
n_blocks = len(block_ids)
idx = async_tensor_h2d(block_ids, device=self.device, dtype=torch.int64)
grid = (n_blocks * n_segs * max_chunks,)
_zero_kv_blocks_kernel[grid](
seg_addrs,
seg_block_strides,
seg_page_sizes,
idx,
n_blocks,
Expand Down