[Test] Fix scheduler fixtures after prefill burst counting (#40411)
Co-authored-by: Mohammad Angkad <mohammad.angkad@radixark.ai>
This commit is contained in:
co-authored by
Mohammad Angkad
parent
99a44c88d4
commit
22f02cc339
@@ -528,6 +528,7 @@ def test_pdmux_split_prefill_schedules_auxiliary_output_copy():
|
|||||||
is_prebuilt=lambda: False,
|
is_prebuilt=lambda: False,
|
||||||
is_split_prefill=lambda: True,
|
is_split_prefill=lambda: True,
|
||||||
),
|
),
|
||||||
|
split_index=0,
|
||||||
reqs=[],
|
reqs=[],
|
||||||
req_pool_indices=torch.tensor([3]),
|
req_pool_indices=torch.tensor([3]),
|
||||||
input_ids=torch.tensor([5]),
|
input_ids=torch.tensor([5]),
|
||||||
|
|||||||
@@ -24,6 +24,8 @@ from sglang.test.test_utils import CustomTestCase
|
|||||||
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||||
|
|
||||||
LAUNCH_TIMESTAMPS = (0.0, 0.125, 1.0, 1.125)
|
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"
|
PP_MODULE = "sglang.srt.managers.scheduler_pp_mixin"
|
||||||
PDMUX_MODULE = "sglang.srt.multiplex.multiplexing_mixin"
|
PDMUX_MODULE = "sglang.srt.multiplex.multiplexing_mixin"
|
||||||
|
|
||||||
@@ -32,6 +34,11 @@ class _BeforeModelForward(Exception):
|
|||||||
pass
|
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():
|
def load_mlx_scheduler_module():
|
||||||
# Run the real Python loop on CPU without importing a Metal runtime. Use a
|
# 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.
|
# private module name so this cannot replace the installed MLX scheduler.
|
||||||
@@ -242,17 +249,61 @@ class TestSchedulerIdleStepCounters(CustomTestCase):
|
|||||||
2 if chained else 0,
|
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):
|
def make_batches(self, mode):
|
||||||
return [
|
return [
|
||||||
ScheduleBatch(
|
self.make_batch(mode, launch_ts=launch_ts)
|
||||||
reqs=[
|
|
||||||
SimpleNamespace(rid="request", finished=Mock(return_value=False))
|
|
||||||
],
|
|
||||||
forward_mode=mode,
|
|
||||||
spec_algorithm=SpeculativeAlgorithm.NONE,
|
|
||||||
launch_ts=launch_ts,
|
|
||||||
extend_num_tokens=1024,
|
|
||||||
)
|
|
||||||
for launch_ts in LAUNCH_TIMESTAMPS
|
for launch_ts in LAUNCH_TIMESTAMPS
|
||||||
]
|
]
|
||||||
|
|
||||||
@@ -261,20 +312,21 @@ class TestSchedulerIdleStepCounters(CustomTestCase):
|
|||||||
observed_iters = []
|
observed_iters = []
|
||||||
|
|
||||||
def run_batch(batch, pp_proxy_tensors=None):
|
def run_batch(batch, pp_proxy_tensors=None):
|
||||||
# Exercise the real timestamp, iteration, and flag handoff. Only
|
# Exercise the real timestamp, iteration, and flag handoff.
|
||||||
# model execution is stopped, at the scripted pre-forward hook.
|
self.launch_batch(
|
||||||
with patch(
|
scheduler, batch, batch.launch_ts, pp_proxy_tensors=pp_proxy_tensors
|
||||||
"sglang.srt.managers.scheduler.time.monotonic",
|
)
|
||||||
return_value=batch.launch_ts,
|
|
||||||
):
|
|
||||||
with self.assertRaises(_BeforeModelForward):
|
|
||||||
Scheduler.run_batch(scheduler, batch, pp_proxy_tensors)
|
|
||||||
return GenerationBatchResult()
|
return GenerationBatchResult()
|
||||||
|
|
||||||
def process_batch_result(batch, result):
|
def process_batch_result(batch, result):
|
||||||
observed_idle_flags.append(batch.after_idle_gap)
|
observed_idle_flags.append(batch.after_idle_gap)
|
||||||
observed_iters.append(batch.forward_iter)
|
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.run_batch = run_batch
|
||||||
scheduler.process_batch_result = process_batch_result
|
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(observed_idle_flags, [False, False, after_idle, False])
|
||||||
self.assertEqual(scheduler.forward_ct, 4)
|
self.assertEqual(scheduler.forward_ct, 4)
|
||||||
self.assertEqual(observed_iters, [1, 2, 3, 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:
|
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(
|
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:
|
else:
|
||||||
self.assertEqual(scheduler.decode_moment_totals[0], expected_samples)
|
# Decode keeps launch-to-launch cadence and drops non-contiguous steps.
|
||||||
self.assertEqual(scheduler.decode_moment_totals[2], expected_busy_us)
|
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):
|
def make_scheduler(self, schedule):
|
||||||
scheduler = Scheduler.__new__(Scheduler)
|
scheduler = Scheduler.__new__(Scheduler)
|
||||||
scheduler._engine_paused = False
|
scheduler._engine_paused = False
|
||||||
scheduler._sched_idled = False
|
scheduler._sched_idled = False
|
||||||
scheduler._prev_step = None
|
scheduler._prev_step = None
|
||||||
|
scheduler._prev_prefill_end_ts = None
|
||||||
scheduler.forward_ct = 0
|
scheduler.forward_ct = 0
|
||||||
scheduler.processed_tokens_counter = 0
|
scheduler.processed_tokens_counter = 0
|
||||||
scheduler.spec_algorithm = SpeculativeAlgorithm.NONE
|
scheduler.spec_algorithm = SpeculativeAlgorithm.NONE
|
||||||
|
|||||||
Reference in New Issue
Block a user