[mem_cache][4/N] refactor: extract MambaTokenToKVPoolAllocator into allocator/ (#27256)
This commit is contained in:
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user