[srt] Batch scheduler cache frees (#33475)
This commit is contained in:
@@ -391,6 +391,60 @@ class TestRadixCache(unittest.TestCase):
|
||||
)
|
||||
self.assertEqual(cache.total_size(), 5)
|
||||
|
||||
def test_cache_unfinished_req_deferred_free_keeps_original_indices(self):
|
||||
class DeferredFreeAllocator:
|
||||
device = torch.device("cpu")
|
||||
|
||||
def __init__(self):
|
||||
self.free_group = []
|
||||
self.freed = None
|
||||
|
||||
def free_group_begin(self):
|
||||
self.free_group = []
|
||||
|
||||
def free_segment(self, free_index, *, start_pos):
|
||||
self.free_group.append(free_index)
|
||||
|
||||
def free_group_end(self):
|
||||
self.freed = torch.cat(self.free_group)
|
||||
|
||||
class ReqToTokenPool:
|
||||
def __init__(self, row):
|
||||
self.req_to_token = row.unsqueeze(0)
|
||||
|
||||
def write(self, indices, values):
|
||||
self.req_to_token[indices] = values
|
||||
|
||||
allocator = DeferredFreeAllocator()
|
||||
cache = RadixCache.create_simulated(mock_allocator=allocator)
|
||||
token_ids = array("q", [1, 2, 3])
|
||||
tree_indices = torch.tensor([10, 11, 12], dtype=torch.int64)
|
||||
request_indices = torch.tensor([20, 21, 22], dtype=torch.int64)
|
||||
cache.insert(
|
||||
InsertParams(
|
||||
key=RadixKey(array("q", token_ids)),
|
||||
value=tree_indices,
|
||||
)
|
||||
)
|
||||
cache.req_to_token_pool = ReqToTokenPool(request_indices.clone())
|
||||
req = unittest.mock.Mock(
|
||||
req_pool_idx=0,
|
||||
cache_protected_len=0,
|
||||
extra_key=None,
|
||||
priority=0,
|
||||
last_node=cache.root_node,
|
||||
)
|
||||
req.get_fill_ids.return_value = token_ids
|
||||
|
||||
allocator.free_group_begin()
|
||||
cache.cache_unfinished_req(req)
|
||||
allocator.free_group_end()
|
||||
|
||||
torch.testing.assert_close(allocator.freed, request_indices)
|
||||
torch.testing.assert_close(
|
||||
cache.req_to_token_pool.req_to_token[0], tree_indices
|
||||
)
|
||||
|
||||
def test_kv_cache_events(self):
|
||||
"""Test KV cache events functionality."""
|
||||
test_cases = [
|
||||
|
||||
@@ -224,6 +224,41 @@ class TestSWA(unittest.TestCase):
|
||||
allocator.free_swa(full_indices[1:2])
|
||||
self.assertEqual(allocator.swa_available_size(), 16)
|
||||
|
||||
def test_free_swa_batches_with_free_group(self):
|
||||
_, allocator, _ = _build_swa_tree(
|
||||
is_eagle=False,
|
||||
kv_size=32,
|
||||
kv_size_swa=32,
|
||||
)
|
||||
index_batches = []
|
||||
for size in (2, 3, 1, 4):
|
||||
indices = _swa_alloc(allocator, size)
|
||||
assert indices is not None
|
||||
index_batches.append(indices)
|
||||
|
||||
available_before_free = allocator.swa_available_size()
|
||||
allocator.free_group_begin()
|
||||
for indices in index_batches:
|
||||
allocator.free_swa(indices)
|
||||
|
||||
self.assertEqual(len(allocator.swa_free_group), len(index_batches))
|
||||
self.assertEqual(allocator.swa_available_size(), available_before_free)
|
||||
|
||||
allocator.free_group_end()
|
||||
|
||||
all_indices = torch.cat(index_batches).to(torch.int64)
|
||||
self.assertEqual(allocator.swa_free_group, [])
|
||||
self.assertTrue(
|
||||
torch.equal(
|
||||
allocator.full_to_swa_index_mapping[all_indices],
|
||||
torch.zeros_like(all_indices),
|
||||
)
|
||||
)
|
||||
self.assertEqual(
|
||||
allocator.swa_available_size(),
|
||||
available_before_free + all_indices.numel(),
|
||||
)
|
||||
|
||||
def test_swa_radix_cache_1(self):
|
||||
# args
|
||||
req_size = 10
|
||||
|
||||
Reference in New Issue
Block a user