From 9dab534b35586ee0668da0ca7718911f9a458121 Mon Sep 17 00:00:00 2001 From: Vedant V Jhaveri Date: Wed, 26 Nov 2025 09:09:21 -0800 Subject: [PATCH] fix spec dec request level metrics (#13754) Co-authored-by: Liangsheng Yin --- python/sglang/srt/managers/tokenizer_manager.py | 9 +++++---- python/sglang/srt/speculative/ngram_info.py | 4 ++++ python/sglang/srt/speculative/ngram_worker.py | 4 ++++ 3 files changed, 13 insertions(+), 4 deletions(-) diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index c27a054da..f1a02ae62 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -1837,17 +1837,18 @@ class TokenizerManager(TokenizerCommunicatorMixin): meta_info["spec_accept_length"] = 0 meta_info["spec_verify_ct"] = recv_obj.spec_verify_ct[i] + # The draft tokens per speculative step (excluding the target-sampled token). + num_guess_tokens = self.server_args.speculative_num_draft_tokens - 1 + if ( recv_obj.spec_verify_ct[i] > 0 - and self.server_args.speculative_num_steps is not None + and num_guess_tokens is not None and not isinstance(recv_obj, BatchEmbeddingOutput) and hasattr(recv_obj, "spec_accepted_tokens") # Checks that `spec_accepted_tokens[i]` will exist. and len(recv_obj.spec_accepted_tokens) > i ): - total_draft_tokens = ( - recv_obj.spec_verify_ct[i] * self.server_args.speculative_num_steps - ) + total_draft_tokens = recv_obj.spec_verify_ct[i] * num_guess_tokens accepted_tokens = recv_obj.spec_accepted_tokens[i] # Calculate per-request acceptance rate and average acceptance length. diff --git a/python/sglang/srt/speculative/ngram_info.py b/python/sglang/srt/speculative/ngram_info.py index 5ba756aa3..637f51c97 100644 --- a/python/sglang/srt/speculative/ngram_info.py +++ b/python/sglang/srt/speculative/ngram_info.py @@ -185,6 +185,10 @@ class NgramVerifyInput(SpecInput): ) raise e req.spec_verify_ct += 1 + req.spec_accepted_tokens += ( + sum(1 for idx in accept_index_row if idx != -1) - 1 + ) + if has_finished: self.accept_length = (self.accept_index != -1).sum(dim=1) - 1 self.accept_index = self.accept_index[self.accept_index != -1] diff --git a/python/sglang/srt/speculative/ngram_worker.py b/python/sglang/srt/speculative/ngram_worker.py index d6ad689c2..5c61ab31a 100644 --- a/python/sglang/srt/speculative/ngram_worker.py +++ b/python/sglang/srt/speculative/ngram_worker.py @@ -295,6 +295,7 @@ class NGRAMWorker: self._prepare_for_speculative_decoding(batch) model_worker_batch = batch.get_model_worker_batch() num_accepted_tokens = 0 + accept_lens = None if model_worker_batch.forward_mode.is_target_verify(): batch_result = self.target_worker.forward_batch_generation( @@ -308,6 +309,8 @@ class NGRAMWorker: logits_output, next_token_ids, num_accepted_tokens = verify_input.verify( batch, logits_output, self.page_size ) + # Store accept_lens for per-request metrics + accept_lens = verify_input.accept_length if batch.return_logprob: self.add_logprob_values(batch, verify_input, logits_output) self._update_ngram_cache(batch) @@ -328,4 +331,5 @@ class NGRAMWorker: next_token_ids=next_token_ids, num_accepted_tokens=num_accepted_tokens, can_run_cuda_graph=can_run_cuda_graph, + accept_lens=accept_lens, )