From ceb1d2e580dd3dda5d6c9a481fe5b15afedd3f82 Mon Sep 17 00:00:00 2001 From: Zhiqiang Xie Date: Fri, 18 Sep 2026 16:04:15 -0700 Subject: [PATCH] [PD] Enable optimistic prefill with buffer-only L3 write-through HiCache (#40043) Co-authored-by: cctry --- python/sglang/srt/arg_groups/serving_hook.py | 17 +++- python/sglang/srt/disaggregation/prefill.py | 14 +++ .../test_disaggregation_optimistic_prefill.py | 97 ++++++++++++++++++- .../unit/server_args/test_server_args.py | 42 ++++++++ 4 files changed, 161 insertions(+), 9 deletions(-) diff --git a/python/sglang/srt/arg_groups/serving_hook.py b/python/sglang/srt/arg_groups/serving_hook.py index 2015ebf0e..fc2ccfead 100644 --- a/python/sglang/srt/arg_groups/serving_hook.py +++ b/python/sglang/srt/arg_groups/serving_hook.py @@ -466,13 +466,20 @@ def handle_other_validations(server_args: Any): "_handle_other_validations", optimistic_prefill_attempts=0, ) - elif cfg.enable_hierarchical_cache and ( - cfg.hicache_storage_backend is not None - or cfg.hicache_write_policy != "write_back" + elif cfg.enable_hierarchical_cache and not ( + ( + cfg.hicache_storage_backend is None + and cfg.hicache_write_policy == "write_back" + ) + or ( + cfg.hicache_storage_backend is not None + and cfg.hicache_host_memory_mode == "buffer_only" + and cfg.hicache_write_policy == "write_through" + ) ): logger.warning( - "Optimistic prefill only supports L2 hierarchical cache " - "with write-back policy" + "Optimistic prefill supports L2 write-back or L3 buffer-only " + "write-through hierarchical cache" ) declare_resolution( server_args, diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index deb2ee921..e0cbfb423 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -1496,6 +1496,17 @@ class SchedulerDisaggregationPrefillMixin: """Release KV cache and requeue an optimistic prefill request.""" max_attempts = get_disagg().optimistic_prefill_attempts maybe_cache_unfinished_req(req, self.tree_cache) + # The cached prefix is evictable once the KV is released. Its length + # (capped at what a retry can match) seeds the retry's storage baseline, + # so an evicted prefix is looked up in L3 once before it is recomputed. + yielded_prefix_len = ( + 0 + if req.skip_radix_cache_insert + else min( + req.kv.cache_protected_len, + req._compute_max_prefix_len(len(req.full_untruncated_fill_ids)), + ) + ) self._release_aborted_request(req) release_kv_cache(req, self.tree_cache) req.reset_for_retract() @@ -1510,6 +1521,9 @@ class SchedulerDisaggregationPrefillMixin: req.pending_bootstrap = True req.time_stats.reset_prefill_retry_time() req.advance_cache_request_handle() + # A fresh lookup budget for the new attempt, as after a retraction. + req.storage_prefetch_retry_attempts = 0 + req.storage_prefetch_last_match_len = yielded_prefix_len or None if req.prefill_attempt_count >= max_attempts: logger.info( f"Req {req.rid} exhausted optimistic prefill attempts " diff --git a/test/registered/disaggregation/test_disaggregation_optimistic_prefill.py b/test/registered/disaggregation/test_disaggregation_optimistic_prefill.py index 33194891f..dec9efdb5 100644 --- a/test/registered/disaggregation/test_disaggregation_optimistic_prefill.py +++ b/test/registered/disaggregation/test_disaggregation_optimistic_prefill.py @@ -1,3 +1,6 @@ +import os +import shutil +import tempfile import time import unittest import uuid @@ -16,7 +19,7 @@ from sglang.test.server_fixtures.disaggregation_fixture import ( ) from sglang.test.test_utils import DEFAULT_MODEL_NAME_FOR_TEST -register_cuda_ci(est_time=193, stage="base-b", runner_config="2-gpu-large") +register_cuda_ci(est_time=300, stage="base-b", runner_config="2-gpu-large") FORCE_RETRY_PROB = 0.1 @@ -37,18 +40,22 @@ def rid_that_forces_retry(prefix: str) -> str: class OptimisticPrefillRetryCounterMixin: - def _get_retry_counter(self) -> float: + def _get_counter_total(self, family_name: str) -> float: + """Sum of a prefill-side Prometheus counter across its label sets.""" response = requests.get(f"{self.prefill_url}/metrics") response.raise_for_status() total = 0.0 for family in text_string_to_metric_families(response.text): - if family.name != "sglang:num_prefill_retries": + if family.name != family_name: continue for sample in family.samples: - if sample.name == "sglang:num_prefill_retries_total": + if sample.name == f"{family_name}_total": total += sample.value return total + def _get_retry_counter(self) -> float: + return self._get_counter_total("sglang:num_prefill_retries") + def assert_retry_counter_increases(self, fn): before_retries = self._get_retry_counter() result = fn() @@ -202,5 +209,87 @@ class TestOptimisticPrefillFailure(PDDisaggregationServerBase): time.sleep(1) # trigger memory check +class TestOptimisticPrefillL3BufferWriteThrough( + OptimisticPrefillRetryCounterMixin, PDDisaggregationServerBase +): + """Optimistic prefill with buffer-only L3 (write-through, file backend). + + Small prefill and decode pools keep yielded prefixes evictable while their + retries wait for decode, so retries that recover the prefix from L3 run + under real load; the gsm8k score is the correctness check.""" + + @classmethod + def setUpClass(cls): + cls.hicache_dir = tempfile.mkdtemp(prefix="sglang-hicache-") + os.environ["SGLANG_HICACHE_FILE_BACKEND_STORAGE_DIR"] = cls.hicache_dir + super().setUpClass() + cls._force_retry_prob_was_set = ( + envs.SGLANG_TEST_FORCE_OPTIMISTIC_PREFILL_RETRY_PROB.is_set() + ) + cls._force_retry_prob_value = ( + envs.SGLANG_TEST_FORCE_OPTIMISTIC_PREFILL_RETRY_PROB.get() + ) + envs.SGLANG_TEST_FORCE_OPTIMISTIC_PREFILL_RETRY_PROB.set(FORCE_RETRY_PROB) + cls.model = DEFAULT_MODEL_NAME_FOR_TEST + cls.extra_prefill_args = [ + "--optimistic-prefill-attempts", + "2", + "--chunked-prefill-size", + "128", + "--max-total-tokens", + "16384", + "--enable-metrics", + "--enable-hierarchical-cache", + "--hicache-size", + "4", + "--hicache-host-memory-mode", + "buffer_only", + "--hicache-write-policy", + "write_through", + "--hicache-storage-backend", + "file", + "--hicache-storage-prefetch-policy", + "wait_complete", + ] + # A small decode pool gates bootstrap, so a yielded request waits long + # enough for the prefill pool above to evict its cached prefix. + cls.extra_decode_args = ["--max-total-tokens", "16384"] + cls.launch_all() + + @classmethod + def tearDownClass(cls): + try: + super().tearDownClass() + finally: + if getattr(cls, "_force_retry_prob_was_set", False): + envs.SGLANG_TEST_FORCE_OPTIMISTIC_PREFILL_RETRY_PROB.set( + cls._force_retry_prob_value + ) + else: + envs.SGLANG_TEST_FORCE_OPTIMISTIC_PREFILL_RETRY_PROB.clear() + os.environ.pop("SGLANG_HICACHE_FILE_BACKEND_STORAGE_DIR", None) + shutil.rmtree(cls.hicache_dir, ignore_errors=True) + + def test_gsm8k(self): + args = SimpleNamespace( + base_url=f"http://{self.base_host}:{self.lb_port}", + eval_name="gsm8k", + api="completion", + max_tokens=512, + num_examples=200, + num_threads=128, + ) + metrics = self.assert_retry_counter_increases(lambda: run_eval(args)) + print(f"Evaluation metrics: {metrics}") + self.assertGreater(metrics["score"], 0.62) + # Write-through published prefixes to L3; report what retries fetched back. + self.assertGreater(self._get_counter_total("sglang:backuped_tokens"), 0) + print( + "L3 prefetch hit tokens: " + f"{self._get_counter_total('sglang:storage_prefetch_hit_tokens')}" + ) + time.sleep(1) # trigger memory check + + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index ec239753f..d73d10a49 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -1961,6 +1961,48 @@ class TestHiCacheArgs(unittest.TestCase): with envs.SGLANG_UNIFIED_RADIX_TREE_CORE_BACKEND.override(backend): handle_hicache(args) + def test_optimistic_prefill_allows_only_exercised_hicache_modes(self): + common = { + "enable_hierarchical_cache": True, + "disaggregation_mode": "prefill", + "optimistic_prefill_attempts": 3, + } + cases = [ + ({"hicache_write_policy": "write_back"}, 3), + ( + { + "hicache_storage_backend": "file", + "hicache_host_memory_mode": "buffer_only", + "hicache_write_policy": "write_through", + }, + 3, + ), + ({"hicache_write_policy": "write_through"}, 0), + ( + { + "hicache_storage_backend": "file", + "hicache_host_memory_mode": "cache", + "hicache_write_policy": "write_through", + }, + 0, + ), + ( + { + "hicache_storage_backend": "file", + "hicache_host_memory_mode": "buffer_only", + "hicache_write_policy": "write_through_selective", + }, + 0, + ), + ] + for overrides, expected in cases: + with self.subTest(overrides=overrides): + args = ServerArgs(model_path="dummy", **common, **overrides) + serving_hook.handle_other_validations(args) + self.assertEqual( + resolution_result(args, "optimistic_prefill_attempts"), expected + ) + def test_hicache_io_backend_and_mem_layout_compatibility(self): cases = [ {