[mem_cache] Make free_swa sync-free on page_size == 1 (#36723)

This commit is contained in:
Liangsheng Yin
2026-09-02 14:18:22 -07:00
committed by GitHub
parent acea43079f
commit 19c7679e9e
4 changed files with 129 additions and 78 deletions
+27 -7
View File
@@ -6,6 +6,7 @@ from sglang.srt.mem_cache.allocator.token import TokenToKVPoolAllocator
from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool
from sglang.srt.utils import is_npu
from sglang.srt.utils.common import get_num_new_pages
from sglang.srt.utils.invariants import Bucket, Invariant, IsTrue, expect
_is_npu = is_npu()
@@ -17,6 +18,13 @@ if _is_npu:
)
# free_swa releases whatever the mapping points at, so an entry that reads as the
# padding slot would push slot 0 into the SWA free list and hand it out twice.
_SWA_PEER_MAPPED = Invariant("swa.peer_mapped", Bucket.FATAL_UNCONTAINABLE, IsTrue())
# free_full leaves the mapping alone, so a live entry would strand its SWA peer.
_SWA_PEER_RELEASED = Invariant("swa.peer_released", Bucket.GUARD, IsTrue())
class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
"""Allocator for SWA hybrid KV cache."""
@@ -355,11 +363,15 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
return
if self.page_size == 1:
# A filter here would make the output shape data-dependent,
# which costs a device-to-host sync.
mapping_indices = free_index
swa_indices = self.full_to_swa_index_mapping[mapping_indices]
expect(_SWA_PEER_MAPPED, swa_indices > 0, msg="caller wants free_full")
else:
mapping_indices = self._expand_to_full_pages(free_index)
swa_indices = self.full_to_swa_index_mapping[mapping_indices]
swa_indices = self.full_to_swa_index_mapping[mapping_indices]
self.clear_full_to_swa_mapping(mapping_indices)
if self.free_group is not None:
@@ -371,18 +383,26 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
self._release_swa(swa_indices)
def _release_swa(self, swa_indices: torch.Tensor):
# One filter per group: its data-dependent shape costs a sync, and
# filtering the batch selects the same slots as filtering per call.
self.swa_attn_allocator.free(swa_indices[swa_indices > 0])
if self.page_size > 1:
# HiCache LOAD_BACK re-pairs a page-aligned full chunk with an offset
# SWA one (commit_hicache_transfer advances by raw token count), so a
# page can hold unmapped slots; one filter per group, not per call.
swa_indices = swa_indices[swa_indices > 0]
self.swa_attn_allocator.free(swa_indices)
assert self.swa_attn_allocator.available_size() <= self.swa_attn_allocator.size
def free_full(self, free_index: torch.Tensor):
if free_index.numel() == 0:
return
# Checked at enqueue: a cache action later in this group may pair the
# slot again, and that new peer is not this call's to judge.
expect(
_SWA_PEER_RELEASED,
self.full_to_swa_index_mapping[free_index] == 0,
msg="caller wants free",
)
if self.free_group is None:
# Full side only: a tombstoned range's mapping entries read as the
# padding slot, so `free` would push slot 0 into the SWA free list.
self.full_attn_allocator.free(free_index)
else:
self.full_free_group.append(self._copy_for_free_group(free_index))
@@ -404,7 +424,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
if self.full_free_group:
full_free_group = self.full_free_group
self.full_free_group = []
self.free_full(torch.cat(full_free_group))
self.full_attn_allocator.free(torch.cat(full_free_group))
assert (
self.full_attn_allocator.available_size() <= self.full_attn_allocator.size
)