perf(disaggregation): slice decode waiting queue for prebuilt batches
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user