[core] Probe set_kv_buffer / set_mla_kv_buffer slot ids for OOB (#27459)
This commit is contained in:
@@ -1126,6 +1126,9 @@ class MHATokenToKVPool(KVCache):
|
|||||||
v_scale: Optional[float] = None,
|
v_scale: Optional[float] = None,
|
||||||
layer_id_override: Optional[int] = None,
|
layer_id_override: Optional[int] = None,
|
||||||
):
|
):
|
||||||
|
# Catch stale slot ids here instead of as illegal-addr / silent KV
|
||||||
|
# corruption in the store_kvcache write (gated on SGLANG_ENABLE_ASYNC_ASSERT).
|
||||||
|
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_kv_buffer (MHA)")
|
||||||
if layer_id_override is not None:
|
if layer_id_override is not None:
|
||||||
layer_id = layer_id_override
|
layer_id = layer_id_override
|
||||||
else:
|
else:
|
||||||
@@ -1422,6 +1425,7 @@ class MHATokenToKVPoolFP4(MHATokenToKVPool):
|
|||||||
v_scale: Optional[float] = None,
|
v_scale: Optional[float] = None,
|
||||||
layer_id_override: Optional[int] = None,
|
layer_id_override: Optional[int] = None,
|
||||||
):
|
):
|
||||||
|
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_kv_buffer (MHA-FP4)")
|
||||||
from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode
|
from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode
|
||||||
|
|
||||||
if layer_id_override is not None:
|
if layer_id_override is not None:
|
||||||
@@ -1812,6 +1816,7 @@ class MLATokenToKVPool(KVCache):
|
|||||||
cache_k: torch.Tensor,
|
cache_k: torch.Tensor,
|
||||||
cache_v: torch.Tensor,
|
cache_v: torch.Tensor,
|
||||||
):
|
):
|
||||||
|
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_kv_buffer (MLA)")
|
||||||
layer_id = layer.layer_id
|
layer_id = layer.layer_id
|
||||||
assert not self.dsa_kv_cache_store_fp8
|
assert not self.dsa_kv_cache_store_fp8
|
||||||
if cache_k.dtype != self.dtype:
|
if cache_k.dtype != self.dtype:
|
||||||
@@ -1831,6 +1836,7 @@ class MLATokenToKVPool(KVCache):
|
|||||||
cache_k_nope: torch.Tensor,
|
cache_k_nope: torch.Tensor,
|
||||||
cache_k_rope: torch.Tensor,
|
cache_k_rope: torch.Tensor,
|
||||||
):
|
):
|
||||||
|
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_mla_kv_buffer (MLA)")
|
||||||
layer_id = layer.layer_id
|
layer_id = layer.layer_id
|
||||||
|
|
||||||
if _is_hip and self.use_dsa and self.dtype == fp8_dtype:
|
if _is_hip and self.use_dsa and self.dtype == fp8_dtype:
|
||||||
@@ -2004,6 +2010,7 @@ class MLATokenToKVPoolFP4(MLATokenToKVPool):
|
|||||||
cache_k: torch.Tensor,
|
cache_k: torch.Tensor,
|
||||||
cache_v: torch.Tensor,
|
cache_v: torch.Tensor,
|
||||||
):
|
):
|
||||||
|
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_kv_buffer (MLA-FP4)")
|
||||||
layer_id = layer.layer_id
|
layer_id = layer.layer_id
|
||||||
assert not self.dsa_kv_cache_store_fp8
|
assert not self.dsa_kv_cache_store_fp8
|
||||||
if cache_k.dtype != self.dtype:
|
if cache_k.dtype != self.dtype:
|
||||||
@@ -2028,6 +2035,9 @@ class MLATokenToKVPoolFP4(MLATokenToKVPool):
|
|||||||
cache_k_nope: torch.Tensor,
|
cache_k_nope: torch.Tensor,
|
||||||
cache_k_rope: torch.Tensor,
|
cache_k_rope: torch.Tensor,
|
||||||
):
|
):
|
||||||
|
maybe_detect_oob(
|
||||||
|
loc, 0, self.size + self.page_size, "set_mla_kv_buffer (MLA-FP4)"
|
||||||
|
)
|
||||||
layer_id = layer.layer_id
|
layer_id = layer.layer_id
|
||||||
|
|
||||||
if self.dsa_kv_cache_store_fp8:
|
if self.dsa_kv_cache_store_fp8:
|
||||||
|
|||||||
Reference in New Issue
Block a user