use super::super::{Qwen35CudaModel, Qwen35CudaState};
use crate::gguf::forward_qwen35::Qwen35Model;
const MODEL_0_8B: &str = "/home/noah/models/Qwen3.5-0.8B-Q4_K_M.gguf";
#[derive(Clone, Copy)]
struct Budget {
cosine: f64,
logits: f32,
state: f32,
}
const F32_BUDGET: Budget = Budget {
cosine: 0.99999,
logits: 1e-4,
state: 1e-4,
};
const FLASH_BUDGET: Budget = Budget {
cosine: 0.99999,
logits: 2e-3,
state: 5e-3,
};
const F16_GEMM_BUDGET: Budget = Budget {
cosine: 0.999_995,
logits: 7.5e-3,
state: 1.25e-2,
};
fn cosine(a: &[f32], b: &[f32]) -> f64 {
let (mut ab, mut aa, mut bb) = (0.0f64, 0.0f64, 0.0f64);
for (x, y) in a.iter().zip(b) {
ab += f64::from(*x) * f64::from(*y);
aa += f64::from(*x) * f64::from(*x);
bb += f64::from(*y) * f64::from(*y);
}
ab / (aa.sqrt() * bb.sqrt()).max(1e-30)
}
fn rel_linf(got: &[f32], want: &[f32], what: &str) -> f32 {
assert_eq!(got.len(), want.len(), "{what}: length");
let scale = want.iter().fold(0.0f32, |m, v| m.max(v.abs()));
assert!(scale > 1e-6, "{what}: the per-token reference is all ~zero");
got.iter()
.zip(want)
.fold(0.0f32, |m, (g, w)| m.max((g - w).abs()))
/ scale
}
fn argmax(v: &[f32]) -> usize {
v.iter()
.enumerate()
.fold((0, f32::NEG_INFINITY), |(bi, bv), (i, &x)| {
if x > bv {
(i, x)
} else {
(bi, bv)
}
})
.0
}
fn assert_logits_agree(batched: &[f32], per_token: &[f32], what: &str, b: Budget) {
let cos = cosine(batched, per_token);
let linf = rel_linf(batched, per_token, what);
let (ab, ap) = (argmax(batched), argmax(per_token));
println!(
"[3596] {what}: argmax batched {ab} / per-token {ap}, cosine {cos:.7}, rel L∞ {linf:.3e}"
);
assert_eq!(ab, ap, "{what}: argmax differs");
assert!(cos >= b.cosine, "{what}: cosine {cos} < {}", b.cosine);
assert!(linf <= b.logits, "{what}: rel L∞ {linf} > {}", b.logits);
}
fn download(buf: &trueno_gpu::driver::GpuBuffer<f32>, elems: usize) -> Vec<f32> {
let mut v = vec![0.0f32; buf.len()];
buf.copy_to_host(&mut v).expect("download state");
v.truncate(elems);
v
}
fn assert_states_agree(
gpu: &mut Qwen35CudaModel<'_>,
batched: &Qwen35CudaState,
per_token: &Qwen35CudaState,
what: &str,
budget: Budget,
) {
gpu.executor_mut().sync_stream().expect("sync");
assert_eq!(batched.kv_len, per_token.kv_len, "{what}: kv_len");
let mut worst = (0.0f32, String::new());
for il in 0..batched.conv.len() {
let mut pairs = Vec::new();
if let (Some((kb, vb)), Some((kp, vp))) = (&batched.kv[il], &per_token.kv[il]) {
let n = per_token.kv_len * per_token.kv_row;
pairs.push(("k", download(kb, n), download(kp, n)));
pairs.push(("v", download(vb, n), download(vp, n)));
} else {
pairs.push((
"conv",
download(&batched.conv[il], batched.conv_len),
download(&per_token.conv[il], per_token.conv_len),
));
pairs.push((
"ssm",
download(&batched.ssm[il], batched.ssm_len),
download(&per_token.ssm[il], per_token.ssm_len),
));
}
for (kind, b, p) in pairs {
let linf = rel_linf(&b, &p, &format!("{what} layer {il} {kind}"));
if linf > worst.0 {
worst = (linf, format!("layer {il} {kind}"));
}
}
}
println!(
"[3596] {what}: worst state rel L∞ {:.3e} at {}",
worst.0, worst.1
);
assert!(
worst.0 <= budget.state,
"{what}: state rel L∞ {} at {} > {}",
worst.0,
worst.1,
budget.state
);
}
fn tokens(n: usize, vocab: usize, seed: u32) -> Vec<u32> {
let mut s = seed;
(0..n)
.map(|_| {
s = s.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
1000 + (s >> 8) % (vocab as u32 - 2000)
})
.collect()
}
fn batched_equals_per_token(model_path: &str, n: usize, attention: super::PrefillAttention) {
batched_equals_per_token_rows(model_path, n, attention, None);
}
fn batched_equals_per_token_rows(
model_path: &str,
n: usize,
attention: super::PrefillAttention,
chunk_rows: Option<usize>,
) {
super::ATTENTION_OVERRIDE.with(|c| c.set(Some(attention)));
let gemm = crate::cuda::QWEN35_PREFILL_GEMM_OVERRIDE.with(|c| {
c.set(Some(c.get().unwrap_or(crate::cuda::Qwen35PrefillGemm::F32)));
c.get()
});
let fallback = F16_FALLBACK.with(std::cell::Cell::get);
let f16 = gemm == Some(crate::cuda::Qwen35PrefillGemm::F16) && !fallback;
let b = match (attention, f16) {
(_, true) => F16_GEMM_BUDGET,
(super::PrefillAttention::CublasF32, _) => F32_BUDGET,
(super::PrefillAttention::FlashF16In, _) => FLASH_BUDGET,
};
if !std::path::Path::new(model_path).exists() {
eprintln!("SKIP: {model_path} is absent");
return;
}
let executor = crate::cuda_executor_or_skip!(0);
let mapped = crate::gguf::MappedGGUFModel::from_path(model_path).expect("map the GGUF");
let base = Qwen35Model::create_base_model(&mapped.model, mapped.data()).expect("base");
let qwen =
Qwen35Model::from_model_and_layers(&base, &mapped.model, mapped.data()).expect("qwen35");
let mut gpu = Qwen35CudaModel::with_max_seq_len(&qwen, executor, n + 2).expect("gpu model");
if let Some(rows) = chunk_rows {
gpu.set_prefill_chunk_rows(rows);
}
if fallback {
gpu.executor.set_qwen35_prefill_f16(false);
}
assert_eq!(
gpu.executor.qwen35_prefill_f16(),
f16,
"f16 prefill GEMM armed"
);
gpu.set_prefill_attention(attention);
assert_eq!(gpu.prefill_attention_mode(), attention);
let vocab = base.config.vocab_size;
let prompt = tokens(n, vocab, 0x3596_0100 ^ n as u32);
let mut per_token = gpu.new_state().expect("state");
let mut want = Vec::new();
for (pos, &t) in prompt.iter().enumerate() {
want = gpu
.forward_single(t, &mut per_token, pos)
.expect("forward_single");
}
let mut batched = gpu.new_state().expect("state");
let rows = gpu.prefill_chunk_rows(n);
let passes = super::attention_rows_for(gpu.dims, n, gpu.prefill_rows);
let got = gpu.prefill(&prompt, &mut batched, 0).expect("prefill");
let what = format!(
"{model_path} n={n} (chunk rows {rows}, attention {}, rows/pass {passes})",
attention.as_str()
);
assert_logits_agree(&got, &want, &format!("{what} last logits"), b);
assert_states_agree(&mut gpu, &batched, &per_token, &what, b);
let cut = n / 3;
let mut split = gpu.new_state().expect("state");
let _ = gpu
.prefill(&prompt[..cut], &mut split, 0)
.expect("prefill part 1");
let got_split = gpu
.prefill(&prompt[cut..], &mut split, cut)
.expect("prefill part 2");
assert_logits_agree(&got_split, &want, &format!("{what} split at {cut}"), b);
assert_states_agree(
&mut gpu,
&split,
&per_token,
&format!("{what} split at {cut}"),
b,
);
let next = argmax(&want) as u32;
let step_b = gpu
.forward_single(next, &mut batched, n)
.expect("decode from batched");
let step_p = gpu
.forward_single(next, &mut per_token, n)
.expect("decode from per-token");
assert_logits_agree(&step_b, &step_p, &format!("{what} decode step"), b);
super::ATTENTION_OVERRIDE.with(|c| c.set(None));
}
#[test]
#[serial_test::serial]
fn qwen35_prefill_equals_per_token_at_64_positions_0_8b() {
batched_equals_per_token(MODEL_0_8B, 64, super::PrefillAttention::CublasF32);
}
#[test]
#[serial_test::serial]
fn qwen35_flash_prefill_equals_per_token_at_64_positions_0_8b() {
batched_equals_per_token(MODEL_0_8B, 64, super::PrefillAttention::FlashF16In);
}
#[test]
#[serial_test::serial]
fn qwen35_flash_prefill_equals_per_token_across_a_chunk_boundary_0_8b() {
batched_equals_per_token(MODEL_0_8B, 600, super::PrefillAttention::FlashF16In);
}
std::thread_local! {
static F16_FALLBACK: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
}
#[test]
fn f16_prewarm_fits_only_with_a_gib_to_spare() {
const GIB: usize = 1 << 30;
assert!(super::f16_prewarm_fits(7 * GIB, 8 * GIB));
assert!(!super::f16_prewarm_fits(7 * GIB + 1, 8 * GIB));
assert!(!super::f16_prewarm_fits(0, GIB - 1));
assert!(!super::f16_prewarm_fits(usize::MAX, usize::MAX));
}
#[test]
#[serial_test::serial]
fn qwen35_f16_mode_falls_back_to_f32_when_the_prewarm_is_not_armed_0_8b() {
F16_FALLBACK.with(|c| c.set(true));
f16_gemm_batched_equals_per_token(64);
F16_FALLBACK.with(|c| c.set(false));
}
fn f16_gemm_batched_equals_per_token(n: usize) {
crate::cuda::QWEN35_PREFILL_GEMM_OVERRIDE
.with(|c| c.set(Some(crate::cuda::Qwen35PrefillGemm::F16)));
batched_equals_per_token(MODEL_0_8B, n, super::PrefillAttention::CublasF32);
crate::cuda::QWEN35_PREFILL_GEMM_OVERRIDE.with(|c| c.set(None));
}
#[test]
#[serial_test::serial]
fn qwen35_f16_gemm_prefill_equals_per_token_at_64_positions_0_8b() {
f16_gemm_batched_equals_per_token(64);
}
#[test]
#[serial_test::serial]
fn qwen35_f16_gemm_prefill_equals_per_token_across_a_chunk_boundary_0_8b() {
f16_gemm_batched_equals_per_token(600);
}
#[test]
#[serial_test::serial]
fn qwen35_prefill_equals_per_token_across_a_chunk_boundary_0_8b() {
batched_equals_per_token(MODEL_0_8B, 600, super::PrefillAttention::CublasF32);
}
#[test]
#[serial_test::serial]
fn qwen35_prefill_equals_per_token_with_many_attention_passes_0_8b() {
let rows = 37usize;
let budget = 4 * 4 * rows * 600;
super::SCORES_BUDGET_OVERRIDE.with(|c| c.set(Some(budget)));
batched_equals_per_token(MODEL_0_8B, 600, super::PrefillAttention::CublasF32);
super::SCORES_BUDGET_OVERRIDE.with(|c| c.set(None));
}
#[test]
#[serial_test::serial]
fn qwen35_prefill_equals_per_token_with_unified_memory_chunk_rows_0_8b() {
batched_equals_per_token_rows(
MODEL_0_8B,
600,
super::PrefillAttention::FlashF16In,
Some(super::UNIFIED_PREFILL_CHUNK_ROWS),
);
batched_equals_per_token_rows(
MODEL_0_8B,
600,
super::PrefillAttention::FlashF16In,
Some(64),
);
}
#[test]
#[serial_test::serial]
fn qwen35_prefill_refuses_an_empty_prompt_and_positions_past_the_cache() {
if !std::path::Path::new(MODEL_0_8B).exists() {
eprintln!("SKIP: {MODEL_0_8B} is absent");
return;
}
let executor = crate::cuda_executor_or_skip!(0);
let mapped = crate::gguf::MappedGGUFModel::from_path(MODEL_0_8B).expect("map");
let base = Qwen35Model::create_base_model(&mapped.model, mapped.data()).expect("base");
let qwen =
Qwen35Model::from_model_and_layers(&base, &mapped.model, mapped.data()).expect("qwen35");
let mut gpu = Qwen35CudaModel::with_max_seq_len(&qwen, executor, 8).expect("gpu model");
let mut state = gpu.new_state().expect("state");
assert!(
gpu.prefill(&[], &mut state, 0).is_err(),
"an empty prompt must be refused"
);
assert!(
gpu.prefill(&[1000; 9], &mut state, 0).is_err(),
"9 positions into an 8-row cache must be refused, not truncated"
);
assert!(
gpu.prefill(&[u32::MAX], &mut state, 0).is_err(),
"a token outside the vocabulary must be refused"
);
assert!(
gpu.prefill_logits_at(&[1000; 4], &mut state, 0, &[2, 1])
.is_err(),
"descending positions must be refused"
);
assert!(
gpu.prefill_logits_at(&[1000; 4], &mut state, 0, &[4])
.is_err(),
"a position past the prompt must be refused"
);
let mut fresh = gpu.new_state().expect("state");
let got = gpu
.prefill_logits_at(&[1000, 1001, 1002, 1003], &mut fresh, 0, &[0, 3])
.expect("two requested rows");
assert_eq!(got.len(), 2, "one logits vector per requested position");
assert!(
gpu.prefill(&[1000], &mut fresh, 0).is_err(),
"rewinding to pos0 0 over 4 written rows must be refused"
);
assert!(
gpu.prefill(&[1000], &mut fresh, 5).is_err(),
"skipping to pos0 5 past 4 written rows must be refused"
);
let mut unused = gpu.new_state().expect("state");
assert!(
gpu.prefill(&[1000], &mut unused, 1).is_err(),
"a fresh state starts at 0"
);
gpu.prefill(&[1000], &mut fresh, 4)
.expect("pos0 == kv_len continues the state");
}
#[test]
fn qwen35_prefill_attention_prefers_f32_then_flash_and_the_environment_pins_one() {
use super::{
attention_candidates,
PrefillAttention::{CublasF32, FlashF16In},
};
let rows: [(Option<&str>, bool, &[super::PrefillAttention]); 7] = [
(None, true, &[CublasF32, FlashF16In]),
(None, false, &[CublasF32]),
(Some("f32"), true, &[CublasF32]),
(Some("flash"), true, &[FlashF16In]),
(Some("flash"), false, &[CublasF32]),
(Some("fast"), true, &[CublasF32, FlashF16In]),
(Some(""), false, &[CublasF32]),
];
for (forced, flash, want) in rows {
assert_eq!(
attention_candidates(forced, flash),
want,
"{forced:?}, flash supported {flash}"
);
}
}