[mem_cache][4/N] refactor: extract MambaTokenToKVPoolAllocator into allocator/ (#27256)

This commit is contained in:
shuwenn
2026-06-07 10:46:29 +08:00
committed by GitHub
parent 5da265de30
commit e57323cae9
18 changed files with 177 additions and 120 deletions
@@ -109,7 +109,7 @@ class TestMamba(unittest.TestCase):
)
assert req_to_token_pool.available_size() == max_num_reqs
assert req_to_token_pool.mamba_pool.available_size() == mamba_cache_size
assert req_to_token_pool.mamba_allocator.available_size() == mamba_cache_size
sampling_params = SamplingParams(
temperature=0,
@@ -125,34 +125,41 @@ class TestMamba(unittest.TestCase):
# alloc req
req_to_token_pool.alloc([req])
assert req_to_token_pool.available_size() == max_num_reqs - 1
assert req_to_token_pool.mamba_pool.available_size() == mamba_cache_size - 1
assert (
req_to_token_pool.mamba_allocator.available_size() == mamba_cache_size - 1
)
# free req
req_to_token_pool.free_mamba_cache(req)
req_to_token_pool.free(req)
assert req_to_token_pool.available_size() == max_num_reqs
assert req_to_token_pool.mamba_pool.available_size() == mamba_cache_size
assert req_to_token_pool.mamba_allocator.available_size() == mamba_cache_size
# alloc req without free mamba cache
req.mamba_pool_idx = None
req_to_token_pool.alloc([req])
req_to_token_pool.free(req)
assert req_to_token_pool.available_size() == max_num_reqs
assert req_to_token_pool.mamba_pool.available_size() == mamba_cache_size - 1
assert (
req_to_token_pool.mamba_allocator.available_size() == mamba_cache_size - 1
)
# alloc again
req_to_token_pool.alloc([req])
assert req_to_token_pool.available_size() == max_num_reqs - 1
assert req_to_token_pool.mamba_pool.available_size() == mamba_cache_size - 1
assert (
req_to_token_pool.mamba_allocator.available_size() == mamba_cache_size - 1
)
def test_mamba_radix_cache_1(self):
tree, allocator, req_to_token_pool, make_dummy_req = (
self._setup_tree_and_allocator()
)
mamba_allocator = req_to_token_pool.mamba_allocator
mamba_pool = req_to_token_pool.mamba_pool
# test
print(
f"[Start] allocator mamba available size: {mamba_pool.available_size()}, full available size: {allocator.available_size()}"
f"[Start] allocator mamba available size: {mamba_allocator.available_size()}, full available size: {allocator.available_size()}"
)
req1 = make_dummy_req()
req1_token_ids, req1_kv_indices = [1, 2, 3], allocator.alloc(3)
@@ -170,7 +177,7 @@ class TestMamba(unittest.TestCase):
)
prefix_len = result.prefix_len
print(
f"req1: prefix_len: {prefix_len}, allocator mamba available size: {mamba_pool.available_size()}, full available size: {allocator.available_size()}"
f"req1: prefix_len: {prefix_len}, allocator mamba available size: {mamba_allocator.available_size()}, full available size: {allocator.available_size()}"
)
req2 = make_dummy_req()
req2_token_ids, req2_kv_indices = [1, 2, 3, 4, 5, 6, 7], allocator.alloc(7)
@@ -188,7 +195,7 @@ class TestMamba(unittest.TestCase):
)
prefix_len = result.prefix_len
print(
f"req2: prefix_len: {prefix_len}, allocator mamba available size: {mamba_pool.available_size()}, full available size: {allocator.available_size()}"
f"req2: prefix_len: {prefix_len}, allocator mamba available size: {mamba_allocator.available_size()}, full available size: {allocator.available_size()}"
)
req3 = make_dummy_req()
@@ -207,7 +214,7 @@ class TestMamba(unittest.TestCase):
)
prefix_len = result.prefix_len
print(
f"req3: prefix_len: {prefix_len}, allocator mamba available size: {mamba_pool.available_size()}, full available size: {allocator.available_size()}"
f"req3: prefix_len: {prefix_len}, allocator mamba available size: {mamba_allocator.available_size()}, full available size: {allocator.available_size()}"
)
req4 = make_dummy_req()
req4_token_ids, req4_kv_indices = [1, 2, 3, 4, 5, 60, 70], allocator.alloc(7)
@@ -225,7 +232,7 @@ class TestMamba(unittest.TestCase):
)
prefix_len = result.prefix_len
print(
f"req4: prefix_len: {prefix_len}, allocator mamba available size: {mamba_pool.available_size()}, full available size: {allocator.available_size()}"
f"req4: prefix_len: {prefix_len}, allocator mamba available size: {mamba_allocator.available_size()}, full available size: {allocator.available_size()}"
)
tree.pretty_print()
@@ -553,7 +560,7 @@ class TestMamba(unittest.TestCase):
_, _, req_to_token_pool, _ = self._setup_tree_and_allocator()
mamba_pool = req_to_token_pool.mamba_pool
n = 3
indices = mamba_pool.alloc(n)
indices = req_to_token_pool.mamba_allocator.alloc(n)
self.assertIsNotNone(indices)
# Write known sentinel values at the allocated slots.
@@ -608,7 +615,7 @@ class TestMamba(unittest.TestCase):
n_tokens = 4
kv_indices = allocator.alloc(n_tokens)
self.assertIsNotNone(kv_indices)
mamba_indices = mamba_pool.alloc(1)
mamba_indices = req_to_token_pool.mamba_allocator.alloc(1)
self.assertIsNotNone(mamba_indices)
# Write sentinel values into KV buffers (all full-attention layers).
@@ -3545,7 +3545,7 @@ class UnifiedRadixCacheSuite:
xfer = tree.components[ComponentType.MAMBA].build_hicache_transfers(
node, CacheTransferPhase.LOAD_BACK
)[0]
new_mamba = req_to_token_pool.mamba_pool.alloc(1)
new_mamba = req_to_token_pool.mamba_allocator.alloc(1)
self.assertIsNotNone(new_mamba)
xfer.device_indices = new_mamba
tree.components[ComponentType.MAMBA].commit_hicache_transfer(