use frink_models::tokenizer::StopTokens;
use crate::generate::{DecodeError, FinishReason, GenerationParams};
pub(crate) type PerTokenProbs = Vec<(usize, Vec<f32>)>;
#[allow(clippy::too_many_arguments)]
pub(crate) fn sample_until_stop(
logits: Vec<f32>,
pos: usize,
prompt_ids: &[usize],
stop_tokens: &StopTokens,
params: &GenerationParams,
mut decode_one: impl FnMut(&[usize]) -> Vec<u8>,
engine: &mut dyn DecodeEngine,
mut emit: impl FnMut(&str),
decode_token: &dyn Fn(usize) -> String,
draft_max: usize,
) -> Result<(FinishReason, Vec<usize>, Vec<f32>, PerTokenProbs), DecodeError> {
let ctx = crate::choice_stream::StepContext {
params,
prompt_ids,
stop_tokens,
decode_token,
};
let mut choice = crate::choice_stream::ChoiceStream::new(logits, pos, params);
while choice.check_budget(params) {
if draft_max > 0
&& speculate(
&mut choice,
&ctx,
engine,
&mut decode_one,
&mut emit,
draft_max,
)?
{
continue;
}
choice.step(&ctx, engine, &mut decode_one, &mut emit)?;
}
choice.flush(&mut emit);
Ok(choice.into_parts())
}
fn speculate(
choice: &mut crate::choice_stream::ChoiceStream,
ctx: &crate::choice_stream::StepContext<'_>,
engine: &mut dyn DecodeEngine,
decode_one: &mut impl FnMut(&[usize]) -> Vec<u8>,
emit: &mut impl FnMut(&str),
draft_max: usize,
) -> Result<bool, DecodeError> {
let remaining = ctx
.params
.max_tokens
.saturating_sub(choice.generated_ids.len());
let room = remaining.saturating_sub(1);
if room == 0 {
return Ok(false);
}
let draft = engine.draft(ctx.prompt_ids, &choice.generated_ids, draft_max.min(room));
if draft.is_empty() {
return Ok(false);
}
let Some(rows) = engine.batch(&draft, choice.pos) else {
return Ok(false);
};
let (state, logits, mut history) = choice.verify_inputs();
let block = verify_block(
state,
logits,
&rows,
&draft,
ctx.params,
ctx.prompt_ids,
&mut history,
ctx.stop_tokens,
ctx.decode_token,
)?;
engine.observe(block.accepted, block.drafted);
if block.accepted < draft.len() {
engine.truncate(choice.pos + block.accepted);
}
choice.pos += block.accepted;
let mut stopped = None;
let mut last: Option<usize> = None;
for (i, &t) in block.tokens.iter().enumerate() {
match choice.commit(t, ctx, decode_one, emit) {
crate::choice_stream::Committed::Continue => last = Some(t),
crate::choice_stream::Committed::Stopped(reason) => {
let kept = choice.pos - block.accepted + i.min(block.accepted);
engine.truncate(kept);
choice.pos = kept;
stopped = Some(reason);
break;
}
}
}
if let Some(reason) = stopped {
choice.stop(reason);
return Ok(true);
}
if block.grammar_complete {
choice.stop(FinishReason::Stop);
return Ok(true);
}
if let Some(t) = last {
choice.logits = engine.step(t, choice.pos);
choice.pos += 1;
}
Ok(true)
}
pub(crate) fn earliest_stop_match<'a>(text: &str, stops: &'a [String]) -> Option<(usize, &'a str)> {
stops
.iter()
.filter(|s| !s.is_empty())
.filter_map(|s| text.find(s.as_str()).map(|at| (at, s.as_str())))
.min_by_key(|(at, s)| (*at, std::cmp::Reverse(s.len())))
}
pub(crate) fn floor_char_boundary(s: &str, idx: usize) -> usize {
crate::policy::detokenize::floor_char_boundary(s, idx)
}
pub(crate) const DEFAULT_DRAFT_MAX: usize = 5;
pub(crate) const DRAFT_NGRAM: usize = 2;
pub(crate) trait DecodeEngine {
fn step(&mut self, token: usize, pos: usize) -> Vec<f32>;
fn draft(&mut self, _prompt: &[usize], _history: &[usize], _max: usize) -> Vec<usize> {
Vec::new()
}
fn batch(&mut self, _tokens: &[usize], _pos: usize) -> Option<Vec<Vec<f32>>> {
None
}
fn truncate(&mut self, _pos: usize) {}
fn observe(&mut self, _accepted: usize, _drafted: usize) {}
}
impl<F: FnMut(usize, usize) -> Vec<f32>> DecodeEngine for F {
fn step(&mut self, token: usize, pos: usize) -> Vec<f32> {
self(token, pos)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[allow(
dead_code,
reason = "wired by the server speculative-decoding row; see above"
)]
pub(crate) struct VerifiedBlock {
pub tokens: Vec<usize>,
pub accepted: usize,
pub drafted: usize,
pub grammar_complete: bool,
}
#[allow(clippy::too_many_arguments)]
#[allow(
dead_code,
reason = "wired by the server speculative-decoding row; see VerifiedBlock"
)]
pub(crate) fn verify_block(
state: &mut crate::sample_step::SampleState,
first_logits: &[f32],
draft_logits: &[Vec<f32>],
draft: &[usize],
params: &GenerationParams,
prompt: &[usize],
history: &mut Vec<usize>,
stop_tokens: &StopTokens,
decode_token: &dyn Fn(usize) -> String,
) -> Result<VerifiedBlock, DecodeError> {
debug_assert_eq!(
draft.len(),
draft_logits.len(),
"one logit row per drafted token, or the rows and the drafts have drifted"
);
let mut out = VerifiedBlock {
tokens: Vec::new(),
accepted: 0,
drafted: draft.len(),
grammar_complete: false,
};
for i in 0..=draft.len() {
let logits = if i == 0 {
first_logits
} else {
&draft_logits[i - 1]
};
match crate::sample_step::sample_next(
state,
logits,
params,
prompt,
history,
stop_tokens,
decode_token,
)? {
crate::sample_step::Step::GrammarComplete => {
out.grammar_complete = true;
return Ok(out);
}
crate::sample_step::Step::Token { id: t, .. } => {
history.push(t);
out.tokens.push(t);
if i == draft.len() {
break;
}
if t == draft[i] {
out.accepted += 1;
} else {
break;
}
}
}
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::sample_step::{sample_next, SampleState, Step};
use frink_models::sampling::SamplingParams;
fn greedy_params() -> GenerationParams {
GenerationParams {
cache_salt: None,
prompt_logprobs: None,
wants_logprobs: false,
n: 1,
interleave_choices: false,
reasoning: None,
max_tokens: 64,
sampling: SamplingParams {
temperature: 0.0,
..SamplingParams::default()
},
seed: 1,
stop: Vec::new(),
stop_token_ids: Vec::new(),
json_object: false,
grammar: None,
cancel: None,
ignore_eos: false,
reasoning_budget: crate::reasoning_budget::ReasoningBudget::Unrestricted,
lora: None,
}
}
fn peaked(vocab: usize, winner: usize) -> Vec<f32> {
let mut v = vec![0.0f32; vocab];
v[winner] = 10.0;
v
}
fn no_stops() -> StopTokens {
StopTokens::default()
}
fn decode(_: usize) -> String {
String::new()
}
#[test]
fn all_drafts_accepted_earn_one_extra_token() {
let params = greedy_params();
let mut state = SampleState::new(params.seed);
let mut history = Vec::new();
let draft = [3usize, 4, 5];
let first = peaked(8, 3);
let rows = vec![peaked(8, 4), peaked(8, 5), peaked(8, 6)];
let b = verify_block(
&mut state,
&first,
&rows,
&draft,
¶ms,
&[],
&mut history,
&no_stops(),
&decode,
)
.expect("verify");
assert_eq!(b.accepted, 3);
assert_eq!(b.drafted, 3);
assert_eq!(b.tokens, vec![3, 4, 5, 6], "three drafts plus the bonus");
assert_eq!(history, vec![3, 4, 5, 6]);
}
#[test]
fn a_wrong_draft_commits_the_samplers_token_and_stops() {
let params = greedy_params();
let mut state = SampleState::new(params.seed);
let mut history = Vec::new();
let draft = [3usize, 99, 5];
let first = peaked(8, 3);
let rows = vec![peaked(8, 4), peaked(8, 7), peaked(8, 1)];
let b = verify_block(
&mut state,
&first,
&rows,
&draft,
¶ms,
&[],
&mut history,
&no_stops(),
&decode,
)
.expect("verify");
assert_eq!(b.accepted, 1, "only the first draft matched");
assert_eq!(b.tokens, vec![3, 4], "the sampler's token, not 99");
assert!(
!b.tokens.contains(&99),
"a draft must never be emitted as itself"
);
assert!(!b.tokens.contains(&7), "rows past the mismatch are invalid");
}
#[test]
fn a_verified_block_commits_what_sequential_sampling_would_have() {
for draft in [
vec![3usize, 4, 5], vec![3, 99, 5], vec![42, 4, 5], ] {
let params = greedy_params();
let first = peaked(8, 3);
let rows = vec![peaked(8, 4), peaked(8, 5), peaked(8, 6)];
let mut spec_state = SampleState::new(params.seed);
let mut spec_history = Vec::new();
let b = verify_block(
&mut spec_state,
&first,
&rows,
&draft,
¶ms,
&[],
&mut spec_history,
&no_stops(),
&decode,
)
.expect("verify");
let mut seq_state = SampleState::new(params.seed);
let mut seq_history: Vec<usize> = Vec::new();
for i in 0..b.tokens.len() {
let logits = if i == 0 { &first } else { &rows[i - 1] };
let Step::Token { id: t, .. } = sample_next(
&mut seq_state,
logits,
¶ms,
&[],
&seq_history,
&no_stops(),
&decode,
)
.expect("sequential") else {
panic!("grammar completed in a grammarless test")
};
seq_history.push(t);
}
assert_eq!(
b.tokens, seq_history,
"draft {draft:?}: a verified block must commit the sequential answer"
);
assert_eq!(
b.tokens.len(),
b.accepted + 1,
"draft {draft:?}: committed {} tokens for {} accepted drafts",
b.tokens.len(),
b.accepted
);
}
}
struct Scripted {
script: Vec<usize>,
vocab: usize,
drafts: Vec<Vec<usize>>,
round: usize,
batches: usize,
accepted: usize,
drafted: usize,
steps: usize,
}
impl Scripted {
fn new(script: &[usize], vocab: usize, drafts: Vec<Vec<usize>>) -> Self {
Scripted {
script: script.to_vec(),
vocab,
drafts,
round: 0,
batches: 0,
accepted: 0,
drafted: 0,
steps: 0,
}
}
fn row(&self, pos: usize) -> Vec<f32> {
peaked(self.vocab, self.script.get(pos).copied().unwrap_or(0))
}
}
impl DecodeEngine for Scripted {
fn step(&mut self, _token: usize, pos: usize) -> Vec<f32> {
self.steps += 1;
self.row(pos + 1)
}
fn draft(&mut self, _p: &[usize], _h: &[usize], max: usize) -> Vec<usize> {
let d = self.drafts.get(self.round).cloned().unwrap_or_default();
self.round += 1;
d.into_iter().take(max).collect()
}
fn batch(&mut self, tokens: &[usize], pos: usize) -> Option<Vec<Vec<f32>>> {
self.batches += 1;
Some((0..tokens.len()).map(|i| self.row(pos + i + 1)).collect())
}
fn truncate(&mut self, _pos: usize) {}
fn observe(&mut self, accepted: usize, drafted: usize) {
self.accepted += accepted;
self.drafted += drafted;
}
}
#[test]
fn speculation_changes_the_forward_count_and_not_the_answer() {
let script = vec![1usize, 2, 3, 4, 5, 6, 7];
let vocab = 16;
let params = GenerationParams {
max_tokens: 6,
..greedy_params()
};
let bytes = |ids: &[usize]| {
ids.iter()
.map(|i| format!("<{i}>"))
.collect::<String>()
.into_bytes()
};
let mut plain = Scripted::new(&script, vocab, Vec::new());
let mut plain_text = String::new();
let (_, plain_ids, _, _probs) = sample_until_stop(
peaked(vocab, script[0]),
0,
&[],
&no_stops(),
¶ms,
bytes,
&mut plain,
|c| plain_text.push_str(c),
&decode,
0,
)
.expect("plain");
for (label, drafts) in [
("always right", vec![vec![1usize, 2, 3], vec![5, 6, 7]]),
("always wrong", vec![vec![9usize, 9, 9], vec![9, 9, 9]]),
(
"mixed, and silent",
vec![vec![1usize, 9, 3], vec![], vec![5, 6]],
),
] {
let mut engine = Scripted::new(&script, vocab, drafts);
let mut spec_text = String::new();
let (_, spec_ids, _, _probs) = sample_until_stop(
peaked(vocab, script[0]),
0,
&[],
&no_stops(),
¶ms,
bytes,
&mut engine,
|c| spec_text.push_str(c),
&decode,
3,
)
.expect("speculative");
println!(
"CASE {label} accepted={} drafted={} spec_steps={} plain_steps={}",
engine.accepted, engine.drafted, engine.steps, plain.steps
);
assert_eq!(
spec_ids, plain_ids,
"{label}: speculation changed the ids (plain={plain_ids:?} spec={spec_ids:?})"
);
assert_eq!(
spec_text, plain_text,
"{label}: speculation changed the text"
);
if label == "always wrong" {
assert_eq!(engine.accepted, 0, "a wrong draft must never be accepted");
} else {
assert!(
engine.accepted > 0,
"{label}: nothing was accepted, so this case proves nothing"
);
assert!(
engine.steps < plain.steps,
"{label}: accepted {} drafts and still took {} forwards against {}",
engine.accepted,
engine.steps,
plain.steps
);
}
}
}
#[test]
fn an_empty_draft_commits_exactly_one_token() {
let params = greedy_params();
let mut state = SampleState::new(params.seed);
let mut history = Vec::new();
let b = verify_block(
&mut state,
&peaked(8, 2),
&[],
&[],
¶ms,
&[],
&mut history,
&no_stops(),
&decode,
)
.expect("verify");
assert_eq!(b.tokens, vec![2]);
assert_eq!(b.accepted, 0);
assert_eq!(b.drafted, 0);
}
#[test]
fn earliest_stop_match_finds_the_leftmost_match_across_multiple_stops() {
assert_eq!(
earliest_stop_match("hello world", &["world".to_string(), "hello".to_string()]),
Some((0, "hello")),
"the leftmost match wins, not the caller's first entry"
);
assert_eq!(
earliest_stop_match("hello world", &["nope".to_string()]),
None
);
}
}