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:
@@ -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)),
|
||||
|
||||
Reference in New Issue
Block a user