use crate::quantize::parallel_k::{
fused_q4k_parallel_matvec, fused_q4k_parallel_matvec_into, fused_q4k_q8k_ffn_up_gate_into,
fused_q4k_q8k_parallel_matvec_into, fused_q4k_tiled_matvec, fused_q5k_parallel_matvec,
fused_q5k_parallel_matvec_into, fused_q6k_parallel_matvec, fused_q6k_parallel_matvec_into,
};
use crate::quantize::types::QK_K;
fn generate_q4k_weights(out_dim: usize, in_dim: usize) -> Vec<u8> {
let super_blocks_per_row = in_dim.div_ceil(QK_K);
let bytes_per_row = super_blocks_per_row * 144;
let total_bytes = out_dim * bytes_per_row;
let mut data = vec![0u8; total_bytes];
for (i, byte) in data.iter_mut().enumerate() {
*byte = ((i * 17 + 31) % 256) as u8;
}
for row in 0..out_dim {
for sb in 0..super_blocks_per_row {
let block_start = row * bytes_per_row + sb * 144;
data[block_start] = 0x00;
data[block_start + 1] = 0x3c;
data[block_start + 2] = 0x00;
data[block_start + 3] = 0x38; }
}
data
}
fn generate_q5k_weights(out_dim: usize, in_dim: usize) -> Vec<u8> {
let super_blocks_per_row = in_dim.div_ceil(QK_K);
let bytes_per_row = super_blocks_per_row * 176;
let total_bytes = out_dim * bytes_per_row;
let mut data = vec![0u8; total_bytes];
for (i, byte) in data.iter_mut().enumerate() {
*byte = ((i * 19 + 37) % 256) as u8;
}
for row in 0..out_dim {
for sb in 0..super_blocks_per_row {
let block_start = row * bytes_per_row + sb * 176;
data[block_start] = 0x00;
data[block_start + 1] = 0x3c;
data[block_start + 2] = 0x00;
data[block_start + 3] = 0x38;
}
}
data
}
fn generate_q6k_weights(out_dim: usize, in_dim: usize) -> Vec<u8> {
let super_blocks_per_row = in_dim.div_ceil(QK_K);
let bytes_per_row = super_blocks_per_row * 210;
let total_bytes = out_dim * bytes_per_row;
let mut data = vec![0u8; total_bytes];
for (i, byte) in data.iter_mut().enumerate() {
*byte = ((i * 23 + 41) % 256) as u8;
}
for row in 0..out_dim {
for sb in 0..super_blocks_per_row {
let block_start = row * bytes_per_row + sb * 210;
data[block_start + 208] = 0x00;
data[block_start + 209] = 0x3c;
}
}
data
}
fn generate_q8k_activations(in_dim: usize) -> (Vec<f32>, Vec<i8>) {
let super_blocks = in_dim.div_ceil(QK_K);
let scales: Vec<f32> = (0..super_blocks).map(|i| 0.1 + (i as f32) * 0.01).collect();
let quants: Vec<i8> = (0..in_dim)
.map(|i| ((i % 256) as i8).wrapping_sub(64))
.collect();
(scales, quants)
}
#[test]
fn test_q4k_tiled_matvec_single_row_pk14() {
let in_dim = 256; let out_dim = 1;
let weights = generate_q4k_weights(out_dim, in_dim);
let activations = vec![1.0f32; in_dim];
let result = fused_q4k_tiled_matvec(&weights, &activations, in_dim, out_dim, None);
assert!(result.is_ok());
let output = result.expect("test value should be present");
assert_eq!(output.len(), out_dim);
assert!(output[0].is_finite());
}
#[test]
fn test_q4k_tiled_matvec_multiple_rows_pk14() {
let in_dim = 512; let out_dim = 32;
let weights = generate_q4k_weights(out_dim, in_dim);
let activations = vec![0.5f32; in_dim];
let result = fused_q4k_tiled_matvec(&weights, &activations, in_dim, out_dim, None);
assert!(result.is_ok());
let output = result.expect("test value should be present");
assert_eq!(output.len(), out_dim);
for val in &output {
assert!(val.is_finite());
}
}
#[test]
fn test_q4k_tiled_matvec_custom_tile_size_pk14() {
let in_dim = 256;
let out_dim = 128;
let weights = generate_q4k_weights(out_dim, in_dim);
let activations = vec![0.25f32; in_dim];
let result = fused_q4k_tiled_matvec(&weights, &activations, in_dim, out_dim, Some(16));
assert!(result.is_ok());
let output = result.expect("test value should be present");
assert_eq!(output.len(), out_dim);
}
#[test]
fn test_q4k_tiled_matvec_partial_last_tile_pk14() {
let in_dim = 256;
let out_dim = 100; let weights = generate_q4k_weights(out_dim, in_dim);
let activations = vec![1.0f32; in_dim];
let result = fused_q4k_tiled_matvec(&weights, &activations, in_dim, out_dim, None);
assert!(result.is_ok());
let output = result.expect("test value should be present");
assert_eq!(output.len(), out_dim);
}
#[test]
fn test_q4k_tiled_matvec_many_tiles_pk14() {
let in_dim = 256;
let out_dim = 256; let weights = generate_q4k_weights(out_dim, in_dim);
let activations = vec![0.1f32; in_dim];
let result = fused_q4k_tiled_matvec(&weights, &activations, in_dim, out_dim, None);
assert!(result.is_ok());
let output = result.expect("test value should be present");
assert_eq!(output.len(), out_dim);
}
#[test]
fn test_q4k_tiled_matvec_weight_too_small_pk14() {
let in_dim = 256;
let out_dim = 10;
let weights = vec![0u8; 100]; let activations = vec![1.0f32; in_dim];
let result = fused_q4k_tiled_matvec(&weights, &activations, in_dim, out_dim, None);
assert!(result.is_err());
let err = result.unwrap_err();
assert!(err.to_string().contains("weight data too small"));
}
#[test]
fn test_q4k_tiled_matvec_activation_mismatch_pk14() {
let in_dim = 256;
let out_dim = 10;
let weights = generate_q4k_weights(out_dim, in_dim);
let activations = vec![1.0f32; 128];
let result = fused_q4k_tiled_matvec(&weights, &activations, in_dim, out_dim, None);
assert!(result.is_err());
let err = result.unwrap_err();
assert!(err.to_string().contains("doesn't match in_dim"));
}
#[test]
fn test_q4k_tiled_matvec_tile_size_1_pk14() {
let in_dim = 256;
let out_dim = 8;
let weights = generate_q4k_weights(out_dim, in_dim);
let activations = vec![1.0f32; in_dim];
let result = fused_q4k_tiled_matvec(&weights, &activations, in_dim, out_dim, Some(1));
assert!(result.is_ok());
let output = result.expect("test value should be present");
assert_eq!(output.len(), out_dim);
}
#[test]
fn test_q4k_tiled_matvec_tile_larger_than_out_pk14() {
let in_dim = 256;
let out_dim = 10;
let weights = generate_q4k_weights(out_dim, in_dim);
let activations = vec![1.0f32; in_dim];
let result = fused_q4k_tiled_matvec(&weights, &activations, in_dim, out_dim, Some(128));
assert!(result.is_ok());
let output = result.expect("test value should be present");
assert_eq!(output.len(), out_dim);
}
#[test]
fn test_q4k_parallel_matvec_sequential_path_pk14() {
let in_dim = 256;
let out_dim = 128; let weights = generate_q4k_weights(out_dim, in_dim);
let activations = vec![1.0f32; in_dim];
let result = fused_q4k_parallel_matvec(&weights, &activations, in_dim, out_dim);
assert!(result.is_ok());
let output = result.expect("test value should be present");
assert_eq!(output.len(), out_dim);
for val in &output {
assert!(val.is_finite());
}
}
#[test]
fn test_q4k_parallel_matvec_at_threshold_pk14() {
let in_dim = 256;
let out_dim = 256; let weights = generate_q4k_weights(out_dim, in_dim);
let activations = vec![0.5f32; in_dim];
let result = fused_q4k_parallel_matvec(&weights, &activations, in_dim, out_dim);
assert!(result.is_ok());
let output = result.expect("test value should be present");
assert_eq!(output.len(), out_dim);
}
#[test]
fn test_q4k_parallel_matvec_parallel_path_pk14() {
let in_dim = 512;
let out_dim = 512; let weights = generate_q4k_weights(out_dim, in_dim);
let activations = vec![0.1f32; in_dim];
let result = fused_q4k_parallel_matvec(&weights, &activations, in_dim, out_dim);
assert!(result.is_ok());
let output = result.expect("test value should be present");
assert_eq!(output.len(), out_dim);
}
#[test]
fn test_q4k_parallel_matvec_single_row_pk14() {
let in_dim = 256;
let out_dim = 1;
let weights = generate_q4k_weights(out_dim, in_dim);
let activations = vec![1.0f32; in_dim];
let result = fused_q4k_parallel_matvec(&weights, &activations, in_dim, out_dim);
assert!(result.is_ok());
let output = result.expect("test value should be present");
assert_eq!(output.len(), 1);
}
#[test]
fn test_q4k_parallel_matvec_weight_error_pk14() {
let in_dim = 256;
let out_dim = 64;
let weights = vec![0u8; 10]; let activations = vec![1.0f32; in_dim];
let result = fused_q4k_parallel_matvec(&weights, &activations, in_dim, out_dim);
assert!(result.is_err());
}
#[test]
fn test_q4k_parallel_matvec_activation_error_pk14() {
let in_dim = 256;
let out_dim = 64;
let weights = generate_q4k_weights(out_dim, in_dim);
let activations = vec![1.0f32; 64];
let result = fused_q4k_parallel_matvec(&weights, &activations, in_dim, out_dim);
assert!(result.is_err());
}
#[test]
fn test_q4k_parallel_matvec_into_sequential_pk14() {
let in_dim = 256;
let out_dim = 128; let weights = generate_q4k_weights(out_dim, in_dim);
let activations = vec![1.0f32; in_dim];
let mut output = vec![0.0f32; out_dim];
let result =
fused_q4k_parallel_matvec_into(&weights, &activations, in_dim, out_dim, &mut output);
assert!(result.is_ok());
for val in &output {
assert!(val.is_finite());
}
}
#[test]
fn test_q4k_parallel_matvec_into_parallel_pk14() {
let in_dim = 256;
let out_dim = 512; let weights = generate_q4k_weights(out_dim, in_dim);
let activations = vec![0.5f32; in_dim];
let mut output = vec![0.0f32; out_dim];
let result =
fused_q4k_parallel_matvec_into(&weights, &activations, in_dim, out_dim, &mut output);
assert!(result.is_ok());
assert_eq!(output.len(), out_dim);
}
#[test]
fn test_q4k_parallel_matvec_into_midi_tile_boundary_pk14() {
let in_dim = 256;
let out_dim = 64; let weights = generate_q4k_weights(out_dim, in_dim);
let activations = vec![1.0f32; in_dim];
let mut output = vec![0.0f32; out_dim];
let result =
fused_q4k_parallel_matvec_into(&weights, &activations, in_dim, out_dim, &mut output);
assert!(result.is_ok());
}
#[test]
fn test_q4k_parallel_matvec_into_partial_midi_tile_pk14() {
let in_dim = 256;
let out_dim = 300; let weights = generate_q4k_weights(out_dim, in_dim);
let activations = vec![0.25f32; in_dim];
let mut output = vec![0.0f32; out_dim];
let result =
fused_q4k_parallel_matvec_into(&weights, &activations, in_dim, out_dim, &mut output);
assert!(result.is_ok());
assert_eq!(output.len(), out_dim);
}
#[test]
fn test_q4k_parallel_matvec_into_output_buffer_too_small_pk14() {
let in_dim = 256;
let out_dim = 128;
let weights = generate_q4k_weights(out_dim, in_dim);
let activations = vec![1.0f32; in_dim];
let mut output = vec![0.0f32; 64];
let result =
fused_q4k_parallel_matvec_into(&weights, &activations, in_dim, out_dim, &mut output);
assert!(result.is_err());
let err = result.unwrap_err();
assert!(err.to_string().contains("Output buffer too small"));
}
include!("q4k_parallel.rs");
include!("q4k_q8k.rs");