[Spec] Fix spec_accept_rate and unify accept/draft naming (#23530)

This commit is contained in:
Liangsheng Yin
2026-04-28 14:40:04 -07:00
committed by GitHub
parent 3e1c5e1b74
commit cf0061da43
16 changed files with 76 additions and 65 deletions
@@ -339,7 +339,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
cached_tokens=recv_obj.cached_tokens,
cached_tokens_details=recv_obj.cached_tokens_details,
spec_verify_ct=recv_obj.spec_verify_ct,
spec_accepted_tokens=recv_obj.spec_accepted_tokens,
spec_accepted_drafts=recv_obj.spec_accepted_drafts,
spec_acceptance_histogram=recv_obj.spec_acceptance_histogram,
input_token_logprobs_val=recv_obj.input_token_logprobs_val,
input_token_logprobs_idx=recv_obj.input_token_logprobs_idx,
+9 -3
View File
@@ -93,8 +93,9 @@ class SpeculativeDecodingMetricsMixin:
# Verify count: number of verification forward passes
spec_verify_ct: List[int]
# Accepted tokens: Number of accepted tokens during speculative decoding
spec_accepted_tokens: List[int]
# Accepted drafts: Number of accepted draft tokens during speculative decoding
# (strict drafts-only count, excludes the bonus token).
spec_accepted_drafts: List[int]
# Acceptance histogram: List of lists, where each inner list represents histogram counts.
# List index = number of accepted tokens in a step, List value = count of steps with that many accepted tokens.
@@ -1915,7 +1916,12 @@ class SpeculativeMetrics:
"""Speculative decoding metrics."""
accept_length: float = field(
metadata={"metric": ("gauge", "Avg accepted tokens per step")}
metadata={
"metric": (
"gauge",
"Mean acceptance length (accepted drafts + bonus token per forward)",
)
}
)
accept_rate: float = field(
metadata={"metric": ("gauge", "Speculative acceptance rate")}
@@ -125,8 +125,8 @@ def _handle_output_by_index(output, i):
new_output = BatchTokenIDOutput(
rids=[output.rids[i]],
spec_verify_ct=_extract_field_by_index(output, "spec_verify_ct", i),
spec_accepted_tokens=_extract_field_by_index(
output, "spec_accepted_tokens", i
spec_accepted_drafts=_extract_field_by_index(
output, "spec_accepted_drafts", i
),
spec_acceptance_histogram=_extract_field_by_index(
output, "spec_acceptance_histogram", i
@@ -213,8 +213,8 @@ def _handle_output_by_index(output, i):
new_output = BatchStrOutput(
rids=[output.rids[i]],
spec_verify_ct=_extract_field_by_index(output, "spec_verify_ct", i),
spec_accepted_tokens=_extract_field_by_index(
output, "spec_accepted_tokens", i
spec_accepted_drafts=_extract_field_by_index(
output, "spec_accepted_drafts", i
),
spec_acceptance_histogram=_extract_field_by_index(
output, "spec_acceptance_histogram", i
+3 -5
View File
@@ -839,13 +839,11 @@ class Req(ReqDllmMixin):
False # Track if breakdown was already computed
)
# The number of verification forward passes in the speculative decoding.
# This is used to compute the average acceptance length per request.
# Per-request count of verification forward passes.
self.spec_verify_ct = 0
# The number of accepted tokens in speculative decoding for this request.
# This is used to compute the acceptance rate and average acceptance length per request.
self.spec_accepted_tokens = 0
# Per-request count of accepted draft tokens (excludes the bonus token).
self.spec_accepted_drafts = 0
# Acceptance histogram for speculative decoding.
# List index = number of accepted tokens in a step, List value = count of steps with that many accepted tokens.
@@ -357,7 +357,7 @@ class SchedulerOutputProcessorMixin:
next_token_ids = result.next_token_ids.tolist()
accept_lens = result.accept_lens.tolist()
result.num_accepted_tokens = sum(accept_lens) - len(batch.reqs)
result.num_accepted_drafts = sum(accept_lens) - len(batch.reqs)
result.accept_length_per_req_cpu = [x - 1 for x in accept_lens]
predict_tokens = []
@@ -372,7 +372,7 @@ class SchedulerOutputProcessorMixin:
req.spec_verify_ct += 1
accepted_draft_tokens = result.accept_length_per_req_cpu[i]
req.spec_accepted_tokens += accepted_draft_tokens
req.spec_accepted_drafts += accepted_draft_tokens
req.update_spec_acceptance_histogram(accepted_draft_tokens)
return predict_tokens
@@ -432,7 +432,7 @@ class SchedulerOutputProcessorMixin:
self.num_generated_tokens += len(batch.reqs)
if not batch.spec_algorithm.is_none():
self.update_spec_metrics(batch.batch_size(), result.num_accepted_tokens)
self.update_spec_metrics(batch.batch_size(), result.num_accepted_drafts)
if self.enable_metrics:
self.metrics_collector.increment_decode_cuda_graph_pass(
value=can_run_cuda_graph
@@ -545,7 +545,7 @@ class SchedulerOutputProcessorMixin:
self.report_decode_stats(
can_run_cuda_graph,
running_batch=batch,
num_accepted_tokens=result.num_accepted_tokens,
num_accepted_drafts=result.num_accepted_drafts,
)
def _handle_finished_req(
@@ -962,7 +962,7 @@ class SchedulerOutputProcessorMixin:
cached_tokens = []
cached_tokens_details = [] # Detailed breakdown by cache source
spec_verify_ct = []
spec_accepted_tokens = []
spec_accepted_drafts = []
spec_acceptance_histogram = []
retraction_counts = []
output_hidden_states = None
@@ -1071,7 +1071,7 @@ class SchedulerOutputProcessorMixin:
if not self.spec_algorithm.is_none():
spec_verify_ct.append(req.spec_verify_ct)
spec_accepted_tokens.append(req.spec_accepted_tokens)
spec_accepted_drafts.append(req.spec_accepted_drafts)
spec_acceptance_histogram.append(req.spec_acceptance_histogram)
if return_logprob:
@@ -1176,7 +1176,7 @@ class SchedulerOutputProcessorMixin:
rids=rids,
http_worker_ipcs=http_worker_ipcs,
spec_verify_ct=spec_verify_ct,
spec_accepted_tokens=spec_accepted_tokens,
spec_accepted_drafts=spec_accepted_drafts,
spec_acceptance_histogram=spec_acceptance_histogram,
time_stats=time_stats,
finished_reasons=finished_reasons,
+14 -11
View File
@@ -2081,23 +2081,26 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
if (
hasattr(recv_obj, "spec_verify_ct")
and recv_obj.spec_verify_ct[i] > 0
and hasattr(recv_obj, "spec_accepted_tokens")
and len(recv_obj.spec_accepted_tokens) > i
and hasattr(recv_obj, "spec_accepted_drafts")
and len(recv_obj.spec_accepted_drafts) > i
):
# The draft tokens per speculative step (excluding the target-sampled token).
num_guess_tokens = self.server_args.speculative_num_draft_tokens - 1
total_draft_tokens = recv_obj.spec_verify_ct[i] * num_guess_tokens
accepted_tokens = recv_obj.spec_accepted_tokens[i]
# Total number of proposed draft tokens per request.
all_drafts = recv_obj.spec_verify_ct[i] * (
self.server_args.speculative_num_draft_tokens - 1
)
accepted_drafts = recv_obj.spec_accepted_drafts[i]
# Calculate per-request acceptance rate and average acceptance length.
if total_draft_tokens > 0:
# Calculate acceptance rate: accepted / (steps * lookahead)
meta_info["spec_accept_rate"] = accepted_tokens / total_draft_tokens
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).
meta_info["spec_accept_length"] = (
recv_obj.completion_tokens[i] / recv_obj.spec_verify_ct[i]
)
meta_info["spec_accept_token_num"] = accepted_tokens
meta_info["spec_draft_token_num"] = total_draft_tokens
meta_info["spec_accepted_drafts"] = accepted_drafts
meta_info["spec_proposed_drafts"] = all_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.
+1 -1
View File
@@ -27,7 +27,7 @@ class GenerationBatchResult:
logits_output: Optional[LogitsProcessorOutput] = None
pp_hidden_states_proxy_tensors: Optional[PPProxyTensors] = None
next_token_ids: Optional[Union[torch.Tensor, List[torch.Tensor]]] = None
num_accepted_tokens: int = 0
num_accepted_drafts: int = 0 # no bonus included
accept_length_per_req_cpu: Optional[List[int]] = None
can_run_cuda_graph: bool = False
@@ -304,13 +304,13 @@ class SchedulerMetricsCollector:
# Speculative decoding
self.spec_accept_length = Gauge(
name="sglang:spec_accept_length",
documentation="The average acceptance length of speculative decoding.",
documentation="Mean acceptance length of speculative decoding (accepted drafts + bonus token per forward).",
labelnames=labels.keys(),
multiprocess_mode="mostrecent",
)
self.spec_accept_rate = Gauge(
name="sglang:spec_accept_rate",
documentation="The average acceptance rate of speculative decoding (`accepted tokens / total draft tokens` in batch).",
documentation="Speculative acceptance rate (`accepted drafts / proposed drafts` in batch).",
labelnames=labels.keys(),
multiprocess_mode="mostrecent",
)
@@ -99,11 +99,12 @@ class SchedulerMetricsMixin:
self.last_input_throughput: float = 0.0
self.step_time_dict = defaultdict(list) # Dict[batch size -> step time]
# The number of accepted tokens and forward ct for the recent `decode_log_interval` batches (for logging)
self.spec_num_accepted_tokens = 0
# Cumulative spec-decoding counters (reset every decode_log_interval).
# Each update adds (num_accepted_drafts + bs, bs).
# `*_accepted_tokens` = drafts + bonus; `*_accepted_drafts` = drafts-only.
self.spec_num_accepted_tokens = 0 # per-log-interval
self.spec_num_forward_ct = 0
# The total number of accepted tokens and forward ct for the whole server lifetime
self.spec_total_num_accepted_tokens = 0
self.spec_total_num_accepted_tokens = 0 # lifetime
self.spec_total_num_forward_ct = 0
# For PD disaggregation
@@ -180,10 +181,12 @@ class SchedulerMetricsMixin:
kv_events_config, self.attn_dp_rank
)
def update_spec_metrics(self: Scheduler, bs: int, num_accepted_tokens: int):
self.spec_num_accepted_tokens += num_accepted_tokens + bs
def update_spec_metrics(self: Scheduler, bs: int, num_accepted_drafts: int):
self.spec_num_accepted_tokens += num_accepted_drafts + bs
self.spec_num_forward_ct += bs
self.num_generated_tokens += num_accepted_tokens
# Bonus tokens updated elsewhere
self.num_generated_tokens += num_accepted_drafts
def _init_estimated_perf_constants(self: Scheduler) -> None:
model_config = self.model_config
@@ -464,13 +467,13 @@ class SchedulerMetricsMixin:
self: Scheduler,
can_run_cuda_graph: bool,
running_batch: ScheduleBatch = None,
num_accepted_tokens: int = 0,
num_accepted_drafts: int = 0,
):
batch = running_batch or self.running_batch
# Every-iteration work: realtime token counting + status logger
if self.current_scheduler_metrics_enabled:
decode_tokens = batch.batch_size() + num_accepted_tokens
decode_tokens = batch.batch_size() + num_accepted_drafts
self.metrics_collector.increment_realtime_tokens(
# TODO unify this w/ the bumping logic in `Scheduler.num_generated_tokens` accumulator
decode_tokens=decode_tokens,
@@ -527,15 +530,16 @@ class SchedulerMetricsMixin:
spec_accept_length = (
self.spec_num_accepted_tokens / self.spec_num_forward_ct
)
# Calculate acceptance rate: accepted tokens / total draft tokens
draft_tokens_fallback = (self.server_args.speculative_num_steps or 0) + 1
num_draft_tokens = (
self.server_args.speculative_num_draft_tokens or draft_tokens_fallback
num_accepted_drafts = (
self.spec_num_accepted_tokens - self.spec_num_forward_ct
)
total_draft_tokens = self.spec_num_forward_ct * num_draft_tokens
if self.server_args.speculative_num_draft_tokens:
draft_per_round = self.server_args.speculative_num_draft_tokens - 1
else:
draft_per_round = self.server_args.speculative_num_steps or 0
total_draft_tokens = self.spec_num_forward_ct * draft_per_round
spec_accept_rate = (
self.spec_num_accepted_tokens / total_draft_tokens
num_accepted_drafts / total_draft_tokens
if total_draft_tokens > 0
else 0
)
+1 -1
View File
@@ -422,7 +422,7 @@ class DFlashVerifyInput(SpecInput):
new_verified_list.append(new_verified_token)
accept_length_per_req_cpu.append(max(0, appended - 1))
req.spec_verify_ct += 1
req.spec_accepted_tokens += accept_length_per_req_cpu[-1]
req.spec_accepted_drafts += accept_length_per_req_cpu[-1]
commit_lens = torch.tensor(commit_lens_cpu, dtype=torch.int32, device=device)
new_verified_id = torch.tensor(
@@ -1178,7 +1178,7 @@ class DFlashWorker:
return GenerationBatchResult(
logits_output=logits_output,
next_token_ids=next_token_ids,
num_accepted_tokens=0,
num_accepted_drafts=0,
can_run_cuda_graph=batch_result.can_run_cuda_graph,
)
@@ -1239,7 +1239,7 @@ class DFlashWorker:
batch.spec_info = draft_input
batch.forward_mode = ForwardMode.DECODE
num_accepted_tokens = sum(accept_length_per_req_cpu)
num_accepted_drafts = sum(accept_length_per_req_cpu)
if not self._logged_first_verify and self.tp_rank == 0:
logger.info(
"DFLASH verify completed. accept_length_per_req=%s",
@@ -1250,7 +1250,7 @@ class DFlashWorker:
return GenerationBatchResult(
logits_output=logits_output,
next_token_ids=new_verified_id,
num_accepted_tokens=num_accepted_tokens,
num_accepted_drafts=num_accepted_drafts,
accept_length_per_req_cpu=accept_length_per_req_cpu,
can_run_cuda_graph=can_run_cuda_graph,
)
+1 -1
View File
@@ -454,7 +454,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
unfinished_accept_index.append(accept_index[i])
req.spec_verify_ct += 1
accepted_draft_tokens = sum(1 for idx in accept_index_row if idx != -1) - 1
req.spec_accepted_tokens += accepted_draft_tokens
req.spec_accepted_drafts += accepted_draft_tokens
req.update_spec_acceptance_histogram(accepted_draft_tokens)
if has_finished:
@@ -463,7 +463,7 @@ class EAGLEWorker(TpModelWorker):
return GenerationBatchResult(
logits_output=logits_output,
next_token_ids=next_token_ids,
num_accepted_tokens=0,
num_accepted_drafts=0,
can_run_cuda_graph=can_run_cuda_graph,
)
else:
@@ -513,7 +513,7 @@ class EAGLEWorker(TpModelWorker):
return GenerationBatchResult(
logits_output=logits_output,
next_token_ids=verify_output.verified_id,
num_accepted_tokens=sum(verify_output.accept_length_per_req_cpu),
num_accepted_drafts=sum(verify_output.accept_length_per_req_cpu),
accept_length_per_req_cpu=verify_output.accept_length_per_req_cpu,
can_run_cuda_graph=can_run_cuda_graph,
)
@@ -264,7 +264,7 @@ class MultiLayerEagleWorker(TpModelWorker):
return GenerationBatchResult(
logits_output=logits_output,
next_token_ids=next_token_ids,
num_accepted_tokens=0,
num_accepted_drafts=0,
can_run_cuda_graph=can_run_cuda_graph,
)
else:
@@ -291,7 +291,7 @@ class MultiLayerEagleWorker(TpModelWorker):
return GenerationBatchResult(
logits_output=logits_output,
next_token_ids=verify_output.verified_id,
num_accepted_tokens=sum(verify_output.accept_length_per_req_cpu),
num_accepted_drafts=sum(verify_output.accept_length_per_req_cpu),
can_run_cuda_graph=can_run_cuda_graph,
)
+3 -3
View File
@@ -192,7 +192,7 @@ class NgramVerifyInput(SpecInput):
raise e
req.spec_verify_ct += 1
accepted_draft_tokens = sum(1 for idx in accept_index_row if idx != -1) - 1
req.spec_accepted_tokens += accepted_draft_tokens
req.spec_accepted_drafts += accepted_draft_tokens
req.update_spec_acceptance_histogram(accepted_draft_tokens)
if has_finished:
@@ -446,14 +446,14 @@ class NgramVerifyInput(SpecInput):
self._fill_requests(batch, logits_output)
accept_length_cpu = self.accept_length.cpu()
num_accepted_tokens = accept_length_cpu.sum().item()
num_accepted_drafts = accept_length_cpu.sum().item()
self._free_cache(batch, page_size, accept_length_cpu)
batch.seq_lens.add_(self.accept_length + 1)
batch.seq_lens_cpu.add_(accept_length_cpu + 1)
return logits_output, self.verified_id, num_accepted_tokens
return logits_output, self.verified_id, num_accepted_drafts
def filter_batch(self, new_indices: torch.Tensor, has_been_filtered: bool = True):
pass
@@ -260,7 +260,7 @@ class NGRAMWorker:
model_worker_batch = batch.get_model_worker_batch()
spec_info = model_worker_batch.spec_info
num_accepted_tokens = 0
num_accepted_drafts = 0
accept_lens = None
accept_length_per_req_cpu = None
@@ -303,7 +303,7 @@ class NGRAMWorker:
# and will be applied to produce wrong results
batch.sampling_info.vocab_mask = None
logits_output, next_token_ids, num_accepted_tokens = verify_input.verify(
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()
@@ -347,7 +347,7 @@ class NGRAMWorker:
return GenerationBatchResult(
logits_output=logits_output,
next_token_ids=next_token_ids,
num_accepted_tokens=num_accepted_tokens,
num_accepted_drafts=num_accepted_drafts,
accept_length_per_req_cpu=accept_length_per_req_cpu,
can_run_cuda_graph=can_run_cuda_graph,
accept_lens=accept_lens,