#[test]
fn test_fused_q4k_dot_single_nonzero_activation() {
let mut scales = [0u8; 12];
for j in 0..4 {
scales[j] = 1;
}
for j in 8..12 {
scales[j] = 0x01;
}
let qs = [0x55u8; 128]; let data = build_q4k_superblock(F16_ONE, F16_ZERO, &scales, &qs);
let mut activations = vec![0.0f32; QK_K];
activations[0] = 1.0;
let result = fused_q4k_dot(&data, &activations).expect("single nonzero");
assert!(result.is_finite());
}
#[test]
fn test_fused_q4k_dot_last_activation_nonzero() {
let mut scales = [0u8; 12];
for j in 0..4 {
scales[j] = 1;
}
for j in 8..12 {
scales[j] = 0x01;
}
let qs = [0x55u8; 128];
let data = build_q4k_superblock(F16_ONE, F16_ZERO, &scales, &qs);
let mut activations = vec![0.0f32; QK_K];
activations[QK_K - 1] = 1.0;
let result = fused_q4k_dot(&data, &activations).expect("last position");
assert!(result.is_finite());
}
#[test]
fn test_fused_q4k_dot_chunk_boundary_activations() {
let mut scales = [0u8; 12];
for j in 0..4 {
scales[j] = 3;
}
for j in 8..12 {
scales[j] = 0x03;
}
let qs = [0xAAu8; 128]; let data = build_q4k_superblock(F16_ONE, F16_ZERO, &scales, &qs);
for boundary in [0, 32, 64, 96, 128, 160, 192, 224] {
let mut activations = vec![0.0f32; QK_K];
activations[boundary] = 1.0;
let result = fused_q4k_dot(&data, &activations).expect(&format!("boundary {boundary}"));
assert!(
result.is_finite(),
"Failed at boundary {boundary}: {result}"
);
}
}
#[test]
fn test_fused_q4k_dot_large_d() {
let mut scales = [0u8; 12];
for j in 0..4 {
scales[j] = 1;
}
for j in 8..12 {
scales[j] = 0x01;
}
let qs = [0x11u8; 128]; let data = build_q4k_superblock(0x7BFF, F16_ZERO, &scales, &qs);
let activations = vec![1.0f32; QK_K];
let result = fused_q4k_dot(&data, &activations).expect("large d");
assert!(
result.is_finite(),
"Large d should still be finite: {result}"
);
assert!(result > 0.0);
}
#[test]
fn test_fused_q4k_dot_subnormal_d() {
let mut scales = [0u8; 12];
for j in 0..4 {
scales[j] = 63; }
for j in 8..12 {
scales[j] = 0x0F; }
let qs = [0xFFu8; 128]; let data = build_q4k_superblock(0x0001, F16_ZERO, &scales, &qs);
let activations = vec![1.0f32; QK_K];
let result = fused_q4k_dot(&data, &activations).expect("subnormal d");
assert!(result.is_finite());
assert!(
result.abs() < 1.0,
"Subnormal d should give small result: {result}"
);
}
#[test]
fn test_fused_q4k_dot_simd_empty() {
let result = fused_q4k_dot_simd(&[], &[]).expect("empty simd");
assert_eq!(result, 0.0);
}
#[test]
fn test_fused_q4k_q8k_dot_simd_empty() {
let result = fused_q4k_q8k_dot_simd(&[], &[], &[]).expect("empty simd q8k");
assert_eq!(result, 0.0);
}
#[test]
fn test_fused_q4k_dot_simd_all_chunks_nonzero() {
let mut scales = [0u8; 12];
for j in 0..4 {
scales[j] = 10 + j as u8;
scales[j + 4] = 3 + j as u8;
}
for j in 8..12 {
scales[j] = 0x3A + (j - 8) as u8;
}
let mut qs = [0u8; 128];
for i in 0..32 {
qs[i] = 0x12; }
for i in 32..64 {
qs[i] = 0x34; }
for i in 64..96 {
qs[i] = 0x56; }
for i in 96..128 {
qs[i] = 0x78; }
let data = build_q4k_superblock(F16_ONE, F16_QUARTER, &scales, &qs);
let activations: Vec<f32> = (0..QK_K).map(|i| (i as f32 * 0.01) + 0.1).collect();
let scalar = fused_q4k_dot(&data, &activations).expect("scalar");
let simd = fused_q4k_dot_simd(&data, &activations).expect("simd");
let abs_diff = (scalar - simd).abs();
let rel_err = if scalar.abs() > 1e-6 {
abs_diff / scalar.abs()
} else {
abs_diff
};
assert!(
rel_err < 0.01,
"All-chunks parity: scalar={scalar}, simd={simd}, rel_err={rel_err}"
);
}
#[test]
fn test_fused_q4k_q8k_dot_simd_all_chunks_nonzero() {
let mut scales = [0u8; 12];
for j in 0..4 {
scales[j] = 12;
scales[j + 4] = 4;
}
for j in 8..12 {
scales[j] = 0x4C;
}
let mut qs = [0u8; 128];
for i in 0..128 {
qs[i] = ((i * 5 + 17) % 256) as u8;
}
let data = build_q4k_superblock(F16_HALF, F16_QUARTER, &scales, &qs);
let q8k_scales = vec![0.8f32; 1];
let q8k_quants: Vec<i8> = (0..QK_K)
.map(|i| {
let v = (i as i16 * 3) % 256 - 128;
v.clamp(-128, 127) as i8
})
.collect();
let scalar = fused_q4k_q8k_dot(&data, &q8k_scales, &q8k_quants).expect("scalar");
let simd = fused_q4k_q8k_dot_simd(&data, &q8k_scales, &q8k_quants).expect("simd");
let abs_diff = (scalar - simd).abs();
let rel_err = if scalar.abs() > 1e-6 {
abs_diff / scalar.abs()
} else {
abs_diff
};
assert!(
rel_err < 0.05,
"Q8K all-chunks parity: scalar={scalar}, simd={simd}, rel_err={rel_err}"
);
}
#[test]
fn test_fused_q4k_dot_blocks_4_through_7() {
let mut scales = [0u8; 12];
scales[0] = 0xC0; scales[1] = 0x80; scales[2] = 0x40; scales[3] = 0x00; scales[8] = 0x0F; scales[9] = 0x0A; scales[10] = 0x05; scales[11] = 0x01;
let qs = [0x88u8; 128]; let data = build_q4k_superblock(F16_ONE, F16_ZERO, &scales, &qs);
let activations = vec![1.0f32; QK_K];
let result = fused_q4k_dot(&data, &activations).expect("blocks 4-7");
assert!(result.is_finite());
assert!(result.abs() > 0.0, "Blocks 4-7 should contribute: {result}");
}
#[test]
fn test_fused_q4k_dot_block_groups_independence() {
let mut scales_low = [0u8; 12];
for j in 0..4 {
scales_low[j] = 10; }
let mut scales_high = [0u8; 12];
scales_high[0] = 0xC0;
scales_high[1] = 0xC0;
scales_high[2] = 0xC0;
scales_high[3] = 0xC0;
for j in 8..12 {
scales_high[j] = 0x0A; }
let qs = [0x55u8; 128];
let data_low = build_q4k_superblock(F16_ONE, F16_ZERO, &scales_low, &qs);
let data_high = build_q4k_superblock(F16_ONE, F16_ZERO, &scales_high, &qs);
let activations = vec![1.0f32; QK_K];
let result_low = fused_q4k_dot(&data_low, &activations).expect("low blocks");
let result_high = fused_q4k_dot(&data_high, &activations).expect("high blocks");
assert!(result_low.is_finite());
assert!(result_high.is_finite());
assert!(result_low.abs() > 0.0);
assert!(result_high.abs() > 0.0);
}
#[test]
fn test_fused_q4k_q8k_dot_per_byte_contribution() {
let mut scales = [0u8; 12];
for j in 0..4 {
scales[j] = 1;
}
for j in 8..12 {
scales[j] = 0x01;
}
let q8k_scales = vec![1.0f32; 1];
let qs_zero = [0x00u8; 128];
let data_zero = build_q4k_superblock(F16_ONE, F16_ZERO, &scales, &qs_zero);
let q8k_quants = vec![1i8; QK_K];
let baseline = fused_q4k_q8k_dot(&data_zero, &q8k_scales, &q8k_quants).expect("baseline");
let mut qs_byte0 = [0x00u8; 128];
qs_byte0[0] = 0xFF;
let data_byte0 = build_q4k_superblock(F16_ONE, F16_ZERO, &scales, &qs_byte0);
let result_byte0 = fused_q4k_q8k_dot(&data_byte0, &q8k_scales, &q8k_quants).expect("byte0");
assert_ne!(baseline, result_byte0, "Byte 0 should affect result");
let mut qs_byte64 = [0x00u8; 128];
qs_byte64[64] = 0xFF;
let data_byte64 = build_q4k_superblock(F16_ONE, F16_ZERO, &scales, &qs_byte64);
let result_byte64 = fused_q4k_q8k_dot(&data_byte64, &q8k_scales, &q8k_quants).expect("byte64");
assert_ne!(baseline, result_byte64, "Byte 64 should affect result");
}
#[test]
fn test_fused_q4k_q8k_dot_accumulator_paths() {
let mut scales = [0u8; 12];
for j in 0..4 {
scales[j] = 5;
scales[j + 4] = 2;
}
for j in 8..12 {
scales[j] = 0x25;
}
let mut qs = [0u8; 128];
for i in 0..32 {
qs[i] = 0x0F; }
for i in 32..64 {
qs[i] = 0xF0; }
for i in 64..96 {
qs[i] = if i % 2 == 0 { 0x0F } else { 0xF0 };
}
for i in 96..128 {
qs[i] = 0x77; }
let data = build_q4k_superblock(F16_ONE, F16_HALF, &scales, &qs);
let q8k_scales = vec![1.0f32; 1];
let q8k_quants: Vec<i8> = (0..QK_K).map(|i| ((i % 64) as i8) - 32).collect();
let result = fused_q4k_q8k_dot(&data, &q8k_scales, &q8k_quants).expect("accumulator paths");
assert!(result.is_finite(), "Should be finite: {result}");
}