[misc] Resolve SWA ownership at enqueue time for grouped free() (#36646)

This commit is contained in:
Liangsheng Yin
2026-09-01 23:06:55 -07:00
committed by GitHub
parent 832d029870
commit 01c3a5f54f
3 changed files with 143 additions and 14 deletions
@@ -94,6 +94,7 @@ class BaseTokenToKVPoolAllocator(abc.ABC):
return self._kvcache
def free_group_begin(self):
assert self.free_group is None, "free groups cannot be nested"
self.free_group = []
def free_group_end(self):
+30 -13
View File
@@ -319,15 +319,10 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
return
# NOTE: the API is not idempotent.
if self.free_group is None:
self.full_attn_allocator.free(free_index)
self.free_swa(free_index)
else:
self.free_group.append(self._copy_for_free_group(free_index))
assert (
self.full_attn_allocator.available_size() <= self.full_attn_allocator.size
)
assert self.swa_attn_allocator.available_size() <= self.swa_attn_allocator.size
# SWA first: it reads the mapping, and a cache action later in this group
# can re-point free_index at a different SWA slot.
self.free_swa(free_index)
self.free_full(free_index)
def set_full_to_swa_mapping(
self, full_indices: torch.Tensor, swa_indices: torch.Tensor
@@ -365,7 +360,6 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
mapping_indices = self._expand_to_full_pages(free_index)
swa_indices = self.full_to_swa_index_mapping[mapping_indices]
swa_indices = swa_indices[swa_indices > 0]
self.clear_full_to_swa_mapping(mapping_indices)
if self.free_group is not None:
@@ -374,7 +368,13 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
self.swa_free_group.append(swa_indices)
return
self.swa_attn_allocator.free(swa_indices)
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])
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:
@@ -400,11 +400,15 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
if self.swa_free_group:
swa_free_group = self.swa_free_group
self.swa_free_group = []
self.swa_attn_allocator.free(torch.cat(swa_free_group))
self._release_swa(torch.cat(swa_free_group))
if self.full_free_group:
full_free_group = self.full_free_group
self.full_free_group = []
self.free_full(torch.cat(full_free_group))
assert (
self.full_attn_allocator.available_size() <= self.full_attn_allocator.size
)
assert self.swa_attn_allocator.available_size() <= self.swa_attn_allocator.size
def _expand_to_full_pages(self, indices: torch.Tensor) -> torch.Tensor:
# Duplicates are kept: deduplicating would be a torch.unique whose
@@ -516,6 +520,19 @@ class PureSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
def translate_loc_from_full_to_swa(self, kv_indices: torch.Tensor):
return kv_indices
def set_full_to_swa_mapping(
self, full_indices: torch.Tensor, swa_indices: torch.Tensor
) -> None:
# Registered with the KV pool and read by the attention kernels.
raise NotImplementedError(
"PureSWATokenToKVPoolAllocator has no full->SWA mapping to rewrite"
)
def clear_full_to_swa_mapping(self, full_indices: torch.Tensor) -> None:
raise NotImplementedError(
"PureSWATokenToKVPoolAllocator has no full->SWA mapping to clear"
)
def alloc(self, need_size: int):
assert self.page_size == 1
return self.swa_attn_allocator.alloc(need_size)
@@ -560,7 +577,7 @@ class PureSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
# Not inherited: the SWA parent's hooks drive swa_free_group and
# full_free_group, which this pure-SWA variant does not have.
def free_group_begin(self):
self.free_group = []
BaseTokenToKVPoolAllocator.free_group_begin(self)
def free_group_end(self):
pending, self.free_group = self.free_group, None
@@ -7,7 +7,10 @@ import torch
from sglang.srt.disaggregation.kv_events import BlockRemoved, BlockStored
from sglang.srt.environ import envs
from sglang.srt.mem_cache.allocator.base import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator
from sglang.srt.mem_cache.allocator.swa import (
PureSWATokenToKVPoolAllocator,
SWATokenToKVPoolAllocator,
)
from sglang.srt.mem_cache.base_prefix_cache import (
BasePrefixCache,
DecLockRefParams,
@@ -106,6 +109,29 @@ def _build_swa_tree(
return tree, allocator, req_to_token_pool
def _build_pure_swa_allocator(size_swa: int = 16):
device = get_device()
kv_pool = SWAKVPool(
size=0,
size_swa=size_swa,
page_size=1,
dtype=torch.bfloat16,
head_num=8,
head_dim=128,
swa_attention_layer_ids=list(range(4)),
full_attention_layer_ids=[],
device=device,
)
return PureSWATokenToKVPoolAllocator(
size_swa=size_swa,
page_size=1,
dtype=torch.bfloat16,
device=device,
kvcache=kv_pool,
need_sort=False,
)
def _swa_alloc(allocator, need_size):
"""SWA-pool alloc that also works for page_size > 1 (built-in alloc asserts page_size == 1)."""
if allocator.page_size == 1:
@@ -335,6 +361,91 @@ class TestSWA(unittest.TestCase):
torch.isin(new_swa, allocator.swa_attn_allocator.free_pages).item()
)
def _build_two_mapped_slots(self, page_size=1):
_, allocator, _ = _build_swa_tree(
is_eagle=False,
page_size=page_size,
kv_size=8 * page_size,
kv_size_swa=8 * page_size,
)
old_full = _swa_alloc(allocator, page_size)
new_full = _swa_alloc(allocator, page_size)
assert old_full is not None and new_full is not None
old_swa = allocator.full_to_swa_index_mapping[old_full].clone()
new_swa = allocator.full_to_swa_index_mapping[new_full].clone()
return allocator, old_full, new_full, old_swa, new_swa
def _swa_slot_is_free(self, allocator, swa_index):
# free_pages holds page ids for page_size > 1 and token ids otherwise,
# so compare in page space (a no-op divide when page_size == 1).
swa_pages = swa_index // allocator.page_size
free_pages = allocator.swa_attn_allocator.free_pages
return bool(torch.isin(swa_pages, free_pages).all().item())
def _run_remap_during_free_group(self, allocator, old_full, new_full, new_swa):
"""Queue a combined free, then transfer another SWA slot onto the same
full slot before the group flushes -- what tombstone recovery does."""
allocator.free_group_begin()
allocator.free(old_full)
allocator.set_full_to_swa_mapping(old_full, new_swa)
allocator.clear_full_to_swa_mapping(new_full)
allocator.free_group_end()
def test_free_group_owns_mapping_at_enqueue_time(self):
for page_size in (1, 4):
with self.subTest(page_size=page_size):
allocator, old_full, new_full, old_swa, new_swa = (
self._build_two_mapped_slots(page_size=page_size)
)
available_before = allocator.swa_available_size()
self._run_remap_during_free_group(
allocator, old_full, new_full, new_swa
)
self.assertTrue(
self._swa_slot_is_free(allocator, old_swa),
"the SWA slot owned at enqueue time leaked",
)
self.assertFalse(
self._swa_slot_is_free(allocator, new_swa),
"the replacement SWA slot was freed while still mapped",
)
self.assertEqual(
allocator.swa_available_size(), available_before + page_size
)
# Everything still in use stays reachable through the mapping.
mapped = allocator.full_to_swa_index_mapping[:-1]
num_mapped = int((mapped > 0).sum().item())
num_in_use = (
allocator.swa_attn_allocator.size - allocator.swa_available_size()
)
self.assertEqual(num_mapped, num_in_use)
def test_free_group_owns_tombstoned_indices(self):
"""free_swa then free of the same full slot must free the SWA slot once."""
allocator, full_indices, _, swa_indices, _ = self._build_two_mapped_slots()
swa_available_before = allocator.swa_available_size()
allocator.free_group_begin()
allocator.free_swa(full_indices)
allocator.free(full_indices)
allocator.free_group_end()
self.assertEqual(allocator.swa_available_size(), swa_available_before + 1)
self.assertTrue(self._swa_slot_is_free(allocator, swa_indices))
def test_pure_swa_rejects_mapping_edits(self):
allocator = _build_pure_swa_allocator()
indices = allocator.alloc(2)
with self.assertRaises(NotImplementedError):
allocator.clear_full_to_swa_mapping(indices)
with self.assertRaises(NotImplementedError):
allocator.set_full_to_swa_mapping(indices, indices)
torch.testing.assert_close(
allocator.full_to_swa_index_mapping[indices], indices
)
def test_swa_radix_cache_1(self):
# args
req_size = 10