fix spec dec request level metrics (#13754)
Co-authored-by: Liangsheng Yin <lsyincs@gmail.com>
This commit is contained in:
@@ -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.
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user