[Metrics] Discount queued prefill load by recent cache hits when waiting-queue matching is off (#35248)

This commit is contained in:
Hanming Lu
2026-08-18 11:27:21 -07:00
committed by GitHub
parent 7dcaf11987
commit 526af15845
7 changed files with 131 additions and 3 deletions
@@ -0,0 +1,70 @@
import unittest
from types import SimpleNamespace
from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.managers.schedule_policy import (
CacheAgnosticPolicy,
CacheAwarePolicy,
SchedulePolicy,
)
from sglang.srt.managers.scheduler_components.load_inquirer import (
SchedulerLoadInquirer,
)
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
class TestSchedulePolicyWaitingQueueMatching(unittest.TestCase):
def make_policy(self, policy, supports_fast_match_prefix):
schedule_policy = object.__new__(SchedulePolicy)
schedule_policy.policy = policy
schedule_policy.tree_cache = SimpleNamespace(
supports_fast_match_prefix=lambda: supports_fast_match_prefix
)
return schedule_policy
def test_cache_agnostic_policy_requires_fast_matching(self):
policy = self.make_policy(CacheAgnosticPolicy.FCFS, False)
self.assertFalse(policy.waiting_queue_prefix_matched([]))
policy.tree_cache = SimpleNamespace(supports_fast_match_prefix=lambda: True)
self.assertTrue(policy.waiting_queue_prefix_matched([]))
def test_lpm_queue_limit_can_disable_matching(self):
policy = self.make_policy(CacheAwarePolicy.LPM, False)
self.assertTrue(policy.waiting_queue_prefix_matched([None] * 128))
self.assertFalse(policy.waiting_queue_prefix_matched([None] * 129))
class TestSchedulerLoadInquirer(unittest.TestCase):
def make_inquirer(self, waiting_queue_prefix_matched):
waiting_req = SimpleNamespace(seqlen=100, num_matched_prefix_tokens=20)
chunked_req = SimpleNamespace(seqlen=50, prefix_indices=range(10))
return SimpleNamespace(
disaggregation_mode=DisaggregationMode.NULL,
get_waiting_queue=lambda: [waiting_req],
waiting_queue_prefix_matched=lambda: waiting_queue_prefix_matched,
get_chunked_req=lambda: chunked_req,
get_recent_cache_hit_rate=lambda: 0.75,
)
def test_waiting_tokens_are_estimated_when_prefix_matching_is_skipped(self):
inquirer = self.make_inquirer(waiting_queue_prefix_matched=False)
self.assertEqual(
SchedulerLoadInquirer.get_num_waiting_uncached_tokens(inquirer),
65,
)
def test_waiting_tokens_use_exact_match_when_prefix_matching_is_done(self):
inquirer = self.make_inquirer(waiting_queue_prefix_matched=True)
self.assertEqual(
SchedulerLoadInquirer.get_num_waiting_uncached_tokens(inquirer),
120,
)
if __name__ == "__main__":
unittest.main()
@@ -13,6 +13,7 @@ from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.managers.scheduler_components.metrics_reporter import (
PrefillStats,
SchedulerMetricsReporter,
_CacheHitRateWindow,
)
from sglang.test.test_utils import CustomTestCase
@@ -136,6 +137,12 @@ class TestForwardPassMetrics(unittest.TestCase):
self.reporter = _make_reporter(self, self.scheduler)
self.scheduler.enable_fpm = True
def test_cache_hit_rate_window_keeps_last_15s_of_tokens(self):
window = _CacheHitRateWindow()
self.assertEqual(window.add(hit_tokens=20, total_tokens=100, now=0.0), 0.2)
self.assertEqual(window.add(hit_tokens=80, total_tokens=100, now=10.0), 0.5)
self.assertEqual(window.add(hit_tokens=90, total_tokens=100, now=15.0), 0.85)
def _make_batch(self, **overrides):
defaults = dict(
forward_mode=_FakeForwardMode(),