use super::{complex, modular, simd_ops::SimdOps};
use hermes_simd_core::scalar::Scalar as ScalarTrait;
use hermes_simd_core::sparse::{
BlockedCooData, CsrData, DenseWithMaskData, SellPData, ValidatedData,
};
use hermes_simd_core::view::SimdError;
#[inline(always)]
pub fn sum<T: SimdOps>(data: &[T]) -> T {
T::sum(data)
}
#[inline(always)]
pub fn min<T: SimdOps>(data: &[T]) -> T {
T::min(data)
}
#[inline(always)]
pub fn max<T: SimdOps>(data: &[T]) -> T {
T::max(data)
}
#[inline(always)]
pub fn abs_sum<T: SimdOps>(data: &[T]) -> T {
T::abs_sum(data)
}
#[inline(always)]
pub fn abs_max<T: SimdOps>(data: &[T]) -> T {
T::abs_max(data)
}
#[inline(always)]
pub fn scale<T: SimdOps>(data: &mut [T], scalar: T) {
T::scale(data, scalar)
}
#[inline(always)]
pub fn argmin<T: SimdOps>(data: &[T]) -> Option<(usize, T)> {
T::argmin(data)
}
#[inline(always)]
pub fn argmax<T: SimdOps>(data: &[T]) -> Option<(usize, T)> {
T::argmax(data)
}
#[inline(always)]
pub fn dot<T: SimdOps>(a: &[T], b: &[T]) -> Result<T, SimdError> {
T::dot(a, b)
}
#[inline(always)]
pub fn axpy<T: SimdOps>(alpha: T, x: &[T], out: &mut [T]) -> Result<(), SimdError> {
T::axpy(alpha, x, out)
}
#[inline(always)]
pub fn axpy_mul<T: SimdOps>(alpha: T, a: &[T], b: &[T], out: &mut [T]) -> Result<(), SimdError> {
T::axpy_mul(alpha, a, b, out)
}
#[inline(always)]
pub fn axpy_rows<T: SimdOps>(
alphas: &[T],
x: &[T],
out: &mut [T],
row_stride: usize,
rows: usize,
cols: usize,
) -> Result<(), SimdError> {
T::axpy_rows(alphas, x, out, row_stride, rows, cols)
}
#[inline(always)]
pub fn axpy_rows_batch<T: SimdOps>(
alphas: &[T],
x_panel: &[T],
out: &mut [T],
row_stride: usize,
rows: usize,
depth: usize,
cols: usize,
) -> Result<(), SimdError> {
T::axpy_rows_batch(alphas, x_panel, out, row_stride, rows, depth, cols)
}
#[inline(always)]
pub fn elementwise_mul<T: SimdOps>(a: &[T], b: &[T], out: &mut [T]) -> Result<(), SimdError> {
T::elementwise_mul(a, b, out)
}
#[inline(always)]
pub fn elementwise_add<T: SimdOps>(a: &[T], b: &[T], out: &mut [T]) -> Result<(), SimdError> {
T::elementwise_add(a, b, out)
}
#[inline(always)]
pub fn elementwise_sub<T: SimdOps>(a: &[T], b: &[T], out: &mut [T]) -> Result<(), SimdError> {
T::elementwise_sub(a, b, out)
}
#[inline(always)]
pub fn elementwise_div<T: SimdOps>(a: &[T], b: &[T], out: &mut [T]) -> Result<(), SimdError> {
T::elementwise_div(a, b, out)
}
#[inline]
pub fn ntt_butterfly_stage_u64(
data: &mut [u64],
stage_len: usize,
twiddles: &[u64],
modulus: u64,
) -> Result<(), SimdError> {
modular::ntt_butterfly_stage_u64(data, stage_len, twiddles, modulus)
}
#[inline(always)]
pub fn masked_sum<T: SimdOps>(data: &[T], mask: &[bool]) -> T {
T::masked_sum(data, mask)
}
#[inline(always)]
pub fn masked_dot<T: SimdOps>(a: &[T], b: &[T], mask: &[bool]) -> Result<T, SimdError> {
T::masked_dot(a, b, mask)
}
#[inline(always)]
pub fn masked_add<T: SimdOps>(
a: &[T],
b: &[T],
mask: &[bool],
out: &mut [T],
) -> Result<(), SimdError> {
T::masked_add(a, b, mask, out)
}
#[inline(always)]
pub fn spmv_csr<T: SimdOps>(data: ValidatedData<CsrData<'_, T>>, x: &[T], y: &mut [T]) {
T::spmv_csr(data, x, y)
}
#[inline(always)]
pub fn spmv_bcoo<T: SimdOps, const BM: usize, const BN: usize>(
data: ValidatedData<BlockedCooData<'_, T, BM, BN>>,
x: &[T],
y: &mut [T],
) {
T::spmv_bcoo::<BM, BN>(data, x, y)
}
#[inline(always)]
pub fn spmv_dense_masked<T: SimdOps>(data: DenseWithMaskData<'_, T>, x: &[T], y: &mut [T]) {
T::spmv_dense_masked(data, x, y)
}
#[inline(always)]
pub fn spmv_sellp<T: SimdOps, const C: usize>(
data: ValidatedData<SellPData<'_, T, C>>,
x: &[T],
y: &mut [T],
) {
T::spmv_sellp::<C>(data, x, y)
}
#[inline(always)]
pub fn tiled_gemm<T: SimdOps>(
a: &[T],
b: &[T],
c: &mut [T],
m: usize,
n: usize,
k: usize,
) -> Result<(), SimdError> {
T::tiled_gemm(a, b, c, m, n, k)
}
#[inline(always)]
pub fn gemv<T: SimdOps>(
a: &[T],
x: &[T],
y: &mut [T],
nrows: usize,
ncols: usize,
) -> Result<(), SimdError> {
T::gemv(a, x, y, nrows, ncols)
}
#[inline(always)]
pub fn gemv_transpose<T: SimdOps>(
a: &[T],
x: &[T],
y: &mut [T],
nrows: usize,
ncols: usize,
) -> Result<(), SimdError> {
T::gemv_transpose(a, x, y, nrows, ncols)
}
#[inline(always)]
pub fn gemv_strided<T: SimdOps>(
a: &[T],
x: &[T],
y: &mut [T],
nrows: usize,
ncols: usize,
lda: usize,
) -> Result<(), SimdError> {
T::gemv_strided(a, x, y, nrows, ncols, lda)
}
#[inline(always)]
pub fn gemv_transpose_strided<T: SimdOps>(
a: &[T],
x: &[T],
y: &mut [T],
nrows: usize,
ncols: usize,
lda: usize,
) -> Result<(), SimdError> {
T::gemv_transpose_strided(a, x, y, nrows, ncols, lda)
}
#[inline]
pub fn interleaved_complex_mul_assign<T, A, const CONJ_B: bool>(
a: &mut [T],
b: &[T],
) -> Result<(), SimdError>
where
T: ScalarTrait + core::ops::Neg<Output = T>,
A: hermes_simd_core::arch::SimdArch + hermes_simd_core::kernel::SimdKernel<T>,
{
complex::interleaved_complex_mul_assign::<T, A, CONJ_B>(a, b)
}
#[inline]
pub fn interleaved_complex_dot<T, A, const CONJ_B: bool>(
a: &[T],
b: &[T],
) -> Result<(T, T), SimdError>
where
T: ScalarTrait + core::ops::Neg<Output = T>,
A: hermes_simd_core::arch::SimdArch + hermes_simd_core::kernel::SimdKernel<T>,
{
complex::interleaved_complex_dot::<T, A, CONJ_B>(a, b)
}
#[inline]
pub fn interleaved_complex_mul_assign_runtime<T, const CONJ_B: bool>(
a: &mut [T],
b: &[T],
) -> Result<(), SimdError>
where
T: SimdOps + core::ops::Neg<Output = T>,
{
T::interleaved_complex_mul_assign::<CONJ_B>(a, b)
}
#[inline]
pub fn interleaved_complex_dot_runtime<T, const CONJ_B: bool>(
a: &[T],
b: &[T],
) -> Result<(T, T), SimdError>
where
T: SimdOps + core::ops::Neg<Output = T>,
{
T::interleaved_complex_dot::<CONJ_B>(a, b)
}
#[inline(always)]
pub fn reduce_popcount<T: SimdOps>(data: &[T]) -> usize {
T::reduce_popcount(data)
}
#[inline(always)]
pub fn reduce_popcount_and<T: SimdOps>(a: &[T], b: &[T]) -> Result<usize, SimdError> {
T::reduce_popcount_and(a, b)
}
#[inline(always)]
pub fn reduce_popcount_or<T: SimdOps>(a: &[T], b: &[T]) -> Result<usize, SimdError> {
T::reduce_popcount_or(a, b)
}
#[inline(always)]
pub fn reduce_popcount_xor<T: SimdOps>(a: &[T], b: &[T]) -> Result<usize, SimdError> {
T::reduce_popcount_xor(a, b)
}