From 8b473aa0bca1da87c403cbcd6b8e8433800d04d1 Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Thu, 21 May 2026 20:15:51 -0700 Subject: [PATCH] [core] step 1: route non-spec `seq_lens` via `FutureMap` with per-mode bootstrap fixes (#25944) --- .../decode_schedule_batch_mixin.py | 13 ++- python/sglang/srt/managers/overlap_utils.py | 79 +++++++++---------- python/sglang/srt/managers/schedule_batch.py | 11 ++- python/sglang/srt/managers/scheduler.py | 15 ++-- python/sglang/srt/managers/tp_worker.py | 2 +- python/sglang/srt/managers/utils.py | 3 + python/sglang/srt/mem_cache/common.py | 16 ++-- python/sglang/srt/speculative/eagle_info.py | 2 - .../sglang/srt/speculative/eagle_worker_v2.py | 13 +-- .../multi_layer_eagle_worker_v2.py | 14 ++-- python/sglang/srt/speculative/spec_info.py | 3 + 11 files changed, 95 insertions(+), 76 deletions(-) diff --git a/python/sglang/srt/disaggregation/decode_schedule_batch_mixin.py b/python/sglang/srt/disaggregation/decode_schedule_batch_mixin.py index 76cc76e93..92fd55ab7 100644 --- a/python/sglang/srt/disaggregation/decode_schedule_batch_mixin.py +++ b/python/sglang/srt/disaggregation/decode_schedule_batch_mixin.py @@ -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) diff --git a/python/sglang/srt/managers/overlap_utils.py b/python/sglang/srt/managers/overlap_utils.py index b634ec65d..1458191d2 100644 --- a/python/sglang/srt/managers/overlap_utils.py +++ b/python/sglang/srt/managers/overlap_utils.py @@ -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 diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index b54e16f7e..70886301f 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -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: diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 55f497107..701a8eab2 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -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): diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py index 773e61f67..f552102b8 100644 --- a/python/sglang/srt/managers/tp_worker.py +++ b/python/sglang/srt/managers/tp_worker.py @@ -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 ( diff --git a/python/sglang/srt/managers/utils.py b/python/sglang/srt/managers/utils.py index 33db3942b..bffe856b9 100644 --- a/python/sglang/srt/managers/utils.py +++ b/python/sglang/srt/managers/utils.py @@ -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 diff --git a/python/sglang/srt/mem_cache/common.py b/python/sglang/srt/mem_cache/common.py index 56fccddc9..ca3733e08 100644 --- a/python/sglang/srt/mem_cache/common.py +++ b/python/sglang/srt/mem_cache/common.py @@ -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) diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py index bcdeaf069..b54f75a81 100644 --- a/python/sglang/srt/speculative/eagle_info.py +++ b/python/sglang/srt/speculative/eagle_info.py @@ -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): diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 3f24eebe4..a3d14af91 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -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], diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py index aebd5d0ab..f5f592aa8 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -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], diff --git a/python/sglang/srt/speculative/spec_info.py b/python/sglang/srt/speculative/spec_info.py index ca2be5666..75b9af39f 100644 --- a/python/sglang/srt/speculative/spec_info.py +++ b/python/sglang/srt/speculative/spec_info.py @@ -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