From bd448e51bd8bb2863253fce4c8c675ddb13cd3e8 Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Wed, 29 Apr 2026 00:02:22 -0700 Subject: [PATCH] [Spec] Split `accept_length` into `num_accepted_drafts` and `num_accepted_tokens` (#23962) --- .../srt/layers/attention/aiter_backend.py | 10 +- .../attention/flashattention_backend.py | 8 +- .../nsa/nsa_backend_mtp_precompute.py | 5 +- .../srt/layers/attention/nsa_backend.py | 4 +- .../srt/layers/attention/triton_backend.py | 4 +- .../layers/attention/trtllm_mha_backend.py | 8 +- .../layers/attention/trtllm_mla_backend.py | 2 +- .../srt/layers/attention/wave_backend.py | 4 +- python/sglang/srt/layers/utils/logprob.py | 8 +- .../scheduler_output_processor_mixin.py | 8 +- .../sglang/srt/managers/tokenizer_manager.py | 2 +- python/sglang/srt/managers/utils.py | 2 +- .../srt/model_executor/forward_batch_info.py | 16 ++- .../srt/speculative/adaptive_runtime_state.py | 6 +- .../srt/speculative/adaptive_spec_params.py | 8 +- python/sglang/srt/speculative/dflash_info.py | 10 +- .../sglang/srt/speculative/dflash_worker.py | 10 +- .../eagle_draft_extend_cuda_graph_runner.py | 43 +++++-- python/sglang/srt/speculative/eagle_info.py | 109 ++++++++++-------- .../sglang/srt/speculative/eagle_info_v2.py | 28 ++--- python/sglang/srt/speculative/eagle_worker.py | 18 +-- .../sglang/srt/speculative/eagle_worker_v2.py | 32 ++--- ...er_eagle_draft_extend_cuda_graph_runner.py | 57 ++++++--- .../speculative/multi_layer_eagle_worker.py | 10 +- .../multi_layer_eagle_worker_v2.py | 10 +- python/sglang/srt/speculative/ngram_info.py | 46 +++++--- python/sglang/srt/speculative/ngram_worker.py | 14 ++- python/sglang/srt/speculative/spec_utils.py | 57 ++++----- .../test/attention/test_trtllm_mla_backend.py | 12 +- python/sglang/test/kits/spec_decoding_kit.py | 4 +- 30 files changed, 330 insertions(+), 225 deletions(-) diff --git a/python/sglang/srt/layers/attention/aiter_backend.py b/python/sglang/srt/layers/attention/aiter_backend.py index e52ae6984..0350038e9 100755 --- a/python/sglang/srt/layers/attention/aiter_backend.py +++ b/python/sglang/srt/layers/attention/aiter_backend.py @@ -878,7 +878,7 @@ class AiterAttnBackend(AttentionBackend): self.indices_updater_prefill.max_kv_len, ) elif forward_batch.forward_mode.is_draft_extend(): - # EAGLE V1: DRAFT_EXTEND mode - uses spec_info.accept_length + # EAGLE V1: DRAFT_EXTEND mode - uses spec_info.num_accepted_tokens if self.use_mla: kv_indices, kv_indptr, qo_indptr, custom_mask = ( spec_info.generate_attn_arg_prefill( @@ -946,7 +946,7 @@ class AiterAttnBackend(AttentionBackend): ) ) kv_indices = kv_indices.to(torch.int64) - draft_max_extend_len = torch.max(spec_info.accept_length).item() + draft_max_extend_len = torch.max(spec_info.num_accepted_tokens).item() self.forward_metadata = ForwardMetadata( kv_indptr, @@ -1956,12 +1956,12 @@ class AiterAttnBackend(AttentionBackend): num_kv_splits=num_kv_splits, ) elif forward_mode.is_draft_extend(): - # EAGLE V1: Uses spec_info.accept_length + # EAGLE V1: Uses spec_info.num_accepted_tokens num_tokens_per_bs = self.speculative_num_steps + 1 seq_lens = seq_lens[:bs] - accept_lens = spec_info.accept_length[:bs] + extend_lens = spec_info.num_accepted_tokens[:bs] qo_indptr = self.qo_indptr[: bs + 1] - qo_indptr[1 : bs + 1] = torch.cumsum(accept_lens, dim=0) + qo_indptr[1 : bs + 1] = torch.cumsum(extend_lens, dim=0) kv_indptr = self.kv_indptr[: bs + 1] kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0) kv_indices = self.cuda_graph_kv_indices diff --git a/python/sglang/srt/layers/attention/flashattention_backend.py b/python/sglang/srt/layers/attention/flashattention_backend.py index 193ce338c..5e2c77e28 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -2144,14 +2144,14 @@ class FlashAttentionBackend(AttentionBackend): metadata.cu_seqlens_k[1:].copy_( torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32) ) - accept_length = spec_info.accept_length[:bs] - if spec_info.accept_length_cpu: - metadata.max_seq_len_q = max(spec_info.accept_length_cpu) + 1 + extend_lens = spec_info.num_accepted_tokens[:bs] + if spec_info.num_accepted_tokens_cpu: + metadata.max_seq_len_q = max(spec_info.num_accepted_tokens_cpu) else: metadata.max_seq_len_q = 1 metadata.cu_seqlens_q[1:].copy_( - torch.cumsum(accept_length, dim=0, dtype=torch.int32) + torch.cumsum(extend_lens, dim=0, dtype=torch.int32) ) max_seq_pages = ( diff --git a/python/sglang/srt/layers/attention/nsa/nsa_backend_mtp_precompute.py b/python/sglang/srt/layers/attention/nsa/nsa_backend_mtp_precompute.py index 846d276e4..af8af4008 100644 --- a/python/sglang/srt/layers/attention/nsa/nsa_backend_mtp_precompute.py +++ b/python/sglang/srt/layers/attention/nsa/nsa_backend_mtp_precompute.py @@ -261,8 +261,9 @@ class NativeSparseAttnBackendMTPPrecomputeMixin: cache_seqlens = seq_lens.to(torch.int32) cu_seqlens_k = compute_cu_seqlens(cache_seqlens) - # Extend seqlens from spec_info - extend_seq_lens = spec_info.accept_length[:bs] + # Extend seqlens from spec_info: num_accepted_tokens already includes + # the bonus token (drafts + 1). + extend_seq_lens = spec_info.num_accepted_tokens[:bs] extend_seq_lens_cpu = extend_seq_lens.tolist() # Page indices (repeated per accept length) diff --git a/python/sglang/srt/layers/attention/nsa_backend.py b/python/sglang/srt/layers/attention/nsa_backend.py index 397932f38..f9de5b465 100644 --- a/python/sglang/srt/layers/attention/nsa_backend.py +++ b/python/sglang/srt/layers/attention/nsa_backend.py @@ -504,7 +504,7 @@ class NativeSparseAttnBackend( page_table, repeats=self.speculative_num_draft_tokens, dim=0 ) else: - # DRAFT_EXTEND (v1): V1 worker extends by (accept_length + 1) per request + # DRAFT_EXTEND (v1): V1 worker extends by (num_accepted_drafts + 1) per request # after verification. Lengths vary per request based on how many tokens # were accepted. page_table = torch.repeat_interleave( @@ -1037,7 +1037,7 @@ class NativeSparseAttnBackend( torch.cumsum(cache_seqlens, dim=0, dtype=torch.int32) ) - extend_seq_lens = spec_info.accept_length[:bs] + extend_seq_lens = spec_info.num_accepted_tokens[:bs] extend_seq_lens_cpu = extend_seq_lens.tolist() page_indices = self.req_to_token[req_pool_indices, :max_seqlen_k] diff --git a/python/sglang/srt/layers/attention/triton_backend.py b/python/sglang/srt/layers/attention/triton_backend.py index d37f9101a..0f6c80051 100644 --- a/python/sglang/srt/layers/attention/triton_backend.py +++ b/python/sglang/srt/layers/attention/triton_backend.py @@ -395,9 +395,9 @@ class TritonAttnBackend(AttentionBackend): kv_indices = kv_indices.to(torch.int64) mask_indptr = None # TODO(FIXME): This will trigger an invalid Eagle tree when using - # `max(spec_info.accept_length_cpu)`. + # `max(spec_info.num_accepted_tokens_cpu)`. # It might have been forgotten to update somewhere. - max_extend_len = torch.max(spec_info.accept_length).item() + max_extend_len = torch.max(spec_info.num_accepted_tokens).item() num_kv_splits = None attn_logits = None attn_lse = None diff --git a/python/sglang/srt/layers/attention/trtllm_mha_backend.py b/python/sglang/srt/layers/attention/trtllm_mha_backend.py index 0fe16379e..61568a295 100644 --- a/python/sglang/srt/layers/attention/trtllm_mha_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mha_backend.py @@ -524,14 +524,14 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): metadata.cu_seqlens_k[1:].copy_( torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32) ) - accept_length = spec_info.accept_length[:bs] - if spec_info.accept_length_cpu: - metadata.max_seq_len_q = max(spec_info.accept_length_cpu) + 1 + extend_lens = spec_info.num_accepted_tokens[:bs] + if spec_info.num_accepted_tokens_cpu: + metadata.max_seq_len_q = max(spec_info.num_accepted_tokens_cpu) else: metadata.max_seq_len_q = 1 metadata.cu_seqlens_q[1:].copy_( - torch.cumsum(accept_length, dim=0, dtype=torch.int32) + torch.cumsum(extend_lens, dim=0, dtype=torch.int32) ) max_seq_pages = ( diff --git a/python/sglang/srt/layers/attention/trtllm_mla_backend.py b/python/sglang/srt/layers/attention/trtllm_mla_backend.py index 65c1cdb54..58d7ab2f2 100755 --- a/python/sglang/srt/layers/attention/trtllm_mla_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mla_backend.py @@ -1006,7 +1006,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): q = q.view(bs, -1, layer.tp_q_head_num, layer.head_dim) needs_unpad = False else: - # draft_extend: handle varying accept_lengths. If total_tokens % bs == 0, + # draft_extend: handle varying num_accepted_drafts_per_req. If total_tokens % bs == 0, # we can directly reshape q; otherwise, pad to max_seq_len_q. total_tokens = q.shape[0] tokens_per_seq = total_tokens // bs if bs > 0 else 0 diff --git a/python/sglang/srt/layers/attention/wave_backend.py b/python/sglang/srt/layers/attention/wave_backend.py index 9669a4568..829877db8 100644 --- a/python/sglang/srt/layers/attention/wave_backend.py +++ b/python/sglang/srt/layers/attention/wave_backend.py @@ -293,9 +293,9 @@ class WaveAttnBackend(AttentionBackend): ) mask_indptr = None # TODO(FIXME): This will trigger an invalid Eagle tree when using - # `max(spec_info.accept_length_cpu)`. + # `max(spec_info.num_accepted_tokens_cpu)`. # It might have been forgotten to update somewhere. - max_extend_len = torch.max(spec_info.accept_length).item() + max_extend_len = torch.max(spec_info.num_accepted_tokens).item() num_kv_splits = None attn_logits = None attn_lse = None diff --git a/python/sglang/srt/layers/utils/logprob.py b/python/sglang/srt/layers/utils/logprob.py index 0bdb0c87f..92739cd9c 100644 --- a/python/sglang/srt/layers/utils/logprob.py +++ b/python/sglang/srt/layers/utils/logprob.py @@ -338,11 +338,11 @@ def add_output_logprobs_for_spec_v1( if logits_output is None: logits_output = res.logits_output - if hasattr(res, "accept_length_per_req_cpu"): - accept_length_per_req_cpu = res.accept_length_per_req_cpu + if hasattr(res, "num_accepted_drafts_per_req_cpu"): + num_accepted_drafts_per_req_cpu = res.num_accepted_drafts_per_req_cpu else: # FIXME: Get a NgramVerifyOutput class and use that instead of this hack. - accept_length_per_req_cpu = res.accept_length.tolist() + num_accepted_drafts_per_req_cpu = res.num_accepted_drafts.tolist() top_logprobs_nums = batch.top_logprobs_nums token_ids_logprobs = batch.token_ids_logprobs @@ -363,7 +363,7 @@ def add_output_logprobs_for_spec_v1( logits_output.next_token_logits / temperatures, dim=-1 ) batch_next_token_ids = res.verified_id - num_tokens_per_req = [accept + 1 for accept in accept_length_per_req_cpu] + num_tokens_per_req = [accept + 1 for accept in num_accepted_drafts_per_req_cpu] # We should repeat top_logprobs_nums to match num_tokens_per_req. top_logprobs_nums_repeat_interleaved = [ diff --git a/python/sglang/srt/managers/scheduler_output_processor_mixin.py b/python/sglang/srt/managers/scheduler_output_processor_mixin.py index 69b268487..f35309826 100644 --- a/python/sglang/srt/managers/scheduler_output_processor_mixin.py +++ b/python/sglang/srt/managers/scheduler_output_processor_mixin.py @@ -358,7 +358,7 @@ class SchedulerOutputProcessorMixin: next_token_ids = result.next_token_ids.tolist() accept_lens = result.accept_lens.tolist() result.num_accepted_drafts = sum(accept_lens) - len(batch.reqs) - result.accept_length_per_req_cpu = [x - 1 for x in accept_lens] + result.num_accepted_drafts_per_req_cpu = [x - 1 for x in accept_lens] predict_tokens = [] stride = self.draft_worker.speculative_num_draft_tokens @@ -371,7 +371,7 @@ class SchedulerOutputProcessorMixin: ) req.spec_verify_ct += 1 - accepted_draft_tokens = result.accept_length_per_req_cpu[i] + accepted_draft_tokens = result.num_accepted_drafts_per_req_cpu[i] req.spec_accepted_drafts += accepted_draft_tokens req.update_spec_acceptance_histogram(accepted_draft_tokens) @@ -603,13 +603,13 @@ class SchedulerOutputProcessorMixin: req.mamba_last_track_seqlen = seq_len elif ( not batch.spec_algorithm.is_none() - and result.accept_length_per_req_cpu is not None + and result.num_accepted_drafts_per_req_cpu is not None ): # for spec decode, update mamba_last_track_seqlen if this iteration crosses a track interval actual_seq_len = req.seqlen - 1 if ( actual_seq_len // mamba_track_interval - != (actual_seq_len - result.accept_length_per_req_cpu[i] - 1) + != (actual_seq_len - result.num_accepted_drafts_per_req_cpu[i] - 1) // mamba_track_interval ): req.mamba_next_track_idx = ( diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 318366c44..acb3c9fcd 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -2098,7 +2098,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): if all_drafts > 0: # accept_rate: accepted_drafts / total_proposed_drafts (strict count, no bonus). meta_info["spec_accept_rate"] = accepted_drafts / all_drafts - # accept_length: accepted_drafts / verify_ct (includes bonus token). + # accept_length: completion_tokens / verify_ct (includes bonus token). meta_info["spec_accept_length"] = ( recv_obj.completion_tokens[i] / recv_obj.spec_verify_ct[i] ) diff --git a/python/sglang/srt/managers/utils.py b/python/sglang/srt/managers/utils.py index 7a5a21cda..81fbc8a05 100644 --- a/python/sglang/srt/managers/utils.py +++ b/python/sglang/srt/managers/utils.py @@ -28,7 +28,7 @@ class GenerationBatchResult: pp_hidden_states_proxy_tensors: Optional[PPProxyTensors] = None next_token_ids: Optional[Union[torch.Tensor, List[torch.Tensor]]] = None num_accepted_drafts: int = 0 # no bonus included - accept_length_per_req_cpu: Optional[List[int]] = None + num_accepted_drafts_per_req_cpu: Optional[List[int]] = None can_run_cuda_graph: bool = False # For output processing diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 73097ce7e..c927aa8e1 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -999,9 +999,12 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): spec_info.topk_index = self._pad_tensor_to_size( spec_info.topk_index, bs ) - if spec_info.accept_length is not None: - spec_info.accept_length = self._pad_tensor_to_size( - spec_info.accept_length, bs + if spec_info.num_accepted_drafts is not None: + spec_info.num_accepted_drafts = self._pad_tensor_to_size( + spec_info.num_accepted_drafts, bs + ) + spec_info.num_accepted_tokens = self._pad_tensor_to_size( + spec_info.num_accepted_tokens, bs ) spec_info.hidden_states = self._pad_tensor_to_size( spec_info.hidden_states, num_tokens @@ -1045,7 +1048,12 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): ] logits_output.hidden_states = logits_output.hidden_states[:num_tokens] elif self.forward_mode.is_draft_extend(): # draft extend - self.spec_info.accept_length = self.spec_info.accept_length[:bs] + self.spec_info.num_accepted_drafts = self.spec_info.num_accepted_drafts[ + :bs + ] + self.spec_info.num_accepted_tokens = self.spec_info.num_accepted_tokens[ + :bs + ] logits_output.next_token_logits = logits_output.next_token_logits[:bs] logits_output.hidden_states = logits_output.hidden_states[:bs] elif self.forward_mode.is_draft_extend_v2(): # draft extend_v2 diff --git a/python/sglang/srt/speculative/adaptive_runtime_state.py b/python/sglang/srt/speculative/adaptive_runtime_state.py index e4b482f41..fc469797b 100644 --- a/python/sglang/srt/speculative/adaptive_runtime_state.py +++ b/python/sglang/srt/speculative/adaptive_runtime_state.py @@ -71,7 +71,7 @@ class AdaptiveController: The worker only needs to: 1. Call ``register()`` for the initial state, then ``init_states()`` once during startup. - 2. Call ``on_verify_complete(accept_lengths)`` after each decode verify. + 2. Call ``on_verify_complete(num_accepted_drafts_per_req)`` after each decode verify. """ def __init__(self, worker: AdaptiveSpecWorker, config_path: str | None = None): @@ -107,9 +107,9 @@ class AdaptiveController: self._states[steps] = state self._activate(self.params.current_steps) - def on_verify_complete(self, accept_lengths: list[int]) -> None: + def on_verify_complete(self, num_accepted_drafts_per_req: list[int]) -> None: """Feed verify results; switch runtime state if EMA warrants it.""" - if self.params.update(accept_lengths): + if self.params.update(num_accepted_drafts_per_req): self._activate(self.params.current_steps) def _activate(self, speculative_num_steps: int) -> None: diff --git a/python/sglang/srt/speculative/adaptive_spec_params.py b/python/sglang/srt/speculative/adaptive_spec_params.py index fbce0160b..0bde5c2e4 100644 --- a/python/sglang/srt/speculative/adaptive_spec_params.py +++ b/python/sglang/srt/speculative/adaptive_spec_params.py @@ -122,16 +122,16 @@ class AdaptiveSpeculativeParams: f"steps={self.current_steps}, candidate_steps={self.candidate_steps}" ) - def update(self, accept_lengths: list[int]) -> bool: + def update(self, num_accepted_drafts_per_req: list[int]) -> bool: """Update EMA with observed accept lengths. Returns True if params changed. Args: - accept_lengths: Per-request accepted draft token counts from last verify. + num_accepted_drafts_per_req: Per-request accepted draft token counts from last verify. """ - if not accept_lengths: + if not num_accepted_drafts_per_req: return False - batch_avg = sum(accept_lengths) / len(accept_lengths) + batch_avg = sum(num_accepted_drafts_per_req) / len(num_accepted_drafts_per_req) self.ema_accept_len = ( 1 - self.ema_alpha ) * self.ema_accept_len + self.ema_alpha * batch_avg diff --git a/python/sglang/srt/speculative/dflash_info.py b/python/sglang/srt/speculative/dflash_info.py index 5c7741a39..6d233e401 100644 --- a/python/sglang/srt/speculative/dflash_info.py +++ b/python/sglang/srt/speculative/dflash_info.py @@ -322,7 +322,7 @@ class DFlashVerifyInput(SpecInput): new_verified_id: int64 tensor [bs] (the new current token per request) commit_lens: int32 tensor [bs] (how many verify-input tokens are committed) next_target_hidden: tensor [sum(commit_lens), feature_dim] - accept_length_per_req_cpu: list[int] (accepted draft tokens per request) + num_accepted_drafts_per_req_cpu: list[int] (accepted draft tokens per request) """ if batch.forward_mode.is_idle(): empty = torch.empty((0,), dtype=torch.int64, device=batch.device) @@ -387,7 +387,7 @@ class DFlashVerifyInput(SpecInput): ).cpu() max_acc = self.draft_token_num - 1 - accept_length_per_req_cpu: List[int] = [] + num_accepted_drafts_per_req_cpu: List[int] = [] commit_lens_cpu: List[int] = [] new_verified_list: List[int] = [] @@ -420,9 +420,9 @@ class DFlashVerifyInput(SpecInput): commit_lens_cpu.append(appended) new_verified_list.append(new_verified_token) - accept_length_per_req_cpu.append(max(0, appended - 1)) + num_accepted_drafts_per_req_cpu.append(max(0, appended - 1)) req.spec_verify_ct += 1 - req.spec_accepted_drafts += accept_length_per_req_cpu[-1] + req.spec_accepted_drafts += num_accepted_drafts_per_req_cpu[-1] commit_lens = torch.tensor(commit_lens_cpu, dtype=torch.int32, device=device) new_verified_id = torch.tensor( @@ -497,5 +497,5 @@ class DFlashVerifyInput(SpecInput): new_verified_id, commit_lens, next_target_hidden, - accept_length_per_req_cpu, + num_accepted_drafts_per_req_cpu, ) diff --git a/python/sglang/srt/speculative/dflash_worker.py b/python/sglang/srt/speculative/dflash_worker.py index 780b83875..9fa1174b5 100644 --- a/python/sglang/srt/speculative/dflash_worker.py +++ b/python/sglang/srt/speculative/dflash_worker.py @@ -1216,7 +1216,7 @@ class DFlashWorker: new_verified_id, commit_lens, next_target_hidden, - accept_length_per_req_cpu, + num_accepted_drafts_per_req_cpu, ) = verify_input.verify( batch=batch, logits_output=logits_output, @@ -1239,11 +1239,11 @@ class DFlashWorker: batch.spec_info = draft_input batch.forward_mode = ForwardMode.DECODE - num_accepted_drafts = sum(accept_length_per_req_cpu) + num_accepted_drafts = sum(num_accepted_drafts_per_req_cpu) if not self._logged_first_verify and self.tp_rank == 0: logger.info( - "DFLASH verify completed. accept_length_per_req=%s", - accept_length_per_req_cpu, + "DFLASH verify completed. num_accepted_drafts_per_req=%s", + num_accepted_drafts_per_req_cpu, ) self._logged_first_verify = True @@ -1251,6 +1251,6 @@ class DFlashWorker: logits_output=logits_output, next_token_ids=new_verified_id, num_accepted_drafts=num_accepted_drafts, - accept_length_per_req_cpu=accept_length_per_req_cpu, + num_accepted_drafts_per_req_cpu=num_accepted_drafts_per_req_cpu, can_run_cuda_graph=can_run_cuda_graph, ) diff --git a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py index 52e6de905..bccd0a09d 100644 --- a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py @@ -49,7 +49,8 @@ class EagleDraftExtendInputBuffers(ForwardInputBuffers): seq_lens: torch.Tensor seq_lens_cpu: torch.Tensor extend_seq_lens: torch.Tensor - accept_length: torch.Tensor + num_accepted_drafts: torch.Tensor + num_accepted_tokens: torch.Tensor next_token_logits_buffer: torch.Tensor global_num_tokens_gpu: Optional[torch.Tensor] global_num_tokens_for_logprob_gpu: Optional[torch.Tensor] @@ -166,7 +167,10 @@ class EAGLEDraftExtendCudaGraphRunner: extend_seq_lens = torch.full( (self.max_bs,), self.num_tokens_per_bs, dtype=torch.int32 ) - accept_length = torch.full( + num_accepted_drafts = torch.full( + (self.max_bs,), self.num_tokens_per_bs, dtype=torch.int32 + ) + num_accepted_tokens = torch.full( (self.max_bs,), self.num_tokens_per_bs, dtype=torch.int32 ) @@ -221,7 +225,8 @@ class EAGLEDraftExtendCudaGraphRunner: seq_lens=seq_lens, seq_lens_cpu=seq_lens_cpu, extend_seq_lens=extend_seq_lens, - accept_length=accept_length, + num_accepted_drafts=num_accepted_drafts, + num_accepted_tokens=num_accepted_tokens, next_token_logits_buffer=next_token_logits_buffer, global_num_tokens_gpu=global_num_tokens_gpu, global_num_tokens_for_logprob_gpu=global_num_tokens_for_logprob_gpu, @@ -299,7 +304,8 @@ class EAGLEDraftExtendCudaGraphRunner: positions = buffers.positions[:num_tokens] mrope_positions = buffers.mrope_positions[:, :num_tokens] hidden_states = buffers.hidden_states[:num_tokens] - accept_length = buffers.accept_length[:bs] + num_accepted_drafts = buffers.num_accepted_drafts[:bs] + num_accepted_tokens = buffers.num_accepted_tokens[:bs] next_token_logits_buffer = buffers.next_token_logits_buffer[ : bs if self.forward_mode == ForwardMode.DRAFT_EXTEND else num_tokens ] @@ -347,7 +353,8 @@ class EAGLEDraftExtendCudaGraphRunner: spec_info = EagleDraftInput( hidden_states=hidden_states, - accept_length=accept_length, + num_accepted_drafts=num_accepted_drafts, + num_accepted_tokens=num_accepted_tokens, ) spec_info.positions = None @@ -453,7 +460,8 @@ class EAGLEDraftExtendCudaGraphRunner: buffers.seq_lens.fill_(self.seq_len_fill_value) buffers.out_cache_loc.zero_() buffers.positions.zero_() - buffers.accept_length.fill_(self.num_tokens_per_bs) + buffers.num_accepted_drafts.fill_(self.num_tokens_per_bs) + buffers.num_accepted_tokens.fill_(self.num_tokens_per_bs) buffers.extend_seq_lens.fill_(self.num_tokens_per_bs) # Common inputs @@ -472,8 +480,13 @@ class EAGLEDraftExtendCudaGraphRunner: buffers.hidden_states[:num_tokens].copy_( forward_batch.spec_info.hidden_states ) - if forward_batch.spec_info.accept_length is not None: - buffers.accept_length[:raw_bs].copy_(forward_batch.spec_info.accept_length) + if forward_batch.spec_info.num_accepted_drafts is not None: + buffers.num_accepted_drafts[:raw_bs].copy_( + forward_batch.spec_info.num_accepted_drafts + ) + buffers.num_accepted_tokens[:raw_bs].copy_( + forward_batch.spec_info.num_accepted_tokens + ) buffers.req_pool_indices[:raw_bs].copy_(forward_batch.req_pool_indices) # TODO(ch-wan): support num_token_non_padded @@ -507,7 +520,12 @@ class EAGLEDraftExtendCudaGraphRunner: if bs != raw_bs: forward_batch.spec_info.positions = buffers.positions[:num_tokens] - forward_batch.spec_info.accept_length = buffers.accept_length[:bs] + forward_batch.spec_info.num_accepted_drafts = buffers.num_accepted_drafts[ + :bs + ] + forward_batch.spec_info.num_accepted_tokens = buffers.num_accepted_tokens[ + :bs + ] self.draft_extend_attn_backend.init_forward_metadata_replay_cuda_graph( bs=bs, @@ -531,7 +549,12 @@ class EAGLEDraftExtendCudaGraphRunner: # DRAFT_EXTEND_V2: all tokens calculations whether accepted or not. unpadding_bs = num_tokens elif bs != raw_bs: - forward_batch.spec_info.accept_length = buffers.accept_length[:raw_bs] + forward_batch.spec_info.num_accepted_drafts = buffers.num_accepted_drafts[ + :raw_bs + ] + forward_batch.spec_info.num_accepted_tokens = buffers.num_accepted_tokens[ + :raw_bs + ] unpadding_bs = raw_bs else: unpadding_bs = None diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py index bf91cdeb7..28b66bcc6 100644 --- a/python/sglang/srt/speculative/eagle_info.py +++ b/python/sglang/srt/speculative/eagle_info.py @@ -37,8 +37,8 @@ from sglang.srt.speculative.spec_utils import ( TREE_SPEC_KERNEL_AVAILABLE, align_evict_mask_to_page_size, assign_req_to_token_pool_func, - create_accept_length_filter, create_extend_after_decode_spec_info, + create_num_accepted_drafts_filter, filter_finished_cache_loc_kernel, generate_simulated_accept_index, get_src_tgt_cache_loc, @@ -249,7 +249,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): ), logits_output=logits_output, verified_id=torch.empty(0, dtype=torch.long, device=batch.device), - accept_length_per_req_cpu=[], + num_accepted_drafts_per_req_cpu=[], accepted_indices=torch.full( (0, self.spec_steps + 1), -1, @@ -268,7 +268,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): accept_index = torch.full( (bs, self.spec_steps + 1), -1, dtype=torch.int32, device=batch.device ) - accept_length = torch.empty((bs,), dtype=torch.int32, device=batch.device) + num_accepted_drafts = torch.empty((bs,), dtype=torch.int32, device=batch.device) if bs != len(sampling_info): sampling_info = copy.deepcopy(sampling_info) @@ -319,10 +319,10 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): if is_all_greedy or not TREE_SPEC_KERNEL_AVAILABLE: target_predict = torch.argmax(logits_output.next_token_logits, dim=-1) target_predict = target_predict.reshape(bs, self.draft_token_num) - predict, accept_index, accept_length = verify_tree_greedy_func( + predict, accept_index, num_accepted_drafts = verify_tree_greedy_func( predicts=predict, # mutable accept_index=accept_index, # mutable - accept_token_num=accept_length, # mutable + accept_token_num=num_accepted_drafts, # mutable candidates=candidates, retrieve_index=self.retrieve_index, retrieve_next_token=self.retrieve_next_token, @@ -370,7 +370,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): tree_speculative_sampling_target_only( predicts=predict, # mutable accept_index=accept_index, # mutable - accept_token_num=accept_length, # mutable + accept_token_num=num_accepted_drafts, # mutable candidates=candidates, # kwarg LHS retained as `retrive_*` to match sgl_kernel op schema. retrive_index=self.retrieve_index, @@ -397,14 +397,14 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): if tp_group.world_size > 1: tp_group.broadcast(predict, src=0) tp_group.broadcast(accept_index, src=0) - tp_group.broadcast(accept_length, src=0) + tp_group.broadcast(num_accepted_drafts, src=0) if SIMULATE_ACC_LEN > 0.0: # Do simulation accept_index = generate_simulated_accept_index( accept_index=accept_index, predict=predict, # mutable - accept_length=accept_length, # mutable + num_accepted_drafts=num_accepted_drafts, # mutable bs=bs, spec_steps=self.spec_steps, ) @@ -458,7 +458,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): req.update_spec_acceptance_histogram(accepted_draft_tokens) if has_finished: - accept_length = (accept_index != -1).sum(dim=1) - 1 + num_accepted_drafts = (accept_index != -1).sum(dim=1) - 1 # Free the KV cache for unaccepted tokens # TODO: fuse them @@ -466,10 +466,12 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): verified_id = predict[accept_index] evict_mask = torch.full_like(self.draft_token, True, dtype=torch.bool) evict_mask[accept_index] = False - accept_length_cpu = accept_length.cpu() + num_accepted_drafts_cpu = num_accepted_drafts.cpu() + num_accepted_tokens_cpu = num_accepted_drafts_cpu + 1 # FIXME: this `tolist()` fixes the numerical calculation consistency # try to unify the tensor representation and list representation - accept_length_list = accept_length_cpu.tolist() + num_accepted_drafts_list = num_accepted_drafts_cpu.tolist() + num_accepted_tokens_list = num_accepted_tokens_cpu.tolist() if page_size == 1: # TODO: boolean array index leads to a device sync. Remove it. @@ -492,7 +494,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): batch.seq_lens, batch.out_cache_loc, accept_index, - accept_length, + num_accepted_drafts, self.draft_token_num, page_size, ) @@ -509,12 +511,12 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): # to_free_slots also needs to be page-aligned without the first partial page # # split each row of out_cache_loc into two parts. - # 1. the first part goes to tgt_cache_loc. length = accept_length[i] + 1 + # 1. the first part goes to tgt_cache_loc. length = num_accepted_drafts[i] + 1 # 2. the second part goes to to_free_slots. get_target_cache_loc[(bs,)]( tgt_cache_loc, to_free_slots, - accept_length, + num_accepted_drafts, to_free_num_slots, batch.out_cache_loc, self.draft_token_num, @@ -538,20 +540,22 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): batch.req_pool_indices, batch.req_to_token_pool.req_to_token, batch.seq_lens, - batch.seq_lens + accept_length + 1, + batch.seq_lens + num_accepted_drafts + 1, batch.out_cache_loc, bs, ) else: batch.out_cache_loc = tgt_cache_loc - batch.seq_lens.add_(accept_length + 1) - batch.seq_lens_cpu.add_(accept_length_cpu + 1) + batch.seq_lens.add_(num_accepted_drafts + 1) + batch.seq_lens_cpu.add_(num_accepted_tokens_cpu) draft_input = EagleDraftInput( hidden_states=batch.spec_info.hidden_states[accept_index], verified_id=verified_id, - accept_length=accept_length, - accept_length_cpu=accept_length_list, + 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, @@ -561,7 +565,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): draft_input=draft_input, logits_output=logits_output, verified_id=verified_id, - accept_length_per_req_cpu=draft_input.accept_length_cpu, + num_accepted_drafts_per_req_cpu=draft_input.num_accepted_drafts_cpu, accepted_indices=accept_index, ) else: @@ -570,51 +574,60 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): batch.req_pool_indices, batch.req_to_token_pool.req_to_token, batch.seq_lens, - batch.seq_lens + accept_length + 1, + batch.seq_lens + num_accepted_drafts + 1, batch.out_cache_loc[accept_index], bs, ) - batch.seq_lens.add_(accept_length + 1) - batch.seq_lens_cpu.add_(accept_length_cpu + 1) + batch.seq_lens.add_(num_accepted_drafts + 1) + batch.seq_lens_cpu.add_(num_accepted_tokens_cpu) if len(unfinished_accept_index) > 0: unfinished_accept_index = torch.cat(unfinished_accept_index) unfinished_index_device = torch.tensor( unfinished_index, dtype=torch.int64, device=predict.device ) - draft_input_accept_length_cpu = [ - accept_length_list[i] for i in unfinished_index + draft_input_num_accepted_drafts_cpu = [ + num_accepted_drafts_list[i] for i in unfinished_index + ] + draft_input_num_accepted_tokens_cpu = [ + num_accepted_tokens_list[i] for i in unfinished_index ] if page_size == 1 or self.topk == 1: batch.out_cache_loc = batch.out_cache_loc[unfinished_accept_index] else: batch.out_cache_loc = torch.empty( - len(unfinished_index) + sum(draft_input_accept_length_cpu), + len(unfinished_index) + + sum(draft_input_num_accepted_drafts_cpu), dtype=torch.int64, device=predict.device, ) - accept_length_filter = create_accept_length_filter( - accept_length, + num_accepted_drafts_filter = create_num_accepted_drafts_filter( + num_accepted_drafts, unfinished_index_device, batch.seq_lens, ) - batch.seq_lens_cpu.add_(accept_length_cpu + 1) + batch.seq_lens_cpu.add_(num_accepted_tokens_cpu) filter_finished_cache_loc_kernel[(bs,)]( batch.out_cache_loc, tgt_cache_loc, - accept_length, - accept_length_filter, + num_accepted_drafts, + num_accepted_drafts_filter, next_power_of_2(bs), next_power_of_2(self.draft_token_num), ) + unfinished_num_accepted_drafts = num_accepted_drafts[ + unfinished_index_device + ] draft_input = EagleDraftInput( hidden_states=batch.spec_info.hidden_states[ unfinished_accept_index ], verified_id=predict[unfinished_accept_index], - accept_length_cpu=draft_input_accept_length_cpu, - accept_length=accept_length[unfinished_index_device], + 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[ @@ -634,7 +647,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): draft_input=draft_input, logits_output=logits_output, verified_id=verified_id, - accept_length_per_req_cpu=accept_length_list, + num_accepted_drafts_per_req_cpu=num_accepted_drafts_list, accepted_indices=accept_index, ) @@ -651,9 +664,14 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): # Inputs for extend # shape: (b,) + # `num_accepted_drafts` and `num_accepted_tokens` are kept in sync: + # `num_accepted_tokens = num_accepted_drafts + 1` (per-req, one bonus per req). + # Storing both avoids repeated `+ 1` at every consumer (attn backends, kernels). verified_id: torch.Tensor = None - accept_length: torch.Tensor = None - accept_length_cpu: List[int] = None + num_accepted_drafts: torch.Tensor = None + num_accepted_tokens: torch.Tensor = None + num_accepted_drafts_cpu: List[int] = None + num_accepted_tokens_cpu: List[int] = None # Inputs for the attention backends # shape: (b + 1,) @@ -713,8 +731,10 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): 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), - accept_length=torch.empty((0,), device=device, dtype=torch.int32), - accept_length_cpu=[], + 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( @@ -727,7 +747,7 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): return batch.input_ids = self.verified_id - batch.extend_lens = [x + 1 for x in batch.spec_info.accept_length_cpu] + 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 @@ -736,14 +756,13 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): batch.return_hidden_states = False self.capture_hidden_mode = CaptureHiddenMode.LAST - self.accept_length.add_(1) self.positions = torch.empty_like(batch.input_ids, dtype=torch.long) - self.verified_id = torch.empty_like(self.accept_length, dtype=torch.int32) + self.verified_id = torch.empty_like(self.num_accepted_tokens, dtype=torch.int32) create_extend_after_decode_spec_info[(len(batch.seq_lens),)]( batch.input_ids, batch.seq_lens, - self.accept_length, + self.num_accepted_tokens, self.positions, self.verified_id, next_power_of_2(max(speculative_num_steps + 1, len(batch.seq_lens))), @@ -757,9 +776,9 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): req_to_token: torch.Tensor, ): device = req_pool_indices.device - bs = self.accept_length.numel() + bs = self.num_accepted_drafts.numel() qo_indptr = torch.zeros((bs + 1,), dtype=torch.int32, device=device) - qo_indptr[1:] = torch.cumsum(self.accept_length, dim=0) + qo_indptr[1:] = torch.cumsum(self.num_accepted_tokens, dim=0) cum_kv_seq_len = torch.zeros((bs + 1,), dtype=torch.int32, device=device) cum_kv_seq_len[1:] = torch.cumsum(paged_kernel_lens, dim=0) @@ -843,6 +862,6 @@ class EagleVerifyOutput: # Accepted token ids including the bonus token verified_id: torch.Tensor # Accepted token length per sequence in a batch in CPU. - accept_length_per_req_cpu: List[int] + num_accepted_drafts_per_req_cpu: List[int] # Accepted indices from logits_output.next_token_logits accepted_indices: torch.Tensor diff --git a/python/sglang/srt/speculative/eagle_info_v2.py b/python/sglang/srt/speculative/eagle_info_v2.py index 75dbf6a15..4c7535ced 100644 --- a/python/sglang/srt/speculative/eagle_info_v2.py +++ b/python/sglang/srt/speculative/eagle_info_v2.py @@ -314,13 +314,13 @@ class EagleVerifyInputV2Mixin: """ if batch.forward_mode.is_idle(): predict = torch.empty(0, dtype=torch.int32, device=batch.input_ids.device) - accept_length = torch.empty( + num_accepted_drafts = torch.empty( 0, dtype=torch.int32, device=batch.input_ids.device ) accept_index = torch.empty( 0, dtype=torch.int32, device=batch.input_ids.device ) - return predict, accept_length, accept_index + return predict, num_accepted_drafts, accept_index bs = len(batch.seq_lens) sampling_info = batch.sampling_info @@ -362,16 +362,16 @@ class EagleVerifyInputV2Mixin: accept_index = torch.full( (bs, self.spec_steps + 1), -1, dtype=torch.int32, device=device ) - accept_length = torch.empty((bs,), dtype=torch.int32, device=device) + num_accepted_drafts = torch.empty((bs,), dtype=torch.int32, device=device) # Sample tokens if sampling_info.is_all_greedy or _is_npu or _is_hip: target_predict = torch.argmax(next_token_logits, dim=-1) target_predict = target_predict.reshape(bs, self.draft_token_num) - predict, accept_index, accept_length = verify_tree_greedy_func( + predict, accept_index, num_accepted_drafts = verify_tree_greedy_func( predicts=predict, # mutable accept_index=accept_index, # mutable - accept_token_num=accept_length, # mutable + accept_token_num=num_accepted_drafts, # mutable candidates=candidates, retrieve_index=self.retrieve_index, retrieve_next_token=self.retrieve_next_token, @@ -413,7 +413,7 @@ class EagleVerifyInputV2Mixin: tree_speculative_sampling_target_only( predicts=predict, # mutable accept_index=accept_index, # mutable - accept_token_num=accept_length, # mutable + accept_token_num=num_accepted_drafts, # mutable candidates=candidates, # kwarg LHS retained as `retrive_*` to match sgl_kernel op schema. retrive_index=self.retrieve_index, @@ -440,22 +440,23 @@ class EagleVerifyInputV2Mixin: if tp_group.world_size > 1: tp_group.broadcast(predict, src=0) tp_group.broadcast(accept_index, src=0) - tp_group.broadcast(accept_length, src=0) + tp_group.broadcast(num_accepted_drafts, src=0) if SIMULATE_ACC_LEN > 0: # Do simulation accept_index = generate_simulated_accept_index( accept_index=accept_index, predict=predict, # mutable - accept_length=accept_length, # mutable + num_accepted_drafts=num_accepted_drafts, # mutable simulate_acc_len=SIMULATE_ACC_LEN, bs=bs, spec_steps=self.spec_steps, ) - # Include the bonus token - accept_length.add_(1) - return predict, accept_length, accept_index + # `num_accepted_drafts` stays drafts-only inside this function; the returned + # tensor includes the trailing/bonus token via out-of-place +1 so the + # name no longer flips semantics mid-function (naming doc C2). + return predict, num_accepted_drafts + 1, accept_index @triton.jit @@ -468,9 +469,10 @@ def fill_new_verified_id( # NOTE: we cannot fuse any in-place operations of `accept_lens` inside this kernel # because this kernel reads accept_lens pid = tl.program_id(axis=0) - accept_length = tl.load(accept_lens + pid) + # `accept_lens` includes the bonus token; the last accepted slot is at -1. + accept_len = tl.load(accept_lens + pid) - verified_id_idx = num_draft_tokens * pid + accept_length - 1 + verified_id_idx = num_draft_tokens * pid + accept_len - 1 verified_id_data = tl.load(verified_id + verified_id_idx) tl.store(new_verified_id + pid, verified_id_data) diff --git a/python/sglang/srt/speculative/eagle_worker.py b/python/sglang/srt/speculative/eagle_worker.py index c63ec9a72..79ea70ed0 100644 --- a/python/sglang/srt/speculative/eagle_worker.py +++ b/python/sglang/srt/speculative/eagle_worker.py @@ -483,7 +483,7 @@ class EAGLEWorker(TpModelWorker): if get_global_tracing_enabled(): for idx, req in enumerate(batch.reqs): - accepted = verify_output.accept_length_per_req_cpu[idx] + accepted = verify_output.num_accepted_drafts_per_req_cpu[idx] req.time_stats.set_spec_verify_end_time(accepted_tokens=accepted) set_time_batch( @@ -508,13 +508,15 @@ class EAGLEWorker(TpModelWorker): controller = getattr(self, "adaptive_controller", None) if controller is not None: - controller.on_verify_complete(verify_output.accept_length_per_req_cpu) + controller.on_verify_complete( + verify_output.num_accepted_drafts_per_req_cpu + ) return GenerationBatchResult( logits_output=logits_output, next_token_ids=verify_output.verified_id, - num_accepted_drafts=sum(verify_output.accept_length_per_req_cpu), - accept_length_per_req_cpu=verify_output.accept_length_per_req_cpu, + num_accepted_drafts=sum(verify_output.num_accepted_drafts_per_req_cpu), + num_accepted_drafts_per_req_cpu=verify_output.num_accepted_drafts_per_req_cpu, can_run_cuda_graph=can_run_cuda_graph, ) @@ -987,7 +989,7 @@ class EAGLEWorker(TpModelWorker): accepted_length = ( torch.tensor( - res.accept_length_per_req_cpu, + res.num_accepted_drafts_per_req_cpu, device=logits_output.hidden_states.device, dtype=torch.int64, ) @@ -1101,7 +1103,8 @@ class EAGLEWorker(TpModelWorker): seq_lens_backup = batch.seq_lens.clone() seq_lens_cpu_backup = batch.seq_lens_cpu.clone() req_pool_indices_backup = batch.req_pool_indices - accept_length_backup = batch.spec_info.accept_length.clone() + num_accepted_drafts_backup = batch.spec_info.num_accepted_drafts.clone() + num_accepted_tokens_backup = batch.spec_info.num_accepted_tokens.clone() return_logprob_backup = batch.return_logprob input_is_idle = batch.forward_mode.is_idle() @@ -1187,7 +1190,8 @@ class EAGLEWorker(TpModelWorker): batch.seq_lens = seq_lens_backup batch.seq_lens_cpu = seq_lens_cpu_backup batch.req_pool_indices = req_pool_indices_backup - batch.spec_info.accept_length = accept_length_backup + batch.spec_info.num_accepted_drafts = num_accepted_drafts_backup + batch.spec_info.num_accepted_tokens = num_accepted_tokens_backup batch.return_logprob = return_logprob_backup def capture_for_decode( diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 4cc46e2c4..6f477f4cf 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -579,8 +579,11 @@ class EagleDraftWorker(BaseDraftWorker): self.plan_stream ) - if forward_batch.spec_info.accept_length is None: - forward_batch.spec_info.accept_length = batch_result.accept_lens + if forward_batch.spec_info.num_accepted_drafts is None: + # `batch_result.accept_lens` already includes the bonus token, so use it + # directly for `num_accepted_tokens` and subtract 1 for `num_accepted_drafts`. + forward_batch.spec_info.num_accepted_drafts = batch_result.accept_lens - 1 + forward_batch.spec_info.num_accepted_tokens = batch_result.accept_lens # Run draft extend batch in the main compute stream can_cuda_graph = ( @@ -838,10 +841,10 @@ class EAGLEWorkerV2(BaseSpecWorker): maybe_detect_nan(logits_output.next_token_logits, "verify: target model logits") ( predict, - accept_length, + accept_lens, accept_index, ) = verify_input.sample(batch, logits_output, vocab_mask) - new_seq_lens = batch.seq_lens + accept_length + new_seq_lens = batch.seq_lens + accept_lens # Update mamba state for hybrid GDN models after verification if ( @@ -849,7 +852,7 @@ class EAGLEWorkerV2(BaseSpecWorker): or self.target_worker.model_runner.mamba2_config is not None ): self._mamba_verify_update( - batch, verify_input, accept_length, accept_index, bs + batch, verify_input, accept_lens, accept_index, bs ) verify_done = torch.get_device_module(self.device).Event() @@ -857,10 +860,10 @@ class EAGLEWorkerV2(BaseSpecWorker): if not batch.forward_mode.is_idle(): all_verified_id = predict[accept_index] - verified_id = torch.empty_like(accept_length, dtype=torch.int32) + verified_id = torch.empty_like(accept_lens, dtype=torch.int32) fill_new_verified_id[(bs,)]( all_verified_id, - accept_length, + accept_lens, verified_id, self.speculative_num_draft_tokens, ) @@ -884,7 +887,7 @@ class EAGLEWorkerV2(BaseSpecWorker): next_token_ids=predict, can_run_cuda_graph=can_run_cuda_graph, next_draft_input=next_draft_input, - accept_lens=accept_length, + accept_lens=accept_lens, routed_experts_output=forward_batch_output.routed_experts_output, ) @@ -892,14 +895,13 @@ class EAGLEWorkerV2(BaseSpecWorker): self, batch: ModelWorkerBatch, verify_input: EagleVerifyInput, - accept_length: torch.Tensor, + accept_lens: torch.Tensor, accept_index: torch.Tensor, bs: int, ): """Update mamba state for hybrid GDN models after verification.""" - # Calculate accepted_steps for mamba state update - # Include the bonus token (+1) - accepted_length_with_bonus = accept_length + # `accept_lens` already includes the bonus token (drafts + 1 per req). + accepted_length_with_bonus = accept_lens if not batch.forward_mode.is_idle() and accept_index.numel() > 0: if verify_input.topk != 1: raise ValueError("Spec v2 currently only supports topk = 1.") @@ -956,7 +958,7 @@ class EAGLEWorkerV2(BaseSpecWorker): self, batch: ModelWorkerBatch, accept_index: torch.Tensor, - accept_length: torch.Tensor, + num_accepted_drafts: torch.Tensor, ): """ Move accepted tokens to the target KV cache. @@ -964,7 +966,7 @@ class EAGLEWorkerV2(BaseSpecWorker): Args: batch: The batch to run. accept_index: The index of the accepted tokens. - accept_length: The length of the accepted tokens. + num_accepted_drafts: The length of the accepted tokens. """ bs = len(batch.seq_lens) size = bs * self.speculative_num_draft_tokens @@ -981,7 +983,7 @@ class EAGLEWorkerV2(BaseSpecWorker): batch.req_pool_indices, self.req_to_token_pool.req_to_token, batch.seq_lens, - batch.seq_lens + accept_length, + batch.seq_lens + num_accepted_drafts, tgt_cache_loc, self.req_to_token_pool.req_to_token.shape[1], next_power_of_2(bs), diff --git a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py index f036ed371..295c1009b 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py @@ -72,7 +72,8 @@ class MultiLayerEagleDraftExtendInputBuffers(ForwardInputBuffers): seq_lens: torch.Tensor seq_lens_cpu: torch.Tensor req_pool_indices: torch.Tensor - accept_length: torch.Tensor + num_accepted_drafts: torch.Tensor + num_accepted_tokens: torch.Tensor # Per-step buffers extend_seq_lens: torch.Tensor extend_start_loc: torch.Tensor @@ -159,7 +160,8 @@ class MultiLayerEagleDraftExtendCudaGraphRunner: # shared states seq_lens = cuda_graph_buffers["seq_lens"] req_pool_indices = cuda_graph_buffers["req_pool_indices"] - accept_length = cuda_graph_buffers["accept_length"] + num_accepted_drafts = cuda_graph_buffers["num_accepted_drafts"] + num_accepted_tokens = cuda_graph_buffers["num_accepted_tokens"] extend_seq_lens = torch.full( (self.max_bs,), @@ -229,7 +231,8 @@ class MultiLayerEagleDraftExtendCudaGraphRunner: seq_lens=seq_lens, seq_lens_cpu=seq_lens_cpu, req_pool_indices=req_pool_indices, - accept_length=accept_length, + num_accepted_drafts=num_accepted_drafts, + num_accepted_tokens=num_accepted_tokens, extend_seq_lens=extend_seq_lens, extend_start_loc=extend_start_loc, mrope_positions=mrope_positions, @@ -301,7 +304,8 @@ class MultiLayerEagleDraftExtendCudaGraphRunner: extend_seq_lens = buffers.extend_seq_lens[:bs] extend_seq_lens_cpu = self.extend_seq_lens_cpu[:bs] extend_start_loc = buffers.extend_start_loc[:bs] - accept_length = buffers.accept_length[:bs] + num_accepted_drafts = buffers.num_accepted_drafts[:bs] + num_accepted_tokens = buffers.num_accepted_tokens[:bs] out_cache_loc = buffers.out_cache_loc[:num_tokens] positions = buffers.positions[:num_tokens] mrope_positions = buffers.mrope_positions[:, :num_tokens] @@ -347,7 +351,8 @@ class MultiLayerEagleDraftExtendCudaGraphRunner: spec_info = EagleDraftInput( hidden_states=hidden_states, - accept_length=accept_length, + num_accepted_drafts=num_accepted_drafts, + num_accepted_tokens=num_accepted_tokens, ) spec_info.positions = None @@ -439,11 +444,12 @@ class MultiLayerEagleDraftExtendCudaGraphRunner: ): buffers.hidden_states[:num_tokens].copy_(ret.hidden_states[:num_tokens]) + # num_accepted_drafts is drafts-only; the last accepted draft sits at index + # `num_accepted_drafts` within the (current_token + drafts) slot range. select_index = ( torch.arange(bs, device=self.model_runner.device) * (self.speculative_num_draft_tokens + self.step) - + buffers.accept_length[:bs] - - 1 + + buffers.num_accepted_drafts[:bs] + self.step ) @@ -452,9 +458,11 @@ class MultiLayerEagleDraftExtendCudaGraphRunner: if self.next_cuda_graph_runner is not None: next_buffers = self.next_cuda_graph_runner.buffers + # rejected drafts = proposed drafts - accepted drafts. + # speculative_num_draft_tokens includes the current-token slot, so -1. padding_lens = ( - self.speculative_num_draft_tokens - buffers.accept_length[:bs] - ) + self.speculative_num_draft_tokens - 1 + ) - buffers.num_accepted_drafts[:bs] assign_new_state_triton( ret.topk_index, buffers.input_ids, @@ -515,8 +523,13 @@ class MultiLayerEagleDraftExtendCudaGraphRunner: buffers.hidden_states[:num_tokens].copy_( forward_batch.spec_info.hidden_states ) - if forward_batch.spec_info.accept_length is not None: - buffers.accept_length[:raw_bs].copy_(forward_batch.spec_info.accept_length) + if forward_batch.spec_info.num_accepted_drafts is not None: + buffers.num_accepted_drafts[:raw_bs].copy_( + forward_batch.spec_info.num_accepted_drafts + ) + buffers.num_accepted_tokens[:raw_bs].copy_( + forward_batch.spec_info.num_accepted_tokens + ) buffers.req_pool_indices[:raw_bs].copy_(forward_batch.req_pool_indices) if forward_batch.seq_lens_cpu is not None: @@ -553,7 +566,8 @@ class MultiLayerEagleDraftExtendCudaGraphRunner: buffers.global_num_tokens_for_logprob_gpu.fill_(bs * self.num_tokens_per_bs) forward_batch.spec_info.hidden_states = buffers.hidden_states[:num_tokens] - forward_batch.spec_info.accept_length = buffers.accept_length[:bs] + forward_batch.spec_info.num_accepted_drafts = buffers.num_accepted_drafts[:bs] + forward_batch.spec_info.num_accepted_tokens = buffers.num_accepted_tokens[:bs] forward_batch.spec_info.num_tokens_per_req = self.num_tokens_per_bs forward_batch.spec_info.num_tokens_for_logprob_per_req = 1 forward_batch.spec_info.positions = buffers.positions[:num_tokens] @@ -583,7 +597,12 @@ class MultiLayerEagleDraftExtendCudaGraphRunner: # DRAFT_EXTEND_V2: all tokens calculations whether accepted or not. unpadding_bs = num_tokens elif bs != raw_bs: - forward_batch.spec_info.accept_length = buffers.accept_length[:raw_bs] + forward_batch.spec_info.num_accepted_drafts = buffers.num_accepted_drafts[ + :raw_bs + ] + forward_batch.spec_info.num_accepted_tokens = buffers.num_accepted_tokens[ + :raw_bs + ] unpadding_bs = raw_bs else: unpadding_bs = None @@ -671,7 +690,10 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner: self.cuda_graph_buffers["req_pool_indices"] = torch.zeros( (self.max_bs,), dtype=torch.int64 ) - self.cuda_graph_buffers["accept_length"] = torch.full( + self.cuda_graph_buffers["num_accepted_drafts"] = torch.full( + (self.max_bs,), 1, dtype=torch.int32 + ) + self.cuda_graph_buffers["num_accepted_tokens"] = torch.full( (self.max_bs,), 1, dtype=torch.int32 ) @@ -704,7 +726,12 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner: self.cuda_graph_buffers["out_cache_loc"].zero_() self.cuda_graph_buffers["swa_out_cache_loc"].zero_() self.cuda_graph_buffers["positions"].zero_() - self.cuda_graph_buffers["accept_length"][: forward_batch.batch_size].copy_( + # `batch_result.accept_lens` is drafts + bonus. + bs = forward_batch.batch_size + self.cuda_graph_buffers["num_accepted_drafts"][:bs].copy_( + batch_result.accept_lens - 1 + ) + self.cuda_graph_buffers["num_accepted_tokens"][:bs].copy_( batch_result.accept_lens ) diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker.py b/python/sglang/srt/speculative/multi_layer_eagle_worker.py index 42bdf3e1c..2cb78d12f 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker.py @@ -291,7 +291,7 @@ class MultiLayerEagleWorker(TpModelWorker): return GenerationBatchResult( logits_output=logits_output, next_token_ids=verify_output.verified_id, - num_accepted_drafts=sum(verify_output.accept_length_per_req_cpu), + num_accepted_drafts=sum(verify_output.num_accepted_drafts_per_req_cpu), can_run_cuda_graph=can_run_cuda_graph, ) @@ -544,7 +544,7 @@ class MultiLayerEagleWorker(TpModelWorker): if self.target_worker.model_runner.hybrid_gdn_config is not None: accepted_length = ( torch.tensor( - res.accept_length_per_req_cpu, + res.num_accepted_drafts_per_req_cpu, device=logits_output.hidden_states.device, dtype=torch.int64, ) @@ -658,7 +658,8 @@ class MultiLayerEagleWorker(TpModelWorker): seq_lens_backup = batch.seq_lens.clone() seq_lens_cpu_backup = batch.seq_lens_cpu.clone() req_pool_indices_backup = batch.req_pool_indices - accept_length_backup = batch.spec_info.accept_length + num_accepted_drafts_backup = batch.spec_info.num_accepted_drafts + num_accepted_tokens_backup = batch.spec_info.num_accepted_tokens return_logprob_backup = batch.return_logprob input_is_idle = batch.forward_mode.is_idle() @@ -755,5 +756,6 @@ class MultiLayerEagleWorker(TpModelWorker): batch.seq_lens = seq_lens_backup batch.seq_lens_cpu = seq_lens_cpu_backup batch.req_pool_indices = req_pool_indices_backup - batch.spec_info.accept_length = accept_length_backup + batch.spec_info.num_accepted_drafts = num_accepted_drafts_backup + batch.spec_info.num_accepted_tokens = num_accepted_tokens_backup batch.return_logprob = return_logprob_backup 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 da2ba93fc..9a5937beb 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -753,19 +753,19 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker): maybe_detect_nan(logits_output.next_token_logits, "verify: target model logits") ( predict, - accept_length, + accept_lens, accept_index, ) = verify_input.sample(batch, logits_output) - new_seq_lens = batch.seq_lens + accept_length + new_seq_lens = batch.seq_lens + accept_lens verify_done = torch.get_device_module(self.device).Event() verify_done.record() if not batch.forward_mode.is_idle(): all_verified_id = predict[accept_index] - verified_id = torch.empty_like(accept_length, dtype=torch.int32) + verified_id = torch.empty_like(accept_lens, dtype=torch.int32) fill_new_verified_id[(bs,)]( all_verified_id, - accept_length, + accept_lens, verified_id, self.speculative_num_draft_tokens, ) @@ -788,7 +788,7 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker): next_token_ids=predict, can_run_cuda_graph=can_run_cuda_graph, next_draft_input=next_draft_input, - accept_lens=accept_length, + accept_lens=accept_lens, routed_experts_output=forward_batch_output.routed_experts_output, ) diff --git a/python/sglang/srt/speculative/ngram_info.py b/python/sglang/srt/speculative/ngram_info.py index cb1ed1810..efb389e5f 100644 --- a/python/sglang/srt/speculative/ngram_info.py +++ b/python/sglang/srt/speculative/ngram_info.py @@ -196,7 +196,7 @@ class NgramVerifyInput(SpecInput): req.update_spec_acceptance_histogram(accepted_draft_tokens) if has_finished: - self.accept_length = (self.accepted_indices != -1).sum(dim=1) - 1 + self.num_accepted_drafts = (self.accepted_indices != -1).sum(dim=1) - 1 self.accepted_indices = self.accepted_indices[self.accepted_indices != -1] logits_output.next_token_logits = logits_output.next_token_logits[ @@ -209,7 +209,10 @@ class NgramVerifyInput(SpecInput): self.verified_id = self.predict[self.accepted_indices] def _free_cache( - self, batch: ScheduleBatch, page_size: int, accept_length_cpu: torch.Tensor + self, + batch: ScheduleBatch, + page_size: int, + num_accepted_drafts_cpu: torch.Tensor, ): bs = batch.batch_size() # Free the KV cache for unaccepted tokens @@ -226,7 +229,7 @@ class NgramVerifyInput(SpecInput): batch.seq_lens, batch.out_cache_loc, self.accepted_indices, - self.accept_length, + self.num_accepted_drafts, self.draft_token_num, page_size, ) @@ -243,12 +246,12 @@ class NgramVerifyInput(SpecInput): # to_free_slots also needs to be page-aligned without the first partial page # # split each row of out_cache_loc into two parts. - # 1. the first part goes to tgt_cache_loc. length = accept_length[i] + 1 + # 1. the first part goes to tgt_cache_loc. length = num_accepted_drafts[i] + 1 # 2. the second part goes to to_free_slots. get_target_cache_loc[(bs,)]( tgt_cache_loc, to_free_slots, - self.accept_length, + self.num_accepted_drafts, to_free_num_slots, batch.out_cache_loc, self.draft_token_num, @@ -265,16 +268,16 @@ class NgramVerifyInput(SpecInput): ) batch.out_cache_loc = tgt_cache_loc - accept_length_list = accept_length_cpu.tolist() + num_accepted_drafts_list = num_accepted_drafts_cpu.tolist() for i, req in enumerate(batch.reqs): - req.kv_committed_len += accept_length_list[i] + 1 + req.kv_committed_len += num_accepted_drafts_list[i] + 1 req.kv_allocated_len = req.kv_committed_len assign_req_to_token_pool[(bs,)]( batch.req_pool_indices, batch.req_to_token_pool.req_to_token, batch.seq_lens, - batch.seq_lens + self.accept_length + 1, + batch.seq_lens + self.num_accepted_tokens, batch.out_cache_loc, batch.req_to_token_pool.req_to_token.shape[1], triton.next_power_of_2(bs), @@ -296,12 +299,14 @@ class NgramVerifyInput(SpecInput): self.accepted_indices = torch.full( (bs, self.draft_token_num), -1, dtype=torch.int32, device=self.device ) - self.accept_length = torch.empty((bs,), dtype=torch.int32, device=self.device) + self.num_accepted_drafts = torch.empty( + (bs,), dtype=torch.int32, device=self.device + ) verify_tree_greedy( predicts=self.predict, # mutable accept_index=self.accepted_indices, # mutable - accept_token_num=self.accept_length, # mutable + accept_token_num=self.num_accepted_drafts, # mutable candidates=candidates, # kwarg LHS retained as `retrive_*` to match sgl_kernel op schema. retrive_index=self.retrieve_index, @@ -324,7 +329,9 @@ class NgramVerifyInput(SpecInput): self.accepted_indices = torch.full( (bs, self.draft_token_num), -1, dtype=torch.int32, device=self.device ) - self.accept_length = torch.empty((bs,), dtype=torch.int32, device=self.device) + self.num_accepted_drafts = torch.empty( + (bs,), dtype=torch.int32, device=self.device + ) # apply temperature and get target probs expanded_temperature = torch.repeat_interleave( sampling_info.temperatures, self.draft_token_num, dim=0 @@ -364,7 +371,7 @@ class NgramVerifyInput(SpecInput): tree_speculative_sampling_target_only( predicts=self.predict, # mutable accept_index=self.accepted_indices, # mutable - accept_token_num=self.accept_length, # mutable + accept_token_num=self.num_accepted_drafts, # mutable candidates=candidates.to(torch.int64), # kwarg LHS retained as `retrive_*` to match sgl_kernel op schema. retrive_index=self.retrieve_index.to(torch.int64), @@ -445,13 +452,18 @@ class NgramVerifyInput(SpecInput): self._fill_requests(batch, logits_output) - accept_length_cpu = self.accept_length.cpu() - num_accepted_drafts = accept_length_cpu.sum().item() + # Sync the bonus-included view after the kernel + `_fill_requests` + # finalize `num_accepted_drafts`. + self.num_accepted_tokens = self.num_accepted_drafts + 1 - self._free_cache(batch, page_size, accept_length_cpu) + num_accepted_drafts_cpu = self.num_accepted_drafts.cpu() + num_accepted_tokens_cpu = num_accepted_drafts_cpu + 1 + num_accepted_drafts = num_accepted_drafts_cpu.sum().item() - batch.seq_lens.add_(self.accept_length + 1) - batch.seq_lens_cpu.add_(accept_length_cpu + 1) + self._free_cache(batch, page_size, num_accepted_drafts_cpu) + + batch.seq_lens.add_(self.num_accepted_tokens) + batch.seq_lens_cpu.add_(num_accepted_tokens_cpu) return logits_output, self.verified_id, num_accepted_drafts diff --git a/python/sglang/srt/speculative/ngram_worker.py b/python/sglang/srt/speculative/ngram_worker.py index 453556ca4..b4b65a0a5 100644 --- a/python/sglang/srt/speculative/ngram_worker.py +++ b/python/sglang/srt/speculative/ngram_worker.py @@ -262,7 +262,7 @@ class NGRAMWorker: spec_info = model_worker_batch.spec_info num_accepted_drafts = 0 accept_lens = None - accept_length_per_req_cpu = None + num_accepted_drafts_per_req_cpu = None if model_worker_batch.forward_mode.is_target_verify(): if batch.has_grammar: @@ -306,19 +306,21 @@ class NGRAMWorker: logits_output, next_token_ids, num_accepted_drafts = verify_input.verify( batch, logits_output, self.page_size, vocab_mask ) - accept_length_per_req_cpu = verify_input.accept_length.cpu().tolist() + num_accepted_drafts_per_req_cpu = ( + verify_input.num_accepted_drafts.cpu().tolist() + ) if get_global_tracing_enabled(): for idx, req in enumerate(batch.reqs): accepted = ( - verify_input.accept_length[idx].item() - if verify_input.accept_length is not None + verify_input.num_accepted_drafts[idx].item() + if verify_input.num_accepted_drafts is not None else 0 ) req.time_stats.set_spec_verify_end_time(accepted_tokens=accepted) # Store accept_lens for per-request metrics - accept_lens = verify_input.accept_length + accept_lens = verify_input.num_accepted_drafts if batch.return_logprob: add_output_logprobs_for_spec_v1(batch, verify_input, logits_output) self._update_ngram_corpus(batch) @@ -348,7 +350,7 @@ class NGRAMWorker: logits_output=logits_output, next_token_ids=next_token_ids, num_accepted_drafts=num_accepted_drafts, - accept_length_per_req_cpu=accept_length_per_req_cpu, + num_accepted_drafts_per_req_cpu=num_accepted_drafts_per_req_cpu, can_run_cuda_graph=can_run_cuda_graph, accept_lens=accept_lens, ) diff --git a/python/sglang/srt/speculative/spec_utils.py b/python/sglang/srt/speculative/spec_utils.py index 3199ed4ce..2833f1970 100644 --- a/python/sglang/srt/speculative/spec_utils.py +++ b/python/sglang/srt/speculative/spec_utils.py @@ -69,16 +69,17 @@ def create_extend_after_decode_spec_info( pid = tl.program_id(axis=0) offsets = tl.arange(0, bs_upper) seq_length = tl.load(seq_lens + pid) - accept_length = tl.load(accept_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_length - tl.store(positions_ptr + offsets, seq_length - accept_length + offsets, mask) + mask = offsets < accept_len + tl.store(positions_ptr + offsets, seq_length - accept_len + offsets, mask) - accept_len_cumsum += accept_length - 1 + accept_len_cumsum += accept_len - 1 verified_id_data = tl.load(verified_id + accept_len_cumsum) tl.store(new_verified_id + pid, verified_id_data) @@ -357,7 +358,7 @@ def align_evict_mask_to_page_size( def get_target_cache_loc( tgt_cache_loc, to_free_slots, - accept_length, + num_accepted_drafts, to_free_num_slots, out_cache_loc, num_verify_tokens: tl.constexpr, @@ -369,9 +370,9 @@ def get_target_cache_loc( bs_offset = tl.arange(0, bs_upper) # write the first part to tgt_cache_loc - accept_len_all = tl.load(accept_length + bs_offset, mask=bs_offset < bid) + accept_len_all = tl.load(num_accepted_drafts + bs_offset, mask=bs_offset < bid) tgt_cache_loc_start = tl.sum(accept_len_all) + bid - copy_len = tl.load(accept_length + bid) + 1 + copy_len = tl.load(num_accepted_drafts + bid) + 1 out_cache_loc_row = tl.load( out_cache_loc + bid * num_verify_tokens + offset, mask=offset < copy_len ) @@ -404,7 +405,7 @@ def get_src_tgt_cache_loc( seq_lens: torch.Tensor, out_cache_loc: torch.Tensor, accept_index: torch.Tensor, - accept_length: torch.Tensor, + num_accepted_drafts: torch.Tensor, draft_token_num: int, page_size: int, ): @@ -412,7 +413,7 @@ def get_src_tgt_cache_loc( tgt_cache_loc = torch.empty_like(src_cache_loc) extended_len = seq_lens + draft_token_num keep_len = torch.minimum( - (seq_lens + accept_length + 1 + page_size - 1) // page_size * page_size, + (seq_lens + num_accepted_drafts + 1 + page_size - 1) // page_size * page_size, extended_len, ) to_free_num_slots = extended_len - keep_len @@ -423,23 +424,25 @@ def get_src_tgt_cache_loc( def filter_finished_cache_loc_kernel( out_cache_loc, tgt_cache_loc, - accept_length, - accept_length_filter, + num_accepted_drafts, + num_accepted_drafts_filter, bs_upper: tl.constexpr, num_verify_tokens_upper: tl.constexpr, ): bid = tl.program_id(0) bs_offset = tl.arange(0, bs_upper) - accept_length_all = tl.load(accept_length + bs_offset, mask=bs_offset < bid) - old_start = tl.sum(accept_length_all) + bid - - accept_length_filter_all = tl.load( - accept_length_filter + bs_offset, mask=bs_offset < bid + num_accepted_drafts_all = tl.load( + num_accepted_drafts + bs_offset, mask=bs_offset < bid ) - new_start = tl.sum(accept_length_filter_all) + old_start = tl.sum(num_accepted_drafts_all) + bid - copy_len = tl.load(accept_length_filter + bid) + num_accepted_drafts_filter_all = tl.load( + num_accepted_drafts_filter + bs_offset, mask=bs_offset < bid + ) + new_start = tl.sum(num_accepted_drafts_filter_all) + + copy_len = tl.load(num_accepted_drafts_filter + bid) copy_offset = tl.arange(0, num_verify_tokens_upper) value = tl.load( tgt_cache_loc + old_start + copy_offset, mask=copy_offset < copy_len @@ -450,17 +453,17 @@ def filter_finished_cache_loc_kernel( @torch.compile(dynamic=True, disable=_is_npu) -def create_accept_length_filter( - accept_length: torch.Tensor, +def create_num_accepted_drafts_filter( + num_accepted_drafts: torch.Tensor, unfinished_index_device: torch.Tensor, seq_lens: torch.Tensor, ): - accept_length_filter = torch.zeros_like(accept_length) - accept_length_filter[unfinished_index_device] = ( - accept_length[unfinished_index_device] + 1 + num_accepted_drafts_filter = torch.zeros_like(num_accepted_drafts) + num_accepted_drafts_filter[unfinished_index_device] = ( + num_accepted_drafts[unfinished_index_device] + 1 ) - seq_lens.add_(accept_length + 1) - return accept_length_filter + seq_lens.add_(num_accepted_drafts + 1) + return num_accepted_drafts_filter @torch.compile(dynamic=True, disable=_is_npu) @@ -517,7 +520,7 @@ def select_top_k_tokens( def generate_simulated_accept_index( accept_index, predict, - accept_length, + num_accepted_drafts, bs, spec_steps, simulate_acc_len: float = SIMULATE_ACC_LEN, @@ -562,7 +565,7 @@ def generate_simulated_accept_index( sim_accept_index[:, :simulate_acc_len] = accept_indx_first_col + torch.arange( simulate_acc_len, device=accept_index.device ) - accept_length.fill_(simulate_acc_len - 1) + num_accepted_drafts.fill_(simulate_acc_len - 1) predict.fill_(100) # some legit token id return sim_accept_index diff --git a/python/sglang/test/attention/test_trtllm_mla_backend.py b/python/sglang/test/attention/test_trtllm_mla_backend.py index 79f196520..69784b84f 100755 --- a/python/sglang/test/attention/test_trtllm_mla_backend.py +++ b/python/sglang/test/attention/test_trtllm_mla_backend.py @@ -1308,7 +1308,7 @@ class TestTRTLLMMLA(CustomTestCase): device = torch.device("cuda") # Create accept lengths (varying lengths for each batch) - accept_lengths = torch.randint( + num_accepted_drafts_per_req = torch.randint( 1, token_per_batch + 1, (batch_size,), device=device, dtype=torch.int32 ) @@ -1316,7 +1316,7 @@ class TestTRTLLMMLA(CustomTestCase): cum_accept_lengths = torch.zeros( batch_size + 1, device=device, dtype=torch.int32 ) - cum_accept_lengths[1:] = torch.cumsum(accept_lengths, dim=0) + cum_accept_lengths[1:] = torch.cumsum(num_accepted_drafts_per_req, dim=0) # Create raw output tensor (batch format) raw_out = torch.randn( @@ -1334,7 +1334,7 @@ class TestTRTLLMMLA(CustomTestCase): total_tokens, tp_q_head_num, v_head_dim, device=device, dtype=dtype ) - return raw_out, output, accept_lengths, cum_accept_lengths + return raw_out, output, num_accepted_drafts_per_req, cum_accept_lengths # Test 1: pad_draft_extend_query_kernel basic functionality with self.subTest(test="pad_kernel_basic"): @@ -1395,7 +1395,7 @@ class TestTRTLLMMLA(CustomTestCase): tp_q_head_num = 16 v_head_dim = 64 - raw_out, output, accept_lengths, cum_accept_lengths = ( + raw_out, output, num_accepted_drafts_per_req, cum_accept_lengths = ( _create_test_output_data( self, batch_size, token_per_batch, tp_q_head_num, v_head_dim ) @@ -1408,7 +1408,7 @@ class TestTRTLLMMLA(CustomTestCase): unpad_draft_extend_output_kernel[grid]( raw_out_ptr=raw_out, output_ptr=output, - accept_length_ptr=accept_lengths, + accept_length_ptr=num_accepted_drafts_per_req, cumsum_ptr=cum_accept_lengths, batch_size=batch_size, token_per_batch=token_per_batch, @@ -1419,7 +1419,7 @@ class TestTRTLLMMLA(CustomTestCase): # Verify the unpadding worked correctly for i in range(batch_size): - accept_len = accept_lengths[i].item() + accept_len = num_accepted_drafts_per_req[i].item() output_start = cum_accept_lengths[i].item() # Check that valid positions are copied correctly diff --git a/python/sglang/test/kits/spec_decoding_kit.py b/python/sglang/test/kits/spec_decoding_kit.py index 4262743de..7ad509bb6 100644 --- a/python/sglang/test/kits/spec_decoding_kit.py +++ b/python/sglang/test/kits/spec_decoding_kit.py @@ -4,7 +4,7 @@ from sglang.test.test_utils import is_in_ci, write_github_step_summary class SpecDecodingMixin: bs_1_speed_thres: float - accept_length_thres: float + num_accepted_drafts_thres: float def test_bs_1_speed(self): args = BenchArgs(port=int(self.base_url.split(":")[-1]), max_new_tokens=2048) @@ -19,5 +19,5 @@ class SpecDecodingMixin: f"{speed=:.2f} token/s\n" ) - self.assertGreater(acc_length, self.accept_length_thres) + self.assertGreater(acc_length, self.num_accepted_drafts_thres) self.assertGreater(speed, self.bs_1_speed_thres)