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 #[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 #[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 #[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 #[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 let pop = mask.popcount() as usize;
264
265 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 #[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 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}