From e1164a6dfc569cd65bcd6fc436ef60e8241af991 Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Thu, 11 Jun 2026 20:59:30 -0700 Subject: [PATCH] [Spec] Remove dead `prepare_for_verify` / `prepare_extend_after_decode` + extend-decode kernel (#27761) --- python/sglang/srt/speculative/eagle_info.py | 126 +----------------- python/sglang/srt/speculative/spec_utils.py | 3 - .../srt/speculative/triton_ops/cache_locs.py | 27 ---- 3 files changed, 6 insertions(+), 150 deletions(-) diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py index eadbdf74c..21be4c096 100644 --- a/python/sglang/srt/speculative/eagle_info.py +++ b/python/sglang/srt/speculative/eagle_info.py @@ -7,25 +7,12 @@ import torch from sglang.srt.constrained.base_grammar_backend import BaseGrammarObject from sglang.srt.environ import envs from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton -from sglang.srt.managers.schedule_batch import ScheduleBatch -from sglang.srt.mem_cache.common import ( - alloc_paged_token_slots_extend, - alloc_token_slots, - get_last_loc, -) from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode -from sglang.srt.server_args import get_global_server_args from sglang.srt.speculative.eagle_info_v2 import ( EagleDraftInputV2Mixin, EagleVerifyInputV2Mixin, ) from sglang.srt.speculative.spec_info import SpecInput, SpecInputType -from sglang.srt.speculative.spec_utils import ( - assign_req_to_token_pool_func, - create_extend_after_decode_spec_info, -) -from sglang.srt.utils import next_power_of_2 -from sglang.srt.utils.async_probe import maybe_detect_oob logger = logging.getLogger(__name__) @@ -94,67 +81,6 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): seq_lens_cpu=torch.empty((0,), dtype=torch.int64), ) - def prepare_for_verify(self, batch: ScheduleBatch, page_size: int): - - if batch.forward_mode.is_idle(): - return - - batch.input_ids = self.draft_token - maybe_detect_oob( - batch.input_ids, - 0, - batch.model_config.vocab_size, - "eagle prepare_for_verify input_ids", - ) - - if page_size == 1: - batch.out_cache_loc = alloc_token_slots( - batch.tree_cache, - len(batch.input_ids), - ) - end_offset = batch.seq_lens + self.draft_token_num - else: - prefix_lens = batch.seq_lens - prefix_lens_cpu = batch.seq_lens_cpu - end_offset = prefix_lens + self.draft_token_num - end_offset_cpu = prefix_lens_cpu + self.draft_token_num - last_loc = get_last_loc( - batch.req_to_token_pool.req_to_token, - batch.req_pool_indices, - prefix_lens, - ) - batch.out_cache_loc = alloc_paged_token_slots_extend( - batch.tree_cache, - prefix_lens, - prefix_lens_cpu, - end_offset, - end_offset_cpu, - last_loc, - len(batch.input_ids), - ) - - bs = batch.batch_size() - assign_req_to_token_pool_func( - batch.req_pool_indices, - batch.req_to_token_pool.req_to_token, - batch.seq_lens, - end_offset, - batch.out_cache_loc, - bs, - ) - - if get_global_server_args().enable_mamba_extra_buffer(): - batch.mamba_track_indices = torch.tensor( - [ - req.mamba_ping_pong_track_buffer[req.mamba_next_track_idx] - for req in batch.reqs - ], - dtype=torch.int64, - device=batch.device, - ) - batch.mamba_track_mask = None - batch.mamba_track_seqlens = None - def generate_attn_arg_prefill( self, req_pool_indices: torch.Tensor, @@ -230,8 +156,7 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): capture_hidden_mode: CaptureHiddenMode = CaptureHiddenMode.FULL # Per-req bonus token (the "+1" target prediction at end of each accept - # chain). Written by `EagleDraftExtendInput.prepare_extend_after_decode`; - # the worker copies it here for next iter's draft. + # chain); the worker copies it here post-extend for next iter's draft. bonus_tokens: torch.Tensor = None # shape: (b + 1,) @@ -365,16 +290,13 @@ class EagleDraftExtendInput(SpecInput): hidden_states: Optional[torch.Tensor] = None # Per-req accept counts. `num_accept_tokens = num_correct_drafts + 1`. - # Both kept for cuda-graph buffer indexing and the - # `create_extend_after_decode_spec_info` kernel. + # Both kept for cuda-graph buffer indexing. num_correct_drafts: torch.Tensor = None num_accept_tokens: torch.Tensor = None # CPU view, read by attention backends during the extend forward. num_accept_tokens_cpu: List[int] = None - # Batch-state slices for the draft-extend forward. Set by verify (sliced to - # reqs continuing into next iter). `prepare_extend_after_decode` copies - # these onto `batch.{input_ids, seq_lens, seq_lens_cpu, req_pool_indices}`. + # Per-req batch-state slices for the draft-extend forward: # - input_ids: accept tokens flat over surviving reqs # - seq_lens / _cpu: per-req sequence length (post-accept) # - req_pool_indices: per-req kv-pool slot @@ -383,10 +305,9 @@ class EagleDraftExtendInput(SpecInput): seq_lens_cpu: torch.Tensor = None req_pool_indices: torch.Tensor = None - # Set by `prepare_extend_after_decode`: - # - positions: kernel-written, shape `[total_accepted]`. - # - bonus_tokens: kernel-written, shape `[bs]`. The worker reads this - # post-extend to populate next iter's `EagleDraftInput.bonus_tokens`. + # - positions: shape `[total_accepted]`. + # - bonus_tokens: shape `[bs]`; read post-extend to populate next iter's + # `EagleDraftInput.bonus_tokens`. positions: Optional[torch.Tensor] = None bonus_tokens: Optional[torch.Tensor] = None @@ -457,41 +378,6 @@ class EagleDraftExtendInput(SpecInput): capture_hidden_mode=capture_hidden_mode, ) - def prepare_extend_after_decode( - self, - batch: ScheduleBatch, - speculative_num_steps: int, - ): - # Caller must have installed `self` as `batch.spec_info` before calling. - assert batch.spec_info is self - if batch.forward_mode.is_idle(): - return - - # The kernel below populates `self.positions` and `self.bonus_tokens`; - # the worker reads `self.bonus_tokens` to construct next iter's - # `EagleDraftInput`. - batch.input_ids = self.input_ids - batch.extend_lens = self.num_accept_tokens_cpu - batch.extend_num_tokens = sum(batch.extend_lens) - batch.seq_lens = self.seq_lens - batch.seq_lens_cpu = self.seq_lens_cpu - batch.req_pool_indices = self.req_pool_indices - batch.return_logprob = False - batch.return_hidden_states = False - - self.capture_hidden_mode = CaptureHiddenMode.LAST - self.positions = torch.empty_like(batch.input_ids, dtype=torch.long) - self.bonus_tokens = torch.empty_like(self.num_accept_tokens, dtype=torch.int32) - - create_extend_after_decode_spec_info[(len(batch.seq_lens),)]( - batch.input_ids, - batch.seq_lens, - self.num_accept_tokens, - self.positions, - self.bonus_tokens, - next_power_of_2(max(speculative_num_steps + 1, len(batch.seq_lens))), - ) - def generate_attn_arg_prefill( self, req_pool_indices: torch.Tensor, diff --git a/python/sglang/srt/speculative/spec_utils.py b/python/sglang/srt/speculative/spec_utils.py index b94b5b3b9..fb3e7b3de 100644 --- a/python/sglang/srt/speculative/spec_utils.py +++ b/python/sglang/srt/speculative/spec_utils.py @@ -28,9 +28,6 @@ from sglang.srt.speculative.triton_ops.cache_locs import ( from sglang.srt.speculative.triton_ops.cache_locs import ( assign_req_to_token_pool_func as assign_req_to_token_pool_func, ) -from sglang.srt.speculative.triton_ops.cache_locs import ( - create_extend_after_decode_spec_info as create_extend_after_decode_spec_info, -) from sglang.srt.speculative.triton_ops.cache_locs import ( filter_finished_cache_loc_kernel as filter_finished_cache_loc_kernel, ) diff --git a/python/sglang/srt/speculative/triton_ops/cache_locs.py b/python/sglang/srt/speculative/triton_ops/cache_locs.py index e16a15ad5..35894e2e1 100644 --- a/python/sglang/srt/speculative/triton_ops/cache_locs.py +++ b/python/sglang/srt/speculative/triton_ops/cache_locs.py @@ -12,33 +12,6 @@ _is_npu = is_npu() _is_musa = is_musa() -@triton.jit -def create_extend_after_decode_spec_info( - accept_tokens, - seq_lens, - accept_lens, - positions, - bonus_tokens_ptr, - bs_upper: tl.constexpr, -): - pid = tl.program_id(axis=0) - offsets = tl.arange(0, bs_upper) - seq_length = tl.load(seq_lens + pid) - # `accept_lens` includes the bonus token; load this req's value. - accept_len = tl.load(accept_lens + pid) - - accept_len_cumsum = tl.sum( - tl.load(accept_lens + offsets, mask=offsets < pid, other=0) - ) - positions_ptr = positions + accept_len_cumsum - mask = offsets < accept_len - tl.store(positions_ptr + offsets, seq_length - accept_len + offsets, mask) - - accept_len_cumsum += accept_len - 1 - bonus_token = tl.load(accept_tokens + accept_len_cumsum) - tl.store(bonus_tokens_ptr + pid, bonus_token) - - @triton.jit def assign_req_to_token_pool( req_pool_indices,