fix spec dec request level metrics (#13754)

Co-authored-by: Liangsheng Yin <lsyincs@gmail.com>
This commit is contained in:
Vedant V Jhaveri
2025-11-26 09:09:21 -08:00
committed by GitHub
parent 262c3c1fde
commit 9dab534b35
3 changed files with 13 additions and 4 deletions

View File

@@ -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.

View File

@@ -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]

View File

@@ -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,
)