Fix overlap prebuilt row reuse race (#35748)

This commit is contained in:
jasonjk-park
2026-08-21 02:00:13 -07:00
committed by GitHub
parent 896acc8860
commit 0db2c53dec
3 changed files with 53 additions and 17 deletions
@@ -2632,6 +2632,10 @@ class SchedulerDisaggregationDecodeMixin:
# construct fake completed prefill
new_batch.prepare_for_prebuilt()
if self.enable_overlap:
# A finished request can still have one redundant forward in flight.
# Drain it before a prebuilt request seeds a potentially reused row.
self.schedule_stream.wait_stream(self.forward_stream)
new_batch.process_prebuilt(self.future_map)
return new_batch
+1 -7
View File
@@ -523,14 +523,8 @@ class FutureMap:
self.confidence_relay.scatter(indices, confidence)
# Only spec_v2 needs the event; it gates the seq_lens D2H on the private stream.
if self.spec_algo.is_some():
device_module = torch.get_device_module(self.device)
if self.publish_ready is None:
self.publish_ready = device_module.Event()
else:
# Chain the records: event fire implies every prior publish is
# visible, so an off-forward-stream publish (PD-decode prebuilt
# seeding) cannot drop the in-flight forward's fence.
device_module.current_stream().wait_event(self.publish_ready)
self.publish_ready = torch.get_device_module(self.device).Event()
self.publish_ready.record()
self._publish_fresh = True
if publish_confidence:
@@ -413,29 +413,37 @@ class TestCommonKVManagerPrefillRecompute(unittest.TestCase):
self.assertEqual(mgr.failure_records, {})
class TestDecodePrebuiltPriority(unittest.TestCase):
def test_waiting_queue_is_sorted_before_prebuilt_selection(self):
class TestDecodePrebuilt(unittest.TestCase):
def _new_scheduler(self, *, enable_overlap: bool) -> Scheduler:
scheduler = Scheduler.__new__(Scheduler)
scheduler.grammar_manager = MagicMock()
scheduler.grammar_manager.has_waiting_grammars.return_value = False
original_waiting_queue = [MagicMock(rid="low"), MagicMock(rid="high")]
scheduler.waiting_queue = original_waiting_queue
scheduler.waiting_queue[0].priority = 1
scheduler.waiting_queue[1].priority = 10
scheduler.enable_priority_scheduling = True
scheduler.waiting_queue = []
scheduler.enable_priority_scheduling = False
scheduler.running_batch = MagicMock()
scheduler.running_batch.batch_size.return_value = 0
scheduler.req_to_token_pool = MagicMock(size=1)
scheduler.token_to_kv_pool_allocator = MagicMock()
scheduler.tree_cache = MagicMock()
scheduler.model_config = MagicMock()
scheduler.enable_overlap = False
scheduler.enable_overlap = enable_overlap
scheduler.spec_algorithm = MagicMock()
scheduler.max_running_requests = 1
scheduler.future_map = MagicMock()
scheduler.policy = MagicMock()
scheduler.policy.calc_priority.side_effect = lambda waiting_queue, _: (
waiting_queue.sort(key=lambda req: -req.priority)
scheduler.schedule_stream = MagicMock()
scheduler.forward_stream = MagicMock()
return scheduler
def test_waiting_queue_is_sorted_before_prebuilt_selection(self):
scheduler = self._new_scheduler(enable_overlap=False)
original_waiting_queue = [MagicMock(rid="low"), MagicMock(rid="high")]
scheduler.waiting_queue = original_waiting_queue
scheduler.waiting_queue[0].priority = 1
scheduler.waiting_queue[1].priority = 10
scheduler.enable_priority_scheduling = True
scheduler.policy.calc_priority.side_effect = (
lambda waiting_queue, _: waiting_queue.sort(key=lambda req: -req.priority)
)
new_batch = MagicMock()
@@ -459,6 +467,36 @@ class TestDecodePrebuiltPriority(unittest.TestCase):
self.assertEqual([req.rid for req in selected_reqs], ["high"])
self.assertEqual([req.rid for req in scheduler.waiting_queue], ["low"])
def test_overlap_waits_for_forward_before_processing_prebuilt(self):
scheduler = self._new_scheduler(enable_overlap=True)
scheduler.waiting_queue = [MagicMock(rid="request")]
call_order = []
new_batch = MagicMock()
new_batch.prepare_for_prebuilt.side_effect = lambda: call_order.append(
"prepare"
)
scheduler.schedule_stream.wait_stream.side_effect = lambda _: call_order.append(
"wait"
)
new_batch.process_prebuilt.side_effect = lambda *_: call_order.append("process")
with patch(
"sglang.srt.disaggregation.decode.ScheduleBatch.init_new",
return_value=new_batch,
), get_context().override_server_args(
disaggregation_decode_enable_radix_cache=False
):
ret = SchedulerDisaggregationDecodeMixin.get_new_prebuilt_batch(
scheduler, scheduler.running_batch
)
self.assertIs(ret, new_batch)
scheduler.schedule_stream.wait_stream.assert_called_once_with(
scheduler.forward_stream
)
self.assertEqual(call_order, ["prepare", "wait", "process"])
if __name__ == "__main__":
unittest.main()