use crate::error::RealizarError;
use crate::quantize::fused_k::{fused_q4k_dot, fused_q4k_q8k_dot};
use crate::quantize::types::QK_K;
const Q4K_SUPER_BLOCK_BYTES: usize = 144;
const Q8K_BLOCK_SIZE: usize = 256;
#[test]
fn test_fused_q4k_dot_empty_data() {
let result = fused_q4k_dot(&[], &[]);
assert!(result.is_ok());
assert_eq!(result.expect("test value should be present"), 0.0);
}
#[test]
fn test_fused_q4k_dot_invalid_block_size() {
let invalid_data = vec![0u8; 143];
let activations = vec![1.0f32; 256];
let result = fused_q4k_dot(&invalid_data, &activations);
assert!(result.is_err());
let err = result.unwrap_err();
match err {
RealizarError::InvalidShape { reason } => {
assert!(reason.contains("not a multiple"));
},
_ => panic!("Expected InvalidShape error"),
}
}
#[test]
fn test_fused_q4k_dot_activation_length_mismatch() {
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES];
let wrong_activations = vec![1.0f32; 128];
let result = fused_q4k_dot(&q4k_data, &wrong_activations);
assert!(result.is_err());
let err = result.unwrap_err();
match err {
RealizarError::InvalidShape { reason } => {
assert!(reason.contains("doesn't match"));
},
_ => panic!("Expected InvalidShape error"),
}
}
#[test]
fn test_fused_q4k_dot_zero_scales() {
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES];
let activations = vec![1.0f32; QK_K];
let result = fused_q4k_dot(&q4k_data, &activations);
assert!(result.is_ok());
assert_eq!(result.expect("test value should be present"), 0.0);
}
#[test]
fn test_fused_q4k_dot_inf_activations() {
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES];
let mut activations = vec![1.0f32; QK_K];
activations[0] = f32::INFINITY;
let result = fused_q4k_dot(&q4k_data, &activations);
assert!(result.is_ok());
}
#[test]
fn test_fused_q4k_dot_neg_inf_activations() {
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES];
let mut activations = vec![1.0f32; QK_K];
activations[128] = f32::NEG_INFINITY;
let result = fused_q4k_dot(&q4k_data, &activations);
assert!(result.is_ok());
}
#[test]
fn test_fused_q4k_dot_nan_activations() {
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES];
let mut activations = vec![1.0f32; QK_K];
activations[64] = f32::NAN;
let result = fused_q4k_dot(&q4k_data, &activations);
assert!(result.is_ok());
}
#[test]
fn test_fused_q4k_dot_multiple_super_blocks() {
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES * 2];
let activations = vec![1.0f32; QK_K * 2];
let result = fused_q4k_dot(&q4k_data, &activations);
assert!(result.is_ok());
}
#[test]
fn test_fused_q4k_dot_three_super_blocks() {
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES * 3];
let activations = vec![0.5f32; QK_K * 3];
let result = fused_q4k_dot(&q4k_data, &activations);
assert!(result.is_ok());
}
#[test]
fn test_fused_q4k_q8k_dot_empty() {
let result = fused_q4k_q8k_dot(&[], &[], &[]);
assert!(result.is_ok());
assert_eq!(result.expect("test value should be present"), 0.0);
}
#[test]
fn test_fused_q4k_q8k_dot_invalid_q4k_length() {
let invalid_q4k = vec![0u8; 100]; let scales = vec![1.0f32; 1];
let quants = vec![0i8; QK_K];
let result = fused_q4k_q8k_dot(&invalid_q4k, &scales, &quants);
assert!(result.is_err());
}
#[test]
fn test_fused_q4k_q8k_dot_scale_length_mismatch() {
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES];
let extra_scales = vec![1.0f32; 5]; let quants = vec![0i8; QK_K];
let result = fused_q4k_q8k_dot(&q4k_data, &extra_scales, &quants);
let _ = result; }
#[test]
fn test_fused_q4k_q8k_dot_quants_length_mismatch() {
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES];
let scales = vec![1.0f32; 1];
let wrong_quants = vec![0i8; 128];
let result = fused_q4k_q8k_dot(&q4k_data, &scales, &wrong_quants);
assert!(result.is_err());
}
#[test]
fn test_fused_q4k_q8k_dot_zero_scales() {
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES];
let scales = vec![0.0f32; 1]; let quants = vec![127i8; QK_K];
let result = fused_q4k_q8k_dot(&q4k_data, &scales, &quants);
assert!(result.is_ok());
assert_eq!(result.expect("test value should be present"), 0.0);
}
#[test]
fn test_fused_q4k_q8k_dot_inf_scale() {
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES];
let scales = vec![f32::INFINITY; 1];
let quants = vec![1i8; QK_K];
let result = fused_q4k_q8k_dot(&q4k_data, &scales, &quants);
assert!(result.is_ok());
}
#[test]
fn test_fused_q4k_q8k_dot_negative_scale() {
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES];
let scales = vec![-1.0f32; 1]; let quants = vec![50i8; QK_K];
let result = fused_q4k_q8k_dot(&q4k_data, &scales, &quants);
assert!(result.is_ok());
}
#[test]
fn test_fused_q4k_q8k_dot_extreme_quants() {
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES];
let scales = vec![1.0f32; 1];
let mut quants = vec![0i8; QK_K];
quants[0] = i8::MAX; quants[1] = i8::MIN; quants[128] = 64;
quants[255] = -64;
let result = fused_q4k_q8k_dot(&q4k_data, &scales, &quants);
assert!(result.is_ok());
}
#[test]
fn test_fused_q4k_q8k_dot_multiple_blocks() {
let num_blocks = 4;
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES * num_blocks];
let scales = vec![1.0f32; num_blocks];
let quants = vec![1i8; QK_K * num_blocks];
let result = fused_q4k_q8k_dot(&q4k_data, &scales, &quants);
assert!(result.is_ok());
}
#[test]
fn test_fused_q4k_dot_single_block_boundary() {
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES];
let activations = vec![0.0f32; QK_K];
let result = fused_q4k_dot(&q4k_data, &activations);
assert!(result.is_ok());
assert_eq!(result.expect("test value should be present"), 0.0);
}
#[test]
fn test_fused_q4k_dot_max_reasonable_blocks() {
let num_blocks = 16;
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES * num_blocks];
let activations = vec![1.0f32; QK_K * num_blocks];
let result = fused_q4k_dot(&q4k_data, &activations);
assert!(result.is_ok());
}
#[test]
fn test_fused_q4k_dot_alternating_activations() {
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES];
let mut activations = vec![0.0f32; QK_K];
for i in 0..QK_K {
activations[i] = if i % 2 == 0 { 1.0 } else { -1.0 };
}
let result = fused_q4k_dot(&q4k_data, &activations);
assert!(result.is_ok());
}
#[test]
fn test_fused_q4k_dot_subnormal_activations() {
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES];
let mut activations = vec![f32::MIN_POSITIVE; QK_K];
activations[0] = f32::MIN_POSITIVE / 2.0;
let result = fused_q4k_dot(&q4k_data, &activations);
assert!(result.is_ok());
}
#[test]
fn test_fused_q4k_dot_max_f32_activations() {
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES];
let activations = vec![f32::MAX / 256.0; QK_K];
let result = fused_q4k_dot(&q4k_data, &activations);
assert!(result.is_ok());
}
#[test]
fn test_fused_q4k_dot_mixed_signs() {
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES];
let mut activations = vec![0.0f32; QK_K];
for i in 0..QK_K / 2 {
activations[i] = 1.0;
}
for i in QK_K / 2..QK_K {
activations[i] = -1.0;
}
let result = fused_q4k_dot(&q4k_data, &activations);
assert!(result.is_ok());
}
#[test]
fn test_fused_q4k_q8k_dot_all_zeros() {
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES];
let scales = vec![0.0f32; 1];
let quants = vec![0i8; QK_K];
let result = fused_q4k_q8k_dot(&q4k_data, &scales, &quants);
assert!(result.is_ok());
assert_eq!(result.expect("test value should be present"), 0.0);
}
#[test]
fn test_fused_q4k_q8k_dot_nan_scale() {
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES];
let scales = vec![f32::NAN; 1];
let quants = vec![1i8; QK_K];
let result = fused_q4k_q8k_dot(&q4k_data, &scales, &quants);
assert!(result.is_ok());
assert!(result.expect("test value should be present").is_nan());
}
#[test]
fn test_fused_q4k_q8k_dot_alternating_quants() {
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES];
let scales = vec![1.0f32; 1];
let mut quants = vec![0i8; QK_K];
for i in 0..QK_K {
quants[i] = if i % 2 == 0 { 127 } else { -128 };
}
let result = fused_q4k_q8k_dot(&q4k_data, &scales, &quants);
assert!(result.is_ok());
}
#[test]
fn test_fused_q4k_dot_large_model_layer() {
let num_blocks = 16;
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES * num_blocks];
let activations = vec![0.1f32; QK_K * num_blocks];
let result = fused_q4k_dot(&q4k_data, &activations);
assert!(result.is_ok());
}
#[test]
fn test_fused_q4k_q8k_dot_large_model_layer() {
let num_blocks = 16;
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES * num_blocks];
let scales = vec![1.0f32; num_blocks];
let quants = vec![1i8; QK_K * num_blocks];
let result = fused_q4k_q8k_dot(&q4k_data, &scales, &quants);
assert!(result.is_ok());
}
#[test]
fn test_fused_q4k_dot_deterministic() {
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES];
let activations = vec![1.0f32; QK_K];
let result1 = fused_q4k_dot(&q4k_data, &activations).expect("test value should be present");
let result2 = fused_q4k_dot(&q4k_data, &activations).expect("test value should be present");
assert_eq!(result1, result2);
}
#[test]
fn test_fused_q4k_q8k_dot_deterministic() {
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES];
let scales = vec![1.0f32; 1];
let quants = vec![42i8; QK_K];
let result1 = fused_q4k_q8k_dot(&q4k_data, &scales, &quants).expect("test value should be present");
let result2 = fused_q4k_q8k_dot(&q4k_data, &scales, &quants).expect("test value should be present");
assert_eq!(result1, result2);
}
#[test]
fn test_fused_q4k_dot_error_message_includes_sizes() {
let invalid_data = vec![0u8; 100];
let activations = vec![1.0f32; 256];
let result = fused_q4k_dot(&invalid_data, &activations);
let err_str = format!("{:?}", result.unwrap_err());
assert!(err_str.contains("100") || err_str.contains("144"));
}
#[test]
fn test_fused_q4k_dot_error_message_includes_expected() {
let q4k_data = vec![0u8; Q4K_SUPER_BLOCK_BYTES];
let wrong_activations = vec![1.0f32; 100];
let result = fused_q4k_dot(&q4k_data, &wrong_activations);
let err_str = format!("{:?}", result.unwrap_err());
assert!(err_str.contains("100") || err_str.contains("256"));
}