use crate::decoder::Decoder;
use crate::sampling::{sampling_distribution, Sampler, SamplingParams};
use ferrox_core::cache::KvCache;
#[derive(Debug, Clone, PartialEq, Default)]
pub struct DraftDist {
support: Vec<(usize, f32)>,
}
impl DraftDist {
pub fn deterministic(token: usize) -> Self {
DraftDist {
support: vec![(token, 1.0)],
}
}
pub fn from_dense(probs: &[f32]) -> Self {
DraftDist {
support: probs
.iter()
.enumerate()
.filter(|&(_, &p)| p > 0.0)
.map(|(i, &p)| (i, p))
.collect(),
}
}
pub fn from_support(support: Vec<(usize, f32)>) -> Self {
DraftDist { support }
}
pub fn support(&self) -> &[(usize, f32)] {
&self.support
}
pub fn prob(&self, token: usize) -> f32 {
self.support
.iter()
.find(|&&(t, _)| t == token)
.map(|&(_, p)| p)
.unwrap_or(0.0)
}
}
#[derive(Debug, Clone, Default, PartialEq)]
pub struct DraftBlock {
tokens: Vec<usize>,
dists: Vec<DraftDist>,
}
impl DraftBlock {
pub fn empty() -> Self {
DraftBlock::default()
}
pub fn new(tokens: Vec<usize>, dists: Vec<DraftDist>) -> Self {
assert_eq!(
tokens.len(),
dists.len(),
"a draft block needs one draft distribution per drafted token"
);
DraftBlock { tokens, dists }
}
pub fn deterministic(tokens: Vec<usize>) -> Self {
let dists = tokens
.iter()
.map(|&t| DraftDist::deterministic(t))
.collect();
DraftBlock { tokens, dists }
}
pub fn tokens(&self) -> &[usize] {
&self.tokens
}
pub fn dists(&self) -> &[DraftDist] {
&self.dists
}
pub fn len(&self) -> usize {
self.tokens.len()
}
pub fn is_empty(&self) -> bool {
self.tokens.is_empty()
}
pub fn truncate(&mut self, len: usize) {
self.tokens.truncate(len);
self.dists.truncate(len);
}
}
pub trait Drafter {
fn propose(&self, history: &[usize], target_hidden: &[f32], max_len: usize) -> DraftBlock;
}
#[derive(Debug, Clone, Copy)]
pub struct PromptLookupSpeculator {
pub ngram_size: usize,
pub max_draft_len: usize,
}
impl PromptLookupSpeculator {
pub fn new(ngram_size: usize, max_draft_len: usize) -> Self {
assert!(ngram_size >= 1, "ngram_size must be at least 1");
assert!(max_draft_len >= 1, "max_draft_len must be at least 1");
PromptLookupSpeculator {
ngram_size,
max_draft_len,
}
}
pub fn propose_tokens(&self, history: &[usize]) -> Vec<usize> {
if history.len() < self.ngram_size + 1 {
return Vec::new();
}
let needle = &history[history.len() - self.ngram_size..];
let last_possible_start = history.len() - self.ngram_size - 1;
for start in (0..=last_possible_start).rev() {
if &history[start..start + self.ngram_size] == needle {
let continuation_start = start + self.ngram_size;
let available = history.len() - continuation_start;
let take = available.min(self.max_draft_len);
return history[continuation_start..continuation_start + take].to_vec();
}
}
Vec::new()
}
}
impl Drafter for PromptLookupSpeculator {
fn propose(&self, history: &[usize], _target_hidden: &[f32], max_len: usize) -> DraftBlock {
let mut tokens = self.propose_tokens(history);
tokens.truncate(max_len.min(self.max_draft_len));
DraftBlock::deterministic(tokens)
}
}
pub fn accept_or_resample(
target: &[f32],
draft: &DraftDist,
token: usize,
rng: &mut Sampler,
) -> Option<usize> {
let p = target.get(token).copied().unwrap_or(0.0);
let q = draft.prob(token);
if q <= 0.0 || p >= q {
return None;
}
if rng.uniform() < p / q {
return None;
}
let mut residual = target.to_vec();
for &(t, qt) in draft.support() {
if let Some(r) = residual.get_mut(t) {
*r = (*r - qt).max(0.0);
}
}
let total: f32 = residual.iter().sum();
if total <= 0.0 {
return Some(rng.sample_from(target));
}
for r in residual.iter_mut() {
*r /= total;
}
Some(rng.sample_from(&residual))
}
#[derive(Debug, Clone, Default)]
pub struct SpeculativeOptions {
pub max_new_tokens: usize,
pub start_pos: usize,
pub sampling: SamplingParams,
pub seed: u64,
}
#[derive(Debug, Clone, Default)]
pub struct SpeculativeDecodeResult {
pub generated_tokens: Vec<usize>,
pub forward_calls: usize,
pub tokens_generated: usize,
pub verification_steps: usize,
pub drafted_tokens: usize,
pub accepted_tokens: usize,
pub evaluated_at_position: Vec<usize>,
pub accepted_at_position: Vec<usize>,
}
impl SpeculativeDecodeResult {
pub fn tokens_per_call(&self) -> f64 {
if self.forward_calls == 0 {
0.0
} else {
self.tokens_generated as f64 / self.forward_calls as f64
}
}
pub fn acceptance_length(&self) -> Option<f64> {
if self.verification_steps == 0 {
None
} else {
Some(self.tokens_generated as f64 / self.verification_steps as f64)
}
}
pub fn accept_rate(&self) -> Option<f64> {
if self.drafted_tokens == 0 {
None
} else {
Some(self.accepted_tokens as f64 / self.drafted_tokens as f64)
}
}
pub fn accept_rate_per_position(&self) -> Vec<f64> {
self.evaluated_at_position
.iter()
.zip(self.accepted_at_position.iter())
.map(|(&seen, &ok)| {
if seen == 0 {
0.0
} else {
ok as f64 / seen as f64
}
})
.collect()
}
fn record_position(&mut self, position: usize, accepted: bool) {
if self.evaluated_at_position.len() <= position {
self.evaluated_at_position.resize(position + 1, 0);
self.accepted_at_position.resize(position + 1, 0);
}
self.evaluated_at_position[position] += 1;
self.drafted_tokens += 1;
if accepted {
self.accepted_at_position[position] += 1;
self.accepted_tokens += 1;
}
}
}
pub fn speculative_decode<D: Drafter + ?Sized>(
decoder: &Decoder,
prompt_tokens: &[usize],
max_new_tokens: usize,
kv_caches: &mut [KvCache],
drafter: &D,
) -> SpeculativeDecodeResult {
speculative_decode_with(
decoder,
prompt_tokens,
kv_caches,
drafter,
&SpeculativeOptions {
max_new_tokens,
..SpeculativeOptions::default()
},
)
}
pub fn speculative_decode_with<D: Drafter + ?Sized>(
decoder: &Decoder,
prompt_tokens: &[usize],
kv_caches: &mut [KvCache],
drafter: &D,
options: &SpeculativeOptions,
) -> SpeculativeDecodeResult {
assert!(!prompt_tokens.is_empty(), "prompt must not be empty");
for cache in kv_caches.iter() {
assert_eq!(
cache.seq_len, options.start_pos,
"start_pos must be the caches' current length: they hold exactly the \
context preceding the prompt"
);
}
let mut result = SpeculativeDecodeResult::default();
if options.max_new_tokens == 0 {
return result;
}
let mut rng = Sampler::new(options.seed);
let mut history: Vec<usize> = prompt_tokens.to_vec();
let mut generated: Vec<usize> = Vec::with_capacity(options.max_new_tokens);
let (prefill_logits, prefill_hidden) =
decoder.forward_batch_with_hidden(prompt_tokens, options.start_pos, kv_caches);
result.forward_calls += 1;
let last = prefill_logits
.last()
.expect("prompt_tokens is non-empty, so forward_batch returns at least one logits vector");
let mut target_hidden = prefill_hidden.last().cloned().unwrap_or_default();
let mut pending = {
let probs = sampling_distribution(last, &options.sampling, &history);
rng.sample_from(&probs)
};
let mut pos = options.start_pos + prompt_tokens.len();
loop {
generated.push(pending);
history.push(pending);
if generated.len() == options.max_new_tokens {
break;
}
let draft_budget = options.max_new_tokens - generated.len() - 1;
let mut draft = drafter.propose(&history, &target_hidden, draft_budget);
draft.truncate(draft_budget);
let mut batch = Vec::with_capacity(1 + draft.len());
batch.push(pending);
batch.extend_from_slice(draft.tokens());
let (batch_logits, batch_hidden) =
decoder.forward_batch_with_hidden(&batch, pos, kv_caches);
result.forward_calls += 1;
result.verification_steps += 1;
let mut accepted = 0usize;
let mut replacement: Option<usize> = None;
for (i, (&token, dist)) in draft.tokens().iter().zip(draft.dists()).enumerate() {
let target = sampling_distribution(&batch_logits[i], &options.sampling, &history);
match accept_or_resample(&target, dist, token, &mut rng) {
None => {
result.record_position(i, true);
accepted += 1;
history.push(token);
generated.push(token);
}
Some(resampled) => {
result.record_position(i, false);
replacement = Some(resampled);
break;
}
}
}
let committed_len = pos + 1 + accepted;
if accepted < draft.len() {
for cache in kv_caches.iter_mut() {
cache.truncate(committed_len);
}
}
debug_assert!(kv_caches.iter().all(|c| c.seq_len == committed_len));
target_hidden = batch_hidden[accepted].clone();
pending = match replacement {
Some(tok) => tok,
None => {
let probs =
sampling_distribution(&batch_logits[accepted], &options.sampling, &history);
rng.sample_from(&probs)
}
};
pos = committed_len;
debug_assert!(generated.len() < options.max_new_tokens);
}
debug_assert_eq!(generated.len(), options.max_new_tokens);
result.tokens_generated = generated.len();
result.generated_tokens = generated;
result
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::glm_5_2;
use crate::ModelConfig;
use std::cell::RefCell;
fn tiny_test_config() -> ModelConfig {
let mut cfg = glm_5_2();
cfg.hidden_dim = 16;
cfg.n_heads = 4;
cfg.n_kv_heads = 2;
cfg.head_dim = 4;
cfg.moe.hidden_dim = 16;
cfg.moe.n_experts = 6;
cfg.moe.n_experts_active = 2;
cfg.moe.n_shared_experts = 1;
cfg.moe.expert_ffn_dim = 8;
cfg
}
fn caches(decoder: &Decoder) -> Vec<KvCache> {
(0..decoder.layers.len())
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect()
}
fn argmax(logits: &[f32]) -> usize {
logits
.iter()
.enumerate()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
.map(|(i, _)| i)
.unwrap_or(0)
}
struct FixedDrafter {
block: DraftBlock,
seen_history: RefCell<Vec<Vec<usize>>>,
seen_hidden_len: RefCell<Vec<usize>>,
}
impl FixedDrafter {
fn new(block: DraftBlock) -> Self {
FixedDrafter {
block,
seen_history: RefCell::new(Vec::new()),
seen_hidden_len: RefCell::new(Vec::new()),
}
}
}
impl Drafter for FixedDrafter {
fn propose(&self, history: &[usize], target_hidden: &[f32], max_len: usize) -> DraftBlock {
self.seen_history.borrow_mut().push(history.to_vec());
self.seen_hidden_len.borrow_mut().push(target_hidden.len());
let mut block = self.block.clone();
block.truncate(max_len);
block
}
}
#[test]
fn proposes_the_continuation_after_a_real_repeat() {
let spec = PromptLookupSpeculator::new(2, 4);
let history = vec![1, 2, 3, 4, 5, 9, 9, 9, 1, 2];
assert_eq!(spec.propose_tokens(&history), vec![3, 4, 5, 9]);
assert_eq!(
spec.propose(&history, &[], 8),
DraftBlock::deterministic(vec![3, 4, 5, 9])
);
}
#[test]
fn respects_max_draft_len() {
let spec = PromptLookupSpeculator::new(2, 2);
let history = vec![1, 2, 3, 4, 5, 6, 7, 1, 2];
assert_eq!(spec.propose_tokens(&history), vec![3, 4]);
}
#[test]
fn returns_empty_when_no_earlier_match_exists() {
let spec = PromptLookupSpeculator::new(2, 4);
let history = vec![1, 2, 3, 4, 5];
assert_eq!(spec.propose_tokens(&history), Vec::<usize>::new());
}
#[test]
fn returns_empty_when_history_too_short() {
let spec = PromptLookupSpeculator::new(3, 4);
let history = vec![1, 2, 3];
assert_eq!(spec.propose_tokens(&history), Vec::<usize>::new());
}
#[test]
fn finds_the_most_recent_match_when_several_exist() {
let spec = PromptLookupSpeculator::new(1, 3);
let history = vec![9, 8, 7, 6, 9, 5, 4, 9];
assert_eq!(spec.propose_tokens(&history), vec![5, 4, 9]);
}
#[test]
fn the_trait_caps_a_block_at_the_callers_budget() {
let spec = PromptLookupSpeculator::new(2, 4);
let history = vec![1, 2, 3, 4, 5, 9, 9, 9, 1, 2];
let block = spec.propose(&history, &[], 2);
assert_eq!(block.tokens(), &[3, 4]);
assert_eq!(block.dists().len(), 2);
}
#[test]
fn a_draft_at_least_as_likely_under_the_target_is_always_accepted() {
let target = vec![0.6f32, 0.3, 0.1];
let draft = DraftDist::from_dense(&[0.5, 0.4, 0.1]);
let mut rng = Sampler::new(1);
for _ in 0..100 {
assert_eq!(accept_or_resample(&target, &draft, 0, &mut rng), None);
}
}
#[test]
fn a_draft_the_target_rules_out_is_always_rejected() {
let target = vec![0.5f32, 0.5, 0.0];
let draft = DraftDist::deterministic(2);
let mut rng = Sampler::new(2);
for _ in 0..50 {
let replacement = accept_or_resample(&target, &draft, 2, &mut rng);
let tok = replacement.expect("p(2) = 0 means token 2 can never be accepted");
assert!(tok < 2, "residual must never resample the rejected token");
}
}
#[test]
fn resampling_reproduces_the_target_distribution() {
let target = vec![0.30f32, 0.25, 0.20, 0.15, 0.07, 0.03];
let drafts = [
DraftDist::from_dense(&[0.02, 0.03, 0.05, 0.10, 0.30, 0.50]),
DraftDist::deterministic(3),
DraftDist::from_support(vec![(0, 0.5), (5, 0.5)]),
DraftDist::from_dense(&target),
];
let draws = 200_000;
for (d, draft) in drafts.iter().enumerate() {
let mut rng = Sampler::new(0xA11CE + d as u64);
let mut counts = vec![0usize; target.len()];
for _ in 0..draws {
let dense = {
let mut v = vec![0.0f32; target.len()];
for &(t, p) in draft.support() {
v[t] = p;
}
v
};
let x = rng.sample_from(&dense);
let out = accept_or_resample(&target, draft, x, &mut rng).unwrap_or(x);
counts[out] += 1;
}
let tv: f64 = counts
.iter()
.enumerate()
.map(|(i, &c)| (c as f64 / draws as f64 - target[i] as f64).abs())
.sum::<f64>()
/ 2.0;
assert!(
tv < 0.01,
"draft {d}: speculative output distribution differs from the target \
(total variation {tv:.4}); counts={counts:?}"
);
}
}
#[test]
fn speculative_decode_matches_greedy_token_for_token() {
let cfg = tiny_test_config();
let vocab = 8;
let prompt = vec![1usize, 2, 3, 4, 1, 2];
let max_new = 6;
let decoder_a = Decoder::new_random_small(cfg.clone(), 2, vocab);
let mut caches_a = caches(&decoder_a);
let speculator = PromptLookupSpeculator::new(2, 3);
let result = speculative_decode(&decoder_a, &prompt, max_new, &mut caches_a, &speculator);
let decoder_b = Decoder::new_random_small(cfg, 2, vocab);
let mut caches_b = caches(&decoder_b);
let mut pending = decoder_b
.forward_batch(&prompt, 0, &mut caches_b)
.pop()
.unwrap();
let mut greedy = Vec::with_capacity(max_new);
for pos in (prompt.len()..).take(max_new) {
let tok = argmax(&pending);
greedy.push(tok);
pending = decoder_b.forward_token(tok, pos, &mut caches_b);
}
assert_eq!(
result.generated_tokens, greedy,
"speculative decode must produce exactly the same tokens as plain greedy decode"
);
}
fn exact_marginals(
decoder: &Decoder,
prompt: &[usize],
params: &SamplingParams,
depth: usize,
vocab: usize,
) -> Vec<Vec<f64>> {
#[allow(clippy::too_many_arguments)]
fn walk(
decoder: &Decoder,
kv: &[KvCache],
logits: &[f32],
history: &mut Vec<usize>,
weight: f64,
level: usize,
depth: usize,
pos: usize,
params: &SamplingParams,
marginals: &mut [Vec<f64>],
) {
let probs = sampling_distribution(logits, params, history);
for (token, &p) in probs.iter().enumerate() {
if p <= 0.0 {
continue;
}
marginals[level][token] += weight * p as f64;
if level + 1 == depth {
continue;
}
let mut branch: Vec<KvCache> = kv.to_vec();
let next = decoder.forward_token(token, pos, &mut branch);
history.push(token);
walk(
decoder,
&branch,
&next,
history,
weight * p as f64,
level + 1,
depth,
pos + 1,
params,
marginals,
);
history.pop();
}
}
let mut marginals = vec![vec![0.0f64; vocab]; depth];
let mut kv = caches(decoder);
let logits = decoder.forward_batch(prompt, 0, &mut kv).pop().unwrap();
let mut history = prompt.to_vec();
walk(
decoder,
&kv,
&logits,
&mut history,
1.0,
0,
depth,
prompt.len(),
params,
&mut marginals,
);
marginals
}
#[test]
fn speculative_decode_at_temperature_matches_plain_sampling() {
let cfg = tiny_test_config();
let vocab = 6;
let prompt = vec![1usize, 2, 3, 1, 2];
let max_new = 3;
let params = SamplingParams {
temperature: 1.0,
..SamplingParams::default()
};
let seeds = 4_000u64;
let decoder = Decoder::new_random_small(cfg, 1, vocab);
let speculator = PromptLookupSpeculator::new(2, 3);
let exact = exact_marginals(&decoder, &prompt, ¶ms, max_new, vocab);
let mut spec_counts = vec![vec![0usize; vocab]; max_new];
for seed in 0..seeds {
let mut kv = caches(&decoder);
let out = speculative_decode_with(
&decoder,
&prompt,
&mut kv,
&speculator,
&SpeculativeOptions {
max_new_tokens: max_new,
sampling: params.clone(),
seed,
..SpeculativeOptions::default()
},
);
for (i, &t) in out.generated_tokens.iter().enumerate() {
spec_counts[i][t] += 1;
}
}
for i in 0..max_new {
let tv: f64 = (0..vocab)
.map(|t| (spec_counts[i][t] as f64 / seeds as f64 - exact[i][t]).abs())
.sum::<f64>()
/ 2.0;
assert!(
tv < 0.03,
"position {i}: speculative sampling drifted from the target's own \
distribution (total variation {tv:.4})\n speculative = {:?}\n exact = {:?}",
spec_counts[i]
.iter()
.map(|&c| c as f64 / seeds as f64)
.collect::<Vec<_>>(),
exact[i]
);
}
}
#[test]
fn speculative_decode_saves_real_calls_when_drafts_hit() {
let cfg = tiny_test_config();
let vocab = 8;
let prompt = vec![1usize, 2, 3, 1, 2];
let max_new = 8;
let decoder = Decoder::new_random_small(cfg, 2, vocab);
let mut kv = caches(&decoder);
let speculator = PromptLookupSpeculator::new(2, 4);
let result = speculative_decode(&decoder, &prompt, max_new, &mut kv, &speculator);
assert_eq!(result.tokens_generated, max_new);
assert!(
result.forward_calls <= max_new,
"speculative decode must never need MORE forward_batch calls than plain \
sequential decode would (calls={}, tokens={})",
result.forward_calls,
max_new
);
}
#[test]
fn speculative_decode_with_no_repeats_falls_back_to_one_token_per_call() {
let cfg = tiny_test_config();
let vocab = 8;
let prompt = vec![1usize, 2, 3];
let max_new = 5;
let decoder = Decoder::new_random_small(cfg, 2, vocab);
let mut kv = caches(&decoder);
let speculator = PromptLookupSpeculator::new(10, 4); let result = speculative_decode(&decoder, &prompt, max_new, &mut kv, &speculator);
assert_eq!(result.tokens_generated, max_new);
assert_eq!(
result.forward_calls,
1 + max_new - 1,
"prefill (1 call) + one call per token, minus the last token, whose KV is \
never needed because generation stopped"
);
assert_eq!(result.drafted_tokens, 0);
assert_eq!(result.accept_rate(), None);
}
#[test]
fn the_drafter_is_asked_to_continue_the_anchor_token() {
let cfg = tiny_test_config();
let decoder = Decoder::new_random_small(cfg, 2, 8);
let mut kv = caches(&decoder);
let prompt = vec![1usize, 2, 3];
let drafter = FixedDrafter::new(DraftBlock::deterministic(vec![5, 6]));
let result = speculative_decode(&decoder, &prompt, 4, &mut kv, &drafter);
let seen = drafter.seen_history.borrow();
assert!(!seen.is_empty(), "the drafter must actually be consulted");
for (round, history) in seen.iter().enumerate() {
assert_eq!(
history.len(),
prompt.len() + round + 1,
"round {round}: history must grow by the committed tokens"
);
assert_eq!(
history[..prompt.len()],
prompt[..],
"the prompt must stay at the front of the drafter's history"
);
}
assert_eq!(seen[0][prompt.len()], result.generated_tokens[0]);
}
#[test]
fn the_drafter_receives_the_targets_hidden_state() {
let cfg = tiny_test_config();
let hidden_dim = cfg.hidden_dim;
let decoder = Decoder::new_random_small(cfg, 2, 8);
let mut kv = caches(&decoder);
let drafter = FixedDrafter::new(DraftBlock::deterministic(vec![5, 6]));
speculative_decode(&decoder, &[1usize, 2, 3], 4, &mut kv, &drafter);
let lens = drafter.seen_hidden_len.borrow();
assert!(!lens.is_empty());
for len in lens.iter() {
assert_eq!(
*len, hidden_dim,
"every round must pass a full target hidden state, not an empty slice"
);
}
}
#[test]
fn resuming_a_warm_cache_gives_the_same_tokens_as_one_fresh_run() {
let cfg = tiny_test_config();
let vocab = 8;
let decoder = Decoder::new_random_small(cfg, 2, vocab);
let speculator = PromptLookupSpeculator::new(2, 3);
let full_prompt = vec![1usize, 2, 3, 4, 1, 2];
let max_new = 6;
let mut fresh = caches(&decoder);
let cold = speculative_decode(&decoder, &full_prompt, max_new, &mut fresh, &speculator);
let split = 4;
let mut warm = caches(&decoder);
decoder.forward_batch(&full_prompt[..split], 0, &mut warm);
let resumed = speculative_decode_with(
&decoder,
&full_prompt[split..],
&mut warm,
&speculator,
&SpeculativeOptions {
max_new_tokens: max_new,
start_pos: split,
..SpeculativeOptions::default()
},
);
assert_eq!(
resumed.generated_tokens, cold.generated_tokens,
"resuming a warm cache must not change the output"
);
}
#[test]
fn rolls_back_to_absolute_positions_on_a_warm_cache() {
let cfg = tiny_test_config();
let decoder = Decoder::new_random_small(cfg, 2, 8);
let drafter = FixedDrafter::new(DraftBlock::deterministic(vec![7, 7, 7]));
let context = vec![1usize, 2, 3, 4];
let prompt = vec![5usize, 6];
let max_new = 6;
let mut kv = caches(&decoder);
decoder.forward_batch(&context, 0, &mut kv);
assert_eq!(kv[0].seq_len, context.len());
let result = speculative_decode_with(
&decoder,
&prompt,
&mut kv,
&drafter,
&SpeculativeOptions {
max_new_tokens: max_new,
start_pos: context.len(),
..SpeculativeOptions::default()
},
);
assert_eq!(result.tokens_generated, max_new);
let expected = context.len() + prompt.len() + result.tokens_generated - 1;
for cache in kv.iter() {
assert_eq!(
cache.seq_len,
expected,
"cache length must be absolute: context {} + prompt {} + generated {} - 1",
context.len(),
prompt.len(),
result.tokens_generated
);
}
}
#[test]
fn a_resumed_run_continues_a_previous_one() {
let cfg = tiny_test_config();
let decoder = Decoder::new_random_small(cfg, 2, 8);
let speculator = PromptLookupSpeculator::new(2, 3);
let prompt = vec![1usize, 2, 3, 4, 1, 2];
let mut one = caches(&decoder);
let long = speculative_decode(&decoder, &prompt, 8, &mut one, &speculator);
let mut kv = caches(&decoder);
let first = speculative_decode(&decoder, &prompt, 4, &mut kv, &speculator);
let resume_prompt = vec![*first.generated_tokens.last().unwrap()];
let start = prompt.len() + first.tokens_generated - 1;
let second = speculative_decode_with(
&decoder,
&resume_prompt,
&mut kv,
&speculator,
&SpeculativeOptions {
max_new_tokens: 5,
start_pos: start,
..SpeculativeOptions::default()
},
);
let mut stitched = first.generated_tokens.clone();
stitched.pop(); stitched.extend_from_slice(&second.generated_tokens);
assert_eq!(
&stitched[..8],
&long.generated_tokens[..],
"a decode split across two calls must equal the same decode in one"
);
}
#[test]
#[should_panic(expected = "start_pos must be the caches' current length")]
fn a_mismatched_start_pos_is_refused_rather_than_silently_wrong() {
let cfg = tiny_test_config();
let decoder = Decoder::new_random_small(cfg, 2, 8);
let mut kv = caches(&decoder);
decoder.forward_batch(&[1usize, 2, 3], 0, &mut kv);
let speculator = PromptLookupSpeculator::new(2, 2);
speculative_decode(&decoder, &[4usize, 5], 2, &mut kv, &speculator);
}
#[test]
fn per_position_accept_rates_expose_suffix_decay() {
let mut result = SpeculativeDecodeResult::default();
for _ in 0..100 {
result.record_position(0, true);
result.record_position(1, false);
}
result.verification_steps = 100;
result.tokens_generated = 200;
assert_eq!(result.accept_rate(), Some(0.5));
assert_eq!(result.accept_rate_per_position(), vec![1.0, 0.0]);
assert_eq!(result.acceptance_length(), Some(2.0));
}
#[test]
fn positions_after_a_rejection_are_not_counted_as_drafted() {
let cfg = tiny_test_config();
let decoder = Decoder::new_random_small(cfg, 2, 8);
let mut kv = caches(&decoder);
let drafter = FixedDrafter::new(DraftBlock::deterministic(vec![7, 7, 7, 7]));
let result = speculative_decode(&decoder, &[1usize, 2, 3], 6, &mut kv, &drafter);
let evaluated = &result.evaluated_at_position;
let accepted = &result.accepted_at_position;
assert!(
result.drafted_tokens > result.accepted_tokens,
"the scenario is pointless unless something was actually rejected \
(drafted {}, accepted {})",
result.drafted_tokens,
result.accepted_tokens
);
for (i, &seen) in evaluated.iter().enumerate().skip(1) {
assert!(
seen <= accepted[i - 1],
"position {i} was evaluated {seen} times but position {} was only \
accepted {} times: evaluated={evaluated:?} accepted={accepted:?}",
i - 1,
accepted[i - 1]
);
}
assert_eq!(
result.drafted_tokens,
evaluated.iter().sum::<usize>(),
"drafted_tokens must be the per-position counts' total"
);
assert_eq!(
result.accepted_tokens,
result.accepted_at_position.iter().sum::<usize>()
);
assert!(result.accepted_tokens <= result.drafted_tokens);
}
#[test]
fn acceptance_length_is_reported_per_verification_step_not_per_call() {
let cfg = tiny_test_config();
let decoder = Decoder::new_random_small(cfg, 2, 8);
let mut kv = caches(&decoder);
let speculator = PromptLookupSpeculator::new(2, 3);
let result = speculative_decode(&decoder, &[1usize, 2, 3, 1, 2], 6, &mut kv, &speculator);
assert_eq!(result.verification_steps, result.forward_calls - 1);
let length = result.acceptance_length().unwrap();
assert!(length >= 1.0, "every verification step commits >= 1 token");
assert!(
length > result.tokens_per_call(),
"acceptance length must not be diluted by the prefill call"
);
}
#[test]
fn an_empty_run_reports_no_acceptance_length_rather_than_zero() {
let result = SpeculativeDecodeResult::default();
assert_eq!(result.acceptance_length(), None);
assert_eq!(result.accept_rate(), None);
assert_eq!(result.tokens_per_call(), 0.0);
}
}