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:
laoyao0822
2026-05-30 04:51:50 +08:00
parent 07c9544737
commit e9c341afe8
5 changed files with 196 additions and 29 deletions

View File

@@ -1175,15 +1175,13 @@ class DecodeTransferQueue:
decode_req.req.output_ids.append(output_id[0].item())
decode_req.req.cached_tokens = cached_tokens[0].item()
if not self.spec_algorithm.is_none():
# ``metadata_buffers.get_buf`` returns views into the reusable
# metadata slots. ``pop_transferred`` frees the slot immediately
# after this commit, while the request can stay in the waiting queue
# before ``process_prebuilt`` consumes the EAGLE state. Keep an
# owned copy so subsequent transfers cannot overwrite the prebuilt
# draft top-k/hidden state.
decode_req.req.output_topk_p = output_topk_p.clone()
decode_req.req.output_topk_index = output_topk_index.clone()
decode_req.req.hidden_states_tensor = output_hidden_states.clone()
# ``metadata_buffers.get_buf`` returns views into reusable metadata
# slots. Keep the slot owned by the request until prebuilt consumes
# the EAGLE state instead of cloning these tensors on the hot path.
decode_req.req.metadata_buffer_index = idx
decode_req.req.output_topk_p = output_topk_p
decode_req.req.output_topk_index = output_topk_index
decode_req.req.hidden_states_tensor = output_hidden_states
_cp_draft_shared_kv_debug(
"decode_transfer_commit rid=%s room=%s metadata_idx=%s cached_tokens=%s "
@@ -1224,7 +1222,7 @@ class DecodeTransferQueue:
)
transferred_reqs = []
completed_decode_reqs = []
decode_reqs_to_free_metadata = []
remaining_queue = []
for decode_req, poll in zip(self.queue, polls):
if rids_to_check is not None and decode_req.req.rid not in rids_to_check:
@@ -1247,14 +1245,13 @@ class DecodeTransferQueue:
)
# release pre-allocated kv cache, but don't insert into the tree since it's failed
release_kv_cache(decode_req.req, self.tree_cache, is_insert=False)
completed_decode_reqs.append(decode_req)
decode_reqs_to_free_metadata.append(decode_req)
if self.scheduler.enable_metrics:
self.scheduler.metrics_collector.increment_transfer_failed_reqs()
continue
elif poll == KVPoll.Success:
should_remove = self._commit_transfer_to_req(decode_req)
if should_remove:
completed_decode_reqs.append(decode_req)
# Check if request was aborted due to corruption
if isinstance(decode_req.req.finished_reason, FINISH_ABORT):
self.scheduler.stream_output(
@@ -1263,10 +1260,13 @@ class DecodeTransferQueue:
release_kv_cache(
decode_req.req, self.tree_cache, is_insert=False
)
decode_reqs_to_free_metadata.append(decode_req)
if self.scheduler.enable_metrics:
self.scheduler.metrics_collector.increment_transfer_failed_reqs()
else:
transferred_reqs.append(decode_req.req)
if self.spec_algorithm.is_none():
decode_reqs_to_free_metadata.append(decode_req)
else:
remaining_queue.append(decode_req)
elif poll in [
@@ -1281,7 +1281,7 @@ class DecodeTransferQueue:
if len(polls) < len(self.queue):
remaining_queue.extend(self.queue[len(polls) :])
for decode_req in completed_decode_reqs:
for decode_req in decode_reqs_to_free_metadata:
idx = decode_req.metadata_buffer_index
assert idx != -1
self.req_to_metadata_buffer_idx_allocator.free(idx)
@@ -1459,7 +1459,11 @@ class SchedulerDisaggregationDecodeMixin:
# construct fake completed prefill
new_batch.prepare_for_prebuilt()
new_batch.process_prebuilt(self.server_args, self.future_map)
try:
new_batch.process_prebuilt(self.server_args, self.future_map)
finally:
for req in can_run_list:
self._free_decode_metadata_index_if_held(req)
return new_batch

View File

@@ -3189,6 +3189,9 @@ class Scheduler(
if self.enable_hisparse:
self.hisparse_coordinator.request_finished(req)
release_kv_cache(req, self.tree_cache)
release_req_to_metadata_buffer(
req, self.req_to_metadata_buffer_idx_allocator
)
# For disaggregation prefill mode, free the metadata buffer index
if self.disaggregation_mode == DisaggregationMode.PREFILL:
release_req_to_metadata_buffer(

View File

@@ -51,6 +51,19 @@ class SchedulerOutputProcessorMixin:
storage_backend_type = type(storage_backend).__name__
return storage_backend_type
def _free_decode_metadata_index_if_held(self: Scheduler, req: Req) -> None:
idx = getattr(req, "metadata_buffer_index", -1)
if idx is None or idx < 0:
return
allocator = getattr(self, "req_to_metadata_buffer_idx_allocator", None)
assert allocator is not None, (
"decode request holds metadata_buffer_index but scheduler has no "
"req_to_metadata_buffer_idx_allocator"
)
allocator.free(idx)
req.metadata_buffer_index = -1
def _maybe_log_eagle_accept_debug(
self: Scheduler,
batch: ScheduleBatch,