[Spec] Rename accepted_drafts -> correct_drafts for unambiguous naming (#24081)

This commit is contained in:
Liangsheng Yin
2026-05-11 22:12:40 -07:00
committed by GitHub
parent 6bfb365c07
commit 4ad63ad02f
7 changed files with 44 additions and 16 deletions
@@ -2125,9 +2125,13 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
recv_obj.completion_tokens[i] / recv_obj.spec_verify_ct[i]
)
meta_info["spec_num_correct_drafts"] = num_correct_drafts
meta_info["spec_num_proposed_drafts"] = num_proposed_drafts
meta_info["spec_verify_ct"] = recv_obj.spec_verify_ct[i]
# FIXME: backward-compat aliases, remove in next release.
meta_info["spec_accepted_drafts"] = num_correct_drafts
meta_info["spec_proposed_drafts"] = num_proposed_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.
if (
@@ -2135,6 +2139,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
and len(recv_obj.spec_correct_drafts_histogram) > i
and recv_obj.spec_correct_drafts_histogram[i]
):
meta_info["spec_correct_drafts_histogram"] = (
recv_obj.spec_correct_drafts_histogram[i]
)
# FIXME: backward-compat alias, remove in next release.
meta_info["spec_accept_histogram"] = (
recv_obj.spec_correct_drafts_histogram[i]
)
@@ -609,7 +609,15 @@ class SchedulerReqTimeStats(ReqTimeStatsBase):
ts = ts or time.perf_counter()
self.spec_verify_start_time = ts
def set_spec_verify_end_time(self, ts=None, accepted_tokens: int = 0):
def set_spec_verify_end_time(
self,
ts=None,
num_correct_drafts: int = 0,
# FIXME: backward-compat alias, remove in next release.
accepted_tokens: Optional[int] = None,
):
if accepted_tokens is not None:
num_correct_drafts = accepted_tokens
ts = ts or time.perf_counter()
if self.trace_ctx.tracing_enable:
@@ -618,7 +626,11 @@ class SchedulerReqTimeStats(ReqTimeStatsBase):
stage,
self.spec_verify_start_time,
ts,
{"accepted_tokens": accepted_tokens},
{
"num_correct_drafts": num_correct_drafts,
# FIXME: backward-compat alias, remove in next release.
"accepted_tokens": num_correct_drafts,
},
)
def set_spec_draft_extend_start_time(self, ts=None):
+3 -3
View File
@@ -426,11 +426,11 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
# Iterate every accepted token and check if req has finished after append the token
# should be checked BEFORE free kv cache slots
for i, (req, accept_index_row) in enumerate(zip(batch.reqs, accept_index_cpu)):
num_accepted = 0
num_accept_tokens = 0
for j, idx in enumerate(accept_index_row):
if idx == -1:
break
num_accepted += 1
num_accept_tokens += 1
id = predict_cpu[idx]
req.output_ids.append(id)
if req.require_reasoning and think_end_id is not None:
@@ -451,7 +451,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
accept_index[i, j + 1 :] = -1
break
# Update KV cache tracking for the accepted tokens
req.kv_committed_len += num_accepted
req.kv_committed_len += num_accept_tokens
req.kv_allocated_len = req.kv_committed_len
if not req.finished():
unfinished_index.append(i)
@@ -491,8 +491,12 @@ class EAGLEWorker(TpModelWorker):
if get_global_tracing_enabled():
for idx, req in enumerate(batch.reqs):
accepted = verify_output.num_correct_drafts_per_req_cpu[idx]
req.time_stats.set_spec_verify_end_time(accepted_tokens=accepted)
num_correct_drafts = verify_output.num_correct_drafts_per_req_cpu[
idx
]
req.time_stats.set_spec_verify_end_time(
num_correct_drafts=num_correct_drafts
)
set_time_batch(
batch.reqs, "set_spec_draft_extend_start_time", trace_only=True
@@ -452,8 +452,10 @@ class FrozenKVMTPWorker(TpModelWorker):
if get_global_tracing_enabled():
for idx, req in enumerate(batch.reqs):
accepted = verify_output.num_correct_drafts_per_req_cpu[idx]
req.time_stats.set_spec_verify_end_time(accepted_tokens=accepted)
num_correct_drafts = verify_output.num_correct_drafts_per_req_cpu[idx]
req.time_stats.set_spec_verify_end_time(
num_correct_drafts=num_correct_drafts
)
set_time_batch(batch.reqs, "set_spec_draft_extend_start_time", trace_only=True)
with self.draft_tp_context(
@@ -321,12 +321,14 @@ class NGRAMWorker:
if get_global_tracing_enabled():
for idx, req in enumerate(batch.reqs):
accepted = (
num_correct_drafts = (
verify_input.num_correct_drafts[idx].item()
if verify_input.num_correct_drafts is not None
else 0
)
req.time_stats.set_spec_verify_end_time(accepted_tokens=accepted)
req.time_stats.set_spec_verify_end_time(
num_correct_drafts=num_correct_drafts
)
# Store accept_lens (with bonus) for per-request metrics; downstream
# subtracts 1 to recover drafts-only counts.
+4 -4
View File
@@ -618,19 +618,19 @@ def traverse_tree(
if curr == 0:
# the first token generated by the target model, and thus it is always
# accepted from the previous iteration
accepted = True
is_accepted = True
else:
parent_bitmask = allocate_token_bitmask[parent_pos]
curr_token_id = draft_tokens[curr]
if vocab_size and curr_token_id >= vocab_size:
accepted = False
is_accepted = False
else:
# 32 boolean bitmask values are packed into 32-bit integers
accepted = (
is_accepted = (
parent_bitmask[curr_token_id // 32] & (1 << (curr_token_id % 32))
) != 0
if accepted:
if is_accepted:
if curr != 0:
# Accept the current token
grammar.accept_token(int(draft_tokens[curr]))