mod common;
use common::assert_close;
use ferrox_core::cache::KvCache;
use ferrox_models::{Decoder, ModelConfig};
const FIXTURE: &str = concat!(
env!("CARGO_MANIFEST_DIR"),
"/tests/fixtures/gptoss_tiny.gguf"
);
const PROMPT: [usize; 6] = [3, 7, 11, 19, 23, 5];
const GOLDEN_LOGITS: [f32; 48] = [
-0.6055459,
0.10276483,
-0.22337931,
-0.10290539,
0.048099145,
0.726398,
-0.051841702,
-0.10426964,
0.682685,
0.10511006,
-0.6360203,
0.09794438,
0.13892847,
-0.0350772,
-0.025142908,
0.2722936,
-0.20271212,
-0.22984387,
0.5637356,
0.22094958,
0.27907962,
-0.34294918,
-0.107179776,
-0.1551839,
-0.35526568,
-0.56609416,
-0.16943386,
0.08455296,
-0.044256665,
-0.21528326,
0.25945616,
-0.4198072,
-0.18812446,
-0.15074603,
-0.36562067,
0.14833681,
-0.05061131,
0.077752486,
0.3013656,
-0.15252003,
0.33605093,
-0.58283293,
-0.6513095,
0.1741069,
0.7685709,
-0.14693882,
-0.13071328,
0.5620597,
];
const TOL: f32 = 1e-6;
fn load() -> Decoder {
let file = ferrox_gguf::GgufFile::open(FIXTURE).expect("fixture opens");
let config = ModelConfig::from_gguf(&file).expect("fixture config parses");
Decoder::from_gguf(FIXTURE, config).expect("fixture loads")
}
#[test]
fn gpt_oss_prefill_matches_llama_cpp() {
let decoder = load();
let mut caches: Vec<KvCache> = decoder
.layers
.iter()
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect();
let logits = decoder.forward_batch_last(&PROMPT, 0, &mut caches);
assert_close(&logits, &GOLDEN_LOGITS, TOL, "prefill (forward_batch_last)");
}
#[test]
fn gpt_oss_decode_matches_llama_cpp() {
let decoder = load();
let mut caches: Vec<KvCache> = decoder
.layers
.iter()
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect();
let mut logits = Vec::new();
for (pos, &tok) in PROMPT.iter().enumerate() {
logits = decoder.forward_token(tok, pos, &mut caches);
}
assert_close(&logits, &GOLDEN_LOGITS, TOL, "decode (forward_token)");
}
#[test]
fn gpt_oss_multi_seq_matches_llama_cpp() {
let decoder = load();
let mut caches: Vec<Vec<KvCache>> = vec![decoder
.layers
.iter()
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect()];
let mut logits = Vec::new();
for (pos, &tok) in PROMPT.iter().enumerate() {
let out = decoder.forward_multi_seq(&[tok], &[pos], &mut caches);
logits = out.into_iter().next().unwrap();
}
assert_close(&logits, &GOLDEN_LOGITS, TOL, "multi-seq");
}
#[test]
fn gpt_oss_loader_wires_the_whole_graph() {
let decoder = load();
let g = decoder
.gpt_oss
.as_ref()
.expect("gpt-oss checkpoint must take the gpt-oss path");
assert_eq!(g.layers.len(), decoder.layers.len());
for layer in &g.layers {
assert_eq!(layer.attn_sinks.len(), decoder.config.n_heads);
assert_eq!(layer.o_bias.len(), decoder.config.hidden_dim);
assert_eq!(layer.router_bias.len(), decoder.config.moe.n_experts);
assert_eq!(layer.expert_bias.len(), decoder.config.moe.n_experts);
}
assert_eq!(decoder.config.sliding_window, Some(4));
assert_eq!(decoder.config.swa_pattern, Some(2));
assert_eq!(decoder.config.layer_sliding_window(0), Some(4));
assert_eq!(decoder.config.layer_sliding_window(1), None);
assert_eq!(
decoder.config.rope_layout,
ferrox_models::config::RopeLayout::Neox
);
assert_eq!(
decoder.config.layer_rope_theta(0),
decoder.config.rope_theta
);
}
#[test]
fn gpt_oss_golden_is_not_vacuous() {
let baseline = {
let decoder = load();
let mut caches: Vec<KvCache> = decoder
.layers
.iter()
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect();
decoder.forward_batch_last(&PROMPT, 0, &mut caches)
};
assert_close(&baseline, &GOLDEN_LOGITS, TOL, "baseline");
let max_delta = |broken: &[f32]| -> f32 {
broken
.iter()
.zip(GOLDEN_LOGITS.iter())
.map(|(a, b)| (a - b).abs())
.fold(0f32, f32::max)
};
{
let mut decoder = load();
for layer in decoder.gpt_oss.as_mut().unwrap().layers.iter_mut() {
layer
.attn_sinks
.iter_mut()
.for_each(|s| *s = f32::NEG_INFINITY);
}
let mut caches: Vec<KvCache> = decoder
.layers
.iter()
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect();
let broken = decoder.forward_batch_last(&PROMPT, 0, &mut caches);
assert!(
max_delta(&broken) > TOL * 10.0,
"removing the attention sinks must move the logits, else the sink term is dead code"
);
}
{
let mut decoder = load();
for layer in decoder.gpt_oss.as_mut().unwrap().layers.iter_mut() {
layer.router_bias.iter_mut().for_each(|b| *b = 0.0);
}
let mut caches: Vec<KvCache> = decoder
.layers
.iter()
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect();
let broken = decoder.forward_batch_last(&PROMPT, 0, &mut caches);
assert!(max_delta(&broken) > TOL * 10.0, "router bias must matter");
}
{
let mut decoder = load();
for layer in decoder.gpt_oss.as_mut().unwrap().layers.iter_mut() {
for b in layer.expert_bias.iter_mut() {
b.gate.iter_mut().for_each(|x| *x = 0.0);
b.up.iter_mut().for_each(|x| *x = 0.0);
b.down.iter_mut().for_each(|x| *x = 0.0);
}
}
let mut caches: Vec<KvCache> = decoder
.layers
.iter()
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect();
let broken = decoder.forward_batch_last(&PROMPT, 0, &mut caches);
assert!(max_delta(&broken) > TOL * 10.0, "expert biases must matter");
}
{
let mut decoder = load();
for layer in decoder.gpt_oss.as_mut().unwrap().layers.iter_mut() {
layer.o_bias.iter_mut().for_each(|b| *b = 0.0);
}
let mut caches: Vec<KvCache> = decoder
.layers
.iter()
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect();
let broken = decoder.forward_batch_last(&PROMPT, 0, &mut caches);
assert!(
max_delta(&broken) > TOL * 10.0,
"attn output bias must matter"
);
}
{
let mut decoder = load();
decoder.config.swa_pattern = None;
let mut caches: Vec<KvCache> = decoder
.layers
.iter()
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect();
let broken = decoder.forward_batch_last(&PROMPT, 0, &mut caches);
assert!(
max_delta(&broken) > TOL * 10.0,
"the alternating SWA pattern must matter, else layer 1 is being windowed too"
);
}
{
let mut decoder = load();
decoder.config.rope_layout = ferrox_models::config::RopeLayout::Norm;
let mut caches: Vec<KvCache> = decoder
.layers
.iter()
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect();
let broken = decoder.forward_batch_last(&PROMPT, 0, &mut caches);
assert!(
max_delta(&broken) > TOL * 10.0,
"RoPE layout must matter; gpt-oss is NEOX"
);
}
}