From 44ec2ee18dc456a18b0f1f6ac6f03472c12cee60 Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Thu, 21 May 2026 13:56:44 -0700 Subject: [PATCH] [core] Unify output_tokens_buf in FutureMap (#25922) --- python/sglang/srt/managers/overlap_utils.py | 26 ++++++++++----------- python/sglang/srt/managers/scheduler.py | 2 +- 2 files changed, 13 insertions(+), 15 deletions(-) diff --git a/python/sglang/srt/managers/overlap_utils.py b/python/sglang/srt/managers/overlap_utils.py index f75576687..b634ec65d 100644 --- a/python/sglang/srt/managers/overlap_utils.py +++ b/python/sglang/srt/managers/overlap_utils.py @@ -55,16 +55,17 @@ class FutureMap: self.spec_algo = spec_algo self.req_pool_size = req_to_token_pool.req_to_token.shape[0] - if self.spec_algo.is_none(): - self.token_ids_buf = torch.empty( - (self.req_pool_size,), dtype=torch.int64, device=self.device - ) - else: + # 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 ) - # Forward-only bufs are lazy (worker-dependent shape). + # Remaining forward-only bufs are lazy (worker-dependent shape). self._forward_buf_initialized = False # Fences schedule-consumed buf fields; lazy device.Event() (cuda/hip-agnostic). @@ -85,9 +86,6 @@ class FutureMap: dtype=topk_index0.dtype, device=self.device, ) - self.bonus_tokens_buf = torch.empty( - (self.req_pool_size,), dtype=torch.int64, device=self.device - ) if spec_need_hidden_states(): hidden_states0 = draft_input.hidden_states[0] self.hidden_states_buf = torch.empty( @@ -98,7 +96,7 @@ class FutureMap: def resolve_future(self, batch: ScheduleBatch): if self.spec_algo.is_none(): - _resolve_future_token_ids(batch.input_ids, self.token_ids_buf) + _resolve_future_token_ids(batch.input_ids, self.output_tokens_buf) else: draft_input: EagleDraftInput = batch.spec_info if draft_input is None: @@ -111,7 +109,7 @@ class FutureMap: 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.bonus_tokens_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 @@ -157,14 +155,14 @@ class FutureMap: if self.spec_algo.is_none(): # next_token_ids is int32; buf is int64. Advanced indexing requires # an explicit cast. - self.token_ids_buf[indices] = payload.to(torch.int64) + self.output_tokens_buf[indices] = payload.to(torch.int64) return draft_input: EagleDraftInput = payload if not self._forward_buf_initialized: self._lazy_init_forward_buf(draft_input) - self.bonus_tokens_buf[indices] = draft_input.bonus_tokens.to( - self.bonus_tokens_buf.dtype + self.output_tokens_buf[indices] = draft_input.bonus_tokens.to( + self.output_tokens_buf.dtype ) self.topk_p_buf[indices] = draft_input.topk_p.to(self.topk_p_buf.dtype) self.topk_index_buf[indices] = draft_input.topk_index.to( diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 3101b833d..55f497107 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -2889,7 +2889,7 @@ class Scheduler( batch_result.future_indices = future_indices # Placeholder for next iter's resolve_future to look up the - # real token from token_ids_buf via the negated indices. + # real token from output_tokens_buf via the negated indices. batch.input_ids = -future_indices.indices if batch.is_spec_v2: