From 0299393758a435c54c79024797e977a5873a6d3c Mon Sep 17 00:00:00 2001 From: Zhangheng Date: Fri, 10 Jul 2026 23:46:38 +0800 Subject: [PATCH] [UnifiedTree]: Sync mamba int8 checkpoint (#30626) --- .../mamba_component.py | 95 +++++++++++++++---- .../test_int8_mamba_checkpoint_e2e.py | 36 ++++++- .../test_unified_radix_cache_unittest.py | 74 ++++++++++++++- 3 files changed, 181 insertions(+), 24 deletions(-) diff --git a/python/sglang/srt/mem_cache/unified_cache_components/mamba_component.py b/python/sglang/srt/mem_cache/unified_cache_components/mamba_component.py index 1828f2974..edfbaa052 100644 --- a/python/sglang/srt/mem_cache/unified_cache_components/mamba_component.py +++ b/python/sglang/srt/mem_cache/unified_cache_components/mamba_component.py @@ -174,7 +174,7 @@ class MambaComponent(TreeComponent): # Device layer if EvictLayer.DEVICE in target and cd.value is not None: - self.cache.req_to_token_pool.mamba_allocator.free(cd.value) + self._free_mamba_value(cd.value) freed = len(cd.value) self.cache.component_evictable_size_[self.component_type] -= freed cd.value = None @@ -294,6 +294,33 @@ class MambaComponent(TreeComponent): assert slot is not None, "Can not alloc mamba cache" return slot + @property + def int8_ckpt_pool(self): + return getattr(self.cache.req_to_token_pool, "mamba_ckpt_pool", None) + + def _alloc_int8_ckpt_slot(self) -> torch.Tensor: + slot = self.int8_ckpt_pool.alloc(1) + if slot is None: + self.cache.evict(EvictParams(num_tokens=0, mamba_num=1)) + slot = self.int8_ckpt_pool.alloc(1) + assert slot is not None, "Can not alloc int8 mamba checkpoint slot" + return slot + + def _commit_int8_checkpoint(self, active_slots: torch.Tensor) -> torch.Tensor: + ckpt_slot = self._alloc_int8_ckpt_slot() + self.int8_ckpt_pool.store_from_active( + self.cache.req_to_token_pool.mamba_pool, + active_slots.view(-1), + ckpt_slot, + ) + return ckpt_slot + + def _free_mamba_value(self, mamba_value: torch.Tensor) -> None: + if self.int8_ckpt_pool is not None: + self.int8_ckpt_pool.free(mamba_value) + else: + self.cache.req_to_token_pool.mamba_allocator.free(mamba_value) + def prepare_for_caching_req( self, req: Req, @@ -326,18 +353,35 @@ class MambaComponent(TreeComponent): keep_idx = self.cache.req_to_token_pool.get_mamba_ping_pong_keep_idx( req ) - mamba_value = ( + active_value = ( req.mamba_ping_pong_track_buffer[keep_idx].unsqueeze(-1).clone() ) else: - mamba_value = req.mamba_pool_idx.unsqueeze(-1).clone() - insert_params.mamba_value = mamba_value + active_value = req.mamba_pool_idx.unsqueeze(-1).clone() + if self.int8_ckpt_pool is not None: + insert_params.mamba_value = self._commit_int8_checkpoint(active_value) + else: + insert_params.mamba_value = active_value return cache_len else: if cache_len is None: return 0 # Donate the mamba index to the radix cache instead of copying. - if self.enable_mamba_extra_buffer: + if self.int8_ckpt_pool is not None: + if self.enable_mamba_extra_buffer: + new_slot = self._alloc_mamba_slot() + src_active = ( + self.cache.req_to_token_pool.donate_mamba_ping_pong_slot( + req, new_slot + ) + ) + mamba_value_donated = self._commit_int8_checkpoint(src_active) + self.cache.req_to_token_pool.mamba_allocator.free(src_active) + else: + mamba_value_donated = self._commit_int8_checkpoint( + req.mamba_pool_idx.view(-1) + ) + elif self.enable_mamba_extra_buffer: new_slot = self._alloc_mamba_slot() mamba_value_donated = ( self.cache.req_to_token_pool.donate_mamba_ping_pong_slot( @@ -364,29 +408,40 @@ class MambaComponent(TreeComponent): insert_params: Optional[InsertParams] = None, ) -> None: if is_finished: - mamba_exist = ( - insert_result.mamba_exist if insert_result is not None else True + mamba_value_inserted = ( + insert_result is not None and not insert_result.mamba_exist ) - if self.enable_mamba_extra_buffer: - keep_idx = self.cache.req_to_token_pool.get_mamba_ping_pong_keep_idx( - req + pool = self.cache.req_to_token_pool + + if self.int8_ckpt_pool is not None: + insert_value_unused = ( + not mamba_value_inserted + and insert_params is not None + and insert_params.mamba_value is not None ) - else: - keep_idx = None - if mamba_exist: - keep_idx = None - free_mamba_cache = True if self.enable_mamba_extra_buffer else mamba_exist - if free_mamba_cache: - self.cache.req_to_token_pool.free_mamba_cache( + if insert_value_unused: + self._free_mamba_value(insert_params.mamba_value) + pool.free_mamba_cache(req) + return + + if self.enable_mamba_extra_buffer: + keep_idx = ( + pool.get_mamba_ping_pong_keep_idx(req) + if mamba_value_inserted + else None + ) + pool.free_mamba_cache( req, mamba_ping_pong_track_buffer_to_keep=keep_idx ) + return + + if not mamba_value_inserted: + pool.free_mamba_cache(req) else: if insert_params.mamba_value is not None and ( insert_result is None or insert_result.mamba_exist ): - self.cache.req_to_token_pool.mamba_allocator.free( - insert_params.mamba_value - ) + self._free_mamba_value(insert_params.mamba_value) req.mamba_last_track_seqlen = None # ---- HiCache Hooks ---- diff --git a/test/registered/radix_cache/test_int8_mamba_checkpoint_e2e.py b/test/registered/radix_cache/test_int8_mamba_checkpoint_e2e.py index 06112611d..ea1a6a973 100644 --- a/test/registered/radix_cache/test_int8_mamba_checkpoint_e2e.py +++ b/test/registered/radix_cache/test_int8_mamba_checkpoint_e2e.py @@ -21,16 +21,24 @@ Usage: python3 -m unittest test_int8_mamba_checkpoint_e2e """ +import time import unittest from types import SimpleNamespace from urllib.parse import urlparse +from sglang.srt.utils import kill_process_tree from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.kits.kl_divergence_kit import KLDivergenceMixin -from sglang.test.server_fixtures.default_fixture import DefaultServerBase -from sglang.test.test_utils import DEFAULT_HYBRID_MAMBA_MODEL_NAME_FOR_TEST +from sglang.test.server_fixtures.default_fixture import ( + DefaultServerBase, + openai_api_env, +) +from sglang.test.test_utils import ( + DEFAULT_HYBRID_MAMBA_MODEL_NAME_FOR_TEST, + popen_launch_server, +) -register_cuda_ci(est_time=400, stage="extra-b", runner_config="4-gpu-h100") +register_cuda_ci(est_time=800, stage="extra-b", runner_config="4-gpu-h100") class TestInt8MambaCheckpointE2E(KLDivergenceMixin, DefaultServerBase): @@ -89,5 +97,27 @@ class TestInt8MambaCheckpointE2E(KLDivergenceMixin, DefaultServerBase): self.assertGreaterEqual(metrics["accuracy"], self.gsm8k_threshold) +class TestUnifiedRadixTreeInt8MambaCheckpointE2E(TestInt8MambaCheckpointE2E): + """Run the same int8 mamba checkpoint checks with UnifiedRadixTree forced on.""" + + @classmethod + def setUpClass(cls): + assert cls.model is not None, "Please set cls.model in subclass" + + with openai_api_env(cls.api_key): + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=cls.timeout, + other_args=cls.other_args, + env={"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1"}, + ) + + @classmethod + def tearDownClass(cls): + kill_process_tree(cls.process.pid, wait_timeout=60) + time.sleep(2) + + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py index 0f724c764..d0d3dff92 100644 --- a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py +++ b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py @@ -91,6 +91,7 @@ class CacheConfig: mamba_head_dim: int = 16 mamba_state_size: int = 16 mamba_conv_kernel: int = 4 + enable_int8_mamba_checkpoint: bool = False # Model / pool kv_size: int = 256 @@ -214,7 +215,11 @@ class TestUnifiedTreeNodeGetPrefixHashValues(CustomTestCase): def build_fixture(cfg: CacheConfig, *, enable_kv_cache_events: bool = False): """Create (tree, allocator, req_to_token_pool) from a CacheConfig.""" - server_args = ServerArgs(model_path="dummy", page_size=cfg.page_size) + server_args = ServerArgs( + model_path="dummy", + page_size=cfg.page_size, + enable_int8_mamba_checkpoint=cfg.enable_int8_mamba_checkpoint, + ) # MambaRadixCache reads mamba_cache_chunk_size, whose property otherwise # loads the HF config for self.model_path — impossible for the dummy model. # Mirror the property's default for a dummy HF config: FLA_CHUNK_SIZE. @@ -4037,6 +4042,73 @@ class UnifiedLRUListBoundedRefreshTest(CustomTestCase): self.assertEqual(self._lru_order(lru), before) +class TestUnifiedRadixCacheInt8MambaCheckpoint(CustomTestCase): + cfg = CacheConfig( + components=(ComponentType.FULL, ComponentType.MAMBA), + enable_mamba_extra_buffer=True, + enable_int8_mamba_checkpoint=True, + mamba_cache_size=8, + kv_size=64, + max_context_len=64, + ) + + def _make_req(self, req_to_token_pool, tokens): + req = Req( + rid="int8-mamba", + origin_input_text="", + origin_input_ids=array("q", tokens), + sampling_params=SamplingParams(temperature=0, max_new_tokens=1), + ) + req_to_token_pool.alloc([req]) + req.output_ids = array("q") + req.kv_committed_len = len(tokens) + req.kv_allocated_len = len(tokens) + req.cache_protected_len = 0 + req.swa_uuid_for_lock = None + req.extra_key = None + req.mamba_last_track_seqlen = len(tokens) + return req + + def _cache_finished(self, cache, allocator, req_to_token_pool, tokens): + req = self._make_req(req_to_token_pool, tokens) + kv_indices = allocator.alloc(len(tokens)) + self.assertIsNotNone(kv_indices) + req_to_token_pool.write((req.req_pool_idx, slice(0, len(tokens))), kv_indices) + req.last_node = cache.root_node + + cache.cache_finished_req(req, is_insert=True) + + def test_finished_req_stores_radix_mamba_state_in_int8_pool(self): + cache, allocator, req_to_token_pool = build_fixture(self.cfg) + ckpt_pool = req_to_token_pool.mamba_ckpt_pool + self.assertIsNotNone(ckpt_pool) + active_initial = req_to_token_pool.mamba_allocator.available_size() + ckpt_initial = ckpt_pool.available_size() + tokens = [1, 2, 3, 4] + + self._cache_finished(cache, allocator, req_to_token_pool, tokens) + self.assertEqual( + req_to_token_pool.mamba_allocator.available_size(), active_initial + ) + self.assertEqual(ckpt_pool.available_size(), ckpt_initial - 1) + self.assertEqual(cache.mamba_evictable_size(), 1) + + self._cache_finished(cache, allocator, req_to_token_pool, tokens) + self.assertEqual( + req_to_token_pool.mamba_allocator.available_size(), active_initial + ) + self.assertEqual(ckpt_pool.available_size(), ckpt_initial - 1) + self.assertEqual(cache.mamba_evictable_size(), 1) + + result = cache.evict(EvictParams(mamba_num=1)) + self.assertEqual(result.mamba_num_evicted, 1) + self.assertEqual( + req_to_token_pool.mamba_allocator.available_size(), active_initial + ) + self.assertEqual(ckpt_pool.available_size(), ckpt_initial) + self.assertEqual(cache.mamba_evictable_size(), 0) + + _CONFIGS: list[CacheConfig] = [ CacheConfig(page_size=1, components=(ComponentType.FULL,)), CacheConfig(page_size=4, components=(ComponentType.FULL,)),