Fix overlap prebuilt row reuse race (#35748)
This commit is contained in:
@@ -2632,6 +2632,10 @@ class SchedulerDisaggregationDecodeMixin:
|
|||||||
|
|
||||||
# construct fake completed prefill
|
# construct fake completed prefill
|
||||||
new_batch.prepare_for_prebuilt()
|
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)
|
new_batch.process_prebuilt(self.future_map)
|
||||||
|
|
||||||
return new_batch
|
return new_batch
|
||||||
|
|||||||
@@ -523,14 +523,8 @@ class FutureMap:
|
|||||||
self.confidence_relay.scatter(indices, confidence)
|
self.confidence_relay.scatter(indices, confidence)
|
||||||
# Only spec_v2 needs the event; it gates the seq_lens D2H on the private stream.
|
# Only spec_v2 needs the event; it gates the seq_lens D2H on the private stream.
|
||||||
if self.spec_algo.is_some():
|
if self.spec_algo.is_some():
|
||||||
device_module = torch.get_device_module(self.device)
|
|
||||||
if self.publish_ready is None:
|
if self.publish_ready is None:
|
||||||
self.publish_ready = device_module.Event()
|
self.publish_ready = torch.get_device_module(self.device).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.record()
|
self.publish_ready.record()
|
||||||
self._publish_fresh = True
|
self._publish_fresh = True
|
||||||
if publish_confidence:
|
if publish_confidence:
|
||||||
|
|||||||
@@ -413,29 +413,37 @@ class TestCommonKVManagerPrefillRecompute(unittest.TestCase):
|
|||||||
self.assertEqual(mgr.failure_records, {})
|
self.assertEqual(mgr.failure_records, {})
|
||||||
|
|
||||||
|
|
||||||
class TestDecodePrebuiltPriority(unittest.TestCase):
|
class TestDecodePrebuilt(unittest.TestCase):
|
||||||
def test_waiting_queue_is_sorted_before_prebuilt_selection(self):
|
def _new_scheduler(self, *, enable_overlap: bool) -> Scheduler:
|
||||||
scheduler = Scheduler.__new__(Scheduler)
|
scheduler = Scheduler.__new__(Scheduler)
|
||||||
scheduler.grammar_manager = MagicMock()
|
scheduler.grammar_manager = MagicMock()
|
||||||
scheduler.grammar_manager.has_waiting_grammars.return_value = False
|
scheduler.grammar_manager.has_waiting_grammars.return_value = False
|
||||||
original_waiting_queue = [MagicMock(rid="low"), MagicMock(rid="high")]
|
scheduler.waiting_queue = []
|
||||||
scheduler.waiting_queue = original_waiting_queue
|
scheduler.enable_priority_scheduling = False
|
||||||
scheduler.waiting_queue[0].priority = 1
|
|
||||||
scheduler.waiting_queue[1].priority = 10
|
|
||||||
scheduler.enable_priority_scheduling = True
|
|
||||||
scheduler.running_batch = MagicMock()
|
scheduler.running_batch = MagicMock()
|
||||||
scheduler.running_batch.batch_size.return_value = 0
|
scheduler.running_batch.batch_size.return_value = 0
|
||||||
scheduler.req_to_token_pool = MagicMock(size=1)
|
scheduler.req_to_token_pool = MagicMock(size=1)
|
||||||
scheduler.token_to_kv_pool_allocator = MagicMock()
|
scheduler.token_to_kv_pool_allocator = MagicMock()
|
||||||
scheduler.tree_cache = MagicMock()
|
scheduler.tree_cache = MagicMock()
|
||||||
scheduler.model_config = MagicMock()
|
scheduler.model_config = MagicMock()
|
||||||
scheduler.enable_overlap = False
|
scheduler.enable_overlap = enable_overlap
|
||||||
scheduler.spec_algorithm = MagicMock()
|
scheduler.spec_algorithm = MagicMock()
|
||||||
scheduler.max_running_requests = 1
|
scheduler.max_running_requests = 1
|
||||||
scheduler.future_map = MagicMock()
|
scheduler.future_map = MagicMock()
|
||||||
scheduler.policy = MagicMock()
|
scheduler.policy = MagicMock()
|
||||||
scheduler.policy.calc_priority.side_effect = lambda waiting_queue, _: (
|
scheduler.schedule_stream = MagicMock()
|
||||||
waiting_queue.sort(key=lambda req: -req.priority)
|
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()
|
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 selected_reqs], ["high"])
|
||||||
self.assertEqual([req.rid for req in scheduler.waiting_queue], ["low"])
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user