From 22f02cc3399dcd59380fc546de8d3c89fba6fa3a Mon Sep 17 00:00:00 2001 From: Mohammad Miadh Angkad <176301910+mmangkad@users.noreply.github.com> Date: Sun, 20 Sep 2026 07:08:04 +0000 Subject: [PATCH] [Test] Fix scheduler fixtures after prefill burst counting (#40411) Co-authored-by: Mohammad Angkad --- .../unit/managers/test_auxiliary_output.py | 1 + .../test_disagg_idle_step_counters.py | 144 ++++++++++++++---- 2 files changed, 115 insertions(+), 30 deletions(-) diff --git a/test/registered/unit/managers/test_auxiliary_output.py b/test/registered/unit/managers/test_auxiliary_output.py index 4fed959b8..0799b41dc 100644 --- a/test/registered/unit/managers/test_auxiliary_output.py +++ b/test/registered/unit/managers/test_auxiliary_output.py @@ -528,6 +528,7 @@ def test_pdmux_split_prefill_schedules_auxiliary_output_copy(): is_prebuilt=lambda: False, is_split_prefill=lambda: True, ), + split_index=0, reqs=[], req_pool_indices=torch.tensor([3]), input_ids=torch.tensor([5]), diff --git a/test/registered/unit/managers/test_disagg_idle_step_counters.py b/test/registered/unit/managers/test_disagg_idle_step_counters.py index 140d0f379..d6aa44aa4 100644 --- a/test/registered/unit/managers/test_disagg_idle_step_counters.py +++ b/test/registered/unit/managers/test_disagg_idle_step_counters.py @@ -24,6 +24,8 @@ from sglang.test.test_utils import CustomTestCase register_cpu_ci(est_time=1, suite="base-a-test-cpu") LAUNCH_TIMESTAMPS = (0.0, 0.125, 1.0, 1.125) +# Prefill busy time is charged launch -> result, so the result clock matters too. +RESULT_TIMESTAMPS = tuple(ts + 0.1 for ts in LAUNCH_TIMESTAMPS) PP_MODULE = "sglang.srt.managers.scheduler_pp_mixin" PDMUX_MODULE = "sglang.srt.multiplex.multiplexing_mixin" @@ -32,6 +34,11 @@ class _BeforeModelForward(Exception): pass +# The scheduler truncates once per recorded step, not once over the total. +def total_us(intervals): + return sum(int(interval * 1e6) for interval in intervals) + + def load_mlx_scheduler_module(): # Run the real Python loop on CPU without importing a Metal runtime. Use a # private module name so this cannot replace the installed MLX scheduler. @@ -242,17 +249,61 @@ class TestSchedulerIdleStepCounters(CustomTestCase): 2 if chained else 0, ) + def test_split_prefill_is_charged_once_from_its_first_chunk(self): + # A split prefill spans several run_batch calls but yields one result, + # so only its first chunk may set the burst's start and its idle flag. + scheduler = self.make_scheduler([]) + scheduler._sched_idled = True + batch = self.make_batch(ForwardMode.SPLIT_PREFILL) + for split_index, launch_ts in enumerate((0.0, 0.2, 0.4)): + batch.split_index = split_index + self.launch_batch(scheduler, batch, launch_ts) + + self.assertEqual(scheduler.forward_ct, 3) + self.assertEqual(batch.forward_iter, 3) + self.assertEqual(batch.split_prefill_start, (1, 0.0)) + self.assertTrue(batch.after_idle_gap) + + self.record_result(scheduler, batch, 0.5) + # One charge for the whole burst: first chunk's launch to the result. + self.assertEqual(scheduler.total_prefill_busy_us, total_us([0.5 - 0.0])) + self.assertEqual(scheduler.total_prefill_uncached_tokens, 1024) + + def test_a_prefill_is_not_charged_for_an_overlapping_earlier_one(self): + # An overlapped launch can precede the previous prefill's result; the + # span already charged to that prefill must not be charged twice. + scheduler = self.make_scheduler([]) + prefill = self.make_batch(ForwardMode.EXTEND) + decode = self.make_batch(ForwardMode.DECODE, extend_num_tokens=None) + next_prefill = self.make_batch(ForwardMode.EXTEND) + # Both later batches launch while the first prefill is still in flight. + self.launch_batch(scheduler, prefill, 0.0) + self.launch_batch(scheduler, decode, 0.5) + self.launch_batch(scheduler, next_prefill, 0.6) + + self.record_result(scheduler, prefill, 1.0) + self.record_result(scheduler, decode, 1.1) + self.record_result(scheduler, next_prefill, 1.5) + + # 0.6 -> 1.0 already belongs to the first prefill, so the second one is + # charged from that result rather than from its own launch. + self.assertEqual(scheduler.total_prefill_busy_us, total_us([1.0, 0.5])) + self.assertEqual(scheduler.total_prefill_uncached_tokens, 2 * 1024) + # The decode broke contiguity, so it contributes no timing sample. + self.assertEqual(scheduler.decode_moment_totals[0], 0) + + def make_batch(self, mode, *, launch_ts=None, extend_num_tokens=1024): + return ScheduleBatch( + reqs=[SimpleNamespace(rid="request", finished=Mock(return_value=False))], + forward_mode=mode, + spec_algorithm=SpeculativeAlgorithm.NONE, + launch_ts=launch_ts, + extend_num_tokens=extend_num_tokens, + ) + def make_batches(self, mode): return [ - ScheduleBatch( - reqs=[ - SimpleNamespace(rid="request", finished=Mock(return_value=False)) - ], - forward_mode=mode, - spec_algorithm=SpeculativeAlgorithm.NONE, - launch_ts=launch_ts, - extend_num_tokens=1024, - ) + self.make_batch(mode, launch_ts=launch_ts) for launch_ts in LAUNCH_TIMESTAMPS ] @@ -261,20 +312,21 @@ class TestSchedulerIdleStepCounters(CustomTestCase): observed_iters = [] def run_batch(batch, pp_proxy_tensors=None): - # Exercise the real timestamp, iteration, and flag handoff. Only - # model execution is stopped, at the scripted pre-forward hook. - with patch( - "sglang.srt.managers.scheduler.time.monotonic", - return_value=batch.launch_ts, - ): - with self.assertRaises(_BeforeModelForward): - Scheduler.run_batch(scheduler, batch, pp_proxy_tensors) + # Exercise the real timestamp, iteration, and flag handoff. + self.launch_batch( + scheduler, batch, batch.launch_ts, pp_proxy_tensors=pp_proxy_tensors + ) return GenerationBatchResult() def process_batch_result(batch, result): observed_idle_flags.append(batch.after_idle_gap) observed_iters.append(batch.forward_iter) - scheduler._record_step_counters(batch, result) + self.record_result( + scheduler, + batch, + RESULT_TIMESTAMPS[batch.forward_iter - 1], + result=result, + ) scheduler.run_batch = run_batch scheduler.process_batch_result = process_batch_result @@ -287,28 +339,60 @@ class TestSchedulerIdleStepCounters(CustomTestCase): self.assertEqual(observed_idle_flags, [False, False, after_idle, False]) self.assertEqual(scheduler.forward_ct, 4) self.assertEqual(observed_iters, [1, 2, 3, 4]) - expected_intervals = [ - LAUNCH_TIMESTAMPS[1] - LAUNCH_TIMESTAMPS[0], - LAUNCH_TIMESTAMPS[3] - LAUNCH_TIMESTAMPS[2], - ] - if not after_idle: - expected_intervals.append(LAUNCH_TIMESTAMPS[2] - LAUNCH_TIMESTAMPS[1]) - expected_samples = len(expected_intervals) - expected_busy_us = round(sum(expected_intervals) * 1_000_000) if mode == ForwardMode.EXTEND: - self.assertEqual(scheduler.total_prefill_busy_us, expected_busy_us) + # A prefill is charged from the previous prefill's result, or from + # its own launch when none applies -- the first one, or after a gap. + expected_intervals = [ + RESULT_TIMESTAMPS[0] - LAUNCH_TIMESTAMPS[0], + RESULT_TIMESTAMPS[1] - RESULT_TIMESTAMPS[0], + RESULT_TIMESTAMPS[2] + - (LAUNCH_TIMESTAMPS[2] if after_idle else RESULT_TIMESTAMPS[1]), + RESULT_TIMESTAMPS[3] - RESULT_TIMESTAMPS[2], + ] self.assertEqual( - scheduler.total_prefill_uncached_tokens, expected_samples * 1024 + scheduler.total_prefill_busy_us, total_us(expected_intervals) + ) + # Every prefill contributes its tokens; only the span is gap-aware. + self.assertEqual( + scheduler.total_prefill_uncached_tokens, + len(LAUNCH_TIMESTAMPS) * 1024, ) else: - self.assertEqual(scheduler.decode_moment_totals[0], expected_samples) - self.assertEqual(scheduler.decode_moment_totals[2], expected_busy_us) + # Decode keeps launch-to-launch cadence and drops non-contiguous steps. + expected_intervals = [ + LAUNCH_TIMESTAMPS[1] - LAUNCH_TIMESTAMPS[0], + LAUNCH_TIMESTAMPS[3] - LAUNCH_TIMESTAMPS[2], + ] + if not after_idle: + expected_intervals.append(LAUNCH_TIMESTAMPS[2] - LAUNCH_TIMESTAMPS[1]) + self.assertEqual(scheduler.decode_moment_totals[0], len(expected_intervals)) + self.assertEqual( + scheduler.decode_moment_totals[2], total_us(expected_intervals) + ) + + def launch_batch(self, scheduler, batch, launch_ts, *, pp_proxy_tensors=None): + # Only model execution is stopped, at the scripted pre-forward hook. + with patch( + "sglang.srt.managers.scheduler.time.monotonic", return_value=launch_ts + ): + with self.assertRaises(_BeforeModelForward): + Scheduler.run_batch(scheduler, batch, pp_proxy_tensors) + + def record_result(self, scheduler, batch, result_ts, *, result=None): + # Prefill accounting reads the clock again when the result lands. + with patch( + "sglang.srt.managers.scheduler.time.monotonic", return_value=result_ts + ): + scheduler._record_step_counters( + batch, GenerationBatchResult() if result is None else result + ) def make_scheduler(self, schedule): scheduler = Scheduler.__new__(Scheduler) scheduler._engine_paused = False scheduler._sched_idled = False scheduler._prev_step = None + scheduler._prev_prefill_end_ts = None scheduler.forward_ct = 0 scheduler.processed_tokens_counter = 0 scheduler.spec_algorithm = SpeculativeAlgorithm.NONE