1#[cfg(feature = "simd_support")]
12use core::simd::SimdElement;
13#[cfg(feature = "simd_support")]
14use core::simd::prelude::*;
15
16pub(crate) trait WideningMultiply<RHS = Self> {
17 type Output;
18
19 fn wmul(self, x: RHS) -> Self::Output;
20}
21
22macro_rules! wmul_impl {
23 ($ty:ty, $wide:ty, $shift:expr) => {
24 impl WideningMultiply for $ty {
25 type Output = ($ty, $ty);
26
27 #[inline(always)]
28 fn wmul(self, x: $ty) -> Self::Output {
29 let tmp = (self as $wide) * (x as $wide);
30 ((tmp >> $shift) as $ty, tmp as $ty)
31 }
32 }
33 };
34
35 ($(($ty:ident, $wide:ty),)+, $shift:expr) => {
37 $(
38 impl WideningMultiply for $ty {
39 type Output = ($ty, $ty);
40
41 #[inline(always)]
42 fn wmul(self, x: $ty) -> Self::Output {
43 let y: $wide = self.cast();
48 let x: $wide = x.cast();
49 let tmp = y * x;
50 let hi: $ty = (tmp >> Simd::splat($shift)).cast();
51 let lo: $ty = tmp.cast();
52 (hi, lo)
53 }
54 }
55 )+
56 };
57}
58impl WideningMultiply for u8 {
type Output = (u8, u8);
#[inline(always)]
fn wmul(self, x: u8) -> Self::Output {
let tmp = (self as u16) * (x as u16);
((tmp >> 8) as u8, tmp as u8)
}
}wmul_impl! { u8, u16, 8 }
59impl WideningMultiply for u16 {
type Output = (u16, u16);
#[inline(always)]
fn wmul(self, x: u16) -> Self::Output {
let tmp = (self as u32) * (x as u32);
((tmp >> 16) as u16, tmp as u16)
}
}wmul_impl! { u16, u32, 16 }
60impl WideningMultiply for u32 {
type Output = (u32, u32);
#[inline(always)]
fn wmul(self, x: u32) -> Self::Output {
let tmp = (self as u64) * (x as u64);
((tmp >> 32) as u32, tmp as u32)
}
}wmul_impl! { u32, u64, 32 }
61impl WideningMultiply for u64 {
type Output = (u64, u64);
#[inline(always)]
fn wmul(self, x: u64) -> Self::Output {
let tmp = (self as u128) * (x as u128);
((tmp >> 64) as u64, tmp as u64)
}
}wmul_impl! { u64, u128, 64 }
62
63macro_rules! wmul_impl_large {
70 ($ty:ty, $half:expr) => {
71 impl WideningMultiply for $ty {
72 type Output = ($ty, $ty);
73
74 #[inline(always)]
75 fn wmul(self, b: $ty) -> Self::Output {
76 const LOWER_MASK: $ty = !0 >> $half;
77 let mut low = (self & LOWER_MASK).wrapping_mul(b & LOWER_MASK);
78 let mut t = low >> $half;
79 low &= LOWER_MASK;
80 t += (self >> $half).wrapping_mul(b & LOWER_MASK);
81 low += (t & LOWER_MASK) << $half;
82 let mut high = t >> $half;
83 t = low >> $half;
84 low &= LOWER_MASK;
85 t += (b >> $half).wrapping_mul(self & LOWER_MASK);
86 low += (t & LOWER_MASK) << $half;
87 high += t >> $half;
88 high += (self >> $half).wrapping_mul(b >> $half);
89
90 (high, low)
91 }
92 }
93 };
94
95 (($($ty:ty,)+) $scalar:ty, $half:expr) => {
97 $(
98 impl WideningMultiply for $ty {
99 type Output = ($ty, $ty);
100
101 #[inline(always)]
102 fn wmul(self, b: $ty) -> Self::Output {
103 let lower_mask = <$ty>::splat(!0 >> $half);
105 let half = <$ty>::splat($half);
106 let mut low = (self & lower_mask) * (b & lower_mask);
107 let mut t = low >> half;
108 low &= lower_mask;
109 t += (self >> half) * (b & lower_mask);
110 low += (t & lower_mask) << half;
111 let mut high = t >> half;
112 t = low >> half;
113 low &= lower_mask;
114 t += (b >> half) * (self & lower_mask);
115 low += (t & lower_mask) << half;
116 high += t >> half;
117 high += (self >> half) * (b >> half);
118
119 (high, low)
120 }
121 }
122 )+
123 };
124}
125impl WideningMultiply for u128 {
type Output = (u128, u128);
#[inline(always)]
fn wmul(self, b: u128) -> Self::Output {
const LOWER_MASK: u128 = !0 >> 64;
let mut low = (self & LOWER_MASK).wrapping_mul(b & LOWER_MASK);
let mut t = low >> 64;
low &= LOWER_MASK;
t += (self >> 64).wrapping_mul(b & LOWER_MASK);
low += (t & LOWER_MASK) << 64;
let mut high = t >> 64;
t = low >> 64;
low &= LOWER_MASK;
t += (b >> 64).wrapping_mul(self & LOWER_MASK);
low += (t & LOWER_MASK) << 64;
high += t >> 64;
high += (self >> 64).wrapping_mul(b >> 64);
(high, low)
}
}wmul_impl_large! { u128, 64 }
126
127macro_rules! wmul_impl_usize {
128 ($ty:ty) => {
129 impl WideningMultiply for usize {
130 type Output = (usize, usize);
131
132 #[inline(always)]
133 fn wmul(self, x: usize) -> Self::Output {
134 let (high, low) = (self as $ty).wmul(x as $ty);
135 (high as usize, low as usize)
136 }
137 }
138 };
139}
140#[cfg(target_pointer_width = "16")]
141wmul_impl_usize! { u16 }
142#[cfg(target_pointer_width = "32")]
143wmul_impl_usize! { u32 }
144#[cfg(target_pointer_width = "64")]
145impl WideningMultiply for usize {
type Output = (usize, usize);
#[inline(always)]
fn wmul(self, x: usize) -> Self::Output {
let (high, low) = (self as u64).wmul(x as u64);
(high as usize, low as usize)
}
}wmul_impl_usize! { u64 }
146
147#[cfg(feature = "simd_support")]
148mod simd_wmul {
149 use super::*;
150 #[cfg(target_arch = "x86")]
151 use core::arch::x86::*;
152 #[cfg(target_arch = "x86_64")]
153 use core::arch::x86_64::*;
154
155 wmul_impl! {
156 (u8x4, u16x4),
157 (u8x8, u16x8),
158 (u8x16, u16x16),
159 (u8x32, u16x32),
160 (u8x64, Simd<u16, 64>),,
161 8
162 }
163
164 wmul_impl! { (u16x2, u32x2),, 16 }
165 wmul_impl! { (u16x4, u32x4),, 16 }
166 #[cfg(not(target_feature = "sse2"))]
167 wmul_impl! { (u16x8, u32x8),, 16 }
168 #[cfg(not(target_feature = "avx2"))]
169 wmul_impl! { (u16x16, u32x16),, 16 }
170 #[cfg(not(target_feature = "avx512bw"))]
171 wmul_impl! { (u16x32, Simd<u32, 32>),, 16 }
172
173 #[allow(unused_macros)]
176 macro_rules! wmul_impl_16 {
177 ($ty:ident, $mulhi:ident, $mullo:ident) => {
178 impl WideningMultiply for $ty {
179 type Output = ($ty, $ty);
180
181 #[inline(always)]
182 #[allow(clippy::undocumented_unsafe_blocks)]
183 fn wmul(self, x: $ty) -> Self::Output {
184 let hi = unsafe { $mulhi(self.into(), x.into()) }.into();
185 let lo = unsafe { $mullo(self.into(), x.into()) }.into();
186 (hi, lo)
187 }
188 }
189 };
190 }
191
192 #[cfg(target_feature = "sse2")]
193 wmul_impl_16! { u16x8, _mm_mulhi_epu16, _mm_mullo_epi16 }
194 #[cfg(target_feature = "avx2")]
195 wmul_impl_16! { u16x16, _mm256_mulhi_epu16, _mm256_mullo_epi16 }
196 #[cfg(target_feature = "avx512bw")]
197 wmul_impl_16! { u16x32, _mm512_mulhi_epu16, _mm512_mullo_epi16 }
198
199 wmul_impl! {
200 (u32x2, u64x2),
201 (u32x4, u64x4),
202 (u32x8, u64x8),
203 (u32x16, Simd<u64, 16>),,
204 32
205 }
206
207 wmul_impl_large! { (u64x2, u64x4, u64x8,) u64, 32 }
208}
209
210pub(crate) trait FloatSIMDUtils {
212 fn all_lt(self, other: Self) -> bool;
218 fn all_le(self, other: Self) -> bool;
219 fn all_finite(self) -> bool;
220
221 type Mask;
222 fn le_mask(self, other: Self) -> Self::Mask;
223
224 fn decrease_masked(self, mask: Self::Mask) -> Self;
229
230 type UInt;
233 fn cast_from_int(i: Self::UInt) -> Self;
234}
235
236#[cfg(test)]
237pub(crate) trait FloatSIMDScalarUtils: FloatSIMDUtils {
238 type Scalar;
239
240 fn replace(self, index: usize, new_value: Self::Scalar) -> Self;
241 fn extract_lane(self, index: usize) -> Self::Scalar;
242}
243
244pub(crate) trait FloatAsSIMD: Sized {
246 #[cfg(test)]
247 const LEN: usize = 1;
248
249 #[inline(always)]
250 fn splat(scalar: Self) -> Self {
251 scalar
252 }
253}
254
255pub(crate) trait IntAsSIMD: Sized {
256 #[inline(always)]
257 fn splat(scalar: Self) -> Self {
258 scalar
259 }
260}
261
262impl IntAsSIMD for u32 {}
263impl IntAsSIMD for u64 {}
264
265pub(crate) trait BoolAsSIMD: Sized {
266 fn all(self) -> bool;
267}
268
269impl BoolAsSIMD for bool {
270 #[inline(always)]
271 fn all(self) -> bool {
272 self
273 }
274}
275
276macro_rules! scalar_float_impl {
277 ($ty:ident, $uty:ident) => {
278 impl FloatSIMDUtils for $ty {
279 type Mask = bool;
280 type UInt = $uty;
281
282 #[inline(always)]
283 fn all_lt(self, other: Self) -> bool {
284 self < other
285 }
286
287 #[inline(always)]
288 fn all_le(self, other: Self) -> bool {
289 self <= other
290 }
291
292 #[inline(always)]
293 fn all_finite(self) -> bool {
294 self.is_finite()
295 }
296
297 #[inline(always)]
298 fn le_mask(self, other: Self) -> Self::Mask {
299 self <= other
300 }
301
302 #[inline(always)]
303 fn decrease_masked(self, mask: Self::Mask) -> Self {
304 debug_assert!(mask, "At least one lane must be set");
305 <$ty>::from_bits(self.to_bits() - 1)
306 }
307
308 #[inline]
309 fn cast_from_int(i: Self::UInt) -> Self {
310 i as $ty
311 }
312 }
313
314 #[cfg(test)]
315 impl FloatSIMDScalarUtils for $ty {
316 type Scalar = $ty;
317
318 #[inline]
319 fn replace(self, index: usize, new_value: Self::Scalar) -> Self {
320 debug_assert_eq!(index, 0);
321 new_value
322 }
323
324 #[inline]
325 fn extract_lane(self, index: usize) -> Self::Scalar {
326 debug_assert_eq!(index, 0);
327 self
328 }
329 }
330
331 impl FloatAsSIMD for $ty {}
332 };
333}
334
335impl FloatSIMDUtils for f32 {
type Mask = bool;
type UInt = u32;
#[inline(always)]
fn all_lt(self, other: Self) -> bool { self < other }
#[inline(always)]
fn all_le(self, other: Self) -> bool { self <= other }
#[inline(always)]
fn all_finite(self) -> bool { self.is_finite() }
#[inline(always)]
fn le_mask(self, other: Self) -> Self::Mask { self <= other }
#[inline(always)]
fn decrease_masked(self, mask: Self::Mask) -> Self {
if true {
if !mask {
{
::core::panicking::panic_fmt(format_args!("At least one lane must be set"));
}
};
};
<f32>::from_bits(self.to_bits() - 1)
}
#[inline]
fn cast_from_int(i: Self::UInt) -> Self { i as f32 }
}
impl FloatAsSIMD for f32 {}scalar_float_impl!(f32, u32);
336impl FloatSIMDUtils for f64 {
type Mask = bool;
type UInt = u64;
#[inline(always)]
fn all_lt(self, other: Self) -> bool { self < other }
#[inline(always)]
fn all_le(self, other: Self) -> bool { self <= other }
#[inline(always)]
fn all_finite(self) -> bool { self.is_finite() }
#[inline(always)]
fn le_mask(self, other: Self) -> Self::Mask { self <= other }
#[inline(always)]
fn decrease_masked(self, mask: Self::Mask) -> Self {
if true {
if !mask {
{
::core::panicking::panic_fmt(format_args!("At least one lane must be set"));
}
};
};
<f64>::from_bits(self.to_bits() - 1)
}
#[inline]
fn cast_from_int(i: Self::UInt) -> Self { i as f64 }
}
impl FloatAsSIMD for f64 {}scalar_float_impl!(f64, u64);
337
338#[cfg(feature = "simd_support")]
339macro_rules! simd_impl {
340 ($fty:ident, $uty:ident) => {
341 impl<const LANES: usize> FloatSIMDUtils for Simd<$fty, LANES> {
342 type Mask = Mask<<$fty as SimdElement>::Mask, LANES>;
343 type UInt = Simd<$uty, LANES>;
344
345 #[inline(always)]
346 fn all_lt(self, other: Self) -> bool {
347 self.simd_lt(other).all()
348 }
349
350 #[inline(always)]
351 fn all_le(self, other: Self) -> bool {
352 self.simd_le(other).all()
353 }
354
355 #[inline(always)]
356 fn all_finite(self) -> bool {
357 self.is_finite().all()
358 }
359
360 #[inline(always)]
361 fn le_mask(self, other: Self) -> Self::Mask {
362 self.simd_le(other)
363 }
364
365 #[inline(always)]
366 fn decrease_masked(self, mask: Self::Mask) -> Self {
367 debug_assert!(mask.any(), "At least one lane must be set");
374 Self::from_bits(self.to_bits() + mask.to_simd().cast())
375 }
376
377 #[inline]
378 fn cast_from_int(i: Self::UInt) -> Self {
379 i.cast()
380 }
381 }
382
383 #[cfg(test)]
384 impl<const LANES: usize> FloatSIMDScalarUtils for Simd<$fty, LANES> {
385 type Scalar = $fty;
386
387 #[inline]
388 fn replace(mut self, index: usize, new_value: Self::Scalar) -> Self {
389 self.as_mut_array()[index] = new_value;
390 self
391 }
392
393 #[inline]
394 fn extract_lane(self, index: usize) -> Self::Scalar {
395 self.as_array()[index]
396 }
397 }
398 };
399}
400
401#[cfg(feature = "simd_support")]
402simd_impl!(f32, u32);
403#[cfg(feature = "simd_support")]
404simd_impl!(f64, u64);
405
406#[cfg(test)]
407mod test {
408 use crate::distr::utils::FloatSIMDUtils;
409 #[cfg(feature = "simd_support")]
410 use std::simd::{Mask, Simd};
411
412 #[test]
413 fn decrease_masked() {
414 assert_eq!((-1.0 - f32::EPSILON).decrease_masked(true), -1.0);
415 assert_eq!((1.0 + f64::EPSILON).decrease_masked(true), 1.0);
416
417 #[cfg(feature = "simd_support")]
418 assert_eq!(
419 Simd::<f32, 4>::splat(1.0 + f32::EPSILON).decrease_masked(Mask::splat(true)),
420 Simd::splat(1.0)
421 );
422
423 #[cfg(feature = "simd_support")]
424 assert_eq!(
425 Simd::<f64, 2>::from_array([-1.0, -1.0 - f64::EPSILON])
426 .decrease_masked(Mask::from_array([false, true])),
427 Simd::splat(-1.0)
428 );
429 }
430
431 #[test]
432 fn decrease_masked_infinity() {
433 assert_eq!(f32::INFINITY.decrease_masked(true), f32::MAX);
434 assert_eq!((-f64::INFINITY).decrease_masked(true), -f64::MAX);
435
436 #[cfg(feature = "simd_support")]
437 assert_eq!(
438 Simd::<f32, 2>::splat(-f32::INFINITY).decrease_masked(Mask::splat(true)),
439 Simd::splat(-f32::MAX)
440 );
441 }
442}