[core] step 1: route non-spec seq_lens via FutureMap with per-mode bootstrap fixes (#25944)

This commit is contained in:
Liangsheng Yin
2026-05-21 20:15:51 -07:00
committed by GitHub
parent cc77c36029
commit 8b473aa0bc
11 changed files with 95 additions and 76 deletions
@@ -173,16 +173,25 @@ class ScheduleBatchDisaggregationDecodeMixin:
topk_index=topk_index,
hidden_states=hidden_states,
bonus_tokens=last_tokens_tensor,
new_seq_lens=self.seq_lens,
)
spec_info.capture_hidden_mode = CaptureHiddenMode.LAST
if self.enable_overlap:
from sglang.srt.managers.overlap_utils import FutureIndices
spec_info.future_indices = FutureIndices(indices=self.req_pool_indices)
future_map.publish(spec_info.future_indices, spec_info.new_seq_lens)
future_map.publish(spec_info.future_indices, self.seq_lens)
future_map.stash(spec_info.future_indices, spec_info)
self.spec_info = spec_info
else:
# Non-spec: input_ids feeds the next decode forward directly.
self.input_ids = last_tokens_tensor
if self.enable_overlap:
from sglang.srt.managers.overlap_utils import FutureIndices
future_indices = FutureIndices(indices=self.req_pool_indices)
# Bootstrap FutureMap so the first DECODE after PREBUILT can
# resolve_future from buf. Non-spec convention: batch.seq_lens
# at decode forward INCLUDES this iter's new token, so publish
# current + 1.
future_map.publish(future_indices, self.seq_lens + 1)
future_map.stash(future_indices, last_tokens_tensor)
+37 -42
View File
@@ -48,28 +48,22 @@ class FutureMap:
spec_algo: SpeculativeAlgorithm,
req_to_token_pool: ReqToTokenPool,
):
# All buffers are indexed by req_pool_idx. Slot 0 mirrors the KV cache
# pool's padding row, so CUDA-graph padded batches (req_pool_idx == 0)
# read/write here harmlessly.
# Bufs indexed by req_pool_idx; slot 0 mirrors KV padding row so
# CUDA-graph padded batches (req_pool_idx == 0) are harmless.
self.device = device
self.spec_algo = spec_algo
self.req_pool_size = req_to_token_pool.req_to_token.shape[0]
# Forward-only token slot, eager (int64 fixed). Both modes use it:
# non-spec stashes next_token_ids; spec stashes bonus_tokens.
self.output_tokens_buf = torch.empty(
(self.req_pool_size,), dtype=torch.int64, device=self.device
)
if not self.spec_algo.is_none():
# Schedule-consumed buf, eager fixed dtype.
self.new_seq_lens_buf = torch.empty(
(self.req_pool_size,), dtype=torch.int64, device=self.device
)
# Remaining forward-only bufs are lazy (worker-dependent shape).
self.new_seq_lens_buf = torch.empty(
(self.req_pool_size,), dtype=torch.int64, device=self.device
)
if self.spec_algo.is_some():
self._forward_buf_initialized = False
# Fences schedule-consumed buf fields; lazy device.Event() (cuda/hip-agnostic).
self.publish_ready = None
self.publish_ready = None # lazy device.Event(); only spec_v2 needs it
def _lazy_init_forward_buf(self, draft_input: EagleDraftInput):
self._forward_buf_initialized = True
@@ -95,29 +89,34 @@ class FutureMap:
)
def resolve_future(self, batch: ScheduleBatch):
if batch.forward_mode.is_decode():
batch.seq_lens = self.new_seq_lens_buf[batch.req_pool_indices]
torch._assert_async((batch.seq_lens > 0).all())
if self.spec_algo.is_none():
_resolve_future_token_ids(batch.input_ids, self.output_tokens_buf)
else:
draft_input: EagleDraftInput = batch.spec_info
if draft_input is None:
# FIXME(lsyin): No future exists, only for prefill batch, not compatible with mixed mode
return
indices = draft_input.future_indices.indices
# FIXME: redundant. `indices` = batch.req_pool_indices, pinned via
# record_batch_in_overlap's attr_snapshot for 2 iters; refcount > 0
# across forward's read, allocator can't reclaim. Safe to remove.
indices.record_stream(torch.get_device_module(self.device).current_stream())
draft_input.topk_p = self.topk_p_buf[indices]
draft_input.topk_index = self.topk_index_buf[indices]
draft_input.bonus_tokens = self.output_tokens_buf[indices]
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]
self._resolve_spec_extras(batch)
def _resolve_spec_extras(self, batch: ScheduleBatch) -> None:
draft_input: EagleDraftInput = batch.spec_info
if draft_input is None:
# FIXME(lsyin): only prefill; not compatible with mixed mode
return
indices = draft_input.future_indices.indices
# FIXME: indices = batch.req_pool_indices, pinned 2 iters via
# record_batch_in_overlap; record_stream here is redundant.
indices.record_stream(torch.get_device_module(self.device).current_stream())
draft_input.topk_p = self.topk_p_buf[indices]
draft_input.topk_index = self.topk_index_buf[indices]
draft_input.bonus_tokens = self.output_tokens_buf[indices]
if spec_need_hidden_states():
draft_input.hidden_states = self.hidden_states_buf[indices]
def invalidate(self, batch: ScheduleBatch, future_indices: FutureIndices) -> None:
sentinel = -future_indices.indices
batch.input_ids = sentinel
batch.seq_lens = sentinel
def resolve_seq_lens_cpu(self, batch: ScheduleBatch) -> None:
fi = batch.spec_info.future_indices if batch.spec_info is not None else None
@@ -131,30 +130,26 @@ class FutureMap:
def publish(
self, future_indices: FutureIndices, new_seq_lens: torch.Tensor
) -> None:
"""Store schedule-consumed fields and signal publish_ready."""
if self.spec_algo.is_none():
return
indices = future_indices.indices
if indices.shape[0] == 0:
return # DP idle
self.new_seq_lens_buf[indices] = new_seq_lens.to(self.new_seq_lens_buf.dtype)
if self.publish_ready is None:
self.publish_ready = torch.get_device_module(self.device).Event()
self.publish_ready.record()
# Fast path: only spec_v2 needs the event (schedule-stream D2H sync).
if self.spec_algo.is_some():
if self.publish_ready is None:
self.publish_ready = torch.get_device_module(self.device).Event()
self.publish_ready.record()
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:
# 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.
self.output_tokens_buf[indices] = payload.to(torch.int64)
return
+9 -2
View File
@@ -2159,6 +2159,14 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
req.fill_ids = req.origin_input_ids + req.output_ids
req.set_extend_input_len(1)
if running_batch.enable_overlap:
# running_batch.seq_lens (GPU) is the FutureMap sentinel between iters;
# restore from CPU shadow before merge so MIXED's seq_lens has real values.
# (resolve_future only restores for is_decode(), not is_mixed().)
running_batch.seq_lens = running_batch.seq_lens_cpu.to(
running_batch.device, non_blocking=True
)
input_ids = torch.cat([self.input_ids, running_batch.input_ids])
out_cache_loc = torch.cat([self.out_cache_loc, running_batch.out_cache_loc])
@@ -2411,8 +2419,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
# Update seq_lens after allocation
if self.enable_overlap:
# Do not use in-place operations in the overlap mode
self.seq_lens = self.seq_lens + 1
# Overlap: GPU seq_lens restored by resolve_future from FutureMap buf.
self.seq_lens_cpu = self.seq_lens_cpu + 1
self.orig_seq_lens = self.orig_seq_lens + 1
else:
+6 -9
View File
@@ -2849,9 +2849,9 @@ class Scheduler(
with self._overlap_forward_isolation(batch):
future_indices = FutureIndices(indices=batch.req_pool_indices)
# Spec_v2 worker fires this between sample-end and
# draft_extend; publish moves the fence to verify-end so
# schedule prep can overlap with draft_extend.
# Spec_v2 fires on_publish mid-worker (between verify and
# draft_extend) so schedule prep can overlap with draft_extend.
# Non-spec has no later work — scheduler publishes after return.
fwd_kwargs = (
{"on_publish": partial(self.future_map.publish, future_indices)}
if batch.is_spec_v2
@@ -2865,6 +2865,8 @@ class Scheduler(
batch_result = self.model_worker.forward_batch_generation(
batch, **fwd_kwargs
)
if not batch.is_spec_v2:
self.future_map.publish(future_indices, batch.seq_lens + 1)
# Park any refs the worker wants kept alive 2 iters
# (cross-stream tensor lifetime; pinned in the same
# ring slot as the SB attr snapshot).
@@ -2888,16 +2890,11 @@ class Scheduler(
else:
batch_result.future_indices = future_indices
# Placeholder for next iter's resolve_future to look up the
# real token from output_tokens_buf via the negated indices.
batch.input_ids = -future_indices.indices
self.future_map.invalidate(batch, future_indices)
if batch.is_spec_v2:
batch.spec_info = batch_result.next_draft_input
batch.spec_info.future_indices = future_indices
# Schedule-stream sentinel between iters; next iter's
# resolve_future reassigns batch.seq_lens from new_seq_lens_buf.
batch.seq_lens = -future_indices.indices
elif self.enable_pdmux and batch.forward_mode.is_split_prefill():
batch_result = self.tp_worker.forward_batch_split_prefill(batch)
if isinstance(batch_result.next_token_ids, torch.Tensor):
+1 -1
View File
@@ -484,7 +484,7 @@ class TpModelWorker(BaseTpWorker):
)
if is_verify:
# Skip sampling and return logits for target forward
# Skip sampling; spec_v2 worker fires its own publish post-verify.
return batch_result
if (
+3
View File
@@ -47,6 +47,9 @@ class GenerationBatchResult:
# sync path: forward stream -> output processor
accept_lens: Optional[torch.Tensor] = None
# Next-iter seq_lens; published via on_publish.
new_seq_lens: Optional[torch.Tensor] = None
# relay path: forward stream -> next step forward
next_draft_input: Optional[EagleDraftInput] = None
+11 -5
View File
@@ -531,7 +531,13 @@ def alloc_for_decode(batch: ScheduleBatch, token_per_req: int) -> torch.Tensor:
batch.maybe_evict_swa()
bs = batch.seq_lens.shape[0]
if batch.enable_overlap:
# batch.seq_lens (GPU) is a sentinel between iters (FutureMap.invalidate);
# materialize from CPU shadow for the allocator. Tensor stays local.
seq_lens_gpu = batch.seq_lens_cpu.to(batch.device, non_blocking=True)
else:
seq_lens_gpu = batch.seq_lens
bs = seq_lens_gpu.shape[0]
if batch.tree_cache.page_size == 1:
# Non-paged allocation
@@ -539,9 +545,9 @@ def alloc_for_decode(batch: ScheduleBatch, token_per_req: int) -> torch.Tensor:
else:
# Paged allocation
last_loc = batch.req_to_token_pool.req_to_token[
batch.req_pool_indices, batch.seq_lens - 1
batch.req_pool_indices, seq_lens_gpu - 1
]
seq_lens_next = batch.seq_lens + token_per_req
seq_lens_next = seq_lens_gpu + token_per_req
out_cache_loc = alloc_paged_token_slots_decode(
tree_cache=batch.tree_cache,
seq_lens=seq_lens_next,
@@ -552,9 +558,9 @@ def alloc_for_decode(batch: ScheduleBatch, token_per_req: int) -> torch.Tensor:
# Write to req_to_token_pool
if batch.model_config.is_encoder_decoder:
locs = batch.encoder_lens + batch.seq_lens
locs = batch.encoder_lens + seq_lens_gpu
else:
locs = batch.seq_lens.clone()
locs = seq_lens_gpu.clone()
batch.req_to_token_pool.write(
(batch.req_pool_indices, locs), out_cache_loc.to(torch.int32)
@@ -695,7 +695,6 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
# V2 overlap worker only
future_indices: Optional[FutureIndices] = None
new_seq_lens: Optional[torch.Tensor] = None
# V2 reuses `EagleDraftInput` across phases (V1 has a separate
# `EagleDraftExtendInput` for these). Set during V2's draft-extend.
num_correct_drafts: Optional[torch.Tensor] = None
@@ -742,7 +741,6 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
topk_p=torch.empty((0, topk), device=device, dtype=torch.float32),
topk_index=torch.empty((0, topk), device=device, dtype=torch.int64),
capture_hidden_mode=capture_hidden_mode,
new_seq_lens=torch.empty((0,), device=device, dtype=torch.int32),
)
def filter_batch(self, new_indices: torch.Tensor, has_been_filtered: bool = True):
@@ -552,7 +552,6 @@ class EagleDraftWorker(BaseDraftWorker):
next_draft_input = EagleDraftInput(
hidden_states=target_hidden_states,
bonus_tokens=next_token_ids,
new_seq_lens=batch.seq_lens,
# draft mode is same with decode mode, only 1 token per req
num_tokens_per_req=1,
num_tokens_for_logprob_per_req=1,
@@ -772,9 +771,12 @@ class EAGLEWorkerV2(BaseSpecWorker):
batch.capture_hidden_mode = target_capture_mode
batch_output = self.target_worker.forward_batch_generation(batch)
# Spec_v2 convention: batch.seq_lens = length BEFORE this iter's tokens.
# Extend processed L prompt tokens; next verify iter expects same L.
batch_output.new_seq_lens = batch.seq_lens
# Publish before draft_extend so the fence is at target-end.
if on_publish is not None:
on_publish(batch.seq_lens)
on_publish(batch_output.new_seq_lens)
# Draft prefill
with (
@@ -820,7 +822,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
batch_output = self.verify(batch)
# Publish before draft_extend so the fence is at verify-end.
if on_publish is not None:
on_publish(batch_output.next_draft_input.new_seq_lens)
on_publish(batch_output.new_seq_lens)
with (
self.draft_worker.draft_tp_context(
self.draft_worker.draft_runner.tp_group
@@ -1097,9 +1099,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
batch, logits_output, predict, accept_index, self.speculative_num_steps
)
next_draft_input = EagleDraftInput(
bonus_tokens=bonus_tokens, new_seq_lens=new_seq_lens
)
next_draft_input = EagleDraftInput(bonus_tokens=bonus_tokens)
# verify_forward_batch transitively holds verify-time GPU tensors
# (draft_token / out_cache_loc / ...) that must outlive the imminent
@@ -1112,6 +1112,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
speculative_num_draft_tokens=self.speculative_num_draft_tokens,
next_draft_input=next_draft_input,
accept_lens=accept_lens,
new_seq_lens=new_seq_lens,
routed_experts_output=forward_batch_output.routed_experts_output,
indexer_topk_output=forward_batch_output.indexer_topk_output,
extra_keep_alive_refs=[verify_forward_batch],
@@ -384,7 +384,6 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker):
next_draft_input = EagleDraftInput(
hidden_states=target_hidden_states,
bonus_tokens=next_token_ids,
new_seq_lens=batch.seq_lens,
# draft mode is same with decode mode, only 1 token per req
num_tokens_per_req=1,
num_tokens_for_logprob_per_req=1,
@@ -674,9 +673,12 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
batch.capture_hidden_mode = target_capture_mode
batch_output = self.target_worker.forward_batch_generation(batch)
# Spec_v2 convention: batch.seq_lens = length BEFORE this iter's tokens.
# Extend processed L prompt tokens; next verify iter expects same L.
batch_output.new_seq_lens = batch.seq_lens
# Publish before draft_extend so the fence is at target-end.
if on_publish is not None:
on_publish(batch.seq_lens)
on_publish(batch_output.new_seq_lens)
# Chain-style MTP needs FULL to get all-token hidden states;
# non-chain only needs LAST (the target model's hidden states).
@@ -706,7 +708,7 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
batch_output = self.verify(batch)
# Publish before draft_extend so the fence is at verify-end.
if on_publish is not None:
on_publish(batch_output.next_draft_input.new_seq_lens)
on_publish(batch_output.new_seq_lens)
self.draft_worker._draft_extend_for_decode(batch, batch_output)
return batch_output
@@ -786,10 +788,7 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
batch, logits_output, predict, accept_index, self.speculative_num_steps
)
next_draft_input = EagleDraftInput(
bonus_tokens=bonus_tokens,
new_seq_lens=new_seq_lens,
)
next_draft_input = EagleDraftInput(bonus_tokens=bonus_tokens)
# verify_forward_batch transitively holds verify-time GPU tensors that
# must outlive the imminent batch.input_ids rebind; scheduler pins it
# in batch_record_buf via extra_keep_alive_refs. See EAGLEWorkerV2.verify.
@@ -800,6 +799,7 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
speculative_num_draft_tokens=self.speculative_num_draft_tokens,
next_draft_input=next_draft_input,
accept_lens=accept_lens,
new_seq_lens=new_seq_lens,
routed_experts_output=forward_batch_output.routed_experts_output,
indexer_topk_output=forward_batch_output.indexer_topk_output,
extra_keep_alive_refs=[verify_forward_batch],
@@ -82,6 +82,9 @@ class SpeculativeAlgorithm(Enum):
spec_class=spec_class,
)
def is_some(self) -> bool:
return self != SpeculativeAlgorithm.NONE
def is_none(self) -> bool:
return self == SpeculativeAlgorithm.NONE