[misc] Resolve SWA ownership at enqueue time for grouped free() (#36646)
This commit is contained in:
@@ -94,6 +94,7 @@ class BaseTokenToKVPoolAllocator(abc.ABC):
|
|||||||
return self._kvcache
|
return self._kvcache
|
||||||
|
|
||||||
def free_group_begin(self):
|
def free_group_begin(self):
|
||||||
|
assert self.free_group is None, "free groups cannot be nested"
|
||||||
self.free_group = []
|
self.free_group = []
|
||||||
|
|
||||||
def free_group_end(self):
|
def free_group_end(self):
|
||||||
|
|||||||
@@ -319,15 +319,10 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
return
|
return
|
||||||
|
|
||||||
# NOTE: the API is not idempotent.
|
# NOTE: the API is not idempotent.
|
||||||
if self.free_group is None:
|
# SWA first: it reads the mapping, and a cache action later in this group
|
||||||
self.full_attn_allocator.free(free_index)
|
# can re-point free_index at a different SWA slot.
|
||||||
self.free_swa(free_index)
|
self.free_swa(free_index)
|
||||||
else:
|
self.free_full(free_index)
|
||||||
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
|
|
||||||
|
|
||||||
def set_full_to_swa_mapping(
|
def set_full_to_swa_mapping(
|
||||||
self, full_indices: torch.Tensor, swa_indices: torch.Tensor
|
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)
|
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]
|
||||||
swa_indices = swa_indices[swa_indices > 0]
|
|
||||||
self.clear_full_to_swa_mapping(mapping_indices)
|
self.clear_full_to_swa_mapping(mapping_indices)
|
||||||
|
|
||||||
if self.free_group is not None:
|
if self.free_group is not None:
|
||||||
@@ -374,7 +368,13 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
self.swa_free_group.append(swa_indices)
|
self.swa_free_group.append(swa_indices)
|
||||||
return
|
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):
|
def free_full(self, free_index: torch.Tensor):
|
||||||
if free_index.numel() == 0:
|
if free_index.numel() == 0:
|
||||||
@@ -400,11 +400,15 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
if self.swa_free_group:
|
if self.swa_free_group:
|
||||||
swa_free_group = self.swa_free_group
|
swa_free_group = self.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:
|
if self.full_free_group:
|
||||||
full_free_group = self.full_free_group
|
full_free_group = self.full_free_group
|
||||||
self.full_free_group = []
|
self.full_free_group = []
|
||||||
self.free_full(torch.cat(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:
|
def _expand_to_full_pages(self, indices: torch.Tensor) -> torch.Tensor:
|
||||||
# Duplicates are kept: deduplicating would be a torch.unique whose
|
# 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):
|
def translate_loc_from_full_to_swa(self, kv_indices: torch.Tensor):
|
||||||
return kv_indices
|
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):
|
def alloc(self, need_size: int):
|
||||||
assert self.page_size == 1
|
assert self.page_size == 1
|
||||||
return self.swa_attn_allocator.alloc(need_size)
|
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
|
# Not inherited: the SWA parent's hooks drive swa_free_group and
|
||||||
# full_free_group, which this pure-SWA variant does not have.
|
# full_free_group, which this pure-SWA variant does not have.
|
||||||
def free_group_begin(self):
|
def free_group_begin(self):
|
||||||
self.free_group = []
|
BaseTokenToKVPoolAllocator.free_group_begin(self)
|
||||||
|
|
||||||
def free_group_end(self):
|
def free_group_end(self):
|
||||||
pending, self.free_group = self.free_group, None
|
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.disaggregation.kv_events import BlockRemoved, BlockStored
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.mem_cache.allocator.base import BaseTokenToKVPoolAllocator
|
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 (
|
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||||
BasePrefixCache,
|
BasePrefixCache,
|
||||||
DecLockRefParams,
|
DecLockRefParams,
|
||||||
@@ -106,6 +109,29 @@ def _build_swa_tree(
|
|||||||
return tree, allocator, req_to_token_pool
|
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):
|
def _swa_alloc(allocator, need_size):
|
||||||
"""SWA-pool alloc that also works for page_size > 1 (built-in alloc asserts page_size == 1)."""
|
"""SWA-pool alloc that also works for page_size > 1 (built-in alloc asserts page_size == 1)."""
|
||||||
if allocator.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()
|
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):
|
def test_swa_radix_cache_1(self):
|
||||||
# args
|
# args
|
||||||
req_size = 10
|
req_size = 10
|
||||||
|
|||||||
Reference in New Issue
Block a user