Skip to main content

hermes_simd/dispatch/
sparse.rs

1//! Runtime-dispatched sparse matrix-vector multiplication.
2
3use hermes_simd_core::arch::SimdArch;
4use hermes_simd_core::kernel::SimdKernel;
5use hermes_simd_core::scalar::Scalar;
6use hermes_simd_core::sparse::{
7    BlockedCoo, BlockedCooData, Csr, CsrData, DenseWithMask, DenseWithMaskData, SellP, SellPData,
8    SparseSpMv, SparseView, Validated, ValidatedData,
9};
10use hermes_simd_macros::runtime_dispatch;
11
12#[runtime_dispatch(avx512f, avx2, neon, scalar)]
13pub(super) fn dispatch_spmv_csr_kernel<T, A>(
14    data: ValidatedData<CsrData<'_, T>>,
15    x: &[T],
16    y: &mut [T],
17) where
18    T: Scalar,
19    A: SimdArch + SimdKernel<T>,
20{
21    SparseView::<T, Validated<Csr>, A>::from_validated_csr(data).spmv(x, y);
22}
23
24#[runtime_dispatch(avx512f, avx2, neon, scalar)]
25pub(super) fn dispatch_spmv_dense_masked_kernel<T, A>(
26    data: DenseWithMaskData<'_, T>,
27    x: &[T],
28    y: &mut [T],
29) where
30    T: Scalar,
31    A: SimdArch + SimdKernel<T>,
32{
33    SparseView::<T, DenseWithMask, A>::from_dense_with_mask(data).spmv(x, y);
34}
35
36#[runtime_dispatch(avx512f, avx2, neon, scalar)]
37pub(super) fn dispatch_spmv_sellp_kernel<T, const C: usize, A>(
38    data: ValidatedData<SellPData<'_, T, C>>,
39    x: &[T],
40    y: &mut [T],
41) where
42    T: Scalar,
43    A: SimdArch + SimdKernel<T>,
44{
45    SparseView::<T, Validated<SellP<C>>, A>::from_validated_sellp(data).spmv(x, y);
46}
47
48#[runtime_dispatch(avx512f, avx2, neon, scalar)]
49pub(super) fn dispatch_spmv_bcoo_kernel<T, const BM: usize, const BN: usize, A>(
50    data: ValidatedData<BlockedCooData<'_, T, BM, BN>>,
51    x: &[T],
52    y: &mut [T],
53) where
54    T: Scalar,
55    A: SimdArch + SimdKernel<T>,
56{
57    SparseView::<T, Validated<BlockedCoo<BM, BN>>, A>::from_validated_blocked_coo(data).spmv(x, y);
58}