Stabilize disaggregated decode burst handoff

Burst decode can sit in disaggregation prealloc/transfer queues while health probes still need an immediate scheduler-alive response, and EAGLE prebuilt metadata must outlive the synthetic prebuilt step until the first real decode consume. The transfer path also needs to send final aux metadata even when there are no new KV pages in the final chunk.

This keeps decode metadata slot ownership narrow instead of cloning EAGLE tensors, adds fail-safe cleanup for prebuilt exceptions/finished prebuilt requests, fixes final empty chunk metadata transfer, and records the investigation ledger for future debugging.

Constraint: bs=1 historically worked, so decode compute/sampling behavior must not be changed without direct evidence.

Rejected: Clone EAGLE metadata tensors on transfer commit | avoids lifetime issues but adds hot-path CPU/GPU memory traffic.

Rejected: Treat prefill AbortReq logs as root cause | current evidence shows they can be downstream of decode/router aborts.

Confidence: medium

Scope-risk: moderate

Directive: Do not move EAGLE metadata slot release earlier than first real decode result processing without proving the H2D/spec_info consume point changed.

Tested: g0034 docker py_compile for touched scheduler/disagg/test files; PYTHONPATH=python python -m pytest -q test/registered/unit/disaggregation/test_decode_queue_compaction.py test/registered/unit/mem_cache/test_req_to_token_pool.py test/registered/unit/managers/test_scheduler_health_check.py -> 24 passed

