[PD] Enable optimistic prefill with buffer-only L3 write-through HiCache (#40043)
Co-authored-by: cctry <csycfl@gmail.com>
This commit is contained in:
@@ -466,13 +466,20 @@ def handle_other_validations(server_args: Any):
|
|||||||
"_handle_other_validations",
|
"_handle_other_validations",
|
||||||
optimistic_prefill_attempts=0,
|
optimistic_prefill_attempts=0,
|
||||||
)
|
)
|
||||||
elif cfg.enable_hierarchical_cache and (
|
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
|
cfg.hicache_storage_backend is not None
|
||||||
or cfg.hicache_write_policy != "write_back"
|
and cfg.hicache_host_memory_mode == "buffer_only"
|
||||||
|
and cfg.hicache_write_policy == "write_through"
|
||||||
|
)
|
||||||
):
|
):
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Optimistic prefill only supports L2 hierarchical cache "
|
"Optimistic prefill supports L2 write-back or L3 buffer-only "
|
||||||
"with write-back policy"
|
"write-through hierarchical cache"
|
||||||
)
|
)
|
||||||
declare_resolution(
|
declare_resolution(
|
||||||
server_args,
|
server_args,
|
||||||
|
|||||||
@@ -1496,6 +1496,17 @@ class SchedulerDisaggregationPrefillMixin:
|
|||||||
"""Release KV cache and requeue an optimistic prefill request."""
|
"""Release KV cache and requeue an optimistic prefill request."""
|
||||||
max_attempts = get_disagg().optimistic_prefill_attempts
|
max_attempts = get_disagg().optimistic_prefill_attempts
|
||||||
maybe_cache_unfinished_req(req, self.tree_cache)
|
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)
|
self._release_aborted_request(req)
|
||||||
release_kv_cache(req, self.tree_cache)
|
release_kv_cache(req, self.tree_cache)
|
||||||
req.reset_for_retract()
|
req.reset_for_retract()
|
||||||
@@ -1510,6 +1521,9 @@ class SchedulerDisaggregationPrefillMixin:
|
|||||||
req.pending_bootstrap = True
|
req.pending_bootstrap = True
|
||||||
req.time_stats.reset_prefill_retry_time()
|
req.time_stats.reset_prefill_retry_time()
|
||||||
req.advance_cache_request_handle()
|
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:
|
if req.prefill_attempt_count >= max_attempts:
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Req {req.rid} exhausted optimistic prefill attempts "
|
f"Req {req.rid} exhausted optimistic prefill attempts "
|
||||||
|
|||||||
@@ -1,3 +1,6 @@
|
|||||||
|
import os
|
||||||
|
import shutil
|
||||||
|
import tempfile
|
||||||
import time
|
import time
|
||||||
import unittest
|
import unittest
|
||||||
import uuid
|
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
|
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
|
FORCE_RETRY_PROB = 0.1
|
||||||
@@ -37,18 +40,22 @@ def rid_that_forces_retry(prefix: str) -> str:
|
|||||||
|
|
||||||
|
|
||||||
class OptimisticPrefillRetryCounterMixin:
|
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 = requests.get(f"{self.prefill_url}/metrics")
|
||||||
response.raise_for_status()
|
response.raise_for_status()
|
||||||
total = 0.0
|
total = 0.0
|
||||||
for family in text_string_to_metric_families(response.text):
|
for family in text_string_to_metric_families(response.text):
|
||||||
if family.name != "sglang:num_prefill_retries":
|
if family.name != family_name:
|
||||||
continue
|
continue
|
||||||
for sample in family.samples:
|
for sample in family.samples:
|
||||||
if sample.name == "sglang:num_prefill_retries_total":
|
if sample.name == f"{family_name}_total":
|
||||||
total += sample.value
|
total += sample.value
|
||||||
return total
|
return total
|
||||||
|
|
||||||
|
def _get_retry_counter(self) -> float:
|
||||||
|
return self._get_counter_total("sglang:num_prefill_retries")
|
||||||
|
|
||||||
def assert_retry_counter_increases(self, fn):
|
def assert_retry_counter_increases(self, fn):
|
||||||
before_retries = self._get_retry_counter()
|
before_retries = self._get_retry_counter()
|
||||||
result = fn()
|
result = fn()
|
||||||
@@ -202,5 +209,87 @@ class TestOptimisticPrefillFailure(PDDisaggregationServerBase):
|
|||||||
time.sleep(1) # trigger memory check
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -1961,6 +1961,48 @@ class TestHiCacheArgs(unittest.TestCase):
|
|||||||
with envs.SGLANG_UNIFIED_RADIX_TREE_CORE_BACKEND.override(backend):
|
with envs.SGLANG_UNIFIED_RADIX_TREE_CORE_BACKEND.override(backend):
|
||||||
handle_hicache(args)
|
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):
|
def test_hicache_io_backend_and_mem_layout_compatibility(self):
|
||||||
cases = [
|
cases = [
|
||||||
{
|
{
|
||||||
|
|||||||
Reference in New Issue
Block a user