[PD] Enable optimistic prefill with buffer-only L3 write-through HiCache (#40043)

Co-authored-by: cctry <csycfl@gmail.com>
This commit is contained in:
Zhiqiang Xie
2026-09-18 16:04:15 -07:00
committed by GitHub
co-authored by cctry
parent 5e4b94b134
commit ceb1d2e580
4 changed files with 161 additions and 9 deletions
+12 -5
View File
@@ -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,
@@ -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 "
@@ -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()
@@ -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 = [
{