perf(disaggregation): slice decode waiting queue for prebuilt batches

This commit is contained in:
wxiwnd
2026-04-08 23:09:03 +08:00
parent 6237271fc9
commit 14c40f4e7d
2 changed files with 125 additions and 20 deletions

View File

@@ -1212,23 +1212,16 @@ class SchedulerDisaggregationDecodeMixin:
num_not_used_batch = batch_size - curr_batch_size
# pop req from waiting queue
can_run_list: List[Req] = []
waiting_queue: List[Req] = []
for i in range(len(self.waiting_queue)):
req = self.waiting_queue[i]
# we can only add at least `num_not_used_batch` new batch to the running queue
if i < num_not_used_batch:
can_run_list.append(req)
req.init_next_round_input(self.tree_cache)
else:
waiting_queue.append(req)
self.waiting_queue = waiting_queue
if len(can_run_list) == 0:
num_new_reqs = min(num_not_used_batch, len(self.waiting_queue))
if num_new_reqs <= 0:
return None
can_run_list = self.waiting_queue[:num_new_reqs]
self.waiting_queue = self.waiting_queue[num_new_reqs:]
for req in can_run_list:
req.init_next_round_input(self.tree_cache)
set_time_batch(can_run_list, "set_forward_entry_time")
# construct a schedule batch with those requests and mark as decode

View File

@@ -10,6 +10,7 @@ from sglang.srt.disaggregation.decode import (
DecodePreallocQueue,
DecodeRequest,
DecodeTransferQueue,
SchedulerDisaggregationDecodeMixin,
)
from sglang.srt.disaggregation.utils import ReqToMetadataIdxAllocator
from sglang.srt.managers.schedule_batch import FINISH_ABORT
@@ -33,12 +34,26 @@ class FakeReq:
self.prealloc_done = False
self.finished_reason = cast(Any, None)
self.cached_tokens = 0
self.init_next_round_calls = []
self.sampling_params = SimpleNamespace(max_new_tokens=0)
self.time_stats = SimpleNamespace(
set_bootstrap_done_time=lambda: None,
set_decode_transfer_queue_entry_time=lambda: None,
set_wait_queue_entry_time=lambda: None,
)
class FakeTimeStats:
def __init__(self):
self.forward_entry_time = None
def set_bootstrap_done_time(self):
return None
def set_decode_transfer_queue_entry_time(self):
return None
def set_wait_queue_entry_time(self):
return None
def set_forward_entry_time(self, ts):
self.forward_entry_time = ts
self.time_stats = FakeTimeStats()
def add_latency(self, stage):
self.latencies.append(stage)
@@ -46,6 +61,9 @@ class FakeReq:
def load_kv_cache(self, req_to_token_pool, token_to_kv_pool_allocator):
self.load_calls.append((req_to_token_pool, token_to_kv_pool_allocator))
def init_next_round_input(self, tree_cache):
self.init_next_round_calls.append(tree_cache)
class FakeReceiver:
def __init__(self, should_fail=False):
@@ -65,6 +83,18 @@ class FakeReceiver:
return None
class FakeBatch:
def __init__(self):
self.prepared = False
self.processed = []
def prepare_for_prebuilt(self):
self.prepared = True
def process_prebuilt(self, server_args, future_map):
self.processed.append((server_args, future_map))
class FakeItem:
def __init__(self, value):
self.value = value
@@ -490,6 +520,88 @@ class TestDecodeQueueCompaction(CustomTestCase):
self.assertEqual(success.metadata_buffer_index, 7)
self.assertEqual(queue.queue, [skipped, waiting, blocked, tail])
def test_get_new_prebuilt_batch_slices_waiting_queue_prefix(self):
scheduler = cast(Any, SimpleNamespace())
scheduler.grammar_manager = SimpleNamespace(
has_waiting_grammars=lambda: False,
get_ready_grammar_requests=lambda: [],
)
scheduler._add_request_to_queue = lambda req: None
scheduler.waiting_queue = [FakeReq(f"req-{i}", i) for i in range(5)]
scheduler.running_batch = SimpleNamespace(batch_size=lambda: 1)
scheduler.req_to_token_pool = SimpleNamespace(size=8)
scheduler.max_running_requests = 4
scheduler.tree_cache = object()
scheduler.token_to_kv_pool_allocator = object()
scheduler.model_config = object()
scheduler.enable_overlap = False
scheduler.spec_algorithm = object()
scheduler.server_args = object()
scheduler.future_map = object()
captured = {}
def fake_init_new(reqs, *args, **kwargs):
captured["reqs"] = list(reqs)
batch = FakeBatch()
captured["batch"] = batch
return batch
with patch(
"sglang.srt.disaggregation.decode.ScheduleBatch.init_new", fake_init_new
):
batch = SchedulerDisaggregationDecodeMixin.get_new_prebuilt_batch(scheduler)
self.assertIs(batch, captured["batch"])
self.assertEqual(
[req.rid for req in captured["reqs"]], ["req-0", "req-1", "req-2"]
)
self.assertEqual(
[req.rid for req in scheduler.waiting_queue], ["req-3", "req-4"]
)
for req in captured["reqs"]:
self.assertEqual(req.init_next_round_calls, [scheduler.tree_cache])
self.assertIsNotNone(req.time_stats.forward_entry_time)
self.assertTrue(captured["batch"].prepared)
self.assertEqual(
captured["batch"].processed,
[(scheduler.server_args, scheduler.future_map)],
)
def test_get_new_prebuilt_batch_keeps_waiting_queue_when_no_capacity(self):
scheduler = cast(Any, SimpleNamespace())
scheduler.grammar_manager = SimpleNamespace(
has_waiting_grammars=lambda: False,
get_ready_grammar_requests=lambda: [],
)
scheduler._add_request_to_queue = lambda req: None
scheduler.waiting_queue = [FakeReq(f"req-{i}", i) for i in range(3)]
scheduler.running_batch = SimpleNamespace(batch_size=lambda: 4)
scheduler.req_to_token_pool = SimpleNamespace(size=8)
scheduler.max_running_requests = 4
scheduler.tree_cache = object()
scheduler.token_to_kv_pool_allocator = object()
scheduler.model_config = object()
scheduler.enable_overlap = False
scheduler.spec_algorithm = object()
scheduler.server_args = object()
scheduler.future_map = object()
with patch(
"sglang.srt.disaggregation.decode.ScheduleBatch.init_new",
side_effect=AssertionError(
"init_new should not be called without capacity"
),
):
batch = SchedulerDisaggregationDecodeMixin.get_new_prebuilt_batch(scheduler)
self.assertIsNone(batch)
self.assertEqual(
[req.rid for req in scheduler.waiting_queue], ["req-0", "req-1", "req-2"]
)
for req in scheduler.waiting_queue:
self.assertEqual(req.init_next_round_calls, [])
if __name__ == "__main__":
unittest.main()