Files
sglang/test/registered/unit/disaggregation/test_decode_queue_compaction.py

496 lines
17 KiB
Python

"""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()