[Spec] Polish FutureMap after #25879: rename callback, async guard, cleanup (#25962)

This commit is contained in:
Liangsheng Yin
2026-05-21 13:56:22 -07:00
committed by GitHub
parent 17dadebd4e
commit c9a0e55eb5
4 changed files with 27 additions and 26 deletions
+13 -5
View File
@@ -1,7 +1,7 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import TYPE_CHECKING, Optional
from typing import TYPE_CHECKING, Union
import torch
@@ -67,8 +67,8 @@ class FutureMap:
# Forward-only bufs are lazy (worker-dependent shape).
self._forward_buf_initialized = False
# Fences the schedule-consumed buf fields.
self.publish_ready: Optional[torch.cuda.Event] = None
# Fences schedule-consumed buf fields; lazy device.Event() (cuda/hip-agnostic).
self.publish_ready = None
def _lazy_init_forward_buf(self, draft_input: EagleDraftInput):
self._forward_buf_initialized = True
@@ -115,6 +115,9 @@ class FutureMap:
draft_input.new_seq_lens = self.new_seq_lens_buf[indices]
# Resolve seq_lens placeholder (-indices) to the post-verify view.
batch.seq_lens = draft_input.new_seq_lens
# Async guard: catches a (-indices) sentinel slipping through if
# publish_ready fencing or buf indexing is wrong.
torch._assert_async((batch.seq_lens > 0).all())
if spec_need_hidden_states():
draft_input.hidden_states = self.hidden_states_buf[indices]
@@ -141,11 +144,16 @@ class FutureMap:
self.publish_ready = torch.get_device_module(self.device).Event()
self.publish_ready.record()
def stash(self, future_indices: FutureIndices, payload) -> None:
def stash(
self,
future_indices: FutureIndices,
payload: Union[torch.Tensor, EagleDraftInput],
) -> None:
"""Store forward-only fields for the next forward batch to pick up."""
indices = future_indices.indices
if indices.shape[0] == 0:
return # DP idle
# DP idle: payload is empty stub; lazy-init shape peek would IndexError.
return
if self.spec_algo.is_none():
# next_token_ids is int32; buf is int64. Advanced indexing requires
# an explicit cast.
+4 -11
View File
@@ -2842,12 +2842,9 @@ class Scheduler(
# Run forward
if self.is_generation:
if self.enable_overlap:
# Spec v2 pre-isolation CPU mirror prep: D2H new_seq_lens_buf
# into batch.seq_lens_cpu + set seq_lens_sum. For non-spec_v2,
# ForwardBatch.init_new lazily computes the sum.
if batch.is_spec_v2:
# FIXME: make this optional to different backends.
self.future_map.resolve_seq_lens_cpu(batch)
# Self-gates on batch.spec_info.future_indices; non-spec_v2
# no-ops (ForwardBatch.init_new lazily computes the sum).
self.future_map.resolve_seq_lens_cpu(batch)
with self._overlap_forward_isolation(batch):
future_indices = FutureIndices(indices=batch.req_pool_indices)
@@ -2856,11 +2853,7 @@ class Scheduler(
# draft_extend; publish moves the fence to verify-end so
# schedule prep can overlap with draft_extend.
fwd_kwargs = (
{
"on_verify_complete": partial(
self.future_map.publish, future_indices
)
}
{"on_publish": partial(self.future_map.publish, future_indices)}
if batch.is_spec_v2
else {}
)