use hermes_simd::{Scalar, SveArch};
use hermes_simd_core::align::Unaligned;
use hermes_simd_core::execution::Unmasked;
use hermes_simd_core::kernel::SimdKernel;
use hermes_simd_core::view::SimdView;
use proptest::prelude::*;
fn lane_bits<A: SimdKernel<f32>>(bm: u64) -> u64 {
bm & ((1u64 << A::LANE_COUNT) - 1)
}
fn check_bitmask_roundtrip<A: SimdKernel<f32>>(bm: u64) {
let bm = lane_bits::<A>(bm);
let roundtrip = unsafe { A::mask_to_bitmask(A::mask_from_bitmask(bm)) };
assert_eq!(
lane_bits::<A>(roundtrip),
bm,
"bitmask round-trip failed for {bm:#b}"
);
}
fn check_compress_expand_identity<A: SimdKernel<f32>>(bm: u64, vals: &[f32]) {
let lanes = A::LANE_COUNT;
let bm = lane_bits::<A>(bm);
let src: Vec<f32> = (0..lanes).map(|i| vals[i % vals.len()]).collect();
const FILL: f32 = -512.5;
let mut out = vec![0.0f32; lanes];
unsafe {
let v = A::load_unaligned(src.as_ptr());
let mask = A::mask_from_bitmask(bm);
let compressed = A::compress(v, mask);
let restored = A::expand(compressed, mask, A::splat(FILL));
A::store_unaligned(out.as_mut_ptr(), restored);
}
for (i, &x) in out.iter().enumerate() {
if (bm >> i) & 1 == 1 {
assert_eq!(x, src[i], "active lane {i} not restored (mask {bm:#b})");
} else {
assert_eq!(x, FILL, "inactive lane {i} not filled (mask {bm:#b})");
}
}
}
fn check_gather_matches_reference<A>(values: &[f32], indices: &[i32])
where
A: hermes_simd_core::arch::SimdArch + SimdKernel<f32>,
{
let view = SimdView::<f32, A, Unaligned, Unmasked, &[f32]>::new(values).unwrap();
let mut out = vec![0.0f32; indices.len()];
view.gather(indices, &mut out).unwrap();
for (k, &idx) in indices.iter().enumerate() {
assert_eq!(out[k], values[idx as usize], "gather mismatch at {k}");
}
}
fn check_leading_k_masked_sum<A: SimdKernel<f32>>() {
let lanes = A::LANE_COUNT;
let vals: Vec<f32> = (0..lanes).map(|i| (i + 1) as f32).collect();
for k in 0..=lanes + 2 {
let total = unsafe {
let v = A::load_unaligned(vals.as_ptr());
A::masked_sum_reduce(v, A::leading_k_mask(k))
};
let expected: f32 = vals[..k.min(lanes)].iter().sum();
assert_eq!(total, expected, "leading_k_mask({k}) sum mismatch");
}
}
fn check_masked_merge_ops<A: SimdKernel<f32>>() {
let lanes = A::LANE_COUNT;
let a_vals: Vec<f32> = (0..lanes).map(|i| (i + 1) as f32).collect();
let b_vals: Vec<f32> = (0..lanes).map(|i| (2 * i + 3) as f32).collect();
let src_vals: Vec<f32> = (0..lanes).map(|i| -((i + 1) as f32)).collect();
let k = lanes / 2; let mut buf = vec![0.0f32; lanes];
unsafe {
let a = A::load_unaligned(a_vals.as_ptr());
let b = A::load_unaligned(b_vals.as_ptr());
let src = A::load_unaligned(src_vals.as_ptr());
let mask = A::leading_k_mask(k);
A::store_unaligned(
buf.as_mut_ptr(),
A::masked_load_unaligned(a_vals.as_ptr(), mask, src),
);
for i in 0..lanes {
let want = if i < k { a_vals[i] } else { src_vals[i] };
assert_eq!(buf[i], want, "masked_load lane {i}");
}
A::store_unaligned(buf.as_mut_ptr(), A::masked_add(a, b, mask, src));
for i in 0..lanes {
let want = if i < k {
a_vals[i] + b_vals[i]
} else {
src_vals[i]
};
assert_eq!(buf[i], want, "masked_add lane {i}");
}
A::store_unaligned(buf.as_mut_ptr(), A::masked_mul(a, b, mask, src));
for i in 0..lanes {
let want = if i < k {
a_vals[i] * b_vals[i]
} else {
src_vals[i]
};
assert_eq!(buf[i], want, "masked_mul lane {i}");
}
A::store_unaligned(buf.as_mut_ptr(), A::masked_fmadd(a, b, src, mask));
for i in 0..lanes {
let want = if i < k {
a_vals[i] * b_vals[i] + src_vals[i]
} else {
src_vals[i]
};
assert_eq!(buf[i], want, "masked_fmadd lane {i}");
}
let mut dst = src_vals.clone();
A::masked_store_unaligned(dst.as_mut_ptr(), mask, a);
for i in 0..lanes {
let want = if i < k { a_vals[i] } else { src_vals[i] };
assert_eq!(dst[i], want, "masked_store lane {i}");
}
}
}
fn check_vector_to_mask_roundtrip<A: SimdKernel<f32>>(bm: u64) {
let bm = lane_bits::<A>(bm);
let roundtrip = unsafe {
A::mask_to_bitmask(A::vector_to_mask(A::mask_to_vector(A::mask_from_bitmask(
bm,
))))
};
assert_eq!(
lane_bits::<A>(roundtrip),
bm,
"vector_to_mask round-trip failed for {bm:#b}"
);
}
fn check_vector_to_mask_matches_cmp<A: SimdKernel<f32>>(vals: &[f32]) {
let lanes = A::LANE_COUNT;
let a_vals: Vec<f32> = (0..lanes).map(|i| vals[i % vals.len()]).collect();
let b_vals: Vec<f32> = a_vals
.iter()
.enumerate()
.map(|(i, &v)| if i % 2 == 0 { v } else { v + 1.0 })
.collect();
let bm = unsafe {
let a = A::load_unaligned(a_vals.as_ptr());
let b = A::load_unaligned(b_vals.as_ptr());
A::mask_to_bitmask(A::vector_to_mask(A::cmp_eq(a, b)))
};
for i in 0..lanes {
let want = a_vals[i] == b_vals[i];
let got = (bm >> i) & 1 == 1;
assert_eq!(got, want, "cmp_eq lane {i}: {} vs {}", a_vals[i], b_vals[i]);
}
}
fn check_cmp_ne_complements_cmp_eq<A: SimdKernel<f32>>(vals: &[f32]) {
let lanes = A::LANE_COUNT;
let mut a_vals: Vec<f32> = (0..lanes).map(|i| vals[i % vals.len()]).collect();
let mut b_vals: Vec<f32> = a_vals
.iter()
.enumerate()
.map(|(i, &v)| if i % 2 == 0 { v } else { v + 1.0 })
.collect();
a_vals[0] = f32::NAN;
b_vals[0] = f32::NAN;
if lanes > 1 {
a_vals[1] = f32::NAN;
}
let (eq, ne) = unsafe {
let a = A::load_unaligned(a_vals.as_ptr());
let b = A::load_unaligned(b_vals.as_ptr());
(
lane_bits::<A>(A::mask_to_bitmask(A::vector_to_mask(A::cmp_eq(a, b)))),
lane_bits::<A>(A::mask_to_bitmask(A::vector_to_mask(A::cmp_ne(a, b)))),
)
};
assert_eq!(
ne,
lane_bits::<A>(!eq),
"cmp_ne must be the complement of cmp_eq (eq {eq:#b}, ne {ne:#b})"
);
for i in 0..lanes {
let want = a_vals[i] != b_vals[i];
assert_eq!(
(ne >> i) & 1 == 1,
want,
"cmp_ne lane {i}: {} vs {}",
a_vals[i],
b_vals[i]
);
}
}
fn check_blend_honors_canonical_mask<A: SimdKernel<f32>>(bm: u64) {
let lanes = A::LANE_COUNT;
let bm = lane_bits::<A>(bm);
let true_vals: Vec<f32> = (0..lanes).map(|i| (i + 1) as f32).collect();
let false_vals: Vec<f32> = (0..lanes).map(|i| -((i + 1) as f32)).collect();
let mut out = vec![0.0f32; lanes];
unsafe {
let selected = A::load_unaligned(true_vals.as_ptr());
let rejected = A::load_unaligned(false_vals.as_ptr());
let mask = A::mask_to_vector(A::mask_from_bitmask(bm));
A::store_unaligned(out.as_mut_ptr(), A::blend(mask, selected, rejected));
}
for (i, &got) in out.iter().enumerate() {
let want = if (bm >> i) & 1 == 1 {
true_vals[i]
} else {
false_vals[i]
};
assert_eq!(got, want, "blend lane {i} (mask {bm:#b})");
}
}
fn check_all_kernel_invariants<A>(bm: u64, vals: &[f32])
where
A: hermes_simd_core::arch::SimdArch + SimdKernel<f32>,
{
check_bitmask_roundtrip::<A>(bm);
check_vector_to_mask_roundtrip::<A>(bm);
check_vector_to_mask_matches_cmp::<A>(vals);
check_cmp_ne_complements_cmp_eq::<A>(vals);
check_blend_honors_canonical_mask::<A>(bm);
check_compress_expand_identity::<A>(bm, vals);
check_leading_k_masked_sum::<A>();
check_masked_merge_ops::<A>();
}
proptest! {
#[test]
fn prop_kernel_invariants_all_backends(
bm in any::<u64>(),
vals in prop::collection::vec(-1000.0f32..1000.0, 1..32),
) {
check_all_kernel_invariants::<Scalar>(bm, &vals);
check_all_kernel_invariants::<SveArch>(bm, &vals);
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
{
if std::is_x86_feature_detected!("avx2") && std::is_x86_feature_detected!("fma") {
check_all_kernel_invariants::<hermes_simd::Avx2>(bm, &vals);
}
if std::is_x86_feature_detected!("avx512f") {
check_all_kernel_invariants::<hermes_simd::Avx512>(bm, &vals);
}
}
#[cfg(target_arch = "aarch64")]
{
check_all_kernel_invariants::<hermes_simd::Neon>(bm, &vals);
}
}
#[test]
fn prop_gather_matches_reference_all_backends(
(values, indices) in prop::collection::vec(-1000.0f32..1000.0, 1..256)
.prop_flat_map(|v| {
let n = v.len();
(Just(v), prop::collection::vec(0..n as i32, 0..64))
}),
) {
check_gather_matches_reference::<Scalar>(&values, &indices);
check_gather_matches_reference::<SveArch>(&values, &indices);
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
{
if std::is_x86_feature_detected!("avx2") && std::is_x86_feature_detected!("fma") {
check_gather_matches_reference::<hermes_simd::Avx2>(&values, &indices);
}
if std::is_x86_feature_detected!("avx512f") {
check_gather_matches_reference::<hermes_simd::Avx512>(&values, &indices);
}
}
#[cfg(target_arch = "aarch64")]
{
check_gather_matches_reference::<hermes_simd::Neon>(&values, &indices);
}
}
}
#[test]
fn gather_rejects_out_of_bounds_indices() {
let values = [1.0f32, 2.0, 3.0];
let view = SimdView::<f32, Scalar, Unaligned, Unmasked, &[f32]>::new(&values).unwrap();
let mut out = [0.0f32; 2];
assert!(matches!(
view.gather(&[0, 3], &mut out),
Err(hermes_simd_core::view::SimdError::IndexOutOfBounds)
));
assert!(matches!(
view.gather(&[-1, 0], &mut out),
Err(hermes_simd_core::view::SimdError::IndexOutOfBounds)
));
}
fn check_recip_sqrt_f32<A: SimdKernel<f32>>() {
let lanes = A::LANE_COUNT;
let inputs: Vec<f32> = (0..lanes).map(|i| 0.3 + 1.7 * i as f32).collect();
let mut out = vec![0.0f32; lanes];
unsafe {
A::store_unaligned(
out.as_mut_ptr(),
A::recip_sqrt(A::load_unaligned(inputs.as_ptr())),
);
}
let tol = 8.0 * f64::from(f32::EPSILON);
for (&y, &x) in out.iter().zip(inputs.iter()) {
let want = 1.0_f64 / f64::from(x).sqrt();
let rel = (f64::from(y) - want).abs() / want;
assert!(
rel <= tol,
"f32 recip_sqrt: x={x} got={y} want={want} rel={rel:e}"
);
}
}
fn check_recip_sqrt_f64<A: SimdKernel<f64>>() {
let lanes = A::LANE_COUNT;
let inputs: Vec<f64> = (0..lanes).map(|i| 0.3 + 1.7 * i as f64).collect();
let mut out = vec![0.0f64; lanes];
unsafe {
A::store_unaligned(
out.as_mut_ptr(),
A::recip_sqrt(A::load_unaligned(inputs.as_ptr())),
);
}
let tol = 4.0 * f64::EPSILON;
for (&y, &x) in out.iter().zip(inputs.iter()) {
let want = 1.0_f64 / x.sqrt();
let rel = (y - want).abs() / want;
assert!(
rel <= tol,
"f64 recip_sqrt: x={x} got={y} want={want} rel={rel:e}"
);
}
}
#[test]
fn recip_sqrt_is_full_precision_all_backends() {
check_recip_sqrt_f32::<Scalar>();
check_recip_sqrt_f64::<Scalar>();
check_recip_sqrt_f32::<SveArch>();
check_recip_sqrt_f64::<SveArch>();
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
{
if std::is_x86_feature_detected!("avx2") && std::is_x86_feature_detected!("fma") {
check_recip_sqrt_f32::<hermes_simd::Avx2>();
check_recip_sqrt_f64::<hermes_simd::Avx2>();
}
if std::is_x86_feature_detected!("avx512f") {
check_recip_sqrt_f32::<hermes_simd::Avx512>();
check_recip_sqrt_f64::<hermes_simd::Avx512>();
}
}
#[cfg(target_arch = "aarch64")]
{
check_recip_sqrt_f32::<hermes_simd::Neon>();
check_recip_sqrt_f64::<hermes_simd::Neon>();
}
}