diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index fc5fb4c6e..07ee75093 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -2125,9 +2125,13 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): recv_obj.completion_tokens[i] / recv_obj.spec_verify_ct[i] ) + meta_info["spec_num_correct_drafts"] = num_correct_drafts + meta_info["spec_num_proposed_drafts"] = num_proposed_drafts + meta_info["spec_verify_ct"] = recv_obj.spec_verify_ct[i] + + # FIXME: backward-compat aliases, remove in next release. meta_info["spec_accepted_drafts"] = num_correct_drafts meta_info["spec_proposed_drafts"] = num_proposed_drafts - meta_info["spec_verify_ct"] = recv_obj.spec_verify_ct[i] # Acceptance histogram: tracks how many decoding steps accepted a certain number of draft tokens. if ( @@ -2135,6 +2139,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): and len(recv_obj.spec_correct_drafts_histogram) > i and recv_obj.spec_correct_drafts_histogram[i] ): + meta_info["spec_correct_drafts_histogram"] = ( + recv_obj.spec_correct_drafts_histogram[i] + ) + # FIXME: backward-compat alias, remove in next release. meta_info["spec_accept_histogram"] = ( recv_obj.spec_correct_drafts_histogram[i] ) diff --git a/python/sglang/srt/observability/req_time_stats.py b/python/sglang/srt/observability/req_time_stats.py index 151205db8..326aace02 100644 --- a/python/sglang/srt/observability/req_time_stats.py +++ b/python/sglang/srt/observability/req_time_stats.py @@ -609,7 +609,15 @@ class SchedulerReqTimeStats(ReqTimeStatsBase): ts = ts or time.perf_counter() self.spec_verify_start_time = ts - def set_spec_verify_end_time(self, ts=None, accepted_tokens: int = 0): + def set_spec_verify_end_time( + self, + ts=None, + num_correct_drafts: int = 0, + # FIXME: backward-compat alias, remove in next release. + accepted_tokens: Optional[int] = None, + ): + if accepted_tokens is not None: + num_correct_drafts = accepted_tokens ts = ts or time.perf_counter() if self.trace_ctx.tracing_enable: @@ -618,7 +626,11 @@ class SchedulerReqTimeStats(ReqTimeStatsBase): stage, self.spec_verify_start_time, ts, - {"accepted_tokens": accepted_tokens}, + { + "num_correct_drafts": num_correct_drafts, + # FIXME: backward-compat alias, remove in next release. + "accepted_tokens": num_correct_drafts, + }, ) def set_spec_draft_extend_start_time(self, ts=None): diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py index d92a99dc3..3edba378e 100644 --- a/python/sglang/srt/speculative/eagle_info.py +++ b/python/sglang/srt/speculative/eagle_info.py @@ -426,11 +426,11 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): # Iterate every accepted token and check if req has finished after append the token # should be checked BEFORE free kv cache slots for i, (req, accept_index_row) in enumerate(zip(batch.reqs, accept_index_cpu)): - num_accepted = 0 + num_accept_tokens = 0 for j, idx in enumerate(accept_index_row): if idx == -1: break - num_accepted += 1 + num_accept_tokens += 1 id = predict_cpu[idx] req.output_ids.append(id) if req.require_reasoning and think_end_id is not None: @@ -451,7 +451,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): accept_index[i, j + 1 :] = -1 break # Update KV cache tracking for the accepted tokens - req.kv_committed_len += num_accepted + req.kv_committed_len += num_accept_tokens req.kv_allocated_len = req.kv_committed_len if not req.finished(): unfinished_index.append(i) diff --git a/python/sglang/srt/speculative/eagle_worker.py b/python/sglang/srt/speculative/eagle_worker.py index aa6e73380..a8b07d09b 100644 --- a/python/sglang/srt/speculative/eagle_worker.py +++ b/python/sglang/srt/speculative/eagle_worker.py @@ -491,8 +491,12 @@ class EAGLEWorker(TpModelWorker): if get_global_tracing_enabled(): for idx, req in enumerate(batch.reqs): - accepted = verify_output.num_correct_drafts_per_req_cpu[idx] - req.time_stats.set_spec_verify_end_time(accepted_tokens=accepted) + num_correct_drafts = verify_output.num_correct_drafts_per_req_cpu[ + idx + ] + req.time_stats.set_spec_verify_end_time( + num_correct_drafts=num_correct_drafts + ) set_time_batch( batch.reqs, "set_spec_draft_extend_start_time", trace_only=True diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_worker.py b/python/sglang/srt/speculative/frozen_kv_mtp_worker.py index e2e477fb3..349f6f216 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_worker.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_worker.py @@ -452,8 +452,10 @@ class FrozenKVMTPWorker(TpModelWorker): if get_global_tracing_enabled(): for idx, req in enumerate(batch.reqs): - accepted = verify_output.num_correct_drafts_per_req_cpu[idx] - req.time_stats.set_spec_verify_end_time(accepted_tokens=accepted) + num_correct_drafts = verify_output.num_correct_drafts_per_req_cpu[idx] + req.time_stats.set_spec_verify_end_time( + num_correct_drafts=num_correct_drafts + ) set_time_batch(batch.reqs, "set_spec_draft_extend_start_time", trace_only=True) with self.draft_tp_context( diff --git a/python/sglang/srt/speculative/ngram_worker.py b/python/sglang/srt/speculative/ngram_worker.py index 19bd77885..49cbe7022 100644 --- a/python/sglang/srt/speculative/ngram_worker.py +++ b/python/sglang/srt/speculative/ngram_worker.py @@ -321,12 +321,14 @@ class NGRAMWorker: if get_global_tracing_enabled(): for idx, req in enumerate(batch.reqs): - accepted = ( + num_correct_drafts = ( verify_input.num_correct_drafts[idx].item() if verify_input.num_correct_drafts is not None else 0 ) - req.time_stats.set_spec_verify_end_time(accepted_tokens=accepted) + req.time_stats.set_spec_verify_end_time( + num_correct_drafts=num_correct_drafts + ) # Store accept_lens (with bonus) for per-request metrics; downstream # subtracts 1 to recover drafts-only counts. diff --git a/python/sglang/srt/speculative/spec_utils.py b/python/sglang/srt/speculative/spec_utils.py index 5c879a978..fa4f2eded 100644 --- a/python/sglang/srt/speculative/spec_utils.py +++ b/python/sglang/srt/speculative/spec_utils.py @@ -618,19 +618,19 @@ def traverse_tree( if curr == 0: # the first token generated by the target model, and thus it is always # accepted from the previous iteration - accepted = True + is_accepted = True else: parent_bitmask = allocate_token_bitmask[parent_pos] curr_token_id = draft_tokens[curr] if vocab_size and curr_token_id >= vocab_size: - accepted = False + is_accepted = False else: # 32 boolean bitmask values are packed into 32-bit integers - accepted = ( + is_accepted = ( parent_bitmask[curr_token_id // 32] & (1 << (curr_token_id % 32)) ) != 0 - if accepted: + if is_accepted: if curr != 0: # Accept the current token grammar.accept_token(int(draft_tokens[curr]))