diff --git a/python/sglang/kernels/jit/csrc/elementwise/kvcache.cuh b/python/sglang/kernels/jit/csrc/elementwise/kvcache.cuh index 50133517d..6be2f04c2 100644 --- a/python/sglang/kernels/jit/csrc/elementwise/kvcache.cuh +++ b/python/sglang/kernels/jit/csrc/elementwise/kvcache.cuh @@ -25,6 +25,7 @@ struct StoreKVCacheParams { int64_t stride_indices; uint32_t batch_size; int64_t size_limit; + int64_t reserved_skip_index; }; constexpr uint32_t kNumWarps = 4; @@ -97,7 +98,7 @@ __global__ void store_kvcache(const __grid_constant__ StoreKVCacheParams params) const auto& [ k_input, v_input, k_cache, v_cache, indices, // ptr stride_k, stride_v, stride_cache, stride_indices, batch_size, // size - size_limit // bound + size_limit, reserved_skip_index // bounds and reserved sink ] = params; if (item_id >= batch_size) return; @@ -113,7 +114,9 @@ __global__ void store_kvcache(const __grid_constant__ StoreKVCacheParams params) const auto k_dst = pointer::offset(k_cache, index * stride_cache, split_id * kSplitSize); const auto v_dst = pointer::offset(v_cache, index * stride_cache, split_id * kSplitSize); - copy_kv_warp(k_src, v_src, k_dst, v_dst); + if (index != reserved_skip_index) { + copy_kv_warp(k_src, v_src, k_dst, v_dst); + } PDLTriggerSecondary(); } @@ -145,7 +148,8 @@ struct StoreKVCacheKernel { const tvm::ffi::TensorView v_cache, const tvm::ffi::TensorView indices, const int num_split, - const int64_t size_limit) { + const int64_t size_limit, + const int64_t reserved_skip_index) { using namespace host; auto B = SymbolicSize{"batch_size"}; auto D = SymbolicSize{"element_size"}; @@ -196,6 +200,7 @@ struct StoreKVCacheKernel { .stride_indices = I.unwrap(), .batch_size = static_cast(B.unwrap()), .size_limit = size_limit, + .reserved_skip_index = reserved_skip_index, }; // select kernel and update num_split if needed const auto use_int32 = indice_dtype.is_type(); diff --git a/python/sglang/kernels/ops/kvcache/kvcache.py b/python/sglang/kernels/ops/kvcache/kvcache.py index 8529a72c3..f4534d5e3 100644 --- a/python/sglang/kernels/ops/kvcache/kvcache.py +++ b/python/sglang/kernels/ops/kvcache/kvcache.py @@ -58,6 +58,7 @@ def store_cache( row_bytes: int = 0, num_split: int = 0, # can be tuned for performance size_limit: int = 0, + reserved_skip_index: int = 0, ) -> None: """Store key and value tensors into KV cache at specified indices. @@ -71,6 +72,9 @@ def store_cache( reserved padding slot); an index outside [0, size_limit) fails fast (device assert) instead of an illegal memory access. Defaults to the cache row count when 0. + reserved_skip_index (int): If nonnegative, writes targeting this index + are skipped. Defaults to the reserved CUDA-graph padding slot 0; + pass -1 to disable skipping. """ row_bytes = row_bytes or k.shape[-1] * k.element_size() module = _jit_kvcache_module(row_bytes) @@ -91,4 +95,5 @@ def store_cache( indices, num_split, size_limit, + reserved_skip_index, ) diff --git a/test/registered/kernels/ops/kvcache/test_store_cache.py b/test/registered/kernels/ops/kvcache/test_store_cache.py index a114797d1..957185ae4 100644 --- a/test/registered/kernels/ops/kvcache/test_store_cache.py +++ b/test/registered/kernels/ops/kvcache/test_store_cache.py @@ -33,7 +33,7 @@ def test_store_cache(batch_size: int, element_dim: int) -> None: v = torch.randn((batch_size, element_dim), dtype=DTYPE, device=DEVICE) k_cache = torch.randn((CACHE_SIZE, element_dim), dtype=DTYPE, device=DEVICE) v_cache = torch.randn((CACHE_SIZE, element_dim), dtype=DTYPE, device=DEVICE) - indices = torch.randperm(CACHE_SIZE, device=DEVICE)[:batch_size] + indices = torch.randperm(CACHE_SIZE - 1, device=DEVICE)[:batch_size] + 1 # AOT store cache store_cache(k, v, k_cache, v_cache, indices) @@ -60,7 +60,7 @@ def test_store_cache_dtypes( v = torch.randn((batch_size, element_dim), dtype=dtype, device=DEVICE) k_cache = torch.randn((SMALL_CACHE, element_dim), dtype=dtype, device=DEVICE) v_cache = torch.randn((SMALL_CACHE, element_dim), dtype=dtype, device=DEVICE) - indices = torch.randperm(SMALL_CACHE, device=DEVICE)[:batch_size] + indices = torch.randperm(SMALL_CACHE - 1, device=DEVICE)[:batch_size] + 1 store_cache(k, v, k_cache, v_cache, indices) @@ -78,7 +78,9 @@ def test_store_cache_int32_indices(batch_size: int, element_dim: int) -> None: k_cache = torch.randn((SMALL_CACHE, element_dim), dtype=DTYPE, device=DEVICE) v_cache = torch.randn((SMALL_CACHE, element_dim), dtype=DTYPE, device=DEVICE) # int32 indices exercise a different CUDA template instantiation than default int64 - indices = torch.randperm(SMALL_CACHE, device=DEVICE)[:batch_size].to(torch.int32) + indices = (torch.randperm(SMALL_CACHE - 1, device=DEVICE)[:batch_size] + 1).to( + torch.int32 + ) store_cache(k, v, k_cache, v_cache, indices) @@ -86,6 +88,55 @@ def test_store_cache_int32_indices(batch_size: int, element_dim: int) -> None: assert torch.all(v_cache[indices.long()] == v) +@pytest.mark.parametrize("index_dtype", [torch.int32, torch.int64]) +@pytest.mark.parametrize("num_split", [1, 2, 4]) +def test_store_cache_reserved_skip_index( + index_dtype: torch.dtype, num_split: int +) -> None: + element_dim = 1024 + k = torch.randn((4, element_dim), dtype=DTYPE, device=DEVICE) + v = torch.randn((4, element_dim), dtype=DTYPE, device=DEVICE) + # Model kernels may leave CUDA-graph padding rows undefined. Reproduce the + # dangerous case directly instead of requiring a full model checkpoint. + k[[0, 2]] = torch.nan + v[[0, 2]] = torch.nan + k_cache = torch.randn((SMALL_CACHE, element_dim), dtype=DTYPE, device=DEVICE) + v_cache = torch.randn((SMALL_CACHE, element_dim), dtype=DTYPE, device=DEVICE) + reserved_k_before = k_cache[0].clone() + reserved_v_before = v_cache[0].clone() + indices = torch.tensor([0, 7, 0, 9], dtype=index_dtype, device=DEVICE) + + store_cache( + k, + v, + k_cache, + v_cache, + indices, + num_split=num_split, + ) + + torch.testing.assert_close(k_cache[0], reserved_k_before, rtol=0.0, atol=0.0) + torch.testing.assert_close(v_cache[0], reserved_v_before, rtol=0.0, atol=0.0) + torch.testing.assert_close(k_cache[indices[1].long()], k[1], rtol=0.0, atol=0.0) + torch.testing.assert_close(v_cache[indices[1].long()], v[1], rtol=0.0, atol=0.0) + torch.testing.assert_close(k_cache[indices[3].long()], k[3], rtol=0.0, atol=0.0) + torch.testing.assert_close(v_cache[indices[3].long()], v[3], rtol=0.0, atol=0.0) + + +def test_store_cache_zero_index_can_be_written_when_skip_disabled() -> None: + element_dim = 64 + k = torch.randn((1, element_dim), dtype=DTYPE, device=DEVICE) + v = torch.randn((1, element_dim), dtype=DTYPE, device=DEVICE) + k_cache = torch.randn((SMALL_CACHE, element_dim), dtype=DTYPE, device=DEVICE) + v_cache = torch.randn((SMALL_CACHE, element_dim), dtype=DTYPE, device=DEVICE) + indices = torch.zeros(1, dtype=torch.int64, device=DEVICE) + + store_cache(k, v, k_cache, v_cache, indices, reserved_skip_index=-1) + + torch.testing.assert_close(k_cache[0], k[0], rtol=0.0, atol=0.0) + torch.testing.assert_close(v_cache[0], v[0], rtol=0.0, atol=0.0) + + def _valid_num_splits(element_dim: int, dtype: torch.dtype) -> list: """Return the list of valid num_split values for a given element_dim/dtype.""" row_bytes = element_dim * dtype.itemsize @@ -114,7 +165,7 @@ def test_store_cache_num_split( v = torch.randn((batch_size, element_dim), dtype=dtype, device=DEVICE) k_cache = torch.randn((SMALL_CACHE, element_dim), dtype=dtype, device=DEVICE) v_cache = torch.randn((SMALL_CACHE, element_dim), dtype=dtype, device=DEVICE) - indices = torch.randperm(SMALL_CACHE, device=DEVICE)[:batch_size] + indices = torch.randperm(SMALL_CACHE - 1, device=DEVICE)[:batch_size] + 1 # Verify each num_split kernel path (1, 2, 4) produces correct results store_cache(k, v, k_cache, v_cache, indices, num_split=num_split)