[spec] Misc defensive guards for EAGLE draft KV indexing (#27486)

This commit is contained in:
Liangsheng Yin
2026-06-07 21:27:46 -07:00
committed by GitHub
parent 6c2770149b
commit 6365d6faee
6 changed files with 36 additions and 2 deletions
@@ -8,6 +8,7 @@
#include <dlpack/dlpack.h>
#include <tvm/ffi/container/tensor.h>
#include <cassert>
#include <cstdint>
namespace {
@@ -23,6 +24,7 @@ struct StoreKVCacheParams {
int64_t stride_cache_bytes;
int64_t stride_indices;
uint32_t batch_size;
int64_t size_limit;
};
constexpr uint32_t kNumWarps = 4;
@@ -94,7 +96,8 @@ __global__ void store_kvcache(const __grid_constant__ StoreKVCacheParams params)
const uint32_t split_id = warp_id % kSplit;
const auto& [
k_input, v_input, k_cache, v_cache, indices, // ptr
stride_k, stride_v, stride_cache, stride_indices, batch_size // size
stride_k, stride_v, stride_cache, stride_indices, batch_size, // size
size_limit // bound
] = params;
if (item_id >= batch_size) return;
@@ -102,6 +105,9 @@ __global__ void store_kvcache(const __grid_constant__ StoreKVCacheParams params)
PDLWaitPrimary<kUsePDL>();
const auto index = *index_ptr;
// A stale/OOB slot id would cause an illegal memory access in the store below;
// fail fast at the culprit instead. always-on (kvcache JIT compiles without NDEBUG).
assert(index >= 0 && index < size_limit);
const auto k_src = pointer::offset(k_input, item_id * stride_k, split_id * kSplitSize);
const auto v_src = pointer::offset(v_input, item_id * stride_v, split_id * kSplitSize);
const auto k_dst = pointer::offset(k_cache, index * stride_cache, split_id * kSplitSize);
@@ -138,7 +144,8 @@ struct StoreKVCacheKernel {
const tvm::ffi::TensorView k_cache,
const tvm::ffi::TensorView v_cache,
const tvm::ffi::TensorView indices,
const int num_split) {
const int num_split,
const int64_t size_limit) {
using namespace host;
auto B = SymbolicSize{"batch_size"};
auto D = SymbolicSize{"element_size"};
@@ -188,6 +195,7 @@ struct StoreKVCacheKernel {
.stride_cache_bytes = S.unwrap() * dtype_size,
.stride_indices = I.unwrap(),
.batch_size = static_cast<uint32_t>(B.unwrap()),
.size_limit = size_limit,
};
// select kernel and update num_split if needed
const auto use_int32 = indice_dtype.is_type<int32_t>();
+8
View File
@@ -57,6 +57,7 @@ def store_cache(
*,
row_bytes: int = 0,
num_split: int = 0, # can be tuned for performance
size_limit: int = 0,
) -> None:
"""Store key and value tensors into KV cache at specified indices.
@@ -66,6 +67,10 @@ def store_cache(
k_cache (torch.Tensor): Key cache tensor of shape (num_pages, H * D).
v_cache (torch.Tensor): Value cache tensor of shape (num_pages, H * D).
indices (torch.Tensor): Indices tensor of shape (batch_size,).
size_limit (int): Valid slot bound (cache row count = real slots + the
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.
"""
row_bytes = row_bytes or k.shape[-1] * k.element_size()
module = _jit_kvcache_module(row_bytes)
@@ -76,6 +81,8 @@ def store_cache(
num_split = 2
else:
num_split = 1
if size_limit <= 0:
size_limit = k_cache.shape[0]
module.store_cache(
k,
v,
@@ -83,4 +90,5 @@ def store_cache(
v_cache,
indices,
num_split,
size_limit,
)
@@ -103,6 +103,7 @@ def _set_kv_buffer_impl(
row_dim: int, # head_num * head_dim
store_dtype: torch.dtype,
device_module: Any,
size_limit: int,
alt_stream: Optional[torch.cuda.Stream] = None,
same_kv_dim: bool = True,
) -> None:
@@ -115,6 +116,7 @@ def _set_kv_buffer_impl(
v_cache.view(-1, row_dim),
indices,
row_bytes=row_bytes,
size_limit=size_limit,
)
if _is_cpu and _cpu_has_amx_support:
@@ -1254,6 +1256,9 @@ class MHATokenToKVPool(KVCache):
row_dim=self.row_dim,
store_dtype=self.store_dtype,
device_module=self.device_module,
# size + page_size = real slots + the reserved padding slot (padded /
# dummy tokens write there); valid index range is [0, size + page_size).
size_limit=self.size + self.page_size,
alt_stream=self.alt_stream,
same_kv_dim=self.same_kv_dim,
)
@@ -89,6 +89,11 @@ class MultiLayerEagleWorker(TpModelWorker):
self.topk = server_args.speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens
assert self.speculative_num_draft_tokens == self.speculative_num_steps + 1, (
"multi-layer EAGLE requires speculative_num_draft_tokens == "
"speculative_num_steps + 1, "
f"got {self.speculative_num_draft_tokens} and {self.speculative_num_steps}"
)
self.gpu_id = gpu_id
self.device = server_args.device
self.target_worker = target_worker
@@ -106,6 +106,11 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker):
self.topk = server_args.speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens
assert self.speculative_num_draft_tokens == self.speculative_num_steps + 1, (
"multi-layer EAGLE requires speculative_num_draft_tokens == "
"speculative_num_steps + 1, "
f"got {self.speculative_num_draft_tokens} and {self.speculative_num_steps}"
)
self.speculative_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm
)
@@ -86,6 +86,9 @@ def draft_kv_indices_buffer_width(
num_seqs * topk branches each attend up to max_context_len KV slots; the topk
factor is mandatory -- dropping it under-allocates and overflows the row (#27338, #27460).
"""
assert (
num_seqs * topk * max_context_len < 2**31
), "kv_indices flat offset would overflow int32; reduce batch/topk/context"
return num_seqs * topk * max_context_len