[Speculative Decoding] Enabling bonus token in speculative decoding for KV cache based models (#5765)
This commit is contained in:
@@ -1,5 +1,6 @@
|
||||
from collections import defaultdict
|
||||
from functools import cached_property
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
from typing import Any, Dict, List, Optional, Set, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
@@ -13,7 +14,7 @@ from vllm.model_executor.layers.typical_acceptance_sampler import (
|
||||
TypicalAcceptanceSampler)
|
||||
from vllm.sequence import (CompletionSequenceGroupOutput, ExecuteModelRequest,
|
||||
HiddenStates, SamplerOutput, SequenceGroupMetadata,
|
||||
get_all_seq_ids)
|
||||
get_all_seq_ids_and_request_ids)
|
||||
from vllm.spec_decode.batch_expansion import BatchExpansionTop1Scorer
|
||||
from vllm.spec_decode.draft_model_runner import TP1DraftModelRunner
|
||||
from vllm.spec_decode.interfaces import (SpeculativeProposals,
|
||||
@@ -112,11 +113,7 @@ class SpecDecodeWorker(LoraNotSupportedWorkerBase):
|
||||
draft_worker_kwargs.pop("ngram_prompt_lookup_max"))
|
||||
ngram_prompt_lookup_min = (
|
||||
draft_worker_kwargs.pop("ngram_prompt_lookup_min"))
|
||||
|
||||
disable_bonus_tokens = True
|
||||
|
||||
if ngram_prompt_lookup_max > 0:
|
||||
disable_bonus_tokens = False
|
||||
proposer_worker = NGramWorker(**draft_worker_kwargs)
|
||||
proposer_worker.set_ngram_window_size(ngram_prompt_lookup_min,
|
||||
ngram_prompt_lookup_max)
|
||||
@@ -128,11 +125,9 @@ class SpecDecodeWorker(LoraNotSupportedWorkerBase):
|
||||
|
||||
if draft_worker_kwargs[
|
||||
"model_config"].hf_config.model_type == "mlp_speculator":
|
||||
disable_bonus_tokens = False
|
||||
proposer_worker = MLPSpeculatorWorker(**draft_worker_kwargs)
|
||||
elif draft_worker_kwargs[
|
||||
"model_config"].hf_config.model_type == "medusa":
|
||||
disable_bonus_tokens = False
|
||||
proposer_worker = MedusaWorker(**draft_worker_kwargs)
|
||||
else:
|
||||
if draft_tp == 1:
|
||||
@@ -149,10 +144,10 @@ class SpecDecodeWorker(LoraNotSupportedWorkerBase):
|
||||
spec_decode_sampler: SpecDecodeBaseSampler = None
|
||||
if draft_token_acceptance_method == "rejection_sampler":
|
||||
spec_decode_sampler = RejectionSampler(
|
||||
disable_bonus_tokens=disable_bonus_tokens, )
|
||||
disable_bonus_tokens=False, )
|
||||
elif draft_token_acceptance_method == "typical_acceptance_sampler":
|
||||
spec_decode_sampler = TypicalAcceptanceSampler(
|
||||
disable_bonus_tokens=disable_bonus_tokens,
|
||||
disable_bonus_tokens=False,
|
||||
posterior_threshold=\
|
||||
typical_acceptance_sampler_posterior_threshold,
|
||||
posterior_alpha=typical_acceptance_sampler_posterior_alpha,
|
||||
@@ -200,6 +195,15 @@ class SpecDecodeWorker(LoraNotSupportedWorkerBase):
|
||||
self._metrics = AsyncMetricsCollector(
|
||||
self.spec_decode_sampler
|
||||
) if metrics_collector is None else metrics_collector
|
||||
# Tracks the sequence IDs that received a bonus token ID in
|
||||
# their last forward pass. Needed only if KV cache is being
|
||||
# used for token generation such as in the case of MultiStepWorker.
|
||||
self._seq_with_bonus_token_in_last_step: Set[int] = set()
|
||||
# Tracks the currently active request ids and the sequence IDs
|
||||
# corresponding to them
|
||||
self._request_id_seq_id_mapping: Dict[str, Set[int]] = defaultdict(set)
|
||||
# Tracks if the proposer worker uses the KV cache or not.
|
||||
|
||||
self.probs_dtype = self.spec_decode_sampler.probs_dtype
|
||||
self.token_id_dtype = self.spec_decode_sampler.token_id_dtype
|
||||
# Lazy initiazliation.
|
||||
@@ -307,6 +311,7 @@ class SpecDecodeWorker(LoraNotSupportedWorkerBase):
|
||||
broadcast_tensor_dict({}, src=0)
|
||||
return []
|
||||
|
||||
self._track_finished_requests(execute_model_req)
|
||||
disable_all_speculation = self._should_disable_all_speculation(
|
||||
execute_model_req)
|
||||
num_lookahead_slots = execute_model_req.num_lookahead_slots
|
||||
@@ -453,7 +458,8 @@ class SpecDecodeWorker(LoraNotSupportedWorkerBase):
|
||||
self.previous_hidden_states = None
|
||||
|
||||
# Generate proposals using draft worker.
|
||||
proposals = self.proposer_worker.get_spec_proposals(execute_model_req)
|
||||
proposals = self.proposer_worker.get_spec_proposals(
|
||||
execute_model_req, self._seq_with_bonus_token_in_last_step)
|
||||
|
||||
proposal_scores = self.scorer.score_proposals(
|
||||
execute_model_req,
|
||||
@@ -585,7 +591,9 @@ class SpecDecodeWorker(LoraNotSupportedWorkerBase):
|
||||
|
||||
# Get the sequence ids and num_logprobs (sampling parameter) in the
|
||||
# batch.
|
||||
seq_ids = get_all_seq_ids(seq_group_metadata_list)
|
||||
seq_ids, request_ids_seq_ids_mapping = get_all_seq_ids_and_request_ids(
|
||||
seq_group_metadata_list)
|
||||
|
||||
num_logprobs_per_seq = get_all_num_logprobs(seq_group_metadata_list)
|
||||
|
||||
# Serialize all tensors to CPU Python lists.
|
||||
@@ -608,7 +616,6 @@ class SpecDecodeWorker(LoraNotSupportedWorkerBase):
|
||||
for sequence_index in range(batch_size):
|
||||
# Each sequence may have a different num_logprobs; retrieve it.
|
||||
num_logprobs = num_logprobs_per_seq[sequence_index]
|
||||
|
||||
step_output_token_ids.append(
|
||||
create_sequence_group_output(
|
||||
token_id=accepted_token_ids_by_step[step_index]
|
||||
@@ -623,18 +630,48 @@ class SpecDecodeWorker(LoraNotSupportedWorkerBase):
|
||||
topk_logprobs=topk_logprobs_by_step[step_index]
|
||||
[sequence_index][:num_logprobs],
|
||||
))
|
||||
|
||||
sampler_output_list.append(
|
||||
SamplerOutput(outputs=step_output_token_ids))
|
||||
|
||||
# Populate the data structures needed to keep track of sequences with
|
||||
# bonus tokens.
|
||||
self._track_sequences_with_bonus_tokens(seq_ids,
|
||||
request_ids_seq_ids_mapping,
|
||||
accepted_token_ids_by_step)
|
||||
maybe_rejsample_metrics = (
|
||||
self._metrics.maybe_collect_rejsample_metrics(k))
|
||||
if maybe_rejsample_metrics is not None:
|
||||
sampler_output_list[
|
||||
0].spec_decode_worker_metrics = maybe_rejsample_metrics
|
||||
|
||||
return sampler_output_list
|
||||
|
||||
def _track_finished_requests(self, execute_model_req: ExecuteModelRequest):
|
||||
"""
|
||||
Removes the finished requests and their associated sequence ids from
|
||||
internal book keeping data structures.
|
||||
"""
|
||||
for finished_request in execute_model_req.finished_requests_ids:
|
||||
for seq_id in self._request_id_seq_id_mapping[finished_request]:
|
||||
self._seq_with_bonus_token_in_last_step.discard(seq_id)
|
||||
del self._request_id_seq_id_mapping[finished_request]
|
||||
|
||||
def _track_sequences_with_bonus_tokens(
|
||||
self, seq_ids: List[int],
|
||||
request_ids_seq_ids_mapping: Dict[str, Set[int]],
|
||||
accepted_token_ids_by_step: List[List[int]]):
|
||||
"""
|
||||
Updates the internal data structures which keep track of sequences
|
||||
which have been assigned bonus tokens in their last forward pass.
|
||||
"""
|
||||
for seq_index, seq_id in enumerate(seq_ids):
|
||||
last_token_id = accepted_token_ids_by_step[-1][seq_index]
|
||||
if last_token_id == -1:
|
||||
self._seq_with_bonus_token_in_last_step.discard(seq_id)
|
||||
else:
|
||||
self._seq_with_bonus_token_in_last_step.add(seq_id)
|
||||
for request_id, sequences in request_ids_seq_ids_mapping.items():
|
||||
self._request_id_seq_id_mapping[request_id].update(sequences)
|
||||
|
||||
@cached_property
|
||||
def _vocab_size(self) -> int:
|
||||
"""Get the vocab size of the model and make sure it's consistent between
|
||||
|
||||
Reference in New Issue
Block a user