perf(disaggregation): compact decode queues in one pass
This commit is contained in:
@@ -0,0 +1,495 @@
|
||||
"""Unit tests for decode queue one-pass compaction."""
|
||||
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from typing import Any, cast
|
||||
from unittest.mock import patch
|
||||
|
||||
from sglang.srt.disaggregation.base import KVPoll
|
||||
from sglang.srt.disaggregation.decode import (
|
||||
DecodePreallocQueue,
|
||||
DecodeRequest,
|
||||
DecodeTransferQueue,
|
||||
)
|
||||
from sglang.srt.disaggregation.utils import ReqToMetadataIdxAllocator
|
||||
from sglang.srt.managers.schedule_batch import FINISH_ABORT
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cpu_ci(est_time=8, suite="stage-a-test-cpu")
|
||||
|
||||
|
||||
class FakeReq:
|
||||
def __init__(self, rid, bootstrap_room):
|
||||
self.rid = rid
|
||||
self.bootstrap_room = bootstrap_room
|
||||
self.bootstrap_host = "host"
|
||||
self.return_logprob = False
|
||||
self.latencies = []
|
||||
self.origin_input_ids = []
|
||||
self.output_ids = []
|
||||
self.is_retracted = True
|
||||
self.load_calls = []
|
||||
self.prealloc_done = False
|
||||
self.finished_reason = cast(Any, None)
|
||||
self.cached_tokens = 0
|
||||
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,
|
||||
)
|
||||
|
||||
def add_latency(self, stage):
|
||||
self.latencies.append(stage)
|
||||
|
||||
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))
|
||||
|
||||
|
||||
class FakeReceiver:
|
||||
def __init__(self, should_fail=False):
|
||||
self.should_fail = should_fail
|
||||
self.init_calls = []
|
||||
self.clear_calls = 0
|
||||
|
||||
def init(self, *args):
|
||||
self.init_calls.append(args)
|
||||
|
||||
def failure_exception(self):
|
||||
if self.should_fail:
|
||||
raise RuntimeError("boom")
|
||||
|
||||
def clear(self):
|
||||
self.clear_calls += 1
|
||||
return None
|
||||
|
||||
|
||||
class FakeItem:
|
||||
def __init__(self, value):
|
||||
self.value = value
|
||||
|
||||
def item(self):
|
||||
return self.value
|
||||
|
||||
|
||||
class FakeTensor:
|
||||
def __init__(self, values):
|
||||
self.values = values
|
||||
|
||||
def __getitem__(self, item):
|
||||
if isinstance(item, tuple):
|
||||
row, col = item
|
||||
return FakeTensor(self.values[row][col])
|
||||
return FakeTensor(self.values[item])
|
||||
|
||||
def cpu(self):
|
||||
return self
|
||||
|
||||
def numpy(self):
|
||||
return self.values
|
||||
|
||||
|
||||
class FakeAllocator(ReqToMetadataIdxAllocator):
|
||||
def __init__(self):
|
||||
super().__init__(size=0)
|
||||
self.freed = []
|
||||
|
||||
def free(self, free_index):
|
||||
self.freed.append(free_index)
|
||||
|
||||
|
||||
class TestDecodeQueueCompaction(CustomTestCase):
|
||||
def test_decode_transfer_queue_compacts_in_one_pass(self):
|
||||
streamed = []
|
||||
released = []
|
||||
committed = []
|
||||
allocator = FakeAllocator()
|
||||
|
||||
queue = DecodeTransferQueue.__new__(DecodeTransferQueue)
|
||||
queue.gloo_group = None
|
||||
queue.req_to_metadata_buffer_idx_allocator = allocator
|
||||
queue.tp_rank = 0
|
||||
queue.metadata_buffers = cast(Any, object())
|
||||
queue.tree_cache = cast(Any, object())
|
||||
queue.spec_algorithm = cast(Any, SimpleNamespace(is_none=lambda: True))
|
||||
queue.scheduler = cast(
|
||||
Any,
|
||||
SimpleNamespace(
|
||||
stream_output=lambda reqs, return_logprob: streamed.extend(
|
||||
req.rid for req in reqs
|
||||
),
|
||||
enable_metrics=False,
|
||||
token_to_kv_pool_allocator=SimpleNamespace(
|
||||
get_kvcache=lambda: SimpleNamespace()
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
keep = DecodeRequest(
|
||||
req=cast(Any, FakeReq("keep", 1)),
|
||||
kv_receiver=cast(Any, FakeReceiver()),
|
||||
metadata_buffer_index=10,
|
||||
)
|
||||
success = DecodeRequest(
|
||||
req=cast(Any, FakeReq("success", 2)),
|
||||
kv_receiver=cast(Any, FakeReceiver()),
|
||||
metadata_buffer_index=11,
|
||||
)
|
||||
failed = DecodeRequest(
|
||||
req=cast(Any, FakeReq("failed", 3)),
|
||||
kv_receiver=cast(Any, FakeReceiver(should_fail=True)),
|
||||
metadata_buffer_index=12,
|
||||
)
|
||||
skipped = DecodeRequest(
|
||||
req=cast(Any, FakeReq("skip", 4)),
|
||||
kv_receiver=cast(Any, FakeReceiver()),
|
||||
metadata_buffer_index=13,
|
||||
)
|
||||
queue.queue = [keep, success, failed, skipped]
|
||||
|
||||
with (
|
||||
patch(
|
||||
"sglang.srt.disaggregation.decode.poll_and_all_reduce",
|
||||
return_value=[
|
||||
KVPoll.Transferring,
|
||||
KVPoll.Success,
|
||||
KVPoll.Failed,
|
||||
KVPoll.Success,
|
||||
],
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.disaggregation.decode.release_kv_cache",
|
||||
lambda req, tree_cache, is_insert=False: released.append(
|
||||
(req.rid, is_insert)
|
||||
),
|
||||
),
|
||||
):
|
||||
queue._commit_transfer_to_req = lambda decode_req: (
|
||||
committed.append(decode_req.req.rid) or True
|
||||
)
|
||||
transferred = queue.pop_transferred(
|
||||
rids_to_check=["keep", "success", "failed"]
|
||||
)
|
||||
|
||||
self.assertEqual([req.rid for req in transferred], ["success"])
|
||||
self.assertEqual(committed, ["success"])
|
||||
self.assertEqual(streamed, ["failed"])
|
||||
self.assertEqual(released, [("failed", False)])
|
||||
self.assertEqual(allocator.freed, [11, 12])
|
||||
self.assertEqual(queue.queue, [keep, skipped])
|
||||
|
||||
def test_decode_transfer_queue_keeps_metadata_waiters(self):
|
||||
allocator = FakeAllocator()
|
||||
|
||||
queue = DecodeTransferQueue.__new__(DecodeTransferQueue)
|
||||
queue.gloo_group = None
|
||||
queue.req_to_metadata_buffer_idx_allocator = allocator
|
||||
queue.tp_rank = 0
|
||||
queue.metadata_buffers = cast(Any, object())
|
||||
queue.tree_cache = cast(Any, object())
|
||||
queue.spec_algorithm = cast(Any, SimpleNamespace(is_none=lambda: True))
|
||||
queue.scheduler = cast(
|
||||
Any,
|
||||
SimpleNamespace(
|
||||
stream_output=lambda reqs, return_logprob: None,
|
||||
enable_metrics=False,
|
||||
token_to_kv_pool_allocator=SimpleNamespace(
|
||||
get_kvcache=lambda: SimpleNamespace()
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
waiting = DecodeRequest(
|
||||
req=cast(Any, FakeReq("waiting", 1)),
|
||||
kv_receiver=cast(Any, FakeReceiver()),
|
||||
metadata_buffer_index=10,
|
||||
)
|
||||
queue.queue = [waiting]
|
||||
|
||||
with patch(
|
||||
"sglang.srt.disaggregation.decode.poll_and_all_reduce",
|
||||
return_value=[KVPoll.Success],
|
||||
):
|
||||
queue._commit_transfer_to_req = lambda decode_req: False
|
||||
transferred = queue.pop_transferred()
|
||||
|
||||
self.assertEqual(transferred, [])
|
||||
self.assertEqual(allocator.freed, [])
|
||||
self.assertEqual(queue.queue, [waiting])
|
||||
|
||||
def test_commit_transfer_to_req_waits_for_real_metadata(self):
|
||||
queue = DecodeTransferQueue.__new__(DecodeTransferQueue)
|
||||
queue.metadata_buffers = cast(
|
||||
Any,
|
||||
SimpleNamespace(
|
||||
get_buf=lambda idx: (
|
||||
[FakeItem(0)],
|
||||
[FakeItem(0)],
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
[FakeItem(0)],
|
||||
)
|
||||
),
|
||||
)
|
||||
queue.scheduler = cast(
|
||||
Any,
|
||||
SimpleNamespace(
|
||||
server_args=SimpleNamespace(disaggregation_transfer_backend="mooncake")
|
||||
),
|
||||
)
|
||||
queue.spec_algorithm = cast(Any, SimpleNamespace(is_none=lambda: True))
|
||||
|
||||
receiver = FakeReceiver()
|
||||
decode_req = DecodeRequest(
|
||||
req=cast(Any, FakeReq("waiting", 3)),
|
||||
kv_receiver=cast(Any, receiver),
|
||||
metadata_buffer_index=9,
|
||||
)
|
||||
|
||||
should_remove = queue._commit_transfer_to_req(decode_req)
|
||||
|
||||
self.assertIs(should_remove, False)
|
||||
self.assertIs(decode_req.kv_receiver, receiver)
|
||||
self.assertEqual(receiver.clear_calls, 0)
|
||||
self.assertEqual(decode_req.req.output_ids, [])
|
||||
|
||||
def test_commit_transfer_to_req_aborts_on_room_mismatch(self):
|
||||
queue = DecodeTransferQueue.__new__(DecodeTransferQueue)
|
||||
queue.metadata_buffers = cast(
|
||||
Any,
|
||||
SimpleNamespace(
|
||||
get_buf=lambda idx: (
|
||||
[FakeItem(0)],
|
||||
[FakeItem(0)],
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
[FakeItem(99)],
|
||||
)
|
||||
),
|
||||
)
|
||||
queue.scheduler = cast(
|
||||
Any,
|
||||
SimpleNamespace(
|
||||
server_args=SimpleNamespace(disaggregation_transfer_backend="mooncake")
|
||||
),
|
||||
)
|
||||
queue.spec_algorithm = cast(Any, SimpleNamespace(is_none=lambda: True))
|
||||
|
||||
receiver = FakeReceiver()
|
||||
decode_req = DecodeRequest(
|
||||
req=cast(Any, FakeReq("corrupt", 3)),
|
||||
kv_receiver=cast(Any, receiver),
|
||||
metadata_buffer_index=9,
|
||||
)
|
||||
|
||||
aborted = []
|
||||
with patch(
|
||||
"sglang.srt.disaggregation.decode.prepare_abort",
|
||||
lambda req, message, status_code: aborted.append(
|
||||
(req.rid, message, status_code)
|
||||
),
|
||||
):
|
||||
should_remove = queue._commit_transfer_to_req(decode_req)
|
||||
|
||||
self.assertIs(should_remove, True)
|
||||
self.assertEqual(receiver.clear_calls, 1)
|
||||
self.assertIsNone(decode_req.kv_receiver)
|
||||
self.assertEqual(len(aborted), 1)
|
||||
self.assertEqual(aborted[0][0], "corrupt")
|
||||
|
||||
def test_resume_retracted_reqs_compacts_queue_in_one_pass(self):
|
||||
prealloc_queue = DecodePreallocQueue.__new__(DecodePreallocQueue)
|
||||
prealloc_queue.req_to_token_pool = cast(
|
||||
Any, SimpleNamespace(available_size=lambda: 1)
|
||||
)
|
||||
prealloc_queue.token_to_kv_pool_allocator = cast(Any, object())
|
||||
prealloc_queue.num_reserved_decode_tokens = 1
|
||||
prealloc_queue._allocatable_tokens = (
|
||||
lambda retractable_tokens=None, count_retracted=False: 8
|
||||
)
|
||||
prealloc_queue._pre_alloc = lambda req: setattr(req, "prealloc_done", True)
|
||||
|
||||
first = FakeReq("resume", 1)
|
||||
first.origin_input_ids = [1, 2]
|
||||
first.output_ids = [3]
|
||||
skipped = FakeReq("skip", 2)
|
||||
skipped.origin_input_ids = [4]
|
||||
blocked = FakeReq("blocked", 3)
|
||||
blocked.origin_input_ids = [5, 6, 7, 8, 9, 10, 11, 12]
|
||||
|
||||
prealloc_queue.retracted_queue = cast(Any, [first, skipped, blocked])
|
||||
|
||||
resumed = prealloc_queue.resume_retracted_reqs(
|
||||
rids_to_check=["resume", "blocked"]
|
||||
)
|
||||
|
||||
self.assertEqual(resumed, [first])
|
||||
self.assertIs(first.is_retracted, False)
|
||||
self.assertIs(first.prealloc_done, True)
|
||||
self.assertEqual(len(first.load_calls), 1)
|
||||
self.assertEqual(prealloc_queue.retracted_queue, [skipped, blocked])
|
||||
|
||||
def test_pop_preallocated_still_removes_failed_reqs_after_block(self):
|
||||
streamed = []
|
||||
queue = DecodePreallocQueue.__new__(DecodePreallocQueue)
|
||||
queue._resolve_pending_reqs = lambda: None
|
||||
queue._update_handshake_waiters = lambda rids_to_check=None: None
|
||||
queue.req_to_token_pool = cast(
|
||||
Any,
|
||||
SimpleNamespace(
|
||||
available_size=lambda: 1,
|
||||
req_to_token=FakeTensor([[1, 2, 3, 4, 5, 6, 7, 8]]),
|
||||
),
|
||||
)
|
||||
queue.req_to_metadata_buffer_idx_allocator = cast(
|
||||
Any, SimpleNamespace(available_size=lambda: 1, alloc=lambda: 7)
|
||||
)
|
||||
queue.token_to_kv_pool_allocator = cast(Any, SimpleNamespace(page_size=1))
|
||||
queue.token_to_kv_pool = cast(Any, object())
|
||||
queue.num_reserved_decode_tokens = 1
|
||||
queue.scheduler = cast(
|
||||
Any,
|
||||
SimpleNamespace(
|
||||
stream_output=lambda reqs, return_logprob: streamed.extend(
|
||||
req.rid for req in reqs
|
||||
),
|
||||
enable_metrics=False,
|
||||
running_batch=SimpleNamespace(reqs=[]),
|
||||
sliding_window_size=4,
|
||||
),
|
||||
)
|
||||
queue._allocatable_tokens = (
|
||||
lambda retractable_tokens=None, count_retracted=True: 2
|
||||
)
|
||||
queue._pre_alloc = lambda req: setattr(req, "req_pool_idx", 0)
|
||||
|
||||
blocked_req = FakeReq("blocked", 1)
|
||||
blocked_req.origin_input_ids = [1, 2, 3]
|
||||
blocked_req.sampling_params = SimpleNamespace(max_new_tokens=8)
|
||||
blocked = DecodeRequest(
|
||||
req=cast(Any, blocked_req),
|
||||
kv_receiver=cast(Any, FakeReceiver()),
|
||||
waiting_for_input=True,
|
||||
)
|
||||
failed_req = FakeReq("failed", 2)
|
||||
failed_req.finished_reason = FINISH_ABORT("boom")
|
||||
failed = DecodeRequest(
|
||||
req=cast(Any, failed_req),
|
||||
kv_receiver=cast(Any, FakeReceiver()),
|
||||
waiting_for_input=True,
|
||||
)
|
||||
tail = DecodeRequest(
|
||||
req=cast(Any, FakeReq("tail", 3)),
|
||||
kv_receiver=cast(Any, FakeReceiver()),
|
||||
waiting_for_input=True,
|
||||
)
|
||||
queue.queue = cast(Any, [blocked, failed, tail])
|
||||
|
||||
preallocated, failed_reqs = queue.pop_preallocated(
|
||||
rids_to_check=["blocked", "failed", "tail"]
|
||||
)
|
||||
|
||||
self.assertEqual(preallocated, [])
|
||||
self.assertEqual([req.req.rid for req in failed_reqs], ["failed"])
|
||||
self.assertEqual(streamed, ["failed"])
|
||||
self.assertEqual(queue.queue, [blocked, tail])
|
||||
|
||||
def test_pop_preallocated_compacts_queue_in_one_pass(self):
|
||||
streamed = []
|
||||
queue = DecodePreallocQueue.__new__(DecodePreallocQueue)
|
||||
queue._resolve_pending_reqs = lambda: None
|
||||
queue._update_handshake_waiters = lambda rids_to_check=None: None
|
||||
queue.req_to_token_pool = cast(
|
||||
Any,
|
||||
SimpleNamespace(
|
||||
available_size=lambda: 1,
|
||||
req_to_token=FakeTensor([[1, 2, 3, 4, 5, 6, 7, 8]]),
|
||||
),
|
||||
)
|
||||
queue.req_to_metadata_buffer_idx_allocator = cast(
|
||||
Any, SimpleNamespace(available_size=lambda: 1, alloc=lambda: 7)
|
||||
)
|
||||
queue.token_to_kv_pool_allocator = cast(Any, SimpleNamespace(page_size=1))
|
||||
queue.token_to_kv_pool = cast(Any, object())
|
||||
queue.num_reserved_decode_tokens = 1
|
||||
queue.scheduler = cast(
|
||||
Any,
|
||||
SimpleNamespace(
|
||||
stream_output=lambda reqs, return_logprob: streamed.extend(
|
||||
req.rid for req in reqs
|
||||
),
|
||||
enable_metrics=False,
|
||||
running_batch=SimpleNamespace(reqs=[]),
|
||||
sliding_window_size=4,
|
||||
),
|
||||
)
|
||||
queue._allocatable_tokens = (
|
||||
lambda retractable_tokens=None, count_retracted=True: 6
|
||||
)
|
||||
queue._pre_alloc = lambda req: setattr(req, "req_pool_idx", 0)
|
||||
|
||||
skipped = DecodeRequest(
|
||||
req=cast(Any, FakeReq("skip", 1)),
|
||||
kv_receiver=cast(Any, FakeReceiver()),
|
||||
waiting_for_input=True,
|
||||
)
|
||||
failed_req = FakeReq("failed", 2)
|
||||
failed_req.finished_reason = FINISH_ABORT("boom")
|
||||
failed = DecodeRequest(
|
||||
req=cast(Any, failed_req),
|
||||
kv_receiver=cast(Any, FakeReceiver()),
|
||||
waiting_for_input=True,
|
||||
)
|
||||
waiting = DecodeRequest(
|
||||
req=cast(Any, FakeReq("wait", 3)),
|
||||
kv_receiver=cast(Any, FakeReceiver()),
|
||||
waiting_for_input=False,
|
||||
)
|
||||
success_req = FakeReq("success", 4)
|
||||
success_req.origin_input_ids = [1, 2]
|
||||
success_req.sampling_params = SimpleNamespace(max_new_tokens=2)
|
||||
success = DecodeRequest(
|
||||
req=cast(Any, success_req),
|
||||
kv_receiver=cast(Any, FakeReceiver()),
|
||||
waiting_for_input=True,
|
||||
)
|
||||
blocked_req = FakeReq("blocked", 5)
|
||||
blocked_req.origin_input_ids = [1, 2, 3, 4, 5, 6]
|
||||
blocked_req.sampling_params = SimpleNamespace(max_new_tokens=8)
|
||||
blocked = DecodeRequest(
|
||||
req=cast(Any, blocked_req),
|
||||
kv_receiver=cast(Any, FakeReceiver()),
|
||||
waiting_for_input=True,
|
||||
)
|
||||
tail = DecodeRequest(
|
||||
req=cast(Any, FakeReq("tail", 6)),
|
||||
kv_receiver=cast(Any, FakeReceiver()),
|
||||
waiting_for_input=True,
|
||||
)
|
||||
|
||||
queue.queue = cast(Any, [skipped, failed, waiting, success, blocked, tail])
|
||||
|
||||
preallocated, failed_reqs = queue.pop_preallocated(
|
||||
rids_to_check=["failed", "wait", "success", "blocked"]
|
||||
)
|
||||
|
||||
self.assertEqual([req.req.rid for req in preallocated], ["success"])
|
||||
self.assertEqual([req.req.rid for req in failed_reqs], ["failed"])
|
||||
self.assertEqual(streamed, ["failed"])
|
||||
self.assertEqual(success.metadata_buffer_index, 7)
|
||||
self.assertEqual(queue.queue, [skipped, waiting, blocked, tail])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user