diff --git a/python/sglang/srt/arg_groups/hicache_hook.py b/python/sglang/srt/arg_groups/hicache_hook.py index cd3c02df6..0dd85b080 100644 --- a/python/sglang/srt/arg_groups/hicache_hook.py +++ b/python/sglang/srt/arg_groups/hicache_hook.py @@ -23,6 +23,19 @@ def handle_hicache(server_args: Any): 2) Storage <-> layout compatibility (may rewrite layout). """ cfg = resolving_view(server_args) + if cfg.enable_unified_cache_external_linker: + if cfg.enable_hierarchical_cache: + raise ValueError( + "--enable-unified-cache-external-linker and " + "--enable-hierarchical-cache are mutually exclusive." + ) + if cfg.hicache_storage_backend is not None: + raise ValueError( + "--enable-unified-cache-external-linker does not use " + "--hicache-storage-backend." + ) + return + # Skip all normalization when neither hicache nor decode-offload path is active. if not ( cfg.enable_hierarchical_cache diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 00a6844f0..9c61510a6 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -481,6 +481,9 @@ class Scheduler( self.enable_hierarchical_cache = get_memory().enable_hierarchical_cache self.enable_session_radix_cache = get_memory().enable_session_radix_cache self.enable_hicache_storage = get_memory().hicache_storage_backend is not None + self.enable_unified_cache_external_linker = ( + get_memory().enable_unified_cache_external_linker + ) self.enable_decode_hicache = ( get_disagg().disaggregation_decode_enable_radix_cache and self.enable_hierarchical_cache @@ -3132,6 +3135,15 @@ class Scheduler( return False return True + def _release_aborted_request(self, rid: str) -> None: + """Drop the cache-side state an aborted request left behind.""" + if ( + self.enable_hierarchical_cache + or self.enable_hicache_storage + or self.enable_unified_cache_external_linker + ): + self.tree_cache.release_aborted_request(rid) + def _abort_on_queued_limit(self, recv_req: Req) -> bool: """Abort an incoming or existing request if the waiting queue is full. Returns True if the incoming request is aborted.""" if ( @@ -3158,9 +3170,7 @@ class Scheduler( direction * recv_req.priority < direction * candidate_req.priority ) if abort_existing_req: - if self.enable_hicache_storage: - # Release prefetch events associated with the request - self.tree_cache.release_aborted_request(candidate_req.rid) + self._release_aborted_request(candidate_req.rid) self.waiting_queue.pop(idx) self.beam_coordinator.retire_group(candidate_req) req_to_abort = candidate_req @@ -3189,9 +3199,7 @@ class Scheduler( for req in self.waiting_queue: entry_time = req.time_stats.wait_queue_entry_time if 0 < entry_time < deadline: - if self.enable_hicache_storage: - # Release prefetch events associated with the request - self.tree_cache.release_aborted_request(req.rid) + self._release_aborted_request(req.rid) self.ipc_channels.send_to_tokenizer.send_output( _make_abort_req( req, @@ -3347,8 +3355,7 @@ class Scheduler( req, self.req_to_metadata_buffer_idx_allocator ) req.pending_bootstrap = False - if self.enable_hicache_storage: - self.tree_cache.release_aborted_request(req.rid) + self._release_aborted_request(req.rid) release_kv_cache(req, self.tree_cache, is_insert=False) self.chunked_req = None @@ -3607,7 +3614,11 @@ class Scheduler( for req in ready_grammar_requests: self._add_request_to_queue(req) - if self.enable_hierarchical_cache or get_memory().enable_flexkv: + if ( + self.enable_hierarchical_cache + or get_memory().enable_flexkv + or self.enable_unified_cache_external_linker + ): self.tree_cache.check_hicache_events() if self.enable_hicache_storage: self._retry_missed_storage_prefetches() @@ -3780,7 +3791,10 @@ class Scheduler( if res != AddReqResult.CONTINUE: if res == AddReqResult.NO_TOKEN: - if self.enable_hierarchical_cache: + if ( + self.enable_hierarchical_cache + or self.enable_unified_cache_external_linker + ): # Set batch_is_full after making sure there are requests that can be served running_batch.batch_is_full = len(adder.can_run_list) > 0 or ( not running_batch.is_empty() @@ -3844,7 +3858,7 @@ class Scheduler( self.chunked_req is None or len(can_run_list) != 1 ) - if self.enable_hierarchical_cache: + if self.enable_hierarchical_cache or self.enable_unified_cache_external_linker: # todo (zhiqiang): disable cuda graph execution if hicache loading triggered new_batch.hicache_consumer_index = ( self.tree_cache.ready_to_load_host_cache() @@ -5052,10 +5066,8 @@ class Scheduler( # This only works for requests that have not started anything. # We still need to send something back to TokenizerManager to clean up the state. req = self.waiting_queue.pop(i) + self._release_aborted_request(req.rid) self.beam_coordinator.retire_group(req) - if self.enable_hicache_storage: - # to release prefetch events associated with the request - self.tree_cache.release_aborted_request(req.rid) self.ipc_channels.send_to_tokenizer.send_output(_make_abort_req(req), req) # For disaggregation decode mode, the request in the waiting queue has KV cache allocated. if self.disaggregation_mode == DisaggregationMode.DECODE: @@ -5086,8 +5098,7 @@ class Scheduler( for req in self.dllm_manager.pop_aborted_reqs( recv_req.abort_all, recv_req.rid ): - if self.enable_hicache_storage: - self.tree_cache.release_aborted_request(req.rid) + self._release_aborted_request(req.rid) self.ipc_channels.send_to_tokenizer.send_output( _make_abort_req(req), req ) @@ -5107,8 +5118,7 @@ class Scheduler( for req in self.disagg_prefill_bootstrap_queue.queue: if recv_req.abort_all or req.rid.startswith(recv_req.rid): logger.debug(f"Abort bootstrap queue request. {req.rid=}") - if self.enable_hicache_storage: - self.tree_cache.release_aborted_request(req.rid) + self._release_aborted_request(req.rid) if hasattr(req.disagg_kv_sender, "abort"): req.disagg_kv_sender.abort() diff --git a/python/sglang/srt/mem_cache/hicache_storage.py b/python/sglang/srt/mem_cache/hicache_storage.py index d0402364d..b573e69f1 100644 --- a/python/sglang/srt/mem_cache/hicache_storage.py +++ b/python/sglang/srt/mem_cache/hicache_storage.py @@ -129,7 +129,7 @@ class PoolTransferResult: extra_pool_hit_pages: dict[str, int] # Pools with TRAILING_PAGES (SWA, Mamba state) only hold a window that ends on an - # offloaded node boundary. + # offloaded node boundary, so 5 can be restorable while 4 and 3 are not. # Each rank owns its own shard and may hold a different set, so reducing a # per-rank maximum would pick a length that is illegal on another rank; the # caller intersects these sets instead. diff --git a/python/sglang/srt/mem_cache/kv_cache_builder.py b/python/sglang/srt/mem_cache/kv_cache_builder.py index e0a806670..ac030ac0e 100644 --- a/python/sglang/srt/mem_cache/kv_cache_builder.py +++ b/python/sglang/srt/mem_cache/kv_cache_builder.py @@ -313,6 +313,8 @@ def build_kv_cache( enable_mamba_extra_buffer_lazy=server_args.enable_mamba_extra_buffer_lazy(), pp_rank=ps.pp_rank, pp_size=ps.pp_size, + attn_cp_rank=ps.attn_cp_rank, + attn_cp_size=ps.attn_cp_size, chunked_prefill_size=effective_chunked_prefill_size, sliding_window_size=sliding_window_size, mtp_draft_device_pools=mtp_draft_device_pools, diff --git a/python/sglang/srt/mem_cache/kv_cache_configurator.py b/python/sglang/srt/mem_cache/kv_cache_configurator.py index 3d4d093ab..6ac0bfe1f 100644 --- a/python/sglang/srt/mem_cache/kv_cache_configurator.py +++ b/python/sglang/srt/mem_cache/kv_cache_configurator.py @@ -105,6 +105,7 @@ def _should_elide_dsa_index_k(*, is_draft_worker: bool) -> bool: not memory_config.enable_hisparse and not is_draft_worker and not memory_config.enable_hierarchical_cache + and not memory_config.enable_unified_cache_external_linker and get_disagg().disaggregation_mode == "null" ) diff --git a/python/sglang/srt/mem_cache/registry.py b/python/sglang/srt/mem_cache/registry.py index a1371e384..8e90b901d 100644 --- a/python/sglang/srt/mem_cache/registry.py +++ b/python/sglang/srt/mem_cache/registry.py @@ -108,6 +108,9 @@ def default_radix_cache_factory(ctx: TreeCacheBuildContext) -> BasePrefixCache: logger.info("Using experimental C++ radix tree implementation.") return RadixCacheCpp(params=params, server_args=server_args) + if server_args.enable_unified_cache_external_linker: + return _create_unified_radix_cache(ctx, server_args, params) + if ctx.is_hybrid_swa and ctx.full_tokens_per_layer == 0: from sglang.srt.mem_cache.pure_swa_radix_cache import PureSWARadixCache @@ -193,6 +196,26 @@ def _create_unified_radix_cache( ctx.tp_worker.register_hicache_layer_transfer_counter( cache.cache_controller.layer_done_counter ) + elif server_args.enable_unified_cache_external_linker: + backend = server_args.unified_cache_external_linker_backend + if backend == "mooncake": + from sglang.srt.mem_cache.storage.mooncake_store.mooncake_direct_linker import ( + MooncakeDirectLinker, + ) + + linker_cls = MooncakeDirectLinker + else: + raise ValueError( + f"Unknown unified cache external linker backend: {backend!r}" + ) + + cache.init_cache_linker( + linker_cls(server_args, params, components=set(cache.components)) + ) + counter = cache.linker.layer_done_counter + kvcache = params.token_to_kv_pool_allocator.get_kvcache() + kvcache.register_layer_transfer_counter(counter) + ctx.tp_worker.register_hicache_layer_transfer_counter(counter) return cache diff --git a/python/sglang/srt/mem_cache/unified_cache/components/full_component.py b/python/sglang/srt/mem_cache/unified_cache/components/full_component.py index 95fecb2ab..e76108615 100644 --- a/python/sglang/srt/mem_cache/unified_cache/components/full_component.py +++ b/python/sglang/srt/mem_cache/unified_cache/components/full_component.py @@ -487,7 +487,10 @@ class FullComponent(TreeComponent): if phase == ExternalLinkerLoadPhase.ABORT: self._full_allocator().free(transfer.device_indices) return None + if phase == ExternalLinkerLoadPhase.PREPARE: + return transfer + assert phase == ExternalLinkerLoadPhase.COMMIT return transfer def free_host_values(self, host_values: list[torch.Tensor]) -> None: diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index b534f9d74..feec0ab26 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -2821,6 +2821,23 @@ class ServerArgs: NS("memory"), ] = 4 + # ------------------------------------------------------------------------- + # Unified Radix Cache + # ------------------------------------------------------------------------- + enable_unified_cache_external_linker: A[ + bool, + "Link UnifiedRadixCache directly to an external KV store (direct L3), with no host cache tier.", + NS("memory"), + ] = False + unified_cache_external_linker_backend: A[ + str, + Arg( + help="Storage backend for --enable-unified-cache-external-linker.", + choices=["mooncake"], + ), + NS("memory"), + ] = "mooncake" + # ------------------------------------------------------------------------- # Hierarchical sparse attention # ------------------------------------------------------------------------- diff --git a/python/sglang/test/mooncake_utils.py b/python/sglang/test/mooncake_utils.py new file mode 100644 index 000000000..f8cee645b --- /dev/null +++ b/python/sglang/test/mooncake_utils.py @@ -0,0 +1,185 @@ +import os +import signal +import socket +import subprocess +import time + +import requests + +from sglang.test.server_fixtures.disaggregation_fixture import ( + get_rdma_devices_args, +) +from sglang.test.test_utils import find_available_port + + +class MooncakeTestServices: + """Lifecycle for a local Mooncake master and external storage client.""" + + def __init__( + self, + *, + protocol: str = "rdma", + store_segment_size: int = 4 * 1024**3, + device: str | None = None, + local_hostname: str | None = None, + ): + self.protocol = protocol + self.store_segment_size = store_segment_size + self.device = device if device is not None else self._default_device(protocol) + self.local_hostname = local_hostname or socket.gethostbyname( + socket.gethostname() + ) + + self.master_port = find_available_port(50051) + self.master_metrics_port = find_available_port(9003) + self.metadata_port = find_available_port(8080) + self.store_port = find_available_port(50052) + self.store_http_port = find_available_port(8081) + + self.metadata_process = None + self.master_process = None + self.store_process = None + + @staticmethod + def _default_device(protocol: str) -> str: + configured = os.environ.get("SGLANG_TEST_MOONCAKE_DEVICE") + if configured is not None: + return configured + if protocol == "rdma": + return get_rdma_devices_args().split(",")[0] + return "" + + def start(self): + self.metadata_process = self._launch( + [ + "python3", + "-m", + "mooncake.http_metadata_server", + "--port", + str(self.metadata_port), + ] + ) + try: + self.master_process = self._launch( + [ + "mooncake_master", + "--port", + str(self.master_port), + "--metrics_port", + str(self.master_metrics_port), + ] + ) + self._wait_for_core_services() + self.store_process = self._launch( + [ + "mooncake_client", + f"--host={self.local_hostname}", + f"--port={self.store_port}", + f"--master_server_address=127.0.0.1:{self.master_port}", + f"--metadata_server=http://127.0.0.1:{self.metadata_port}/metadata", + f"--protocol={self.protocol}", + f"--device_names={self.device}", + f"--global_segment_size={self.store_segment_size}", + "--enable_http_server=true", + f"--http_port={self.store_http_port}", + ], + env={**os.environ, "MC_MS_AUTO_DISC": "0"}, + ) + self._wait_for_store() + except Exception: + self.stop() + raise + + def stop(self): + for name in ("store_process", "master_process", "metadata_process"): + process = getattr(self, name) + if process is None: + continue + self._stop_process_group(process) + setattr(self, name, None) + + def server_env(self) -> dict[str, str]: + return { + "MOONCAKE_MASTER": f"127.0.0.1:{self.master_port}", + "MOONCAKE_PROTOCOL": self.protocol, + "MC_MS_AUTO_DISC": "0", + "MOONCAKE_DEVICE": self.device, + "MOONCAKE_LOCAL_HOSTNAME": self.local_hostname, + "MOONCAKE_TE_META_DATA_SERVER": ( + f"http://127.0.0.1:{self.metadata_port}/metadata" + ), + "MOONCAKE_GLOBAL_SEGMENT_SIZE": "0", + } + + def master_metric(self, name: str) -> float: + response = requests.get( + f"http://127.0.0.1:{self.master_metrics_port}/metrics", + timeout=5, + ) + response.raise_for_status() + prefix = f"{name} " + for line in response.text.splitlines(): + if line.startswith(prefix): + return float(line.split()[1]) + raise AssertionError(f"{name} is missing from Mooncake master metrics") + + @staticmethod + def _launch(command, env=None): + return subprocess.Popen( + command, + stdout=subprocess.DEVNULL, + stderr=subprocess.STDOUT, + preexec_fn=os.setsid, + env=env, + ) + + def _wait_for_core_services(self, timeout: int = 30): + deadline = time.monotonic() + timeout + master_ready_at = time.monotonic() + 3 + while time.monotonic() < deadline: + self._raise_if_exited(self.metadata_process, "metadata service") + self._raise_if_exited(self.master_process, "master service") + try: + requests.get( + f"http://127.0.0.1:{self.metadata_port}/metadata", + timeout=2, + ) + if time.monotonic() >= master_ready_at: + return + except requests.RequestException: + pass + time.sleep(1) + raise TimeoutError("Timed out waiting for Mooncake metadata and master") + + def _wait_for_store(self, timeout: int = 90): + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + self._raise_if_exited(self.store_process, "store service") + try: + with socket.create_connection( + (self.local_hostname, self.store_port), timeout=2 + ): + return + except OSError: + time.sleep(1) + raise TimeoutError("Timed out waiting for Mooncake store") + + @staticmethod + def _raise_if_exited(process, name: str): + returncode = process.poll() + if returncode is not None: + raise RuntimeError(f"Mooncake {name} exited with code {returncode}") + + @staticmethod + def _stop_process_group(process): + try: + os.killpg(os.getpgid(process.pid), signal.SIGTERM) + process.wait(timeout=10) + except ProcessLookupError: + return + except (subprocess.TimeoutExpired, OSError): + try: + os.killpg(os.getpgid(process.pid), signal.SIGKILL) + process.wait(timeout=5) + except (ProcessLookupError, subprocess.TimeoutExpired, OSError): + pass diff --git a/test/registered/radix_cache/unified_radix_tree/linker/test_unified_cache_linker_kl_dsv4.py b/test/registered/radix_cache/unified_radix_tree/linker/test_unified_cache_linker_kl_dsv4.py new file mode 100644 index 000000000..59839147e --- /dev/null +++ b/test/registered/radix_cache/unified_radix_tree/linker/test_unified_cache_linker_kl_dsv4.py @@ -0,0 +1,146 @@ +"""DeepSeek-V4 Flash UnifiedRadixCache direct-linker load-back KL tests.""" + +import json +import os +import unittest + +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kits.unified_radix_cache_kit import UnifiedRadixTreeTestMixin +from sglang.test.kl_multiturn_utils import get_input_ids +from sglang.test.mooncake_utils import MooncakeTestServices +from sglang.test.test_utils import ( + CustomTestCase, + find_available_port, + popen_launch_server, + terminate_and_kill_process_tree, +) + +DSV4_FLASH_MODEL = os.environ.get( + "SGLANG_LINKER_DSV4_FLASH_MODEL", "sgl-project/DeepSeek-V4-Flash-FP8" +) +DSV4_FLASH_LAUNCH_TIMEOUT = 3600 + +register_cuda_ci(est_time=1500, stage="extra-b", runner_config="4-gpu-h100") + + +class TestDeepSeekV4FlashUnifiedCacheLinkerKL( + UnifiedRadixTreeTestMixin, CustomTestCase +): + page_size = 256 + kl_threshold = 0.01 + sampling_temperature = 0 + max_new_tokens = 64 + prefix_len = 2048 + decode_hit_request_batch_size = 3 + decode_hit_inter_batch_delay_s = 0.5 + + @classmethod + def setUpClass(cls): + cls.model = DSV4_FLASH_MODEL + cls.base_url = f"http://127.0.0.1:{find_available_port(30000)}" + cls.mooncake = MooncakeTestServices() + cls.mooncake.start() + cls.process = None + try: + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DSV4_FLASH_LAUNCH_TIMEOUT, + other_args=[ + "--trust-remote-code", + "--tp-size", + "4", + "--attention-backend", + "compressed", + "--page-size", + str(cls.page_size), + "--chunked-prefill-size", + "8192", + "--mem-fraction-static", + "0.92", + "--disable-shared-experts-fusion", + "--swa-full-tokens-ratio", + "0.25", + "--max-total-tokens", + "8192", + "--max-running-requests", + "1", + "--enable-cache-report", + "--enable-unified-cache-external-linker", + "--hicache-storage-backend-extra-config", + json.dumps({"enable_group_semantics": True}), + ], + env={ + **cls.mooncake.server_env(), + "SGLANG_DSV4_FP4_EXPERTS": "0", + "SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1", + }, + ) + cls.input_ids = get_input_ids(cls.model, num_samples=18) + except Exception: + try: + if cls.process is not None: + terminate_and_kill_process_tree(cls.process) + finally: + cls.mooncake.stop() + raise + + @classmethod + def tearDownClass(cls): + try: + if cls.process is not None: + terminate_and_kill_process_tree(cls.process) + finally: + cls.mooncake.stop() + + @unittest.skip("Linker CI targets Direct load-back KL accuracy") + def test_gsm8k(self): + pass + + @unittest.skip("Linker CI targets Direct load-back KL accuracy") + def test_mmlu(self): + pass + + def prefill_cache_assert(self, result, prefix_len, label): + self._record_cache_result(result, prefix_len, label) + + def decode_cache_assert(self, result, history_len, output_len, label): + self._record_cache_result(result, history_len + output_len, label) + + def _record_cache_result(self, result, expected_cached_tokens, label): + meta_info = result["meta_info"] + cached_tokens = int(meta_info["cached_tokens"]) + minimum = max(0, expected_cached_tokens - self.page_size) + self.assertGreaterEqual( + cached_tokens, + minimum, + f"{label}: expected cached_tokens >= {minimum}, got {cached_tokens}", + ) + details = meta_info.get("cached_tokens_details") or {} + remote_tokens = int(details.get("host", 0)) + self._direct_remote_tokens += remote_tokens + if remote_tokens: + print(f"{label}: Direct load-back confirmed for {remote_tokens} tokens") + + def _run_linker_kl_case(self, test_case): + self._direct_remote_tokens = 0 + test_case() + print(f"Direct load-back total: {self._direct_remote_tokens} tokens") + self.assertGreater( + self._direct_remote_tokens, + 0, + "Expected this KL case to load KV through the Mooncake Direct Linker", + ) + + def test_multiturn_logprobs_match(self): + self._run_linker_kl_case(super().test_multiturn_logprobs_match) + + def test_multiturn_prefill_cache_hit_branching(self): + self._run_linker_kl_case(super().test_multiturn_prefill_cache_hit_branching) + + def test_multiturn_decode_cache_hit_branching(self): + self._run_linker_kl_case(super().test_multiturn_decode_cache_hit_branching) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/radix_cache/unified_radix_tree/linker/test_unified_cache_linker_kl_glm52.py b/test/registered/radix_cache/unified_radix_tree/linker/test_unified_cache_linker_kl_glm52.py new file mode 100644 index 000000000..55b6d0b91 --- /dev/null +++ b/test/registered/radix_cache/unified_radix_tree/linker/test_unified_cache_linker_kl_glm52.py @@ -0,0 +1,136 @@ +"""GLM-5.2 UnifiedRadixCache direct-linker load-back KL tests.""" + +import json +import os +import unittest + +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kits.unified_radix_cache_kit import UnifiedRadixTreeTestMixin +from sglang.test.kl_multiturn_utils import get_input_ids +from sglang.test.mooncake_utils import MooncakeTestServices +from sglang.test.test_utils import ( + CustomTestCase, + find_available_port, + popen_launch_server, + terminate_and_kill_process_tree, +) + +GLM52_MODEL = os.environ.get("SGLANG_LINKER_GLM52_MODEL", "zai-org/GLM-5.2-FP8") +GLM52_LAUNCH_TIMEOUT = 3600 + +register_cuda_ci(est_time=1200, stage="extra-b", runner_config="8-gpu-h200") + + +class TestGLM52UnifiedCacheLinkerKL(UnifiedRadixTreeTestMixin, CustomTestCase): + page_size = 64 + kl_threshold = 0.03 + sampling_temperature = 0 + max_new_tokens = 64 + prefix_len = 2048 + decode_hit_request_batch_size = 3 + decode_hit_inter_batch_delay_s = 0.5 + + @classmethod + def setUpClass(cls): + cls.model = GLM52_MODEL + cls.base_url = f"http://127.0.0.1:{find_available_port(30000)}" + cls.mooncake = MooncakeTestServices() + cls.mooncake.start() + cls.process = None + try: + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=GLM52_LAUNCH_TIMEOUT, + other_args=[ + "--trust-remote-code", + "--tp-size", + "8", + "--page-size", + str(cls.page_size), + "--mem-fraction-static", + "0.8", + "--model-loader-extra-config", + '{"enable_multithread_load": true, "num_threads": 64}', + "--max-total-tokens", + "12000", + "--max-running-requests", + "1", + "--enable-cache-report", + "--enable-unified-cache-external-linker", + "--hicache-storage-backend-extra-config", + json.dumps({"enable_group_semantics": True}), + ], + env={ + **cls.mooncake.server_env(), + "SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1", + }, + ) + cls.input_ids = get_input_ids(cls.model, num_samples=18) + except Exception: + try: + if cls.process is not None: + terminate_and_kill_process_tree(cls.process) + finally: + cls.mooncake.stop() + raise + + @classmethod + def tearDownClass(cls): + try: + if cls.process is not None: + terminate_and_kill_process_tree(cls.process) + finally: + cls.mooncake.stop() + + @unittest.skip("Linker CI targets Direct load-back KL accuracy") + def test_gsm8k(self): + pass + + @unittest.skip("Linker CI targets Direct load-back KL accuracy") + def test_mmlu(self): + pass + + def prefill_cache_assert(self, result, prefix_len, label): + self._record_cache_result(result, prefix_len, label) + + def decode_cache_assert(self, result, history_len, output_len, label): + self._record_cache_result(result, history_len + output_len, label) + + def _record_cache_result(self, result, expected_cached_tokens, label): + meta_info = result["meta_info"] + cached_tokens = int(meta_info["cached_tokens"]) + minimum = max(0, expected_cached_tokens - self.page_size) + self.assertGreaterEqual( + cached_tokens, + minimum, + f"{label}: expected cached_tokens >= {minimum}, got {cached_tokens}", + ) + details = meta_info.get("cached_tokens_details") or {} + remote_tokens = int(details.get("host", 0)) + self._direct_remote_tokens += remote_tokens + if remote_tokens: + print(f"{label}: Direct load-back confirmed for {remote_tokens} tokens") + + def _run_linker_kl_case(self, test_case): + self._direct_remote_tokens = 0 + test_case() + print(f"Direct load-back total: {self._direct_remote_tokens} tokens") + self.assertGreater( + self._direct_remote_tokens, + 0, + "Expected this KL case to load KV through the Mooncake Direct Linker", + ) + + def test_multiturn_logprobs_match(self): + self._run_linker_kl_case(super().test_multiturn_logprobs_match) + + def test_multiturn_prefill_cache_hit_branching(self): + self._run_linker_kl_case(super().test_multiturn_prefill_cache_hit_branching) + + def test_multiturn_decode_cache_hit_branching(self): + self._run_linker_kl_case(super().test_multiturn_decode_cache_hit_branching) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/managers/test_scheduler_timeouts.py b/test/registered/unit/managers/test_scheduler_timeouts.py index 4fd27f1c4..d9bdc99a4 100644 --- a/test/registered/unit/managers/test_scheduler_timeouts.py +++ b/test/registered/unit/managers/test_scheduler_timeouts.py @@ -51,12 +51,42 @@ def _req( def _scheduler(waiting_queue): s = Scheduler.__new__(Scheduler) s.waiting_queue = waiting_queue + s.enable_hierarchical_cache = False s.enable_hicache_storage = False + s.enable_unified_cache_external_linker = False s.ipc_channels = SimpleNamespace(send_to_tokenizer=MagicMock()) s.beam_coordinator = MagicMock() return s +class TestQueuedLimitAbort(CustomTestCase): + def setUp(self): + patcher = patch( + "sglang.srt.managers.scheduler.get_serving", + return_value=SimpleNamespace(weight_version="v0"), + ) + patcher.start() + self.addCleanup(patcher.stop) + + def test_hicache_without_storage_uses_common_abort_cleanup(self): + candidate = _req("candidate", wait_entry=1.0) + candidate.priority = 0 + incoming = _req("incoming", wait_entry=2.0) + incoming.priority = 1 + + s = _scheduler([candidate]) + s.max_queued_requests = 1 + s.enable_priority_scheduling = True + s.schedule_low_priority_values_first = False + s.enable_hierarchical_cache = True + s.tree_cache = MagicMock(spec=["release_aborted_request"]) + + self.assertFalse(s._abort_on_queued_limit(incoming)) + + s.tree_cache.release_aborted_request.assert_called_once_with("candidate") + self.assertEqual(s.waiting_queue, []) + + class TestWaitingTimeout(CustomTestCase): def setUp(self): patcher = patch( diff --git a/test/registered/unit/mem_cache/test_registry.py b/test/registered/unit/mem_cache/test_registry.py index 80a613ada..5f9f903fb 100644 --- a/test/registered/unit/mem_cache/test_registry.py +++ b/test/registered/unit/mem_cache/test_registry.py @@ -52,6 +52,7 @@ def _make_ctx( enable_streaming_session=enable_streaming, enable_lmcache=enable_lmcache, enable_flexkv=False, + enable_unified_cache_external_linker=False, ) return TreeCacheBuildContext( server_args=server_args,