Skip to main content

hermes_simd_core/view/
masked.rs

1use crate::align::Alignment;
2use crate::arch::SimdArch;
3use crate::execution::ExecutionMode;
4use crate::kernel::{SimdKernel, MAX_SIMD_LANES};
5use crate::mask::BitMask;
6use crate::scalar::Scalar;
7use crate::view::{SimdError, SimdView};
8
9impl<'a, T: 'a, Arch: SimdArch + SimdKernel<T>, Align: Alignment, Mode: ExecutionMode, Ref: 'a>
10    SimdView<'a, T, Arch, Align, Mode, Ref>
11where
12    T: Scalar,
13{
14    /// Add elementwise where mask is active, writing to `out`. Inactive lanes copy from `self`.
15    #[inline(always)]
16    pub fn masked_add<ORef, const N: usize>(
17        &self,
18        other: &SimdView<'_, T, Arch, Align, Mode, ORef>,
19        mask: &BitMask<N>,
20        out: &mut [T],
21    ) -> Result<(), SimdError>
22    where
23        ORef: 'a,
24    {
25        debug_assert_eq!(N, Arch::LANE_COUNT);
26        super::check_lengths_equal(self.len(), other.len())?;
27        super::check_output_length(self.len(), out.len())?;
28
29        let len = self.len();
30        let lane_count = Arch::LANE_COUNT;
31        let simd_len = (len / lane_count) * lane_count;
32
33        let mut ptr1 = self.as_slice().as_ptr();
34        let mut ptr2 = other.as_slice().as_ptr();
35        let mut ptr_out = out.as_mut_ptr();
36
37        unsafe {
38            let load = |p| {
39                if crate::align::is_aligned_for_arch::<Arch, Align>() {
40                    Arch::load_aligned(p)
41                } else {
42                    Arch::load_unaligned(p)
43                }
44            };
45
46            let store = |p, val| {
47                let is_out_aligned = crate::align::is_aligned_for_arch::<Arch, Align>()
48                    && (p as usize) % Align::ALIGN_BYTES == 0;
49
50                if is_out_aligned {
51                    Arch::store_aligned(p, val);
52                } else {
53                    Arch::store_unaligned(p, val);
54                }
55            };
56
57            let native_mask = mask.to_native_mask::<T, Arch>();
58
59            for _ in 0..(simd_len / lane_count) {
60                let v1 = load(ptr1);
61                let v2 = load(ptr2);
62                let res = Arch::masked_add(v1, v2, native_mask, v1);
63                store(ptr_out, res);
64
65                ptr1 = ptr1.add(lane_count);
66                ptr2 = ptr2.add(lane_count);
67                ptr_out = ptr_out.add(lane_count);
68            }
69        }
70
71        let s_slice = self.as_slice();
72        let o_slice = other.as_slice();
73        for i in simd_len..len {
74            let lane_idx = i - simd_len;
75            out[i] = if mask.is_lane_active(lane_idx) {
76                s_slice[i] + o_slice[i]
77            } else {
78                s_slice[i]
79            };
80        }
81
82        Ok(())
83    }
84
85    /// Multiply elementwise where mask is active, writing to `out`. Inactive lanes copy from `self`.
86    #[inline(always)]
87    pub fn masked_mul<ORef, const N: usize>(
88        &self,
89        other: &SimdView<'_, T, Arch, Align, Mode, ORef>,
90        mask: &BitMask<N>,
91        out: &mut [T],
92    ) -> Result<(), SimdError>
93    where
94        ORef: 'a,
95    {
96        debug_assert_eq!(N, Arch::LANE_COUNT);
97        super::check_lengths_equal(self.len(), other.len())?;
98        super::check_output_length(self.len(), out.len())?;
99
100        let len = self.len();
101        let lane_count = Arch::LANE_COUNT;
102        let simd_len = (len / lane_count) * lane_count;
103
104        let mut ptr1 = self.as_slice().as_ptr();
105        let mut ptr2 = other.as_slice().as_ptr();
106        let mut ptr_out = out.as_mut_ptr();
107
108        unsafe {
109            let load = |p| {
110                if crate::align::is_aligned_for_arch::<Arch, Align>() {
111                    Arch::load_aligned(p)
112                } else {
113                    Arch::load_unaligned(p)
114                }
115            };
116
117            let store = |p, val| {
118                let is_out_aligned = crate::align::is_aligned_for_arch::<Arch, Align>()
119                    && (p as usize) % Align::ALIGN_BYTES == 0;
120
121                if is_out_aligned {
122                    Arch::store_aligned(p, val);
123                } else {
124                    Arch::store_unaligned(p, val);
125                }
126            };
127
128            let native_mask = mask.to_native_mask::<T, Arch>();
129
130            for _ in 0..(simd_len / lane_count) {
131                let v1 = load(ptr1);
132                let v2 = load(ptr2);
133                let res = Arch::masked_mul(v1, v2, native_mask, v1);
134                store(ptr_out, res);
135
136                ptr1 = ptr1.add(lane_count);
137                ptr2 = ptr2.add(lane_count);
138                ptr_out = ptr_out.add(lane_count);
139            }
140        }
141
142        let s_slice = self.as_slice();
143        let o_slice = other.as_slice();
144        for i in simd_len..len {
145            let lane_idx = i - simd_len;
146            out[i] = if mask.is_lane_active(lane_idx) {
147                s_slice[i] * o_slice[i]
148            } else {
149                s_slice[i]
150            };
151        }
152
153        Ok(())
154    }
155
156    /// Fused multiply-add where mask is active: `(self * b) + c`, writing to `out`. Inactive lanes copy from `c`.
157    #[inline(always)]
158    pub fn masked_fmadd<ORef1, ORef2, const N: usize>(
159        &self,
160        b: &SimdView<'_, T, Arch, Align, Mode, ORef1>,
161        c: &SimdView<'_, T, Arch, Align, Mode, ORef2>,
162        mask: &BitMask<N>,
163        out: &mut [T],
164    ) -> Result<(), SimdError>
165    where
166        ORef1: 'a,
167        ORef2: 'a,
168    {
169        debug_assert_eq!(N, Arch::LANE_COUNT);
170        super::check_lengths_equal(self.len(), b.len())?;
171        super::check_lengths_equal(self.len(), c.len())?;
172        super::check_output_length(self.len(), out.len())?;
173
174        let len = self.len();
175        let lane_count = Arch::LANE_COUNT;
176        let simd_len = (len / lane_count) * lane_count;
177
178        let mut ptr_a = self.as_slice().as_ptr();
179        let mut ptr_b = b.as_slice().as_ptr();
180        let mut ptr_c = c.as_slice().as_ptr();
181        let mut ptr_out = out.as_mut_ptr();
182
183        unsafe {
184            let load = |p| {
185                if crate::align::is_aligned_for_arch::<Arch, Align>() {
186                    Arch::load_aligned(p)
187                } else {
188                    Arch::load_unaligned(p)
189                }
190            };
191
192            let store = |p, val| {
193                let is_out_aligned = crate::align::is_aligned_for_arch::<Arch, Align>()
194                    && (p as usize) % Align::ALIGN_BYTES == 0;
195
196                if is_out_aligned {
197                    Arch::store_aligned(p, val);
198                } else {
199                    Arch::store_unaligned(p, val);
200                }
201            };
202
203            let native_mask = mask.to_native_mask::<T, Arch>();
204
205            for _ in 0..(simd_len / lane_count) {
206                let va = load(ptr_a);
207                let vb = load(ptr_b);
208                let vc = load(ptr_c);
209                let res = Arch::masked_fmadd(va, vb, vc, native_mask);
210                store(ptr_out, res);
211
212                ptr_a = ptr_a.add(lane_count);
213                ptr_b = ptr_b.add(lane_count);
214                ptr_c = ptr_c.add(lane_count);
215                ptr_out = ptr_out.add(lane_count);
216            }
217        }
218
219        let a_slice = self.as_slice();
220        let b_slice = b.as_slice();
221        let c_slice = c.as_slice();
222        for i in simd_len..len {
223            let lane_idx = i - simd_len;
224            out[i] = if mask.is_lane_active(lane_idx) {
225                a_slice[i].scalar_fmadd(b_slice[i], c_slice[i])
226            } else {
227                c_slice[i]
228            };
229        }
230
231        Ok(())
232    }
233
234    /// Compress: pack elements where `mask` is active contiguously into `out`. Returns the number of elements written.
235    #[inline(always)]
236    pub fn compress<const N: usize>(
237        &self,
238        mask: &BitMask<N>,
239        out: &mut [T],
240    ) -> Result<usize, SimdError> {
241        debug_assert_eq!(N, Arch::LANE_COUNT);
242        super::check_output_length(self.len(), out.len())?;
243
244        let len = self.len();
245        let lane_count = Arch::LANE_COUNT;
246        let simd_len = (len / lane_count) * lane_count;
247
248        let mut ptr = self.as_slice().as_ptr();
249        let mut ptr_out = out.as_mut_ptr();
250        let mut total_written = 0;
251
252        unsafe {
253            let load = |p| {
254                if crate::align::is_aligned_for_arch::<Arch, Align>() {
255                    Arch::load_aligned(p)
256                } else {
257                    Arch::load_unaligned(p)
258                }
259            };
260
261            let native_mask = mask.to_native_mask::<T, Arch>();
262            // The same mask applies to every chunk, so its popcount is loop-invariant.
263            let pop = mask.popcount() as usize;
264
265            // Scratch for one compacted vector, hoisted out of the loop. The
266            // store writes `lane_count` lanes and the copy reads only `pop ≤
267            // lane_count`, so no lane is read before the store initializes it —
268            // `MaybeUninit` avoids re-zeroing a full `MAX_SIMD_LANES` buffer every
269            // chunk (the previous `[T::ZERO; 64]` per iteration). The compile-time
270            // `LANE_BOUND_CHECK` guarantees `lane_count ≤ MAX_SIMD_LANES`, so the
271            // store stays in bounds (referencing it forces the per-backend assert).
272            let _ = <Arch as SimdKernel<T>>::LANE_BOUND_CHECK;
273            let mut temp = [core::mem::MaybeUninit::<T>::uninit(); MAX_SIMD_LANES];
274
275            for _ in 0..(simd_len / lane_count) {
276                let v = load(ptr);
277                let compressed = Arch::compress(v, native_mask);
278
279                Arch::store_unaligned(temp.as_mut_ptr() as *mut T, compressed);
280                core::ptr::copy_nonoverlapping(temp.as_ptr() as *const T, ptr_out, pop);
281
282                ptr = ptr.add(lane_count);
283                ptr_out = ptr_out.add(pop);
284                total_written += pop;
285            }
286        }
287
288        let s_slice = self.as_slice();
289        for i in simd_len..len {
290            let lane_idx = i - simd_len;
291            if mask.is_lane_active(lane_idx) {
292                out[total_written] = s_slice[i];
293                total_written += 1;
294            }
295        }
296
297        Ok(total_written)
298    }
299
300    /// Expand: scatter the active elements of `self` into `out` at mask positions, filling inactive positions with `fill`.
301    #[inline(always)]
302    pub fn expand<ORef, const N: usize>(
303        &self,
304        mask: &BitMask<N>,
305        fill: &SimdView<'_, T, Arch, Align, Mode, ORef>,
306        out: &mut [T],
307    ) -> Result<(), SimdError>
308    where
309        ORef: 'a,
310    {
311        debug_assert_eq!(N, Arch::LANE_COUNT);
312        let out_len = out.len();
313        let lane_count = Arch::LANE_COUNT;
314        let simd_len = (out_len / lane_count) * lane_count;
315        let pop = mask.popcount() as usize;
316
317        // Calculate the number of source elements required from `self`
318        let num_simd_chunks = simd_len / lane_count;
319        let mut required_len = num_simd_chunks * pop;
320        for i in simd_len..out_len {
321            let lane_idx = i - simd_len;
322            if mask.is_lane_active(lane_idx) {
323                required_len += 1;
324            }
325        }
326
327        if self.len() < required_len {
328            return Err(SimdError::LengthMismatch);
329        }
330        if fill.len() < out_len {
331            return Err(SimdError::InsufficientOutputLength);
332        }
333
334        let src_len = self.len();
335        let mut ptr_src = self.as_slice().as_ptr();
336        let mut ptr_fill = fill.as_slice().as_ptr();
337        let mut ptr_out = out.as_mut_ptr();
338
339        unsafe {
340            let load = |p| {
341                if crate::align::is_aligned_for_arch::<Arch, Align>() {
342                    Arch::load_aligned(p)
343                } else {
344                    Arch::load_unaligned(p)
345                }
346            };
347
348            let store = |p, val| {
349                let is_out_aligned = crate::align::is_aligned_for_arch::<Arch, Align>()
350                    && (p as usize) % Align::ALIGN_BYTES == 0;
351
352                if is_out_aligned {
353                    Arch::store_aligned(p, val);
354                } else {
355                    Arch::store_unaligned(p, val);
356                }
357            };
358
359            let load_safe = |p: *const T, remaining: usize| {
360                if remaining >= lane_count {
361                    load(p)
362                } else {
363                    let mut temp = [T::ZERO; 64];
364                    core::ptr::copy_nonoverlapping(p, temp.as_mut_ptr(), remaining);
365                    load(temp.as_ptr())
366                }
367            };
368
369            let native_mask = mask.to_native_mask::<T, Arch>();
370
371            for chunk_idx in 0..num_simd_chunks {
372                let offset = chunk_idx * pop;
373                let remaining = src_len - offset;
374                let src_v = load_safe(ptr_src, remaining);
375                let fill_v = load(ptr_fill);
376                let res = Arch::expand(src_v, native_mask, fill_v);
377                store(ptr_out, res);
378
379                ptr_src = ptr_src.add(pop);
380                ptr_fill = ptr_fill.add(lane_count);
381                ptr_out = ptr_out.add(lane_count);
382            }
383        }
384
385        let s_slice = self.as_slice();
386        let f_slice = fill.as_slice();
387        let mut src_idx = num_simd_chunks * pop;
388        for i in simd_len..out_len {
389            let lane_idx = i - simd_len;
390            if mask.is_lane_active(lane_idx) {
391                out[i] = s_slice[src_idx];
392                src_idx += 1;
393            } else {
394                out[i] = f_slice[i];
395            }
396        }
397
398        Ok(())
399    }
400}