[Unified Cache][5/N]: Integrate external linker mode end to end (#37381)

Co-authored-by: 晟海 <huangtingwei.htw@antgroup.com>
This commit is contained in:
Zhangheng
2026-09-04 02:02:58 +08:00
committed by GitHub
co-authored by 晟海
parent 619ab2bcce
commit abed680320
13 changed files with 586 additions and 19 deletions
@@ -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()
@@ -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()
@@ -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(
@@ -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,