#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ArgmaxCapture {
Last,
All,
}
#[derive(Debug, thiserror::Error)]
pub enum VerifierError {
#[error("verifier: empty input — at least 1 token required")]
EmptyInput,
#[error("verifier: too many tokens for verify pass (got {got}, max {max})")]
TooManyTokens { got: usize, max: usize },
#[error("verifier: model error: {0}")]
Model(#[from] anyhow::Error),
}
pub type VerifyLogits = Vec<Vec<f32>>;
pub trait Verifier {
fn verify(&mut self, tokens: &[u32]) -> Result<VerifyLogits, VerifierError>;
fn rollback_kv_to(&mut self, seq_pos: usize) -> Result<(), VerifierError>;
}
pub fn accept_prefix(drafts: &[u32], logits_per_pos: &VerifyLogits) -> (usize, u32) {
if logits_per_pos.is_empty() {
return (0, 0);
}
let mut accept_count = 0;
for i in 0..drafts.len() {
if i >= logits_per_pos.len() {
break;
}
let argmax = argmax_u32(&logits_per_pos[i]);
if argmax == drafts[i] {
accept_count += 1;
} else {
return (accept_count, argmax);
}
}
if accept_count < logits_per_pos.len() {
let model_token = argmax_u32(&logits_per_pos[accept_count]);
(accept_count, model_token)
} else {
let model_token = argmax_u32(logits_per_pos.last().unwrap());
(accept_count, model_token)
}
}
fn argmax_u32(logits: &[f32]) -> u32 {
let mut best_idx = 0u32;
let mut best_val = f32::MIN;
for (i, &v) in logits.iter().enumerate() {
if v > best_val {
best_val = v;
best_idx = i as u32;
}
}
best_idx
}
pub fn accept_prefix_argmax(drafts: &[u32], model_argmaxes: &[u32]) -> (usize, u32) {
if model_argmaxes.is_empty() {
return (0, 0);
}
let mut accept_count = 0;
for i in 0..drafts.len() {
if i >= model_argmaxes.len() {
break;
}
if model_argmaxes[i] == drafts[i] {
accept_count += 1;
} else {
return (accept_count, model_argmaxes[i]);
}
}
if accept_count < model_argmaxes.len() {
(accept_count, model_argmaxes[accept_count])
} else {
(accept_count, *model_argmaxes.last().unwrap())
}
}
pub fn rollback_kv_state(
write_pos: usize,
seq_len: usize,
capacity: usize,
is_sliding: bool,
trim: usize,
) -> (usize, usize) {
let trim = trim.min(seq_len);
let new_seq_len = seq_len - trim;
let new_write_pos = if is_sliding {
if capacity == 0 {
0
} else {
(write_pos + capacity - (trim % capacity)) % capacity
}
} else {
write_pos.saturating_sub(trim)
};
(new_write_pos, new_seq_len)
}
#[cfg(test)]
mod tests {
use super::*;
fn one_hot(vocab: usize, target: u32) -> Vec<f32> {
let mut v = vec![0.0_f32; vocab];
v[target as usize] = 1.0;
v
}
#[test]
fn accept_prefix_full_accept() {
let drafts = vec![10u32, 20, 30];
let logits = vec![
one_hot(100, 10),
one_hot(100, 20),
one_hot(100, 30),
one_hot(100, 40), ];
let (accept, tok) = accept_prefix(&drafts, &logits);
assert_eq!(accept, 3, "all 3 drafts should be accepted");
assert_eq!(tok, 40, "model_token should be the K+1th argmax");
}
#[test]
fn accept_prefix_partial_accept() {
let drafts = vec![10u32, 20, 30];
let logits = vec![
one_hot(100, 10),
one_hot(100, 20),
one_hot(100, 99),
one_hot(100, 40),
];
let (accept, tok) = accept_prefix(&drafts, &logits);
assert_eq!(accept, 2, "first 2 accepted, 3rd mismatches");
assert_eq!(tok, 99, "model_token = model's argmax at mismatch position");
}
#[test]
fn accept_prefix_zero_accept() {
let drafts = vec![10u32, 20, 30];
let logits = vec![
one_hot(100, 99),
one_hot(100, 20),
one_hot(100, 30),
one_hot(100, 40),
];
let (accept, tok) = accept_prefix(&drafts, &logits);
assert_eq!(accept, 0);
assert_eq!(tok, 99);
}
#[test]
fn accept_prefix_empty_drafts_uses_first_logits() {
let drafts: Vec<u32> = Vec::new();
let logits = vec![one_hot(100, 42)];
let (accept, tok) = accept_prefix(&drafts, &logits);
assert_eq!(accept, 0);
assert_eq!(tok, 42, "model_token = argmax at position 0");
}
#[test]
fn accept_prefix_empty_logits_returns_zero_zero() {
let drafts = vec![10u32, 20];
let logits: VerifyLogits = Vec::new();
let (accept, tok) = accept_prefix(&drafts, &logits);
assert_eq!(accept, 0);
assert_eq!(tok, 0);
}
struct MockVerifier {
scripted: VerifyLogits,
expected_input_len: Option<usize>,
verify_inputs: Vec<Vec<u32>>,
rollbacks: Vec<usize>,
}
impl MockVerifier {
fn new(scripted: VerifyLogits) -> Self {
Self {
scripted,
expected_input_len: None,
verify_inputs: Vec::new(),
rollbacks: Vec::new(),
}
}
fn with_expected_input_len(mut self, n: usize) -> Self {
self.expected_input_len = Some(n);
self
}
}
impl Verifier for MockVerifier {
fn verify(&mut self, tokens: &[u32]) -> Result<VerifyLogits, VerifierError> {
if tokens.is_empty() {
return Err(VerifierError::EmptyInput);
}
if let Some(n) = self.expected_input_len {
assert_eq!(
tokens.len(),
n,
"MockVerifier: expected {n} input tokens, got {}",
tokens.len()
);
}
self.verify_inputs.push(tokens.to_vec());
Ok(self.scripted.clone())
}
fn rollback_kv_to(&mut self, seq_pos: usize) -> Result<(), VerifierError> {
self.rollbacks.push(seq_pos);
Ok(())
}
}
fn ground_truth_next(seq: &[u32], vocab: u32) -> u32 {
let s: u64 = seq.iter().map(|&t| t as u64).sum();
let last = *seq.last().unwrap_or(&0) as u64;
((s.wrapping_mul(31).wrapping_add(last)) % vocab as u64) as u32
}
struct GroundTruthVerifier {
vocab: u32,
prefix: Vec<u32>,
rollbacks: Vec<usize>,
}
impl GroundTruthVerifier {
fn new(vocab: u32, initial_prefix: Vec<u32>) -> Self {
Self {
vocab,
prefix: initial_prefix,
rollbacks: Vec::new(),
}
}
}
impl Verifier for GroundTruthVerifier {
fn verify(&mut self, tokens: &[u32]) -> Result<VerifyLogits, VerifierError> {
let mut logits = Vec::with_capacity(tokens.len());
for i in 0..tokens.len() {
let mut seq = self.prefix.clone();
if i > 0 {
seq.extend_from_slice(&tokens[1..=i]);
}
let next = ground_truth_next(&seq, self.vocab);
logits.push(one_hot(self.vocab as usize, next));
}
Ok(logits)
}
fn rollback_kv_to(&mut self, seq_pos: usize) -> Result<(), VerifierError> {
self.prefix.truncate(seq_pos);
self.rollbacks.push(seq_pos);
Ok(())
}
}
fn default_decode(prompt: &[u32], vocab: u32, n_tokens: usize) -> Vec<u32> {
let mut gen = prompt.to_vec();
for _ in 0..n_tokens {
let next = ground_truth_next(&gen, vocab);
gen.push(next);
}
gen
}
fn spec_decode_loop(
prompt: &[u32],
vocab: u32,
n_tokens: usize,
cfg: &super::super::ngram_proposer::NgramConfig,
) -> (Vec<u32>, GroundTruthVerifier) {
let mut gen = prompt.to_vec();
let mut verifier = GroundTruthVerifier::new(vocab, prompt.to_vec());
let target_len = prompt.len() + n_tokens;
while gen.len() < target_len {
let drafts = super::super::ngram_proposer::propose(&gen, cfg);
let last = *gen.last().unwrap();
let mut input = vec![last];
input.extend_from_slice(&drafts);
let logits = verifier.verify(&input).unwrap();
let (accept, model_tok) = accept_prefix(&drafts, &logits);
gen.extend_from_slice(&drafts[..accept]);
gen.push(model_tok);
verifier.prefix = gen.clone();
verifier.rollback_kv_to(gen.len()).unwrap();
if gen.len() >= target_len + cfg.k {
break; }
}
(gen, verifier)
}
#[test]
fn spec_decode_byte_identity_vs_default_decode() {
let prompt = vec![1u32, 2, 3, 1, 2, 3, 4]; let vocab = 256u32;
let n_tokens = 30;
let cfg = super::super::ngram_proposer::NgramConfig {
min_ngram: 1,
max_ngram: 3,
k: 3,
max_model_len: 4096,
};
let default_out = default_decode(&prompt, vocab, n_tokens);
let (spec_out, _v) = spec_decode_loop(&prompt, vocab, n_tokens, &cfg);
let cmp_len = prompt.len() + n_tokens;
assert!(spec_out.len() >= cmp_len);
assert_eq!(
&spec_out[..cmp_len], &default_out[..cmp_len],
"spec-decode must match default decode under greedy WHEN BOTH PATHS USE THE SAME VERIFIER (synthetic GroundTruthVerifier here; real-model claim is falsified — see ADR-034 G3)"
);
}
#[test]
fn spec_decode_at_k_zero_calls_verifier_with_single_token() {
let prompt = vec![5u32, 6, 7];
let vocab = 100u32;
let cfg = super::super::ngram_proposer::NgramConfig {
min_ngram: 1,
max_ngram: 3,
k: 0,
max_model_len: 4096,
};
let mut verifier = GroundTruthVerifier::new(vocab, prompt.clone());
let mut gen = prompt.clone();
for _ in 0..5 {
let drafts = super::super::ngram_proposer::propose(&gen, &cfg);
assert!(drafts.is_empty(), "K=0 must always return empty drafts");
let last = *gen.last().unwrap();
let logits = verifier.verify(&[last]).unwrap();
assert_eq!(logits.len(), 1, "K=0 verify produces 1 logits row");
let (accept, tok) = accept_prefix(&drafts, &logits);
assert_eq!(accept, 0);
gen.push(tok);
verifier.prefix = gen.clone();
verifier.rollback_kv_to(gen.len()).unwrap();
}
let default_out = default_decode(&prompt, vocab, 5);
assert_eq!(gen, default_out);
}
#[test]
fn accept_prefix_invariants_under_random_inputs() {
let vocab = 50usize;
let mut state: u64 = 0xCAFE_BEEF;
let next_rand = |s: &mut u64| -> u64 {
*s = s
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
*s >> 33
};
for _ in 0..500 {
let n_drafts = (next_rand(&mut state) % 6) as usize;
let n_logits = (next_rand(&mut state) % 8) as usize;
let drafts: Vec<u32> = (0..n_drafts)
.map(|_| (next_rand(&mut state) % vocab as u64) as u32)
.collect();
let logits: VerifyLogits = (0..n_logits)
.map(|_| {
let target = (next_rand(&mut state) % vocab as u64) as u32;
one_hot(vocab, target)
})
.collect();
let (accept, tok) = accept_prefix(&drafts, &logits);
assert!(accept <= drafts.len(), "accept_count > drafts.len()");
assert!(
accept <= logits.len() || logits.is_empty(),
"accept_count > logits.len() (got {accept} vs {})",
logits.len()
);
for i in 0..accept {
assert_eq!(drafts[i], argmax_u32(&logits[i]),
"accepted draft[{i}] = {} doesn't match argmax {} (logits len {}, drafts len {})",
drafts[i], argmax_u32(&logits[i]), logits.len(), drafts.len());
}
if accept < drafts.len() && accept < logits.len() {
assert_ne!(
drafts[accept],
argmax_u32(&logits[accept]),
"first rejected draft equals model argmax — should have been accepted"
);
}
assert!(
(tok as usize) < vocab,
"model_token {tok} out of vocab range {vocab}"
);
}
}
#[test]
fn spec_decode_loop_full_accept_advances_seq_pos_by_k_plus_1() {
let drafts = vec![10u32, 20, 30];
let scripted = vec![
one_hot(100, 10),
one_hot(100, 20),
one_hot(100, 30),
one_hot(100, 40),
];
let mut mock = MockVerifier::new(scripted).with_expected_input_len(4);
let seq_pos_before: usize = 7;
let last = 5u32;
let mut input = vec![last];
input.extend_from_slice(&drafts);
let logits = mock.verify(&input).unwrap();
let (accept, tok) = accept_prefix(&drafts, &logits);
let seq_pos_after = seq_pos_before + accept + 1;
assert_eq!(accept, 3);
assert_eq!(tok, 40);
assert_eq!(seq_pos_after, 11);
assert_eq!(
mock.verify_inputs,
vec![vec![5u32, 10, 20, 30]],
"verifier saw [last] ++ drafts in correct order"
);
mock.rollback_kv_to(seq_pos_after).unwrap();
assert_eq!(mock.rollbacks, vec![11]);
}
#[test]
fn spec_decode_loop_partial_accept_rolls_back_rejected() {
let drafts = vec![10u32, 20, 30];
let scripted = vec![
one_hot(100, 10),
one_hot(100, 20),
one_hot(100, 99),
one_hot(100, 40),
];
let mut mock = MockVerifier::new(scripted).with_expected_input_len(4);
let seq_pos_before: usize = 7;
let logits = mock.verify(&[5u32, 10, 20, 30]).unwrap();
let (accept, tok) = accept_prefix(&drafts, &logits);
let seq_pos_after = seq_pos_before + accept + 1;
assert_eq!(accept, 2);
assert_eq!(tok, 99);
assert_eq!(seq_pos_after, 10);
assert_eq!(mock.verify_inputs, vec![vec![5u32, 10, 20, 30]]);
mock.rollback_kv_to(seq_pos_after).unwrap();
assert_eq!(
mock.rollbacks,
vec![10],
"rollback to seq_pos_after = before + accept + 1"
);
}
#[test]
fn mock_verifier_rejects_empty_input_per_contract() {
let mut mock = MockVerifier::new(vec![one_hot(100, 0)]);
let result = mock.verify(&[]);
assert!(matches!(result, Err(VerifierError::EmptyInput)));
}
#[test]
fn mock_verifier_input_len_validation_catches_caller_bug() {
let mut mock = MockVerifier::new(vec![one_hot(100, 0)]).with_expected_input_len(4);
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
let _ = mock.verify(&[1u32, 2, 3]);
}));
assert!(
result.is_err(),
"MockVerifier with expected_input_len=4 should panic on 3 tokens"
);
}
#[test]
fn ground_truth_decode_is_deterministic() {
let seq = vec![1u32, 2, 3, 4, 5];
let a = ground_truth_next(&seq, 256);
let b = ground_truth_next(&seq, 256);
let c = ground_truth_next(&seq, 256);
assert_eq!(a, b);
assert_eq!(b, c);
let d = ground_truth_next(&[1u32, 2, 3, 4, 6], 256);
assert_ne!(a, d, "ground_truth_next must be input-sensitive");
}
#[test]
fn accept_prefix_argmax_full_accept() {
let drafts = vec![10u32, 20, 30];
let argmaxes = vec![10u32, 20, 30, 40];
let (accept, tok) = accept_prefix_argmax(&drafts, &argmaxes);
assert_eq!(accept, 3);
assert_eq!(tok, 40);
}
#[test]
fn accept_prefix_argmax_partial_accept() {
let drafts = vec![10u32, 20, 30];
let argmaxes = vec![10u32, 20, 99, 40];
let (accept, tok) = accept_prefix_argmax(&drafts, &argmaxes);
assert_eq!(accept, 2);
assert_eq!(tok, 99);
}
#[test]
fn accept_prefix_argmax_zero_accept() {
let drafts = vec![10u32, 20, 30];
let argmaxes = vec![99u32, 20, 30, 40];
let (accept, tok) = accept_prefix_argmax(&drafts, &argmaxes);
assert_eq!(accept, 0);
assert_eq!(tok, 99);
}
#[test]
fn accept_prefix_argmax_matches_logits_variant() {
let drafts = vec![10u32, 20, 30];
let argmaxes = vec![10u32, 20, 99, 40];
let logits: VerifyLogits = argmaxes.iter().map(|&t| one_hot(100, t)).collect();
let (a1, t1) = accept_prefix(&drafts, &logits);
let (a2, t2) = accept_prefix_argmax(&drafts, &argmaxes);
assert_eq!((a1, t1), (a2, t2));
}
#[test]
fn accept_prefix_argmax_empty() {
let (a, t) = accept_prefix_argmax(&[], &[]);
assert_eq!((a, t), (0, 0));
}
#[test]
fn rollback_full_attention_subtracts() {
let (wp, sl) = rollback_kv_state(100, 100, 4096, false, 3);
assert_eq!((wp, sl), (97, 97));
}
#[test]
fn rollback_full_attention_zero_trim() {
let (wp, sl) = rollback_kv_state(100, 100, 4096, false, 0);
assert_eq!((wp, sl), (100, 100));
}
#[test]
fn rollback_full_attention_clamps_at_zero() {
let (wp, sl) = rollback_kv_state(100, 100, 4096, false, 200);
assert_eq!((wp, sl), (0, 0));
}
#[test]
fn rollback_sliding_wraps_no_wrap() {
let (wp, sl) = rollback_kv_state(10, 50, 100, true, 3);
assert_eq!((wp, sl), (7, 47));
}
#[test]
fn rollback_sliding_wraps_through_zero() {
let (wp, sl) = rollback_kv_state(2, 10, 10, true, 3);
assert_eq!((wp, sl), (9, 7));
}
#[test]
fn rollback_sliding_wraps_full_circle() {
let (wp, sl) = rollback_kv_state(2, 10, 10, true, 10);
assert_eq!((wp, sl), (2, 0));
}
#[test]
fn rollback_sliding_zero_capacity_safe() {
let (wp, sl) = rollback_kv_state(0, 0, 0, true, 5);
assert_eq!((wp, sl), (0, 0));
}
#[test]
fn rollback_invariant_seq_len_le_capacity() {
for cap in [1usize, 8, 64, 256, 1024] {
for sl in 0..=cap {
for wp in 0..=cap.saturating_sub(1).max(0) {
for is_sliding in [false, true] {
for trim in [0usize, 1, sl, sl / 2, sl + 1] {
let (_nwp, nsl) = rollback_kv_state(wp, sl, cap, is_sliding, trim);
assert!(nsl <= cap, "seq_len > cap: cap={cap} wp={wp} sl={sl} sliding={is_sliding} trim={trim} → nsl={nsl}");
assert!(
nsl <= sl,
"seq_len grew: cap={cap} wp={wp} sl={sl} trim={trim} → nsl={nsl}"
);
}
}
}
}
}
}
}