use super::sampler::{
self, DEFAULT_NO_REPEAT_NGRAM_SIZE, NGRAM_WINDOW_SINGLE, argmax_row,
masked_sliding_window_logits_if_needed,
};
use super::tensor::Mat;
use crate::error::FocrResult;
pub(crate) const SPEC_DRAFT_MAX: usize = 4;
pub(crate) const SPEC_DRAFT_NGRAM: usize = 3;
#[allow(dead_code)] fn greedy_from_row(
row: &[f32],
sequence: &[u32],
ngram_size: usize,
window: usize,
) -> FocrResult<u32> {
if ngram_size == 0 || sequence.len() < ngram_size {
return argmax_row(row);
}
if let Some(masked) =
masked_sliding_window_logits_if_needed(row, sequence, ngram_size, window, &[])
{
return argmax_row(&masked);
}
argmax_row(row)
}
#[allow(dead_code)] pub(crate) fn accept_longest(
history: &[u32],
draft: &[u32],
verify_logits: &[&[f32]],
eos_id: u32,
) -> usize {
let mut sequence: Vec<u32> = history.to_vec();
let mut accepted = 0usize;
for (&token, &row) in draft.iter().zip(verify_logits.iter()) {
let Ok(greedy_token) = greedy_from_row(
row,
&sequence,
DEFAULT_NO_REPEAT_NGRAM_SIZE,
NGRAM_WINDOW_SINGLE,
) else {
break;
};
if token != greedy_token {
break;
}
accepted += 1;
if eos_id == greedy_token {
break;
}
sequence.push(token);
}
accepted
}
pub(crate) struct RoundEmit {
pub accepted: usize,
pub correction: Option<sampler::DecodeOutput>,
}
pub(crate) fn resolve_round(
generated: &[u32],
draft: &[u32],
verify_logits: &[Mat],
params: &sampler::DecodeParams,
) -> FocrResult<RoundEmit> {
let rows: Vec<&[f32]> = verify_logits.iter().map(|m| m.row(0)).collect();
let accepted = accept_longest(generated, draft, &rows, params.eos_token_id);
if accepted > 0 && params.eos_token_id == draft[accepted - 1] {
return Ok(RoundEmit {
accepted,
correction: None,
});
}
let Some(correction_row) = verify_logits.get(accepted) else {
return Ok(RoundEmit {
accepted,
correction: None,
});
};
let mut context = generated.to_vec();
context.extend_from_slice(&draft[..accepted]);
let correction = sampler::decode_step(correction_row, &context, params)?;
Ok(RoundEmit {
accepted,
correction: Some(correction),
})
}
#[allow(dead_code)] pub(crate) fn draft_ngram(seq: &[u32], max_draft: usize, ngram: usize) -> Vec<u32> {
if ngram == 0 || max_draft == 0 || seq.len() <= ngram {
return Vec::new();
}
let n = seq.len();
let needle = &seq[n - ngram..];
let Some(start) = seq
.windows(ngram)
.take(n - ngram)
.rposition(|window| window == needle)
else {
return Vec::new();
};
let from = start + ngram;
let to = (from + max_draft).min(n);
seq[from..to].to_vec()
}
#[cfg(test)]
mod tests {
use super::{SPEC_DRAFT_MAX, SPEC_DRAFT_NGRAM, accept_longest, draft_ngram, resolve_round};
use crate::native_engine::sampler::{self, DecodeParams};
use crate::native_engine::tensor::Mat;
const V: usize = 64;
const EOS: u32 = 1;
fn row_peaked(peak: u32, runner_up: u32) -> Vec<f32> {
let mut r = vec![0.0f32; V];
r[peak as usize] = 10.0;
r[runner_up as usize] = 9.0;
r
}
fn row_argmax(peak: u32) -> Vec<f32> {
let mut r = vec![0.0f32; V];
r[peak as usize] = 8.0;
r
}
fn greedy(row: &[f32], history: &[u32]) -> u32 {
let m = Mat::from_vec(1, row.len(), row.to_vec());
sampler::sample(&m, history, &DecodeParams::single_image()).expect("greedy chooser")
}
fn ref_seq_greedy(history: &[u32], rows: &[&[f32]], eos: u32, max_steps: usize) -> Vec<u32> {
let mut seq = history.to_vec();
let mut out = Vec::new();
for &row in rows.iter().take(max_steps) {
let g = greedy(row, &seq);
out.push(g);
seq.push(g);
if g == eos {
break;
}
}
out
}
fn spec_stream(history: &[u32], draft: &[u32], rows: &[&[f32]], eos: u32) -> Vec<u32> {
let k = accept_longest(history, draft, rows, eos);
let mut out = draft[..k].to_vec();
let ended_at_eos = k > 0 && draft[k - 1] == eos;
if !ended_at_eos {
let mut seq = history.to_vec();
seq.extend_from_slice(&draft[..k]);
out.push(greedy(rows[k], &seq));
}
out
}
fn assert_parity(history: &[u32], draft: &[u32], rows: &[&[f32]], eos: u32) {
let spec = spec_stream(history, draft, rows, eos);
let reference = ref_seq_greedy(history, rows, eos, spec.len());
assert_eq!(
spec, reference,
"speculative accept+correct stream must equal sequential greedy"
);
}
fn history_banning_token_7() -> Vec<u32> {
let prefix: Vec<u32> = (20u32..54).collect();
assert_eq!(prefix.len(), 34);
let mut h = Vec::with_capacity(69);
h.extend_from_slice(&prefix); h.push(7); h.extend_from_slice(&prefix); h
}
#[test]
fn full_accept_returns_whole_draft_and_matches_sequential() {
let l0 = row_argmax(3);
let l1 = row_argmax(4);
let l2 = row_argmax(2);
let bonus = row_argmax(5); let rows: [&[f32]; 4] = [&l0, &l1, &l2, &bonus];
let history: [u32; 0] = [];
let draft = [3u32, 4, 2];
assert_eq!(accept_longest(&history, &draft, &rows, EOS), 3);
assert_parity(&history, &draft, &rows, EOS);
}
#[test]
fn mid_mismatch_truncates_at_first_divergence() {
let l0 = row_argmax(3);
let l1 = row_argmax(4); let l2 = row_argmax(2);
let rows: [&[f32]; 4] = [&l0, &l1, &l2, &l2];
let history: [u32; 0] = [];
let draft = [3u32, 7, 2];
assert_eq!(accept_longest(&history, &draft, &rows, EOS), 1);
assert_parity(&history, &draft, &rows, EOS);
}
#[test]
fn eos_in_draft_accepts_through_eos_and_stops() {
let l0 = row_argmax(5);
let l1 = row_argmax(EOS); let l2 = row_argmax(6); let rows: [&[f32]; 4] = [&l0, &l1, &l2, &l2];
let history: [u32; 0] = [];
let draft = [5u32, EOS, 6];
assert_eq!(accept_longest(&history, &draft, &rows, EOS), 2);
assert_parity(&history, &draft, &rows, EOS);
}
#[test]
fn ngram35_ban_flips_the_verified_token() {
let history = history_banning_token_7();
let l0 = row_peaked(7, 6);
let l1 = row_argmax(5);
let rows: [&[f32]; 2] = [&l0, &l1];
let raw_argmax = sampler::argmax_row(&l0).unwrap();
assert_eq!(raw_argmax, 7, "raw argmax (no ban) is token 7");
assert_eq!(greedy(&l0, &history), 6, "ban flips greedy to token 6");
assert_eq!(accept_longest(&history, &[6], &rows, EOS), 1);
assert_eq!(accept_longest(&history, &[7], &rows, EOS), 0);
assert_parity(&history, &[6], &rows, EOS);
}
#[test]
fn short_verify_logits_stops_without_panic() {
let l0 = row_argmax(3);
let rows: [&[f32]; 1] = [&l0]; let history: [u32; 0] = [];
let draft = [3u32, 4, 2];
assert_eq!(accept_longest(&history, &draft, &rows, EOS), 1);
}
fn assert_valid_draft(seq: &[u32], max_draft: usize, ngram: usize, draft: &[u32]) {
assert!(
draft.len() <= max_draft,
"draft must respect the max_draft budget"
);
if draft.is_empty() {
return;
}
let n = seq.len();
let needle = &seq[n - ngram..];
let backed_by_history = (0..n - ngram).rev().any(|s| {
&seq[s..s + ngram] == needle
&& seq.get(s + ngram..s + ngram + draft.len()) == Some(draft)
});
assert!(
backed_by_history,
"every proposed token must be a verbatim continuation of a matched suffix"
);
}
#[test]
fn draft_replays_continuation_of_repeated_ngram() {
let seq = [5u32, 6, 7, 8, 5, 6];
let draft = draft_ngram(&seq, 2, 2);
assert_eq!(
draft,
vec![7, 8],
"predicts the tokens that followed earlier [5,6]"
);
assert_valid_draft(&seq, 2, 2, &draft);
}
#[test]
fn draft_picks_the_most_recent_earlier_occurrence() {
let seq = [5u32, 6, 7, 5, 6, 9, 5, 6];
let draft = draft_ngram(&seq, 1, 2);
assert_eq!(
draft,
vec![9],
"most recent earlier match wins over the older one"
);
assert_valid_draft(&seq, 1, 2, &draft);
}
#[test]
fn draft_empty_when_suffix_never_recurs() {
let seq = [1u32, 2, 3, 4, 5];
assert!(
draft_ngram(&seq, 4, 2).is_empty(),
"no earlier match -> empty proposal"
);
}
#[test]
fn draft_truncates_to_max_draft() {
let seq = [5u32, 6, 7, 8, 9, 5, 6];
let draft = draft_ngram(&seq, 3, 2);
assert_eq!(draft, vec![7, 8, 9], "continuation truncated to max_draft");
assert_eq!(draft.len(), 3, "never proposes more than max_draft tokens");
assert_valid_draft(&seq, 3, 2, &draft);
}
#[test]
fn draft_empty_when_needle_longer_than_history() {
assert!(
draft_ngram(&[1u32, 2], 4, 5).is_empty(),
"needle longer than history -> empty"
);
assert!(
draft_ngram(&[1u32, 2, 3], 4, 3).is_empty(),
"len == ngram -> empty"
);
}
#[test]
fn draft_empty_on_degenerate_inputs_without_panic() {
let seq = [5u32, 6, 7, 5, 6];
assert!(draft_ngram(&seq, 0, 2).is_empty(), "zero budget -> empty");
assert!(
draft_ngram(&seq, 4, 0).is_empty(),
"zero-length needle -> empty"
);
assert!(draft_ngram(&[], 4, 2).is_empty(), "empty seq -> empty");
assert!(
draft_ngram(&[7u32], 4, 1).is_empty(),
"len == ngram (single token) -> empty"
);
}
#[test]
fn draft_is_a_pure_proposal_never_panics_and_stays_valid() {
let seqs: [&[u32]; 5] = [
&[],
&[7],
&[1, 1, 1, 1],
&[5, 6, 7, 8, 5, 6, 7],
&[2, 3, 2, 3, 2, 3, 2, 3],
];
for seq in seqs {
for ngram in 0..=4usize {
for max_draft in 0..=4usize {
let draft = draft_ngram(seq, max_draft, ngram);
if ngram == 0 || max_draft == 0 || seq.len() <= ngram {
assert!(
draft.is_empty(),
"degenerate inputs must yield an empty draft"
);
} else {
assert_valid_draft(seq, max_draft, ngram, &draft);
}
}
}
}
}
const LV: usize = 16;
fn peak_row(token: u32) -> Mat {
let mut r = vec![0.0f32; LV];
r[token as usize] = 10.0;
Mat::from_vec(1, LV, r)
}
fn params_single(max_length: usize) -> DecodeParams {
let mut p = DecodeParams::single_image();
p.max_length = max_length;
p
}
fn content_logits(seq: &[u32]) -> Mat {
let start = seq.len().saturating_sub(3);
let mut h: u64 = 0xcbf2_9ce4_8422_2325;
for &t in &seq[start..] {
h ^= u64::from(t);
h = h.wrapping_mul(0x0000_0100_0000_01b3);
}
let pick = if seq.len() >= 4 && (h & 7) == 0 {
EOS
} else {
2 + (h % 5) as u32
};
peak_row(pick)
}
#[derive(Default)]
struct Coverage {
empty_drafts: usize,
full_accepts: usize,
partial_accepts: usize,
corrections: usize,
eos_stops: usize,
max_stops: usize,
}
fn seq_generate(
oracle: &dyn Fn(&[u32]) -> Mat,
prompt: &[u32],
params: &DecodeParams,
) -> Vec<u32> {
let mut generated = prompt.to_vec();
let mut emitted = Vec::new();
while emitted.len() < params.max_length {
let logits = oracle(&generated);
let step = sampler::decode_step(&logits, &generated, params).expect("seq decode_step");
generated.push(step.token_id);
emitted.push(step.token_id);
if step.is_eos {
break;
}
}
emitted
}
fn spec_generate(
oracle: &dyn Fn(&[u32]) -> Mat,
prompt: &[u32],
params: &DecodeParams,
max_draft: usize,
ngram: usize,
cov: &mut Coverage,
) -> Vec<u32> {
let mut generated = prompt.to_vec();
let mut emitted = Vec::new();
let mut eos = false;
while emitted.len() < params.max_length {
let draft = draft_ngram(&generated, max_draft, ngram);
if draft.is_empty() {
cov.empty_drafts += 1;
let logits = oracle(&generated);
let step =
sampler::decode_step(&logits, &generated, params).expect("spec fallback step");
generated.push(step.token_id);
emitted.push(step.token_id);
if step.is_eos {
eos = true;
break;
}
continue;
}
let mut verify_logits: Vec<Mat> = Vec::with_capacity(draft.len() + 1);
for i in 0..=draft.len() {
let mut ctx = generated.clone();
ctx.extend_from_slice(&draft[..i]);
verify_logits.push(oracle(&ctx));
}
let emit =
resolve_round(&generated, &draft, &verify_logits, params).expect("resolve_round");
if emit.accepted == draft.len() {
cov.full_accepts += 1;
} else {
cov.partial_accepts += 1;
}
let mut stopped = false;
for &token in &draft[..emit.accepted] {
generated.push(token);
emitted.push(token);
if params.eos_token_id == token {
eos = true;
stopped = true;
break;
}
if emitted.len() >= params.max_length {
stopped = true;
break;
}
}
if stopped {
break;
}
match emit.correction {
None => break,
Some(c) => {
cov.corrections += 1;
generated.push(c.token_id);
emitted.push(c.token_id);
if c.is_eos {
eos = true;
break;
}
}
}
}
if eos {
cov.eos_stops += 1;
} else if emitted.len() >= params.max_length {
cov.max_stops += 1;
}
emitted
}
fn run_length_case(
target: &[u32],
prompt: &[u32],
max_length: usize,
max_draft: usize,
ngram: usize,
cov: &mut Coverage,
) -> Vec<u32> {
let params = params_single(max_length);
let oracle = |s: &[u32]| {
let t = target.get(s.len()).copied().unwrap_or(EOS);
peak_row(t)
};
let seq = seq_generate(&oracle, prompt, ¶ms);
let spec = spec_generate(&oracle, prompt, ¶ms, max_draft, ngram, cov);
assert_eq!(
spec, seq,
"spec != sequential greedy (length oracle) target={target:?} prompt={prompt:?} \
md={max_draft} ng={ngram} ml={max_length}"
);
seq
}
#[test]
fn spec_loop_is_byte_identical_to_sequential_greedy() {
let mut cov = Coverage::default();
let target_a = [8u32, 9, 8, 9, 8, 9, 8, 9, EOS];
let prompt_a = [8u32, 9];
let full = run_length_case(&target_a, &prompt_a, 100, 3, 2, &mut cov);
assert_eq!(
full,
vec![8, 9, 8, 9, 8, 9, EOS],
"EOS-terminated greedy stream"
);
let capped = run_length_case(&target_a, &prompt_a, 3, 3, 2, &mut cov);
assert_eq!(capped, vec![8u32, 9, 8], "max_length=3 cutoff (mid-round)");
run_length_case(
&target_a,
&prompt_a,
100,
SPEC_DRAFT_MAX,
SPEC_DRAFT_NGRAM,
&mut cov,
);
let target_b = [3u32, 4, 5, 6, 3, 4, 5, 6, 3, 4, 5, 7, EOS];
let prompt_b = [3u32, 4, 5];
run_length_case(&target_b, &prompt_b, 100, 4, 3, &mut cov);
run_length_case(&target_b, &prompt_b, 100, 3, 2, &mut cov);
let oracle: fn(&[u32]) -> Mat = content_logits;
let mut rng: u64 = 0x1234_5678_9abc_def0;
let mut next = || {
rng ^= rng << 13;
rng ^= rng >> 7;
rng ^= rng << 17;
rng
};
for _ in 0..48 {
let plen = 3 + (next() % 3) as usize;
let mut prompt = Vec::with_capacity(plen);
for _ in 0..plen {
prompt.push(2 + (next() % 5) as u32);
}
for &(md, ng) in &[(3usize, 2usize), (4, 3), (4, 2), (2, 1)] {
for &ml in &[10usize, 18, 24] {
let params = params_single(ml);
let seq = seq_generate(&oracle, &prompt, ¶ms);
let spec = spec_generate(&oracle, &prompt, ¶ms, md, ng, &mut cov);
assert_eq!(
spec, seq,
"spec != sequential greedy (content oracle) prompt={prompt:?} \
md={md} ng={ng} ml={ml}"
);
}
}
}
assert!(cov.empty_drafts > 0, "empty-draft fallback never exercised");
assert!(cov.full_accepts > 0, "full-accept round never exercised");
assert!(
cov.partial_accepts > 0,
"reject+correction round never exercised"
);
assert!(cov.corrections > 0, "correction token never emitted");
assert!(cov.eos_stops > 0, "EOS halt never exercised");
assert!(cov.max_stops > 0, "max_length cutoff never exercised");
}
#[test]
fn resolve_round_full_accept_appends_bonus_correction() {
let params = DecodeParams::single_image();
let rows = vec![peak_row(3), peak_row(4), peak_row(5)];
let emit = resolve_round(&[], &[3, 4], &rows, ¶ms).unwrap();
assert_eq!(emit.accepted, 2);
let c = emit.correction.expect("bonus correction after full accept");
assert_eq!(c.token_id, 5);
assert!(!c.is_eos);
}
#[test]
fn resolve_round_mid_reject_corrects_from_divergent_row() {
let params = DecodeParams::single_image();
let rows = vec![peak_row(3), peak_row(4), peak_row(9)];
let emit = resolve_round(&[], &[3, 7], &rows, ¶ms).unwrap();
assert_eq!(emit.accepted, 1);
let c = emit.correction.expect("correction at first divergence");
assert_eq!(c.token_id, 4);
}
#[test]
fn resolve_round_accepted_eos_has_no_correction() {
let params = DecodeParams::single_image();
let rows = vec![peak_row(5), peak_row(EOS), peak_row(6)];
let emit = resolve_round(&[], &[5, EOS], &rows, ¶ms).unwrap();
assert_eq!(emit.accepted, 2);
assert!(
emit.correction.is_none(),
"accepted EOS halts with no correction"
);
}
}