[mem_cache] Move mamba state and retraction_backup into ReqKvInfo (#37164)

This commit is contained in:
Liangsheng Yin
2026-08-30 22:03:01 -07:00
committed by GitHub
parent 9a9e167179
commit 5d12ad4fd7
30 changed files with 375 additions and 410 deletions
@@ -7,6 +7,7 @@ import unittest
from collections import deque
from types import SimpleNamespace
from sglang.srt.managers.schedule_batch import ReqKvInfo
from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
@@ -708,7 +709,7 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
req = FakeRequest()
runner._req_to_token_pool.alloc([req])
runner._req_to_token_pool.auxiliary_state_pool.store_cache(
req.mamba_pool_idx,
req.kv.mamba_pool_idx,
[FakeNativeCache(mx.array([42.0], dtype=mx.float32)), None],
[0],
)
@@ -731,7 +732,7 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
self.assertIsInstance(pending.cache[1], ContiguousAttentionKVCache)
restored = [FakeNativeCache(), None]
runner._req_to_token_pool.auxiliary_state_pool.restore_cache(
req.mamba_pool_idx, restored, [0]
req.kv.mamba_pool_idx, restored, [0]
)
self.assertEqual(restored[0].state[0].tolist(), [1.0])
@@ -782,11 +783,11 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
runner.prefill_finalize(pending)
tracked = [FakeNativeCache(), None]
runner._req_to_token_pool.auxiliary_state_pool.restore_cache(
req.mamba_ping_pong_track_buffer[0], tracked, [0]
req.kv.mamba_ping_pong_track_buffer[0], tracked, [0]
)
self.assertEqual([len(x[0]) for x in runner.model.seen_inputs], [64, 6])
self.assertEqual(req.mamba_last_track_seqlen, 64)
self.assertEqual(req.kv.mamba_last_track_seqlen, 64)
self.assertEqual(tracked[0].state[0].tolist(), [64.0])
self.assertEqual(pending.synced_offset, 70)
@@ -825,7 +826,7 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
req = FakeRequest()
runner._req_to_token_pool.alloc([req])
runner._req_to_token_pool.auxiliary_state_pool.store_cache(
req.mamba_pool_idx,
req.kv.mamba_pool_idx,
[FakeNativeCache(mx.array([64.0], dtype=mx.float32)), None],
[0],
)
@@ -844,12 +845,12 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
runner.prefill_finalize(pending)
tracked = [FakeNativeCache(), None]
runner._req_to_token_pool.auxiliary_state_pool.restore_cache(
req.mamba_ping_pong_track_buffer[0], tracked, [0]
req.kv.mamba_ping_pong_track_buffer[0], tracked, [0]
)
self.assertEqual([len(x[0]) for x in runner.model.seen_inputs], [192, 1])
self.assertEqual(runner.model.seen_auxiliary_states, [[64.0], [192.0]])
self.assertEqual(req.mamba_last_track_seqlen, 256)
self.assertEqual(req.kv.mamba_last_track_seqlen, 256)
self.assertEqual(tracked[0].state[0].tolist(), [192.0])
self.assertEqual(pending.synced_offset, 257)
@@ -926,11 +927,11 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
self.assertIn(req_indices[0], range(1, pool.size + 1))
self.assertIsNotNone(auxiliary_state_idx)
self.assertIsNone(req.kv.req_pool_idx)
self.assertIsNotNone(req.mamba_pool_idx)
self.assertIsNotNone(req.kv.mamba_pool_idx)
self.assertIs(pool.mamba_allocator, pool.mamba_pool)
self.assertEqual(pool.auxiliary_state_pool.available_size(), 3)
pool.free_auxiliary_state_cache(req)
self.assertIsNone(req.mamba_pool_idx)
self.assertIsNone(req.kv.mamba_pool_idx)
self.assertEqual(pool.available_size(), 2)
self.assertEqual(pool.auxiliary_state_pool.available_size(), 4)
@@ -944,13 +945,13 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
)
req = FakeRequest()
pool.alloc([req])
req.mamba_ping_pong_track_buffer = pool.auxiliary_state_pool.alloc(1)
req.mamba_next_track_idx = 0
req.kv.mamba_ping_pong_track_buffer = pool.auxiliary_state_pool.alloc(1)
req.kv.mamba_next_track_idx = 0
pool.free_auxiliary_state_cache(req, track_buffer_to_keep=0)
self.assertIsNone(req.mamba_pool_idx)
self.assertIsNone(req.mamba_ping_pong_track_buffer)
self.assertIsNone(req.kv.mamba_pool_idx)
self.assertIsNone(req.kv.mamba_ping_pong_track_buffer)
self.assertEqual(pool.auxiliary_state_pool.available_size(), 3)
def test_auxiliary_state_component_inserts_tracked_slot_and_frees_live_slot(self):
@@ -963,9 +964,9 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
)
req = FakeRequest()
pool.alloc([req])
req.mamba_ping_pong_track_buffer = pool.auxiliary_state_pool.alloc(1)
req.mamba_next_track_idx = 0
req.mamba_last_track_seqlen = 64
req.kv.mamba_ping_pong_track_buffer = pool.auxiliary_state_pool.alloc(1)
req.kv.mamba_next_track_idx = 0
req.kv.mamba_last_track_seqlen = 64
component = MlxAuxiliaryStateComponent(
SimpleNamespace(req_to_token_pool=pool),
SimpleNamespace(enable_mamba_extra_buffer=False),
@@ -988,9 +989,9 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
self.assertEqual(cache_len, 64)
self.assertTrue(getattr(insert_params, "mlx_auxiliary_state_uses_track_slot"))
self.assertEqual(insert_params.mamba_value.tolist(), [2])
self.assertIsNone(req.mamba_pool_idx)
self.assertIsNone(req.mamba_ping_pong_track_buffer)
self.assertIsNone(req.mamba_last_track_seqlen)
self.assertIsNone(req.kv.mamba_pool_idx)
self.assertIsNone(req.kv.mamba_ping_pong_track_buffer)
self.assertIsNone(req.kv.mamba_last_track_seqlen)
self.assertEqual(pool.auxiliary_state_pool.available_size(), 3)
def test_auxiliary_state_component_unfinished_frees_tracked_source_slot(self):
@@ -1003,9 +1004,9 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
)
req = FakeRequest()
pool.alloc([req])
req.mamba_ping_pong_track_buffer = pool.auxiliary_state_pool.alloc(1)
req.mamba_next_track_idx = 0
req.mamba_last_track_seqlen = 64
req.kv.mamba_ping_pong_track_buffer = pool.auxiliary_state_pool.alloc(1)
req.kv.mamba_next_track_idx = 0
req.kv.mamba_last_track_seqlen = 64
component = MlxAuxiliaryStateComponent(
SimpleNamespace(req_to_token_pool=pool),
SimpleNamespace(enable_mamba_extra_buffer=False),
@@ -1027,9 +1028,9 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
self.assertEqual(cache_len, 64)
self.assertEqual(insert_params.mamba_value.tolist(), [3])
self.assertIsNotNone(req.mamba_pool_idx)
self.assertIsNone(req.mamba_ping_pong_track_buffer)
self.assertIsNone(req.mamba_last_track_seqlen)
self.assertIsNotNone(req.kv.mamba_pool_idx)
self.assertIsNone(req.kv.mamba_ping_pong_track_buffer)
self.assertIsNone(req.kv.mamba_last_track_seqlen)
self.assertEqual(pool.auxiliary_state_pool.available_size(), 2)
def test_auxiliary_state_component_frees_stale_track_slot_when_live_slot_inserted(
@@ -1044,8 +1045,8 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
)
req = FakeRequest()
pool.alloc([req])
req.mamba_ping_pong_track_buffer = pool.auxiliary_state_pool.alloc(1)
req.mamba_next_track_idx = 0
req.kv.mamba_ping_pong_track_buffer = pool.auxiliary_state_pool.alloc(1)
req.kv.mamba_next_track_idx = 0
component = MlxAuxiliaryStateComponent(
SimpleNamespace(req_to_token_pool=pool),
SimpleNamespace(enable_mamba_extra_buffer=False),
@@ -1068,9 +1069,9 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
self.assertEqual(cache_len, 7)
self.assertFalse(getattr(insert_params, "mlx_auxiliary_state_uses_track_slot"))
self.assertEqual(insert_params.mamba_value.tolist(), [1])
self.assertIsNone(req.mamba_pool_idx)
self.assertIsNone(req.mamba_ping_pong_track_buffer)
self.assertIsNone(req.mamba_next_track_idx)
self.assertIsNone(req.kv.mamba_pool_idx)
self.assertIsNone(req.kv.mamba_ping_pong_track_buffer)
self.assertIsNone(req.kv.mamba_next_track_idx)
self.assertEqual(pool.auxiliary_state_pool.available_size(), 3)
def test_auxiliary_state_component_frees_duplicate_live_slot(self):
@@ -1102,7 +1103,7 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase):
insert_params=insert_params,
)
self.assertIsNone(req.mamba_pool_idx)
self.assertIsNone(req.kv.mamba_pool_idx)
self.assertEqual(pool.auxiliary_state_pool.available_size(), 4)
@@ -1518,8 +1519,7 @@ if _HAS_MLX:
class FakeRequest:
def __init__(self):
self.kv = SimpleNamespace(req_pool_idx=None)
self.mamba_pool_idx = None
self.kv = ReqKvInfo()
self.inflight_middle_chunks = 0
class FakeTpWorker:
@@ -21,6 +21,7 @@ import unittest
from types import SimpleNamespace
from unittest import mock
from sglang.srt.managers.schedule_batch import ReqKvInfo
from sglang.srt.runtime_context import get_context
from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci
from sglang.test.test_utils import CustomTestCase
@@ -120,9 +121,7 @@ def _hybrid_stub_for_initialize(
def _fake_req():
return SimpleNamespace(
inflight_middle_chunks=0,
kv=SimpleNamespace(req_pool_idx=None),
mamba_pool_idx=None,
mamba_ping_pong_track_buffer=None,
kv=ReqKvInfo(),
)
@@ -284,13 +283,13 @@ class TestMlxHybridInitializeAllocation(CustomTestCase):
req = _fake_req()
self.assertIsNotNone(pool.alloc([req]))
pool.free(req) # as release_kv_cache does after ChunkCache
self.assertIsNone(req.mamba_pool_idx)
self.assertIsNone(req.kv.mamba_pool_idx)
self.assertEqual(pool.auxiliary_state_pool.available_size(), aux_capacity)
def test_radix_enabled_free_does_not_touch_aux_slot(self):
# Retention contract: with the radix cache enabled the tree component
# owns auxiliary release (it frees or adopts the slot and nulls
# req.mamba_pool_idx BEFORE the row is freed). pool.free(req) must
# req.kv.mamba_pool_idx BEFORE the row is freed). pool.free(req) must
# therefore never release auxiliary slots itself -- even if called
# while mamba_pool_idx is still set -- or a tree-owned snapshot slot
# could be recycled under a live radix node.
@@ -306,7 +305,7 @@ class TestMlxHybridInitializeAllocation(CustomTestCase):
req = _fake_req()
pool.alloc([req])
pool.free(req)
self.assertIsNotNone(req.mamba_pool_idx) # slot NOT released by free()
self.assertIsNotNone(req.kv.mamba_pool_idx) # slot NOT released by free()
self.assertEqual(pool.auxiliary_state_pool.available_size(), free_before - 1)
def test_default_aux_sizing_uses_shared_ratio(self):
@@ -38,6 +38,7 @@ from types import SimpleNamespace
import torch
from sglang.srt.managers.schedule_batch import ReqKvInfo
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.runtime_context import get_context
from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci
@@ -169,7 +170,7 @@ class _FakeReq:
self.rid = rid
self.prefix_indices = torch.empty(0, dtype=torch.long)
self.fill_ids = [0]
self.kv = SimpleNamespace(req_pool_idx=req_pool_idx)
self.kv = ReqKvInfo(req_pool_idx=req_pool_idx)
# Mirrors Req's chunk-finality contract read by
# MlxTpModelWorker._chunk_needs_logits: extend_range=None means
# "not truncated" (final chunk / plain prefill).
@@ -45,8 +45,8 @@ def _track_seqlen(*, tree_page: int, prefix_len: int, extend_len: int) -> int:
)
req.prefix_indices = torch.arange(prefix_len, dtype=torch.int64)
req.set_extend_range(prefix_len, prefix_len + extend_len)
req.mamba_ping_pong_track_buffer = torch.tensor([0, 1], dtype=torch.int64)
req.mamba_next_track_idx = 0
req.kv.mamba_ping_pong_track_buffer = torch.tensor([0, 1], dtype=torch.int64)
req.kv.mamba_next_track_idx = 0
req.mamba_branching_seqlen = None
batch = ScheduleBatch(reqs=[req])
@@ -58,7 +58,7 @@ def _track_seqlen(*, tree_page: int, prefix_len: int, extend_len: int) -> int:
batch.req_to_token_pool.get_mamba_ping_pong_other_idx.return_value = 1
batch._mamba_radix_cache_v2_req_prepare_for_extend(req)
return req.mamba_last_track_seqlen
return req.kv.mamba_last_track_seqlen
class TestMambaCheckpointDepth(unittest.TestCase):
@@ -47,7 +47,7 @@ def _make_req(
req.positional_embed_overrides = None
req.extra_key = None
req.cache_salt = None
req.mamba_pool_idx = None
req.kv.mamba_pool_idx = None
req.sampling_params = SimpleNamespace(max_new_tokens=128, ignore_eos=False)
return req
@@ -142,7 +142,7 @@ class TestMamba(unittest.TestCase):
assert req_to_token_pool.mamba_allocator.available_size() == mamba_cache_size
# alloc req without free mamba cache
req.mamba_pool_idx = None
req.kv.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
@@ -233,7 +233,7 @@ class TestMamba(unittest.TestCase):
InsertParams(
key=key,
value=req1_kv_indices[: len(key)],
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
mamba_value=req1.kv.mamba_pool_idx.unsqueeze(0),
)
)
prefix_len = result.prefix_len
@@ -251,7 +251,7 @@ class TestMamba(unittest.TestCase):
InsertParams(
key=key,
value=req2_kv_indices[: len(key)],
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
mamba_value=req2.kv.mamba_pool_idx.unsqueeze(0),
)
)
prefix_len = result.prefix_len
@@ -270,7 +270,7 @@ class TestMamba(unittest.TestCase):
InsertParams(
key=key,
value=req3_kv_indices[: len(key)],
mamba_value=req3.mamba_pool_idx.unsqueeze(0),
mamba_value=req3.kv.mamba_pool_idx.unsqueeze(0),
)
)
prefix_len = result.prefix_len
@@ -288,7 +288,7 @@ class TestMamba(unittest.TestCase):
InsertParams(
key=key,
value=req4_kv_indices[: len(key)],
mamba_value=req4.mamba_pool_idx.unsqueeze(0),
mamba_value=req4.kv.mamba_pool_idx.unsqueeze(0),
)
)
prefix_len = result.prefix_len
@@ -372,13 +372,13 @@ class TestMamba(unittest.TestCase):
)
)
kv_indices, last_node = result.device_indices, result.last_device_node
assert req9.mamba_pool_idx is not None
assert req9.kv.holds_mamba
assert torch.all(
mamba_pool.mamba_cache.conv[0][:, req9.mamba_pool_idx]
mamba_pool.mamba_cache.conv[0][:, req9.kv.mamba_pool_idx]
== mamba_pool.mamba_cache.conv[0][:, last_node.mamba_value]
)
assert torch.all(
mamba_pool.mamba_cache.temporal[:, req9.mamba_pool_idx]
mamba_pool.mamba_cache.temporal[:, req9.kv.mamba_pool_idx]
== mamba_pool.mamba_cache.temporal[:, last_node.mamba_value]
)
@@ -404,7 +404,7 @@ class TestMamba(unittest.TestCase):
InsertParams(
key=RadixKey(array("q", token_ids)),
value=kv,
mamba_value=req.mamba_pool_idx.unsqueeze(0),
mamba_value=req.kv.mamba_pool_idx.unsqueeze(0),
)
)
@@ -459,7 +459,7 @@ class TestMamba(unittest.TestCase):
InsertParams(
key=key1,
value=allocator.alloc(3)[: len(key1)],
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
mamba_value=req1.kv.mamba_pool_idx.unsqueeze(0),
)
)
events = tree.take_events()
@@ -476,7 +476,7 @@ class TestMamba(unittest.TestCase):
InsertParams(
key=key2,
value=allocator.alloc(5)[: len(key2)],
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
mamba_value=req2.kv.mamba_pool_idx.unsqueeze(0),
)
)
events = tree.take_events()
@@ -515,7 +515,7 @@ class TestMamba(unittest.TestCase):
InsertParams(
key=key1,
value=allocator.alloc(4)[: len(key1)],
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
mamba_value=req1.kv.mamba_pool_idx.unsqueeze(0),
)
)
first_insert_events = [
@@ -530,7 +530,7 @@ class TestMamba(unittest.TestCase):
InsertParams(
key=key2,
value=allocator.alloc(4)[: len(key2)],
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
mamba_value=req2.kv.mamba_pool_idx.unsqueeze(0),
)
)
second_insert_events = [
@@ -771,7 +771,7 @@ class TestMamba(unittest.TestCase):
InsertParams(
key=key1,
value=allocator.alloc(3)[: len(key1)],
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
mamba_value=req1.kv.mamba_pool_idx.unsqueeze(0),
)
)
assert allocator.available_size() == initial_avail - 3
@@ -784,7 +784,7 @@ class TestMamba(unittest.TestCase):
InsertParams(
key=key2,
value=allocator.alloc(7)[: len(key2)],
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
mamba_value=req2.kv.mamba_pool_idx.unsqueeze(0),
prev_prefix_len=0,
)
)
@@ -802,7 +802,7 @@ class TestMamba(unittest.TestCase):
InsertParams(
key=key3,
value=allocator.alloc(8)[: len(key3)],
mamba_value=req3.mamba_pool_idx.unsqueeze(0),
mamba_value=req3.kv.mamba_pool_idx.unsqueeze(0),
prev_prefix_len=2,
)
)
@@ -819,7 +819,7 @@ class TestMamba(unittest.TestCase):
InsertParams(
key=key4,
value=allocator.alloc(9)[: len(key4)],
mamba_value=req4.mamba_pool_idx.unsqueeze(0),
mamba_value=req4.kv.mamba_pool_idx.unsqueeze(0),
prev_prefix_len=8,
)
)
@@ -39,7 +39,6 @@ class _StubReq:
self.best_match_node = None
self.host_hit_length = None
self.num_matched_prefix_tokens = 0
self.mamba_branching_seqlen = None
self.kv = SimpleNamespace(cache_protected_len=None)
def _compute_max_prefix_len(self, input_len):
@@ -42,7 +42,7 @@ def _req_and_pool():
req.kv = ReqKvInfo(req_pool_idx=0)
req.origin_input_ids = [1, 2]
req.output_ids = [3]
req.mamba_pool_idx = torch.tensor(1)
req.kv.mamba_pool_idx = torch.tensor(1)
pool = object.__new__(HybridReqToTokenPool)
pool.req_to_token = torch.zeros(1, 8, dtype=torch.int64)
@@ -59,7 +59,7 @@ class TestRetractionMambaBackup(unittest.TestCase):
allocator = _Allocator(carries_mamba=False)
req.offload_kv_cache(pool, allocator)
self.assertIs(req.retraction_backup.mamba_cpu, MAMBA_STATE)
self.assertIs(req.kv.retraction_backup.mamba_cpu, MAMBA_STATE)
req.load_kv_cache(pool, allocator)
self.assertIs(pool.mamba_pool.loaded, MAMBA_STATE)
@@ -69,7 +69,7 @@ class TestRetractionMambaBackup(unittest.TestCase):
allocator = _Allocator(carries_mamba=True)
req.offload_kv_cache(pool, allocator)
self.assertIsNone(req.retraction_backup.mamba_cpu)
self.assertIsNone(req.kv.retraction_backup.mamba_cpu)
req.load_kv_cache(pool, allocator)
self.assertIsNone(pool.mamba_pool.loaded)
@@ -81,12 +81,6 @@ class _FakeReq:
self.last_node = None
self.swa_uuid_for_lock = None
self.skip_lock_node_ids = {}
self.mamba_pool_idx = None
self.mamba_ping_pong_track_buffer = None
self.mamba_next_track_idx = None
self.mamba_last_track_idx = None
self.mamba_last_track_seqlen = None
self.mamba_branching_seqlen = None
self.to_finish = None
self.finished_reason = None
self.finished_len = None
@@ -97,11 +91,12 @@ class _FakeReq:
def test_session_slot_round_trip_preserves_mamba_state():
# The mamba state rides in the shared ReqKvInfo record. mamba_branching_seqlen
# is a per-turn match observation on the Req and is not preserved by the slot.
req = _FakeReq("session-a", req_pool_idx=0, committed=4, allocated=4)
req.mamba_next_track_idx = 1
req.mamba_last_track_idx = 0
req.mamba_last_track_seqlen = 3
req.mamba_branching_seqlen = 2
req.kv.mamba_next_track_idx = 1
req.kv.mamba_last_track_idx = 0
req.kv.mamba_last_track_seqlen = 3
slot = SessionSlot()
slot.save_from_req(req, is_first=True)
@@ -109,10 +104,9 @@ def test_session_slot_round_trip_preserves_mamba_state():
next_req = _FakeReq("session-a", req_pool_idx=1, committed=0, allocated=0)
slot.restore_to_req(next_req)
assert next_req.mamba_next_track_idx == 1
assert next_req.mamba_last_track_idx == 0
assert next_req.mamba_last_track_seqlen == 3
assert next_req.mamba_branching_seqlen == 2
assert next_req.kv.mamba_next_track_idx == 1
assert next_req.kv.mamba_last_track_idx == 0
assert next_req.kv.mamba_last_track_seqlen == 3
def test_preabort_detaches_session_and_preserves_slot():
@@ -334,7 +334,7 @@ def _insert_seq(env, seq):
mamba_val = None
if env.has_mamba:
req = env.make_req()
mamba_val = req.mamba_pool_idx.unsqueeze(0)
mamba_val = req.kv.mamba_pool_idx.unsqueeze(0)
key = RadixKey(array("q", seq))
env.tree.insert(InsertParams(key=key, value=v[: len(key)], mamba_value=mamba_val))
return True
@@ -356,7 +356,7 @@ def _fill_no_evict(env):
mamba_val = None
if env.has_mamba:
req = env.make_req()
mamba_val = req.mamba_pool_idx.unsqueeze(0)
mamba_val = req.kv.mamba_pool_idx.unsqueeze(0)
key = RadixKey(array("q", seq))
env.tree.insert(
InsertParams(key=key, value=v[: len(key)], mamba_value=mamba_val)
@@ -601,7 +601,7 @@ class TestUnifiedRadixAllocationEvictionRealComponents(CustomTestCase):
sampling_params=SamplingParams(temperature=0, max_new_tokens=1),
)
req_to_token_pool.alloc([req])
params.mamba_value = req.mamba_pool_idx.unsqueeze(0)
params.mamba_value = req.kv.mamba_pool_idx.unsqueeze(0)
cache.insert(params)
def _build_internal_chain(self, component_type, enable_session_radix_cache):
@@ -1075,7 +1075,7 @@ class UnifiedRadixCacheSuite:
params = InsertParams(key=key, value=value[: len(key)], priority=priority)
if self.cfg.has_mamba:
req = self._make_req(req_to_token_pool)
params.mamba_value = req.mamba_pool_idx.unsqueeze(0)
params.mamba_value = req.kv.mamba_pool_idx.unsqueeze(0)
return cache.insert(params)
def test_insert_and_match_basic(self):
@@ -1189,7 +1189,7 @@ class UnifiedRadixCacheSuite:
)
if self.cfg.has_mamba:
req = self._make_req(req_to_token_pool)
params.mamba_value = req.mamba_pool_idx.unsqueeze(0)
params.mamba_value = req.kv.mamba_pool_idx.unsqueeze(0)
result = cache.insert(params)
self.assertEqual(result.prefix_len, len(seq_1p))
self.assertEqual(
@@ -1213,7 +1213,7 @@ class UnifiedRadixCacheSuite:
)
if self.cfg.has_mamba:
req = self._make_req(req_to_token_pool)
params.mamba_value = req.mamba_pool_idx.unsqueeze(0)
params.mamba_value = req.kv.mamba_pool_idx.unsqueeze(0)
result = cache.insert(params)
self.assertEqual(result.prefix_len, len(seq_2p))
self.assertEqual(
@@ -1246,7 +1246,7 @@ class UnifiedRadixCacheSuite:
len(req.prefix_indices), len(req.full_untruncated_fill_ids)
)
if self.cfg.has_mamba:
req.mamba_last_track_seqlen = kv_len
req.kv.mamba_last_track_seqlen = kv_len
cache.cache_finished_req(
req, is_insert=True, kv_len_to_handle=req.effective_kv_committed_len()
@@ -1283,7 +1283,7 @@ class UnifiedRadixCacheSuite:
req.swa_uuid_for_lock = None
req.extra_key = None
if self.cfg.has_mamba:
req.mamba_last_track_seqlen = kv_len
req.kv.mamba_last_track_seqlen = kv_len
req.reasoning_tokens = 1
# cache_finished_req reads get_serving().strip_thinking_cache
@@ -1362,7 +1362,7 @@ class UnifiedRadixCacheSuite:
req.swa_uuid_for_lock = None
req.extra_key = None
if self.cfg.has_mamba:
req.mamba_last_track_seqlen = kv_len
req.kv.mamba_last_track_seqlen = kv_len
cache.cache_unfinished_req(req)
@@ -1507,7 +1507,7 @@ class UnifiedRadixCacheSuite:
len(req.prefix_indices), len(req.full_untruncated_fill_ids)
)
if self.cfg.has_mamba:
req.mamba_last_track_seqlen = kv_len
req.kv.mamba_last_track_seqlen = kv_len
avail_before = allocator.available_size()
cache.cache_finished_req(
@@ -1577,12 +1577,12 @@ class UnifiedRadixCacheSuite:
MatchPrefixParams(key=RadixKey(array("q", seq)), cow_mamba=True, req=req2)
)
self.assertEqual(len(m.device_indices), len(seq))
self.assertIsNotNone(req2.mamba_pool_idx)
self.assertIsNotNone(req2.kv.mamba_pool_idx)
src_value = _device_value(cache, m.last_device_node, ComponentType.MAMBA)
self.assertTrue(
torch.all(
mamba_pool.mamba_cache.conv[0][:, req2.mamba_pool_idx]
mamba_pool.mamba_cache.conv[0][:, req2.kv.mamba_pool_idx]
== mamba_pool.mamba_cache.conv[0][:, src_value]
)
)
@@ -5469,7 +5469,7 @@ class UnifiedRadixCacheSuite:
# Simulate a request without its own mamba slot so load-back allocates one
# (that allocation is what a called-off load-back must free + not publish).
req.mamba_pool_idx = None
req.kv.mamba_pool_idx = None
avail_before = req_to_token_pool.mamba_allocator.available_size()
new_indices, new_node = cache.init_load_back(
InitLoadBackParams(
@@ -5485,7 +5485,7 @@ class UnifiedRadixCacheSuite:
self.assertIsNone(_device_value(cache, leaf, ComponentType.FULL))
self.assertIsNone(_device_value(cache, leaf, ComponentType.MAMBA))
# A failed load-back must roll back the pre-allocated mamba slot.
self.assertIsNone(req.mamba_pool_idx)
self.assertIsNone(req.kv.mamba_pool_idx)
self.assertEqual(
req_to_token_pool.mamba_allocator.available_size(), avail_before
)
@@ -5510,7 +5510,7 @@ class UnifiedRadixCacheSuite:
self._apply_match_to_req(req, match)
# Simulate a request without its own mamba slot so load-back allocates one.
req.mamba_pool_idx = None
req.kv.mamba_pool_idx = None
avail_before = req_to_token_pool.mamba_allocator.available_size()
# H->D load fails after the mamba slot is pre-allocated -> must free it.
with mock.patch.object(cache.cache_controller, "load", return_value=None):
@@ -5524,7 +5524,7 @@ class UnifiedRadixCacheSuite:
)
self.assertEqual(len(new_indices), 0)
self.assertIsNone(req.mamba_pool_idx)
self.assertIsNone(req.kv.mamba_pool_idx)
self.assertEqual(
req_to_token_pool.mamba_allocator.available_size(), avail_before
)
@@ -5549,7 +5549,7 @@ class UnifiedRadixCacheSuite:
self._apply_match_to_req(req, match)
# Simulate a request without its own mamba slot so load-back allocates one.
req.mamba_pool_idx = None
req.kv.mamba_pool_idx = None
avail_before = req_to_token_pool.mamba_allocator.available_size()
# No device room and eviction frees nothing -> load-back bails after the
# mamba pre-alloc, which must still be freed.
@@ -5571,7 +5571,7 @@ class UnifiedRadixCacheSuite:
)
self.assertEqual(len(new_indices), 0)
self.assertIsNone(req.mamba_pool_idx)
self.assertIsNone(req.kv.mamba_pool_idx)
self.assertEqual(
req_to_token_pool.mamba_allocator.available_size(), avail_before
)
@@ -5592,8 +5592,8 @@ class UnifiedRadixCacheSuite:
# A request whose mamba slot was released: load_back's CoW arm allocates one.
req = self._make_req(req_to_token_pool)
req_to_token_pool.mamba_allocator.free(req.mamba_pool_idx.unsqueeze(0))
req.mamba_pool_idx = None
req_to_token_pool.mamba_allocator.free(req.kv.mamba_pool_idx.unsqueeze(0))
req.kv.mamba_pool_idx = None
mamba_avail = req_to_token_pool.mamba_allocator.available_size()
# Impossible quota -> load_back aborts after building the transfers.
@@ -5601,7 +5601,7 @@ class UnifiedRadixCacheSuite:
self.assertFalse(loaded)
# the aborted call must return its slot and not leave req pointing at it
self.assertIsNone(req.mamba_pool_idx)
self.assertIsNone(req.kv.mamba_pool_idx)
self.assertEqual(
req_to_token_pool.mamba_allocator.available_size(), mamba_avail
)
@@ -5621,8 +5621,8 @@ class UnifiedRadixCacheSuite:
self.assertTrue(cache.tree_core.is_full_device_evicted(leaf))
req = self._make_req(req_to_token_pool)
req_to_token_pool.mamba_allocator.free(req.mamba_pool_idx.unsqueeze(0))
req.mamba_pool_idx = None
req_to_token_pool.mamba_allocator.free(req.kv.mamba_pool_idx.unsqueeze(0))
req.kv.mamba_pool_idx = None
mamba_avail = req_to_token_pool.mamba_allocator.available_size()
# cache_controller.load() failing (device alloc / transfer resolution)
@@ -5631,7 +5631,7 @@ class UnifiedRadixCacheSuite:
loaded = cache.load_back(leaf, req=req)
self.assertFalse(loaded)
self.assertIsNone(req.mamba_pool_idx)
self.assertIsNone(req.kv.mamba_pool_idx)
self.assertEqual(
req_to_token_pool.mamba_allocator.available_size(), mamba_avail
)
@@ -5652,14 +5652,14 @@ class UnifiedRadixCacheSuite:
# The request already owns its slot: an aborted load-back must not free it.
req = self._make_req(req_to_token_pool)
preexisting_slot = req.mamba_pool_idx
preexisting_slot = req.kv.mamba_pool_idx
self.assertIsNotNone(preexisting_slot)
mamba_avail = req_to_token_pool.mamba_allocator.available_size()
loaded = cache.load_back(leaf, mem_quota=-(10**9), req=req)
self.assertFalse(loaded)
self.assertIs(req.mamba_pool_idx, preexisting_slot)
self.assertIs(req.kv.mamba_pool_idx, preexisting_slot)
self.assertEqual(
req_to_token_pool.mamba_allocator.available_size(), mamba_avail
)
@@ -5679,15 +5679,15 @@ class UnifiedRadixCacheSuite:
self.assertTrue(cache.tree_core.is_full_device_evicted(leaf))
req = self._make_req(req_to_token_pool)
req_to_token_pool.mamba_allocator.free(req.mamba_pool_idx.unsqueeze(0))
req.mamba_pool_idx = None
req_to_token_pool.mamba_allocator.free(req.kv.mamba_pool_idx.unsqueeze(0))
req.kv.mamba_pool_idx = None
mamba_avail = req_to_token_pool.mamba_allocator.available_size()
loaded = cache.load_back(leaf, req=req)
self.assertTrue(loaded)
# the successful load must keep the freshly allocated slot published
self.assertIsNotNone(req.mamba_pool_idx)
self.assertIsNotNone(req.kv.mamba_pool_idx)
self.assertIsNotNone(_device_value(cache, leaf, ComponentType.MAMBA))
# one slot restores the node's mamba value, one is the request's CoW slot
self.assertEqual(
@@ -5717,17 +5717,17 @@ class UnifiedRadixCacheSuite:
self.assertTrue(cache.tree_core.is_full_device_evicted(leaf))
req = self._make_req(req_to_token_pool)
req_to_token_pool.mamba_allocator.free(req.mamba_pool_idx.unsqueeze(0))
req.mamba_pool_idx = None
req_to_token_pool.mamba_allocator.free(req.kv.mamba_pool_idx.unsqueeze(0))
req.kv.mamba_pool_idx = None
loaded = cache.load_back(leaf, req=req)
self.assertTrue(loaded)
self.assertIsNotNone(req.mamba_pool_idx)
self.assertIsNotNone(req.kv.mamba_pool_idx)
self._finish_pending_loads(cache)
# The CoW slot must actually hold the backed-up mamba state, not merely exist.
actual_temporal, actual_conv = self._snapshot_mamba_state(
req_to_token_pool, req.mamba_pool_idx.unsqueeze(0)
req_to_token_pool, req.kv.mamba_pool_idx.unsqueeze(0)
)
self.assertTrue(torch.equal(actual_temporal, expected_temporal))
self.assertEqual(len(actual_conv), len(expected_conv))
@@ -5747,12 +5747,12 @@ class UnifiedRadixCacheSuite:
self._backup_node(cache, leaf)
req = self._make_req(req_to_token_pool)
req_to_token_pool.mamba_allocator.free(req.mamba_pool_idx.unsqueeze(0))
req.mamba_pool_idx = None
req_to_token_pool.mamba_allocator.free(req.kv.mamba_pool_idx.unsqueeze(0))
req.kv.mamba_pool_idx = None
# device value still present -> nothing to prepare even though host-backed
self.assertIsNone(comp.prepare_load_back(leaf, req=req).allocated_mamba_slot)
self.assertIsNone(req.mamba_pool_idx)
self.assertIsNone(req.kv.mamba_pool_idx)
cache.evict(EvictParams(num_tokens=_node_key_length(cache, leaf)))
self.assertTrue(cache.tree_core.is_full_device_evicted(leaf))
@@ -5769,16 +5769,16 @@ class UnifiedRadixCacheSuite:
# fresh request + host-only mamba -> allocates and publishes onto req
prep = comp.prepare_load_back(leaf, req=req)
self.assertIsNotNone(prep.allocated_mamba_slot)
self.assertEqual(int(req.mamba_pool_idx), int(prep.allocated_mamba_slot[0]))
self.assertEqual(int(req.kv.mamba_pool_idx), int(prep.allocated_mamba_slot[0]))
# node without host-backed mamba -> nothing to prepare
req2 = self._make_req(req_to_token_pool)
req_to_token_pool.mamba_allocator.free(req2.mamba_pool_idx.unsqueeze(0))
req2.mamba_pool_idx = None
req_to_token_pool.mamba_allocator.free(req2.kv.mamba_pool_idx.unsqueeze(0))
req2.kv.mamba_pool_idx = None
root = cache.root_node_handle()
self.assertIsNone(_host_value(cache, root, ComponentType.MAMBA))
self.assertIsNone(comp.prepare_load_back(root, req=req2).allocated_mamba_slot)
self.assertIsNone(req2.mamba_pool_idx)
self.assertIsNone(req2.kv.mamba_pool_idx)
def test_prepare_load_back_skips_device_present_node(self):
if not self.cfg.has_mamba or self.cfg.has_swa or self.cfg.page_size != 1:
@@ -5796,12 +5796,12 @@ class UnifiedRadixCacheSuite:
self.assertIsNotNone(_host_value(cache, leaf, ComponentType.MAMBA))
req = self._make_req(req_to_token_pool)
req_to_token_pool.mamba_allocator.free(req.mamba_pool_idx.unsqueeze(0))
req.mamba_pool_idx = None
req_to_token_pool.mamba_allocator.free(req.kv.mamba_pool_idx.unsqueeze(0))
req.kv.mamba_pool_idx = None
mamba_avail = req_to_token_pool.mamba_allocator.available_size()
self.assertIsNone(comp.prepare_load_back(leaf, req=req).allocated_mamba_slot)
self.assertIsNone(req.mamba_pool_idx)
self.assertIsNone(req.kv.mamba_pool_idx)
self.assertEqual(
req_to_token_pool.mamba_allocator.available_size(), mamba_avail
)
@@ -5820,8 +5820,8 @@ class UnifiedRadixCacheSuite:
cache.evict(EvictParams(num_tokens=_node_key_length(cache, leaf)))
req = self._make_req(req_to_token_pool)
req_to_token_pool.mamba_allocator.free(req.mamba_pool_idx.unsqueeze(0))
req.mamba_pool_idx = None
req_to_token_pool.mamba_allocator.free(req.kv.mamba_pool_idx.unsqueeze(0))
req.kv.mamba_pool_idx = None
retry_slot = req_to_token_pool.mamba_allocator.alloc(1)
# first alloc fails -> prepare must evict a mamba slot and retry
@@ -5838,7 +5838,7 @@ class UnifiedRadixCacheSuite:
prep = comp.prepare_load_back(leaf, req=req)
evict_for_alloc.assert_called_once_with(EvictParams(num_tokens=0, mamba_num=1))
self.assertIs(prep.allocated_mamba_slot, retry_slot)
self.assertEqual(int(req.mamba_pool_idx), int(retry_slot[0]))
self.assertEqual(int(req.kv.mamba_pool_idx), int(retry_slot[0]))
def test_hicache_swa_load_back_min_suffix(self):
"""LOAD_BACK collects only the suffix nodes needed to cover sliding_window_size."""
@@ -6694,7 +6694,7 @@ class TestUnifiedMambaLRUMatchRefresh(CustomTestCase):
InsertParams(
key=RadixKey(array("q", tokens)),
value=value[: len(tokens)],
mamba_value=req.mamba_pool_idx.unsqueeze(0),
mamba_value=req.kv.mamba_pool_idx.unsqueeze(0),
)
)
@@ -6785,7 +6785,7 @@ class TestUnifiedRadixCacheInt8MambaCheckpoint(CustomTestCase):
req.kv.cache_protected_len = 0
req.swa_uuid_for_lock = None
req.extra_key = None
req.mamba_last_track_seqlen = len(tokens)
req.kv.mamba_last_track_seqlen = len(tokens)
return req
def _cache_finished(self, cache, allocator, req_to_token_pool, tokens):