[Spec] Fix spec_accept_rate and unify accept/draft naming (#23530)
This commit is contained in:
@@ -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,
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
@@ -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
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
Reference in New Issue
Block a user