use super::*;
#[test]
fn argmax_matches_two_pass_idiom() {
fn naive(x: &[f32]) -> (usize, f32) {
let m = simd_max_f32(x);
(x.iter().position(|&v| v == m).unwrap_or(0), m)
}
let cases: &[&[f32]] = &[
&[3.0],
&[1.0, 2.0, 3.0, 2.0, 1.0],
&[5.0, 5.0, 5.0], &[-1.0, -2.0, -0.5, -9.0], &[0.0, 1.0, 1.0, 0.5, 1.0], ];
for c in cases {
assert_eq!(simd_argmax_f32(c), naive(c), "mismatch on {c:?}");
}
let mut buf = vec![0.0f32; 4096];
for (i, v) in buf.iter_mut().enumerate() {
*v = ((i * 2654435761) % 997) as f32;
}
buf[1234] = 10_000.0;
assert_eq!(simd_argmax_f32(&buf), (1234, 10_000.0));
let mut state = 0x2545_f491_4f6c_dd1du64;
let mut rng = || {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
state
};
for len in 1..=130usize {
let v: Vec<f32> = (0..len).map(|_| (rng() % 7) as f32).collect();
assert_eq!(simd_argmax_f32(&v), naive(&v), "len={len} v={v:?}");
}
}
#[test]
fn argmax_empty_slice() {
assert_eq!(simd_argmax_f32(&[]), (0, f32::NEG_INFINITY));
}
#[test]
fn simd_level_matches_platform() {
let level = simd_level();
#[cfg(target_arch = "aarch64")]
assert_eq!(level, SimdLevel::Neon);
#[cfg(target_arch = "x86_64")]
assert!(matches!(level, SimdLevel::Avx2 | SimdLevel::Scalar));
#[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
assert_eq!(level, SimdLevel::Scalar);
}
#[test]
fn simd_sigmoid_inplace_matches_fast_sigmoid_within_tolerance() {
let mut rng = fastrand::Rng::with_seed(2026);
for len in 0..=32 {
let mut input: Vec<f32> = (0..len)
.map(|_| (rng.f32() * 80.0) - 40.0) .collect();
let reference: Vec<f32> = input.iter().map(|&x| fast_sigmoid(x)).collect();
simd_sigmoid_inplace(&mut input);
assert_eq!(input.len(), reference.len(), "length changed");
let mut max_diff = 0.0f32;
for (got, want) in input.iter().zip(reference.iter()) {
assert!(*got >= 0.0 && *got <= 1.0, "sigmoid out of [0,1]: {got}");
assert!(
*want >= 0.0 && *want <= 1.0,
"reference out of [0,1]: {want}"
);
max_diff = max_diff.max((got - want).abs());
}
assert!(
max_diff < 5e-6,
"len={len}: max_diff={max_diff:e} exceeds Cephes tolerance"
);
}
}
#[test]
fn simd_sigmoid_inplace_handles_boundaries() {
let mut empty: Vec<f32> = vec![];
simd_sigmoid_inplace(&mut empty);
assert!(empty.is_empty());
let mut extremes = [60.0f32, -60.0, 0.0, 0.0001, -0.0001];
simd_sigmoid_inplace(&mut extremes);
assert!(
(extremes[0] - 1.0).abs() < 1e-6,
"σ(60) ≈ 1, got {}",
extremes[0]
);
assert!(
(extremes[1] - 0.0).abs() < 1e-6,
"σ(-60) ≈ 0, got {}",
extremes[1]
);
assert!(
(extremes[2] - 0.5).abs() < 1e-6,
"σ(0) = 0.5, got {}",
extremes[2]
);
assert!((extremes[3] - 0.5).abs() < 1e-3);
assert!((extremes[4] - 0.5).abs() < 1e-3);
}
#[test]
fn simd_tanh_inplace_matches_fast_tanh_within_fma_tolerance() {
let mut rng = fastrand::Rng::with_seed(2026);
for len in 0..=33usize {
let mut input: Vec<f32> = (0..len)
.map(|_| (rng.f32() * 12.0) - 6.0) .collect();
let reference: Vec<f32> = input.iter().map(|&x| fast_tanh(x)).collect();
simd_tanh_inplace(&mut input);
assert_eq!(input.len(), reference.len(), "length changed");
for (i, (got, want)) in input.iter().zip(reference.iter()).enumerate() {
let abs_tol = 2e-7_f32;
let rel_tol = 2e-7_f32 * want.abs();
let tol = abs_tol.max(rel_tol);
assert!(
(got - want).abs() <= tol,
"len={len} idx={i}: SIMD={got} != scalar={want} (diff={}, tol={tol})",
(got - want).abs()
);
}
}
}
#[test]
fn simd_tanh_inplace_handles_boundaries() {
let mut empty: Vec<f32> = vec![];
simd_tanh_inplace(&mut empty);
assert!(empty.is_empty());
let mut extremes = [10.0f32, -10.0, 5.0, -5.0, 3.5, -3.5];
simd_tanh_inplace(&mut extremes);
assert_eq!(extremes[0], 1.0, "tanh(10) saturates to 1");
assert_eq!(extremes[1], -1.0, "tanh(-10) saturates to -1");
assert_eq!(extremes[2], 1.0, "tanh(5) saturates to 1");
assert_eq!(extremes[3], -1.0, "tanh(-5) saturates to -1");
let mut zero = [0.0f32];
simd_tanh_inplace(&mut zero);
assert_eq!(zero[0], 0.0, "tanh(0) = 0");
let mut rng = fastrand::Rng::with_seed(42);
let mut buf: Vec<f32> = (0..100).map(|_| (rng.f32() * 20.0) - 10.0).collect();
simd_tanh_inplace(&mut buf);
for (i, &v) in buf.iter().enumerate() {
assert!((-1.0..=1.0).contains(&v), "idx={i}: tanh output {v} out of [-1,1]");
}
}
#[test]
fn dot_product_aligned_len_8() {
let a = [1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
let b = [0.5f32, 1.0, 1.5, 2.0, 2.5, 3.0, 3.5, 4.0];
let scalar = scalar_dot_f32(&a, &b, 8);
let simd = simd_dot_f32(&a, &b, 8);
assert!((scalar - simd).abs() < 1e-4, "scalar={scalar}, simd={simd}");
assert!((simd - 102.0).abs() < 1e-4, "simd={simd}");
}
#[test]
fn dot_product_non_aligned_len() {
let a = [1.0f32, 2.0, 3.0, 4.0, 5.0];
let b = [1.0f32, 1.0, 1.0, 1.0, 1.0];
let scalar = scalar_dot_f32(&a, &b, 5);
let simd = simd_dot_f32(&a, &b, 5);
assert!((scalar - simd).abs() < 1e-4, "scalar={scalar}, simd={simd}");
assert!((simd - 15.0).abs() < 1e-4);
}
#[test]
fn dot_product_len_4() {
let a = [1.0f32, 2.0, 3.0, 4.0];
let b = [1.0f32, 0.5, 0.25, 0.125];
let expected = 1.0 + 1.0 + 0.75 + 0.5;
let simd = simd_dot_f32(&a, &b, 4);
assert!((simd - expected).abs() < 1e-4);
}
#[test]
fn dot_product_len_32() {
let a: Vec<f32> = (0..32).map(|i| (i as f32 + 1.0) * 0.1).collect();
let b: Vec<f32> = (0..32).map(|i| (i as f32 + 1.0) * 0.05).collect();
let scalar = scalar_dot_f32(&a, &b, 32);
let simd = simd_dot_f32(&a, &b, 32);
assert!((scalar - simd).abs() < 1e-3, "scalar={scalar}, simd={simd}");
}
#[test]
fn dot_product_zero_length() {
let simd = simd_dot_f32(&[], &[], 0);
assert!((simd - 0.0).abs() < 1e-6);
}
#[test]
fn outer_product_4x4_matches_scalar() {
let m = 4;
let n = 4;
let a = [1.0f32, 2.0, 3.0, 4.0];
let b = [0.5f32, 1.0, 1.5, 2.0];
let mut acc_scalar = vec![0.0f32; m * n];
let mut acc_simd = vec![0.0f32; m * n];
scalar_outer_product_acc(&mut acc_scalar, &a, &b, m, n);
simd_outer_product_acc(&mut acc_simd, &a, &b, m, n);
for i in 0..m * n {
assert!(
(acc_scalar[i] - acc_simd[i]).abs() < 1e-4,
"mismatch at {i}: scalar={}, simd={}",
acc_scalar[i],
acc_simd[i]
);
}
}
#[test]
fn outer_product_8x8_matches_scalar() {
let m = 8;
let n = 8;
let a: Vec<f32> = (0..m).map(|i| (i + 1) as f32 * 0.1).collect();
let b: Vec<f32> = (0..n).map(|j| (j + 1) as f32 * 0.2).collect();
let mut acc_scalar = vec![0.0f32; m * n];
let mut acc_simd = vec![0.0f32; m * n];
scalar_outer_product_acc(&mut acc_scalar, &a, &b, m, n);
simd_outer_product_acc(&mut acc_simd, &a, &b, m, n);
for i in 0..m * n {
assert!(
(acc_scalar[i] - acc_simd[i]).abs() < 1e-4,
"mismatch at {i}: scalar={}, simd={}",
acc_scalar[i],
acc_simd[i]
);
}
}
#[test]
fn outer_product_accumulates() {
let m = 4;
let n = 4;
let a = [1.0f32, 0.0, 0.0, 0.0];
let b = [0.0f32, 0.0, 0.0, 1.0];
let mut acc = vec![0.0f32; m * n];
simd_outer_product_acc(&mut acc, &a, &b, m, n);
assert!((acc[3] - 1.0).abs() < 1e-5);
for (i, &val) in acc.iter().enumerate() {
if i != 3 {
assert!(val.abs() < 1e-6, "acc[{i}] should be 0, got {val}");
}
}
}
#[test]
fn matvec_matches_scalar() {
let rows = 3;
let cols = 4;
let mat = [
1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0f32,
];
let vec = [1.0, 0.0, 1.0, 0.0f32];
let mut acc_scalar = vec![0.0f32; rows];
let mut acc_simd = vec![0.0f32; rows];
for r in 0..rows {
let mut sum = 0.0f32;
for c in 0..cols {
sum += mat[r * cols + c] * vec[c];
}
acc_scalar[r] = sum;
}
simd_matvec(&mut acc_simd, &mat, &vec, rows, cols);
for r in 0..rows {
assert!(
(acc_scalar[r] - acc_simd[r]).abs() < 1e-4,
"mismatch at row {r}: scalar={}, simd={}",
acc_scalar[r],
acc_simd[r]
);
}
}
#[test]
fn matmul_rows_identity() {
let rows = 4;
let cols = 4;
let weight = [
1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 1.0,
];
let input = [1.0, 2.0, 3.0, 4.0f32];
let mut output = vec![0.0f32; rows];
simd_matmul_rows(&mut output, &weight, &input, rows, cols);
assert!((output[0] - 1.0).abs() < 1e-5);
assert!((output[1] - 2.0).abs() < 1e-5);
assert!((output[2] - 3.0).abs() < 1e-5);
assert!((output[3] - 4.0).abs() < 1e-5);
}
#[test]
fn matmul_relu_clamps_negative() {
let rows = 2;
let cols = 2;
let weight = [-1.0, 0.0, 1.0, 1.0];
let input = [1.0, 1.0];
let mut output = vec![0.0f32; rows];
simd_matmul_relu_rows(&mut output, &weight, &input, rows, cols);
assert!((output[0]).abs() < 1e-5, "negative should clamp to 0");
assert!((output[1] - 2.0).abs() < 1e-5);
}
#[test]
fn fma_row_matches_dot() {
let a = [1.0f32, 2.0, 3.0, 4.0];
let b = [0.5f32, 1.0, 1.5, 2.0];
let dot = simd_dot_f32(&a, &b, 4);
let fma = simd_fma_row(&a, &b, 4);
assert!((dot - fma).abs() < 1e-6);
}
#[test]
fn sparse_dot_matches_scalar_dense() {
let weight = [1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
let indices: Vec<usize> = (0..8).collect();
let values = [0.5f32, 1.0, 1.5, 2.0, 2.5, 3.0, 3.5, 4.0];
let sparse = simd_sparse_dot_f32(&weight, 0, &indices, &values, 8);
let dense = simd_dot_f32(&weight, &values, 8);
assert!(
(sparse - dense).abs() < 1e-4,
"sparse={sparse}, dense={dense}"
);
}
#[test]
fn sparse_dot_matches_scalar_sparse() {
let mut weight = vec![0.0f32; 64];
for (i, w) in weight.iter_mut().enumerate() {
*w = (i as f32 + 1.0) * 0.01;
}
let indices: Vec<usize> = vec![0, 3, 7, 12, 15, 20, 25, 31, 38, 45, 50, 56, 63];
let values: Vec<f32> = indices.iter().map(|&i| weight[i] * 2.0).collect();
let simd_result = simd_sparse_dot_f32(&weight, 0, &indices, &values, 13);
let scalar_result = scalar_sparse_dot_f32(&weight, 0, &indices, &values, 13);
assert!(
(simd_result - scalar_result).abs() < 1e-4,
"simd={simd_result}, scalar={scalar_result}"
);
}
#[test]
fn sparse_dot_small_alive_uses_scalar() {
let weight = [1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
let indices = vec![0usize, 3, 7];
let values = [0.5f32, 1.0, 1.5];
let result = simd_sparse_dot_f32(&weight, 0, &indices, &values, 3);
let expected = 1.0 * 0.5 + 4.0 * 1.0 + 8.0 * 1.5;
assert!(
(result - expected).abs() < 1e-4,
"result={result}, expected={expected}"
);
}
#[test]
fn sparse_dot_zero_alive() {
let weight = [1.0f32, 2.0, 3.0, 4.0];
let indices: Vec<usize> = vec![];
let values: Vec<f32> = vec![];
let result = simd_sparse_dot_f32(&weight, 0, &indices, &values, 0);
assert!(result.abs() < 1e-6, "expected 0.0, got {result}");
}
#[test]
fn sparse_dot_with_row_offset() {
let mut weight = [0.0f32; 12]; weight[4] = 1.0;
weight[5] = 2.0;
weight[6] = 3.0;
weight[7] = 4.0;
weight[8] = 5.0;
weight[9] = 6.0;
weight[10] = 7.0;
weight[11] = 8.0;
let weight = weight;
let indices: Vec<usize> = (0..8).collect();
let values = [1.0f32; 8];
let result = simd_sparse_dot_f32(&weight, 4, &indices, &values, 8);
assert!((result - 36.0).abs() < 1e-4, "result={result}");
}
#[test]
fn sparse_dot_alive_5_triggers_simd() {
let weight = [1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
let indices: Vec<usize> = (0..8).collect();
let values = [1.0f32, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0];
let simd_result = simd_sparse_dot_f32(&weight, 0, &indices, &values, 5);
let expected = 1.0 + 2.0 + 3.0 + 4.0 + 5.0;
assert!(
(simd_result - expected).abs() < 1e-4,
"simd={simd_result}, expected={expected}"
);
}
#[test]
fn sparse_matmul_rows_matches_scalar() {
let rows = 4;
let cols = 8;
let weight: Vec<f32> = (0..rows * cols)
.map(|i| {
let r = i / cols;
let c = i % cols;
if r == c { 1.0 } else { 0.1 }
})
.collect();
let indices = vec![1usize, 3, 5];
let values = vec![2.0f32, 3.0, 4.0];
let mut output_scalar = vec![0.0f32; rows];
let mut output_simd = vec![0.0f32; rows];
for (r, out) in output_scalar.iter_mut().enumerate() {
*out = scalar_sparse_dot_f32(&weight, r * cols, &indices, &values, 3);
}
simd_sparse_matmul_rows(&mut output_simd, &weight, &indices, &values, rows, cols, 3);
for (r, (scalar, simd)) in output_scalar.iter().zip(output_simd.iter()).enumerate() {
assert!(
(scalar - simd).abs() < 1e-4,
"row {r}: scalar={scalar}, simd={simd}"
);
}
}
#[test]
fn sparse_matmul_rows_game_config() {
let rows = 32;
let cols = 128;
let weight: Vec<f32> = (0..rows * cols).map(|i| (i % 100) as f32 * 0.01).collect();
let alive = 26;
let indices: Vec<usize> = (0..alive).map(|i| i * (cols / alive)).collect();
let values: Vec<f32> = (0..alive).map(|i| (i as f32 + 1.0) * 0.1).collect();
let mut output_scalar = vec![0.0f32; rows];
let mut output_simd = vec![0.0f32; rows];
for (r, out) in output_scalar.iter_mut().enumerate() {
*out = scalar_sparse_dot_f32(&weight, r * cols, &indices, &values, alive);
}
simd_sparse_matmul_rows(
&mut output_simd,
&weight,
&indices,
&values,
rows,
cols,
alive,
);
for r in 0..rows {
assert!(
(output_scalar[r] - output_simd[r]).abs() < 1e-3,
"row {r}: scalar={}, simd={}",
output_scalar[r],
output_simd[r]
);
}
}
#[test]
fn scale_aligned_len_8() {
let mut x = [2.0f32, 4.0, 6.0, 8.0, 10.0, 12.0, 14.0, 16.0];
simd_scale_inplace(&mut x, 0.5);
let expected = [1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
for i in 0..8 {
assert!((x[i] - expected[i]).abs() < 1e-6, "x[{i}]={}", x[i]);
}
}
#[test]
fn scale_non_aligned_len_13() {
let mut x = [1.0f32; 13];
simd_scale_inplace(&mut x, 3.0);
for (i, &val) in x.iter().enumerate() {
assert!((val - 3.0).abs() < 1e-6, "x[{i}]={val}");
}
}
#[test]
fn scale_empty() {
let mut x: [f32; 0] = [];
simd_scale_inplace(&mut x, 2.0); }
#[test]
fn scale_single_element() {
let mut x = [5.0f32];
simd_scale_inplace(&mut x, 0.2);
assert!((x[0] - 1.0).abs() < 1e-6);
}
#[test]
fn scale_zero() {
let mut x = [1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
simd_scale_inplace(&mut x, 0.0);
for val in &x {
assert!(*val == 0.0, "expected 0.0, got {val}");
}
}
#[test]
fn scale_matches_scalar() {
let mut x_simd: Vec<f32> = (0..97).map(|i| (i as f32 * 0.1).sin()).collect();
let mut x_scalar = x_simd.clone();
let scale = 0.42f32;
simd_scale_inplace(&mut x_simd, scale);
scalar_scale_inplace(&mut x_scalar, scale);
for i in 0..x_simd.len() {
assert!(
(x_simd[i] - x_scalar[i]).abs() < 1e-6,
"x[{i}]: simd={}, scalar={}",
x_simd[i],
x_scalar[i]
);
}
}
#[test]
fn add_scalar_aligned_len_8() {
let mut x = [1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
simd_add_scalar_inplace(&mut x, -10.0);
let expected = [-9.0, -8.0, -7.0, -6.0, -5.0, -4.0, -3.0, -2.0];
for i in 0..8 {
assert!((x[i] - expected[i]).abs() < 1e-6, "x[{i}]={}", x[i]);
}
}
#[test]
fn add_scalar_non_aligned_len_13() {
let mut x = [1.0f32; 13];
simd_add_scalar_inplace(&mut x, 2.0);
for (i, &val) in x.iter().enumerate() {
assert!((val - 3.0).abs() < 1e-6, "x[{i}]={val}");
}
}
#[test]
fn add_scalar_empty() {
let mut x: [f32; 0] = [];
simd_add_scalar_inplace(&mut x, 1.0); }
#[test]
fn add_scalar_matches_scalar_impl() {
let mut x_simd: Vec<f32> = (0..97).map(|i| (i as f32 * 0.1).sin()).collect();
let mut x_scalar = x_simd.clone();
let val = -std::f32::consts::PI;
simd_add_scalar_inplace(&mut x_simd, val);
scalar_add_scalar_inplace(&mut x_scalar, val);
for i in 0..x_simd.len() {
assert!(
(x_simd[i] - x_scalar[i]).abs() < 1e-6,
"x[{i}]: simd={}, scalar={}",
x_simd[i],
x_scalar[i]
);
}
}
#[test]
fn sum_aligned_len_8() {
let x = [1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
let result = simd_sum_f32(&x);
assert!((result - 36.0).abs() < 1e-4, "expected 36.0, got {result}");
}
#[test]
fn sum_non_aligned_len_13() {
let x = [1.0f32; 13];
let result = simd_sum_f32(&x);
assert!((result - 13.0).abs() < 1e-4, "expected 13.0, got {result}");
}
#[test]
fn sum_empty() {
let x: [f32; 0] = [];
let result = simd_sum_f32(&x);
assert!((result - 0.0).abs() < 1e-6, "expected 0.0, got {result}");
}
#[test]
fn sum_single_element() {
let x = [42.0f32];
let result = simd_sum_f32(&x);
assert!((result - 42.0).abs() < 1e-4, "expected 42.0, got {result}");
}
#[test]
fn sum_matches_scalar_impl() {
let x: Vec<f32> = (0..97).map(|i| (i as f32 * 0.1).sin()).collect();
let simd_result = simd_sum_f32(&x);
let scalar_result = scalar_sum_f32(&x);
assert!(
(simd_result - scalar_result).abs() < 1e-4,
"simd={simd_result}, scalar={scalar_result}"
);
}
#[test]
fn add_inplace_aligned_len_8() {
let mut dst = [1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
let src = [0.1f32, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8];
simd_add_inplace(&mut dst, &src);
for (i, val) in dst.iter().enumerate() {
let expected = (1.0 + i as f32) + (i + 1) as f32 * 0.1;
assert!((val - expected).abs() < 1e-6, "mismatch at {i}");
}
}
#[test]
fn add_inplace_non_aligned_len_13() {
let mut dst = [0.0f32; 13];
let src = [1.0f32; 13];
for (i, val) in dst.iter_mut().enumerate() {
*val = i as f32;
}
simd_add_inplace(&mut dst, &src);
for (i, val) in dst.iter().enumerate() {
assert!((val - (i as f32 + 1.0)).abs() < 1e-6, "mismatch at {i}");
}
}
#[test]
fn add_inplace_empty() {
let mut dst: [f32; 0] = [];
let src: [f32; 0] = [];
simd_add_inplace(&mut dst, &src);
}
#[test]
fn add_inplace_single_element() {
let mut dst = [3.0f32];
let src = [7.0f32];
simd_add_inplace(&mut dst, &src);
assert!((dst[0] - 10.0).abs() < 1e-6);
}
#[test]
fn add_inplace_matches_scalar() {
let mut dst_simd = [0.0f32; 37];
let mut dst_scalar = [0.0f32; 37];
for i in 0..37 {
dst_simd[i] = i as f32 * 0.7;
dst_scalar[i] = i as f32 * 0.7;
}
let src: Vec<f32> = (0..37).map(|i| (i as f32 * 0.3).sin()).collect();
simd_add_inplace(&mut dst_simd, &src);
scalar_add_inplace(&mut dst_scalar, &src);
for i in 0..37 {
assert!(
(dst_simd[i] - dst_scalar[i]).abs() < 1e-5,
"mismatch at {i}"
);
}
}
#[test]
fn add_into_aligned_len_8() {
let a = [1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
let b = [8.0f32, 7.0, 6.0, 5.0, 4.0, 3.0, 2.0, 1.0];
let mut dst = [0.0f32; 8];
simd_add_into(&mut dst, &a, &b);
for val in &dst {
assert!((val - 9.0).abs() < 1e-6);
}
}
#[test]
fn add_into_non_aligned_len_13() {
let a: Vec<f32> = (0..13).map(|i| i as f32).collect();
let b = [1.0f32; 13];
let mut dst = [0.0f32; 13];
simd_add_into(&mut dst, &a, &b);
for (i, val) in dst.iter().enumerate() {
assert!((val - (i as f32 + 1.0)).abs() < 1e-6, "mismatch at {i}");
}
}
#[test]
fn add_into_empty() {
let a: [f32; 0] = [];
let b: [f32; 0] = [];
let mut dst: [f32; 0] = [];
simd_add_into(&mut dst, &a, &b);
}
#[test]
fn add_into_matches_scalar() {
let a: Vec<f32> = (0..37).map(|i| (i as f32 * 0.7).sin()).collect();
let b: Vec<f32> = (0..37).map(|i| (i as f32 * 0.3).cos()).collect();
let mut dst_simd = [0.0f32; 37];
let mut dst_scalar = [0.0f32; 37];
simd_add_into(&mut dst_simd, &a, &b);
scalar_add_into(&mut dst_scalar, &a, &b);
for i in 0..37 {
assert!(
(dst_simd[i] - dst_scalar[i]).abs() < 1e-5,
"mismatch at {i}"
);
}
}
#[test]
fn max_aligned_len_8() {
let x = [1.0f32, 5.0, 3.0, 8.0, 2.0, 7.0, 4.0, 6.0];
let max = simd_max_f32(&x);
assert!((max - 8.0).abs() < 1e-6);
}
#[test]
fn max_non_aligned_len_13() {
let x: Vec<f32> = (0..13).map(|i| (i as f32 * 1.7).sin()).collect();
let max = simd_max_f32(&x);
let expected = x.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
assert!((max - expected).abs() < 1e-5);
}
#[test]
fn max_empty() {
let x: [f32; 0] = [];
let max = simd_max_f32(&x);
assert!(max.is_infinite() && max.is_sign_negative());
}
#[test]
fn max_single_element() {
let x = [42.0f32];
let max = simd_max_f32(&x);
assert!((max - 42.0).abs() < 1e-6);
}
#[test]
fn max_negative_values() {
let x = [-5.0f32, -3.0, -8.0, -1.0, -4.0];
let max = simd_max_f32(&x);
assert!((max - (-1.0)).abs() < 1e-6);
}
#[test]
fn max_matches_scalar() {
let x: Vec<f32> = (0..37).map(|i| (i as f32 * 0.97 - 18.0).sin()).collect();
let max_simd = simd_max_f32(&x);
let max_scalar = scalar_max_f32(&x);
assert!((max_simd - max_scalar).abs() < 1e-5);
}
#[test]
fn fused_decay_write_aligned_len_8() {
let mut dst = [1.0f32; 8];
let src = [2.0f32; 8];
let decay = 0.5f32;
let write = 0.5f32;
simd_fused_decay_write(&mut dst, decay, &src, write);
for val in &dst {
assert!((val - 1.5).abs() < 1e-5);
}
}
#[test]
fn fused_decay_write_zero_decay() {
let mut dst = [1.0f32, 2.0, 3.0, 4.0];
let src = [10.0f32, 20.0, 30.0, 40.0];
let decay = 0.0f32;
let write = 1.0f32;
simd_fused_decay_write(&mut dst, decay, &src, write);
for i in 0..4 {
assert!((dst[i] - src[i]).abs() < 1e-5, "mismatch at {i}");
}
}
#[test]
fn fused_decay_write_zero_write() {
let mut dst = [1.0f32, 2.0, 3.0, 4.0];
let src = [10.0f32, 20.0, 30.0, 40.0];
let decay = 1.0f32;
let write = 0.0f32;
simd_fused_decay_write(&mut dst, decay, &src, write);
assert!((dst[0] - 1.0).abs() < 1e-5);
assert!((dst[1] - 2.0).abs() < 1e-5);
assert!((dst[2] - 3.0).abs() < 1e-5);
assert!((dst[3] - 4.0).abs() < 1e-5);
}
#[test]
fn fused_decay_write_empty() {
let mut dst: [f32; 0] = [];
let src: [f32; 0] = [];
simd_fused_decay_write(&mut dst, 0.5, &src, 0.5);
}
#[test]
fn fused_decay_write_matches_scalar() {
let mut dst_simd: Vec<f32> = (0..37).map(|i| i as f32 * 0.7).collect();
let mut dst_scalar: Vec<f32> = (0..37).map(|i| i as f32 * 0.7).collect();
let src: Vec<f32> = (0..37).map(|i| (i as f32 * 0.3).sin()).collect();
let decay = 0.9f32;
let write = 0.1f32;
simd_fused_decay_write(&mut dst_simd, decay, &src, write);
scalar_fused_decay_write(&mut dst_scalar, decay, &src, write);
for i in 0..37 {
assert!(
(dst_simd[i] - dst_scalar[i]).abs() < 1e-4,
"mismatch at {i}: simd={}, scalar={}",
dst_simd[i],
dst_scalar[i]
);
}
}
fn scalar_dot_f16_f32_ref(w: &[half::f16], x: &[f32], len: usize) -> f32 {
let mut sum = 0.0f32;
for i in 0..len {
sum += w[i].to_f32() * x[i];
}
sum
}
#[test]
fn dot_f16_f32_aligned_len_8() {
let w: Vec<half::f16> = (0..8)
.map(|i| half::f16::from_f32(i as f32 * 0.1))
.collect();
let x: Vec<f32> = (0..8).map(|i| i as f32 * 0.2).collect();
let result = simd_dot_f16_f32(&w, &x, 8);
let expected = scalar_dot_f16_f32_ref(&w, &x, 8);
assert!(
(result - expected).abs() < 1e-4,
"f16 dot aligned: got {result}, expected {expected}"
);
}
#[test]
fn dot_f16_f32_non_aligned_len_13() {
let w: Vec<half::f16> = (0..13)
.map(|i| half::f16::from_f32(i as f32 + 1.0))
.collect();
let x: Vec<f32> = (0..13).map(|i| i as f32 * 0.3).collect();
let result = simd_dot_f16_f32(&w, &x, 13);
let expected = scalar_dot_f16_f32_ref(&w, &x, 13);
assert!(
(result - expected).abs() < 1e-3,
"f16 dot non-aligned: got {result}, expected {expected}"
);
}
#[test]
fn dot_f16_f32_len_4() {
let w: Vec<half::f16> = vec![1.0f32, 2.0, 3.0, 4.0]
.into_iter()
.map(half::f16::from_f32)
.collect();
let x: Vec<f32> = vec![0.25, 0.5, 0.75, 1.0];
let result = simd_dot_f16_f32(&w, &x, 4);
let expected = scalar_dot_f16_f32_ref(&w, &x, 4);
assert!(
(result - expected).abs() < 1e-4,
"f16 dot len 4: got {result}, expected {expected}"
);
}
#[test]
fn dot_f16_f32_zero_length() {
let w: Vec<half::f16> = Vec::new();
let x: Vec<f32> = Vec::new();
let result = simd_dot_f16_f32(&w, &x, 0);
assert_eq!(result, 0.0, "f16 dot zero-length should be 0.0");
}
#[test]
fn matmul_f16_f32_identity() {
let w: Vec<half::f16> = vec![1.0f32, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0]
.into_iter()
.map(half::f16::from_f32)
.collect();
let x: Vec<f32> = vec![2.0, 3.0, 4.0];
let mut out = vec![0.0f32; 3];
simd_matmul_f16_f32_rows(&mut out, &w, &x, 3, 3);
assert!(
(out[0] - 2.0).abs() < 1e-4 && (out[1] - 3.0).abs() < 1e-4 && (out[2] - 4.0).abs() < 1e-4,
"f16 identity matmul: got {out:?}"
);
}
#[test]
fn matmul_f16_f32_matches_f32() {
let rows = 4;
let cols = 6;
let weight_f32: Vec<f32> = (0..rows * cols).map(|i| i as f32 * 0.01 - 0.1).collect();
let weight_f16: Vec<half::f16> = weight_f32.iter().map(|&v| half::f16::from_f32(v)).collect();
let input: Vec<f32> = (0..cols).map(|i| i as f32 * 0.05).collect();
let mut out_f32 = vec![0.0f32; rows];
let mut out_f16 = vec![0.0f32; rows];
simd_matmul_rows(&mut out_f32, &weight_f32, &input, rows, cols);
simd_matmul_f16_f32_rows(&mut out_f16, &weight_f16, &input, rows, cols);
for i in 0..rows {
let diff = (out_f32[i] - out_f16[i]).abs();
assert!(
diff < 0.01,
"f16 vs f32 matmul mismatch at row {i}: f32={}, f16={}, diff={diff}",
out_f32[i],
out_f16[i]
);
}
}
#[cfg(feature = "maxsim")]
fn maxsim_naive(queries: &[f32], documents: &[f32], lq: usize, ld: usize, dim: usize) -> f32 {
let mut score = 0.0f32;
for i in 0..lq {
let q_row = &queries[i * dim..(i + 1) * dim];
let mut my_max = f32::NEG_INFINITY;
for j in 0..ld {
let d_row = &documents[j * dim..(j + 1) * dim];
let mut dot = 0.0f32;
for d in 0..dim {
dot += q_row[d] * d_row[d];
}
my_max = my_max.max(dot);
}
score += my_max;
}
score
}
#[cfg(feature = "maxsim")]
mod maxsim_tests {
use super::*;
#[test]
fn maxsim_matches_naive() {
let lq = 8;
let ld = 16;
let dim = 32;
let mut queries = vec![0.0f32; lq * dim];
let mut documents = vec![0.0f32; ld * dim];
for q in queries.iter_mut() {
*q = fastrand::f32() * 2.0 - 1.0;
}
for d in documents.iter_mut() {
*d = fastrand::f32() * 2.0 - 1.0;
}
let naive = maxsim_naive(&queries, &documents, lq, ld, dim);
let fused = maxsim_score(&queries, &documents, lq, ld, dim);
assert!((naive - fused).abs() < 1e-4, "naive={naive}, fused={fused}");
}
#[test]
fn maxsim_single_query_token() {
let dim = 16;
let queries = (0..dim).map(|i| i as f32).collect::<Vec<f32>>();
let documents = (0..3 * dim)
.map(|i| (i as f32 * 0.1).sin())
.collect::<Vec<f32>>();
let result = maxsim_score(&queries, &documents, 1, 3, dim);
let mut expected = f32::NEG_INFINITY;
for j in 0..3 {
let d_row = &documents[j * dim..(j + 1) * dim];
let dot = simd_dot_f32(&queries, d_row, dim);
expected = expected.max(dot);
}
assert!(
(result - expected).abs() < 1e-5,
"result={result}, expected={expected}"
);
}
#[test]
fn maxsim_single_doc_token() {
let dim = 16;
let lq = 4;
let queries = (0..lq * dim)
.map(|i| (i as f32 * 0.2).cos())
.collect::<Vec<f32>>();
let documents = (0..dim).map(|i| i as f32 * 0.5).collect::<Vec<f32>>();
let result = maxsim_score(&queries, &documents, lq, 1, dim);
let mut expected = 0.0f32;
for i in 0..lq {
let q_row = &queries[i * dim..(i + 1) * dim];
expected += simd_dot_f32(q_row, &documents, dim);
}
assert!(
(result - expected).abs() < 1e-4,
"result={result}, expected={expected}"
);
}
#[test]
fn maxsim_symmetry_breaking() {
let dim = 8;
let lq = 4;
let ld = 4;
let queries = (0..lq * dim).map(|i| i as f32).collect::<Vec<f32>>();
let documents = (0..ld * dim)
.map(|i| (i as f32 * 0.3).sin())
.collect::<Vec<f32>>();
let maxsim = maxsim_score(&queries, &documents, lq, ld, dim);
let mut diagonal = 0.0f32;
for i in 0..lq.min(ld) {
let q_row = &queries[i * dim..(i + 1) * dim];
let d_row = &documents[i * dim..(i + 1) * dim];
diagonal += simd_dot_f32(q_row, d_row, dim);
}
assert!(
(maxsim - diagonal).abs() > 1e-3,
"maxsim={maxsim} should differ from diagonal={diagonal}"
);
}
#[test]
fn maxsim_empty_doc() {
let dim = 16;
let queries = vec![1.0f32; dim];
let documents: Vec<f32> = vec![];
let result = maxsim_score(&queries, &documents, 1, 0, dim);
assert_eq!(result, 0.0, "empty doc should return 0.0");
}
#[test]
fn maxsim_large_dim_aligned() {
let dim = 128;
let lq = 4;
let ld = 8;
let queries: Vec<f32> = (0..lq * dim).map(|i| (i as f32 * 0.01).sin()).collect();
let documents: Vec<f32> = (0..ld * dim).map(|i| (i as f32 * 0.01).cos()).collect();
let naive = maxsim_naive(&queries, &documents, lq, ld, dim);
let fused = maxsim_score(&queries, &documents, lq, ld, dim);
assert!((naive - fused).abs() < 1e-3, "naive={naive}, fused={fused}");
}
#[test]
fn maxsim_packed_matches_sequential() {
let dim = 16;
let q1: Vec<f32> = (0..2 * dim).map(|i| i as f32).collect();
let q2: Vec<f32> = (0..3 * dim).map(|i| (i as f32 * 0.5).sin()).collect();
let d1: Vec<f32> = (0..4 * dim).map(|i| (i as f32 * 0.3).cos()).collect();
let d2: Vec<f32> = (0..2 * dim).map(|i| i as f32 * 0.1).collect();
let d3: Vec<f32> = (0..5 * dim).map(|i| (i as f32 * 0.7).sin()).collect();
let queries: Vec<f32> = [q1.clone(), q2.clone()].concat();
let documents: Vec<f32> = [d1.clone(), d2.clone(), d3.clone()].concat();
let query_offsets = [0, q1.len(), q1.len() + q2.len()];
let doc_offsets = [
0,
d1.len(),
d1.len() + d2.len(),
d1.len() + d2.len() + d3.len(),
];
let pair_q_ids = [0usize, 0, 1];
let pair_d_ids = [0usize, 2, 1];
let mut packed = vec![0.0f32; pair_q_ids.len()];
maxsim_score_packed(
&queries,
&query_offsets,
&documents,
&doc_offsets,
&pair_q_ids,
&pair_d_ids,
dim,
&mut packed,
);
let s0 = maxsim_score(&q1, &d1, 2, 4, dim);
let s1 = maxsim_score(&q1, &d3, 2, 5, dim);
let s2 = maxsim_score(&q2, &d2, 3, 2, dim);
assert!(
(packed[0] - s0).abs() < 1e-4,
"pair 0: packed={}, sequential={}",
packed[0],
s0
);
assert!(
(packed[1] - s1).abs() < 1e-4,
"pair 1: packed={}, sequential={}",
packed[1],
s1
);
assert!(
(packed[2] - s2).abs() < 1e-4,
"pair 2: packed={}, sequential={}",
packed[2],
s2
);
}
}
#[cfg(feature = "sigmoid_margin")]
mod sigmoid_margin_tests {
use super::*;
#[test]
fn proof1_loss_matches_manual() {
let n_rows = 2;
let n_cols = 3;
let scores: Vec<f32> = vec![
0.8, 0.2, -0.5, -0.3, 0.9, 0.1, ];
let adjacency: Vec<f32> = vec![
1.0, 0.0, 0.0, 0.0, 1.0, 0.0, ];
let loss = sigmoid_margin_loss(&scores, &adjacency, 1.0, 0.0, n_rows, n_cols);
let sp = |x: f32| -> f32 { (1.0f32 + x.exp()).ln() };
let expected = (sp(-0.8) + sp(0.2) + sp(-0.5) + sp(-0.3) + sp(-0.9) + sp(0.1)) / 6.0;
assert!(
(loss - expected).abs() < 1e-4,
"loss={loss}, expected={expected}"
);
}
#[test]
fn proof1_loss_with_bias_and_temperature() {
let scores = vec![1.0, 0.0];
let adjacency = vec![1.0, 0.0];
let loss = sigmoid_margin_loss(&scores, &adjacency, 2.0, 0.5, 1, 2);
let sp_neg1 = (1.0f32 + (-1.0f32).exp()).ln(); let expected = sp_neg1; assert!(
(loss - expected).abs() < 1e-4,
"loss={loss}, expected={expected}"
);
}
#[test]
fn proof1_loss_perfect_separation() {
let scores = vec![100.0, -100.0];
let adjacency = vec![1.0, 0.0];
let loss = sigmoid_margin_loss(&scores, &adjacency, 1.0, 0.0, 1, 2);
assert!(
loss < 1e-10,
"loss={loss} should be near 0 for perfect separation"
);
}
#[test]
fn proof2_margin_positive_for_separated_embeddings() {
let dim = 8;
let n_queries = 3;
let n_docs = 6;
let k = 2;
let mut queries = vec![0.0f32; n_queries * dim];
let mut documents = vec![0.0f32; n_docs * dim];
let mut neighborhoods = Vec::with_capacity(n_queries * k);
for i in 0..n_queries {
queries[i * dim + i] = 1.0;
documents[(2 * i) * dim + i] = 0.9;
documents[(2 * i + 1) * dim + i] = 0.8;
neighborhoods.push(2 * i);
neighborhoods.push(2 * i + 1);
}
let (pos_min, neg_max, margin) = compute_retrieval_margin(
&queries,
&documents,
&neighborhoods,
dim,
n_queries,
n_docs,
k,
);
assert!(
(pos_min - 0.8).abs() < 1e-5,
"pos_min={pos_min}, expected 0.8"
);
assert!(neg_max.abs() < 1e-5, "neg_max={neg_max}, expected 0.0");
assert!((margin - 0.4).abs() < 1e-5, "margin={margin}, expected 0.4");
assert!(margin > 0.0, "margin should be positive");
}
#[test]
fn proof2_margin_negative_for_mixed_embeddings() {
let dim = 4;
let n_queries = 1;
let n_docs = 3;
let k = 1;
let queries = vec![1.0, 0.0, 0.0, 0.0]; let d0 = vec![0.1, 0.0, 0.0, 0.0];
let d1 = vec![0.9, 0.0, 0.0, 0.0];
let d2 = vec![0.0, 1.0, 0.0, 0.0];
let documents: Vec<f32> = [d0, d1, d2].concat();
let neighborhoods = vec![0];
let (pos_min, neg_max, margin) = compute_retrieval_margin(
&queries,
&documents,
&neighborhoods,
dim,
n_queries,
n_docs,
k,
);
assert!((pos_min - 0.1).abs() < 1e-5, "pos_min={pos_min}");
assert!((neg_max - 0.9).abs() < 1e-5, "neg_max={neg_max}");
assert!(margin < 0.0, "margin should be negative: {margin}");
}
#[test]
fn proof3_bound_scales_as_k_log_n() {
let b1 = dim_sufficiency_bound(2, 100);
assert!(b1 <= 20, "k=2, n=100: bound={b1}, should be ≤ 20");
assert!(b1 >= 10, "k=2, n=100: bound={b1}, should be ≥ 10");
let b2 = dim_sufficiency_bound(4, 1000);
assert!(b2 <= 60, "k=4, n=1000: bound={b2}, should be ≤ 60");
assert!(b2 >= 30, "k=4, n=1000: bound={b2}, should be ≥ 30");
}
#[test]
fn proof3_bound_edge_cases() {
assert_eq!(dim_sufficiency_bound(0, 100), 1, "k=0 → trivial");
assert_eq!(dim_sufficiency_bound(2, 1), 1, "n=1 → trivial");
assert_eq!(dim_sufficiency_bound(2, 2), 3, "n=2 → minimal");
}
#[test]
fn proof3_bound_monotonic() {
let b1 = dim_sufficiency_bound(2, 50);
let b2 = dim_sufficiency_bound(2, 100);
let b3 = dim_sufficiency_bound(2, 200);
assert!(b1 < b2, "bound should increase with n: {b1} < {b2}");
assert!(b2 < b3, "bound should increase with n: {b2} < {b3}");
let bk1 = dim_sufficiency_bound(2, 100);
let bk2 = dim_sufficiency_bound(4, 100);
assert!(bk1 < bk2, "bound should increase with k: {bk1} < {bk2}");
}
#[test]
fn proof4_loss_gradient_pushes_to_positive_margin() {
let dim = 8;
let n = 4; let k = 2; let n_queries = 2;
let neighborhoods: Vec<usize> = vec![0, 1, 2, 3];
let mut queries = vec![0.0f32; n_queries * dim];
let mut documents = vec![0.0f32; n * dim];
queries[0] = 0.3;
queries[dim + 1] = 0.3;
documents[0] = 0.2;
documents[dim] = 0.15;
documents[2 * dim + 1] = 0.2;
documents[3 * dim + 1] = 0.15;
documents[1] = 0.02;
documents[2 * dim] = 0.02;
let adjacency: Vec<f32> = vec![1.0, 1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 1.0];
let (_, _, initial_margin) =
compute_retrieval_margin(&queries, &documents, &neighborhoods, dim, n_queries, n, k);
let t = 10.0f32;
let lr = 0.1;
let mut q = queries.clone();
let mut d = documents.clone();
for _step in 0..100 {
let mut scores = vec![0.0f32; n_queries * n];
for i in 0..n_queries {
for j in 0..n {
scores[i * n + j] =
simd_dot_f32(&q[i * dim..(i + 1) * dim], &d[j * dim..(j + 1) * dim], dim);
}
}
let mut score_grads = vec![0.0f32; n_queries * n];
for i in 0..n_queries {
for j in 0..n {
let idx = i * n + j;
let sign = if adjacency[idx] > 0.5 {
-1.0f32
} else {
1.0f32
};
let x = t * (scores[idx]) * sign;
let sigmoid_x = 1.0 / (1.0 + (-x).exp());
score_grads[idx] = t * sign * sigmoid_x;
}
}
let mut q_grads = vec![0.0f32; n_queries * dim];
for i in 0..n_queries {
for j in 0..n {
let g = score_grads[i * n + j];
for dd in 0..dim {
q_grads[i * dim + dd] += g * d[j * dim + dd];
}
}
}
let mut d_grads = vec![0.0f32; n * dim];
for i in 0..n_queries {
for j in 0..n {
let g = score_grads[i * n + j];
for dd in 0..dim {
d_grads[j * dim + dd] += g * q[i * dim + dd];
}
}
}
for idx in 0..q.len() {
q[idx] -= lr * q_grads[idx];
}
for idx in 0..d.len() {
d[idx] -= lr * d_grads[idx];
}
}
let (_, _, final_margin) =
compute_retrieval_margin(&q, &d, &neighborhoods, dim, n_queries, n, k);
assert!(
final_margin > 0.0,
"final_margin={final_margin} should be > 0 after training"
);
assert!(
final_margin > initial_margin,
"margin should improve: initial={initial_margin}, final={final_margin}"
);
}
#[test]
#[cfg(feature = "maxsim")]
fn proof5_margin_correlates_with_maxsim() {
let dim = 16;
let n_docs = 4;
let lq = 2;
let ld = n_docs;
let k = 1;
let mut queries = vec![0.0f32; 2 * lq * dim]; let mut documents = vec![0.0f32; n_docs * dim];
documents[0] = 1.0;
documents[dim + 1] = 0.1;
documents[2 * dim + 2] = 0.1;
documents[3 * dim + 3] = 0.1;
queries[0] = 1.0;
queries[dim] = 0.9;
let neighborhoods = vec![0];
let (pos_min, neg_max, margin) = compute_retrieval_margin(
&queries[..lq * dim],
&documents,
&neighborhoods,
dim,
1,
n_docs,
k,
);
let ms = maxsim_score(&queries[..lq * dim], &documents, lq, ld, dim);
assert!(margin > 0.0, "margin={margin} should be positive");
assert!(
ms > 0.0,
"maxsim={ms} should be positive for high-margin setup"
);
assert!(
pos_min > neg_max,
"pos_min={pos_min} should exceed neg_max={neg_max}"
);
}
#[test]
#[cfg(feature = "maxsim")]
fn proof6_no_maxsim_regression() {
let dim = 16;
let lq = 4;
let ld = 8;
let queries: Vec<f32> = (0..lq * dim).map(|i| (i as f32 * 0.01).sin()).collect();
let documents: Vec<f32> = (0..ld * dim).map(|i| (i as f32 * 0.01).cos()).collect();
let mut expected = 0.0f32;
for i in 0..lq {
let q_row = &queries[i * dim..(i + 1) * dim];
let mut my_max = f32::NEG_INFINITY;
for j in 0..ld {
let d_row = &documents[j * dim..(j + 1) * dim];
let mut dot = 0.0f32;
for d in 0..dim {
dot += q_row[d] * d_row[d];
}
my_max = my_max.max(dot);
}
expected += my_max;
}
let result = maxsim_score(&queries, &documents, lq, ld, dim);
assert!(
(result - expected).abs() < 1e-3,
"maxsim={result}, expected={expected}"
);
}
#[test]
fn proof7_feature_gate_functions_exist() {
let _loss = sigmoid_margin_loss(&[0.5, -0.5], &[1.0, 0.0], 1.0, 0.0, 1, 2);
let (pm, _nm, m) = compute_retrieval_margin(
&[1.0, 0.0, 0.0, 1.0], &[1.0, 0.0, 0.0, 1.0], &[0, 1], 2,
2,
2,
1,
);
assert!(pm >= 0.0);
assert!(m >= 0.0);
let bound = dim_sufficiency_bound(2, 100);
assert!(bound > 0);
assert!(bound <= 20);
}
}
mod gram_tests {
use super::*;
#[test]
fn test_gram_identity() {
let seq_len = 3;
let d_h = 3;
let x: Vec<f32> = vec![
1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0, ];
let mut gram = vec![0.0f32; seq_len * seq_len];
simd_gram_f32(&x, seq_len, d_h, &mut gram);
let expected = [1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0];
for (i, (&g, &e)) in gram.iter().zip(expected.iter()).enumerate() {
assert!((g - e).abs() < 1e-5, "gram[{i}]={g}, expected={e}");
}
}
#[test]
fn test_gram_ones() {
let seq_len = 4;
let d_h = 8;
let x = vec![1.0f32; seq_len * d_h];
let mut gram = vec![0.0f32; seq_len * seq_len];
simd_gram_f32(&x, seq_len, d_h, &mut gram);
for (i, &g) in gram.iter().enumerate() {
assert!(
(g - d_h as f32).abs() < 1e-4,
"gram[{i}]={g}, expected={}",
d_h
);
}
}
#[test]
fn test_gram_symmetric() {
let seq_len = 5;
let d_h = 8;
let x: Vec<f32> = (0..seq_len * d_h).map(|i| (i as f32 * 0.1).sin()).collect();
let mut gram = vec![0.0f32; seq_len * seq_len];
simd_gram_f32(&x, seq_len, d_h, &mut gram);
for i in 0..seq_len {
for j in 0..seq_len {
let g_ij = gram[i * seq_len + j];
let g_ji = gram[j * seq_len + i];
assert!(
(g_ij - g_ji).abs() < 1e-5,
"G[{i}][{j}]={g_ij} != G[{j}][{i}]={g_ji}"
);
}
}
}
#[test]
fn test_gram_upper_triangle_mirror() {
let seq_len = 3;
let d_h = 4;
let x: Vec<f32> = vec![
1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0, ];
let mut gram = vec![0.0f32; seq_len * seq_len];
simd_gram_f32(&x, seq_len, d_h, &mut gram);
assert!((gram[1] - 70.0).abs() < 1e-4, "G[0][1]={}", gram[1]);
assert!((gram[3] - 70.0).abs() < 1e-4, "G[1][0]={}", gram[3]);
assert!((gram[2] - 110.0).abs() < 1e-4, "G[0][2]={}", gram[2]);
assert!((gram[6] - 110.0).abs() < 1e-4, "G[2][0]={}", gram[6]);
assert!((gram[5] - 278.0).abs() < 1e-4, "G[1][2]={}", gram[5]);
assert!((gram[7] - 278.0).abs() < 1e-4, "G[2][1]={}", gram[7]);
}
#[test]
fn test_gram_2x3() {
let seq_len = 2;
let d_h = 3;
let x: Vec<f32> = vec![1.0, 0.0, 2.0, 3.0, 1.0, 0.0];
let mut gram = vec![0.0f32; seq_len * seq_len];
simd_gram_f32(&x, seq_len, d_h, &mut gram);
assert!((gram[0] - 5.0).abs() < 1e-5, "G[0][0]={}", gram[0]);
assert!((gram[1] - 3.0).abs() < 1e-5, "G[0][1]={}", gram[1]);
assert!((gram[2] - 3.0).abs() < 1e-5, "G[1][0]={}", gram[2]);
assert!((gram[3] - 10.0).abs() < 1e-5, "G[1][1]={}", gram[3]);
}
#[test]
fn test_gram_matches_outer_product() {
let seq_len = 4;
let d_h = 8;
let x: Vec<f32> = (0..seq_len * d_h)
.map(|i| (i as f32 * 0.17).sin() * 0.5)
.collect();
let mut gram = vec![0.0f32; seq_len * seq_len];
simd_gram_f32(&x, seq_len, d_h, &mut gram);
let mut reference = vec![0.0f32; seq_len * seq_len];
for i in 0..seq_len {
for j in 0..seq_len {
let mut sum = 0.0f32;
for k in 0..d_h {
sum += x[i * d_h + k] * x[j * d_h + k];
}
reference[i * seq_len + j] = sum;
}
}
for i in 0..seq_len {
for j in 0..seq_len {
let idx = i * seq_len + j;
assert!(
(gram[idx] - reference[idx]).abs() < 1e-4,
"G[{i}][{j}]: simd={}, reference={}",
gram[idx],
reference[idx]
);
}
}
}
}
#[test]
fn sum_abs_mixed_values() {
let data: Vec<f32> = vec![1.0, -2.0, 3.0, -4.0, 5.0, -6.0, 7.0, -8.0];
let expected: f32 = data.iter().map(|v| v.abs()).sum();
let result = crate::simd::simd_sum_abs_f32(&data);
assert!(
(result - expected).abs() < 1e-6,
"got {result}, expected {expected}"
);
}
#[test]
fn sum_abs_non_aligned_len() {
let data: Vec<f32> = vec![1.0, -2.0, 3.0, -4.0, 5.0];
let expected: f32 = data.iter().map(|v| v.abs()).sum();
let result = crate::simd::simd_sum_abs_f32(&data);
assert!(
(result - expected).abs() < 1e-6,
"got {result}, expected {expected}"
);
}
#[test]
fn sum_abs_empty() {
let data: Vec<f32> = vec![];
let result = crate::simd::simd_sum_abs_f32(&data);
assert_eq!(result, 0.0);
}
#[test]
fn sum_abs_single_element() {
assert_eq!(crate::simd::simd_sum_abs_f32(&[-42.0]), 42.0);
assert_eq!(crate::simd::simd_sum_abs_f32(&[42.0]), 42.0);
assert_eq!(crate::simd::simd_sum_abs_f32(&[0.0]), 0.0);
}
#[test]
fn l_inf_distance_matches_scalar_across_lengths() {
let mut state: u64 = 0x1234_5678_9ABC_DEF0;
let next_f32 = |s: &mut u64| -> f32 {
*s ^= *s << 13;
*s ^= *s >> 7;
*s ^= *s << 17;
(((*s & 0xFFFFFF) as f32) / ((0x1000000) as f32) - 0.5) * 8.0 };
for &len in &[
1usize, 2, 3, 4, 5, 7, 8, 9, 15, 16, 17, 31, 32, 33, 63, 64, 65, 127, 128, 129, 255, 256,
257,
] {
let a: Vec<f32> = (0..len).map(|_| next_f32(&mut state)).collect();
let b: Vec<f32> = (0..len).map(|_| next_f32(&mut state)).collect();
let simd = crate::simd::simd_l_inf_distance_f32(&a, &b, len);
let reference = scalar_l_inf_distance_f32(&a, &b, len);
assert_eq!(
simd.to_bits(),
reference.to_bits(),
"len={len}: SIMD {simd:?} != scalar {reference:?} (a={a:?}, b={b:?})"
);
}
}
#[test]
fn l_inf_distance_known_values() {
assert_eq!(
crate::simd::simd_l_inf_distance_f32(&[1.0, 2.0, 3.0], &[0.0, 0.0, 0.0], 3),
3.0
);
assert_eq!(
crate::simd::simd_l_inf_distance_f32(&[-5.0, 1.0, 1.0], &[1.0, 1.0, 1.0], 3),
6.0
);
assert_eq!(
crate::simd::simd_l_inf_distance_f32(&[0.0, 10.0, 0.0], &[0.0, 0.0, 0.0], 3),
10.0
);
assert_eq!(
crate::simd::simd_l_inf_distance_f32(&[1.5, -2.25, 3.0], &[1.5, -2.25, 3.0], 3),
0.0
);
}
#[test]
fn l_inf_distance_empty_and_single() {
assert_eq!(
crate::simd::simd_l_inf_distance_f32(&[] as &[f32], &[] as &[f32], 0),
0.0
);
assert_eq!(crate::simd::simd_l_inf_distance_f32(&[7.5], &[2.5], 1), 5.0);
assert_eq!(
crate::simd::simd_l_inf_distance_f32(&[-7.5], &[2.5], 1),
10.0
);
}
#[test]
fn l_inf_distance_matches_reference_at_k8() {
let a = [0.125_f32, 0.25, 0.375, 0.5, 0.625, 0.75, 0.875, 1.0];
let b = [1.0_f32, 0.875, 0.75, 0.625, 0.5, 0.375, 0.25, 0.125];
let simd = crate::simd::simd_l_inf_distance_f32(&a, &b, 8);
let reference = scalar_l_inf_distance_f32(&a, &b, 8);
assert_eq!(simd.to_bits(), reference.to_bits());
assert_eq!(simd, 0.875);
}
#[test]
fn simd_sum_sq_quartic_matches_scalar() {
let mut state: u64 = 0x9E3779B97F4A7C15;
let next_f32 = |s: &mut u64| -> f32 {
*s ^= *s << 13;
*s ^= *s >> 7;
*s ^= *s << 17;
(((*s & 0xFFFFFF) as f32) / ((0x1000000) as f32) - 0.5) * 4.0 };
for &len in &[
1usize, 2, 3, 4, 5, 7, 8, 12, 15, 16, 17, 31, 32, 33, 63, 64, 65, 127, 256, 257, 1023, 1024,
] {
let data: Vec<f32> = (0..len).map(|_| next_f32(&mut state)).collect();
let (sim_sq, sim_qu) = crate::simd::simd_sum_sq_quartic(&data);
let (ref_sq, ref_qu) = scalar_sum_sq_quartic(&data);
let tol_sq = (ref_sq.abs() * 1e-5).max(1e-6);
let tol_qu = (ref_qu.abs() * 1e-5).max(1e-6);
assert!(
(sim_sq - ref_sq).abs() <= tol_sq,
"len={len}: sum_sq simd={sim_sq} scalar={ref_sq} (tol {tol_sq})"
);
assert!(
(sim_qu - ref_qu).abs() <= tol_qu,
"len={len}: sum_quartic simd={sim_qu} scalar={ref_qu} (tol {tol_qu})"
);
}
}
#[test]
fn simd_sum_sq_quartic_zero_input() {
assert_eq!(crate::simd::simd_sum_sq_quartic(&[]), (0.0, 0.0));
assert_eq!(crate::simd::simd_sum_sq_quartic(&[0.0]), (0.0, 0.0));
assert_eq!(crate::simd::simd_sum_sq_quartic(&[0.0; 4]), (0.0, 0.0));
assert_eq!(crate::simd::simd_sum_sq_quartic(&[0.0; 17]), (0.0, 0.0));
assert_eq!(crate::simd::simd_sum_sq_quartic(&[0.0; 64]), (0.0, 0.0));
}
#[test]
fn simd_sum_sq_quartic_short_input() {
assert_eq!(crate::simd::simd_sum_sq_quartic(&[3.0]), (9.0, 81.0));
assert_eq!(crate::simd::simd_sum_sq_quartic(&[2.0, -1.0]), (5.0, 17.0));
assert_eq!(
crate::simd::simd_sum_sq_quartic(&[1.0, 2.0, -2.0]),
(9.0, 33.0)
);
}
#[test]
fn test_entropy_uniform() {
let probs: Vec<f32> = vec![0.25, 0.25, 0.25, 0.25];
let logprobs: Vec<f32> = probs.iter().map(|&p| p.ln()).collect();
let h = entropy_f32(&logprobs);
let expected = 4.0f32.ln(); assert!(
(h - expected).abs() < 0.01,
"uniform entropy should be ln(4)≈1.386, got {h}"
);
}
#[test]
fn test_entropy_peaked() {
let probs: Vec<f32> = vec![0.99, 0.003, 0.004, 0.003];
let logprobs: Vec<f32> = probs.iter().map(|&p| p.ln()).collect();
let h = entropy_f32(&logprobs);
assert!(
h < 0.1,
"peaked distribution should have near-zero entropy, got {h}"
);
}
#[test]
fn test_entropy_empty() {
assert_eq!(entropy_f32(&[]), 0.0);
}
#[test]
fn test_entropy_with_neg_inf_logprobs() {
let logprobs: Vec<f32> = vec![
(0.5f32).ln(),
(0.5f32).ln(),
f32::NEG_INFINITY,
f32::NEG_INFINITY,
];
let h = entropy_f32(&logprobs);
let expected = 2.0f32.ln();
assert!(
h.is_finite(),
"entropy must be finite when some logp = -∞, got {h}"
);
assert!(
(h - expected).abs() < 0.01,
"two-token uniform entropy should be ln(2)≈0.693, got {h}"
);
}
#[test]
fn test_coincidence_full_match() {
let top_k = vec![0, 1, 2, 3];
let parent = vec![0, 1, 2, 3];
let score = coincidence_score(&top_k, &parent, 4);
assert!(
(score - 1.0).abs() < 1e-6,
"full match should give 1.0, got {score}"
);
}
#[test]
fn test_coincidence_no_match() {
let top_k = vec![10, 11, 12, 13];
let parent = vec![0, 1, 2, 3];
let score = coincidence_score(&top_k, &parent, 4);
assert!(score.abs() < 1e-6, "no match should give 0.0, got {score}");
}
#[test]
fn test_coincidence_partial_match() {
let top_k = vec![0, 5, 2, 9];
let parent = vec![0, 1, 2, 3];
let score = coincidence_score(&top_k, &parent, 4);
assert!(
(score - 0.5).abs() < 1e-6,
"2 of 4 match should give 0.5, got {score}"
);
}
#[test]
fn test_coincidence_empty_slices() {
assert_eq!(coincidence_score(&[], &[1, 2], 4), 0.0);
assert_eq!(coincidence_score(&[1], &[], 4), 0.0);
assert_eq!(coincidence_score(&[1], &[2], 0), 0.0);
}
#[test]
fn exp_sum_matches_separate_exp_plus_sum() {
let cases: &[&[f32]] = &[
&[0.0],
&[1.0],
&[0.0, 1.0, 2.0, 3.0],
&[-5.0, -1.0, 0.0, 1.0, 5.0, 10.0],
&(0..32).map(|i| (i as f32 - 16.0) * 0.1).collect::<Vec<_>>(),
&(0..17).map(|i| (i as f32 - 8.0) * 0.1).collect::<Vec<_>>(),
&(0..33).map(|i| (i as f32 - 16.0) * 0.1).collect::<Vec<_>>(),
&(0..100)
.map(|i| (i as f32 - 50.0) * 0.05)
.collect::<Vec<_>>(),
];
for case in cases {
let mut fused = case.to_vec();
let mut sep = case.to_vec();
let fused_sum = simd_exp_sum_inplace(&mut fused);
simd_exp_inplace(&mut sep);
let sep_sum = simd_sum_f32(&sep);
for (i, (a, b)) in fused.iter().zip(sep.iter()).enumerate() {
assert!(
(a - b).abs() < 1e-6,
"exp mismatch at {i}: fused={a}, separate={b}, input={}",
case[i]
);
}
let rel_err = (fused_sum - sep_sum).abs() / sep_sum.max(1e-30);
assert!(
rel_err < 1e-5,
"sum mismatch: fused={fused_sum}, separate={sep_sum}, rel_err={rel_err}"
);
}
}
#[test]
fn exp_sum_empty() {
let mut x: Vec<f32> = vec![];
assert_eq!(simd_exp_sum_inplace(&mut x), 0.0);
}
#[test]
fn exp_sum_known_value() {
let mut x = vec![0.0f32, 1.0, 2.0];
let sum = simd_exp_sum_inplace(&mut x);
let expected = 1.0 + std::f32::consts::E + std::f32::consts::E.powi(2);
assert!(
(sum - expected).abs() < 1e-4,
"got {sum}, expected {expected}"
);
assert!((x[0] - 1.0).abs() < 1e-6);
assert!((x[1] - std::f32::consts::E).abs() < 1e-4);
assert!((x[2] - std::f32::consts::E.powi(2)).abs() < 1e-4);
}
#[test]
fn simd_exp_matches_f32_exp_truth_referenced() {
let inputs: Vec<f32> = (-150..=150).map(|i| i as f32 * 0.1).collect();
let mut x = inputs.clone();
simd_exp_inplace(&mut x);
let mut worst_rel: f32 = 0.0;
let mut worst_at: f32 = 0.0;
for (i, &xi) in inputs.iter().enumerate() {
let expected = xi.exp();
let got = x[i];
let denom = expected.abs().max(1e-30);
let rel_err = (got - expected).abs() / denom;
if rel_err > worst_rel {
worst_rel = rel_err;
worst_at = xi;
}
assert!(
rel_err < 5e-4,
"exp({xi}) = {got} vs true {expected}, rel_err = {rel_err:.3e} (worst so far: {worst_rel:.3e} at {worst_at})"
);
}
eprintln!("simd_exp truth-referenced worst rel_err = {worst_rel:.3e} at x = {worst_at}");
}
#[test]
fn simd_exp_sum_matches_f32_exp_truth_referenced() {
for &len in &[1usize, 3, 4, 8, 12, 16, 17, 31, 32, 33, 100] {
let inputs: Vec<f32> = (0..len)
.map(|i| (i as f32 - (len as f32) * 0.5) * 0.3)
.collect();
let mut x = inputs.clone();
let got_sum = simd_exp_sum_inplace(&mut x);
let mut expected_sum = 0.0f32;
for (i, &xi) in inputs.iter().enumerate() {
let expected = xi.exp();
expected_sum += expected;
let denom = expected.abs().max(1e-30);
let rel_err = (x[i] - expected).abs() / denom;
assert!(
rel_err < 5e-4,
"len={len} exp({xi}) = {} vs true {expected}, rel_err={rel_err:.3e}",
x[i]
);
}
let sum_rel = (got_sum - expected_sum).abs() / expected_sum.abs().max(1e-30);
assert!(
sum_rel < 5e-4,
"len={len} sum mismatch: got {got_sum}, exp {expected_sum}, rel={sum_rel:.3e}"
);
}
}
fn ref_sigmoid_tanh_clamp(a: &[f32], q: &[f32], clamp: f32) -> Vec<f32> {
a.iter()
.zip(q.iter())
.map(|(&ai, &qi)| (2.0 * fast_sigmoid(ai + qi) - 1.0).clamp(-clamp, clamp))
.collect()
}
#[test]
fn simd_sigmoid_tanh_clamp_matches_scalar_reference() {
let cases = [-40.0f32, -10.0, -1.0, 0.0, 1.0, 10.0, 40.0];
let zeros = vec![0.0f32; cases.len()];
let clamp = 6.0f32;
let mut out = vec![0.0f32; cases.len()];
simd_sigmoid_tanh_clamp_inplace(&mut out, &cases, &zeros, clamp);
let expected = ref_sigmoid_tanh_clamp(&cases, &zeros, clamp);
for (i, (got, want)) in out.iter().zip(expected.iter()).enumerate() {
assert!(
(got - want).abs() < 1e-6,
"mismatch at i={i} (a={}): simd={got}, scalar={want}, diff={}",
cases[i],
(got - want).abs()
);
}
}
#[test]
fn simd_sigmoid_tanh_clamp_output_in_range_with_outliers() {
let a = [
100.0f32, -100.0, 50.0, -50.0, 0.0, 1.5, -2.3, 10.0, -10.0, 2.5, -2.5, 0.001, 25.0, -25.0,
80.0, -80.0,
];
let q = [0.0f32; 16];
let clamp = 6.0f32;
let mut out = [0.0f32; 16];
simd_sigmoid_tanh_clamp_inplace(&mut out, &a, &q, clamp);
for (i, &v) in out.iter().enumerate() {
assert!(
v > -clamp && v < clamp,
"out-of-range at i={i}: {v} not in ({}, {})",
-clamp,
clamp
);
}
assert!((out[0] - 1.0).abs() < 1e-6, "a=100 → ~+1, got {}", out[0]);
assert!((out[1] + 1.0).abs() < 1e-6, "a=-100 → ~-1, got {}", out[1]);
}
#[test]
fn simd_sigmoid_tanh_clamp_saturation_at_clamp_boundary() {
let a = [100.0f32];
let q = [0.0f32];
let clamp = 0.5f32;
let mut out = [0.0f32];
simd_sigmoid_tanh_clamp_inplace(&mut out, &a, &q, clamp);
assert_eq!(out[0], 0.5, "clamp saturation must give exactly +clamp");
let a_neg = [-100.0f32];
let mut out_neg = [0.0f32];
simd_sigmoid_tanh_clamp_inplace(&mut out_neg, &a_neg, &q, clamp);
assert_eq!(
out_neg[0], -0.5,
"clamp saturation must give exactly -clamp"
);
}
#[test]
fn simd_sigmoid_tanh_clamp_length_33_matches_length_32_prefix() {
let mut rng = fastrand::Rng::with_seed(1234);
let a33: Vec<f32> = (0..33).map(|_| rng.f32() * 20.0 - 10.0).collect();
let q33: Vec<f32> = (0..33).map(|_| rng.f32() * 2.0 - 1.0).collect();
let clamp = 6.0f32;
let mut out33 = vec![0.0f32; 33];
simd_sigmoid_tanh_clamp_inplace(&mut out33, &a33, &q33, clamp);
let mut out32 = vec![0.0f32; 32];
simd_sigmoid_tanh_clamp_inplace(&mut out32, &a33[..32], &q33[..32], clamp);
for i in 0..32 {
assert_eq!(
out33[i], out32[i],
"prefix mismatch at i={i}: len33={}, len32={}",
out33[i], out32[i]
);
}
let expected_tail = (2.0 * fast_sigmoid(a33[32] + q33[32]) - 1.0).clamp(-clamp, clamp);
assert_eq!(
out33[32], expected_tail,
"scalar tail mismatch: simd={}, scalar={}",
out33[32], expected_tail
);
}