[Bugfix] Fix race condition in KVBlockZeroer (#48085)

Signed-off-by: Benjamin Chislett <bchislett@nvidia.com>
Signed-off-by: Benjamin Chislett <chislett.ben@gmail.com>
Co-authored-by: Wentao Ye <44945378+yewentao256@users.noreply.github.com>
This commit is contained in:
Benjamin Chislett
2026-07-09 09:18:19 +00:00
committed by GitHub
co-authored by Wentao Ye
parent 0206f10871
commit 1cd75b3dd4
4 changed files with 83 additions and 20 deletions
+44
View File
@@ -0,0 +1,44 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import pytest
import torch
from vllm.v1.worker.utils import KVBlockZeroer
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
def test_block_ids_are_not_overwritten_while_copy_is_in_flight():
device = torch.device("cuda")
num_blocks = 4
page_size_el = 4
storage = torch.ones((num_blocks, page_size_el), dtype=torch.int32, device=device)
# Build the minimal zeroer state directly so the test can focus on ID-buffer
# lifetime without constructing model attention groups.
zeroer = KVBlockZeroer.__new__(KVBlockZeroer)
zeroer.device = device
zeroer.pin_memory = True
zeroer.max_concurrency = 2
zeroer._id_cap = 8
zeroer._allocate_id_buffers()
zeroer._meta = (
torch.tensor([storage.data_ptr()], dtype=torch.uint64, device=device),
page_size_el,
page_size_el,
1,
)
stream = torch.cuda.Stream()
with torch.cuda.stream(stream):
# Keep the first nonblocking H2D copy pending while the host submits the
# second call. A single shared pinned source would be overwritten here.
torch.cuda._sleep(10_000_000)
zeroer.zero_block_ids([1])
zeroer.zero_block_ids([2])
stream.synchronize()
assert torch.all(storage[0] == 1)
assert torch.all(storage[1] == 0)
assert torch.all(storage[2] == 0)
assert torch.all(storage[3] == 1)
+1
View File
@@ -498,6 +498,7 @@ class GPUModelRunner(LoRAModelRunnerMixin):
kernel_block_sizes=self.kernel_block_sizes,
cache_dtype=self.cache_config.cache_dtype,
static_forward_context=self.compilation_config.static_forward_context,
max_concurrency=self.vllm_config.max_concurrent_batches,
)
@torch.inference_mode()
+1
View File
@@ -1123,6 +1123,7 @@ class GPUModelRunner(
cache_dtype=self.cache_config.cache_dtype,
runner_only_attn_layers=self.runner_only_attn_layers,
static_forward_context=self.compilation_config.static_forward_context,
max_concurrency=self.vllm_config.max_concurrent_batches,
)
def _zero_block_ids(self, block_ids: list[int]) -> None:
+37 -20
View File
@@ -94,6 +94,7 @@ class KVBlockZeroer:
cache_dtype: str,
static_forward_context: dict[str, Any],
runner_only_attn_layers: set[str] | None = None,
max_concurrency: int = 1,
) -> None:
"""Precompute the absolute-address table for the Triton zeroing kernel.
@@ -109,10 +110,14 @@ class KVBlockZeroer:
"""
self.device = device
self.pin_memory = pin_memory
if max_concurrency < 1:
raise ValueError("max_concurrency must be at least 1")
self.max_concurrency = max_concurrency
self._meta: tuple[torch.Tensor, int, int, int] | None = None
self._id_cap: int = 0
self._ids_pinned: torch.Tensor | None = None
self._ids_gpu: torch.Tensor | None = None
self._ids_pinned: list[torch.Tensor] = []
self._ids_gpu: list[torch.Tensor] = []
self._id_buffer_index = 0
if runner_only_attn_layers is None:
runner_only_attn_layers = set()
@@ -175,12 +180,7 @@ class KVBlockZeroer:
blk_size = min(largest_power_of_2_divisor(page_size_el), 1024)
self._id_cap = 8192
self._ids_pinned = torch.empty(
self._id_cap,
dtype=torch.int64,
pin_memory=self.pin_memory,
)
self._ids_gpu = torch.empty(self._id_cap, dtype=torch.int64, device=self.device)
self._allocate_id_buffers()
self._meta = (
torch.tensor(seg_addrs, dtype=torch.uint64, device=self.device),
page_size_el,
@@ -188,6 +188,21 @@ class KVBlockZeroer:
len(seg_addrs),
)
def _allocate_id_buffers(self) -> None:
self._ids_pinned = [
torch.empty(
self._id_cap,
dtype=torch.int64,
pin_memory=self.pin_memory,
)
for _ in range(self.max_concurrency)
]
self._ids_gpu = [
torch.empty(self._id_cap, dtype=torch.int64, device=self.device)
for _ in range(self.max_concurrency)
]
self._id_buffer_index = 0
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:
@@ -195,19 +210,21 @@ class KVBlockZeroer:
seg_addrs, page_size_el, blk_size, n_segs = self._meta
n_blocks = len(block_ids)
if n_blocks > self._id_cap:
# The old pinned buffers may still be the source of an in-flight
# nonblocking copy. Growing is rare, so we don't mind the sync overhead
torch.accelerator.synchronize()
self._id_cap = n_blocks * 2
self._ids_pinned = torch.empty(
self._id_cap,
dtype=torch.int64,
pin_memory=self.pin_memory,
)
self._ids_gpu = torch.empty(
self._id_cap, dtype=torch.int64, device=self.device
)
assert self._ids_pinned is not None and self._ids_gpu is not None
self._ids_pinned[:n_blocks].numpy()[:] = block_ids
idx = self._ids_gpu[:n_blocks]
idx.copy_(self._ids_pinned[:n_blocks], non_blocking=True)
self._allocate_id_buffers()
# The H2D copy is nonblocking, so its pinned source must not be mutated
# while this batch is in flight. Rotate through as many buffers as concurrent
# in-flight batches, to avoid collisions.
buffer_index = self._id_buffer_index
self._id_buffer_index = (buffer_index + 1) % self.max_concurrency
ids_pinned = self._ids_pinned[buffer_index]
ids_pinned[:n_blocks].numpy()[:] = block_ids
idx = self._ids_gpu[buffer_index][:n_blocks]
idx.copy_(ids_pinned[:n_blocks], non_blocking=True)
grid = (n_blocks * n_segs * (page_size_el // blk_size),)
_zero_kv_blocks_kernel[grid](
seg_addrs,