pub const MAX_SIMD_LANES: usize = 64;
pub trait SimdKernel<T: crate::scalar::Scalar>:
crate::private::Sealed + Send + Sync + Sized + 'static
{
type Vector: Copy + Send + Sync + 'static;
type Mask: Copy + Send + Sync + 'static;
type IndexVector: Copy + Send + Sync + 'static;
const LANE_COUNT: usize;
const LANE_BOUND_CHECK: () = assert!(
Self::LANE_COUNT <= MAX_SIMD_LANES,
"SimdKernel::LANE_COUNT exceeds MAX_SIMD_LANES; widen the scalar-fallback stack buffers"
);
const UNROLL_FACTOR: usize = 4;
unsafe fn load_aligned(ptr: *const T) -> Self::Vector;
unsafe fn load_unaligned(ptr: *const T) -> Self::Vector;
unsafe fn store_aligned(ptr: *mut T, val: Self::Vector);
unsafe fn store_unaligned(ptr: *mut T, val: Self::Vector);
const SUPPORTS_NT_STORE: bool = false;
#[inline(always)]
unsafe fn store_streaming(ptr: *mut T, val: Self::Vector) {
Self::store_aligned(ptr, val);
}
#[inline(always)]
fn stream_write_barrier() {}
unsafe fn add(a: Self::Vector, b: Self::Vector) -> Self::Vector;
unsafe fn mul(a: Self::Vector, b: Self::Vector) -> Self::Vector;
unsafe fn sub(a: Self::Vector, b: Self::Vector) -> Self::Vector {
crate::kernel_helpers::generic_binary_op::<T, Self, _>(a, b, |x, y| x - y)
}
unsafe fn fmadd(a: Self::Vector, b: Self::Vector, c: Self::Vector) -> Self::Vector;
unsafe fn sum_reduce(v: Self::Vector) -> T;
unsafe fn masked_load_unaligned(
ptr: *const T,
mask: Self::Mask,
src: Self::Vector,
) -> Self::Vector {
crate::kernel_helpers::generic_masked_load::<T, Self>(ptr, mask, src)
}
unsafe fn masked_store_unaligned(ptr: *mut T, mask: Self::Mask, val: Self::Vector) {
crate::kernel_helpers::generic_masked_store::<T, Self>(ptr, mask, val)
}
unsafe fn masked_add(
a: Self::Vector,
b: Self::Vector,
mask: Self::Mask,
src: Self::Vector,
) -> Self::Vector {
Self::blend(Self::mask_to_vector(mask), Self::add(a, b), src)
}
unsafe fn masked_mul(
a: Self::Vector,
b: Self::Vector,
mask: Self::Mask,
src: Self::Vector,
) -> Self::Vector {
Self::blend(Self::mask_to_vector(mask), Self::mul(a, b), src)
}
unsafe fn masked_fmadd(
a: Self::Vector,
b: Self::Vector,
c: Self::Vector,
mask: Self::Mask,
) -> Self::Vector {
Self::blend(Self::mask_to_vector(mask), Self::fmadd(a, b, c), c)
}
unsafe fn masked_sum_reduce(v: Self::Vector, mask: Self::Mask) -> T {
Self::sum_reduce(Self::blend(Self::mask_to_vector(mask), v, Self::zero()))
}
unsafe fn compress(src: Self::Vector, mask: Self::Mask) -> Self::Vector;
unsafe fn expand(src: Self::Vector, mask: Self::Mask, fill: Self::Vector) -> Self::Vector;
unsafe fn gather(base: *const T, indices: Self::IndexVector) -> Self::Vector;
unsafe fn gather_masked(
base: *const T,
indices: Self::IndexVector,
mask: Self::Mask,
src: Self::Vector,
) -> Self::Vector;
unsafe fn mask_from_bools(bits: &[bool]) -> Self::Mask;
unsafe fn leading_k_mask(k: usize) -> Self::Mask;
unsafe fn mask_from_bitmask(bm: u64) -> Self::Mask {
crate::kernel_helpers::generic_mask_from_bitmask::<T, Self>(bm)
}
unsafe fn mask_to_vector(mask: Self::Mask) -> Self::Vector;
unsafe fn vector_to_mask(v: Self::Vector) -> Self::Mask;
#[inline(always)]
unsafe fn scan_vector<Op: crate::ops::ScanOp<T>, SMode: crate::ops::ScanMode>(
v: Self::Vector,
mut carry: T,
) -> (Self::Vector, T) {
const { Self::LANE_BOUND_CHECK };
let mut buf = [core::mem::MaybeUninit::<T>::uninit(); MAX_SIMD_LANES];
let lanes = Self::LANE_COUNT;
Self::store_unaligned(buf.as_mut_ptr() as *mut T, v);
let mut out_buf = [core::mem::MaybeUninit::<T>::uninit(); MAX_SIMD_LANES];
if SMode::IS_INCLUSIVE {
for j in 0..lanes {
let temp = buf[j].assume_init();
carry = Op::combine(carry, temp);
out_buf[j].write(carry);
}
} else {
for j in 0..lanes {
let temp = buf[j].assume_init();
out_buf[j].write(carry);
carry = Op::combine(carry, temp);
}
}
(Self::load_unaligned(out_buf.as_ptr() as *const T), carry)
}
unsafe fn zero() -> Self::Vector {
Self::splat(T::ZERO)
}
unsafe fn splat(val: T) -> Self::Vector;
unsafe fn div(a: Self::Vector, b: Self::Vector) -> Self::Vector {
crate::kernel_helpers::generic_binary_op::<T, Self, _>(a, b, |x, y| x / y)
}
unsafe fn bitand(a: Self::Vector, b: Self::Vector) -> Self::Vector {
crate::kernel_helpers::generic_binary_op::<T, Self, _>(a, b, |x, y| x.bitand(y))
}
unsafe fn bitor(a: Self::Vector, b: Self::Vector) -> Self::Vector {
crate::kernel_helpers::generic_binary_op::<T, Self, _>(a, b, |x, y| x.bitor(y))
}
unsafe fn bitxor(a: Self::Vector, b: Self::Vector) -> Self::Vector {
crate::kernel_helpers::generic_binary_op::<T, Self, _>(a, b, |x, y| x.bitxor(y))
}
unsafe fn abs(a: Self::Vector) -> Self::Vector {
crate::kernel_helpers::generic_unary_op::<T, Self, _>(a, |x| x.abs())
}
unsafe fn min(a: Self::Vector, b: Self::Vector) -> Self::Vector {
crate::kernel_helpers::generic_binary_op::<T, Self, _>(
a,
b,
|x, y| if x < y { x } else { y },
)
}
unsafe fn max(a: Self::Vector, b: Self::Vector) -> Self::Vector {
crate::kernel_helpers::generic_binary_op::<T, Self, _>(
a,
b,
|x, y| if x > y { x } else { y },
)
}
unsafe fn sqrt(a: Self::Vector) -> Self::Vector {
crate::kernel_helpers::generic_unary_op::<T, Self, _>(a, |x| x.sqrt())
}
unsafe fn recip_sqrt(a: Self::Vector) -> Self::Vector {
crate::kernel_helpers::generic_unary_op::<T, Self, _>(a, |x| T::ONE / x.sqrt())
}
unsafe fn cmp_eq(a: Self::Vector, b: Self::Vector) -> Self::Vector {
crate::kernel_helpers::generic_binary_op::<T, Self, _>(a, b, |x, y| {
if x == y {
T::ALL_ONES
} else {
T::ZERO
}
})
}
unsafe fn cmp_ne(a: Self::Vector, b: Self::Vector) -> Self::Vector {
crate::kernel_helpers::generic_binary_op::<T, Self, _>(a, b, |x, y| {
if x != y {
T::ALL_ONES
} else {
T::ZERO
}
})
}
unsafe fn cmp_lt(a: Self::Vector, b: Self::Vector) -> Self::Vector {
crate::kernel_helpers::generic_binary_op::<T, Self, _>(a, b, |x, y| {
if x < y {
T::ALL_ONES
} else {
T::ZERO
}
})
}
unsafe fn cmp_le(a: Self::Vector, b: Self::Vector) -> Self::Vector {
crate::kernel_helpers::generic_binary_op::<T, Self, _>(a, b, |x, y| {
if x <= y {
T::ALL_ONES
} else {
T::ZERO
}
})
}
unsafe fn cmp_gt(a: Self::Vector, b: Self::Vector) -> Self::Vector {
crate::kernel_helpers::generic_binary_op::<T, Self, _>(a, b, |x, y| {
if x > y {
T::ALL_ONES
} else {
T::ZERO
}
})
}
unsafe fn cmp_ge(a: Self::Vector, b: Self::Vector) -> Self::Vector {
crate::kernel_helpers::generic_binary_op::<T, Self, _>(a, b, |x, y| {
if x >= y {
T::ALL_ONES
} else {
T::ZERO
}
})
}
unsafe fn blend(
mask: Self::Vector,
true_val: Self::Vector,
false_val: Self::Vector,
) -> Self::Vector {
crate::kernel_helpers::generic_blend::<T, Self>(mask, true_val, false_val)
}
#[inline(always)]
unsafe fn neg(a: Self::Vector) -> Self::Vector {
Self::bitxor(a, Self::splat(T::SIGN_MASK))
}
#[inline(always)]
unsafe fn bitnot(a: Self::Vector) -> Self::Vector {
Self::bitxor(a, Self::splat(T::ALL_ONES))
}
unsafe fn mask_to_bitmask(mask: Self::Mask) -> u64;
unsafe fn min_reduce(v: Self::Vector) -> T {
crate::kernel_helpers::generic_horizontal_reduce::<T, Self>(v, T::MAX_VALUE, |a, b| {
a.min_scalar(b)
})
}
unsafe fn max_reduce(v: Self::Vector) -> T {
crate::kernel_helpers::generic_horizontal_reduce::<T, Self>(v, T::MIN_VALUE, |a, b| {
a.max_scalar(b)
})
}
unsafe fn popcount(a: Self::Vector) -> Self::Vector {
crate::kernel_helpers::generic_unary_op::<T, Self, _>(a, |x| {
T::cast_from(x.count_ones() as i32)
})
}
unsafe fn horizontal_bitwise_and(v: Self::Vector) -> T {
crate::kernel_helpers::generic_horizontal_reduce::<T, Self>(v, T::ALL_ONES, |a, b| {
a.bitand(b)
})
}
unsafe fn horizontal_bitwise_or(v: Self::Vector) -> T {
crate::kernel_helpers::generic_horizontal_reduce::<T, Self>(v, T::ZERO, |a, b| a.bitor(b))
}
unsafe fn horizontal_bitwise_xor(v: Self::Vector) -> T {
crate::kernel_helpers::generic_horizontal_reduce::<T, Self>(v, T::ZERO, |a, b| a.bitxor(b))
}
#[inline(always)]
unsafe fn swap_adjacent(v: Self::Vector) -> Self::Vector {
const { Self::LANE_BOUND_CHECK };
let mut buf = [core::mem::MaybeUninit::<T>::uninit(); MAX_SIMD_LANES];
let lanes = Self::LANE_COUNT;
Self::store_unaligned(buf.as_mut_ptr() as *mut T, v);
let mut i = 0usize;
while i + 1 < lanes {
buf.swap(i, i + 1);
i += 2;
}
Self::load_unaligned(buf.as_ptr() as *const T)
}
#[inline(always)]
unsafe fn dup_even(v: Self::Vector) -> Self::Vector {
const { Self::LANE_BOUND_CHECK };
let mut buf = [core::mem::MaybeUninit::<T>::uninit(); MAX_SIMD_LANES];
let lanes = Self::LANE_COUNT;
Self::store_unaligned(buf.as_mut_ptr() as *mut T, v);
let mut out = [core::mem::MaybeUninit::<T>::uninit(); MAX_SIMD_LANES];
for i in 0..lanes {
let src_val = buf[i & !1].assume_init();
out[i].write(src_val);
}
Self::load_unaligned(out.as_ptr() as *const T)
}
#[inline(always)]
unsafe fn dup_odd(v: Self::Vector) -> Self::Vector {
const { Self::LANE_BOUND_CHECK };
let mut buf = [core::mem::MaybeUninit::<T>::uninit(); MAX_SIMD_LANES];
let lanes = Self::LANE_COUNT;
Self::store_unaligned(buf.as_mut_ptr() as *mut T, v);
let mut out = [core::mem::MaybeUninit::<T>::uninit(); MAX_SIMD_LANES];
for i in 0..lanes {
let src_val = buf[(i | 1).min(lanes - 1)].assume_init();
out[i].write(src_val);
}
Self::load_unaligned(out.as_ptr() as *const T)
}
#[inline(always)]
unsafe fn fmaddsub(a: Self::Vector, b: Self::Vector, c: Self::Vector) -> Self::Vector {
crate::kernel_helpers::generic_alternating_fma::<T, Self, false>(a, b, c)
}
#[inline(always)]
unsafe fn fmsubadd(a: Self::Vector, b: Self::Vector, c: Self::Vector) -> Self::Vector {
crate::kernel_helpers::generic_alternating_fma::<T, Self, true>(a, b, c)
}
}