use crate::quantize::{
dequantize_q8_blocks, extract_scale_min,
fused_q4_0_q8_0_dot_scalar, fused_q4_0_q8_0_parallel_matvec,
fused_q4_0_q8_0_parallel_matvec_into, fused_q8_0_q8_0_dot_scalar,
fused_q8_0_q8_0_parallel_matvec, fused_q8_0_q8_0_parallel_matvec_into,
quantize_activations_q8k_into, quantize_to_q8_blocks, InterleavedQ4K, Q8KSuperBlock, Q8_0Block,
};
#[test]
fn test_interleaved_q4k_mod_dot_single_superblock_zeros() {
let data = vec![0u8; 144];
let interleaved = InterleavedQ4K::from_q4k(&data).expect("valid");
let activations = vec![1.0f32; 256];
let result = interleaved.dot(&activations).expect("dot should work");
assert!(result.is_finite(), "Result should be finite: {}", result);
}
#[test]
fn test_interleaved_q4k_mod_dot_single_superblock_ones() {
let mut data = vec![0u8; 144];
data[0..2].copy_from_slice(&0x3C00u16.to_le_bytes());
data[2..4].copy_from_slice(&0x0000u16.to_le_bytes());
data[4] = 1;
for i in 16..144 {
data[i] = 0x11;
}
let interleaved = InterleavedQ4K::from_q4k(&data).expect("valid");
let activations = vec![1.0f32; 256];
let result = interleaved.dot(&activations).expect("dot should work");
assert!(result.is_finite());
}
#[test]
fn test_interleaved_q4k_mod_dot_two_superblocks() {
let mut data = vec![0u8; 288];
data[0..2].copy_from_slice(&0x3C00u16.to_le_bytes());
data[2..4].copy_from_slice(&0x0000u16.to_le_bytes());
data[144..146].copy_from_slice(&0x4000u16.to_le_bytes());
data[146..148].copy_from_slice(&0x0000u16.to_le_bytes());
let interleaved = InterleavedQ4K::from_q4k(&data).expect("valid");
let activations = vec![0.5f32; 512];
let result = interleaved.dot(&activations).expect("dot should work");
assert!(result.is_finite());
}
#[test]
fn test_interleaved_q4k_mod_dot_length_error() {
let data = vec![0u8; 144]; let interleaved = InterleavedQ4K::from_q4k(&data).expect("valid");
let activations = vec![1.0f32; 128];
let result = interleaved.dot(&activations);
assert!(result.is_err(), "Should error on length mismatch");
let activations = vec![1.0f32; 300];
let result = interleaved.dot(&activations);
assert!(result.is_err(), "Should error on length mismatch");
}
#[test]
fn test_interleaved_q4k_mod_dot_varied_activations() {
let mut data = vec![0u8; 144];
data[0..2].copy_from_slice(&0x3C00u16.to_le_bytes());
for i in 4..16 {
data[i] = 10;
}
for i in 16..144 {
data[i] = ((i - 16) % 16) as u8 | (((i - 16 + 1) % 16) << 4) as u8;
}
let interleaved = InterleavedQ4K::from_q4k(&data).expect("valid");
let activations: Vec<f32> = (0..256).map(|i| (i as f32 - 128.0) * 0.01).collect();
let result = interleaved.dot(&activations).expect("dot should work");
assert!(result.is_finite());
}
#[test]
fn test_interleaved_q4k_mod_dot_negative_activations() {
let mut data = vec![0u8; 144];
data[0..2].copy_from_slice(&0x3C00u16.to_le_bytes());
data[4] = 32;
let interleaved = InterleavedQ4K::from_q4k(&data).expect("valid");
let activations = vec![-1.0f32; 256];
let result = interleaved.dot(&activations).expect("dot should work");
assert!(result.is_finite());
}
#[test]
fn test_interleaved_q4k_mod_dot_with_dmin() {
let mut data = vec![0u8; 144];
data[0..2].copy_from_slice(&0x3C00u16.to_le_bytes()); data[2..4].copy_from_slice(&0x3C00u16.to_le_bytes());
data[4] = 0; data[8] = 10;
let interleaved = InterleavedQ4K::from_q4k(&data).expect("valid");
let activations = vec![1.0f32; 256];
let result = interleaved.dot(&activations).expect("dot should work");
assert!(result.is_finite());
}
#[test]
fn test_extract_scale_min_mod_all_blocks() {
let scales: [u8; 12] = [
0b00_000001, 0b00_000010, 0b00_000011, 0b00_000100, 0b00_000101, 0b00_000110, 0b00_000111, 0b00_001000, 0b0010_0001, 0b0100_0011, 0b0110_0101, 0b1000_0111, ];
let (s0, m0) = extract_scale_min(&scales, 0);
assert_eq!(s0, 1.0);
assert_eq!(m0, 5.0);
let (s1, m1) = extract_scale_min(&scales, 1);
assert_eq!(s1, 2.0);
assert_eq!(m1, 6.0);
let (s2, m2) = extract_scale_min(&scales, 2);
assert_eq!(s2, 3.0);
assert_eq!(m2, 7.0);
let (s3, m3) = extract_scale_min(&scales, 3);
assert_eq!(s3, 4.0);
assert_eq!(m3, 8.0);
let (s4, m4) = extract_scale_min(&scales, 4);
assert_eq!(s4, 1.0);
assert_eq!(m4, 2.0);
let (s5, m5) = extract_scale_min(&scales, 5);
assert_eq!(s5, 3.0);
assert_eq!(m5, 4.0);
let (s6, m6) = extract_scale_min(&scales, 6);
assert_eq!(s6, 5.0);
assert_eq!(m6, 6.0);
let (s7, m7) = extract_scale_min(&scales, 7);
assert_eq!(s7, 7.0);
assert_eq!(m7, 8.0);
}
#[test]
fn test_extract_scale_min_mod_high_bits_contribution() {
let mut scales: [u8; 12] = [0; 12];
scales[0] = 0b11_000000; scales[8] = 0b0000_0001;
let (s4, _) = extract_scale_min(&scales, 4);
assert_eq!(s4, 49.0, "Scale4 with high bits contribution");
}
#[test]
fn test_fused_q4_0_q8_0_dot_scalar_mod_empty() {
let q4_data: Vec<u8> = vec![];
let q8_scales: Vec<f32> = vec![];
let q8_quants: Vec<i8> = vec![];
let result = fused_q4_0_q8_0_dot_scalar(&q4_data, &q8_scales, &q8_quants, 0);
assert_eq!(result, 0.0);
}
#[test]
fn test_fused_q4_0_q8_0_dot_scalar_mod_partial_block() {
let mut q4_data = vec![0u8; 10];
q4_data[0..2].copy_from_slice(&0x3C00u16.to_le_bytes());
let q8_scales = vec![1.0f32];
let q8_quants = vec![1i8; 32];
let result = fused_q4_0_q8_0_dot_scalar(&q4_data, &q8_scales, &q8_quants, 32);
assert!(result.is_finite());
}
#[test]
fn test_fused_q4_0_q8_0_dot_scalar_mod_two_complete_blocks() {
let mut q4_data = vec![0u8; 36];
q4_data[0..2].copy_from_slice(&0x3C00u16.to_le_bytes());
for i in 2..18 {
q4_data[i] = 0x88;
}
q4_data[18..20].copy_from_slice(&0x4000u16.to_le_bytes());
for i in 20..36 {
q4_data[i] = 0x88;
}
let q8_scales = vec![1.0f32, 1.0f32];
let q8_quants = vec![1i8; 64];
let result = fused_q4_0_q8_0_dot_scalar(&q4_data, &q8_scales, &q8_quants, 64);
assert!(result.abs() < 1.0, "Expected near 0, got {}", result);
}
#[test]
fn test_fused_q4_0_q8_0_dot_scalar_mod_negative_quants() {
let mut q4_data = vec![0u8; 18];
q4_data[0..2].copy_from_slice(&0x3C00u16.to_le_bytes()); for i in 2..18 {
q4_data[i] = 0x00;
}
let q8_scales = vec![1.0f32];
let q8_quants = vec![10i8; 32];
let result = fused_q4_0_q8_0_dot_scalar(&q4_data, &q8_scales, &q8_quants, 32);
assert!(result < 0.0, "Should be negative with 0x00 q4 quants");
}
#[test]
fn test_fused_q4_0_q8_0_dot_scalar_mod_large_scale() {
let mut q4_data = vec![0u8; 18];
q4_data[0..2].copy_from_slice(&0x63D0u16.to_le_bytes());
for i in 2..18 {
q4_data[i] = 0xFF; }
let q8_scales = vec![1.0f32];
let q8_quants = vec![1i8; 32];
let result = fused_q4_0_q8_0_dot_scalar(&q4_data, &q8_scales, &q8_quants, 32);
assert!(result.is_finite());
}
#[test]
fn test_fused_q8_0_q8_0_dot_scalar_mod_empty() {
let q8_weight_data: Vec<u8> = vec![];
let q8_act_scales: Vec<f32> = vec![];
let q8_act_quants: Vec<i8> = vec![];
let result = fused_q8_0_q8_0_dot_scalar(&q8_weight_data, &q8_act_scales, &q8_act_quants, 0);
assert_eq!(result, 0.0);
}
#[test]
fn test_fused_q8_0_q8_0_dot_scalar_mod_partial_block() {
let mut q8_weight_data = vec![0u8; 20]; q8_weight_data[0..2].copy_from_slice(&0x3C00u16.to_le_bytes());
let q8_act_scales = vec![1.0f32];
let q8_act_quants = vec![1i8; 32];
let result = fused_q8_0_q8_0_dot_scalar(&q8_weight_data, &q8_act_scales, &q8_act_quants, 32);
assert!(result.is_finite());
}
#[test]
fn test_fused_q8_0_q8_0_dot_scalar_mod_two_blocks() {
let mut q8_weight_data = vec![0u8; 68];
q8_weight_data[0..2].copy_from_slice(&0x3C00u16.to_le_bytes());
for i in 2..34 {
q8_weight_data[i] = 5u8;
}
q8_weight_data[34..36].copy_from_slice(&0x3800u16.to_le_bytes());
for i in 36..68 {
q8_weight_data[i] = 10u8;
}
let q8_act_scales = vec![1.0f32, 1.0f32];
let q8_act_quants = vec![2i8; 64];
let result = fused_q8_0_q8_0_dot_scalar(&q8_weight_data, &q8_act_scales, &q8_act_quants, 64);
assert!(
(result - 640.0).abs() < 10.0,
"Expected ~640, got {}",
result
);
}
#[test]
fn test_fused_q8_0_q8_0_dot_scalar_mod_mixed_signs() {
let mut q8_weight_data = vec![0u8; 34];
q8_weight_data[0..2].copy_from_slice(&0x3C00u16.to_le_bytes());
for i in 2..18 {
q8_weight_data[i] = 10u8; }
for i in 18..34 {
q8_weight_data[i] = (-10i8) as u8; }
let q8_act_scales = vec![1.0f32];
let q8_act_quants = vec![5i8; 32];
let result = fused_q8_0_q8_0_dot_scalar(&q8_weight_data, &q8_act_scales, &q8_act_quants, 32);
assert!(
result.abs() < 10.0,
"Mixed signs should nearly cancel: {}",
result
);
}
include!("fused_02.rs");
include!("fused_02_02.rs");