#![cfg(feature = "mmap")]
mod common;
use common::dense_model_or_skip;
fn argmax(v: &[f32]) -> usize {
v.iter()
.enumerate()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
.map(|(i, _)| i)
.unwrap()
}
fn cosine(a: &[f32], b: &[f32]) -> f32 {
let dot: f32 = a.iter().zip(b).map(|(x, y)| x * y).sum();
let na: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
let nb: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
dot / (na * nb)
}
struct Collect(Vec<u32>);
impl cera::ModalitySink for Collect {
fn on_text_tokens(&mut self, t: &[u32]) {
self.0.extend_from_slice(t);
}
fn on_done(&mut self, _r: cera::FinishReason) {}
}
fn build_spec_session(path: &std::path::Path, max_seq_len: Option<u32>) -> cera::Session {
use std::sync::Arc;
let gguf = cera::gguf::GgufFile::open(path).unwrap();
let tokenizer = cera::tokenizer::BpeTokenizer::from_gguf(&gguf).unwrap();
let model: Arc<dyn cera::model::Model> =
Arc::from(cera::model::load_model(gguf, None, 8192).unwrap());
cera::Session::new(
model,
Arc::new(tokenizer),
cera::ModalityCapabilities::text_only(),
cera::SessionConfig {
kv_compression: cera::kv_cache::KvCompression::None,
seed: None,
ubatch_size: 0,
max_seq_len,
..Default::default()
},
)
.unwrap()
}
#[test]
#[ignore = "needs a real dense GGUF; set CERA_DENSE_MODEL"]
fn all_logits_argmax_matches_per_token_forward() {
let Some(path) = dense_model_or_skip() else {
return;
};
let gguf_a = cera::gguf::GgufFile::open(&path).unwrap();
let model_a = cera::model::load_model(gguf_a, None, 8192).unwrap();
assert!(
model_a.supports_all_logits(),
"dense model must support forward_prefill_logits_all"
);
let cfg = model_a.config();
let vocab = cfg.vocab_size;
let tokens: Vec<u32> = vec![1, 15043, 3186, 297, 4223, 29889, 306, 626];
let n = tokens.len();
let mut state_a = cera::kv_cache::InferenceState::from_config(cfg).unwrap();
let all = model_a.forward_prefill_logits_all(&tokens, 0, &mut state_a);
assert_eq!(all.len(), n * vocab, "all-logits must be [n * vocab]");
assert_eq!(state_a.seq_len, n, "KV must be appended for all n tokens");
let gguf_b = cera::gguf::GgufFile::open(&path).unwrap();
let model_b = cera::model::load_model(gguf_b, None, 8192).unwrap();
let mut state_b = cera::kv_cache::InferenceState::from_config(cfg).unwrap();
for (i, &tok) in tokens.iter().enumerate() {
let per_tok = model_b.forward(&[tok], i, &mut state_b);
let row = &all[i * vocab..(i + 1) * vocab];
let a1 = argmax(row);
let a2 = argmax(&per_tok);
let cos = cosine(row, &per_tok);
println!("pos {i}: argmax batched={a1} per-token={a2} cosine={cos:.5}");
assert_eq!(
a1, a2,
"position {i}: batched all-logits argmax ({a1}) != per-token forward argmax ({a2})"
);
assert!(
cos > 0.99,
"position {i}: cosine {cos} too low (batched vs per-token drift)"
);
}
}
fn max_abs_diff(a: &[f32], b: &[f32]) -> f32 {
a.iter()
.zip(b)
.map(|(x, y)| (x - y).abs())
.fold(0.0f32, f32::max)
}
#[test]
#[ignore = "needs a real dense GGUF; set CERA_DENSE_MODEL"]
fn truncate_to_restores_kv_exactly() {
let Some(path) = dense_model_or_skip() else {
return;
};
let tokens: Vec<u32> = vec![1, 15043, 3186, 297, 4223, 29889, 306, 626];
let (l, r) = (4usize, tokens.len());
let ga = cera::gguf::GgufFile::open(&path).unwrap();
let ma = cera::model::load_model(ga, None, 8192).unwrap();
let cfg = ma.config();
let mut sa = cera::kv_cache::InferenceState::from_config(cfg).unwrap();
ma.forward_prefill(&tokens[..l], 0, &mut sa);
let ref_logits = ma.forward_prefill(&tokens[l..], l, &mut sa);
let gb = cera::gguf::GgufFile::open(&path).unwrap();
let mb = cera::model::load_model(gb, None, 8192).unwrap();
let mut sb = cera::kv_cache::InferenceState::from_config(cfg).unwrap();
mb.forward_prefill(&tokens, 0, &mut sb);
assert_eq!(sb.seq_len, r);
sb.truncate_to(l);
assert_eq!(
sb.seq_len, l,
"seq_len must be reset to the truncation length"
);
let test_logits = mb.forward_prefill(&tokens[l..], l, &mut sb);
assert_eq!(
sb.seq_len, r,
"re-prefill must grow the cache back to full length"
);
let d = max_abs_diff(&ref_logits, &test_logits);
println!(
"truncate_to: max_abs_diff = {d:.3e}, argmax ref={} test={}",
argmax(&ref_logits),
argmax(&test_logits)
);
assert_eq!(
argmax(&ref_logits),
argmax(&test_logits),
"truncate+re-prefill changed the argmax"
);
assert!(
d < 1e-3,
"truncate_to did not restore the KV exactly (max_abs_diff = {d})"
);
}
fn greedy_reference(
model: &dyn cera::model::Model,
state: &mut cera::kv_cache::InferenceState,
prompt: &[u32],
max_new: usize,
) -> Vec<u32> {
let mut next = model.forward_prefill(prompt, 0, state);
let mut out: Vec<u32> = Vec::new();
while out.len() < max_new {
let t = cera::sampler::argmax(&next);
out.push(t);
if out.len() >= max_new {
break;
}
next = model.forward(&[t], state.seq_len, state);
}
out
}
#[test]
#[ignore = "needs a real dense GGUF; set CERA_DENSE_MODEL"]
fn greedy_spec_matches_greedy_within_tie_tolerance() {
let Some(path) = dense_model_or_skip() else {
return;
};
let prompt: Vec<u32> = vec![
1, 450, 6635, 3290, 373, 278, 1775, 29889, 450, 6635, 3290, 373, 278,
];
let max_new = 64usize;
let gr = cera::gguf::GgufFile::open(&path).unwrap();
let mr = cera::model::load_model(gr, None, 8192).unwrap();
let cfg = mr.config();
let mut sr = cera::kv_cache::InferenceState::from_config(cfg).unwrap();
let reference = greedy_reference(mr.as_ref(), &mut sr, &prompt, max_new);
let gs = cera::gguf::GgufFile::open(&path).unwrap();
let ms = cera::model::load_model(gs, None, 8192).unwrap();
let mut ss0 = cera::kv_cache::InferenceState::from_config(cfg).unwrap();
let (nodraft, s0) =
cera::spec::greedy_generate_spec(ms.as_ref(), &mut ss0, &prompt, max_new, &[], 999, 6);
assert_eq!(s0.accepted, 0, "ngram=999 should draft nothing");
assert_eq!(
reference, nodraft,
"no-draft spec-decode must equal per-token greedy exactly (orchestration bug)"
);
let mut ss = cera::kv_cache::InferenceState::from_config(cfg).unwrap();
let (spec, stats) =
cera::spec::greedy_generate_spec(ms.as_ref(), &mut ss, &prompt, max_new, &[], 2, 6);
println!(
"spec: {} tokens, {} rounds, {}/{} drafts accepted ({:.0}% acceptance)",
spec.len(),
stats.rounds,
stats.accepted,
stats.drafted,
stats.acceptance_rate() * 100.0
);
assert_eq!(reference.len(), spec.len(), "length mismatch");
assert!(
stats.accepted > 0,
"expected some drafts to be accepted on a repetitive prompt (else the \
verify/truncate path is untested)"
);
if reference != spec {
let i = reference
.iter()
.zip(&spec)
.position(|(a, b)| a != b)
.unwrap();
println!(
"diverge at gen index {i}: per-token greedy={} vs spec={}",
reference[i], spec[i]
);
let mut seq = prompt.clone();
seq.extend_from_slice(&reference[..i]);
let gv = cera::gguf::GgufFile::open(&path).unwrap();
let mv = cera::model::load_model(gv, None, 8192).unwrap();
let mut sv = cera::kv_cache::InferenceState::from_config(cfg).unwrap();
let all = mv.forward_prefill_logits_all(&seq, 0, &mut sv);
let vocab = cfg.vocab_size;
let last = &all[(seq.len() - 1) * vocab..];
let batched_pred = argmax(last) as u32;
let g_spec = last[spec[i] as usize];
let g_ref = last[reference[i] as usize];
let gap = (g_spec - g_ref).abs();
println!(
"batched verifier argmax at pos {i} = {batched_pred}; logit(spec)={g_spec:.4} logit(ref)={g_ref:.4} gap={gap:.4}"
);
assert!(
gap < 0.05,
"spec[{i}]={} sits {gap:.4} below reference[{i}]={} under the batched \
re-forward — too large for a near-tie flip; accept/truncate likely buggy",
spec[i],
reference[i]
);
} else {
println!("greedy-spec matched per-token greedy exactly (no near-tie flips)");
}
}
#[test]
#[ignore = "needs a real dense GGUF; set CERA_DENSE_MODEL"]
fn session_spec_matches_standalone_driver() {
use cera::{GenerateOpts, SpecDecode};
let Some(path) = dense_model_or_skip() else {
return;
};
let prompt: Vec<u32> = vec![
1, 450, 6635, 3290, 373, 278, 1775, 29889, 450, 6635, 3290, 373, 278,
];
let max_new = 64u32;
let sd = SpecDecode { ngram: 2, k: 6 };
let gr = cera::gguf::GgufFile::open(&path).unwrap();
let mr = cera::model::load_model(gr, None, 8192).unwrap();
let cfg = mr.config();
let mut sr = cera::kv_cache::InferenceState::from_config(cfg).unwrap();
let (driver, dstats) = cera::spec::greedy_generate_spec(
mr.as_ref(),
&mut sr,
&prompt,
max_new as usize,
&[], sd.ngram,
sd.k,
);
assert!(
dstats.accepted > 0,
"expected accepted drafts (else the shared path is untested)"
);
let mut session = build_spec_session(&path, None);
session.append_tokens(&prompt).unwrap();
let mut sink = Collect(Vec::new());
let summary = session
.generate(
&GenerateOpts {
max_tokens: max_new,
temperature: 0.0,
ignore_eos: true,
spec: Some(sd),
..Default::default()
},
&mut sink,
)
.unwrap();
println!(
"session spec: {} tokens (finish {:?}); driver: {} tokens",
sink.0.len(),
summary.finish_reason,
driver.len()
);
assert_eq!(
summary.tokens_generated as usize,
driver.len(),
"session must generate the same count as the driver"
);
assert_eq!(
sink.0, driver,
"Session spec-decode output must equal the standalone driver token-for-token"
);
assert_eq!(
session.position() as usize,
prompt.len() + driver.len(),
"current_pos must equal prompt + generated after a spec run"
);
}
#[test]
#[ignore = "needs a real dense GGUF; set CERA_DENSE_MODEL"]
fn session_spec_honors_stop_without_emitting_it() {
use cera::{FinishReason, GenerateOpts, SpecDecode};
let Some(path) = dense_model_or_skip() else {
return;
};
let prompt: Vec<u32> = vec![
1, 450, 6635, 3290, 373, 278, 1775, 29889, 450, 6635, 3290, 373, 278,
];
let gr = cera::gguf::GgufFile::open(&path).unwrap();
let mr = cera::model::load_model(gr, None, 8192).unwrap();
let cfg = mr.config();
let mut sr = cera::kv_cache::InferenceState::from_config(cfg).unwrap();
let prefill = mr.forward_prefill(&prompt, 0, &mut sr);
let first = argmax(&prefill) as u32;
let mut session = build_spec_session(&path, None);
session.append_tokens(&prompt).unwrap();
let mut sink = Collect(Vec::new());
let summary = session
.generate(
&GenerateOpts {
max_tokens: 64,
temperature: 0.0,
ignore_eos: false,
stop_tokens: vec![first],
spec: Some(SpecDecode { ngram: 2, k: 6 }),
..Default::default()
},
&mut sink,
)
.unwrap();
assert!(
matches!(summary.finish_reason, FinishReason::Stop),
"expected Stop, got {:?}",
summary.finish_reason
);
assert_eq!(
summary.tokens_generated, 0,
"stop token must not be counted"
);
assert!(sink.0.is_empty(), "stop token must not be streamed");
assert_eq!(
session.position() as usize,
prompt.len(),
"stop token's KV must not be appended (position stays at the prompt)"
);
}
#[test]
#[ignore = "needs a real dense GGUF; set CERA_DENSE_MODEL"]
fn session_spec_respects_max_seq_len() {
use cera::{FinishReason, GenerateOpts, SpecDecode};
let Some(path) = dense_model_or_skip() else {
return;
};
let prompt: Vec<u32> = vec![
1, 450, 6635, 3290, 373, 278, 1775, 29889, 450, 6635, 3290, 373, 278,
];
let cap = prompt.len() + 3;
let mut session = build_spec_session(&path, Some(cap as u32));
session.append_tokens(&prompt).unwrap();
let mut sink = Collect(Vec::new());
let summary = session
.generate(
&GenerateOpts {
max_tokens: 256, temperature: 0.0,
ignore_eos: true,
spec: Some(SpecDecode { ngram: 2, k: 6 }),
..Default::default()
},
&mut sink,
)
.unwrap();
assert!(
session.position() as usize <= cap,
"spec decode overshot max_seq_len: position {} > cap {cap}",
session.position()
);
assert!(
matches!(summary.finish_reason, FinishReason::ContextFull),
"expected ContextFull at the bound, got {:?}",
summary.finish_reason
);
}
#[test]
#[ignore = "needs a real dense GGUF; set CERA_DENSE_MODEL"]
#[cfg(all(any(target_arch = "aarch64", target_arch = "x86_64"), not(has_blas)))]
fn batched_all_logits_reaches_the_lm_head_gemm() {
let Some(path) = dense_model_or_skip() else {
return;
};
let (_model, detail) = common::dense_gemm_head_fixture(&path);
println!("reaches the batched projection: {detail}");
}