Not-tested: Fresh end-to-end burst traffic after restart; decode node_rank=1 persistent log capture.
This commit is contained in:
laoyao0822
2026-06-05 19:59:04 +08:00
parent 5eecc5f9c5
commit 6eea77e5e9
9 changed files with 687 additions and 28 deletions
@@ -14,7 +14,10 @@ from sglang.srt.disaggregation.decode import (
DecodeTransferQueue,
SchedulerDisaggregationDecodeMixin,
)
from sglang.srt.disaggregation.utils import ReqToMetadataIdxAllocator
from sglang.srt.disaggregation.utils import (
DisaggregationMode,
ReqToMetadataIdxAllocator,
)
from sglang.srt.managers.scheduler_output_processor_mixin import (
SchedulerOutputProcessorMixin,
)
@@ -41,6 +44,13 @@ class FakeReq:
self.cached_tokens = 0
self.init_next_round_calls = []
self.sampling_params = SimpleNamespace(max_new_tokens=0)
self.return_hidden_states = False
self.grammar = None
self.token_ids_logprob = None
self.top_logprobs_num = 0
self.multimodal_inputs = None
self.mamba_ping_pong_track_buffer = None
self.to_finish = None
class FakeTimeStats:
def __init__(self):
@@ -64,6 +74,12 @@ class FakeReq:
def set_quick_finish_time(self):
return None
def set_last_decode_finish_time(self):
return None
def set_completion_time(self):
return None
self.time_stats = FakeTimeStats()
def add_latency(self, stage):
@@ -75,7 +91,7 @@ class FakeReq:
def init_next_round_input(self, tree_cache):
self.init_next_round_calls.append(tree_cache)
def check_finished(self):
def check_finished(self, *args):
return None
def finished(self):
@@ -146,6 +162,18 @@ class FakeAllocator(ReqToMetadataIdxAllocator):
self.freed.append(free_index)
class FakeTokenToKVAllocator:
def __init__(self):
self.begin_calls = 0
self.end_calls = 0
def free_group_begin(self):
self.begin_calls += 1
def free_group_end(self):
self.end_calls += 1
class TestDecodeQueueCompaction(CustomTestCase):
def test_decode_transfer_queue_compacts_in_one_pass(self):
streamed = []
@@ -692,12 +720,127 @@ class TestDecodeQueueCompaction(CustomTestCase):
captured["batch"].processed,
[(scheduler.server_args, scheduler.future_map)],
)
self.assertEqual(allocator.freed, [20, 21, 22])
# EAGLE metadata slots are intentionally held past process_prebuilt().
# The initial draft state is consumed by the first real decode forward,
# so releasing here would let burst transfers overwrite the pinned CPU
# source views before the GPU copy/consume is complete.
self.assertEqual(allocator.freed, [])
for req in captured["reqs"]:
self.assertEqual(req.metadata_buffer_index, -1)
self.assertIsNone(getattr(req, "output_topk_p", None))
self.assertIsNone(getattr(req, "output_topk_index", None))
self.assertIsNone(getattr(req, "hidden_states_tensor", None))
self.assertGreaterEqual(req.metadata_buffer_index, 20)
def test_get_new_prebuilt_batch_frees_metadata_on_prebuilt_error(self):
allocator = FakeAllocator()
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
req0 = FakeReq("req-0", 0)
scheduler.waiting_queue = [req0]
scheduler.running_batch = SimpleNamespace(batch_size=lambda: 0)
scheduler.req_to_token_pool = SimpleNamespace(size=8)
scheduler.max_running_requests = 8
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()
scheduler.req_to_metadata_buffer_idx_allocator = allocator
scheduler._free_decode_metadata_index_if_held = (
SchedulerOutputProcessorMixin._free_decode_metadata_index_if_held.__get__(
scheduler
)
)
req0.metadata_buffer_index = 42
class FailingBatch(FakeBatch):
def process_prebuilt(self, server_args, future_map):
raise RuntimeError("prebuilt failed")
with patch(
"sglang.srt.disaggregation.decode.ScheduleBatch.init_new",
lambda reqs, *args, **kwargs: FailingBatch(),
):
with self.assertRaisesRegex(RuntimeError, "prebuilt failed"):
SchedulerDisaggregationDecodeMixin.get_new_prebuilt_batch(scheduler)
self.assertEqual(allocator.freed, [42])
self.assertEqual(req0.metadata_buffer_index, -1)
def test_process_batch_result_prebuilt_frees_finished_metadata(self):
allocator = FakeAllocator()
scheduler = SchedulerOutputProcessorMixin.__new__(SchedulerOutputProcessorMixin)
scheduler.disaggregation_mode = DisaggregationMode.DECODE
scheduler.req_to_metadata_buffer_idx_allocator = allocator
scheduler.tree_cache = object()
scheduler.stream_output = lambda reqs, return_logprob: None
req = FakeReq("finished", 0)
req.metadata_buffer_index = 43
req.output_topk_p = torch.ones((1,), dtype=torch.float32)
req.output_topk_index = torch.ones((1,), dtype=torch.int64)
req.hidden_states_tensor = torch.ones((4,), dtype=torch.float32)
req.finished_reason = FINISH_ABORT("done")
batch = SimpleNamespace(reqs=[req], return_logprob=False)
with patch(
"sglang.srt.managers.scheduler_output_processor_mixin.release_kv_cache",
lambda *args, **kwargs: None,
):
scheduler.process_batch_result_prebuilt(batch)
self.assertEqual(allocator.freed, [43])
self.assertEqual(req.metadata_buffer_index, -1)
def test_process_batch_result_decode_releases_prebuilt_metadata_after_consume(self):
allocator = FakeAllocator()
token_allocator = FakeTokenToKVAllocator()
scheduler = SchedulerOutputProcessorMixin.__new__(SchedulerOutputProcessorMixin)
scheduler.req_to_metadata_buffer_idx_allocator = allocator
scheduler.server_args = SimpleNamespace(
disaggregation_decode_enable_offload_kvcache=False
)
scheduler.enable_hisparse = False
scheduler.enable_overlap = False
scheduler.enable_metrics = False
scheduler.token_to_kv_pool_allocator = token_allocator
scheduler.tree_cache = object()
scheduler.forward_ct_decode = 0
scheduler.num_generated_tokens = 0
scheduler.stream_output = lambda reqs, return_logprob: None
scheduler.report_decode_stats = lambda *args, **kwargs: None
scheduler.update_spec_metrics = lambda *args, **kwargs: None
scheduler._maybe_log_eagle_accept_debug = lambda *args, **kwargs: None
req = FakeReq("decode", 0)
req.metadata_buffer_index = 44
req.output_topk_p = torch.ones((1,), dtype=torch.float32)
req.output_topk_index = torch.ones((1,), dtype=torch.int64)
req.hidden_states_tensor = torch.ones((4,), dtype=torch.float32)
batch = SimpleNamespace(
reqs=[req],
return_logprob=False,
spec_algorithm=SimpleNamespace(is_none=lambda: True),
is_spec_v2=False,
)
result = SimpleNamespace(
copy_done=None,
logits_output=SimpleNamespace(hidden_states=None, customized_info=None),
next_token_ids=torch.tensor([5], dtype=torch.int64),
can_run_cuda_graph=True,
num_accepted_tokens=1,
)
scheduler.process_batch_result_decode(batch, result)
self.assertEqual(allocator.freed, [44])
self.assertEqual(req.metadata_buffer_index, -1)
self.assertEqual(req.output_ids, [5])
self.assertEqual(token_allocator.begin_calls, 1)
self.assertEqual(token_allocator.end_calls, 1)
def test_get_new_prebuilt_batch_keeps_waiting_queue_when_no_capacity(self):
scheduler = cast(Any, SimpleNamespace())
@@ -0,0 +1,88 @@
import unittest
from collections import deque
from types import SimpleNamespace
from unittest.mock import MagicMock
from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.managers.io_struct import HealthCheckOutput
from sglang.srt.managers.scheduler import Scheduler
from sglang.srt.managers.tokenizer_manager import TokenizerManager
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=2, suite="stage-a-test-cpu")
class TestSchedulerHealthCheck(CustomTestCase):
def _idle_scheduler(self, *, disaggregation_mode=DisaggregationMode.NULL):
return SimpleNamespace(
running_batch=SimpleNamespace(is_empty=lambda: True, reqs=[]),
chunked_req=None,
dllm_manager=SimpleNamespace(any_staging_reqs=lambda: False),
last_batch=None,
cur_batch=None,
enable_overlap=False,
result_queue=deque(),
pp_size=1,
waiting_queue=[],
grammar_manager=SimpleNamespace(grammar_queue=[]),
disaggregation_mode=disaggregation_mode,
enable_hierarchical_cache=False,
)
def test_decode_prealloc_queue_is_busy_for_health_check(self):
scheduler = self._idle_scheduler(disaggregation_mode=DisaggregationMode.DECODE)
scheduler.disagg_decode_prealloc_queue = SimpleNamespace(queue=[object()])
scheduler.disagg_decode_transfer_queue = SimpleNamespace(queue=[])
scheduler.server_args = SimpleNamespace(
disaggregation_decode_enable_offload_kvcache=False
)
self.assertFalse(Scheduler.is_fully_idle(scheduler, for_health_check=True))
def test_busy_health_check_sends_immediate_signal(self):
sent = []
scheduler = SimpleNamespace(
session_controller=SimpleNamespace(maybe_reap=MagicMock()),
is_fully_idle=lambda for_health_check=False: False,
send_to_tokenizer=SimpleNamespace(
send_output=lambda output, recv_obj=None: sent.append((output, recv_obj))
),
return_health_check_ipcs=deque(),
_request_dispatcher=MagicMock(),
)
req = SimpleNamespace(rid="HEALTH_CHECK_unit", http_worker_ipc="ipc://unit")
Scheduler.process_input_requests(scheduler, [req])
scheduler._request_dispatcher.assert_not_called()
self.assertEqual(len(sent), 1)
self.assertIsInstance(sent[0][0], HealthCheckOutput)
self.assertEqual(sent[0][0].http_worker_ipc, "ipc://unit")
self.assertIs(sent[0][1], req)
def test_shutdown_pending_request_without_response_ts_aborts_cleanly(self):
event = SimpleNamespace(set=MagicMock())
state = SimpleNamespace(
finished=False,
obj=SimpleNamespace(stream=True),
text="",
output_ids=[],
out_list=[],
event=event,
)
tokenizer_manager = SimpleNamespace(rid_to_state={"rid": state})
TokenizerManager._finish_all_pending_requests_on_shutdown(tokenizer_manager)
self.assertTrue(state.finished)
self.assertEqual(len(state.out_list), 1)
self.assertEqual(
state.out_list[0]["meta_info"]["finish_reason"]["type"], "abort"
)
event.set.assert_called_once()
if __name__ == "__main__":
unittest.main()
@@ -13,6 +13,7 @@ from sglang.srt.disaggregation.decode import (
_kv_locs_to_page_indices_cpu,
)
from sglang.srt.disaggregation.prefill import (
SchedulerDisaggregationPrefillMixin,
_kv_locs_to_page_indices_cpu as _prefill_kv_locs_to_page_indices_cpu,
)
from sglang.srt.managers.schedule_batch import Req
@@ -200,6 +201,61 @@ class TestDecodePreallocQueue(CustomTestCase):
self.assertEqual(page_indices.dtype.name, "int32")
self.assertEqual(page_indices.tolist(), [32, 33, 34])
def test_prefill_final_empty_kv_chunk_still_sends_aux_metadata(self):
class FakeAllocator:
page_size = 4
def get_kvcache(self):
return object()
class FakeMetadataBuffers:
def __init__(self):
self.requests = []
def set_buf(self, req):
self.requests.append(req)
class FakeSender:
def __init__(self):
self.calls = []
def send(self, page_indices, state_indices):
self.calls.append((page_indices, state_indices))
fake_scheduler = SimpleNamespace(
token_to_kv_pool_allocator=FakeAllocator(),
req_to_token_pool=SimpleNamespace(
req_to_token=torch.arange(100, 112, dtype=torch.int64).reshape(1, -1)
),
disagg_metadata_buffers=FakeMetadataBuffers(),
disagg_prefill_bootstrap_queue=SimpleNamespace(draft_token_to_kv_pool=None),
)
sender = FakeSender()
req = SimpleNamespace(
rid="r-final-empty",
bootstrap_room=123,
start_send_idx=8,
fill_ids=list(range(8)),
origin_input_ids=list(range(8)),
req_pool_idx=0,
disagg_kv_sender=sender,
prefix_indices=[],
host_hit_length=0,
)
SchedulerDisaggregationPrefillMixin.send_kv_chunk(
fake_scheduler, req, last_chunk=True
)
self.assertEqual(req.start_send_idx, 8)
self.assertEqual(fake_scheduler.disagg_metadata_buffers.requests, [req])
self.assertEqual(len(sender.calls), 1)
page_indices, state_indices = sender.calls[0]
self.assertEqual(page_indices.dtype.name, "int32")
self.assertEqual(page_indices.tolist(), [])
self.assertIsNone(state_indices)
def test_prefill_kv_locs_to_page_indices_cpu_copies_only_page_starts(self):
kv_locs = torch.tensor(
[