use crate::decoder::Decoder;
use ferrox_core::cache::KvCache;
#[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(&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()
}
}
#[derive(Debug, Clone)]
pub struct SpeculativeDecodeResult {
pub generated_tokens: Vec<usize>,
pub forward_calls: usize,
pub tokens_generated: 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 speculative_decode(
decoder: &Decoder,
prompt_tokens: &[usize],
max_new_tokens: usize,
kv_caches: &mut [KvCache],
speculator: &PromptLookupSpeculator,
) -> SpeculativeDecodeResult {
assert!(!prompt_tokens.is_empty(), "prompt must not be empty");
let mut history: Vec<usize> = prompt_tokens.to_vec();
let mut generated = Vec::with_capacity(max_new_tokens);
let mut forward_calls = 0usize;
let prefill_logits = decoder.forward_batch(prompt_tokens, 0, kv_caches);
forward_calls += 1;
let mut pending_logits = prefill_logits
.last()
.expect("prompt_tokens is non-empty, so forward_batch returns at least one logits vector")
.clone();
let mut pos = prompt_tokens.len();
while generated.len() < max_new_tokens {
let real_tok = argmax(&pending_logits);
let remaining_budget = max_new_tokens - generated.len() - 1; let mut guesses = speculator.propose(&history);
guesses.truncate(remaining_budget);
if guesses.is_empty() {
let logits = decoder.forward_batch(&[real_tok], pos, kv_caches);
forward_calls += 1;
generated.push(real_tok);
history.push(real_tok);
pending_logits = logits.into_iter().next().unwrap();
pos += 1;
continue;
}
let mut batch = Vec::with_capacity(1 + guesses.len());
batch.push(real_tok);
batch.extend_from_slice(&guesses);
let batch_logits = decoder.forward_batch(&batch, pos, kv_caches);
forward_calls += 1;
let mut accepted_count = 0usize;
for (i, &guess) in guesses.iter().enumerate() {
if argmax(&batch_logits[i]) == guess {
accepted_count += 1;
} else {
break;
}
}
if accepted_count < guesses.len() {
let committed_len = pos + 1 + accepted_count;
for cache in kv_caches.iter_mut() {
cache.truncate(committed_len);
}
}
generated.push(real_tok);
generated.extend_from_slice(&guesses[..accepted_count]);
history.push(real_tok);
history.extend_from_slice(&guesses[..accepted_count]);
pending_logits = batch_logits[accepted_count].clone();
pos += 1 + accepted_count;
}
generated.truncate(max_new_tokens);
SpeculativeDecodeResult {
tokens_generated: generated.len(),
generated_tokens: generated,
forward_calls,
}
}
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)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::glm_5_2;
use crate::ModelConfig;
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
}
#[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(&history), 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(&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(&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(&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(&history), vec![5, 4, 9]);
}
#[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: Vec<KvCache> = (0..2)
.map(|_| KvCache::new(decoder_a.config.n_kv_heads, decoder_a.config.head_dim))
.collect();
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: Vec<KvCache> = (0..2)
.map(|_| KvCache::new(decoder_b.config.n_kv_heads, decoder_b.config.head_dim))
.collect();
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"
);
}
#[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 caches: Vec<KvCache> = (0..2)
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect();
let speculator = PromptLookupSpeculator::new(2, 4);
let result = speculative_decode(&decoder, &prompt, max_new, &mut caches, &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 caches: Vec<KvCache> = (0..2)
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect();
let speculator = PromptLookupSpeculator::new(10, 4); let result = speculative_decode(&decoder, &prompt, max_new, &mut caches, &speculator);
assert_eq!(result.tokens_generated, max_new);
assert_eq!(
result.forward_calls,
1 + max_new,
"prefill (1 call) + one call per token when nothing ever matches"
);
}
}