pub trait Drafter: Send + Sync {
fn clone_drafter(&self) -> Box<dyn Drafter>;
fn reset(&mut self);
fn draft(&mut self, tokens: &[u32], max_k: usize) -> Vec<u32>;
fn suggested_k(&self) -> Option<usize> {
None
}
}
#[derive(Debug, Clone, Copy)]
pub struct PromptLookupDrafter {
pub ngram: usize,
}
impl PromptLookupDrafter {
pub fn new(ngram: usize) -> Self {
Self { ngram }
}
}
impl Drafter for PromptLookupDrafter {
fn clone_drafter(&self) -> Box<dyn Drafter> {
Box::new(*self)
}
fn reset(&mut self) {}
fn draft(&mut self, tokens: &[u32], max_k: usize) -> Vec<u32> {
prompt_lookup_draft(tokens, self.ngram, max_k)
}
}
pub fn prompt_lookup_draft(tokens: &[u32], ngram: usize, k: usize) -> Vec<u32> {
let n = tokens.len();
if ngram == 0 || k == 0 || n <= ngram {
return Vec::new();
}
let pattern = &tokens[n - ngram..];
for start in (0..n - ngram).rev() {
if &tokens[start..start + ngram] == pattern {
let follow = start + ngram;
let end = (follow + k).min(n);
return tokens[follow..end].to_vec();
}
}
Vec::new()
}
#[derive(Debug, Default, Clone, Copy)]
pub struct SpecStats {
pub drafted: usize,
pub accepted: usize,
pub rounds: usize,
}
impl SpecStats {
pub fn acceptance_rate(&self) -> f32 {
if self.drafted == 0 {
0.0
} else {
self.accepted as f32 / self.drafted as f32
}
}
}
pub struct VerifyResult {
pub accepted: Vec<u32>,
pub follow_logits: Vec<f32>,
}
pub fn verify_draft(
model: &dyn crate::model::Model,
state: &mut crate::kv_cache::InferenceState,
guaranteed: u32,
draft: &[u32],
vocab: usize,
) -> VerifyResult {
use crate::sampler::argmax;
let old = state.seq_len;
let mut batch = Vec::with_capacity(1 + draft.len());
batch.push(guaranteed);
batch.extend_from_slice(draft);
let all = model.forward_prefill_logits_all(&batch, old, state);
assert_eq!(
all.len(),
batch.len() * vocab,
"forward_prefill_logits_all returned {} logits for {} positions x vocab {vocab}; \
the row-major [n x vocab] contract is what lets `verify_draft` index row j",
all.len(),
batch.len()
);
let mut accepted = Vec::new();
for (j, &q) in draft.iter().enumerate() {
if argmax(&all[j * vocab..(j + 1) * vocab]) != q {
break;
}
accepted.push(q);
}
let m = accepted.len();
model.truncate_kv(state, old + 1 + m);
let follow_logits = all[m * vocab..(m + 1) * vocab].to_vec();
VerifyResult {
accepted,
follow_logits,
}
}
pub fn greedy_generate_spec(
model: &dyn crate::model::Model,
state: &mut crate::kv_cache::InferenceState,
prompt: &[u32],
max_new: usize,
eos: &[u32],
ngram: usize,
k: usize,
) -> (Vec<u32>, SpecStats) {
use crate::sampler::argmax;
let mut out: Vec<u32> = Vec::new();
let mut stats = SpecStats::default();
assert!(
model.supports_all_logits(),
"greedy_generate_spec requires forward_prefill_logits_all support"
);
assert_eq!(
state.seq_len, 0,
"greedy_generate_spec requires a fresh InferenceState (seq_len == 0); \
this one holds {} positions, so the prompt would prefill after them and \
every position would be offset",
state.seq_len
);
assert!(
!state.is_compressed(),
"greedy_generate_spec requires an uncompressed KV cache: verification \
rewinds the KV, which TurboQuant's packed layout cannot do"
);
if prompt.is_empty() || max_new == 0 {
return (out, stats);
}
let mut next_logits = model.forward_prefill(prompt, 0, state);
let vocab = next_logits.len();
let mut history: Vec<u32> = prompt.to_vec();
let is_eos = |t: u32| eos.contains(&t);
loop {
if out.len() >= max_new {
break;
}
let t = argmax(&next_logits);
if is_eos(t) {
break;
}
out.push(t);
history.push(t);
if out.len() >= max_new {
break;
}
let draft = prompt_lookup_draft(&history, ngram, k);
if draft.is_empty() {
next_logits = model.forward(&[t], state.seq_len, state);
continue;
}
let old = state.seq_len;
stats.rounds += 1;
stats.drafted += draft.len();
let vr = verify_draft(model, state, t, &draft, vocab);
let mut kept = 0usize;
let mut stopped = false;
for &q in &vr.accepted {
if out.len() >= max_new || is_eos(q) {
stopped = true;
break;
}
out.push(q);
history.push(q);
kept += 1;
stats.accepted += 1;
}
if kept < vr.accepted.len() {
model.truncate_kv(state, old + 1 + kept);
}
if stopped {
break;
}
next_logits = vr.follow_logits;
}
(out, stats)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn draft_predicts_repeated_continuation() {
let toks = [10u32, 20, 30, 10, 20];
assert_eq!(prompt_lookup_draft(&toks, 2, 3), vec![30, 10, 20]);
assert_eq!(prompt_lookup_draft(&toks, 2, 1), vec![30]);
}
#[test]
fn draft_uses_most_recent_match() {
let toks = [7u32, 8, 9, 7, 8, 7, 8];
assert_eq!(prompt_lookup_draft(&toks, 2, 2), vec![7, 8]);
}
#[test]
fn draft_empty_when_no_match() {
assert!(prompt_lookup_draft(&[1, 2, 3, 4, 5], 2, 3).is_empty());
}
#[test]
fn draft_empty_on_degenerate_args() {
assert!(prompt_lookup_draft(&[1, 2, 3], 0, 3).is_empty());
assert!(prompt_lookup_draft(&[1, 2, 3], 2, 0).is_empty());
assert!(prompt_lookup_draft(&[1, 2], 2, 3).is_empty()); assert!(prompt_lookup_draft(&[], 2, 3).is_empty());
}
#[test]
fn draft_respects_sequence_end() {
let toks = [1u32, 2, 9, 5, 1, 2];
assert_eq!(prompt_lookup_draft(&toks, 2, 5), vec![9, 5, 1, 2]);
}
mod rewind_goes_through_the_model {
use super::*;
use crate::kv_cache::InferenceState;
use crate::model::{BlockType, Model, ModelConfig};
use std::sync::Mutex;
const VOCAB: u32 = 8;
fn dense_config() -> ModelConfig {
let n_layers = 2;
ModelConfig {
architecture: "llama".into(),
n_layers,
hidden_size: 8,
intermediate_size: 16,
n_heads: 2,
n_kv_heads: 2,
head_dim: 4,
vocab_size: VOCAB as usize,
max_seq_len: 64,
rope_theta: 10_000.0,
rms_norm_eps: 1e-5,
block_types: vec![BlockType::Attention; n_layers],
conv_kernel_size: None,
kv_heads_per_layer: vec![2; n_layers],
scalars: crate::model::ScalarMultipliers::default(),
moe: None,
is_causal: true,
class_labels: Vec::new(),
}
}
struct CyclingStub {
config: ModelConfig,
rewinds: Mutex<Vec<usize>>,
}
impl CyclingStub {
fn new() -> Self {
Self {
config: dense_config(),
rewinds: Mutex::new(Vec::new()),
}
}
fn row(t: u32) -> Vec<f32> {
let mut v = vec![0.0f32; VOCAB as usize];
v[((t + 1) % VOCAB) as usize] = 1.0;
v
}
fn recorded(&self) -> Vec<usize> {
self.rewinds.lock().unwrap().clone()
}
}
impl Model for CyclingStub {
fn config(&self) -> &ModelConfig {
&self.config
}
fn forward(&self, tokens: &[u32], _pos: usize, state: &mut InferenceState) -> Vec<f32> {
state.seq_len += tokens.len();
Self::row(*tokens.last().unwrap())
}
fn forward_prefill_logits_all(
&self,
tokens: &[u32],
_start_pos: usize,
state: &mut InferenceState,
) -> Vec<f32> {
state.seq_len += tokens.len();
tokens.iter().flat_map(|&t| Self::row(t)).collect()
}
fn supports_all_logits(&self) -> bool {
true
}
fn truncate_kv(&self, state: &mut InferenceState, len: usize) {
self.rewinds.lock().unwrap().push(len);
state.truncate_to(len);
}
}
#[test]
fn verify_draft_rewinds_through_the_model() {
let model = CyclingStub::new();
let mut state = InferenceState::from_config(model.config()).unwrap();
model.forward_prefill(&[0, 1, 2], 0, &mut state);
let old = state.seq_len;
assert_eq!(old, 3);
let vr = verify_draft(&model, &mut state, 3, &[4, 5, 0], VOCAB as usize);
assert_eq!(vr.accepted, vec![4, 5]);
assert_eq!(
model.recorded(),
vec![old + 1 + 2],
"verify_draft must rewind via Model::truncate_kv, once, to the \
guaranteed token plus the accepted drafts"
);
assert_eq!(state.seq_len, old + 1 + 2);
}
#[test]
fn the_default_truncate_kv_really_rewinds() {
struct DefaultStub(ModelConfig);
impl Model for DefaultStub {
fn config(&self) -> &ModelConfig {
&self.0
}
fn forward(&self, _: &[u32], _: usize, _: &mut InferenceState) -> Vec<f32> {
unimplemented!("DefaultStub is driven through forward_prefill_logits_all")
}
fn forward_prefill_logits_all(
&self,
tokens: &[u32],
_start_pos: usize,
state: &mut InferenceState,
) -> Vec<f32> {
for layer in &mut state.layers {
if let crate::kv_cache::LayerState::Attention {
key_cache,
value_cache,
..
} = layer
{
for &t in tokens {
key_cache.push(t as f32);
value_cache.push(t as f32);
}
}
}
state.seq_len += tokens.len();
tokens.iter().flat_map(|&t| CyclingStub::row(t)).collect()
}
fn supports_all_logits(&self) -> bool {
true
}
}
let model = DefaultStub(dense_config());
let mut state = InferenceState::from_config(model.config()).unwrap();
model.forward_prefill_logits_all(&[0, 1, 2], 0, &mut state);
assert_eq!(state.seq_len, 3);
let vr = verify_draft(&model, &mut state, 3, &[4, 5, 0], VOCAB as usize);
assert_eq!(vr.accepted, vec![4, 5]);
assert_eq!(state.seq_len, 6);
for layer in &state.layers {
let crate::kv_cache::LayerState::Attention {
key_cache,
value_cache,
..
} = layer
else {
unreachable!("dense config has only attention layers")
};
assert_eq!(
key_cache,
&[0.0, 1.0, 2.0, 3.0, 4.0, 5.0],
"the default truncate_kv did not cut the rejected tail out of \
the key cache"
);
assert_eq!(value_cache.len(), 6);
}
}
#[test]
fn early_stop_rewind_goes_through_the_model() {
let model = CyclingStub::new();
let mut state = InferenceState::from_config(model.config()).unwrap();
let prompt: Vec<u32> = vec![0, 1, 2, 3, 4, 5, 6, 7, 0, 1];
let (out, stats) = greedy_generate_spec(&model, &mut state, &prompt, 3, &[], 2, 3);
assert_eq!(out, vec![2, 3, 4], "the stub decodes the cycle");
assert_eq!(stats.accepted, 2, "budget cut the third accepted draft");
let old = prompt.len();
assert_eq!(
model.recorded(),
vec![old + 1 + 3, old + 1 + 2],
"expected the per-round rewind (all 3 drafts accepted) followed \
by the early-stop rewind (only 2 kept), both through \
Model::truncate_kv"
);
assert_eq!(state.seq_len, old + 1 + 2);
}
}
}