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