Skip to main content

hermes_simd/dispatch/
masked.rs

1//! Runtime-dispatched masked SIMD operations.
2
3use hermes_simd_core::arch::SimdArch;
4use hermes_simd_core::kernel::SimdKernel;
5use hermes_simd_core::{view::SimdError, Scalar as ScalarTrait};
6use hermes_simd_macros::runtime_dispatch;
7
8// ---------------------------------------------------------------------------
9// Internal generic kernel implementations
10// ---------------------------------------------------------------------------
11
12/// Generic masked sum using `leading_k_mask` for the scalar tail.
13#[inline]
14unsafe fn masked_sum_impl<T, Arch>(data: &[T], bool_mask: &[bool]) -> T
15where
16    T: ScalarTrait,
17    Arch: SimdArch + SimdKernel<T>,
18{
19    assert_eq!(
20        data.len(),
21        bool_mask.len(),
22        "data and mask lengths must match"
23    );
24    let len = data.len();
25    let lane_count = Arch::LANE_COUNT;
26    let simd_len = (len / lane_count) * lane_count;
27
28    let mut total = T::ZERO;
29    let mut i = 0usize;
30
31    while i < simd_len {
32        let v = Arch::load_unaligned(data.as_ptr().add(i));
33        let msk = Arch::mask_from_bools(&bool_mask[i..i + lane_count]);
34        total += Arch::masked_sum_reduce(v, msk);
35        i += lane_count;
36    }
37
38    // Scalar tail
39    while i < len {
40        if bool_mask[i] {
41            total += data[i];
42        }
43        i += 1;
44    }
45
46    total
47}
48
49/// Generic masked elementwise add: `out[i] = if mask[i] { a[i] + b[i] } else { a[i] }`.
50#[inline]
51unsafe fn masked_add_impl<T, Arch>(
52    a: &[T],
53    b: &[T],
54    bool_mask: &[bool],
55    out: &mut [T],
56) -> Result<(), SimdError>
57where
58    T: ScalarTrait,
59    Arch: SimdArch + SimdKernel<T>,
60{
61    if a.len() != b.len() || a.len() != bool_mask.len() || a.len() > out.len() {
62        return Err(SimdError::LengthMismatch);
63    }
64    let len = a.len();
65    let lane_count = Arch::LANE_COUNT;
66    let simd_len = (len / lane_count) * lane_count;
67
68    let mut i = 0usize;
69    while i < simd_len {
70        let va = Arch::load_unaligned(a.as_ptr().add(i));
71        let vb = Arch::load_unaligned(b.as_ptr().add(i));
72        let msk = Arch::mask_from_bools(&bool_mask[i..i + lane_count]);
73        let src = va;
74        let result = Arch::masked_add(va, vb, msk, src);
75        Arch::store_unaligned(out.as_mut_ptr().add(i), result);
76        i += lane_count;
77    }
78
79    // Scalar tail
80    while i < len {
81        out[i] = if bool_mask[i] { a[i] + b[i] } else { a[i] };
82        i += 1;
83    }
84
85    Ok(())
86}
87
88/// Generic masked dot product: sum of `a[i] * b[i]` where `mask[i]`.
89#[inline]
90unsafe fn masked_dot_impl<T, Arch>(a: &[T], b: &[T], bool_mask: &[bool]) -> Result<T, SimdError>
91where
92    T: ScalarTrait,
93    Arch: SimdArch + SimdKernel<T>,
94{
95    if a.len() != b.len() || a.len() != bool_mask.len() {
96        return Err(SimdError::LengthMismatch);
97    }
98    let len = a.len();
99    let lane_count = Arch::LANE_COUNT;
100    let simd_len = (len / lane_count) * lane_count;
101
102    let mut acc = Arch::zero();
103    let mut i = 0usize;
104
105    while i < simd_len {
106        let va = Arch::load_unaligned(a.as_ptr().add(i));
107        let vb = Arch::load_unaligned(b.as_ptr().add(i));
108        let msk = Arch::mask_from_bools(&bool_mask[i..i + lane_count]);
109        acc = Arch::masked_fmadd(va, vb, acc, msk);
110        i += lane_count;
111    }
112
113    let mut total = Arch::sum_reduce(acc);
114
115    // Scalar tail
116    while i < len {
117        if bool_mask[i] {
118            total += a[i] * b[i];
119        }
120        i += 1;
121    }
122
123    Ok(total)
124}
125
126// ---------------------------------------------------------------------------
127// Public dispatcher functions
128// ---------------------------------------------------------------------------
129
130#[runtime_dispatch(avx512f, avx2, neon, scalar)]
131pub(super) fn dispatch_masked_sum_kernel<T, A>(data: &[T], mask: &[bool]) -> T
132where
133    T: ScalarTrait,
134    A: SimdArch + SimdKernel<T>,
135{
136    unsafe { masked_sum_impl::<T, A>(data, mask) }
137}
138
139#[runtime_dispatch(avx512f, avx2, neon, scalar)]
140pub(super) fn dispatch_masked_dot_kernel<T, A>(
141    a: &[T],
142    b: &[T],
143    mask: &[bool],
144) -> Result<T, SimdError>
145where
146    T: ScalarTrait,
147    A: SimdArch + SimdKernel<T>,
148{
149    unsafe { masked_dot_impl::<T, A>(a, b, mask) }
150}
151
152#[runtime_dispatch(avx512f, avx2, neon, scalar)]
153pub(super) fn dispatch_masked_add_kernel<T, A>(
154    a: &[T],
155    b: &[T],
156    mask: &[bool],
157    out: &mut [T],
158) -> Result<(), SimdError>
159where
160    T: ScalarTrait,
161    A: SimdArch + SimdKernel<T>,
162{
163    unsafe { masked_add_impl::<T, A>(a, b, mask, out) }
164}