From 78da0d3106741e1c0f58637a5054e1fb27575b30 Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Fri, 8 May 2026 23:24:18 -0700 Subject: [PATCH] [Spec] Move `accept_tokens` off `EagleDraftInput`; pass via method arg (#24735) --- python/sglang/srt/speculative/dflash_info.py | 1 - python/sglang/srt/speculative/eagle_info.py | 170 +++++++++++------- python/sglang/srt/speculative/eagle_worker.py | 21 ++- .../srt/speculative/frozen_kv_mtp_info.py | 2 +- .../srt/speculative/frozen_kv_mtp_worker.py | 22 +-- .../speculative/multi_layer_eagle_worker.py | 21 ++- python/sglang/srt/speculative/ngram_info.py | 1 - 7 files changed, 144 insertions(+), 94 deletions(-) diff --git a/python/sglang/srt/speculative/dflash_info.py b/python/sglang/srt/speculative/dflash_info.py index 41162425c..9cbba1faa 100644 --- a/python/sglang/srt/speculative/dflash_info.py +++ b/python/sglang/srt/speculative/dflash_info.py @@ -213,7 +213,6 @@ class DFlashVerifyInput(SpecInput): last_loc, len(batch.input_ids), ) - self.last_loc = last_loc bs = batch.batch_size() assign_req_to_token_pool_func( diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py index 8e1030392..205c0610a 100644 --- a/python/sglang/srt/speculative/eagle_info.py +++ b/python/sglang/srt/speculative/eagle_info.py @@ -140,7 +140,6 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): last_loc, len(batch.input_ids), ) - self.last_loc = last_loc bs = batch.batch_size() assign_req_to_token_pool_func( @@ -241,23 +240,18 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): accepted token logits. """ if batch.forward_mode.is_idle(): - return EagleVerifyOutput( - draft_input=EagleDraftInput.create_idle_input( - device=batch.device, - hidden_size=batch.model_config.spec_hidden_size, - dtype=batch.model_config.dtype, - topk=self.topk, - capture_hidden_mode=CaptureHiddenMode.LAST, - ), + next_draft_input = EagleDraftInput.create_idle_input( + device=batch.device, + hidden_size=batch.model_config.spec_hidden_size, + dtype=batch.model_config.dtype, + topk=self.topk, + capture_hidden_mode=CaptureHiddenMode.LAST, + ) + return EagleVerifyOutput.create_idle( + next_draft_input=next_draft_input, logits_output=logits_output, - accept_tokens=torch.empty(0, dtype=torch.long, device=batch.device), - num_accepted_drafts_per_req_cpu=[], - accepted_indices=torch.full( - (0, self.spec_steps + 1), - -1, - dtype=torch.int32, - device=batch.device, - ), + device=batch.device, + spec_steps=self.spec_steps, ) bs = self.retrieve_index.shape[0] @@ -551,23 +545,22 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): batch.seq_lens.add_(num_accepted_drafts + 1) batch.seq_lens_cpu.add_(num_accepted_tokens_cpu) - draft_input = EagleDraftInput( + next_draft_input = EagleDraftInput( hidden_states=batch.spec_info.hidden_states[accept_index], - accept_tokens=accept_tokens, num_accepted_drafts=num_accepted_drafts, num_accepted_tokens=num_accepted_drafts + 1, - num_accepted_drafts_cpu=num_accepted_drafts_list, num_accepted_tokens_cpu=num_accepted_tokens_list, - seq_lens_for_draft_extend=batch.seq_lens, - seq_lens_for_draft_extend_cpu=batch.seq_lens_cpu, - req_pool_indices_for_draft_extend=batch.req_pool_indices, ) return EagleVerifyOutput( - draft_input=draft_input, + next_draft_input=next_draft_input, logits_output=logits_output, accept_tokens=accept_tokens, - num_accepted_drafts_per_req_cpu=draft_input.num_accepted_drafts_cpu, + unfinished_accept_tokens=accept_tokens, + seq_lens_for_draft_extend=batch.seq_lens, + seq_lens_for_draft_extend_cpu=batch.seq_lens_cpu, + req_pool_indices_for_draft_extend=batch.req_pool_indices, + num_accepted_drafts_per_req_cpu=num_accepted_drafts_list, accepted_indices=accept_index, ) else: @@ -621,23 +614,36 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): unfinished_num_accepted_drafts = num_accepted_drafts[ unfinished_index_device ] - draft_input = EagleDraftInput( + unfinished_accept_tokens = predict[unfinished_accept_index] + seq_lens_for_draft_extend = batch.seq_lens[unfinished_index_device] + seq_lens_for_draft_extend_cpu = batch.seq_lens_cpu[unfinished_index] + req_pool_indices_for_draft_extend = batch.req_pool_indices[ + unfinished_index_device + ] + next_draft_input = EagleDraftInput( hidden_states=batch.spec_info.hidden_states[ unfinished_accept_index ], - accept_tokens=predict[unfinished_accept_index], - num_accepted_drafts_cpu=draft_input_num_accepted_drafts_cpu, num_accepted_tokens_cpu=draft_input_num_accepted_tokens_cpu, num_accepted_drafts=unfinished_num_accepted_drafts, num_accepted_tokens=unfinished_num_accepted_drafts + 1, - seq_lens_for_draft_extend=batch.seq_lens[unfinished_index_device], - seq_lens_for_draft_extend_cpu=batch.seq_lens_cpu[unfinished_index], - req_pool_indices_for_draft_extend=batch.req_pool_indices[ - unfinished_index_device - ], ) else: - draft_input = EagleDraftInput.create_idle_input( + unfinished_accept_tokens = torch.empty( + (0,), dtype=accept_tokens.dtype, device=accept_tokens.device + ) + seq_lens_for_draft_extend = torch.empty( + (0,), dtype=batch.seq_lens.dtype, device=batch.seq_lens.device + ) + seq_lens_for_draft_extend_cpu = torch.empty( + (0,), dtype=batch.seq_lens_cpu.dtype + ) + req_pool_indices_for_draft_extend = torch.empty( + (0,), + dtype=batch.req_pool_indices.dtype, + device=batch.req_pool_indices.device, + ) + next_draft_input = EagleDraftInput.create_idle_input( device=batch.device, hidden_size=batch.model_config.spec_hidden_size, dtype=batch.model_config.dtype, @@ -646,9 +652,13 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): ) return EagleVerifyOutput( - draft_input=draft_input, + next_draft_input=next_draft_input, logits_output=logits_output, accept_tokens=accept_tokens, + unfinished_accept_tokens=unfinished_accept_tokens, + seq_lens_for_draft_extend=seq_lens_for_draft_extend, + seq_lens_for_draft_extend_cpu=seq_lens_for_draft_extend_cpu, + req_pool_indices_for_draft_extend=req_pool_indices_for_draft_extend, num_accepted_drafts_per_req_cpu=num_accepted_drafts_list, accepted_indices=accept_index, ) @@ -660,7 +670,10 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): # shape: (b, topk) topk_p: torch.Tensor = None topk_index: torch.Tensor = None - # shape: (b, hidden_size) + # shape: (b, hidden_size) when consumed by `draft` forward (one hidden per req); + # shape: (total_accepted, hidden_size) when consumed by `draft_extend` forward + # (one hidden per accepted token). Workers maintain this invariant locally; + # there is no type-level guard. Don't add new readers without checking phase. hidden_states: torch.Tensor = None capture_hidden_mode: CaptureHiddenMode = CaptureHiddenMode.FULL @@ -670,16 +683,10 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): # `num_accepted_tokens = num_accepted_drafts + 1` (per-req, one bonus per req). # Storing both avoids repeated `+ 1` at every consumer (attn backends, kernels). bonus_tokens: torch.Tensor = None - # Flat accepted-token tensor for draft-extend, shape `[sum_accepted]`. - # Set right after verify and consumed by `prepare_extend_after_decode` as - # the extend batch's `input_ids`. Dead after that method returns. - # TODO: drop this field and pass `accept_tokens` directly to - # `prepare_extend_after_decode` as a method arg. Its lifetime is bounded - # by verify -> prepare_extend, no need to live on the dataclass. - accept_tokens: torch.Tensor = None num_accepted_drafts: torch.Tensor = None num_accepted_tokens: torch.Tensor = None - num_accepted_drafts_cpu: List[int] = None + # Read by attention backends during draft-extend forward; kept on the + # dataclass because the backends access it via `forward_batch.spec_info`. num_accepted_tokens_cpu: List[int] = None # Inputs for the attention backends @@ -691,12 +698,6 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): num_tokens_per_req: int = -1 num_tokens_for_logprob_per_req: int = -1 - # Inputs for draft extend - # shape: (b,) - seq_lens_for_draft_extend: torch.Tensor = None - seq_lens_for_draft_extend_cpu: torch.Tensor = None - req_pool_indices_for_draft_extend: torch.Tensor = None - # Inputs for V2 overlap worker future_indices: Optional[FutureIndices] = None new_seq_lens: Optional[torch.Tensor] = None @@ -735,7 +736,6 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): ): return cls( bonus_tokens=torch.empty((0,), device=device, dtype=torch.int32), - accept_tokens=torch.empty((0,), device=device, dtype=torch.int32), hidden_states=torch.empty((0, hidden_size), device=device, dtype=dtype), topk_p=torch.empty((0, topk), device=device, dtype=torch.float32), topk_index=torch.empty((0, topk), device=device, dtype=torch.int64), @@ -743,29 +743,29 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): new_seq_lens=torch.empty((0,), device=device, dtype=torch.int32), num_accepted_drafts=torch.empty((0,), device=device, dtype=torch.int32), num_accepted_tokens=torch.empty((0,), device=device, dtype=torch.int32), - num_accepted_drafts_cpu=[], num_accepted_tokens_cpu=[], ) def prepare_extend_after_decode( self, batch: ScheduleBatch, + verify_output: "EagleVerifyOutput", speculative_num_steps: int, ): if batch.forward_mode.is_idle(): return - # `self.accept_tokens` is the flat accepted-token tensor set by - # `EagleVerifyInput.verify`; use it as the extend batch's `input_ids`. - # The kernel below populates `self.bonus_tokens` ([bs] per-req) for - # the next decode round. - batch.input_ids = self.accept_tokens + # All transient verify->extend handoff state is read off `verify_output`, + # not from `self`. The kernel below populates `self.bonus_tokens` + # ([bs] per-req) for the next decode round; that is the only state on + # `self` that survives past this method. + batch.input_ids = verify_output.unfinished_accept_tokens batch.extend_lens = batch.spec_info.num_accepted_tokens_cpu batch.extend_num_tokens = sum(batch.extend_lens) - batch.seq_lens = batch.spec_info.seq_lens_for_draft_extend - batch.seq_lens_cpu = batch.spec_info.seq_lens_for_draft_extend_cpu - batch.req_pool_indices = batch.spec_info.req_pool_indices_for_draft_extend + batch.seq_lens = verify_output.seq_lens_for_draft_extend + batch.seq_lens_cpu = verify_output.seq_lens_for_draft_extend_cpu + batch.req_pool_indices = verify_output.req_pool_indices_for_draft_extend batch.return_logprob = False batch.return_hidden_states = False @@ -873,13 +873,53 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): @dataclass class EagleVerifyOutput: - # Draft input batch - draft_input: EagleDraftInput - # Logit outputs from target worker + # Next iter's persistent draft state, ready to be installed as `batch.spec_info`. + next_draft_input: EagleDraftInput + # Logit outputs from target worker. logits_output: LogitsProcessorOutput - # Accepted token ids including the bonus token (flat, [sum_accepted]) + # All accepted tokens flat across all reqs incl. those that finished this + # step. Includes the bonus token. Used for output processing. accept_tokens: torch.Tensor - # Accepted token length per sequence in a batch in CPU. + # Below are transient handoff fields for the next iter's draft-extend pass. + # They are scoped to the verify -> prepare_extend_after_decode window only; + # `prepare_extend_after_decode` reads them off this object via method arg + # rather than smuggling them through `EagleDraftInput`. + # + # Subset of `accept_tokens` for reqs continuing into next iter's draft-extend + # forward (= `accept_tokens` when no req finished; flat over unfinished + # reqs only otherwise). Becomes `batch.input_ids` for that forward pass. + unfinished_accept_tokens: torch.Tensor + # `batch.seq_lens` / `batch.seq_lens_cpu` / `batch.req_pool_indices` to + # use for the next iter's draft-extend forward; sliced to surviving reqs. + seq_lens_for_draft_extend: torch.Tensor + seq_lens_for_draft_extend_cpu: torch.Tensor + req_pool_indices_for_draft_extend: torch.Tensor + # Accepted token length per sequence in a batch in CPU (full set). num_accepted_drafts_per_req_cpu: List[int] # Accepted indices from logits_output.next_token_logits accepted_indices: torch.Tensor + + @classmethod + def create_idle( + cls, + *, + next_draft_input: EagleDraftInput, + logits_output: LogitsProcessorOutput, + device: torch.device, + spec_steps: int, + ) -> "EagleVerifyOutput": + return cls( + next_draft_input=next_draft_input, + logits_output=logits_output, + accept_tokens=torch.empty(0, dtype=torch.long, device=device), + unfinished_accept_tokens=torch.empty(0, dtype=torch.long, device=device), + seq_lens_for_draft_extend=torch.empty(0, dtype=torch.int32, device=device), + seq_lens_for_draft_extend_cpu=torch.empty(0, dtype=torch.int32), + req_pool_indices_for_draft_extend=torch.empty( + 0, dtype=torch.int64, device=device + ), + num_accepted_drafts_per_req_cpu=[], + accepted_indices=torch.full( + (0, spec_steps + 1), -1, dtype=torch.int32, device=device + ), + ) diff --git a/python/sglang/srt/speculative/eagle_worker.py b/python/sglang/srt/speculative/eagle_worker.py index 51a5cd35d..3c480e9a1 100644 --- a/python/sglang/srt/speculative/eagle_worker.py +++ b/python/sglang/srt/speculative/eagle_worker.py @@ -505,10 +505,10 @@ class EAGLEWorker(TpModelWorker): # when DP attention is enabled, but it is slow. Skip it for now. if ( self.server_args.enable_dp_attention - or batch.spec_info.accept_tokens.shape[0] > 0 + or verify_output.unfinished_accept_tokens.shape[0] > 0 ): # decode is not finished - self.forward_draft_extend_after_decode(batch) + self.forward_draft_extend_after_decode(batch, verify_output) set_time_batch( batch.reqs, "set_spec_draft_extend_end_time", trace_only=True @@ -527,8 +527,10 @@ class EAGLEWorker(TpModelWorker): can_run_cuda_graph=can_run_cuda_graph, ) - def check_forward_draft_extend_after_decode(self, batch: ScheduleBatch): - local_need_forward = batch.spec_info.accept_tokens.shape[0] > 0 + def check_forward_draft_extend_after_decode( + self, batch: ScheduleBatch, verify_output: EagleVerifyOutput + ): + local_need_forward = verify_output.unfinished_accept_tokens.shape[0] > 0 if not self.server_args.enable_dp_attention: return local_need_forward @@ -977,7 +979,7 @@ class EAGLEWorker(TpModelWorker): batch.forward_mode = ( ForwardMode.DECODE if not batch.forward_mode.is_idle() else ForwardMode.IDLE ) - batch.spec_info = res.draft_input + batch.spec_info = res.next_draft_input return logits_output, res, model_worker_batch, can_run_cuda_graph @@ -1104,7 +1106,9 @@ class EAGLEWorker(TpModelWorker): assert forward_batch.spec_info is batch.spec_info self.capture_for_decode(logits_output, forward_batch.spec_info) - def forward_draft_extend_after_decode(self, batch: ScheduleBatch): + def forward_draft_extend_after_decode( + self, batch: ScheduleBatch, verify_output: EagleVerifyOutput + ): assert isinstance(batch.spec_info, EagleDraftInput) # Backup fields that will be modified in-place seq_lens_backup = batch.seq_lens.clone() @@ -1116,7 +1120,7 @@ class EAGLEWorker(TpModelWorker): input_is_idle = batch.forward_mode.is_idle() - if not input_is_idle and batch.spec_info.accept_tokens.numel() == 0: + if not input_is_idle and verify_output.unfinished_accept_tokens.numel() == 0: batch = batch.copy() batch.prepare_for_idle() hidden_size = ( @@ -1137,7 +1141,8 @@ class EAGLEWorker(TpModelWorker): batch.spec_info.num_tokens_for_logprob_per_req = 1 batch.spec_info.prepare_extend_after_decode( batch, - self.speculative_num_steps, + verify_output=verify_output, + speculative_num_steps=self.speculative_num_steps, ) batch.forward_mode = ( ForwardMode.DRAFT_EXTEND diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_info.py b/python/sglang/srt/speculative/frozen_kv_mtp_info.py index 27a7249b0..d09244616 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_info.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_info.py @@ -62,7 +62,7 @@ class FrozenKVMTPVerifyInput(EagleVerifyInput): def verify(self, *args, **kwargs) -> EagleVerifyOutput: output = super().verify(*args, **kwargs) - output.draft_input = _to_frozen_kv_mtp_draft_input(output.draft_input) + output.next_draft_input = _to_frozen_kv_mtp_draft_input(output.next_draft_input) return output diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_worker.py b/python/sglang/srt/speculative/frozen_kv_mtp_worker.py index 64d7da42c..9039577cc 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_worker.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_worker.py @@ -43,6 +43,7 @@ from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig from sglang.srt.observability.req_time_stats import set_time_batch from sglang.srt.observability.trace import get_global_tracing_enabled from sglang.srt.server_args import ServerArgs +from sglang.srt.speculative.eagle_info import EagleVerifyOutput from sglang.srt.speculative.eagle_utils import ( build_tree_kernel_efficient, organize_draft_results, @@ -463,7 +464,7 @@ class FrozenKVMTPWorker(TpModelWorker): self.server_args.enable_dp_attention or batch.spec_info.bonus_tokens.numel() ): - self.forward_draft_extend_after_decode(batch) + self.forward_draft_extend_after_decode(batch, verify_output) set_time_batch(batch.reqs, "set_spec_draft_extend_end_time", trace_only=True) return GenerationBatchResult( @@ -504,7 +505,9 @@ class FrozenKVMTPWorker(TpModelWorker): mm_input_embeds=mm_input_embeds, ) - def forward_draft_extend_after_decode(self, batch: ScheduleBatch) -> None: + def forward_draft_extend_after_decode( + self, batch: ScheduleBatch, verify_output: EagleVerifyOutput + ) -> None: assert isinstance(batch.spec_info, FrozenKVMTPDraftInput) input_is_idle = batch.forward_mode.is_idle() if not input_is_idle and batch.spec_info.bonus_tokens.numel() == 0: @@ -527,19 +530,18 @@ class FrozenKVMTPWorker(TpModelWorker): req_pool_indices_backup = batch.req_pool_indices try: - if draft_input.seq_lens_for_draft_extend is not None: - # Verify may leave finished requests in ScheduleBatch; seed only - # the unfinished requests carried by draft_input. - batch.seq_lens = draft_input.seq_lens_for_draft_extend - batch.seq_lens_cpu = draft_input.seq_lens_for_draft_extend_cpu - batch.req_pool_indices = draft_input.req_pool_indices_for_draft_extend + # Verify may leave finished requests in ScheduleBatch; seed only + # the unfinished requests carried by `verify_output`. + batch.seq_lens = verify_output.seq_lens_for_draft_extend + batch.seq_lens_cpu = verify_output.seq_lens_for_draft_extend_cpu + batch.req_pool_indices = verify_output.req_pool_indices_for_draft_extend last_token_ids, last_hidden = self._select_last_verified_seed(draft_input) self._run_assistant_seed_step( batch, last_token_ids, last_hidden, - seq_lens_cpu=draft_input.seq_lens_for_draft_extend_cpu, + seq_lens_cpu=verify_output.seq_lens_for_draft_extend_cpu, draft_input=draft_input, ) finally: @@ -766,7 +768,7 @@ class FrozenKVMTPWorker(TpModelWorker): batch.forward_mode = ( ForwardMode.DECODE if not batch.forward_mode.is_idle() else ForwardMode.IDLE ) - batch.spec_info = res.draft_input + batch.spec_info = res.next_draft_input del seq_lens_pre_verify return logits_output, res, model_worker_batch, can_run_cuda_graph diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker.py b/python/sglang/srt/speculative/multi_layer_eagle_worker.py index 366000fd9..b03492905 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker.py @@ -283,10 +283,10 @@ class MultiLayerEagleWorker(TpModelWorker): # when DP attention is enabled, but it is slow. Skip it for now. if ( self.server_args.enable_dp_attention - or batch.spec_info.accept_tokens.shape[0] > 0 + or verify_output.unfinished_accept_tokens.shape[0] > 0 ): # decode is not finished - self.forward_draft_extend_after_decode(batch) + self.forward_draft_extend_after_decode(batch, verify_output) return GenerationBatchResult( logits_output=logits_output, @@ -295,8 +295,10 @@ class MultiLayerEagleWorker(TpModelWorker): can_run_cuda_graph=can_run_cuda_graph, ) - def check_forward_draft_extend_after_decode(self, batch: ScheduleBatch): - local_need_forward = batch.spec_info.accept_tokens.shape[0] > 0 + def check_forward_draft_extend_after_decode( + self, batch: ScheduleBatch, verify_output: EagleVerifyOutput + ): + local_need_forward = verify_output.unfinished_accept_tokens.shape[0] > 0 if not self.server_args.enable_dp_attention: return local_need_forward @@ -589,7 +591,7 @@ class MultiLayerEagleWorker(TpModelWorker): batch.forward_mode = ( ForwardMode.DECODE if not batch.forward_mode.is_idle() else ForwardMode.IDLE ) - batch.spec_info = res.draft_input + batch.spec_info = res.next_draft_input return logits_output, res, model_worker_batch, can_run_cuda_graph @@ -652,7 +654,9 @@ class MultiLayerEagleWorker(TpModelWorker): forward_batch.spec_info.topk_p = torch.cat(topk_p_list, dim=1) forward_batch.spec_info.topk_index = torch.cat(topk_index_list, dim=1) - def forward_draft_extend_after_decode(self, batch: ScheduleBatch): + def forward_draft_extend_after_decode( + self, batch: ScheduleBatch, verify_output: EagleVerifyOutput + ): assert isinstance(batch.spec_info, EagleDraftInput) # Backup fields that will be modified in-place seq_lens_backup = batch.seq_lens.clone() @@ -664,7 +668,7 @@ class MultiLayerEagleWorker(TpModelWorker): input_is_idle = batch.forward_mode.is_idle() - if not input_is_idle and batch.spec_info.accept_tokens.numel() == 0: + if not input_is_idle and verify_output.unfinished_accept_tokens.numel() == 0: batch = batch.copy() batch.prepare_for_idle() hidden_size = ( @@ -684,7 +688,8 @@ class MultiLayerEagleWorker(TpModelWorker): batch.spec_info.num_tokens_for_logprob_per_req = 1 batch.spec_info.prepare_extend_after_decode( batch, - self.speculative_num_steps, + verify_output=verify_output, + speculative_num_steps=self.speculative_num_steps, ) batch.forward_mode = ( ForwardMode.DRAFT_EXTEND diff --git a/python/sglang/srt/speculative/ngram_info.py b/python/sglang/srt/speculative/ngram_info.py index 5b9a0da54..0b60cfe86 100644 --- a/python/sglang/srt/speculative/ngram_info.py +++ b/python/sglang/srt/speculative/ngram_info.py @@ -106,7 +106,6 @@ class NgramVerifyInput(SpecInput): last_loc, len(batch.input_ids), ) - self.last_loc = last_loc bs = batch.batch_size() assign_req_to_token_pool[(bs,)](