[Refactor] Encapsulate SWA loc translation inside SWAKVPool with per-batch cache invalidation (#25824)

Co-authored-by: Claude Sonnet 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
Cheng Wan
2026-05-20 21:26:32 -07:00
committed by GitHub
co-authored by Claude Sonnet 4.6
parent 90efa9c83f
commit 79b937aefb
18 changed files with 298 additions and 174 deletions
@@ -1057,9 +1057,6 @@ class TritonAttnBackend(AttentionBackend):
prefix_kv_indices = self.forward_metadata.kv_indices prefix_kv_indices = self.forward_metadata.kv_indices
window_start_pos = None window_start_pos = None
# For SWA layers, mirror SWAKVPool.set_kv_buffer: read from the
# precomputed pool.swa_loc. Translate out_cache_loc to SWA-pool index space
# as a fallback when pool.swa_loc is not pre-populated.
extend_kv_indices = forward_batch.out_cache_loc extend_kv_indices = forward_batch.out_cache_loc
pool = forward_batch.token_to_kv_pool pool = forward_batch.token_to_kv_pool
if ( if (
@@ -1068,12 +1065,7 @@ class TritonAttnBackend(AttentionBackend):
and isinstance(pool, SWAKVPool) and isinstance(pool, SWAKVPool)
and pool.layers_mapping[layer.layer_id][1] and pool.layers_mapping[layer.layer_id][1]
): ):
if pool.swa_loc is not None: extend_kv_indices = pool.translate_loc_from_full_to_swa(extend_kv_indices)
extend_kv_indices = pool.swa_loc
else:
extend_kv_indices = pool.translate_loc_from_full_to_swa(
extend_kv_indices
)
# Handle cases where extend_seq_lens or extend_start_loc might not be set # Handle cases where extend_seq_lens or extend_start_loc might not be set
# In speculative decoding, we can infer these from spec_info or compute them # In speculative decoding, we can infer these from spec_info or compute them
@@ -187,8 +187,6 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
if self.use_sliding_window_kv_pool: if self.use_sliding_window_kv_pool:
_, is_swa = self._swa_kv_pool.layers_mapping[layer.layer_id] _, is_swa = self._swa_kv_pool.layers_mapping[layer.layer_id]
if is_swa: if is_swa:
if forward_batch.out_cache_loc_swa is not None:
return forward_batch.out_cache_loc_swa
return self._swa_kv_pool.translate_loc_from_full_to_swa( return self._swa_kv_pool.translate_loc_from_full_to_swa(
forward_batch.out_cache_loc forward_batch.out_cache_loc
) )
@@ -179,16 +179,9 @@ def unified_attention_with_output(
kwargs["sinks"] = sinks kwargs["sinks"] = sinks
original_out_cache_loc = forward_batch.out_cache_loc original_out_cache_loc = forward_batch.out_cache_loc
original_out_cache_loc_swa = forward_batch.out_cache_loc_swa
token_to_kv_pool = forward_batch.token_to_kv_pool
original_swa_loc = getattr(token_to_kv_pool, "swa_loc", None)
# Keep the original ForwardBatch object and only narrow cache locations for # Keep the original ForwardBatch object and only narrow cache locations for
# this backend call so model/backend state is still written to the same batch. # this backend call so model/backend state is still written to the same batch.
forward_batch.out_cache_loc = original_out_cache_loc[:real_num_tokens] forward_batch.out_cache_loc = original_out_cache_loc[:real_num_tokens]
if original_out_cache_loc_swa is not None:
forward_batch.out_cache_loc_swa = original_out_cache_loc_swa[:real_num_tokens]
if hasattr(token_to_kv_pool, "set_swa_loc"):
token_to_kv_pool.set_swa_loc(forward_batch.out_cache_loc_swa)
# Store pre-allocated output for FA backend to write directly into. # Store pre-allocated output for FA backend to write directly into.
# Must slice to real_num_tokens to match the narrowed query shape — # Must slice to real_num_tokens to match the narrowed query shape —
@@ -205,11 +198,6 @@ def unified_attention_with_output(
**kwargs, **kwargs,
) )
forward_batch.out_cache_loc = original_out_cache_loc forward_batch.out_cache_loc = original_out_cache_loc
forward_batch.out_cache_loc_swa = original_out_cache_loc_swa
if original_out_cache_loc_swa is not None and hasattr(
token_to_kv_pool, "set_swa_loc"
):
token_to_kv_pool.set_swa_loc(original_swa_loc)
if ret.data_ptr() != output.data_ptr(): if ret.data_ptr() != output.data_ptr():
output[:real_num_tokens].view(ret.shape).copy_(ret) output[:real_num_tokens].view(ret.shape).copy_(ret)
@@ -120,16 +120,9 @@ def unified_linear_attention_with_output(
real_num_tokens = forward_batch.num_token_non_padded_cpu real_num_tokens = forward_batch.num_token_non_padded_cpu
original_out_cache_loc = forward_batch.out_cache_loc original_out_cache_loc = forward_batch.out_cache_loc
original_out_cache_loc_swa = forward_batch.out_cache_loc_swa
token_to_kv_pool = forward_batch.token_to_kv_pool
original_swa_loc = getattr(token_to_kv_pool, "swa_loc", None)
# Keep the original ForwardBatch object and only narrow cache locations for # Keep the original ForwardBatch object and only narrow cache locations for
# this backend call so model/backend state is still written to the same batch. # this backend call so model/backend state is still written to the same batch.
forward_batch.out_cache_loc = original_out_cache_loc[:real_num_tokens] forward_batch.out_cache_loc = original_out_cache_loc[:real_num_tokens]
if original_out_cache_loc_swa is not None:
forward_batch.out_cache_loc_swa = original_out_cache_loc_swa[:real_num_tokens]
if hasattr(token_to_kv_pool, "set_swa_loc"):
token_to_kv_pool.set_swa_loc(forward_batch.out_cache_loc_swa)
ret = forward_batch.attn_backend.forward( ret = forward_batch.attn_backend.forward(
layer=attention_layer, layer=attention_layer,
@@ -139,11 +132,6 @@ def unified_linear_attention_with_output(
b=b[:real_num_tokens], b=b[:real_num_tokens],
) )
forward_batch.out_cache_loc = original_out_cache_loc forward_batch.out_cache_loc = original_out_cache_loc
forward_batch.out_cache_loc_swa = original_out_cache_loc_swa
if original_out_cache_loc_swa is not None and hasattr(
token_to_kv_pool, "set_swa_loc"
):
token_to_kv_pool.set_swa_loc(original_swa_loc)
output[:, :real_num_tokens].copy_(ret) output[:, :real_num_tokens].copy_(ret)
return return
@@ -16,6 +16,9 @@ class BaseSWAKVPool(KVCache):
swa_kv_pool: KVCache swa_kv_pool: KVCache
def invalidate_loc_cache(self) -> None:
pass
@abc.abstractmethod @abc.abstractmethod
def register_mapping(self, full_to_swa_index_mapping: torch.Tensor) -> None: def register_mapping(self, full_to_swa_index_mapping: torch.Tensor) -> None:
raise NotImplementedError() raise NotImplementedError()
@@ -24,10 +27,6 @@ class BaseSWAKVPool(KVCache):
def translate_loc_from_full_to_swa(self, kv_indices: torch.Tensor) -> torch.Tensor: def translate_loc_from_full_to_swa(self, kv_indices: torch.Tensor) -> torch.Tensor:
raise NotImplementedError() raise NotImplementedError()
@abc.abstractmethod
def set_swa_loc(self, loc: torch.Tensor) -> None:
raise NotImplementedError()
@abc.abstractmethod @abc.abstractmethod
def get_state_buf_infos(self) -> Tuple[List[int], List[int], List[int]]: def get_state_buf_infos(self) -> Tuple[List[int], List[int], List[int]]:
raise NotImplementedError() raise NotImplementedError()
@@ -503,13 +503,6 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
return self.full_to_swa_index_mapping[kv_indices].to(torch.int32) return self.full_to_swa_index_mapping[kv_indices].to(torch.int32)
def set_swa_loc(self, loc: torch.Tensor) -> None:
# No-op: SWAKVPool's set_swa_loc precomputes SWA-translated loc once per
# forward batch for set_kv_buffer to read via self.swa_loc. DSV4 has its
# own equivalent cache via `_should_cache_swa + cached_loc` (in
# set_swa_key_buffer_radix_fused), so we ignore main's precomputed loc.
pass
def get_contiguous_buf_infos(self) -> Tuple[List[int], List[int], List[int]]: def get_contiguous_buf_infos(self) -> Tuple[List[int], List[int], List[int]]:
data_ptrs: List[int] = [] data_ptrs: List[int] = []
data_lens: List[int] = [] data_lens: List[int] = []
+31 -15
View File
@@ -55,7 +55,6 @@ class SWAKVPool(BaseSWAKVPool):
self.layer_num = self.full_layer_nums + self.swa_layer_nums self.layer_num = self.full_layer_nums + self.swa_layer_nums
self.start_layer = 0 self.start_layer = 0
self.page_size = page_size self.page_size = page_size
self.swa_loc = None
self.layer_transfer_counter = None self.layer_transfer_counter = None
kwargs["page_size"] = page_size kwargs["page_size"] = page_size
@@ -93,6 +92,8 @@ class SWAKVPool(BaseSWAKVPool):
for swa_layer_id, global_layer_id in enumerate(swa_attention_layer_ids): for swa_layer_id, global_layer_id in enumerate(swa_attention_layer_ids):
self.layers_mapping[global_layer_id] = (swa_layer_id, True) self.layers_mapping[global_layer_id] = (swa_layer_id, True)
self.full_to_swa_index_mapping: Optional[torch.Tensor] = None self.full_to_swa_index_mapping: Optional[torch.Tensor] = None
self._cached_swa_loc: Optional[torch.Tensor] = None
self._cached_loc_key: Optional[tuple] = None
k_size, v_size = self.get_kv_size_bytes() k_size, v_size = self.get_kv_size_bytes()
self.mem_usage = (k_size + v_size) / GB self.mem_usage = (k_size + v_size) / GB
@@ -102,6 +103,11 @@ class SWAKVPool(BaseSWAKVPool):
def register_mapping(self, full_to_swa_index_mapping: torch.Tensor): def register_mapping(self, full_to_swa_index_mapping: torch.Tensor):
self.full_to_swa_index_mapping = full_to_swa_index_mapping self.full_to_swa_index_mapping = full_to_swa_index_mapping
self.invalidate_loc_cache()
def invalidate_loc_cache(self) -> None:
self._cached_swa_loc = None
self._cached_loc_key = None
def register_layer_transfer_counter(self, layer_transfer_counter): def register_layer_transfer_counter(self, layer_transfer_counter):
# Wait happens at this wrapper. Inner pools must not wait again. # Wait happens at this wrapper. Inner pools must not wait again.
@@ -159,15 +165,23 @@ class SWAKVPool(BaseSWAKVPool):
else: else:
return self.full_kv_pool.get_kv_buffer(layer_id_pool) return self.full_kv_pool.get_kv_buffer(layer_id_pool)
def set_swa_loc(self, loc: torch.Tensor): def translate_loc_from_full_to_swa(self, kv_indices: torch.Tensor) -> torch.Tensor:
self.swa_loc = loc
def translate_loc_from_full_to_swa(self, kv_indices: torch.Tensor):
assert self.full_to_swa_index_mapping is not None assert self.full_to_swa_index_mapping is not None
# data_ptr() (not untyped_storage().data_ptr()) encodes the offset, so
# Note: kv_indices could have -1 values (from alloc_extend), which will be mapped to -1 # views at different positions within the same storage get distinct keys.
# since the last item of full_to_swa_index_mapping is -1. # -1 in kv_indices maps to -1 via the sentinel appended to the mapping.
return self.full_to_swa_index_mapping[kv_indices].to(torch.int32) key = (kv_indices.data_ptr(), kv_indices.numel())
if key != self._cached_loc_key:
if self._cached_loc_key is not None:
logger.warning(
"translate_loc_from_full_to_swa: loc tensor changed mid-forward "
"without invalidate_loc_cache() — possible missing call site"
)
self._cached_swa_loc = self.full_to_swa_index_mapping[kv_indices].to(
torch.int32
)
self._cached_loc_key = key
return self._cached_swa_loc
def set_kv_buffer( def set_kv_buffer(
self, self,
@@ -182,12 +196,7 @@ class SWAKVPool(BaseSWAKVPool):
layer_id = layer.layer_id layer_id = layer.layer_id
layer_id_pool, is_swa_layer = self.layers_mapping[layer_id] layer_id_pool, is_swa_layer = self.layers_mapping[layer_id]
if is_swa_layer: if is_swa_layer:
if self.swa_loc is not None:
loc = self.swa_loc
else:
if self.full_to_swa_index_mapping is not None:
loc = self.translate_loc_from_full_to_swa(loc) loc = self.translate_loc_from_full_to_swa(loc)
self.swa_kv_pool.set_kv_buffer( self.swa_kv_pool.set_kv_buffer(
None, None,
loc, loc,
@@ -372,8 +381,8 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
self.is_not_in_free_group = True self.is_not_in_free_group = True
self.free_group = [] self.free_group = []
self.clear()
self._kvcache = kvcache self._kvcache = kvcache
self.clear()
self._kvcache.register_mapping(self.full_to_swa_index_mapping) self._kvcache.register_mapping(self.full_to_swa_index_mapping)
def available_size(self): def available_size(self):
@@ -416,6 +425,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
return self._kvcache.translate_loc_from_full_to_swa(kv_indices) return self._kvcache.translate_loc_from_full_to_swa(kv_indices)
def alloc(self, need_size: int): def alloc(self, need_size: int):
self._kvcache.invalidate_loc_cache()
assert self.page_size == 1 assert self.page_size == 1
if need_size > self.full_attn_allocator.available_size(): if need_size > self.full_attn_allocator.available_size():
return None return None
@@ -444,6 +454,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
last_loc: torch.Tensor, # last_loc for full layers last_loc: torch.Tensor, # last_loc for full layers
extend_num_tokens: int, extend_num_tokens: int,
): ):
self._kvcache.invalidate_loc_cache()
assert self.page_size > 1 assert self.page_size > 1
num_new_pages = get_num_new_pages( num_new_pages = get_num_new_pages(
@@ -496,6 +507,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
extend_num_tokens: int, extend_num_tokens: int,
swa_tail_len: int, swa_tail_len: int,
): ):
self._kvcache.invalidate_loc_cache()
"""Allocate full KV for the whole extend and SWA KV only for the tail. """Allocate full KV for the whole extend and SWA KV only for the tail.
This is used by disaggregated decode preallocation: decode receives full This is used by disaggregated decode preallocation: decode receives full
@@ -557,6 +569,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
seq_lens_cpu: torch.Tensor, seq_lens_cpu: torch.Tensor,
last_loc: torch.Tensor, # last_loc for full layers last_loc: torch.Tensor, # last_loc for full layers
): ):
self._kvcache.invalidate_loc_cache()
assert self.page_size > 1 assert self.page_size > 1
swa_last_loc = self.translate_loc_from_full_to_swa(last_loc) swa_last_loc = self.translate_loc_from_full_to_swa(last_loc)
@@ -604,6 +617,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
if full_indices.numel() == 0: if full_indices.numel() == 0:
return return
assert full_indices.numel() == swa_indices.numel() assert full_indices.numel() == swa_indices.numel()
self._kvcache.invalidate_loc_cache()
if _is_npu: if _is_npu:
self.full_to_swa_index_mapping[full_indices.to(torch.int64)] = ( self.full_to_swa_index_mapping[full_indices.to(torch.int64)] = (
swa_indices.to(torch.int64) swa_indices.to(torch.int64)
@@ -612,6 +626,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
self.full_to_swa_index_mapping[full_indices] = swa_indices self.full_to_swa_index_mapping[full_indices] = swa_indices
def free_swa(self, free_index: torch.Tensor): def free_swa(self, free_index: torch.Tensor):
self._kvcache.invalidate_loc_cache()
swa_indices = self.full_to_swa_index_mapping[free_index] swa_indices = self.full_to_swa_index_mapping[free_index]
swa_indices = swa_indices[swa_indices > 0] swa_indices = swa_indices[swa_indices > 0]
self.swa_attn_allocator.free(swa_indices) self.swa_attn_allocator.free(swa_indices)
@@ -629,6 +644,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
self.swa_attn_allocator.restore_state(state[1]) self.swa_attn_allocator.restore_state(state[1])
def clear(self): def clear(self):
self._kvcache.invalidate_loc_cache()
self.swa_attn_allocator.clear() self.swa_attn_allocator.clear()
self.full_attn_allocator.clear() self.full_attn_allocator.clear()
# Note: the last item is -1, we don't clear it, see the comment in __init__ # Note: the last item is -1, we don't clear it, see the comment in __init__
@@ -179,11 +179,6 @@ class BreakableCudaGraphRunner:
(self.max_num_tokens,), (self.max_num_tokens,),
dtype=torch.int64 if not is_npu() else torch.int32, dtype=torch.int64 if not is_npu() else torch.int32,
) )
out_cache_loc_swa = (
torch.zeros((self.max_num_tokens,), dtype=torch.int64)
if model_runner.is_hybrid_swa
else None
)
positions = torch.zeros((self.max_num_tokens,), dtype=torch.int64) positions = torch.zeros((self.max_num_tokens,), dtype=torch.int64)
if self.is_multimodal: if self.is_multimodal:
input_embeds = torch.zeros( input_embeds = torch.zeros(
@@ -210,7 +205,6 @@ class BreakableCudaGraphRunner:
self.buffers = PrefillInputBuffers( self.buffers = PrefillInputBuffers(
input_ids=input_ids, input_ids=input_ids,
out_cache_loc=out_cache_loc, out_cache_loc=out_cache_loc,
out_cache_loc_swa=out_cache_loc_swa,
mamba_track_indices=None, mamba_track_indices=None,
mamba_track_mask=None, mamba_track_mask=None,
mamba_track_seqlens=None, mamba_track_seqlens=None,
@@ -302,11 +296,6 @@ class BreakableCudaGraphRunner:
token_to_kv_pool=self.model_runner.token_to_kv_pool, token_to_kv_pool=self.model_runner.token_to_kv_pool,
attn_backend=self.model_runner.attn_backend, attn_backend=self.model_runner.attn_backend,
out_cache_loc=buffers.out_cache_loc[:num_tokens], out_cache_loc=buffers.out_cache_loc[:num_tokens],
out_cache_loc_swa=(
buffers.out_cache_loc_swa[:num_tokens]
if buffers.out_cache_loc_swa is not None
else None
),
seq_lens_sum=num_tokens, seq_lens_sum=num_tokens,
mamba_track_indices=None, mamba_track_indices=None,
mamba_track_mask=None, mamba_track_mask=None,
@@ -400,6 +389,9 @@ class BreakableCudaGraphRunner:
self.model_runner.attn_backend.init_forward_metadata(forward_batch) self.model_runner.attn_backend.init_forward_metadata(forward_batch)
def run_once(): def run_once():
# Invalidate SWA loc cache — same fix as in cuda_graph_runner.run_once.
if self.model_runner.is_hybrid_swa:
self.model_runner.token_to_kv_pool.invalidate_loc_cache()
return self._run_forward(forward_batch, num_tokens) return self._run_forward(forward_batch, num_tokens)
for _ in range(2): for _ in range(2):
@@ -135,7 +135,6 @@ class DecodeInputBuffers(ForwardInputBuffers):
seq_lens: torch.Tensor seq_lens: torch.Tensor
seq_lens_cpu: torch.Tensor seq_lens_cpu: torch.Tensor
out_cache_loc: torch.Tensor out_cache_loc: torch.Tensor
out_cache_loc_swa: Optional[torch.Tensor]
positions: torch.Tensor positions: torch.Tensor
mrope_positions: torch.Tensor mrope_positions: torch.Tensor
num_token_non_padded: torch.Tensor num_token_non_padded: torch.Tensor
@@ -169,7 +168,6 @@ class DecodeInputBuffers(ForwardInputBuffers):
cache_loc_dtype: torch.dtype, cache_loc_dtype: torch.dtype,
enable_mamba_track: bool, enable_mamba_track: bool,
ne_token_table: Optional[torch.Tensor] = None, ne_token_table: Optional[torch.Tensor] = None,
is_hybrid_swa: bool = False,
hc_hidden_size: Optional[int] = None, hc_hidden_size: Optional[int] = None,
) -> "DecodeInputBuffers": ) -> "DecodeInputBuffers":
with torch.device(device): with torch.device(device):
@@ -178,11 +176,6 @@ class DecodeInputBuffers(ForwardInputBuffers):
req_pool_indices = torch.zeros((max_bs,), dtype=torch.int64) req_pool_indices = torch.zeros((max_bs,), dtype=torch.int64)
seq_lens = torch.full((max_bs,), seq_len_fill_value, dtype=torch.int32) seq_lens = torch.full((max_bs,), seq_len_fill_value, dtype=torch.int32)
out_cache_loc = torch.zeros((max_num_token,), dtype=cache_loc_dtype) out_cache_loc = torch.zeros((max_num_token,), dtype=cache_loc_dtype)
out_cache_loc_swa = (
torch.zeros((max_num_token,), dtype=torch.int32)
if is_hybrid_swa
else None
)
positions = torch.zeros((max_num_token,), dtype=torch.int64) positions = torch.zeros((max_num_token,), dtype=torch.int64)
mrope_positions = torch.zeros((3, max_num_token), dtype=torch.int64) mrope_positions = torch.zeros((3, max_num_token), dtype=torch.int64)
num_token_non_padded = torch.zeros((1,), dtype=torch.int32) num_token_non_padded = torch.zeros((1,), dtype=torch.int32)
@@ -260,7 +253,6 @@ class DecodeInputBuffers(ForwardInputBuffers):
seq_lens=seq_lens, seq_lens=seq_lens,
seq_lens_cpu=seq_lens_cpu, seq_lens_cpu=seq_lens_cpu,
out_cache_loc=out_cache_loc, out_cache_loc=out_cache_loc,
out_cache_loc_swa=out_cache_loc_swa,
positions=positions, positions=positions,
mrope_positions=mrope_positions, mrope_positions=mrope_positions,
num_token_non_padded=num_token_non_padded, num_token_non_padded=num_token_non_padded,
@@ -292,12 +284,6 @@ class DecodeInputBuffers(ForwardInputBuffers):
if bs != raw_bs: if bs != raw_bs:
self.seq_lens.fill_(seq_len_fill_value) self.seq_lens.fill_(seq_len_fill_value)
self.out_cache_loc.zero_() self.out_cache_loc.zero_()
# Padded SWA indices left over from a previous replay would point
# into real SWA slots, so set_kv_buffer on padded tokens would
# corrupt active requests' KV. Zero the whole buffer so padded
# positions map to the sentinel slot (matches piecewise runner).
if self.out_cache_loc_swa is not None:
self.out_cache_loc_swa.zero_()
if self.mamba_track_indices is not None: if self.mamba_track_indices is not None:
self.mamba_track_indices.zero_() self.mamba_track_indices.zero_()
if self.mamba_track_mask is not None: if self.mamba_track_mask is not None:
@@ -374,14 +360,6 @@ class DecodeInputBuffers(ForwardInputBuffers):
dsts.append(buf[:dim]) dsts.append(buf[:dim])
srcs.append(src) srcs.append(src)
# SWA cache location (int32, separate from the int64 batch above).
if (
self.out_cache_loc_swa is not None
and forward_batch.out_cache_loc_swa is not None
):
dsts.append(self.out_cache_loc_swa[:raw_num_token])
srcs.append(forward_batch.out_cache_loc_swa[:raw_num_token])
# Batch all GPU copies, grouped by dtype pair. # Batch all GPU copies, grouped by dtype pair.
_grouped_foreach_copy_(dsts, srcs) _grouped_foreach_copy_(dsts, srcs)
@@ -696,7 +674,6 @@ class CudaGraphRunner:
ne_token_table=( ne_token_table=(
model_runner.token_table if self.use_ngram_embedding else None model_runner.token_table if self.use_ngram_embedding else None
), ),
is_hybrid_swa=model_runner.is_hybrid_swa,
hc_hidden_size=getattr( hc_hidden_size=getattr(
self.model_runner.model_config, "hc_hidden_size", None self.model_runner.model_config, "hc_hidden_size", None
), ),
@@ -1089,6 +1066,11 @@ class CudaGraphRunner:
# Run and capture # Run and capture
def run_once(): def run_once():
# Without this, warmup-1 caches the translation; the capture run gets
# a hit, skips the gather, and replay reuses stale SWA locations.
if self.model_runner.is_hybrid_swa:
self.model_runner.token_to_kv_pool.invalidate_loc_cache()
# Clean intermediate result cache for DP attention # Clean intermediate result cache for DP attention
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None
set_dp_buffer_len( set_dp_buffer_len(
@@ -1123,15 +1105,6 @@ class CudaGraphRunner:
self.deepep_adapter.capture(is_extend_in_batch=False) self.deepep_adapter.capture(is_extend_in_batch=False)
# swa_loc must be set before capture so that set_kv_buffer's
# Python branch (if self.swa_loc is not None) takes the fast path,
# and the graph records GPU ops using this buffer instead of the
# per-layer translate_loc_from_full_to_swa fallback.
if self.buffers.out_cache_loc_swa is not None:
self.model_runner.token_to_kv_pool.set_swa_loc(
self.buffers.out_cache_loc_swa[:num_tokens]
)
for _ in range(2): for _ in range(2):
self.device_module.synchronize() self.device_module.synchronize()
self.model_runner.tp_group.barrier() self.model_runner.tp_group.barrier()
@@ -296,8 +296,6 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
# The original sequence length without being chunked. Qwen-1M related. # The original sequence length without being chunked. Qwen-1M related.
orig_seq_lens: Optional[torch.Tensor] = None orig_seq_lens: Optional[torch.Tensor] = None
# The indices of output tokens in the token_to_kv_pool_swa
out_cache_loc_swa: Optional[torch.Tensor] = None
# The indices to track mamba state with # The indices to track mamba state with
mamba_track_indices: Optional[torch.Tensor] = None # shape: [b], int64 mamba_track_indices: Optional[torch.Tensor] = None # shape: [b], int64
# The mask to track mamba state if needed # The mask to track mamba state if needed
@@ -655,14 +653,6 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
else: else:
ret._compute_mrope_positions(model_runner, batch) ret._compute_mrope_positions(model_runner, batch)
# Precompute SWA cache location once for all SWA layers
if model_runner.is_hybrid_swa and ret.out_cache_loc is not None:
ret.out_cache_loc_swa = (
model_runner.token_to_kv_pool_allocator.translate_loc_from_full_to_swa(
ret.out_cache_loc
)
)
# Init lora information # Init lora information
if model_runner.server_args.enable_lora: if model_runner.server_args.enable_lora:
# In the non-LoRA overlap loading case, we fetch LoRA adapters into the memory pool # In the non-LoRA overlap loading case, we fetch LoRA adapters into the memory pool
@@ -1015,10 +1005,6 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
) )
self.out_cache_loc = self._pad_tensor_to_size(self.out_cache_loc, num_tokens) self.out_cache_loc = self._pad_tensor_to_size(self.out_cache_loc, num_tokens)
if self.out_cache_loc_swa is not None:
self.out_cache_loc_swa = self._pad_tensor_to_size(
self.out_cache_loc_swa, num_tokens
)
if self.encoder_lens is not None: if self.encoder_lens is not None:
self.encoder_lens = self._pad_tensor_to_size(self.encoder_lens, bs) self.encoder_lens = self._pad_tensor_to_size(self.encoder_lens, bs)
self.positions = self._pad_tensor_to_size(self.positions, num_tokens) self.positions = self._pad_tensor_to_size(self.positions, num_tokens)
@@ -3250,9 +3250,8 @@ class ModelRunner(ModelRunnerKVCacheMixin):
server_args=self.server_args, server_args=self.server_args,
) )
# Use precomputed SWA cache location if self.is_hybrid_swa:
if forward_batch.out_cache_loc_swa is not None: self.token_to_kv_pool.invalidate_loc_cache()
self.token_to_kv_pool.set_swa_loc(forward_batch.out_cache_loc_swa)
# Hisparse coordinator # Hisparse coordinator
forward_batch.hisparse_coordinator = self.hisparse_coordinator forward_batch.hisparse_coordinator = self.hisparse_coordinator
@@ -81,7 +81,6 @@ _is_musa = is_musa()
class PrefillInputBuffers(ForwardInputBuffers): class PrefillInputBuffers(ForwardInputBuffers):
input_ids: torch.Tensor input_ids: torch.Tensor
out_cache_loc: torch.Tensor out_cache_loc: torch.Tensor
out_cache_loc_swa: Optional[torch.Tensor]
mamba_track_indices: Optional[torch.Tensor] mamba_track_indices: Optional[torch.Tensor]
mamba_track_mask: Optional[torch.Tensor] mamba_track_mask: Optional[torch.Tensor]
mamba_track_seqlens: Optional[torch.Tensor] mamba_track_seqlens: Optional[torch.Tensor]
@@ -247,11 +246,6 @@ class PiecewiseCudaGraphRunner:
out_cache_loc = torch.zeros( out_cache_loc = torch.zeros(
(self.max_num_tokens,), dtype=self._cache_loc_dtype() (self.max_num_tokens,), dtype=self._cache_loc_dtype()
) )
out_cache_loc_swa = (
torch.zeros((self.max_num_tokens,), dtype=torch.int32)
if model_runner.is_hybrid_swa
else None
)
mamba_track_indices = ( mamba_track_indices = (
torch.zeros((self.max_bs,), dtype=torch.int64) torch.zeros((self.max_bs,), dtype=torch.int64)
if self.mamba_track_enabled if self.mamba_track_enabled
@@ -291,7 +285,6 @@ class PiecewiseCudaGraphRunner:
self.buffers = PrefillInputBuffers( self.buffers = PrefillInputBuffers(
input_ids=input_ids, input_ids=input_ids,
out_cache_loc=out_cache_loc, out_cache_loc=out_cache_loc,
out_cache_loc_swa=out_cache_loc_swa,
mamba_track_indices=mamba_track_indices, mamba_track_indices=mamba_track_indices,
mamba_track_mask=mamba_track_mask, mamba_track_mask=mamba_track_mask,
mamba_track_seqlens=mamba_track_seqlens, mamba_track_seqlens=mamba_track_seqlens,
@@ -368,11 +361,6 @@ class PiecewiseCudaGraphRunner:
buffers.mrope_positions[:, :num_tokens] if self.is_multimodal else None buffers.mrope_positions[:, :num_tokens] if self.is_multimodal else None
) )
out_cache_loc = buffers.out_cache_loc[:num_tokens] out_cache_loc = buffers.out_cache_loc[:num_tokens]
out_cache_loc_swa = (
buffers.out_cache_loc_swa[:num_tokens]
if buffers.out_cache_loc_swa is not None
else None
)
mamba_track_indices = ( mamba_track_indices = (
buffers.mamba_track_indices[:1] buffers.mamba_track_indices[:1]
if buffers.mamba_track_indices is not None if buffers.mamba_track_indices is not None
@@ -403,7 +391,6 @@ class PiecewiseCudaGraphRunner:
token_to_kv_pool=self.model_runner.token_to_kv_pool, token_to_kv_pool=self.model_runner.token_to_kv_pool,
attn_backend=self.model_runner.attn_backend, attn_backend=self.model_runner.attn_backend,
out_cache_loc=out_cache_loc, out_cache_loc=out_cache_loc,
out_cache_loc_swa=out_cache_loc_swa,
seq_lens_sum=num_tokens, seq_lens_sum=num_tokens,
mamba_track_indices=mamba_track_indices, mamba_track_indices=mamba_track_indices,
mamba_track_mask=mamba_track_mask, mamba_track_mask=mamba_track_mask,
@@ -527,11 +514,6 @@ class PiecewiseCudaGraphRunner:
input_embeds = buffers.input_embeds[:num_tokens] if self.is_multimodal else None input_embeds = buffers.input_embeds[:num_tokens] if self.is_multimodal else None
out_cache_loc = buffers.out_cache_loc[:num_tokens] out_cache_loc = buffers.out_cache_loc[:num_tokens]
out_cache_loc_swa = (
buffers.out_cache_loc_swa[:num_tokens]
if buffers.out_cache_loc_swa is not None
else None
)
mamba_track_indices = ( mamba_track_indices = (
buffers.mamba_track_indices[:bs] buffers.mamba_track_indices[:bs]
if buffers.mamba_track_indices is not None if buffers.mamba_track_indices is not None
@@ -576,7 +558,6 @@ class PiecewiseCudaGraphRunner:
token_to_kv_pool=self.model_runner.token_to_kv_pool, token_to_kv_pool=self.model_runner.token_to_kv_pool,
attn_backend=self.model_runner.attn_backend, attn_backend=self.model_runner.attn_backend,
out_cache_loc=out_cache_loc, out_cache_loc=out_cache_loc,
out_cache_loc_swa=out_cache_loc_swa,
seq_lens_sum=num_tokens, seq_lens_sum=num_tokens,
mamba_track_indices=mamba_track_indices, mamba_track_indices=mamba_track_indices,
mamba_track_mask=mamba_track_mask, mamba_track_mask=mamba_track_mask,
@@ -614,6 +595,10 @@ class PiecewiseCudaGraphRunner:
# Run and capture # Run and capture
def run_once(): def run_once():
# Invalidate SWA loc cache — same fix as in cuda_graph_runner.run_once.
if self.model_runner.is_hybrid_swa:
self.model_runner.token_to_kv_pool.invalidate_loc_cache()
# Clean intermediate result cache for DP attention # Clean intermediate result cache for DP attention
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None
set_dp_buffer_len( set_dp_buffer_len(
@@ -662,8 +647,6 @@ class PiecewiseCudaGraphRunner:
self.raw_num_tokens = num_tokens self.raw_num_tokens = num_tokens
if static_num_tokens != num_tokens: if static_num_tokens != num_tokens:
buffers.out_cache_loc.zero_() buffers.out_cache_loc.zero_()
if buffers.out_cache_loc_swa is not None:
buffers.out_cache_loc_swa.zero_()
buffers.input_ids[num_tokens:static_num_tokens].zero_() buffers.input_ids[num_tokens:static_num_tokens].zero_()
buffers.positions[num_tokens:static_num_tokens].zero_() buffers.positions[num_tokens:static_num_tokens].zero_()
if self.is_multimodal: if self.is_multimodal:
@@ -676,12 +659,6 @@ class PiecewiseCudaGraphRunner:
buffers.input_ids[:num_tokens].copy_(forward_batch.input_ids) buffers.input_ids[:num_tokens].copy_(forward_batch.input_ids)
buffers.positions[:num_tokens].copy_(forward_batch.positions) buffers.positions[:num_tokens].copy_(forward_batch.positions)
buffers.out_cache_loc[:num_tokens].copy_(forward_batch.out_cache_loc) buffers.out_cache_loc[:num_tokens].copy_(forward_batch.out_cache_loc)
if buffers.out_cache_loc_swa is not None:
buffers.out_cache_loc_swa[: self.raw_num_tokens].copy_(
self.model_runner.token_to_kv_pool_allocator.translate_loc_from_full_to_swa(
forward_batch.out_cache_loc
)
)
if ( if (
buffers.mamba_track_indices is not None buffers.mamba_track_indices is not None
@@ -703,12 +680,6 @@ class PiecewiseCudaGraphRunner:
positions = buffers.positions[:static_num_tokens] positions = buffers.positions[:static_num_tokens]
out_cache_loc = buffers.out_cache_loc[:static_num_tokens] out_cache_loc = buffers.out_cache_loc[:static_num_tokens]
out_cache_loc_swa = (
buffers.out_cache_loc_swa[:static_num_tokens]
if buffers.out_cache_loc_swa is not None
else None
)
mamba_track_indices = ( mamba_track_indices = (
buffers.mamba_track_indices[:bs] buffers.mamba_track_indices[:bs]
if buffers.mamba_track_indices is not None if buffers.mamba_track_indices is not None
@@ -766,7 +737,6 @@ class PiecewiseCudaGraphRunner:
token_to_kv_pool=self.model_runner.token_to_kv_pool, token_to_kv_pool=self.model_runner.token_to_kv_pool,
attn_backend=self.model_runner.attn_backend, attn_backend=self.model_runner.attn_backend,
out_cache_loc=out_cache_loc, out_cache_loc=out_cache_loc,
out_cache_loc_swa=out_cache_loc_swa,
seq_lens_sum=forward_batch.seq_lens_sum, seq_lens_sum=forward_batch.seq_lens_sum,
mamba_track_indices=mamba_track_indices, mamba_track_indices=mamba_track_indices,
mamba_track_mask=mamba_track_mask, mamba_track_mask=mamba_track_mask,
@@ -807,9 +777,6 @@ class PiecewiseCudaGraphRunner:
), ),
) )
if out_cache_loc_swa is not None:
self.model_runner.token_to_kv_pool.set_swa_loc(out_cache_loc_swa)
return static_forward_batch return static_forward_batch
def replay( def replay(
@@ -355,6 +355,9 @@ class EAGLEDraftCudaGraphRunner:
# Run and capture # Run and capture
def run_once(): def run_once():
if self.model_runner.is_hybrid_swa:
self.model_runner.token_to_kv_pool.invalidate_loc_cache()
# Clean intermediate result cache for DP attention # Clean intermediate result cache for DP attention
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None
set_dp_buffer_len( set_dp_buffer_len(
@@ -395,6 +395,10 @@ class EAGLEDraftExtendCudaGraphRunner:
# Run and capture # Run and capture
def run_once(): def run_once():
# model.forward() bypasses _forward_raw(), so invalidate manually.
if self.model_runner.is_hybrid_swa:
self.model_runner.token_to_kv_pool.invalidate_loc_cache()
# Clean intermediate result cache for DP attention # Clean intermediate result cache for DP attention
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None
set_dp_buffer_len( set_dp_buffer_len(
@@ -288,6 +288,9 @@ class FrozenKVMTPCudaGraphRunner:
) )
def run_once(): def run_once():
if self.model_runner.is_hybrid_swa:
self.model_runner.token_to_kv_pool.invalidate_loc_cache()
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None
set_dp_buffer_len( set_dp_buffer_len(
global_dp_buffer_len, global_dp_buffer_len,
@@ -66,7 +66,6 @@ class MultiLayerEagleDraftExtendInputBuffers(ForwardInputBuffers):
# Sliced from shared parent buffers # Sliced from shared parent buffers
input_ids: torch.Tensor input_ids: torch.Tensor
out_cache_loc: torch.Tensor out_cache_loc: torch.Tensor
swa_out_cache_loc: torch.Tensor
positions: torch.Tensor positions: torch.Tensor
# Shared from parent # Shared from parent
seq_lens: torch.Tensor seq_lens: torch.Tensor
@@ -150,9 +149,6 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
out_cache_loc = cuda_graph_buffers["out_cache_loc"][ out_cache_loc = cuda_graph_buffers["out_cache_loc"][
offset : offset + self.max_num_token offset : offset + self.max_num_token
] ]
swa_out_cache_loc = cuda_graph_buffers["swa_out_cache_loc"][
offset : offset + self.max_num_token
]
positions = cuda_graph_buffers["positions"][ positions = cuda_graph_buffers["positions"][
offset : offset + self.max_num_token offset : offset + self.max_num_token
] ]
@@ -229,7 +225,6 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
self.buffers = MultiLayerEagleDraftExtendInputBuffers( self.buffers = MultiLayerEagleDraftExtendInputBuffers(
input_ids=input_ids, input_ids=input_ids,
out_cache_loc=out_cache_loc, out_cache_loc=out_cache_loc,
swa_out_cache_loc=swa_out_cache_loc,
positions=positions, positions=positions,
seq_lens=seq_lens, seq_lens=seq_lens,
seq_lens_cpu=seq_lens_cpu, seq_lens_cpu=seq_lens_cpu,
@@ -424,6 +419,10 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
# Run and capture # Run and capture
def run_once(): def run_once():
# model.forward() bypasses _forward_raw(), so invalidate manually.
if self.model_runner.is_hybrid_swa:
self.model_runner.token_to_kv_pool.invalidate_loc_cache()
# Clean intermediate result cache for DP attention # Clean intermediate result cache for DP attention
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None
set_dp_buffer_len( set_dp_buffer_len(
@@ -494,12 +493,6 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
forward_batch.req_to_token_pool.req_to_token, forward_batch.req_to_token_pool.req_to_token,
self.eagle_worker.req_to_hidden_states_pool, self.eagle_worker.req_to_hidden_states_pool,
) )
next_buffers.swa_out_cache_loc.copy_(
self.model_runner.token_to_kv_pool.translate_loc_from_full_to_swa(
next_buffers.out_cache_loc
)
)
forward_batch.out_cache_loc = output_cache_loc_backup forward_batch.out_cache_loc = output_cache_loc_backup
forward_batch.spec_info.hidden_states = hidden_states_backup forward_batch.spec_info.hidden_states = hidden_states_backup
return ret return ret
@@ -683,9 +676,6 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
self.cuda_graph_buffers["out_cache_loc"] = torch.ones( self.cuda_graph_buffers["out_cache_loc"] = torch.ones(
(self.offsets[-1],), dtype=torch.int64 (self.offsets[-1],), dtype=torch.int64
) )
self.cuda_graph_buffers["swa_out_cache_loc"] = torch.ones(
(self.offsets[-1],), dtype=torch.int64
)
self.cuda_graph_buffers["positions"] = torch.zeros( self.cuda_graph_buffers["positions"] = torch.zeros(
(self.offsets[-1],), dtype=torch.int64 (self.offsets[-1],), dtype=torch.int64
) )
@@ -733,7 +723,6 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
self.cuda_graph_buffers["input_ids"].zero_() self.cuda_graph_buffers["input_ids"].zero_()
self.cuda_graph_buffers["seq_lens"].fill_(self.seq_len_fill_value) self.cuda_graph_buffers["seq_lens"].fill_(self.seq_len_fill_value)
self.cuda_graph_buffers["out_cache_loc"].zero_() self.cuda_graph_buffers["out_cache_loc"].zero_()
self.cuda_graph_buffers["swa_out_cache_loc"].zero_()
self.cuda_graph_buffers["positions"].zero_() self.cuda_graph_buffers["positions"].zero_()
# `batch_result.accept_lens` is drafts + bonus. # `batch_result.accept_lens` is drafts + bonus.
bs = forward_batch.batch_size bs = forward_batch.batch_size
@@ -0,0 +1,233 @@
"""Manual tests for SWAKVPool.translate_loc_from_full_to_swa cache behaviour.
These tests cover three properties introduced by PR #25824:
1. Cache key uses data_ptr() — correctly distinguishes views at different
offsets within the same storage (untyped_storage().data_ptr() would not).
2. Allocator mutations invalidate the cache — alloc/free/clear/
set_full_to_swa_mapping each call invalidate_loc_cache() so the next
translation sees the fresh mapping.
3. BaseSWAKVPool.invalidate_loc_cache is a no-op default — subclasses that
don't cache (e.g. DSV4) can be called safely without AttributeError.
Run with:
python -m pytest test/manual/core/test_swa_loc_translation_cache.py -v
"""
import unittest
import torch
from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool, SWATokenToKVPoolAllocator
from sglang.srt.utils import get_device
from sglang.test.test_utils import CustomTestCase
def _build_pool(
kv_size: int = 32,
kv_size_swa: int = 32,
page_size: int = 1,
):
device = get_device()
num_layers = 8
full_layer_ids = [0, 4]
swa_layer_ids = [i for i in range(num_layers) if i not in set(full_layer_ids)]
pool = SWAKVPool(
size=kv_size,
size_swa=kv_size_swa,
page_size=page_size,
dtype=torch.bfloat16,
head_num=4,
head_dim=64,
swa_attention_layer_ids=swa_layer_ids,
full_attention_layer_ids=full_layer_ids,
enable_kvcache_transpose=False,
device=device,
)
allocator = SWATokenToKVPoolAllocator(
size=kv_size,
size_swa=kv_size_swa,
page_size=page_size,
dtype=torch.bfloat16,
device=device,
kvcache=pool,
need_sort=False,
)
return pool, allocator, device
class TestCacheKeyDataPtr(CustomTestCase):
"""Cache key uses data_ptr(), which encodes the storage offset."""
def test_same_offset_view_is_cache_hit(self):
"""Two different Python objects pointing to the same base are a hit."""
pool, allocator, device = _build_pool()
loc = allocator.alloc(4)
self.assertIsNotNone(loc)
# Create two slice objects at offset 0 — same data_ptr, same numel.
view_a = loc[:4]
view_b = loc[:4]
self.assertIsNot(view_a, view_b) # different Python objects
self.assertEqual(view_a.data_ptr(), view_b.data_ptr())
result_a = pool.translate_loc_from_full_to_swa(view_a)
result_b = pool.translate_loc_from_full_to_swa(view_b)
# Both should return the identical tensor (cache hit).
self.assertIs(result_a, result_b)
def test_different_offset_view_is_cache_miss(self):
"""Views at different offsets produce different data_ptr → cache miss."""
pool, allocator, device = _build_pool(kv_size=32, kv_size_swa=32)
loc = allocator.alloc(10)
self.assertIsNotNone(loc)
self.assertGreaterEqual(loc.numel(), 10)
view_lo = loc[0:5]
view_hi = loc[5:10]
self.assertEqual(view_lo.numel(), view_hi.numel()) # same numel
# Different data_ptr (different storage offset).
self.assertNotEqual(view_lo.data_ptr(), view_hi.data_ptr())
# Prime the cache with view_lo.
result_lo = pool.translate_loc_from_full_to_swa(view_lo)
# view_hi should be a cache miss and produce a distinct translation.
result_hi = pool.translate_loc_from_full_to_swa(view_hi)
# They should NOT be the same object (different cache entries).
self.assertIsNot(result_lo, result_hi)
# And the content must differ (different full indices → different swa).
self.assertFalse(torch.equal(result_lo, result_hi))
def test_storage_base_ptr_would_collide(self):
"""Demonstrate that untyped_storage().data_ptr() WOULD collide for the
two views above — confirming data_ptr() is the right key."""
t = torch.arange(20, device=get_device())
a, b = t[0:10], t[5:15]
# Same storage base — old key would collide.
self.assertEqual(a.untyped_storage().data_ptr(), b.untyped_storage().data_ptr())
self.assertEqual(a.numel(), b.numel())
# But data_ptr differs — new key is safe.
self.assertNotEqual(a.data_ptr(), b.data_ptr())
class TestAllocatorMutationInvalidation(CustomTestCase):
"""Each allocator method that writes the mapping calls invalidate_loc_cache."""
def _prime_and_check_invalidation(self, pool, allocator, mutate_fn):
"""Helper: prime cache, mutate, assert fresh translation."""
loc = allocator.alloc(4)
self.assertIsNotNone(loc)
# Prime the cache.
first = pool.translate_loc_from_full_to_swa(loc)
self.assertIsNotNone(pool._cached_loc_key)
# Mutate — should invalidate.
mutate_fn(allocator, loc)
# Cache must be cleared after mutation.
self.assertIsNone(pool._cached_loc_key)
self.assertIsNone(pool._cached_swa_loc)
def test_alloc_invalidates(self):
pool, allocator, _ = _build_pool()
loc = allocator.alloc(4)
pool.translate_loc_from_full_to_swa(loc)
self.assertIsNotNone(pool._cached_loc_key)
# Another alloc should invalidate.
allocator.alloc(4)
self.assertIsNone(pool._cached_loc_key)
def test_free_swa_invalidates(self):
pool, allocator, _ = _build_pool()
loc = allocator.alloc(4)
pool.translate_loc_from_full_to_swa(loc)
self.assertIsNotNone(pool._cached_loc_key)
allocator.free_swa(loc)
self.assertIsNone(pool._cached_loc_key)
def test_clear_invalidates(self):
pool, allocator, _ = _build_pool()
loc = allocator.alloc(4)
pool.translate_loc_from_full_to_swa(loc)
self.assertIsNotNone(pool._cached_loc_key)
allocator.clear()
self.assertIsNone(pool._cached_loc_key)
def test_set_full_to_swa_mapping_invalidates(self):
"""HiCache load-back path: set_full_to_swa_mapping must invalidate."""
pool, allocator, device = _build_pool(kv_size=32, kv_size_swa=32)
loc = allocator.alloc(4)
pool.translate_loc_from_full_to_swa(loc)
self.assertIsNotNone(pool._cached_loc_key)
# Simulate HiCache rebuild with new swa indices.
new_swa = torch.arange(4, dtype=torch.int64, device=device)
allocator.set_full_to_swa_mapping(loc, new_swa)
self.assertIsNone(pool._cached_loc_key)
# Translation after rebuild should reflect the new mapping.
result = pool.translate_loc_from_full_to_swa(loc)
self.assertEqual(result.tolist(), new_swa.tolist())
class TestBaseClassNoOp(CustomTestCase):
"""BaseSWAKVPool.invalidate_loc_cache is a no-op default — must not raise."""
def test_noop_does_not_raise(self):
# BaseSWAKVPool is abstract; instantiate via SWAKVPool which inherits.
pool, _, _ = _build_pool()
# Calling on the concrete class uses the override — that's fine.
pool.invalidate_loc_cache() # must not raise
pool.invalidate_loc_cache() # idempotent
def test_base_class_noop_directly(self):
"""Call the base-class method directly to verify it's a true no-op."""
pool, _, _ = _build_pool()
# Prime the cache first.
loc = pool.full_to_swa_index_mapping # any tensor
pool._cached_loc_key = ("dummy", 1)
pool._cached_swa_loc = torch.zeros(1)
# Call the BASE class method directly — should not clear the cache
# (it's a no-op; the concrete override is what clears).
BaseSWAKVPool.invalidate_loc_cache(pool)
# base no-op: cache untouched
self.assertIsNotNone(pool._cached_loc_key)
class TestExplicitInvalidationCycle(CustomTestCase):
"""Simulates the per-forward-pass invalidation done by model_runner."""
def test_fresh_translation_after_explicit_invalidation(self):
"""After invalidate_loc_cache(), a new alloc produces the right mapping."""
pool, allocator, device = _build_pool(kv_size=32, kv_size_swa=32)
# First "forward pass": alloc 4 tokens, translate.
loc1 = allocator.alloc(4)
trans1 = pool.translate_loc_from_full_to_swa(loc1).clone()
# Simulate start of next forward pass: model_runner calls invalidate.
pool.invalidate_loc_cache()
self.assertIsNone(pool._cached_loc_key)
# Alloc 4 more (mapping changes), translate loc1 again.
loc2 = allocator.alloc(4)
# Alloc already invalidated; translate loc1 with fresh mapping.
trans1_after = pool.translate_loc_from_full_to_swa(loc1)
# loc1's SWA mapping hasn't changed (same full→swa assignment),
# so result should be equal — but it must have been recomputed
# (cache key was None before this call).
self.assertEqual(trans1.tolist(), trans1_after.tolist())
# loc2 should have different translation than loc1.
trans2 = pool.translate_loc_from_full_to_swa(loc2)
# They have different indices, so translation differs.
self.assertFalse(torch.equal(trans1_after, trans2))
if __name__ == "__main__":
unittest.main()
@@ -33,6 +33,7 @@ def _make_self(*, page_size: int, full_available: int, swa_available: int):
), ),
translate_loc_from_full_to_swa=lambda last_loc: last_loc, translate_loc_from_full_to_swa=lambda last_loc: last_loc,
full_to_swa_index_mapping=torch.zeros(64, dtype=torch.int64), full_to_swa_index_mapping=torch.zeros(64, dtype=torch.int64),
_kvcache=SimpleNamespace(invalidate_loc_cache=lambda: None),
) )