use proptest::prelude::*;
use crate::quantize::dequant::dequantize_q4_k;
use crate::quantize::fused_k::{
fused_q4k_dot, fused_q4k_dot_simd, fused_q4k_q8k_dot, fused_q4k_q8k_dot_simd,
};
use crate::quantize::types::QK_K;
fn gen_q4k_superblock() -> impl Strategy<Value = Vec<u8>> {
prop::collection::vec(any::<u8>(), 144..=144)
}
fn gen_activations(num_values: usize) -> impl Strategy<Value = Vec<f32>> {
prop::collection::vec(-10.0f32..10.0f32, num_values..=num_values)
}
fn naive_dot(weights: &[f32], activations: &[f32]) -> f32 {
weights
.iter()
.zip(activations.iter())
.map(|(w, a)| w * a)
.sum()
}
proptest! {
#![proptest_config(ProptestConfig::with_cases(50))]
#[test]
fn prop_fused_q4k_dot_matches_naive(
q4k_data in gen_q4k_superblock(),
activations in gen_activations(QK_K)
) {
let dequantized = match dequantize_q4_k(&q4k_data) {
Ok(d) => d,
Err(_) => return Ok(()),
};
if dequantized.iter().any(|v| !v.is_finite()) {
return Ok(());
}
let naive_result = naive_dot(&dequantized, &activations);
if !naive_result.is_finite() {
return Ok(());
}
let fused_result = fused_q4k_dot(&q4k_data, &activations)?;
if !fused_result.is_finite() {
return Ok(());
}
let tolerance = (naive_result.abs() * 1e-3).max(1e-4);
prop_assert!(
(fused_result - naive_result).abs() <= tolerance,
"Q4_K dot mismatch: fused={}, naive={}, diff={}, tolerance={}",
fused_result, naive_result, (fused_result - naive_result).abs(), tolerance
);
}
#[test]
fn prop_fused_q4k_dot_multi_block(
blocks in prop::collection::vec(gen_q4k_superblock(), 1..=4)
) {
let q4k_data: Vec<u8> = blocks.iter().flatten().copied().collect();
let num_values = blocks.len() * QK_K;
let activations: Vec<f32> = (0..num_values).map(|i| (i as f32 * 0.01) - 0.5).collect();
let dequantized = match dequantize_q4_k(&q4k_data) {
Ok(d) => d,
Err(_) => return Ok(()),
};
if dequantized.iter().any(|v| !v.is_finite()) {
return Ok(());
}
let naive_result = naive_dot(&dequantized, &activations);
if !naive_result.is_finite() {
return Ok(());
}
let fused_result = fused_q4k_dot(&q4k_data, &activations)?;
if !fused_result.is_finite() {
return Ok(());
}
let tolerance = (naive_result.abs() * 1e-3).max(1e-3);
prop_assert!(
(fused_result - naive_result).abs() <= tolerance,
"Multi-block Q4_K dot mismatch: fused={}, naive={}, diff={}",
fused_result, naive_result, (fused_result - naive_result).abs()
);
}
}
#[test]
fn test_fused_q4k_dot_zero_activations() {
let q4k_data = vec![0u8; 144];
let activations = vec![0.0f32; QK_K];
let result = fused_q4k_dot(&q4k_data, &activations).expect("Should succeed");
assert!(
result.abs() < 1e-10,
"Zero activations should give zero result: {}",
result
);
}
#[test]
fn test_fused_q4k_dot_all_ones_activations() {
let q4k_data = vec![0u8; 144];
let activations = vec![1.0f32; QK_K];
let result = fused_q4k_dot(&q4k_data, &activations).expect("Should succeed");
let dequantized = dequantize_q4_k(&q4k_data).expect("Dequant should succeed");
let expected: f32 = dequantized.iter().sum();
let tolerance = expected.abs() * 1e-4 + 1e-6;
assert!(
(result - expected).abs() <= tolerance,
"All-ones result mismatch: got {}, expected {}",
result,
expected
);
}
#[test]
fn test_fused_q4k_dot_invalid_length() {
let q4k_data = vec![0u8; 100]; let activations = vec![1.0f32; QK_K];
let result = fused_q4k_dot(&q4k_data, &activations);
assert!(result.is_err(), "Invalid length should error");
}
#[test]
fn test_fused_q4k_dot_activation_mismatch() {
let q4k_data = vec![0u8; 144];
let activations = vec![1.0f32; 128];
let result = fused_q4k_dot(&q4k_data, &activations);
assert!(result.is_err(), "Activation mismatch should error");
}
#[test]
fn test_fused_q4k_dot_deterministic() {
let q4k_data: Vec<u8> = (0..144).map(|i| (i * 17) as u8).collect();
let activations: Vec<f32> = (0..QK_K).map(|i| (i as f32) * 0.01).collect();
let result1 = fused_q4k_dot(&q4k_data, &activations).expect("Should succeed");
let result2 = fused_q4k_dot(&q4k_data, &activations).expect("Should succeed");
assert_eq!(
result1, result2,
"Fused dot should be deterministic: {} vs {}",
result1, result2
);
}
#[test]
fn test_fused_q4k_dot_scale_sensitivity() {
let mut q4k_data1 = vec![0u8; 144];
let mut q4k_data2 = vec![0u8; 144];
q4k_data1[0..2].copy_from_slice(&0x3C00u16.to_le_bytes());
q4k_data2[0..2].copy_from_slice(&0x4000u16.to_le_bytes());
for i in 16..144 {
q4k_data1[i] = 0x55; q4k_data2[i] = 0x55;
}
for i in 4..16 {
q4k_data1[i] = 0x3F; q4k_data2[i] = 0x3F;
}
let activations = vec![1.0f32; QK_K];
let result1 = fused_q4k_dot(&q4k_data1, &activations).expect("Should succeed");
let result2 = fused_q4k_dot(&q4k_data2, &activations).expect("Should succeed");
if result1.abs() > 1e-3 || result2.abs() > 1e-3 {
assert!(
(result1 - result2).abs() > 1e-6 || (result1.abs() < 1e-6 && result2.abs() < 1e-6),
"Different scales with non-zero quants should give different results: {} vs {}",
result1,
result2
);
}
}
#[test]
fn test_fused_q4k_dot_simd_zero_activations() {
let q4k_data = vec![0u8; 144];
let activations = vec![0.0f32; QK_K];
let result = fused_q4k_dot_simd(&q4k_data, &activations).expect("Should succeed");
assert!(
result.abs() < 1e-10,
"SIMD: Zero activations should give zero result: {}",
result
);
}
#[test]
fn test_fused_q4k_dot_simd_matches_scalar() {
let mut q4k_data = vec![0u8; 144];
q4k_data[0..2].copy_from_slice(&0x3C00u16.to_le_bytes());
for i in 16..144 {
q4k_data[i] = 0xAA; }
let activations: Vec<f32> = (0..QK_K).map(|i| (i as f32) * 0.01).collect();
let scalar_result = fused_q4k_dot(&q4k_data, &activations).expect("Scalar should succeed");
let simd_result = fused_q4k_dot_simd(&q4k_data, &activations).expect("SIMD should succeed");
let tolerance = scalar_result.abs() * 1e-3 + 1e-4;
assert!(
(scalar_result - simd_result).abs() <= tolerance,
"SIMD should match scalar: scalar={}, simd={}, diff={}",
scalar_result,
simd_result,
(scalar_result - simd_result).abs()
);
}
#[test]
fn test_fused_q4k_dot_simd_invalid_length() {
let q4k_data = vec![0u8; 100]; let activations = vec![1.0f32; QK_K];
let result = fused_q4k_dot_simd(&q4k_data, &activations);
assert!(result.is_err(), "SIMD: Invalid length should error");
}
#[test]
fn test_fused_q4k_dot_simd_activation_mismatch() {
let q4k_data = vec![0u8; 144];
let activations = vec![1.0f32; 128];
let result = fused_q4k_dot_simd(&q4k_data, &activations);
assert!(result.is_err(), "SIMD: Activation mismatch should error");
}
#[test]
fn test_fused_q4k_dot_simd_deterministic() {
let q4k_data: Vec<u8> = (0..144).map(|i| (i * 17) as u8).collect();
let activations: Vec<f32> = (0..QK_K).map(|i| (i as f32) * 0.01).collect();
let result1 = fused_q4k_dot_simd(&q4k_data, &activations).expect("Should succeed");
let result2 = fused_q4k_dot_simd(&q4k_data, &activations).expect("Should succeed");
assert_eq!(
result1, result2,
"SIMD should be deterministic: {} vs {}",
result1, result2
);
}
fn gen_q8k_scales() -> Vec<f32> {
vec![1.0f32; 8]
}
fn gen_q8k_quants() -> Vec<i8> {
vec![0i8; QK_K]
}
#[test]
fn test_fused_q4k_q8k_dot_zero_inputs() {
let q4k_data = vec![0u8; 144];
let q8k_scales = gen_q8k_scales();
let q8k_quants = gen_q8k_quants();
let result = fused_q4k_q8k_dot(&q4k_data, &q8k_scales, &q8k_quants).expect("Should succeed");
assert!(
result.abs() < 1e-10,
"Q4K×Q8K: Zero inputs should give zero result: {}",
result
);
}
#[test]
fn test_fused_q4k_q8k_dot_with_values() {
let mut q4k_data = vec![0u8; 144];
q4k_data[0..2].copy_from_slice(&0x3C00u16.to_le_bytes());
for i in 16..144 {
q4k_data[i] = 0x55;
}
let q8k_scales = vec![1.0f32; 8];
let q8k_quants: Vec<i8> = (0..QK_K).map(|i| ((i % 15) as i8) - 7).collect();
let result = fused_q4k_q8k_dot(&q4k_data, &q8k_scales, &q8k_quants);
assert!(result.is_ok(), "Q4K×Q8K dot should succeed");
}
#[test]
fn test_fused_q4k_q8k_dot_invalid_q4k_length() {
let q4k_data = vec![0u8; 100]; let q8k_scales = gen_q8k_scales();
let q8k_quants = gen_q8k_quants();
let result = fused_q4k_q8k_dot(&q4k_data, &q8k_scales, &q8k_quants);
assert!(result.is_err(), "Q4K×Q8K: Invalid Q4K length should error");
}
#[test]
fn test_fused_q4k_q8k_dot_invalid_q8k_quants_length() {
let q4k_data = vec![0u8; 144];
let q8k_scales = gen_q8k_scales();
let q8k_quants = vec![0i8; 128];
let result = fused_q4k_q8k_dot(&q4k_data, &q8k_scales, &q8k_quants);
assert!(
result.is_err(),
"Q4K×Q8K: Invalid Q8K quants length should error"
);
}
#[test]
fn test_fused_q4k_q8k_dot_deterministic() {
let q4k_data: Vec<u8> = (0..144).map(|i| (i * 23) as u8).collect();
let q8k_scales: Vec<f32> = (0..8).map(|i| (i as f32) * 0.1 + 0.5).collect();
let q8k_quants: Vec<i8> = (0..QK_K).map(|i| ((i % 127) as i8) - 64).collect();
let result1 = fused_q4k_q8k_dot(&q4k_data, &q8k_scales, &q8k_quants).expect("Should succeed");
let result2 = fused_q4k_q8k_dot(&q4k_data, &q8k_scales, &q8k_quants).expect("Should succeed");
assert_eq!(
result1, result2,
"Q4K×Q8K should be deterministic: {} vs {}",
result1, result2
);
}
include!("fused_q4k_04.rs");
include!("fused_q4k_02_02.rs");