use crate::quantize::activation::{
fused_swiglu_scalar, quantize_activations_q8_0, quantize_rmsnorm_q8_0_into,
quantize_rmsnorm_q8_0_scalar, softmax_scalar,
};
use crate::quantize::{
dequantize_q8_blocks, quantize_activations_q8k_into, quantize_to_q8_blocks, InterleavedQ4K,
};
#[test]
fn test_q8k_into_not_multiple_of_256() {
let activations = vec![1.0f32; 100]; let mut scales = vec![0.0f32; 1];
let mut quants = vec![0i8; 256];
let result = quantize_activations_q8k_into(&activations, &mut scales, &mut quants);
assert!(result.is_err());
let err = result.unwrap_err().to_string();
assert!(
err.contains("multiple of 256"),
"Expected multiple-of-256 error, got: {}",
err
);
}
#[test]
fn test_q8k_into_scales_buffer_too_small() {
let activations = vec![1.0f32; 256];
let mut scales = vec![0.0f32; 0]; let mut quants = vec![0i8; 256];
let result = quantize_activations_q8k_into(&activations, &mut scales, &mut quants);
assert!(result.is_err());
let err = result.unwrap_err().to_string();
assert!(
err.contains("too small"),
"Expected buffer-too-small error, got: {}",
err
);
}
#[test]
fn test_q8k_into_quants_buffer_too_small() {
let activations = vec![1.0f32; 256];
let mut scales = vec![0.0f32; 1];
let mut quants = vec![0i8; 100]; let result = quantize_activations_q8k_into(&activations, &mut scales, &mut quants);
assert!(result.is_err());
let err = result.unwrap_err().to_string();
assert!(
err.contains("too small"),
"Expected buffer-too-small error, got: {}",
err
);
}
#[test]
fn test_q8k_into_success_single_block() {
let activations = vec![1.0f32; 256];
let mut scales = vec![0.0f32; 1];
let mut quants = vec![0i8; 256];
let result = quantize_activations_q8k_into(&activations, &mut scales, &mut quants);
assert!(result.is_ok());
assert!(scales[0] > 0.0);
let first = quants[0];
for &q in &quants {
assert_eq!(q, first);
}
}
#[test]
fn test_q8k_into_success_multiple_blocks() {
let activations: Vec<f32> = (0..512).map(|i| (i as f32 - 256.0) / 100.0).collect();
let mut scales = vec![0.0f32; 2];
let mut quants = vec![0i8; 512];
let result = quantize_activations_q8k_into(&activations, &mut scales, &mut quants);
assert!(result.is_ok());
assert!(scales[0] > 0.0);
assert!(scales[1] > 0.0);
}
#[test]
fn test_q8k_into_zero_activations() {
let activations = vec![0.0f32; 256];
let mut scales = vec![0.0f32; 1];
let mut quants = vec![0i8; 256];
let result = quantize_activations_q8k_into(&activations, &mut scales, &mut quants);
assert!(result.is_ok());
assert!(scales[0] > 0.0); for &q in &quants {
assert_eq!(q, 0);
}
}
#[test]
fn test_q8_blocks_not_multiple_of_32() {
let values = vec![1.0f32; 50]; let result = quantize_to_q8_blocks(&values);
assert!(result.is_err());
let err = result.unwrap_err().to_string();
assert!(
err.contains("multiple of 32"),
"Expected multiple-of-32 error, got: {}",
err
);
}
#[test]
fn test_q8_blocks_roundtrip_uniform() {
let values = vec![42.0f32; 64]; let blocks = quantize_to_q8_blocks(&values).expect("test value should be present");
assert_eq!(blocks.len(), 2);
let dequantized = dequantize_q8_blocks(&blocks);
assert_eq!(dequantized.len(), 64);
for (orig, deq) in values.iter().zip(dequantized.iter()) {
assert!(
(orig - deq).abs() < 1.0,
"Roundtrip error: orig={}, deq={}",
orig,
deq
);
}
}
#[test]
fn test_q8_blocks_roundtrip_mixed() {
let values: Vec<f32> = (0..96).map(|i| (i as f32 - 48.0) * 2.0).collect();
let blocks = quantize_to_q8_blocks(&values).expect("test value should be present");
assert_eq!(blocks.len(), 3);
let dequantized = dequantize_q8_blocks(&blocks);
assert_eq!(dequantized.len(), 96);
for (orig, deq) in values.iter().zip(dequantized.iter()) {
let diff = (orig - deq).abs();
assert!(diff < 2.0, "Roundtrip error too large: diff={}", diff);
}
}
#[test]
fn test_q8_blocks_roundtrip_zeros() {
let values = vec![0.0f32; 32];
let blocks = quantize_to_q8_blocks(&values).expect("test value should be present");
let dequantized = dequantize_q8_blocks(&blocks);
for deq in &dequantized {
assert!((deq - 0.0).abs() < 0.01);
}
}
#[test]
fn test_q8_blocks_empty() {
let values: Vec<f32> = vec![];
let blocks = quantize_to_q8_blocks(&values).expect("test value should be present");
assert!(blocks.is_empty());
let dequantized = dequantize_q8_blocks(&blocks);
assert!(dequantized.is_empty());
}
#[test]
fn test_interleaved_q4k_dot_empty() {
let data = vec![];
let iq = InterleavedQ4K::from_q4k(&data).expect("test value should be present");
let activations = vec![];
let result = iq.dot(&activations);
assert!(result.is_ok());
assert_eq!(result.expect("test value should be present"), 0.0);
}
#[test]
fn test_interleaved_q4k_dot_mismatch() {
let data = vec![0u8; 144]; let iq = InterleavedQ4K::from_q4k(&data).expect("test value should be present");
let activations = vec![1.0f32; 128]; let result = iq.dot(&activations);
assert!(result.is_err());
}
#[test]
fn test_interleaved_q4k_dot_zero_weights() {
let data = vec![0u8; 144]; let iq = InterleavedQ4K::from_q4k(&data).expect("test value should be present");
let activations = vec![1.0f32; 256];
let result = iq.dot(&activations).expect("test value should be present");
assert_eq!(result, 0.0); }
#[test]
fn test_interleaved_q4k_dot_nonzero() {
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());
for i in 0..12 {
data[4 + i] = 0x01;
}
for i in 0..128 {
data[16 + i] = 0x11;
}
let iq = InterleavedQ4K::from_q4k(&data).expect("test value should be present");
let activations = vec![1.0f32; 256];
let result = iq.dot(&activations).expect("test value should be present");
assert!(result.abs() > 0.0, "Expected non-zero dot product");
}
#[test]
fn test_swiglu_scalar_zeros() {
let mut gate = vec![0.0f32; 8];
let up = vec![1.0f32; 8];
fused_swiglu_scalar(&mut gate, &up);
for &g in &gate {
assert!((g - 0.0).abs() < 1e-6);
}
}
#[test]
fn test_swiglu_scalar_positive() {
let mut gate = vec![2.0f32; 4];
let up = vec![1.0f32; 4];
fused_swiglu_scalar(&mut gate, &up);
for &g in &gate {
assert!((g - 1.7616).abs() < 0.01, "got {}", g);
}
}
#[test]
fn test_swiglu_scalar_negative() {
let mut gate = vec![-5.0f32; 4];
let up = vec![1.0f32; 4];
fused_swiglu_scalar(&mut gate, &up);
for &g in &gate {
assert!((g - (-0.0337)).abs() < 0.01, "got {}", g);
}
}
#[test]
fn test_swiglu_scalar_with_up_scaling() {
let mut gate = vec![1.0f32; 4];
let up = vec![3.0f32; 4];
fused_swiglu_scalar(&mut gate, &up);
for &g in &gate {
assert!((g - 2.1932).abs() < 0.01, "got {}", g);
}
}
#[test]
fn test_swiglu_scalar_empty() {
let mut gate: Vec<f32> = vec![];
let up: Vec<f32> = vec![];
fused_swiglu_scalar(&mut gate, &up);
assert!(gate.is_empty());
}
#[test]
fn test_softmax_scalar_uniform() {
let mut x = vec![1.0f32; 4];
softmax_scalar(&mut x);
for &v in &x {
assert!((v - 0.25).abs() < 1e-5, "got {}", v);
}
}
#[test]
fn test_softmax_scalar_sums_to_one() {
let mut x = vec![1.0, 2.0, 3.0, 4.0, 5.0];
softmax_scalar(&mut x);
let sum: f32 = x.iter().sum();
assert!((sum - 1.0).abs() < 1e-5, "sum should be 1.0, got {}", sum);
}
#[test]
fn test_softmax_scalar_monotone() {
let mut x = vec![1.0, 2.0, 3.0, 4.0];
softmax_scalar(&mut x);
for i in 1..x.len() {
assert!(x[i] >= x[i - 1], "softmax should be monotone");
}
}
#[test]
fn test_softmax_scalar_single_element() {
let mut x = vec![42.0f32];
softmax_scalar(&mut x);
assert!((x[0] - 1.0).abs() < 1e-6);
}
#[test]
fn test_softmax_scalar_large_values() {
let mut x = vec![1000.0, 1001.0, 999.0];
softmax_scalar(&mut x);
let sum: f32 = x.iter().sum();
assert!((sum - 1.0).abs() < 1e-5, "sum should be 1.0, got {}", sum);
assert!(x[1] > x[0]); assert!(x[0] > x[2]); }
#[test]
fn test_softmax_scalar_negative_values() {
let mut x = vec![-1.0, -2.0, -3.0];
softmax_scalar(&mut x);
let sum: f32 = x.iter().sum();
assert!((sum - 1.0).abs() < 1e-5);
assert!(x[0] > x[1]); }
#[test]
fn test_q8_0_activation_roundtrip() {
let activations: Vec<f32> = (0..64).map(|i| (i as f32 - 32.0) / 10.0).collect();
let (scales, quants) = quantize_activations_q8_0(&activations);
assert_eq!(scales.len(), 2); assert_eq!(quants.len(), 64);
for block in 0..2 {
let scale = scales[block];
for i in 0..32 {
let idx = block * 32 + i;
let dequantized = quants[idx] as f32 * scale;
let diff = (activations[idx] - dequantized).abs();
assert!(
diff < scale * 2.0,
"Block {}, idx {}: diff={}",
block,
i,
diff
);
}
}
}
#[test]
fn test_q8_0_activation_partial_block() {
let activations: Vec<f32> = (0..40).map(|i| i as f32).collect();
let (scales, quants) = quantize_activations_q8_0(&activations);
assert_eq!(scales.len(), 2); assert_eq!(quants.len(), 64);
for i in 40..64 {
assert_eq!(quants[i], 0, "Padding at index {} should be 0", i);
}
}
#[test]
fn test_q8_0_activation_near_zero() {
let activations = vec![1e-12f32; 32];
let (scales, _quants) = quantize_activations_q8_0(&activations);
assert_eq!(scales.len(), 1);
assert!(scales[0] > 0.0);
}
#[test]
fn test_rmsnorm_q8_scalar_identity() {
let input = vec![1.0f32; 32];
let norm_weight = vec![1.0f32; 32];
let eps = 1e-5;
let (scales, quants) = quantize_rmsnorm_q8_0_scalar(&input, &norm_weight, eps);
assert_eq!(scales.len(), 1);
assert_eq!(quants.len(), 32);
for (i, &q) in quants.iter().enumerate() {
let dequantized = q as f32 * scales[0];
assert!(
(dequantized - 1.0).abs() < 0.1,
"Element {}: dequant={}, expected ~1.0",
i,
dequantized
);
}
}
#[test]
fn test_rmsnorm_q8_scalar_zeros() {
let input = vec![0.0f32; 32];
let norm_weight = vec![1.0f32; 32];
let eps = 1e-5;
let (scales, quants) = quantize_rmsnorm_q8_0_scalar(&input, &norm_weight, eps);
assert_eq!(scales.len(), 1);
for &q in &quants {
assert_eq!(q, 0, "Zero input should give zero quants");
}
}
include!("rmsnorm.rs");