From 8bb106cee94c1ccb5625cdb78ef0ee3e0cad3918 Mon Sep 17 00:00:00 2001 From: Wei Yang Date: Wed, 19 Aug 2026 02:55:51 +0800 Subject: [PATCH] Fix NIXL cleaner grouping for hybrid cache keys (#35130) Co-authored-by: Wei Yang --- .../mem_cache/storage/nixl/nixl_cleaner.py | 13 ++- .../mem_cache/test_hicache_nixl_cleaner.py | 80 +++++++++++++++---- 2 files changed, 75 insertions(+), 18 deletions(-) diff --git a/python/sglang/srt/mem_cache/storage/nixl/nixl_cleaner.py b/python/sglang/srt/mem_cache/storage/nixl/nixl_cleaner.py index 0cd2a019d..171edc24b 100644 --- a/python/sglang/srt/mem_cache/storage/nixl/nixl_cleaner.py +++ b/python/sglang/srt/mem_cache/storage/nixl/nixl_cleaner.py @@ -11,6 +11,7 @@ import time from dataclasses import dataclass, field from typing import Iterable, Optional +from sglang.srt.mem_cache.hicache_storage import PoolName from sglang.srt.mem_cache.storage.nixl.nixl_routing import BUCKET_HEX_CHARS logger = logging.getLogger(__name__) @@ -21,6 +22,12 @@ _DEFAULT_HIGH_WATERMARK = 80.0 _DEFAULT_LOW_WATERMARK = 70.0 _RANK_SUFFIX_RE = re.compile(r"_(\d+)_(\d+)$") _KV_SUFFIXES = ("_k", "_v") +_POOL_NAME_PATTERN = "|".join( + sorted((re.escape(pool.value) for pool in PoolName), key=len, reverse=True) +) +_HYBRID_COMPONENT_SUFFIX_RE = re.compile( + rf"_(?:{_POOL_NAME_PATTERN})(?:_(?:temporal|conv_\d+|[kv]|\d+))?$" +) _BUCKET_NAME_RE = re.compile(rf"^[0-9a-f]{{{BUCKET_HEX_CHARS}}}$") @@ -40,10 +47,12 @@ class _GroupInfo: def _parse_group_key(name: str) -> str: """Return the logical cache-key group for one NIXL FILE object name. - The physical names are produced by ``HiCacheNixl._get_suffixed_key`` and - ``HiCacheNixl._get_key_list_from_meta``. + The physical names are produced by ``HiCacheNixl._get_suffixed_key``, + ``HiCacheNixl._get_key_list_from_meta``, and + ``HiCacheNixl._get_hybrid_component_keys``. """ stem = name + stem = _HYBRID_COMPONENT_SUFFIX_RE.sub("", stem) for suffix in _KV_SUFFIXES: if stem.endswith(suffix): stem = stem[: -len(suffix)] diff --git a/test/registered/unit/mem_cache/test_hicache_nixl_cleaner.py b/test/registered/unit/mem_cache/test_hicache_nixl_cleaner.py index b627c7977..2f7e4bb13 100644 --- a/test/registered/unit/mem_cache/test_hicache_nixl_cleaner.py +++ b/test/registered/unit/mem_cache/test_hicache_nixl_cleaner.py @@ -40,20 +40,7 @@ class TestHiCacheL3Cleaner(CustomTestCase): os.utime(path, (mtime, mtime)) return path - def test_parse_group_key_strips_rank_and_kv_suffix(self): - """Keys for TP ranks and zero-copy K/V files share one cleanup group.""" - self.assertEqual(_parse_group_key("page-a_model_0_8"), "page-a_model") - self.assertEqual(_parse_group_key("page-a_model_7_8_k"), "page-a_model") - self.assertEqual(_parse_group_key("page-a_model_7_8_v"), "page-a_model") - self.assertEqual(_parse_group_key("page-a_model_k"), "page-a_model") - - def test_tick_deletes_oldest_group_across_bucketed_dirs(self): - """A cleaner batch deletes all files in the oldest logical key group.""" - old_keys = ["page-old_model_0_2", "page-old_model_1_2"] - new_keys = ["page-new_model_0_2", "page-new_model_1_2"] - old_paths = [self._write_key(key, mtime=100.0) for key in old_keys] - new_paths = [self._write_key(key, mtime=200.0) for key in new_keys] - + def _run_single_group_cleanup(self) -> None: cleaner = HiCacheL3Cleaner( self.base_dirs, tp_rank=0, @@ -62,7 +49,6 @@ class TestHiCacheL3Cleaner(CustomTestCase): recheck_groups=1, unlink_workers=1, ) - usage_calls: dict[str, int] = {} def fake_usage(path: str) -> float: @@ -70,8 +56,70 @@ class TestHiCacheL3Cleaner(CustomTestCase): return 90.0 if usage_calls[path] == 1 else 60.0 cleaner._disk_usage_pct = fake_usage - self.assertTrue(cleaner._tick()) + + def test_parse_group_key_strips_rank_and_kv_suffix(self): + """Keys for TP ranks and zero-copy K/V files share one cleanup group.""" + self.assertEqual(_parse_group_key("page-a_model_0_8"), "page-a_model") + self.assertEqual(_parse_group_key("page-a_model_7_8_k"), "page-a_model") + self.assertEqual(_parse_group_key("page-a_model_7_8_v"), "page-a_model") + self.assertEqual(_parse_group_key("page-a_model_k"), "page-a_model") + + def test_parse_group_key_strips_hybrid_component_suffix(self): + """All hybrid component shapes share the logical page's cleanup group.""" + names = [ + "page-a_model_7_8_kv_k", + "page-a_model_7_8_swa_k", + "page-a_model_7_8_swa_v", + "page-a_model_7_8_mamba_temporal", + "page-a_model_7_8_mamba_conv_0", + "page-a_model_7_8_indexer_2", + "page-a_model_7_8_draft_swa", + "page-a_model_deepseek_v4_c4_indexer_state_2", + ] + + for name in names: + with self.subTest(name=name): + self.assertEqual(_parse_group_key(name), "page-a_model") + + def test_tick_deletes_oldest_group_across_bucketed_dirs(self): + """A cleaner batch deletes all files in the oldest logical key group.""" + old_keys = ["page-old_model_0_2", "page-old_model_1_2"] + new_keys = ["page-new_model_0_2", "page-new_model_1_2"] + old_paths = [self._write_key(key, mtime=100.0) for key in old_keys] + new_paths = [self._write_key(key, mtime=200.0) for key in new_keys] + + self._run_single_group_cleanup() + self.assertFalse(any(os.path.exists(path) for path in old_paths)) + self.assertTrue(all(os.path.exists(path) for path in new_paths)) + + def test_tick_deletes_hybrid_components_atomically(self): + """Evict every pool component and TP rank for one logical page.""" + physical_suffixes = [ + "", + "_k", + "_v", + "_kv_k", + "_kv_v", + "_swa_k", + "_swa_v", + "_mamba_temporal", + "_mamba_conv_0", + ] + old_keys = [ + f"page-old_model_{rank}_2{suffix}" + for rank in range(2) + for suffix in physical_suffixes + ] + new_keys = [ + f"page-new_model_{rank}_2{suffix}" + for rank in range(2) + for suffix in physical_suffixes + ] + old_paths = [self._write_key(key, mtime=100.0) for key in old_keys] + new_paths = [self._write_key(key, mtime=200.0) for key in new_keys] + + self._run_single_group_cleanup() self.assertFalse(any(os.path.exists(path) for path in old_paths)) self.assertTrue(all(os.path.exists(path) for path in new_paths))