#[derive(Debug, Clone, Copy)]
pub struct NgramConfig {
pub min_ngram: usize,
pub max_ngram: usize,
pub k: usize,
pub max_model_len: usize,
}
impl NgramConfig {
pub fn default_for_decode(max_model_len: usize) -> Self {
Self {
min_ngram: 1,
max_ngram: 3,
k: 3,
max_model_len,
}
}
}
pub fn propose(tokens: &[u32], cfg: &NgramConfig) -> Vec<u32> {
let total = tokens.len();
if total < cfg.min_ngram {
return Vec::new();
}
let k_room = cfg.max_model_len.saturating_sub(total);
let k_capped = cfg.k.min(k_room);
if k_capped == 0 {
return Vec::new();
}
if cfg.max_ngram == 0 || cfg.min_ngram > cfg.max_ngram {
return Vec::new();
}
let rev = |i: usize| -> u32 { tokens[total - 1 - i] };
let lps_len = cfg.max_ngram;
let mut lps = vec![0u32; lps_len];
let mut longest_ngram: usize = 0;
let mut position: usize = 0;
let mut prev_lps: usize = 0;
let mut i: usize = 1;
while i < total {
if rev(prev_lps) == rev(i) {
prev_lps += 1;
if prev_lps >= longest_ngram {
longest_ngram = prev_lps;
position = i;
}
if i < lps_len {
lps[i] = prev_lps as u32;
}
if prev_lps == cfg.max_ngram {
prev_lps = lps[cfg.max_ngram - 1] as usize;
}
i += 1;
} else if prev_lps != 0 {
prev_lps = lps[prev_lps - 1] as usize;
} else {
i += 1;
}
}
if longest_ngram < cfg.min_ngram {
return Vec::new();
}
let start = total - 1 - position + longest_ngram;
let drafts_room = total.saturating_sub(start);
let n = k_capped.min(drafts_room);
if n == 0 {
return Vec::new();
}
tokens[start..start + n].to_vec()
}
#[cfg(test)]
mod tests {
use super::*;
fn cfg(min_n: usize, max_n: usize, k: usize) -> NgramConfig {
NgramConfig {
min_ngram: min_n,
max_ngram: max_n,
k,
max_model_len: 4096,
}
}
#[test]
fn propose_empty_when_below_min_ngram() {
assert!(propose(&[], &cfg(1, 3, 3)).is_empty());
assert!(propose(&[7], &cfg(2, 3, 3)).is_empty());
}
#[test]
fn propose_empty_when_no_match() {
let drafts = propose(&[1, 2, 3, 4, 5, 6], &cfg(2, 3, 3));
assert!(drafts.is_empty(), "expected no drafts, got {:?}", drafts);
}
#[test]
fn propose_basic_repetition() {
let tokens = vec![10u32, 20, 30, 99, 88, 10, 20, 30];
let drafts = propose(&tokens, &cfg(1, 3, 3));
assert_eq!(drafts, vec![99, 88, 10]);
}
#[test]
fn propose_respects_k_truncation() {
let tokens = vec![10u32, 20, 30, 99, 88, 10, 20, 30];
let drafts = propose(&tokens, &cfg(1, 3, 2));
assert_eq!(drafts, vec![99, 88]);
}
#[test]
fn propose_respects_max_ngram_cap() {
let tokens = vec![1u32, 2, 3, 4, 5, 99, 1, 2, 3, 4, 5];
let drafts = propose(&tokens, &cfg(2, 2, 3));
assert_eq!(drafts, vec![99, 1, 2]);
}
#[test]
fn propose_picks_earliest_occurrence_on_tie() {
let tokens = vec![10u32, 20, 100, 10, 20, 200, 10, 20];
let drafts = propose(&tokens, &cfg(1, 3, 3));
assert_eq!(drafts, vec![100, 10, 20]);
}
#[test]
fn propose_caps_k_at_max_model_len() {
let cfg = NgramConfig {
min_ngram: 1,
max_ngram: 3,
k: 5,
max_model_len: 10, };
let tokens = vec![10u32, 20, 30, 99, 88, 10, 20, 30];
let drafts = propose(&tokens, &cfg);
assert_eq!(drafts.len(), 2, "expected k clamped to max_model_len - len");
assert_eq!(drafts, vec![99, 88]);
}
#[test]
fn propose_handles_longest_match_at_seq_end() {
let tokens = vec![1u32, 2, 3, 1, 2, 3];
let drafts = propose(&tokens, &cfg(1, 3, 3));
assert_eq!(drafts, vec![1, 2, 3]);
}
#[test]
fn propose_zero_max_ngram_returns_empty() {
let bad_cfg = NgramConfig {
min_ngram: 0,
max_ngram: 0,
k: 3,
max_model_len: 4096,
};
assert!(propose(&[1, 2, 3], &bad_cfg).is_empty());
}
#[test]
fn propose_k_zero_returns_empty() {
let bad_cfg = NgramConfig {
min_ngram: 1,
max_ngram: 3,
k: 0,
max_model_len: 4096,
};
assert!(propose(&[1, 2, 3], &bad_cfg).is_empty());
}
#[test]
fn default_config_is_reasonable() {
let cfg = NgramConfig::default_for_decode(4096);
assert_eq!(cfg.k, 3);
assert_eq!(cfg.min_ngram, 1);
assert_eq!(cfg.max_ngram, 3);
assert_eq!(cfg.max_model_len, 4096);
}
fn rand_tokens(seed: u64, n: usize, vocab: u32) -> Vec<u32> {
let mut state = seed;
(0..n)
.map(|_| {
state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((state >> 33) as u32) % vocab
})
.collect()
}
#[test]
#[ignore]
fn bench_ngram_proposer_at_realistic_decode_lengths() {
use std::time::Instant;
let cfg = NgramConfig {
min_ngram: 1,
max_ngram: 3,
k: 3,
max_model_len: 16_384,
};
let lengths = [128usize, 512, 1024, 2048, 4096, 8192];
for &n in &lengths {
let tokens = rand_tokens(0xCAFE_BEEF, n, 256);
for _ in 0..100 {
let _ = propose(&tokens, &cfg);
}
let mut samples: Vec<u128> = Vec::with_capacity(1000);
for _ in 0..1000 {
let t0 = Instant::now();
let _ = propose(&tokens, &cfg);
samples.push(t0.elapsed().as_nanos());
}
samples.sort();
let p50 = samples[500];
let p99 = samples[990];
eprintln!(
"[BENCH iter-115] propose len={:5} p50={:6} ns p99={:6} ns",
n, p50, p99
);
assert!(
(p50 as usize) < 100_000,
"propose at len={n} took {p50} ns p50 — too slow for hot path (target <100 µs)"
);
}
}
}