Fix overlap prebuilt row reuse race (#35748)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user