Hold decode EAGLE metadata until prebuilt consumption
Transferred EAGLE metadata buffers are reusable slot views. The previous clone-based mitigation protected correctness but added copies on the transfer hot path and hid the actual lifetime contract. This change makes the transferred request own the metadata slot while it sits in the decode waiting queue, then releases it immediately after process_prebuilt has consumed top-k and hidden state into the prebuilt batch. Abort paths also release any held decode metadata slot. Constraint: Decode disaggregation metadata buffers are reusable slot views consumed later by process_prebuilt. Rejected: Clone transferred EAGLE tensors at commit time | correct but less efficient and masks the ownership contract. Rejected: Release in process_batch_result_prebuilt | holds slots across forward longer than needed. Confidence: medium Scope-risk: moderate Directive: Do not free successful EAGLE transfer metadata in pop_transferred unless process_prebuilt consumption is also moved earlier. Tested: Remote py_compile for decode.py, scheduler.py, scheduler_output_processor_mixin.py, and test_decode_queue_compaction.py. Tested: Remote focused lifecycle tests passed: 3 passed. Tested: Remote full test_decode_queue_compaction.py passed: 11 passed, 5 warnings. Not-tested: Fresh ETE runtime validation of EAGLE accept-length recovery after the C48 sync.
This commit is contained in:
@@ -15,6 +15,9 @@ from sglang.srt.disaggregation.decode import (
|
||||
SchedulerDisaggregationDecodeMixin,
|
||||
)
|
||||
from sglang.srt.disaggregation.utils import ReqToMetadataIdxAllocator
|
||||
from sglang.srt.managers.scheduler_output_processor_mixin import (
|
||||
SchedulerOutputProcessorMixin,
|
||||
)
|
||||
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
|
||||
@@ -55,6 +58,12 @@ class FakeReq:
|
||||
def set_forward_entry_time(self, ts):
|
||||
self.forward_entry_time = ts
|
||||
|
||||
def set_decode_prebuilt_finish_time(self):
|
||||
return None
|
||||
|
||||
def set_quick_finish_time(self):
|
||||
return None
|
||||
|
||||
self.time_stats = FakeTimeStats()
|
||||
|
||||
def add_latency(self, stage):
|
||||
@@ -66,6 +75,12 @@ class FakeReq:
|
||||
def init_next_round_input(self, tree_cache):
|
||||
self.init_next_round_calls.append(tree_cache)
|
||||
|
||||
def check_finished(self):
|
||||
return None
|
||||
|
||||
def finished(self):
|
||||
return self.finished_reason is not None
|
||||
|
||||
|
||||
class FakeReceiver:
|
||||
def __init__(self, should_fail=False):
|
||||
@@ -340,11 +355,15 @@ class TestDecodeQueueCompaction(CustomTestCase):
|
||||
self.assertEqual(len(aborted), 1)
|
||||
self.assertEqual(aborted[0][0], "corrupt")
|
||||
|
||||
def test_commit_transfer_clones_eagle_metadata_before_reusing_slot(self):
|
||||
def test_pop_transferred_holds_eagle_metadata_slot_until_prebuilt_consumes(self):
|
||||
queue = DecodeTransferQueue.__new__(DecodeTransferQueue)
|
||||
allocator = FakeAllocator()
|
||||
output_topk_p = torch.arange(16, dtype=torch.float32)
|
||||
output_topk_index = torch.arange(16, dtype=torch.int64)
|
||||
output_hidden_states = torch.arange(8, dtype=torch.float32)
|
||||
queue.gloo_group = None
|
||||
queue.req_to_metadata_buffer_idx_allocator = allocator
|
||||
queue.tp_rank = 0
|
||||
queue.metadata_buffers = cast(
|
||||
Any,
|
||||
SimpleNamespace(
|
||||
@@ -362,10 +381,13 @@ class TestDecodeQueueCompaction(CustomTestCase):
|
||||
)
|
||||
),
|
||||
)
|
||||
queue.tree_cache = cast(Any, object())
|
||||
queue.scheduler = cast(
|
||||
Any,
|
||||
SimpleNamespace(
|
||||
server_args=SimpleNamespace(disaggregation_transfer_backend="mooncake")
|
||||
server_args=SimpleNamespace(disaggregation_transfer_backend="mooncake"),
|
||||
stream_output=lambda reqs, return_logprob: None,
|
||||
enable_metrics=False,
|
||||
),
|
||||
)
|
||||
queue.spec_algorithm = cast(Any, SimpleNamespace(is_none=lambda: False))
|
||||
@@ -376,33 +398,46 @@ class TestDecodeQueueCompaction(CustomTestCase):
|
||||
kv_receiver=cast(Any, receiver),
|
||||
metadata_buffer_index=9,
|
||||
)
|
||||
queue.queue = [decode_req]
|
||||
|
||||
should_remove = queue._commit_transfer_to_req(decode_req)
|
||||
self.assertIs(should_remove, True)
|
||||
with patch(
|
||||
"sglang.srt.disaggregation.decode.poll_and_all_reduce",
|
||||
return_value=[KVPoll.Success],
|
||||
):
|
||||
transferred = queue.pop_transferred()
|
||||
|
||||
output_topk_p.fill_(-1)
|
||||
output_topk_index.fill_(-2)
|
||||
output_hidden_states.fill_(-3)
|
||||
self.assertEqual(transferred, [decode_req.req])
|
||||
self.assertEqual(queue.queue, [])
|
||||
self.assertEqual(allocator.freed, [])
|
||||
|
||||
self.assertEqual(decode_req.req.output_ids, [7])
|
||||
self.assertEqual(decode_req.req.cached_tokens, 320)
|
||||
self.assertEqual(decode_req.req.output_topk_p.tolist(), list(range(16)))
|
||||
self.assertEqual(decode_req.req.output_topk_index.tolist(), list(range(16)))
|
||||
self.assertEqual(decode_req.req.metadata_buffer_index, 9)
|
||||
self.assertEqual(
|
||||
decode_req.req.hidden_states_tensor.tolist(),
|
||||
[float(i) for i in range(8)],
|
||||
)
|
||||
self.assertNotEqual(
|
||||
decode_req.req.output_topk_p.data_ptr(), output_topk_p.data_ptr()
|
||||
)
|
||||
self.assertNotEqual(
|
||||
self.assertEqual(
|
||||
decode_req.req.output_topk_index.data_ptr(), output_topk_index.data_ptr()
|
||||
)
|
||||
self.assertNotEqual(
|
||||
self.assertEqual(
|
||||
decode_req.req.hidden_states_tensor.data_ptr(),
|
||||
output_hidden_states.data_ptr(),
|
||||
)
|
||||
|
||||
def test_free_decode_metadata_index_if_held_releases_once(self):
|
||||
allocator = FakeAllocator()
|
||||
req = FakeReq("eagle", 3)
|
||||
req.metadata_buffer_index = 9
|
||||
|
||||
scheduler = SchedulerOutputProcessorMixin.__new__(SchedulerOutputProcessorMixin)
|
||||
scheduler.req_to_metadata_buffer_idx_allocator = allocator
|
||||
|
||||
scheduler._free_decode_metadata_index_if_held(req)
|
||||
scheduler._free_decode_metadata_index_if_held(req)
|
||||
|
||||
self.assertEqual(allocator.freed, [9])
|
||||
self.assertEqual(req.metadata_buffer_index, -1)
|
||||
|
||||
def test_resume_retracted_reqs_compacts_queue_in_one_pass(self):
|
||||
prealloc_queue = DecodePreallocQueue.__new__(DecodePreallocQueue)
|
||||
prealloc_queue.req_to_token_pool = cast(
|
||||
@@ -596,6 +631,7 @@ class TestDecodeQueueCompaction(CustomTestCase):
|
||||
self.assertEqual(queue.queue, [skipped, waiting, blocked, tail])
|
||||
|
||||
def test_get_new_prebuilt_batch_slices_waiting_queue_prefix(self):
|
||||
allocator = FakeAllocator()
|
||||
scheduler = cast(Any, SimpleNamespace())
|
||||
scheduler.grammar_manager = SimpleNamespace(
|
||||
has_waiting_grammars=lambda: False,
|
||||
@@ -613,6 +649,14 @@ class TestDecodeQueueCompaction(CustomTestCase):
|
||||
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
|
||||
)
|
||||
)
|
||||
for i, req in enumerate(scheduler.waiting_queue[:3]):
|
||||
req.metadata_buffer_index = 20 + i
|
||||
|
||||
captured = {}
|
||||
|
||||
@@ -642,6 +686,9 @@ class TestDecodeQueueCompaction(CustomTestCase):
|
||||
captured["batch"].processed,
|
||||
[(scheduler.server_args, scheduler.future_map)],
|
||||
)
|
||||
self.assertEqual(allocator.freed, [20, 21, 22])
|
||||
for req in captured["reqs"]:
|
||||
self.assertEqual(req.metadata_buffer_index, -1)
|
||||
|
||||
def test_get_new_prebuilt_batch_keeps_waiting_queue_when_no_capacity(self):
|
||||
scheduler = cast(Any, SimpleNamespace())
|
||||
|
||||
Reference in New Issue
Block a user