Preserve transferred EAGLE state past metadata-slot reuse

Decode committed EAGLE top-k and hidden-state tensors as views into reusable metadata-buffer rows. The metadata index is freed immediately after transfer commit, while the request may wait before process_prebuilt consumes the draft state. Under concurrent cache-hit traffic a later transfer can overwrite the same row, leaving output_id copied correctly but EAGLE draft state corrupted, which matches low accept length despite successful KV/state registration.

Constraint: Metadata slots are intentionally recycled right after transfer commit for throughput.

Rejected: Hold metadata slots until process_prebuilt | larger lifetime change and reduces transfer capacity; cloning the small prebuilt EAGLE state is narrower.

Confidence: high

Scope-risk: narrow

Directive: Do not store reusable metadata-buffer views on Req unless the slot lifetime is extended through all consumers.

Tested: Local py_compile for decode.py and test_decode_queue_compaction.py.

Tested: Remote g0034 container py_compile for decode.py and test_decode_queue_compaction.py.

Tested: Remote g0034 focused clone-lifetime test: 1 passed.

Tested: Remote g0034 test_decode_queue_compaction.py: 10 passed, 5 warnings.

Not-tested: ETE cache-hit accept-length validation after restarting prefill/decode with this synced code.
This commit is contained in:
laoyao0822
2026-05-30 03:36:42 +08:00
parent b7364d23f9
commit 07c9544737
3 changed files with 121 additions and 5 deletions
@@ -5,6 +5,8 @@ from types import SimpleNamespace
from typing import Any, cast
from unittest.mock import patch
import torch
from sglang.srt.disaggregation.base import KVPoll
from sglang.srt.disaggregation.decode import (
DecodePreallocQueue,
@@ -338,6 +340,69 @@ 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):
queue = DecodeTransferQueue.__new__(DecodeTransferQueue)
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.metadata_buffers = cast(
Any,
SimpleNamespace(
get_buf=lambda idx: (
torch.tensor([7], dtype=torch.int32),
torch.tensor([320], dtype=torch.int32),
None,
None,
None,
None,
output_topk_p,
output_topk_index,
output_hidden_states,
torch.tensor([3], dtype=torch.int64),
)
),
)
queue.scheduler = cast(
Any,
SimpleNamespace(
server_args=SimpleNamespace(disaggregation_transfer_backend="mooncake")
),
)
queue.spec_algorithm = cast(Any, SimpleNamespace(is_none=lambda: False))
receiver = FakeReceiver()
decode_req = DecodeRequest(
req=cast(Any, FakeReq("eagle", 3)),
kv_receiver=cast(Any, receiver),
metadata_buffer_index=9,
)
should_remove = queue._commit_transfer_to_req(decode_req)
self.assertIs(should_remove, True)
output_topk_p.fill_(-1)
output_topk_index.fill_(-2)
output_hidden_states.fill_(-3)
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.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(
decode_req.req.output_topk_index.data_ptr(), output_topk_index.data_ptr()
)
self.assertNotEqual(
decode_req.req.hidden_states_tensor.data_ptr(),
output_hidden_states.data_ptr(),
)
def test_resume_retracted_reqs_compacts_queue_in_one_pass(self):
prealloc_queue = DecodePreallocQueue.__new__(DecodePreallocQueue)
prealloc_queue.req_to_token_pool = cast(
@@ -380,6 +445,7 @@ class TestDecodeQueueCompaction(CustomTestCase):
SimpleNamespace(
available_size=lambda: 1,
req_to_token=FakeTensor([[1, 2, 3, 4, 5, 6, 7, 8]]),
write=lambda *args, **kwargs: None,
),
)
queue.req_to_metadata_buffer_idx_allocator = cast(
@@ -387,6 +453,7 @@ class TestDecodeQueueCompaction(CustomTestCase):
)
queue.token_to_kv_pool_allocator = cast(Any, SimpleNamespace(page_size=1))
queue.token_to_kv_pool = cast(Any, object())
queue.draft_token_to_kv_pool = None
queue.num_reserved_decode_tokens = 1
queue.scheduler = cast(
Any,
@@ -402,7 +469,10 @@ class TestDecodeQueueCompaction(CustomTestCase):
queue._allocatable_tokens = (
lambda retractable_tokens=None, count_retracted=True: 2
)
queue._pre_alloc = lambda req: setattr(req, "req_pool_idx", 0)
queue._pre_alloc = lambda req, **kwargs: (
setattr(req, "req_pool_idx", 0),
torch.arange(len(req.origin_input_ids), dtype=torch.int64),
)[1]
blocked_req = FakeReq("blocked", 1)
blocked_req.origin_input_ids = [1, 2, 3]
@@ -445,6 +515,7 @@ class TestDecodeQueueCompaction(CustomTestCase):
SimpleNamespace(
available_size=lambda: 1,
req_to_token=FakeTensor([[1, 2, 3, 4, 5, 6, 7, 8]]),
write=lambda *args, **kwargs: None,
),
)
queue.req_to_metadata_buffer_idx_allocator = cast(
@@ -452,6 +523,7 @@ class TestDecodeQueueCompaction(CustomTestCase):
)
queue.token_to_kv_pool_allocator = cast(Any, SimpleNamespace(page_size=1))
queue.token_to_kv_pool = cast(Any, object())
queue.draft_token_to_kv_pool = None
queue.num_reserved_decode_tokens = 1
queue.scheduler = cast(
Any,
@@ -467,7 +539,10 @@ class TestDecodeQueueCompaction(CustomTestCase):
queue._allocatable_tokens = (
lambda retractable_tokens=None, count_retracted=True: 6
)
queue._pre_alloc = lambda req: setattr(req, "req_pool_idx", 0)
queue._pre_alloc = lambda req, **kwargs: (
setattr(req, "req_pool_idx", 0),
torch.arange(len(req.origin_input_ids), dtype=torch.int64),
)[1]
skipped = DecodeRequest(
req=cast(Any, FakeReq("skip", 1)),