#![cfg(feature = "vulkan")]
use hanzo_ml::quantized::{GgmlDType, QMatMul, QStorage, QTensor};
use hanzo_ml::{Device, Module, Tensor};
use std::sync::Arc;
static ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
fn pseudo(i: usize) -> f32 {
let mut z = (i as u64).wrapping_add(0x9E37_79B9_7F4A_7C15);
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^= z >> 31;
((z >> 40) as f32 / (1u32 << 24) as f32) * 2.0 - 1.0
}
struct ErrStats {
max_abs: f32,
max_rel: f32,
rms: f32,
quant_max_abs: f32, }
fn quantizable(dtype: GgmlDType) -> bool {
!matches!(
dtype,
GgmlDType::IQ4_NL
| GgmlDType::IQ4_XS
| GgmlDType::TQ2_0
| GgmlDType::IQ2_XXS
| GgmlDType::IQ2_S
| GgmlDType::IQ3_XXS
| GgmlDType::IQ3_S
| GgmlDType::IQ1_S
| GgmlDType::IQ1_M
| GgmlDType::IQ2_XS
)
}
fn pbyte(i: usize) -> u8 {
(((pseudo(i) * 0.5 + 0.5) * 256.0) as i32).clamp(0, 255) as u8
}
fn synth_decode_only(dtype: GgmlDType, nout: usize, k: usize) -> Vec<u8> {
use half::f16;
let mut out: Vec<u8> = Vec::new();
let mut c = 0usize;
match dtype {
GgmlDType::IQ4_NL => {
for _ in 0..nout * (k / 32) {
out.extend_from_slice(&f16::from_f32(pseudo(c) * 0.003).to_le_bytes());
c += 1;
for _ in 0..16 {
out.push(pbyte(c));
c += 1;
}
}
}
GgmlDType::IQ4_XS => {
for _ in 0..nout * (k / 256) {
out.extend_from_slice(&f16::from_f32(pseudo(c) * 1e-4).to_le_bytes());
c += 1;
for _ in 0..6 {
out.push(pbyte(c)); c += 1;
}
for _ in 0..128 {
out.push(pbyte(c)); c += 1;
}
}
}
GgmlDType::TQ2_0 => {
for _ in 0..nout * (k / 256) {
for _ in 0..64 {
out.push(pbyte(c)); c += 1;
}
out.extend_from_slice(&f16::from_f32(pseudo(c) * 0.3).to_le_bytes());
c += 1;
}
}
GgmlDType::IQ2_XXS => {
for _ in 0..nout * (k / 256) {
out.extend_from_slice(&f16::from_f32(pseudo(c) * 0.05).to_le_bytes());
c += 1;
for _ in 0..64 {
out.push(pbyte(c)); c += 1;
}
}
}
GgmlDType::IQ2_XS => {
for _ in 0..nout * (k / 256) {
out.extend_from_slice(&f16::from_f32(pseudo(c) * 0.05).to_le_bytes());
c += 1;
for _ in 0..72 {
out.push(pbyte(c)); c += 1;
}
}
}
GgmlDType::IQ2_S => {
for _ in 0..nout * (k / 256) {
out.extend_from_slice(&f16::from_f32(pseudo(c) * 0.05).to_le_bytes());
c += 1;
for _ in 0..80 {
out.push(pbyte(c));
c += 1;
}
}
}
GgmlDType::IQ3_XXS => {
for _ in 0..nout * (k / 256) {
out.extend_from_slice(&f16::from_f32(pseudo(c) * 0.02).to_le_bytes());
c += 1;
for _ in 0..96 {
out.push(pbyte(c));
c += 1;
}
}
}
GgmlDType::IQ3_S => {
for _ in 0..nout * (k / 256) {
out.extend_from_slice(&f16::from_f32(pseudo(c) * 0.004).to_le_bytes());
c += 1;
for _ in 0..108 {
out.push(pbyte(c));
c += 1;
}
}
}
GgmlDType::IQ1_S => {
for _ in 0..nout * (k / 256) {
out.extend_from_slice(&f16::from_f32(pseudo(c) * 0.05).to_le_bytes());
c += 1;
for _ in 0..48 {
out.push(pbyte(c));
c += 1;
}
}
}
GgmlDType::IQ1_M => {
for _ in 0..nout * (k / 256) {
for _ in 0..48 {
out.push(pbyte(c)); c += 1;
}
out.extend_from_slice(&[0u8, 0u8, 0u8, 0u8, 0u8, 0xC0u8, 0u8, 0x20u8]);
}
}
_ => panic!("synth_decode_only: {dtype:?} is not a decode-only type"),
}
out
}
fn weight_bytes(
dtype: GgmlDType,
nout: usize,
k: usize,
) -> hanzo_ml::Result<(Vec<u8>, Vec<f32>, Vec<f32>)> {
let cpu = Device::Cpu;
if quantizable(dtype) {
let w_host: Vec<f32> = (0..nout * k).map(|i| pseudo(i) * 0.5).collect();
let w_t = Tensor::from_vec(w_host.clone(), (nout, k), &cpu)?;
let q = QTensor::quantize(&w_t, dtype)?;
let w_deq: Vec<f32> = q.dequantize(&cpu)?.flatten_all()?.to_vec1::<f32>()?;
Ok((q.data()?.into_owned(), w_deq, w_host))
} else {
let raw = synth_decode_only(dtype, nout, k);
let qs = QStorage::from_data(std::borrow::Cow::Owned(raw.clone()), &cpu, dtype)?;
let q = QTensor::new(qs, (nout, k))?;
let w_deq: Vec<f32> = q.dequantize(&cpu)?.flatten_all()?.to_vec1::<f32>()?;
Ok((raw, w_deq.clone(), w_deq))
}
}
fn run_case(dev: &Device, dtype: GgmlDType, nout: usize, k: usize) -> hanzo_ml::Result<ErrStats> {
let x_host: Vec<f32> = (0..k).map(|i| pseudo(i + 1_000_003)).collect();
let (raw, w_deq, w_host) = weight_bytes(dtype, nout, k)?;
let vk = dev.as_vulkan_device()?;
let wq = match dtype {
GgmlDType::Q6K => vk.quantize_q6k(&raw, nout, k)?,
GgmlDType::Q3K => vk.quantize_q3k(&raw, nout, k)?,
GgmlDType::IQ2_XXS => vk.quantize_iq2xxs(&raw, nout, k)?,
GgmlDType::IQ2_XS => vk.quantize_iq2xs(&raw, nout, k)?,
GgmlDType::IQ1_M => vk.quantize_iq1m(&raw, nout, k)?,
GgmlDType::IQ1_S => vk.quantize_iq1s(&raw, nout, k)?,
GgmlDType::IQ3_S => vk.quantize_iq3s(&raw, nout, k)?,
GgmlDType::IQ3_XXS => vk.quantize_iq3xxs(&raw, nout, k)?,
GgmlDType::IQ2_S => vk.quantize_iq2s(&raw, nout, k)?,
_ => vk.upload_qweight(&raw)?,
};
let y_gpu: Vec<f32> = match dtype {
GgmlDType::Q4_0 => vk.matvec_q4_0(&wq, &x_host, nout, k)?,
GgmlDType::Q8_0 => vk.matvec_q8_0(&wq, &x_host, nout, k)?,
GgmlDType::Q4K => vk.matvec_q4k_scalar(&wq, &x_host, nout, k)?,
GgmlDType::Q5K => vk.matvec_q5k(&wq, &x_host, nout, k)?,
GgmlDType::Q6K => vk.matvec_q6k(&wq, &x_host, nout, k)?,
GgmlDType::Q2K => vk.matvec_q2k(&wq, &x_host, nout, k)?,
GgmlDType::Q3K => vk.matvec_q3k(&wq, &x_host, nout, k)?,
GgmlDType::IQ4_XS => vk.matvec_iq4xs(&wq, &x_host, nout, k)?,
GgmlDType::IQ4_NL => vk.matvec_iq4nl(&wq, &x_host, nout, k)?,
GgmlDType::IQ2_XXS => vk.matvec_iq2xxs(&wq, &x_host, nout, k)?,
GgmlDType::IQ2_XS => vk.matvec_iq2xs(&wq, &x_host, nout, k)?,
GgmlDType::IQ1_M => vk.matvec_iq1m(&wq, &x_host, nout, k)?,
GgmlDType::IQ1_S => vk.matvec_iq1s(&wq, &x_host, nout, k)?,
GgmlDType::IQ3_S => vk.matvec_iq3s(&wq, &x_host, nout, k)?,
GgmlDType::IQ3_XXS => vk.matvec_iq3xxs(&wq, &x_host, nout, k)?,
GgmlDType::IQ2_S => vk.matvec_iq2s(&wq, &x_host, nout, k)?,
GgmlDType::TQ2_0 => vk.matvec_tq2_0(&wq, &x_host, nout, k)?,
_ => panic!("unsupported dtype in run_case: {dtype:?}"),
};
assert_eq!(y_gpu.len(), nout);
let mut max_abs = 0f32;
let mut sse = 0f64;
let mut ref_sq = 0f64; let mut quant_max_abs = 0f32;
for n in 0..nout {
let mut ref_deq = 0f64;
let mut ref_orig = 0f64;
for j in 0..k {
ref_deq += w_deq[n * k + j] as f64 * x_host[j] as f64;
ref_orig += w_host[n * k + j] as f64 * x_host[j] as f64;
}
let g = y_gpu[n] as f64;
max_abs = max_abs.max((g - ref_deq).abs() as f32);
sse += (g - ref_deq) * (g - ref_deq);
ref_sq += ref_deq * ref_deq;
quant_max_abs = quant_max_abs.max((g - ref_orig).abs() as f32);
}
let ref_rms = (ref_sq / nout as f64).sqrt();
let err_rms = (sse / nout as f64).sqrt();
let max_rel = (err_rms / ref_rms.max(1e-9)) as f32;
Ok(ErrStats {
max_abs,
max_rel,
rms: err_rms as f32,
quant_max_abs,
})
}
fn gpu() -> Option<Device> {
match Device::new_vulkan(0) {
Ok(d) => Some(d),
Err(e) => {
eprintln!("[vulkan_quant_tests] no Vulkan GPU ({e}); skipping");
None
}
}
}
const SHAPES: &[(usize, usize)] = &[
(2048, 2048), (4096, 2048), (2048, 4096), (512, 256), ];
#[test]
fn vulkan_matvec_q4_0_matches_cpu() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else { return Ok(()) };
for &(nout, k) in SHAPES {
let s = run_case(&dev, GgmlDType::Q4_0, nout, k)?;
println!(
"Q4_0 nout={nout:5} k={k:5} max_abs={:.3e} max_rel={:.3e} rms={:.3e} (quant err vs f32: {:.3e})",
s.max_abs, s.max_rel, s.rms, s.quant_max_abs
);
assert!(
s.max_rel < 1e-3 && s.max_abs < 1e-3,
"Q4_0 GPU/CPU mismatch too large: max_abs={} max_rel={}",
s.max_abs,
s.max_rel
);
}
Ok(())
}
#[test]
fn vulkan_matvec_q8_0_matches_cpu() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else { return Ok(()) };
for &(nout, k) in SHAPES {
let s = run_case(&dev, GgmlDType::Q8_0, nout, k)?;
println!(
"Q8_0 nout={nout:5} k={k:5} max_abs={:.3e} max_rel={:.3e} rms={:.3e} (quant err vs f32: {:.3e})",
s.max_abs, s.max_rel, s.rms, s.quant_max_abs
);
assert!(
s.max_rel < 1e-3 && s.max_abs < 1e-3,
"Q8_0 GPU/CPU mismatch too large: max_abs={} max_rel={}",
s.max_abs,
s.max_rel
);
}
Ok(())
}
#[test]
fn vulkan_matvec_q4k_matches_cpu() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else { return Ok(()) };
for &(nout, k) in SHAPES {
let s = run_case(&dev, GgmlDType::Q4K, nout, k)?;
println!(
"Q4_K nout={nout:5} k={k:5} max_abs={:.3e} max_rel={:.3e} rms={:.3e} (quant err vs f32: {:.3e})",
s.max_abs, s.max_rel, s.rms, s.quant_max_abs
);
assert!(
s.max_rel < 1e-3 && s.max_abs < 1e-3,
"Q4_K GPU/CPU mismatch too large: max_abs={} max_rel={}",
s.max_abs,
s.max_rel
);
}
Ok(())
}
#[test]
fn vulkan_matvec_q5k_matches_cpu() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else { return Ok(()) };
for &(nout, k) in SHAPES {
let s = run_case(&dev, GgmlDType::Q5K, nout, k)?;
println!(
"Q5_K nout={nout:5} k={k:5} max_abs={:.3e} max_rel={:.3e} rms={:.3e} (quant err vs f32: {:.3e})",
s.max_abs, s.max_rel, s.rms, s.quant_max_abs
);
assert!(
s.max_rel < 1e-3 && s.max_abs < 1e-3,
"Q5_K GPU/CPU mismatch too large: max_abs={} max_rel={}",
s.max_abs,
s.max_rel
);
}
Ok(())
}
#[test]
fn vulkan_matvec_q6k_matches_cpu() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else { return Ok(()) };
for &(nout, k) in SHAPES {
let s = run_case(&dev, GgmlDType::Q6K, nout, k)?;
println!(
"Q6_K nout={nout:5} k={k:5} max_abs={:.3e} max_rel={:.3e} rms={:.3e} (quant err vs f32: {:.3e})",
s.max_abs, s.max_rel, s.rms, s.quant_max_abs
);
assert!(
s.max_rel < 1e-3 && s.max_abs < 1e-3,
"Q6_K GPU/CPU mismatch too large: max_abs={} max_rel={}",
s.max_abs,
s.max_rel
);
}
Ok(())
}
#[test]
fn vulkan_matvec_q2k_matches_cpu() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else { return Ok(()) };
for &(nout, k) in SHAPES {
let s = run_case(&dev, GgmlDType::Q2K, nout, k)?;
println!(
"Q2_K nout={nout:5} k={k:5} max_abs={:.3e} max_rel={:.3e} rms={:.3e} (quant err vs f32: {:.3e})",
s.max_abs, s.max_rel, s.rms, s.quant_max_abs
);
assert!(
s.max_rel < 1e-3 && s.max_abs < 1e-3,
"Q2_K GPU/CPU mismatch too large: max_abs={} max_rel={}",
s.max_abs,
s.max_rel
);
}
Ok(())
}
#[test]
fn vulkan_matvec_q3k_matches_cpu() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else { return Ok(()) };
for &(nout, k) in SHAPES {
let s = run_case(&dev, GgmlDType::Q3K, nout, k)?;
println!(
"Q3_K nout={nout:5} k={k:5} max_abs={:.3e} max_rel={:.3e} rms={:.3e} (quant err vs f32: {:.3e})",
s.max_abs, s.max_rel, s.rms, s.quant_max_abs
);
assert!(
s.max_rel < 1e-3 && s.max_abs < 1e-3,
"Q3_K GPU/CPU mismatch too large: max_abs={} max_rel={}",
s.max_abs,
s.max_rel
);
}
Ok(())
}
#[test]
fn vulkan_matvec_iq4xs_matches_cpu() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else { return Ok(()) };
for &(nout, k) in SHAPES {
let s = run_case(&dev, GgmlDType::IQ4_XS, nout, k)?;
println!(
"IQ4_XS nout={nout:5} k={k:5} max_abs={:.3e} max_rel={:.3e} rms={:.3e} (quant err vs f32: {:.3e})",
s.max_abs, s.max_rel, s.rms, s.quant_max_abs
);
assert!(
s.max_rel < 1e-3 && s.max_abs < 1e-3,
"IQ4_XS GPU/CPU mismatch too large: max_abs={} max_rel={}",
s.max_abs,
s.max_rel
);
}
Ok(())
}
#[test]
fn vulkan_matvec_iq2xxs_matches_cpu() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else { return Ok(()) };
for &(nout, k) in SHAPES {
let s = run_case(&dev, GgmlDType::IQ2_XXS, nout, k)?;
println!(
"IQ2_XXS nout={nout:5} k={k:5} max_abs={:.3e} max_rel={:.3e} rms={:.3e} (quant err vs f32: {:.3e})",
s.max_abs, s.max_rel, s.rms, s.quant_max_abs
);
assert!(
s.max_rel < 1e-3 && s.max_abs < 1e-3,
"IQ2_XXS GPU/CPU mismatch too large: max_abs={} max_rel={}",
s.max_abs,
s.max_rel
);
}
Ok(())
}
#[test]
fn vulkan_matvec_iq2xs_matches_cpu() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else { return Ok(()) };
for &(nout, k) in SHAPES {
let s = run_case(&dev, GgmlDType::IQ2_XS, nout, k)?;
println!(
"IQ2_XS nout={nout:5} k={k:5} max_abs={:.3e} max_rel={:.3e} rms={:.3e} (quant err vs f32: {:.3e})",
s.max_abs, s.max_rel, s.rms, s.quant_max_abs
);
assert!(
s.max_rel < 1e-3 && s.max_abs < 1e-3,
"IQ2_XS GPU/CPU mismatch too large: max_abs={} max_rel={}",
s.max_abs,
s.max_rel
);
}
Ok(())
}
#[test]
fn vulkan_matvec_iq4nl_matches_cpu() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else { return Ok(()) };
for &(nout, k) in SHAPES {
let s = run_case(&dev, GgmlDType::IQ4_NL, nout, k)?;
println!(
"IQ4_NL nout={nout:5} k={k:5} max_abs={:.3e} max_rel={:.3e} rms={:.3e} (quant err vs f32: {:.3e})",
s.max_abs, s.max_rel, s.rms, s.quant_max_abs
);
assert!(
s.max_rel < 1e-3 && s.max_abs < 1e-3,
"IQ4_NL GPU/CPU mismatch too large: max_abs={} max_rel={}",
s.max_abs,
s.max_rel
);
}
Ok(())
}
#[test]
fn vulkan_matvec_tq2_0_matches_cpu() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else { return Ok(()) };
for &(nout, k) in SHAPES {
let s = run_case(&dev, GgmlDType::TQ2_0, nout, k)?;
println!(
"TQ2_0 nout={nout:5} k={k:5} max_abs={:.3e} max_rel={:.3e} rms={:.3e} (quant err vs f32: {:.3e})",
s.max_abs, s.max_rel, s.rms, s.quant_max_abs
);
assert!(
s.max_rel < 1e-3 && s.max_abs < 1e-3,
"TQ2_0 GPU/CPU mismatch too large: max_abs={} max_rel={}",
s.max_abs,
s.max_rel
);
}
Ok(())
}
#[test]
fn vulkan_matvec_iq2s_matches_cpu() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else { return Ok(()) };
for &(nout, k) in SHAPES {
let s = run_case(&dev, GgmlDType::IQ2_S, nout, k)?;
println!(
"IQ2_S nout={nout:5} k={k:5} max_abs={:.3e} max_rel={:.3e} rms={:.3e}",
s.max_abs, s.max_rel, s.rms
);
assert!(
s.max_rel < 1e-3 && s.max_abs < 1e-3,
"IQ2_S GPU/CPU mismatch: max_abs={} max_rel={}",
s.max_abs,
s.max_rel
);
}
Ok(())
}
#[test]
fn vulkan_matvec_iq3xxs_matches_cpu() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else { return Ok(()) };
for &(nout, k) in SHAPES {
let s = run_case(&dev, GgmlDType::IQ3_XXS, nout, k)?;
println!(
"IQ3_XXS nout={nout:5} k={k:5} max_abs={:.3e} max_rel={:.3e} rms={:.3e}",
s.max_abs, s.max_rel, s.rms
);
assert!(
s.max_rel < 1e-3 && s.max_abs < 1e-3,
"IQ3_XXS GPU/CPU mismatch: max_abs={} max_rel={}",
s.max_abs,
s.max_rel
);
}
Ok(())
}
#[test]
fn vulkan_matvec_iq3s_matches_cpu() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else { return Ok(()) };
for &(nout, k) in SHAPES {
let s = run_case(&dev, GgmlDType::IQ3_S, nout, k)?;
println!(
"IQ3_S nout={nout:5} k={k:5} max_abs={:.3e} max_rel={:.3e} rms={:.3e}",
s.max_abs, s.max_rel, s.rms
);
assert!(
s.max_rel < 1e-3 && s.max_abs < 1e-3,
"IQ3_S GPU/CPU mismatch: max_abs={} max_rel={}",
s.max_abs,
s.max_rel
);
}
Ok(())
}
#[test]
fn vulkan_matvec_iq1s_matches_cpu() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else { return Ok(()) };
for &(nout, k) in SHAPES {
let s = run_case(&dev, GgmlDType::IQ1_S, nout, k)?;
println!(
"IQ1_S nout={nout:5} k={k:5} max_abs={:.3e} max_rel={:.3e} rms={:.3e}",
s.max_abs, s.max_rel, s.rms
);
assert!(
s.max_rel < 1e-3 && s.max_abs < 1e-3,
"IQ1_S GPU/CPU mismatch: max_abs={} max_rel={}",
s.max_abs,
s.max_rel
);
}
Ok(())
}
#[test]
fn vulkan_matvec_iq1m_matches_cpu() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else { return Ok(()) };
for &(nout, k) in SHAPES {
let s = run_case(&dev, GgmlDType::IQ1_M, nout, k)?;
println!(
"IQ1_M nout={nout:5} k={k:5} max_abs={:.3e} max_rel={:.3e} rms={:.3e}",
s.max_abs, s.max_rel, s.rms
);
assert!(
s.max_rel < 1e-3 && s.max_abs < 1e-3,
"IQ1_M GPU/CPU mismatch: max_abs={} max_rel={}",
s.max_abs,
s.max_rel
);
}
Ok(())
}
fn end_to_end_case(dev: &Device, dtype: GgmlDType, nout: usize, k: usize) -> hanzo_ml::Result<f32> {
let x_host: Vec<f32> = (0..k).map(|i| pseudo(i + 7)).collect();
let (bytes, w_deq, _) = weight_bytes(dtype, nout, k)?;
let mut y_ref = vec![0f64; nout];
for n in 0..nout {
let mut acc = 0f64;
for j in 0..k {
acc += w_deq[n * k + j] as f64 * x_host[j] as f64;
}
y_ref[n] = acc;
}
let qs_vk = QStorage::from_data(std::borrow::Cow::Owned(bytes), dev, dtype)?;
let q_vk = QTensor::new(qs_vk, (nout, k))?;
let qm_vk = QMatMul::from_qtensor(q_vk)?;
let x_vk = Tensor::from_vec(x_host, (1, k), dev)?;
let y_vk = qm_vk.forward(&x_vk)?.flatten_all()?.to_vec1::<f32>()?;
assert_eq!(y_vk.len(), nout);
let mut max_abs = 0f32;
for n in 0..nout {
max_abs = max_abs.max((y_vk[n] as f64 - y_ref[n]).abs() as f32);
}
Ok(max_abs)
}
#[test]
fn vulkan_qmatmul_forward_matches_cpu() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else { return Ok(()) };
let _env_guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
std::env::set_var("VK_DP4A_DECODE_OFF", "1");
for &(nout, k) in &[(2048usize, 2048usize), (4096, 2048), (512, 256)] {
for dt in [
GgmlDType::Q4_0,
GgmlDType::Q8_0,
GgmlDType::Q4K,
GgmlDType::Q5K,
GgmlDType::Q6K,
GgmlDType::Q2K,
GgmlDType::Q3K,
GgmlDType::IQ4_XS,
GgmlDType::IQ4_NL,
GgmlDType::TQ2_0,
] {
let max_abs = end_to_end_case(&dev, dt, nout, k)?;
println!("QMatMul::forward {dt:?} nout={nout:5} k={k:5} GPU-vs-(dequant ref) max_abs={max_abs:.3e}");
assert!(
max_abs < 1e-3,
"QMatMul::forward {dt:?} GPU/ref mismatch too large: {max_abs}"
);
}
}
std::env::remove_var("VK_DP4A_DECODE_OFF");
Ok(())
}
fn prefill_case(
dev: &Device,
dtype: GgmlDType,
m: usize,
nout: usize,
k: usize,
) -> hanzo_ml::Result<f32> {
let cpu = Device::Cpu;
let w_host: Vec<f32> = (0..nout * k).map(|i| pseudo(i) * 0.5).collect();
let x_host: Vec<f32> = (0..m * k).map(|i| pseudo(i + 7)).collect();
let w_t = Tensor::from_vec(w_host, (nout, k), &cpu)?;
let q_cpu = Arc::new(QTensor::quantize(&w_t, dtype)?);
let bytes = q_cpu.data()?.into_owned();
let w_deq: Vec<f32> = q_cpu.dequantize(&cpu)?.flatten_all()?.to_vec1::<f32>()?;
let mut y_ref = vec![0f64; m * nout];
for mi in 0..m {
for n in 0..nout {
let mut acc = 0f64;
for j in 0..k {
acc += w_deq[n * k + j] as f64 * x_host[mi * k + j] as f64;
}
y_ref[mi * nout + n] = acc;
}
}
let qs_vk = QStorage::from_data(std::borrow::Cow::Owned(bytes), dev, dtype)?;
let q_vk = QTensor::new(qs_vk, (nout, k))?;
let qm_vk = QMatMul::from_qtensor(q_vk)?;
let x_vk = Tensor::from_vec(x_host, (m, k), dev)?;
let y_vk = qm_vk.forward(&x_vk)?.flatten_all()?.to_vec1::<f32>()?;
assert_eq!(y_vk.len(), m * nout);
let mut max_abs = 0f32;
for i in 0..m * nout {
max_abs = max_abs.max((y_vk[i] as f64 - y_ref[i]).abs() as f32);
}
Ok(max_abs)
}
#[test]
fn vulkan_qmatmul_prefill_matches_cpu() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else { return Ok(()) };
let _env_guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
std::env::set_var("VK_Q4K_LEGACY", "1");
for &m in &[5usize, 8, 17, 100, 200] {
for &(nout, k) in &[(2048usize, 2048usize), (4096, 2048), (512, 256)] {
for dt in [
GgmlDType::Q4_0,
GgmlDType::Q8_0,
GgmlDType::Q4K,
GgmlDType::Q5K,
GgmlDType::Q6K,
] {
let max_abs = prefill_case(&dev, dt, m, nout, k)?;
println!(
"QMatMul::prefill {dt:?} M={m:3} nout={nout:5} k={k:5} GPU-vs-(dequant ref) max_abs={max_abs:.3e}"
);
assert!(
max_abs < 1e-3,
"QMatMul::prefill {dt:?} M={m} GPU/ref mismatch too large: {max_abs}"
);
}
}
}
std::env::remove_var("VK_Q4K_LEGACY");
Ok(())
}
fn moe_case(
dev: &Device,
dtype: GgmlDType,
e_cnt: usize,
n: usize,
k: usize,
t: usize,
topk: usize,
) -> hanzo_ml::Result<f32> {
let cpu = Device::Cpu;
let bank_host: Vec<f32> = (0..e_cnt * n * k).map(|i| pseudo(i) * 0.5).collect();
let bank_t = Tensor::from_vec(bank_host, (e_cnt, n, k), &cpu)?;
let q_bank = QTensor::quantize(&bank_t.reshape((e_cnt * n, k))?, dtype)?; let bank_deq: Vec<f32> = q_bank.dequantize(&cpu)?.flatten_all()?.to_vec1::<f32>()?; let bytes = q_bank.data()?.into_owned();
let x_host: Vec<f32> = (0..t * topk * k).map(|i| pseudo(i + 11) * 0.7).collect();
let ids_host: Vec<u32> = (0..t * topk)
.map(|i| ((i * 7 + 3) % e_cnt) as u32)
.collect();
let qs_vk = QStorage::from_data(std::borrow::Cow::Owned(bytes), dev, dtype)?;
let q_vk = QTensor::new(qs_vk, (e_cnt, n, k))?;
let x_vk = Tensor::from_vec(x_host.clone(), (t, topk, k), dev)?;
let ids_vk = Tensor::from_vec(ids_host.clone(), (t, topk), dev)?;
let y_vk = q_vk
.indexed_moe_forward(&x_vk, &ids_vk)?
.reshape((t * topk, n))?
.to_vec2::<f32>()?;
let mut max_abs = 0f32;
for slot in 0..t * topk {
let e = ids_host[slot] as usize;
for r in 0..n {
let wbase = (e * n + r) * k;
let xbase = slot * k;
let mut acc = 0f64;
for j in 0..k {
acc += bank_deq[wbase + j] as f64 * x_host[xbase + j] as f64;
}
max_abs = max_abs.max((y_vk[slot][r] as f64 - acc).abs() as f32);
}
}
Ok(max_abs)
}
fn flash_case(
dev: &Device,
bh: usize,
lq: usize,
lk: usize,
d: usize,
causal: bool,
) -> hanzo_ml::Result<f32> {
let scale = 1.0f32 / (d as f32).sqrt();
let q: Vec<f32> = (0..bh * lq * d).map(|i| pseudo(i) * 0.5).collect();
let k: Vec<f32> = (0..bh * lk * d).map(|i| pseudo(i + 5) * 0.5).collect();
let v: Vec<f32> = (0..bh * lk * d).map(|i| pseudo(i + 9) * 0.5).collect();
let vk = dev.as_vulkan_device()?;
let out = vk.flash_attn(&q, &k, &v, bh, lq, lk, d, scale, causal)?;
assert_eq!(out.len(), bh * lq * d);
let mut max_abs = 0f32;
for b in 0..bh {
for qi in 0..lq {
let last = if causal { qi + (lk - lq) + 1 } else { lk };
let last = last.min(lk);
let mut sc = vec![0f64; last];
let mut mx = f64::NEG_INFINITY;
for (j, scj) in sc.iter_mut().enumerate() {
let mut s = 0f64;
for t in 0..d {
s += q[(b * lq + qi) * d + t] as f64 * k[(b * lk + j) * d + t] as f64;
}
*scj = s * scale as f64;
mx = mx.max(*scj);
}
let mut denom = 0f64;
for scj in sc.iter_mut() {
*scj = (*scj - mx).exp();
denom += *scj;
}
for t in 0..d {
let mut acc = 0f64;
for (j, &p) in sc.iter().enumerate() {
acc += p * v[(b * lk + j) * d + t] as f64;
}
acc /= denom;
let g = out[(b * lq + qi) * d + t] as f64;
max_abs = max_abs.max((g - acc).abs() as f32);
}
}
}
Ok(max_abs)
}
#[test]
fn vulkan_flash_attn_matches_cpu() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else { return Ok(()) };
let cases = [
(8usize, 16usize, 16usize, 128usize, false), (8, 16, 16, 128, true), (8, 1, 64, 128, true), (4, 7, 13, 64, false), ];
for &(bh, lq, lk, d, causal) in &cases {
let max_abs = flash_case(&dev, bh, lq, lk, d, causal)?;
println!(
"FlashAttn bh={bh} lq={lq:3} lk={lk:3} d={d:3} causal={causal} GPU-vs-(eager ref) max_abs={max_abs:.3e}"
);
assert!(
max_abs < 1e-4,
"FlashAttn GPU/ref mismatch too large: {max_abs}"
);
}
Ok(())
}
#[test]
fn vulkan_moe_forward_matches_cpu() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else { return Ok(()) };
let cases = [
(8usize, 256usize, 256usize, 2usize, 2usize),
(16, 512, 256, 3, 4),
(4, 768, 512, 1, 2), ];
for &(e_cnt, n, k, t, topk) in &cases {
for dt in [GgmlDType::Q4_0, GgmlDType::Q8_0, GgmlDType::Q4K] {
let max_abs = moe_case(&dev, dt, e_cnt, n, k, t, topk)?;
println!(
"MoE {dt:?} E={e_cnt:3} n={n:4} k={k:4} t={t} topk={topk} GPU-vs-(dequant ref) max_abs={max_abs:.3e}"
);
assert!(
max_abs < 1e-3,
"MoE {dt:?} GPU/ref mismatch too large: {max_abs}"
);
}
}
Ok(())
}
#[test]
fn vulkan_q4k_tiled2d_matches_default() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else {
return Ok(());
};
let _env_guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let vk = dev.as_vulkan_device()?;
for &(nout, k) in &[(512usize, 256usize), (2048, 2048), (256, 4096), (320, 1024)] {
for &m in &[1usize, 7, 64, 65, 512] {
let (raw, _, _) = weight_bytes(GgmlDType::Q4K, nout, k)?;
let wq = vk.upload_qweight(&raw)?;
let x: Vec<f32> = (0..m * k).map(|i| pseudo(i + 11)).collect();
std::env::set_var("VK_Q4K_LEGACY", "1");
let y_ref = vk.matmul_q4k(&wq, &x, m, nout, k)?;
std::env::remove_var("VK_Q4K_LEGACY");
std::env::set_var("VK_Q4K_TILED2D", "1");
let y_2d = vk.matmul_q4k(&wq, &x, m, nout, k)?;
std::env::remove_var("VK_Q4K_TILED2D");
assert_eq!(y_2d.len(), m * nout);
let mut max_abs = 0f32;
let mut max_ref = 0f32;
for (a, b) in y_ref.iter().zip(y_2d.iter()) {
max_abs = max_abs.max((a - b).abs());
max_ref = max_ref.max(a.abs());
}
let rel = if max_ref > 0.0 {
max_abs / max_ref
} else {
max_abs
};
assert!(
rel < 2e-3,
"2d-tiled != default (m={m}, nout={nout}, k={k}): max_abs={max_abs}, rel={rel}"
);
}
}
Ok(())
}
#[test]
#[ignore]
fn vulkan_q4k_tiled2d_bench() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else {
return Ok(());
};
let vk = dev.as_vulkan_device()?;
let (m, nout, k) = (512usize, 4096usize, 4096usize);
let (raw, _, _) = weight_bytes(GgmlDType::Q4K, nout, k)?;
let wq = vk.upload_qweight(&raw)?;
let x: Vec<f32> = (0..m * k).map(|i| pseudo(i + 7)).collect();
let iters = 20;
let bench = |var: Option<&str>| -> hanzo_ml::Result<f64> {
std::env::remove_var("VK_Q4K_TILED2D");
if let Some(v) = var {
std::env::set_var(v, "1");
}
let _ = vk.matmul_q4k(&wq, &x, m, nout, k)?;
let t = std::time::Instant::now();
for _ in 0..iters {
let _ = vk.matmul_q4k(&wq, &x, m, nout, k)?;
}
if let Some(v) = var {
std::env::remove_var(v);
}
Ok(t.elapsed().as_secs_f64() * 1e3 / iters as f64)
};
let def_ms = bench(Some("VK_Q4K_LEGACY"))?;
let tiled2d_ms = bench(Some("VK_Q4K_TILED2D"))?;
eprintln!(
"[L1 2D bench] m={m} nout={nout} k={k}: default={def_ms:.3} ms tiled2d={tiled2d_ms:.3} ms speedup={:.2}x",
def_ms / tiled2d_ms
);
Ok(())
}
#[test]
fn vulkan_q4k_dp4a2d_matches_default() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else {
return Ok(());
};
let _env_guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let vk = dev.as_vulkan_device()?;
for &(nout, k) in &[(512usize, 256usize), (2048, 2048), (256, 4096), (320, 1024)] {
for &m in &[1usize, 7, 64, 65, 512] {
let (raw, _, _) = weight_bytes(GgmlDType::Q4K, nout, k)?;
let wq = vk.upload_qweight(&raw)?;
let x: Vec<f32> = (0..m * k).map(|i| pseudo(i + 13)).collect();
std::env::set_var("VK_Q4K_LEGACY", "1");
let y_ref = vk.matmul_q4k(&wq, &x, m, nout, k)?;
std::env::remove_var("VK_Q4K_LEGACY");
std::env::set_var("VK_Q4K_DP4A", "1");
let y_dp = vk.matmul_q4k(&wq, &x, m, nout, k)?;
std::env::remove_var("VK_Q4K_DP4A");
assert_eq!(y_dp.len(), m * nout);
let mut max_abs = 0f32;
let mut max_ref = 0f32;
for (a, b) in y_ref.iter().zip(y_dp.iter()) {
max_abs = max_abs.max((a - b).abs());
max_ref = max_ref.max(a.abs());
}
let rel = if max_ref > 0.0 {
max_abs / max_ref
} else {
max_abs
};
assert!(
rel < 1.5e-2,
"dp4a-2d != default (m={m}, nout={nout}, k={k}): max_abs={max_abs}, max_ref={max_ref}, rel={rel}"
);
}
}
Ok(())
}
#[test]
#[ignore]
fn vulkan_q4k_dp4a2d_bench() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else {
return Ok(());
};
let vk = dev.as_vulkan_device()?;
let (m, nout, k) = (512usize, 4096usize, 4096usize);
let (raw, _, _) = weight_bytes(GgmlDType::Q4K, nout, k)?;
let wq = vk.upload_qweight(&raw)?;
let x: Vec<f32> = (0..m * k).map(|i| pseudo(i + 7)).collect();
let iters = 20;
let bench = |var: Option<&str>| -> hanzo_ml::Result<f64> {
for v in ["VK_Q4K_TILED2D", "VK_Q4K_DP4A"] {
std::env::remove_var(v);
}
if let Some(v) = var {
std::env::set_var(v, "1");
}
let _ = vk.matmul_q4k(&wq, &x, m, nout, k)?;
let t = std::time::Instant::now();
for _ in 0..iters {
let _ = vk.matmul_q4k(&wq, &x, m, nout, k)?;
}
if let Some(v) = var {
std::env::remove_var(v);
}
Ok(t.elapsed().as_secs_f64() * 1e3 / iters as f64)
};
let def_ms = bench(Some("VK_Q4K_LEGACY"))?;
let f32_ms = bench(Some("VK_Q4K_TILED2D"))?;
let dp4a_ms = bench(Some("VK_Q4K_DP4A"))?;
eprintln!(
"[L1 dp4a bench] m={m} nout={nout} k={k}: default={def_ms:.3} 2d_f32={f32_ms:.3} 2d_dp4a={dp4a_ms:.3} ms | dp4a vs default={:.2}x vs 2d_f32={:.2}x",
def_ms / dp4a_ms,
f32_ms / dp4a_ms
);
Ok(())
}
#[test]
fn vulkan_q4k_coopmat_matches_default() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else {
return Ok(());
};
let _env_guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let vk = dev.as_vulkan_device()?;
if vk.coopmat_info().is_none() {
eprintln!("[vulkan_quant_tests] no coopmat; skipping coopmat gate");
return Ok(());
}
for &(nout, k) in &[(512usize, 256usize), (2048, 2048), (256, 4096), (320, 1024)] {
for &m in &[2usize, 16, 17, 64, 512] {
let (raw, _, _) = weight_bytes(GgmlDType::Q4K, nout, k)?;
let wq = vk.upload_qweight(&raw)?;
let x: Vec<f32> = (0..m * k).map(|i| pseudo(i + 17)).collect();
std::env::set_var("VK_Q4K_LEGACY", "1");
let y_ref = vk.matmul_q4k(&wq, &x, m, nout, k)?;
std::env::remove_var("VK_Q4K_LEGACY");
std::env::set_var("VK_Q4K_COOPMAT", "1");
let y_cm = vk.matmul_q4k(&wq, &x, m, nout, k)?;
std::env::remove_var("VK_Q4K_COOPMAT");
let mut max_abs = 0f32;
let mut max_ref = 0f32;
for (a, b) in y_ref.iter().zip(y_cm.iter()) {
max_abs = max_abs.max((a - b).abs());
max_ref = max_ref.max(a.abs());
}
let rel = if max_ref > 0.0 {
max_abs / max_ref
} else {
max_abs
};
assert!(
rel < 1e-2,
"coopmat != legacy (m={m}, nout={nout}, k={k}): max_abs={max_abs}, rel={rel}"
);
}
}
Ok(())
}
#[test]
#[ignore]
fn vulkan_q4k_coopmat_bench() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else {
return Ok(());
};
let vk = dev.as_vulkan_device()?;
if vk.coopmat_info().is_none() {
eprintln!("[vulkan_quant_tests] no coopmat; skipping");
return Ok(());
}
let (m, nout, k) = (512usize, 4096usize, 4096usize);
let (raw, _, _) = weight_bytes(GgmlDType::Q4K, nout, k)?;
let wq = vk.upload_qweight(&raw)?;
let x: Vec<f32> = (0..m * k).map(|i| pseudo(i + 7)).collect();
let iters = 20;
let bench = |var: Option<&str>| -> hanzo_ml::Result<f64> {
for v in ["VK_Q4K_COOPMAT", "VK_Q4K_LEGACY"] {
std::env::remove_var(v);
}
if let Some(v) = var {
std::env::set_var(v, "1");
}
let _ = vk.matmul_q4k(&wq, &x, m, nout, k)?;
let t = std::time::Instant::now();
for _ in 0..iters {
let _ = vk.matmul_q4k(&wq, &x, m, nout, k)?;
}
if let Some(v) = var {
std::env::remove_var(v);
}
Ok(t.elapsed().as_secs_f64() * 1e3 / iters as f64)
};
let dp4a_ms = bench(None)?; let legacy_ms = bench(Some("VK_Q4K_LEGACY"))?;
let cm_ms = bench(Some("VK_Q4K_COOPMAT"))?;
eprintln!(
"[L4 coopmat bench] m={m} nout={nout} k={k}: legacy={legacy_ms:.3} dp4a={dp4a_ms:.3} coopmat={cm_ms:.3} ms | coopmat vs dp4a={:.2}x vs legacy={:.2}x",
dp4a_ms / cm_ms,
legacy_ms / cm_ms
);
Ok(())
}
#[test]
#[ignore]
fn vulkan_q4k_kernel_bench() -> hanzo_ml::Result<()> {
let Some(dev) = gpu() else {
return Ok(());
};
let vk = dev.as_vulkan_device()?;
let (m, nout, k) = (512usize, 4096usize, 4096usize);
let (raw, _, _) = weight_bytes(GgmlDType::Q4K, nout, k)?;
let wq = vk.upload_qweight(&raw)?;
let x: Vec<f32> = (0..m * k).map(|i| pseudo(i + 7)).collect();
let bench = |var: Option<&str>| -> hanzo_ml::Result<f64> {
for v in [
"VK_Q4K_LEGACY",
"VK_Q4K_TILED2D",
"VK_Q4K_DP4A",
"VK_Q4K_COOPMAT",
] {
std::env::remove_var(v);
}
if let Some(v) = var {
std::env::set_var(v, "1");
}
let r = vk.bench_matmul_q4k(&wq, &x, m, nout, k, 50);
if let Some(v) = var {
std::env::remove_var(v);
}
r
};
let leg = bench(Some("VK_Q4K_LEGACY"))?;
let f32t = bench(Some("VK_Q4K_TILED2D"))?;
let dp = bench(None)?; let cm = if vk.coopmat_info().is_some() {
bench(Some("VK_Q4K_COOPMAT"))?
} else {
0.0
};
let gflop = 2.0 * m as f64 * nout as f64 * k as f64 / 1e9;
eprintln!(
"[KERNEL bench] m={m} nout={nout} k={k} ({gflop:.1} GFLOP): legacy={leg:.3} f32_2d={f32t:.3} dp4a={dp:.3} coopmat={cm:.3} ms | dp4a={:.0} GFLOP/s, {:.1}x vs legacy",
gflop / (dp / 1e3),
leg / dp
);
Ok(())
}
#[test]
fn vulkan_rope_norm_matches_unfused() {
let Some(dev) = gpu() else {
return;
};
let vk = dev.as_vulkan_device().unwrap();
for &(b, h, t, d) in &[
(1usize, 4usize, 1usize, 128usize),
(1, 8, 5, 64),
(2, 4, 3, 128),
] {
let hd = d / 2;
let n = b * h * t * d;
let x: Vec<f32> = (0..n).map(|i| pseudo(i + 7)).collect();
let weight: Vec<f32> = (0..d).map(|i| 0.5 + pseudo(i + 99).abs()).collect();
let cs: Vec<f32> = (0..t * hd).map(|i| pseudo(i + 3).cos()).collect();
let sn: Vec<f32> = (0..t * hd).map(|i| pseudo(i + 3).sin()).collect();
let eps = 1e-6f32;
let mut refv = vec![0f32; n];
for row in 0..b * h * t {
let base = row * d;
let mut ss = 0f32;
for i in 0..d {
let v = x[base + i];
ss += v * v;
}
let denom = (ss / d as f32 + eps).sqrt();
let i_t = row % t;
for i_d in 0..hd {
let (i1, i2) = (base + i_d, base + i_d + hd);
let x1 = x[i1] / denom * weight[i_d];
let x2 = x[i2] / denom * weight[i_d + hd];
let (c, s) = (cs[i_t * hd + i_d], sn[i_t * hd + i_d]);
refv[i1] = x1 * c - x2 * s;
refv[i2] = x1 * s + x2 * c;
}
}
let got = vk
.rope_norm_f32(&x, &weight, &cs, &sn, b, h, t, d, eps)
.unwrap();
assert_eq!(got.len(), n);
let mut maxabs = 0f32;
for i in 0..n {
maxabs = maxabs.max((got[i] - refv[i]).abs());
}
eprintln!("rope_norm b{b} h{h} t{t} d{d}: maxabs={maxabs:.3e}");
assert!(
maxabs < 1e-4,
"rope_norm (b{b} h{h} t{t} d{d}) mismatch maxabs={maxabs}"
);
}
}
#[test]
fn vulkan_add_rmsnorm_matches_unfused() {
let Some(dev) = gpu() else {
return;
};
let vk = dev.as_vulkan_device().unwrap();
for &(nrows, m) in &[(1usize, 2048usize), (5, 64), (3, 4096)] {
let n = nrows * m;
let x: Vec<f32> = (0..n).map(|i| pseudo(i + 11)).collect();
let res: Vec<f32> = (0..n).map(|i| pseudo(i + 222)).collect();
let alpha: Vec<f32> = (0..m).map(|i| 0.5 + pseudo(i + 9).abs()).collect();
let eps = 1e-5f32;
let mut s_ref = vec![0f32; n];
let mut y_ref = vec![0f32; n];
for row in 0..nrows {
let base = row * m;
let mut ss = 0f32;
for i in 0..m {
let v = x[base + i] + res[base + i];
s_ref[base + i] = v;
ss += v * v;
}
let denom = (ss / m as f32 + eps).sqrt();
for i in 0..m {
y_ref[base + i] = s_ref[base + i] / denom * alpha[i];
}
}
let (s_gpu, y_gpu) = vk.add_rmsnorm_f32(&x, &res, &alpha, nrows, m, eps).unwrap();
let (mut ms, mut my) = (0f32, 0f32);
for i in 0..n {
ms = ms.max((s_gpu[i] - s_ref[i]).abs());
my = my.max((y_gpu[i] - y_ref[i]).abs());
}
eprintln!("add_rmsnorm r{nrows} m{m}: s_maxabs={ms:.3e} y_maxabs={my:.3e}");
assert!(
ms < 1e-5 && my < 1e-4,
"add_rmsnorm r{nrows} m{m} mismatch s={ms} y={my}"
);
}
}
#[test]
fn vulkan_q4k_dp4a_decode_matches_scalar() {
let Some(dev) = gpu() else {
return;
};
let vk = dev.as_vulkan_device().unwrap();
for &(nout, k) in SHAPES {
let x: Vec<f32> = (0..k).map(|i| pseudo(i + 5)).collect();
let (raw, _w_deq, _w_host) = weight_bytes(GgmlDType::Q4K, nout, k).unwrap();
let wq = vk.upload_qweight(&raw).unwrap();
let scalar = vk.matvec_q4k_scalar(&wq, &x, nout, k).unwrap();
let column = vk.matvec_q4k_dp4a(&wq, &x, nout, k).unwrap();
for (name, got) in [("column", &column)] {
let (mut sse, mut refsq) = (0f64, 0f64);
for i in 0..nout {
let d = (got[i] - scalar[i]) as f64;
sse += d * d;
refsq += (scalar[i] as f64) * (scalar[i] as f64);
}
let rel = (sse / refsq.max(1e-9)).sqrt();
eprintln!("dp4a-{name} vs scalar nout={nout} k={k}: rel={rel:.3e}");
assert!(
rel < 2e-2,
"dp4a {name} rel err {rel} too large (nout={nout} k={k})"
);
}
}
}