hermes_simd/dispatch/
masked.rs1use 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#[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 while i < len {
40 if bool_mask[i] {
41 total += data[i];
42 }
43 i += 1;
44 }
45
46 total
47}
48
49#[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 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#[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 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#[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}