pub mod codebook;
pub mod kquant;
pub mod legacy;
#[derive(Debug, Clone, Copy)]
pub struct MatvecKind {
pub name: &'static str,
pub module_name: &'static str,
pub fn_name: &'static str,
pub src: &'static str,
}
pub const KINDS: &[MatvecKind] = &[
MatvecKind {
name: "Q8_0",
module_name: "ferrox_q8_0",
fn_name: "q8_0_matvec",
src: legacy::Q8_0_MATVEC_KERNEL_SRC,
},
MatvecKind {
name: "Q4_0",
module_name: "ferrox_q4_0",
fn_name: "q4_0_matvec",
src: legacy::Q4_0_MATVEC_KERNEL_SRC,
},
MatvecKind {
name: "Q5_0",
module_name: "ferrox_q5_0",
fn_name: "q5_0_matvec",
src: legacy::Q5_0_MATVEC_KERNEL_SRC,
},
MatvecKind {
name: "Q2_K",
module_name: "ferrox_q2_k",
fn_name: "q2_k_matvec",
src: kquant::Q2_K_MATVEC_KERNEL_SRC,
},
MatvecKind {
name: "Q3_K",
module_name: "ferrox_q3_k",
fn_name: "q3_k_matvec",
src: kquant::Q3_K_MATVEC_KERNEL_SRC,
},
MatvecKind {
name: "Q4_K",
module_name: "ferrox_q4_k",
fn_name: "q4_k_matvec",
src: kquant::Q4_K_MATVEC_KERNEL_SRC,
},
MatvecKind {
name: "Q5_K",
module_name: "ferrox_q5_k",
fn_name: "q5_k_matvec",
src: kquant::Q5_K_MATVEC_KERNEL_SRC,
},
MatvecKind {
name: "Q6_K",
module_name: "ferrox_q6_k",
fn_name: "q6_k_matvec",
src: kquant::Q6_K_MATVEC_KERNEL_SRC,
},
MatvecKind {
name: "IQ4_NL",
module_name: "ferrox_iq4_nl",
fn_name: "iq4_nl_matvec",
src: codebook::IQ4_NL_MATVEC_KERNEL_SRC,
},
MatvecKind {
name: "IQ4_XS",
module_name: "ferrox_iq4_xs",
fn_name: "iq4_xs_matvec",
src: codebook::IQ4_XS_MATVEC_KERNEL_SRC,
},
MatvecKind {
name: "MXFP4",
module_name: "ferrox_mxfp4",
fn_name: "mxfp4_matvec",
src: codebook::MXFP4_MATVEC_KERNEL_SRC,
},
];
pub fn kind_by_name(name: &str) -> Option<&'static MatvecKind> {
KINDS.iter().find(|k| k.name == name)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn every_row_names_an_entry_point_its_source_defines() {
for k in KINDS {
assert!(
k.src.contains(&format!("void {}(", k.fn_name)),
"{}: the source in {} does not define {}",
k.name,
k.module_name,
k.fn_name
);
assert!(
k.src.contains("int n_blocks_per_row"),
"{}: does not take the launch geometry every caller supplies",
k.name
);
}
for (i, a) in KINDS.iter().enumerate() {
for b in &KINDS[i + 1..] {
assert_ne!(
a.module_name, b.module_name,
"{} and {} collide in the process-wide NVRTC module cache",
a.name, b.name
);
assert_ne!(a.fn_name, b.fn_name, "{} vs {}", a.name, b.name);
}
}
}
#[test]
fn a_kind_with_no_matvec_does_not_resolve() {
for absent in ["Q5_1", "Q4_1", "Q8_1", "IQ1_S", "IQ2_XXS", "IQ3_S"] {
assert!(
kind_by_name(absent).is_none(),
"{absent} resolved to a CUDA matvec that does not exist"
);
}
}
#[test]
fn every_matvec_strides_by_the_real_block_geometry() {
for k in KINDS {
let mm = crate::mul_mm::kind_by_name(k.name)
.unwrap_or_else(|| panic!("{}: a matvec with no mul_mm row", k.name));
let row_step = k
.src
.lines()
.find(|l| l.contains("row_ptr +"))
.unwrap_or_else(|| panic!("{}: no row-pointer arithmetic", k.name));
assert!(
row_step.contains(&format!("* {};", mm.block_bytes)),
"{}: strides the row by something other than {} bytes: {}",
k.name,
mm.block_bytes,
row_step.trim()
);
let x_step = k
.src
.lines()
.find(|l| l.contains("base = ") && l.trim_end().ends_with(';'))
.unwrap_or_else(|| panic!("{}: no activation-base arithmetic", k.name));
assert!(
x_step.contains(&format!("* {};", mm.block_elems)),
"{}: steps the activation by something other than {} elements: {}",
k.name,
mm.block_elems,
x_step.trim()
);
}
}
#[test]
fn every_embedded_codebook_is_the_mul_mm_codebook() {
let mut checked = 0usize;
for mm in crate::mul_mm::KINDS {
let Some(cb) = mm.codebook else { continue };
let Some(mv) = kind_by_name(mm.name) else {
panic!(
"{}: has a mul_mm codebook kernel but no matvec, which is the \
split that decomposes a prefill into one launch per position",
mm.name
);
};
let decl = format!("__constant__ float {}[16] = {{", cb.c_name);
let at = mv.src.find(&decl).unwrap_or_else(|| {
panic!("{}: the matvec does not declare {}", mm.name, cb.c_name)
});
let body = &mv.src[at + decl.len()..];
let body = &body[..body.find('}').expect("unterminated codebook")];
let got: Vec<f32> = body
.split(',')
.map(|t| {
t.trim()
.trim_end_matches('f')
.parse::<f32>()
.unwrap_or_else(|e| panic!("{}: {t:?}: {e}", mm.name))
})
.collect();
assert_eq!(got.len(), 16, "{}: codebook is not 16 entries", mm.name);
for (i, (g, w)) in got.iter().zip(cb.values.iter()).enumerate() {
assert_eq!(
g.to_bits(),
w.to_bits(),
"{}: codebook entry {i}: matvec has {g}, mul_mm emits {w}",
mm.name
);
}
checked += 1;
}
assert!(checked >= 3, "the codebook kinds stopped being checked");
}
#[test]
fn the_matvec_table_and_the_mul_mm_table_name_the_same_kinds() {
let mm: Vec<&str> = crate::mul_mm::KINDS.iter().map(|k| k.name).collect();
let mv: Vec<&str> = KINDS.iter().map(|k| k.name).collect();
assert_eq!(mm, mv, "the CUDA decode and prefill tables have diverged");
}
}
#[cfg(all(test, feature = "cuda"))]
mod hardware_tests {
#[test]
#[ignore = "requires real CUDA hardware. Q2_K, Q3_K, IQ4_NL, IQ4_XS and MXFP4 have NEVER executed on a GPU; Q5_0 has not either. Run with --ignored on a CUDA-capable machine and record the result before any doc claims those kinds decode on CUDA"]
fn every_cuda_matvec_matches_the_cpu_reference() {
type Dot = fn(&[u8], &[f32]) -> f32;
type Launch =
fn(&[u8], &[f32], usize, usize, usize) -> Result<Vec<f32>, crate::gpu::CudaError>;
let oracles: &[(&str, Dot, Launch)] = &[
(
"Q8_0",
ferrox_quant::dot_q8_0_f32_scalar,
crate::gpu::launch_q8_0_matvec,
),
(
"Q4_0",
ferrox_quant::dot_q4_0_f32_scalar,
crate::gpu::launch_q4_0_matvec,
),
(
"Q5_0",
ferrox_quant::dot_q5_0_f32_scalar,
crate::gpu::launch_q5_0_matvec,
),
(
"Q2_K",
ferrox_quant::dot_q2_k_f32_scalar,
crate::gpu::launch_q2_k_matvec,
),
(
"Q3_K",
ferrox_quant::dot_q3_k_f32_scalar,
crate::gpu::launch_q3_k_matvec,
),
(
"Q4_K",
ferrox_quant::dot_q4_k_f32_scalar,
crate::gpu::launch_q4_k_matvec,
),
(
"Q5_K",
ferrox_quant::dot_q5_k_f32_scalar,
crate::gpu::launch_q5_k_matvec,
),
(
"Q6_K",
ferrox_quant::dot_q6_k_f32_scalar,
crate::gpu::launch_q6_k_matvec,
),
(
"IQ4_NL",
ferrox_quant::dot_iq4_nl_f32_scalar,
crate::gpu::launch_iq4_nl_matvec,
),
(
"IQ4_XS",
ferrox_quant::dot_iq4_xs_f32_scalar,
crate::gpu::launch_iq4_xs_matvec,
),
(
"MXFP4",
ferrox_quant::dot_mxfp4_gguf_f32_scalar,
crate::gpu::launch_mxfp4_matvec,
),
];
for mm in crate::mul_mm::KINDS {
let (_, dot, launch) = oracles
.iter()
.find(|(name, _, _)| *name == mm.name)
.unwrap_or_else(|| panic!("{}: a CUDA matvec with no CPU oracle", mm.name));
let rows = 7;
let cols = mm.block_elems * 2;
let blocks_per_row = cols / mm.block_elems;
let row_bytes = blocks_per_row * mm.block_bytes;
let weights = crate::mul_mm_ref::fixtures::weights(mm, rows, cols, 31337);
let x: Vec<f32> = (0..cols).map(|i| ((i as f32) * 0.09).sin()).collect();
let expected: Vec<f32> = (0..rows)
.map(|r| dot(&weights[r * row_bytes..(r + 1) * row_bytes], &x))
.collect();
let got = launch(&weights, &x, rows, row_bytes, blocks_per_row)
.expect("kernel launch must succeed on real CUDA hardware");
assert_eq!(got.len(), expected.len());
for (i, (g, w)) in got.iter().zip(expected.iter()).enumerate() {
let scale = w.abs().max(1.0);
assert!(
(g - w).abs() <= 1e-3 * scale,
"{} row {i}: GPU={g} CPU reference={w}",
mm.name
);
}
}
}
}