#[test]
fn test_fused_q4k_dot_multiple_super_blocks() {
let num_super_blocks = 4;
let mut q4k_data = Vec::with_capacity(num_super_blocks * 144);
for sb_idx in 0..num_super_blocks {
let d = 0.5 + (sb_idx as f32) * 0.1;
q4k_data.extend_from_slice(&half::f16::from_f32(d).to_bits().to_le_bytes());
q4k_data.extend_from_slice(&half::f16::from_f32(0.1).to_bits().to_le_bytes());
for i in 0..12 {
q4k_data.push(((sb_idx * 7 + i) % 64) as u8);
}
for i in 0..128 {
q4k_data.push(((sb_idx * 13 + i) % 256) as u8);
}
}
let activations: Vec<f32> = (0..1024).map(|i| (i as f32 * 0.017).sin() * 2.0).collect();
let dequantized = dequantize_q4_k(&q4k_data).expect("test");
let reference = naive_dot_product(&dequantized, &activations);
let fused = fused_q4k_dot(&q4k_data, &activations).expect("test");
assert_ulp_eq(fused, reference, 4, "fused_q4k_dot multiple super-blocks");
}
#[test]
fn test_fused_q4k_dot_edge_values() {
let mut q4k_zeros = Vec::new();
q4k_zeros.extend_from_slice(&half::f16::from_f32(0.0).to_bits().to_le_bytes());
q4k_zeros.extend_from_slice(&half::f16::from_f32(0.0).to_bits().to_le_bytes());
q4k_zeros.extend_from_slice(&[0u8; 12]); q4k_zeros.extend_from_slice(&[0u8; 128]);
let activations_zeros: Vec<f32> = vec![1.0; 256];
let fused_zeros = fused_q4k_dot(&q4k_zeros, &activations_zeros).expect("test");
assert!(
fused_zeros.abs() < 1e-6,
"Zero weights should produce zero dot product"
);
let mut q4k_max = Vec::new();
q4k_max.extend_from_slice(&half::f16::from_f32(1.0).to_bits().to_le_bytes());
q4k_max.extend_from_slice(&half::f16::from_f32(0.0).to_bits().to_le_bytes());
q4k_max.extend_from_slice(&[0xFF; 12]); q4k_max.extend_from_slice(&[0xFF; 128]);
let activations_ones: Vec<f32> = vec![1.0; 256];
let dequantized_max = dequantize_q4_k(&q4k_max).expect("test");
let reference_max = naive_dot_product(&dequantized_max, &activations_ones);
let fused_max = fused_q4k_dot(&q4k_max, &activations_ones).expect("test");
assert_ulp_eq(fused_max, reference_max, 4, "fused_q4k_dot max values");
let activations_neg: Vec<f32> = (0..256).map(|i| -((i as f32) * 0.01)).collect();
let dequantized_neg = dequantize_q4_k(&q4k_max).expect("test");
let reference_neg = naive_dot_product(&dequantized_neg, &activations_neg);
let fused_neg = fused_q4k_dot(&q4k_max, &activations_neg).expect("test");
assert_ulp_eq(
fused_neg,
reference_neg,
4,
"fused_q4k_dot negative activations",
);
}
#[test]
fn test_fused_q4k_dot_length_mismatch() {
let q4k_data = vec![0u8; 144]; let activations = vec![0.0f32; 128];
let result = fused_q4k_dot(&q4k_data, &activations);
assert!(
result.is_err(),
"Should error on activation length mismatch"
);
}
#[test]
fn test_fused_q4k_dot_invalid_data_length() {
let q4k_data = vec![0u8; 143]; let activations = vec![0.0f32; 256];
let result = fused_q4k_dot(&q4k_data, &activations);
assert!(result.is_err(), "Should error on invalid Q4_K data length");
}
#[test]
fn test_fused_q4k_dot_no_intermediate_allocation() {
let q4k_data = vec![0u8; 144];
let activations = vec![0.0f32; 256];
let result: Result<f32> = fused_q4k_dot(&q4k_data, &activations);
assert!(result.is_ok());
}
#[test]
fn test_fused_q6k_dot_basic() {
let mut q6k_data = Vec::new();
for i in 0..128 {
q6k_data.push((i % 16) as u8 | (((i + 1) % 16) as u8) << 4);
}
for i in 0..64 {
q6k_data.push((i % 4) as u8 | (((i + 1) % 4) as u8) << 2);
}
for i in 0..16 {
q6k_data.push((i as i8 - 8) as u8);
}
q6k_data.extend_from_slice(&half::f16::from_f32(1.0).to_bits().to_le_bytes());
let activations: Vec<f32> = (0..256).map(|i| (i as f32) * 0.01).collect();
let dequantized = dequantize_q6_k(&q6k_data).expect("test");
let reference = naive_dot_product(&dequantized, &activations);
let fused = fused_q6k_dot(&q6k_data, &activations).expect("test");
assert_ulp_eq(fused, reference, 4, "fused_q6k_dot basic");
}
#[test]
fn test_fused_q6k_dot_multiple_super_blocks() {
let num_super_blocks = 4;
let mut q6k_data = Vec::with_capacity(num_super_blocks * 210);
for sb_idx in 0..num_super_blocks {
for i in 0..128 {
q6k_data.push(((sb_idx * 7 + i) % 256) as u8);
}
for i in 0..64 {
q6k_data.push(((sb_idx * 11 + i) % 256) as u8);
}
for i in 0..16 {
#[allow(clippy::cast_possible_wrap)]
let scale = ((sb_idx * 3 + i) % 128) as i8;
q6k_data.push(scale as u8);
}
let d = 0.5 + (sb_idx as f32) * 0.2;
q6k_data.extend_from_slice(&half::f16::from_f32(d).to_bits().to_le_bytes());
}
let activations: Vec<f32> = (0..1024).map(|i| (i as f32 * 0.023).cos() * 1.5).collect();
let dequantized = dequantize_q6_k(&q6k_data).expect("test");
let reference = naive_dot_product(&dequantized, &activations);
let fused = fused_q6k_dot(&q6k_data, &activations).expect("test");
assert_ulp_eq(fused, reference, 4, "fused_q6k_dot multiple super-blocks");
}
#[test]
fn test_fused_q6k_dot_length_mismatch() {
let q6k_data = vec![0u8; 210]; let activations = vec![0.0f32; 128];
let result = fused_q6k_dot(&q6k_data, &activations);
assert!(
result.is_err(),
"Should error on activation length mismatch"
);
}
#[test]
fn test_fused_q4k_dot_simd_matches_scalar() {
let num_super_blocks = 4;
let mut q4k_data = Vec::with_capacity(num_super_blocks * 144);
for sb_idx in 0..num_super_blocks {
let d = 0.5 + (sb_idx as f32) * 0.1;
q4k_data.extend_from_slice(&half::f16::from_f32(d).to_bits().to_le_bytes());
q4k_data.extend_from_slice(&half::f16::from_f32(0.1).to_bits().to_le_bytes());
for i in 0..12 {
q4k_data.push(((sb_idx * 7 + i) % 64) as u8);
}
for i in 0..128 {
q4k_data.push(((sb_idx * 13 + i) % 256) as u8);
}
}
let activations: Vec<f32> = (0..1024).map(|i| (i as f32 * 0.017).sin() * 2.0).collect();
let scalar_result = fused_q4k_dot(&q4k_data, &activations).expect("test");
let simd_result = fused_q4k_dot_simd(&q4k_data, &activations).expect("test");
assert_ulp_eq(
simd_result,
scalar_result,
8,
"SIMD result should match scalar within 8 ULPs",
);
}
#[test]
fn test_fused_q4k_dot_simd_error_handling() {
let bad_data = vec![0u8; 143]; let activations = vec![0.0f32; 256];
assert!(fused_q4k_dot_simd(&bad_data, &activations).is_err());
let good_data = vec![0u8; 144];
let bad_activations = vec![0.0f32; 128];
assert!(fused_q4k_dot_simd(&good_data, &bad_activations).is_err());
}
#[test]
fn test_fused_q4k_dot_simd_large_input() {
let num_super_blocks = 16;
let mut q4k_data = Vec::with_capacity(num_super_blocks * 144);
for sb_idx in 0..num_super_blocks {
let d = 1.0 + (sb_idx as f32) * 0.05;
q4k_data.extend_from_slice(&half::f16::from_f32(d).to_bits().to_le_bytes());
q4k_data.extend_from_slice(&half::f16::from_f32(0.0).to_bits().to_le_bytes());
for i in 0..12 {
q4k_data.push(((sb_idx + i) % 64) as u8);
}
for i in 0..128 {
q4k_data.push(((sb_idx * 17 + i * 3) % 256) as u8);
}
}
let activations: Vec<f32> = (0..4096).map(|i| (i as f32 * 0.001).cos()).collect();
let dequantized = dequantize_q4_k(&q4k_data).expect("test");
let reference = naive_dot_product(&dequantized, &activations);
let simd_result = fused_q4k_dot_simd(&q4k_data, &activations).expect("test");
let ulp_d = ulp_diff(simd_result, reference);
assert!(
ulp_d <= 16,
"Large input SIMD result should match reference: simd={}, ref={}, ulp_diff={}",
simd_result,
reference,
ulp_d
);
}
#[test]
fn test_fused_q4k_tiled_matvec_basic() {
use crate::quantize::fused_q4k_tiled_matvec;
let in_dim = 256;
let out_dim = 4;
let mut weight_data = Vec::with_capacity(out_dim * 144);
for row in 0..out_dim {
let d = 0.5 + (row as f32) * 0.1;
weight_data.extend_from_slice(&half::f16::from_f32(d).to_bits().to_le_bytes());
weight_data.extend_from_slice(&half::f16::from_f32(0.05).to_bits().to_le_bytes());
for i in 0..12 {
weight_data.push(((row * 7 + i) % 64) as u8);
}
for i in 0..128 {
weight_data.push(((row * 13 + i) % 256) as u8);
}
}
let activations: Vec<f32> = (0..in_dim).map(|i| (i as f32 * 0.01).sin()).collect();
let mut reference = Vec::with_capacity(out_dim);
for row in 0..out_dim {
let row_start = row * 144;
let row_data = &weight_data[row_start..row_start + 144];
let dot = fused_q4k_dot_simd(row_data, &activations).expect("test");
reference.push(dot);
}
let tiled =
fused_q4k_tiled_matvec(&weight_data, &activations, in_dim, out_dim, None).expect("test");
assert_eq!(tiled.len(), out_dim);
for i in 0..out_dim {
assert_ulp_eq(
tiled[i],
reference[i],
4,
&format!("tiled_matvec output {}", i),
);
}
}