[Speculative Decoding] Enabling bonus token in speculative decoding for KV cache based models (#5765)

This commit is contained in:
sroy745
2024-07-10 16:02:47 -07:00
committed by GitHub
parent 44cc76610d
commit ae151d73be
14 changed files with 645 additions and 80 deletions

View File

@@ -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