use super::{delta_rule_recurrence, delta_rule_recurrence_gqa};
struct Lcg(u64);
impl Lcg {
fn new(seed: u64) -> Self {
Self(seed)
}
fn next_f32(&mut self) -> f32 {
self.0 = self
.0
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
let bits = (self.0 >> 33) as u32;
(f32::from(u16::try_from(bits & 0xFFFF).unwrap_or(0)) / 32768.0) - 1.0
}
fn vec(&mut self, n: usize) -> Vec<f32> {
(0..n).map(|_| self.next_f32()).collect()
}
}
fn tiled_heads(x: &[f32], head_dim: usize, num_k_heads: usize, num_v_heads: usize) -> Vec<f32> {
let mut out = Vec::with_capacity(num_v_heads * head_dim);
for h in 0..num_v_heads {
let src = h % num_k_heads;
out.extend_from_slice(&x[src * head_dim..(src + 1) * head_dim]);
}
out
}
fn repeat_interleave_heads(x: &[f32], head_dim: usize, ratio: usize) -> Vec<f32> {
let num_heads = x.len() / head_dim;
let mut out = Vec::with_capacity(num_heads * ratio * head_dim);
for h in 0..num_heads * ratio {
let src = h / ratio;
out.extend_from_slice(&x[src * head_dim..(src + 1) * head_dim]);
}
out
}
#[test]
fn qwen35_gqa_recurrence_tiles_q_and_k_over_the_value_heads() {
let (num_k_heads, num_v_heads, d) = (2usize, 4usize, 8usize);
let ratio = num_v_heads / num_k_heads;
let mut rng = Lcg::new(0x3477);
let q = rng.vec(num_k_heads * d);
let k = rng.vec(num_k_heads * d);
let v = rng.vec(num_v_heads * d);
let beta = rng.vec(num_v_heads);
let gate: Vec<f32> = rng.vec(num_v_heads).iter().map(|g| g * 0.1).collect();
let state0 = rng.vec(num_v_heads * d * d);
let mut state_gqa = state0.clone();
let mut out_gqa = vec![0.0f32; num_v_heads * d];
delta_rule_recurrence_gqa(
&q,
&k,
&v,
&beta,
&gate,
&mut state_gqa,
&mut out_gqa,
num_k_heads,
d,
num_v_heads,
d,
);
let q_exp = tiled_heads(&q, d, num_k_heads, num_v_heads);
let k_exp = tiled_heads(&k, d, num_k_heads, num_v_heads);
let mut state_ref = state0.clone();
let mut out_ref = vec![0.0f32; num_v_heads * d];
delta_rule_recurrence(
&q_exp,
&k_exp,
&v,
&beta,
&gate,
&mut state_ref,
&mut out_ref,
num_v_heads,
d,
);
assert_eq!(
out_gqa, out_ref,
"the GQA recurrence output is not the tiled (ggml_repeat) reference"
);
assert_eq!(
state_gqa, state_ref,
"the GQA recurrence left a different recurrent state than the tiled reference"
);
let q_interleaved = repeat_interleave_heads(&q, d, ratio);
let k_interleaved = repeat_interleave_heads(&k, d, ratio);
let mut state_interleaved = state0;
let mut out_interleaved = vec![0.0f32; num_v_heads * d];
delta_rule_recurrence(
&q_interleaved,
&k_interleaved,
&v,
&beta,
&gate,
&mut state_interleaved,
&mut out_interleaved,
num_v_heads,
d,
);
assert_ne!(
out_interleaved, out_ref,
"the fixture does not discriminate: repeat_interleave gives the same output as the tiled \
mapping, so this test would pass for a wrong mapping"
);
}
#[test]
fn qwen35_gqa_ratio_one_is_byte_identical_to_the_legacy_recurrence() {
let (num_heads, d) = (3usize, 4usize);
let mut rng = Lcg::new(0x0808);
let q = rng.vec(num_heads * d);
let k = rng.vec(num_heads * d);
let v = rng.vec(num_heads * d);
let beta = rng.vec(num_heads);
let gate: Vec<f32> = rng.vec(num_heads).iter().map(|g| g * 0.1).collect();
let state0 = rng.vec(num_heads * d * d);
let mut state_new = state0.clone();
let mut out_new = vec![0.0f32; num_heads * d];
delta_rule_recurrence_gqa(
&q,
&k,
&v,
&beta,
&gate,
&mut state_new,
&mut out_new,
num_heads,
d,
num_heads,
d,
);
let mut state_old = state0;
let mut out_old = vec![0.0f32; num_heads * d];
delta_rule_recurrence(
&q,
&k,
&v,
&beta,
&gate,
&mut state_old,
&mut out_old,
num_heads,
d,
);
assert_eq!(out_new, out_old, "ratio 1 changed the output bit pattern");
assert_eq!(state_new, state_old, "ratio 1 changed the recurrent state");
}
#[test]
fn qwen35_gqa_state_is_rectangular_k_dim_by_v_dim() {
let (head_k_dim, head_v_dim) = (2usize, 3usize);
let q = [1.0f32, 2.0];
let k = [0.5f32, 0.5];
let v = [1.0f32, -1.0, 2.0];
let beta = [0.5f32];
let gate = [0.0f32]; let mut state = [
1.0f32, 0.0, 0.0, 1.0, 1.0, 1.0, ];
let mut output = [0.0f32; 3];
delta_rule_recurrence_gqa(
&q,
&k,
&v,
&beta,
&gate,
&mut state,
&mut output,
1,
head_k_dim,
1,
head_v_dim,
);
let s2 = 2.0f32.sqrt();
let want = [
(1.125 + 2.0 * 0.125) / s2,
(-0.375 + 2.0 * 0.625) / s2,
(1.25 + 2.0 * 1.25) / s2,
];
for (j, w) in want.iter().enumerate() {
assert!(
(output[j] - w).abs() < 1e-6,
"output[{j}] = {}, want {w} (state: {state:?})",
output[j]
);
}
assert!((state[0] - 1.125).abs() < 1e-6, "{state:?}");
assert!((state[1] - 0.125).abs() < 1e-6, "{state:?}");
assert!((state[4] - 1.25).abs() < 1e-6, "{state:?}");
assert!((state[5] - 1.25).abs() < 1e-6, "{state:?}");
}
#[test]
fn qwen35_gqa_4b_matches_llama_cpp_greedy_tokens() {
const MODEL_PATH: &str = "/home/noah/models/Qwen3.5-4B-Q4_K_M.gguf";
if !std::path::Path::new(MODEL_PATH).exists() {
eprintln!("SKIP: {MODEL_PATH} is absent");
return;
}
let mapped = crate::gguf::MappedGGUFModel::from_path(MODEL_PATH).expect("map the 4B GGUF");
let base =
super::Qwen35Model::create_base_model(&mapped.model, mapped.data()).expect("4B base model");
let qwen = super::Qwen35Model::from_model_and_layers(&base, &mapped.model, mapped.data())
.expect("4B hybrid layers");
assert!(
qwen.num_v_heads > qwen.num_k_heads,
"this file is not a GQA DeltaNet file (num_v_heads {} vs num_k_heads {}), so it cannot \
falsify the head mapping",
qwen.num_v_heads,
qwen.num_k_heads
);
assert_eq!(
qwen.num_v_heads % qwen.num_k_heads,
0,
"num_v_heads {} is not a multiple of num_k_heads {}",
qwen.num_v_heads,
qwen.num_k_heads
);
if let Some(super::Qwen35OwnedLayer::DeltaNet(d)) = qwen
.layers
.iter()
.find(|l| matches!(l, super::Qwen35OwnedLayer::DeltaNet(_)))
{
assert_eq!(
d.ssm_norm_weight.len(),
qwen.head_v_dim,
"ssm_norm_weight is {} wide but head_v_dim is {}",
d.ssm_norm_weight.len(),
qwen.head_v_dim
);
assert_eq!(
d.ssm_a.len(),
qwen.num_v_heads,
"ssm_a (A) is per value head"
);
}
assert_eq!(
qwen.layers.len(),
mapped
.model
.num_layers()
.expect("4B GGUF carries qwen35.block_count"),
"the hybrid loader built {} layers for a {:?}-block file",
qwen.layers.len(),
mapped.model.num_layers()
);
let prompt = [760u32, 6511, 314, 9338, 369];
const WANT: [u32; 8] = [11751, 13, 198, 32, 13, 2912, 198, 33];
let mut state = qwen.new_state(prompt.len() + WANT.len());
let mut logits = Vec::new();
for (pos, &token) in prompt.iter().enumerate() {
logits = qwen
.forward_single_qwen35(token, &mut state, pos)
.expect("4B forward");
assert_eq!(logits.len(), qwen.base.config.vocab_size);
assert!(
logits.iter().all(|l| l.is_finite()),
"position {pos}: the 4B forward produced a non-finite logit"
);
}
let mut got = Vec::with_capacity(WANT.len());
for step in 0..WANT.len() {
let next = crate::gguf::ops::argmax(&logits);
got.push(next);
logits = qwen
.forward_single_qwen35(next, &mut state, prompt.len() + step)
.expect("4B forward");
}
assert_eq!(
got,
WANT.to_vec(),
"the 4B CPU forward does not follow llama.cpp greedily (head mapping / Gated DeltaNet \
arithmetic): got {got:?}, want {:?} — num_k_heads {}, num_v_heads {}, head_k_dim {}, \
head_v_dim {}",
WANT,
qwen.num_k_heads,
qwen.num_v_heads,
qwen.head_k_dim,
qwen.head_v_dim
);
}
#[test]
fn qwen35_0_8b_matches_llama_cpp_greedy_tokens() {
const MODEL_PATH: &str = "/home/noah/models/Qwen3.5-0.8B-Q4_K_M.gguf";
if !std::path::Path::new(MODEL_PATH).exists() {
eprintln!("SKIP: {MODEL_PATH} is absent");
return;
}
let mapped = crate::gguf::MappedGGUFModel::from_path(MODEL_PATH).expect("map the 0.8B GGUF");
let base = super::Qwen35Model::create_base_model(&mapped.model, mapped.data())
.expect("0.8B base model");
let qwen = super::Qwen35Model::from_model_and_layers(&base, &mapped.model, mapped.data())
.expect("0.8B hybrid layers");
assert_eq!(
qwen.num_v_heads, qwen.num_k_heads,
"0.8B is supposed to be the ratio-1 control"
);
let prompt = [760u32, 6511, 314, 9338, 369];
const WANT: [u32; 4] = [279, 6511, 314, 279];
let mut state = qwen.new_state(prompt.len() + WANT.len());
let mut logits = Vec::new();
for (pos, &token) in prompt.iter().enumerate() {
logits = qwen
.forward_single_qwen35(token, &mut state, pos)
.expect("0.8B forward");
}
let mut got = Vec::with_capacity(WANT.len());
for step in 0..WANT.len() {
let next = crate::gguf::ops::argmax(&logits);
got.push(next);
logits = qwen
.forward_single_qwen35(next, &mut state, prompt.len() + step)
.expect("0.8B forward");
}
assert_eq!(
got,
WANT.to_vec(),
"the 0.8B CPU forward does not follow llama.cpp greedily: got {got:?}"
);
}