Skip to main content

zenith_float_num/ieee_soft/
simd.rs

1//! Integer SIMD for software IEEE arrays. Lanes are `u32` / `u64` bit patterns.
2//! Hardware IEEE arithmetic is not used.
3
4/// Number of `u32` lanes in one integer SIMD vector (128-bit register).
5/// Binary64 uses `IEEE_SIMD_LANE_WIDTH / 2` `u64` lanes. Scalar fallback
6/// uses the same width so wrappers can size buffers without `cfg`.
7pub const IEEE_SIMD_LANE_WIDTH: usize = 4;
8
9use super::arith::{
10    add_bits, div_bits, fma_add, fma_bits, isqrt, mul_bits, normalize_mul, pack_finite, pack_mag,
11    round_rne, sqrt_bits, sub_bits, unpack, Class, Format, BIN32, BIN64,
12};
13
14/// Four binary32 bit patterns.
15pub type Bin32x4 = [u32; 4];
16/// Two binary64 bit patterns.
17pub type Bin64x2 = [u64; 2];
18
19#[derive(Clone, Copy)]
20struct U32x4([u32; 4]);
21
22impl U32x4 {
23    fn load(v: [u32; 4]) -> Self {
24        Self(v)
25    }
26
27    fn splat(x: u32) -> Self {
28        Self([x; 4])
29    }
30
31    fn to_array(self) -> [u32; 4] {
32        self.0
33    }
34
35    fn and(self, o: Self) -> Self {
36        self.bop(o, core::ops::BitAnd::bitand)
37    }
38
39    fn or(self, o: Self) -> Self {
40        self.bop(o, core::ops::BitOr::bitor)
41    }
42
43    fn xor(self, o: Self) -> Self {
44        self.bop(o, core::ops::BitXor::bitxor)
45    }
46
47    fn wrapping_add(self, o: Self) -> Self {
48        #[cfg(target_arch = "x86_64")]
49        unsafe {
50            use core::arch::x86_64::{__m128i, _mm_add_epi32, _mm_loadu_si128, _mm_storeu_si128};
51            let a = _mm_loadu_si128(self.0.as_ptr() as *const __m128i);
52            let b = _mm_loadu_si128(o.0.as_ptr() as *const __m128i);
53            let s = _mm_add_epi32(a, b);
54            let mut out = [0u32; 4];
55            _mm_storeu_si128(out.as_mut_ptr() as *mut __m128i, s);
56            return Self(out);
57        }
58        #[cfg(target_arch = "aarch64")]
59        unsafe {
60            use core::arch::aarch64::{uint32x4_t, vaddq_u32, vld1q_u32, vst1q_u32};
61            let a = vld1q_u32(self.0.as_ptr());
62            let b = vld1q_u32(o.0.as_ptr());
63            let s: uint32x4_t = vaddq_u32(a, b);
64            let mut out = [0u32; 4];
65            vst1q_u32(out.as_mut_ptr(), s);
66            return Self(out);
67        }
68        #[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))]
69        {
70            self.bop(o, u32::wrapping_add)
71        }
72    }
73
74    fn srli(self, n: u32) -> Self {
75        Self([self.0[0] >> n, self.0[1] >> n, self.0[2] >> n, self.0[3] >> n])
76    }
77
78    fn eq_mask(self, o: Self) -> bool {
79        self.0 == o.0
80    }
81
82    fn lane_eq(self, o: Self) -> [bool; 4] {
83        [
84            self.0[0] == o.0[0],
85            self.0[1] == o.0[1],
86            self.0[2] == o.0[2],
87            self.0[3] == o.0[3],
88        ]
89    }
90
91    fn bop(self, o: Self, f: fn(u32, u32) -> u32) -> Self {
92        Self([
93            f(self.0[0], o.0[0]),
94            f(self.0[1], o.0[1]),
95            f(self.0[2], o.0[2]),
96            f(self.0[3], o.0[3]),
97        ])
98    }
99}
100
101fn all_normal_bin32(bits: U32x4) -> bool {
102    let exp = bits.srli(23).and(U32x4::splat(0xff));
103    !exp.lane_eq(U32x4::splat(0)).iter().any(|&z| z)
104        && !exp.lane_eq(U32x4::splat(0xff)).iter().any(|&z| z)
105}
106
107fn mul_u32x4_integer(a: [u32; 4], b: [u32; 4]) -> [u64; 4] {
108    #[cfg(target_arch = "x86_64")]
109    unsafe {
110        use core::arch::x86_64::{
111            __m128i, _mm_cvtsi128_si64, _mm_loadu_si128, _mm_mul_epu32, _mm_srli_si128,
112        };
113        let va = _mm_loadu_si128(a.as_ptr() as *const __m128i);
114        let vb = _mm_loadu_si128(b.as_ptr() as *const __m128i);
115        let p_even = _mm_mul_epu32(va, vb);
116        let p_odd = _mm_mul_epu32(_mm_srli_si128(va, 4), _mm_srli_si128(vb, 4));
117        return [
118            _mm_cvtsi128_si64(p_even) as u64,
119            _mm_cvtsi128_si64(p_odd) as u64,
120            _mm_cvtsi128_si64(_mm_srli_si128(p_even, 8)) as u64,
121            _mm_cvtsi128_si64(_mm_srli_si128(p_odd, 8)) as u64,
122        ];
123    }
124    #[cfg(not(target_arch = "x86_64"))]
125    {
126        [
127            a[0] as u64 * b[0] as u64,
128            a[1] as u64 * b[1] as u64,
129            a[2] as u64 * b[2] as u64,
130            a[3] as u64 * b[3] as u64,
131        ]
132    }
133}
134
135/// Software IEEE add on four binary32 lanes. Integer SIMD on the equal-exp
136/// same-sign normal path; otherwise the scalar integer kernel.
137pub fn add_bin32_x4(a: Bin32x4, b: Bin32x4) -> Bin32x4 {
138    let va = U32x4::load(a);
139    let vb = U32x4::load(b);
140    if all_normal_bin32(va) && all_normal_bin32(vb) {
141        let sign_a = va.srli(31);
142        let sign_b = vb.srli(31);
143        let exp_a = va.srli(23).and(U32x4::splat(0xff));
144        let exp_b = vb.srli(23).and(U32x4::splat(0xff));
145        if sign_a.eq_mask(sign_b) && exp_a.eq_mask(exp_b) {
146            let hidden = U32x4::splat(1 << 23);
147            let frac = U32x4::splat(0x7f_ffff);
148            let sa = va.and(frac).or(hidden);
149            let sb = vb.and(frac).or(hidden);
150            let sa3 = U32x4([sa.0[0] << 3, sa.0[1] << 3, sa.0[2] << 3, sa.0[3] << 3]);
151            let sb3 = U32x4([sb.0[0] << 3, sb.0[1] << 3, sb.0[2] << 3, sb.0[3] << 3]);
152            let sum = sa3.wrapping_add(sb3);
153            let mut out = [0u32; 4];
154            for i in 0..4 {
155                let (exp, core) = round_rne(sum.0[i] as u128, exp_a.0[i] as i32, false, BIN32);
156                out[i] = pack_finite(sign_a.0[i] != 0, exp, core, BIN32) as u32;
157            }
158            debug_assert_bit_eq32(&out, a, b, add_bits);
159            return out;
160        }
161    }
162    [
163        add_bits(a[0] as u64, b[0] as u64, BIN32) as u32,
164        add_bits(a[1] as u64, b[1] as u64, BIN32) as u32,
165        add_bits(a[2] as u64, b[2] as u64, BIN32) as u32,
166        add_bits(a[3] as u64, b[3] as u64, BIN32) as u32,
167    ]
168}
169
170fn debug_assert_bit_eq32(got: &Bin32x4, a: Bin32x4, b: Bin32x4, op: fn(u64, u64, Format) -> u64) {
171    for i in 0..4 {
172        debug_assert_eq!(
173            got[i],
174            op(a[i] as u64, b[i] as u64, BIN32) as u32,
175            "SIMD lane bits must match the scalar integer kernel"
176        );
177    }
178}
179
180/// Software IEEE mul on four binary32 lanes. Integer SIMD for significand products
181/// when every lane is normal.
182pub fn mul_bin32_x4(a: Bin32x4, b: Bin32x4) -> Bin32x4 {
183    let va = U32x4::load(a);
184    let vb = U32x4::load(b);
185    if all_normal_bin32(va) && all_normal_bin32(vb) {
186        let sign = va.srli(31).xor(vb.srli(31));
187        let exp_a = va.srli(23).and(U32x4::splat(0xff));
188        let exp_b = vb.srli(23).and(U32x4::splat(0xff));
189        let frac = U32x4::splat(0x7f_ffff);
190        let hidden = U32x4::splat(1 << 23);
191        let sa = va.and(frac).or(hidden).to_array();
192        let sb = vb.and(frac).or(hidden).to_array();
193        let prod = mul_u32x4_integer(sa, sb);
194        let mut out = [0u32; 4];
195        for i in 0..4 {
196            let e = exp_a.0[i] as i32 + exp_b.0[i] as i32 - BIN32.bias;
197            out[i] = normalize_mul(sign.0[i] != 0, e, prod[i] as u128, BIN32) as u32;
198        }
199        debug_assert_bit_eq32(&out, a, b, mul_bits);
200        return out;
201    }
202    [
203        mul_bits(a[0] as u64, b[0] as u64, BIN32) as u32,
204        mul_bits(a[1] as u64, b[1] as u64, BIN32) as u32,
205        mul_bits(a[2] as u64, b[2] as u64, BIN32) as u32,
206        mul_bits(a[3] as u64, b[3] as u64, BIN32) as u32,
207    ]
208}
209
210fn add_u64x2_integer(a: [u64; 2], b: [u64; 2]) -> [u64; 2] {
211    #[cfg(target_arch = "x86_64")]
212    unsafe {
213        use core::arch::x86_64::{__m128i, _mm_add_epi64, _mm_loadu_si128, _mm_storeu_si128};
214        let va = _mm_loadu_si128(a.as_ptr() as *const __m128i);
215        let vb = _mm_loadu_si128(b.as_ptr() as *const __m128i);
216        let s = _mm_add_epi64(va, vb);
217        let mut out = [0u64; 2];
218        _mm_storeu_si128(out.as_mut_ptr() as *mut __m128i, s);
219        return out;
220    }
221    #[cfg(not(target_arch = "x86_64"))]
222    {
223        [a[0].wrapping_add(b[0]), a[1].wrapping_add(b[1])]
224    }
225}
226
227/// Software IEEE add on two binary64 lanes. Integer SIMD on the equal-exp
228/// same-sign normal path.
229pub fn add_bin64_x2(a: Bin64x2, b: Bin64x2) -> Bin64x2 {
230    let ua0 = unpack(a[0], BIN64);
231    let ua1 = unpack(a[1], BIN64);
232    let ub0 = unpack(b[0], BIN64);
233    let ub1 = unpack(b[1], BIN64);
234    if ua0.class == Class::Norm
235        && ua1.class == Class::Norm
236        && ub0.class == Class::Norm
237        && ub1.class == Class::Norm
238        && ua0.sign == ub0.sign
239        && ua1.sign == ub1.sign
240        && ua0.exp == ub0.exp
241        && ua1.exp == ub1.exp
242    {
243        let sa = [ua0.sig << 3, ua1.sig << 3];
244        let sb = [ub0.sig << 3, ub1.sig << 3];
245        let sum = add_u64x2_integer(sa, sb);
246        let (e0, c0) = round_rne(sum[0] as u128, ua0.exp, false, BIN64);
247        let (e1, c1) = round_rne(sum[1] as u128, ua1.exp, false, BIN64);
248        let out = [pack_finite(ua0.sign, e0, c0, BIN64), pack_finite(ua1.sign, e1, c1, BIN64)];
249        debug_assert_eq!(out[0], add_bits(a[0], b[0], BIN64));
250        debug_assert_eq!(out[1], add_bits(a[1], b[1], BIN64));
251        return out;
252    }
253    [add_bits(a[0], b[0], BIN64), add_bits(a[1], b[1], BIN64)]
254}
255
256/// Software IEEE mul on two binary64 lanes.
257pub fn mul_bin64_x2(a: Bin64x2, b: Bin64x2) -> Bin64x2 {
258    if let (Some(a0), Some(b0)) = (normal_sig64(a[0]), normal_sig64(b[0])) {
259        if let (Some(a1), Some(b1)) = (normal_sig64(a[1]), normal_sig64(b[1])) {
260            let p0 = mul_u64_integer(a0.sig, b0.sig);
261            let p1 = mul_u64_integer(a1.sig, b1.sig);
262            let out = [
263                normalize_mul(a0.sign ^ b0.sign, a0.exp + b0.exp - BIN64.bias, p0, BIN64),
264                normalize_mul(a1.sign ^ b1.sign, a1.exp + b1.exp - BIN64.bias, p1, BIN64),
265            ];
266            debug_assert_eq!(out[0], mul_bits(a[0], b[0], BIN64));
267            debug_assert_eq!(out[1], mul_bits(a[1], b[1], BIN64));
268            return out;
269        }
270    }
271    [mul_bits(a[0], b[0], BIN64), mul_bits(a[1], b[1], BIN64)]
272}
273
274struct Norm64 {
275    sign: bool,
276    exp: i32,
277    sig: u64,
278}
279
280fn normal_sig64(bits: u64) -> Option<Norm64> {
281    let u = unpack(bits, BIN64);
282    if u.class != Class::Norm {
283        return None;
284    }
285    Some(Norm64 {
286        sign: u.sign,
287        exp: u.exp,
288        sig: u.sig,
289    })
290}
291
292fn mul_u64_integer(a: u64, b: u64) -> u128 {
293    #[cfg(target_arch = "x86_64")]
294    unsafe {
295        use core::arch::x86_64::{_mm_cvtsi128_si64, _mm_mul_epu32, _mm_set_epi64x};
296        // 53-bit significands: split 32+21 and use integer `mul_epu32` on halves.
297        let alo = a as u32;
298        let ahi = (a >> 32) as u32;
299        let blo = b as u32;
300        let bhi = (b >> 32) as u32;
301        let p0 = _mm_mul_epu32(_mm_set_epi64x(0, alo as i64), _mm_set_epi64x(0, blo as i64));
302        let p1 = _mm_mul_epu32(_mm_set_epi64x(0, alo as i64), _mm_set_epi64x(0, bhi as i64));
303        let p2 = _mm_mul_epu32(_mm_set_epi64x(0, ahi as i64), _mm_set_epi64x(0, blo as i64));
304        let p3 = _mm_mul_epu32(_mm_set_epi64x(0, ahi as i64), _mm_set_epi64x(0, bhi as i64));
305        let lo = _mm_cvtsi128_si64(p0) as u128;
306        let m1 = _mm_cvtsi128_si64(p1) as u128;
307        let m2 = _mm_cvtsi128_si64(p2) as u128;
308        let hi = _mm_cvtsi128_si64(p3) as u128;
309        return lo + ((m1 + m2) << 32) + (hi << 64);
310    }
311    #[cfg(not(target_arch = "x86_64"))]
312    {
313        a as u128 * b as u128
314    }
315}
316
317fn map_pairs_u32(
318    a: &[u32],
319    b: &[u32],
320    chunk: fn(Bin32x4, Bin32x4) -> Bin32x4,
321    tail: fn(u64, u64, Format) -> u64,
322) -> alloc::vec::Vec<u32> {
323    debug_assert_eq!(a.len(), b.len());
324    let mut out = alloc::vec::Vec::with_capacity(a.len());
325    let mut i = 0;
326    while i + 4 <= a.len() {
327        let ca = [a[i], a[i + 1], a[i + 2], a[i + 3]];
328        let cb = [b[i], b[i + 1], b[i + 2], b[i + 3]];
329        out.extend_from_slice(&chunk(ca, cb));
330        i += 4;
331    }
332    while i < a.len() {
333        out.push(tail(a[i] as u64, b[i] as u64, BIN32) as u32);
334        i += 1;
335    }
336    out
337}
338
339fn map_pairs_u64(
340    a: &[u64],
341    b: &[u64],
342    chunk: fn(Bin64x2, Bin64x2) -> Bin64x2,
343    tail: fn(u64, u64, Format) -> u64,
344) -> alloc::vec::Vec<u64> {
345    debug_assert_eq!(a.len(), b.len());
346    let mut out = alloc::vec::Vec::with_capacity(a.len());
347    let mut i = 0;
348    while i + 2 <= a.len() {
349        let r = chunk([a[i], a[i + 1]], [b[i], b[i + 1]]);
350        out.extend_from_slice(&r);
351        i += 2;
352    }
353    while i < a.len() {
354        out.push(tail(a[i], b[i], BIN64));
355        i += 1;
356    }
357    out
358}
359
360/// Elementwise binary32 add over slices of equal length.
361pub fn add_u32_lanes(a: &[u32], b: &[u32]) -> alloc::vec::Vec<u32> {
362    map_pairs_u32(a, b, add_bin32_x4, add_bits)
363}
364
365/// Elementwise binary32 mul over slices of equal length.
366pub fn mul_u32_lanes(a: &[u32], b: &[u32]) -> alloc::vec::Vec<u32> {
367    map_pairs_u32(a, b, mul_bin32_x4, mul_bits)
368}
369
370/// Elementwise binary64 add over slices of equal length.
371pub fn add_u64_lanes(a: &[u64], b: &[u64]) -> alloc::vec::Vec<u64> {
372    map_pairs_u64(a, b, add_bin64_x2, add_bits)
373}
374
375/// Elementwise binary64 mul over slices of equal length.
376pub fn mul_u64_lanes(a: &[u64], b: &[u64]) -> alloc::vec::Vec<u64> {
377    map_pairs_u64(a, b, mul_bin64_x2, mul_bits)
378}
379
380/// Software IEEE sub: flip the sign bit then add.
381pub fn sub_bin32_x4(a: Bin32x4, b: Bin32x4) -> Bin32x4 {
382    let sign = U32x4::splat(0x8000_0000);
383    let nb = U32x4::load(b).xor(sign).to_array();
384    add_bin32_x4(a, nb)
385}
386
387/// Software IEEE sub on two binary64 lanes.
388pub fn sub_bin64_x2(a: Bin64x2, b: Bin64x2) -> Bin64x2 {
389    add_bin64_x2(a, [b[0] ^ BIN64.sign_mask(), b[1] ^ BIN64.sign_mask()])
390}
391
392/// Software IEEE div on four binary32 lanes. Integer SIMD unpack when every
393/// lane is normal; significand quotient is integer `/` per lane (SSE2 has no
394/// integer divide). Otherwise the scalar integer kernel.
395pub fn div_bin32_x4(a: Bin32x4, b: Bin32x4) -> Bin32x4 {
396    let va = U32x4::load(a);
397    let vb = U32x4::load(b);
398    if all_normal_bin32(va) && all_normal_bin32(vb) {
399        let sign = va.srli(31).xor(vb.srli(31));
400        let exp_a = va.srli(23).and(U32x4::splat(0xff));
401        let exp_b = vb.srli(23).and(U32x4::splat(0xff));
402        let frac = U32x4::splat(0x7f_ffff);
403        let hidden = U32x4::splat(1 << 23);
404        let sa = va.and(frac).or(hidden).to_array();
405        let sb = vb.and(frac).or(hidden).to_array();
406        let extra = BIN32.frac + 4;
407        let mut out = [0u32; 4];
408        for i in 0..4 {
409            let num = (sa[i] as u128) << extra;
410            let den = sb[i] as u128;
411            let q = num / den;
412            let r = num % den;
413            let scale = exp_a.0[i] as i32 - exp_b.0[i] as i32 - extra as i32;
414            out[i] = pack_mag(sign.0[i] != 0, scale, q, r != 0, BIN32) as u32;
415        }
416        debug_assert_bit_eq32(&out, a, b, div_bits);
417        return out;
418    }
419    [
420        div_bits(a[0] as u64, b[0] as u64, BIN32) as u32,
421        div_bits(a[1] as u64, b[1] as u64, BIN32) as u32,
422        div_bits(a[2] as u64, b[2] as u64, BIN32) as u32,
423        div_bits(a[3] as u64, b[3] as u64, BIN32) as u32,
424    ]
425}
426
427/// Software IEEE div on two binary64 lanes.
428pub fn div_bin64_x2(a: Bin64x2, b: Bin64x2) -> Bin64x2 {
429    if let (Some(a0), Some(b0)) = (normal_sig64(a[0]), normal_sig64(b[0])) {
430        if let (Some(a1), Some(b1)) = (normal_sig64(a[1]), normal_sig64(b[1])) {
431            let extra = BIN64.frac + 4;
432            let out = [
433                div_normal_lane(a0, b0, extra, BIN64),
434                div_normal_lane(a1, b1, extra, BIN64),
435            ];
436            debug_assert_eq!(out[0], div_bits(a[0], b[0], BIN64));
437            debug_assert_eq!(out[1], div_bits(a[1], b[1], BIN64));
438            return out;
439        }
440    }
441    [div_bits(a[0], b[0], BIN64), div_bits(a[1], b[1], BIN64)]
442}
443
444fn div_normal_lane(a: Norm64, b: Norm64, extra: u32, f: Format) -> u64 {
445    let num = (a.sig as u128) << extra;
446    let den = b.sig as u128;
447    let q = num / den;
448    let r = num % den;
449    let scale = a.exp - b.exp - extra as i32;
450    pack_mag(a.sign ^ b.sign, scale, q, r != 0, f)
451}
452
453/// Software IEEE sqrt on four binary32 lanes. Integer SIMD unpack when every
454/// lane is a non-negative normal; `isqrt` per lane. Otherwise the scalar kernel.
455pub fn sqrt_bin32_x4(a: Bin32x4) -> Bin32x4 {
456    let va = U32x4::load(a);
457    if all_normal_bin32(va) && va.srli(31).eq_mask(U32x4::splat(0)) {
458        let exp = va.srli(23).and(U32x4::splat(0xff));
459        let frac = U32x4::splat(0x7f_ffff);
460        let hidden = U32x4::splat(1 << 23);
461        let sa = va.and(frac).or(hidden).to_array();
462        let mut out = [0u32; 4];
463        for i in 0..4 {
464            out[i] = sqrt_normal_lane(false, exp.0[i] as i32, sa[i] as u64, BIN32) as u32;
465        }
466        debug_assert_unary32(&out, a, sqrt_bits);
467        return out;
468    }
469    [
470        sqrt_bits(a[0] as u64, BIN32) as u32,
471        sqrt_bits(a[1] as u64, BIN32) as u32,
472        sqrt_bits(a[2] as u64, BIN32) as u32,
473        sqrt_bits(a[3] as u64, BIN32) as u32,
474    ]
475}
476
477/// Software IEEE sqrt on two binary64 lanes.
478pub fn sqrt_bin64_x2(a: Bin64x2) -> Bin64x2 {
479    if let (Some(a0), Some(a1)) = (normal_sig64(a[0]), normal_sig64(a[1])) {
480        if !a0.sign && !a1.sign {
481            let out = [
482                sqrt_normal_lane(a0.sign, a0.exp, a0.sig, BIN64),
483                sqrt_normal_lane(a1.sign, a1.exp, a1.sig, BIN64),
484            ];
485            debug_assert_eq!(out[0], sqrt_bits(a[0], BIN64));
486            debug_assert_eq!(out[1], sqrt_bits(a[1], BIN64));
487            return out;
488        }
489    }
490    [sqrt_bits(a[0], BIN64), sqrt_bits(a[1], BIN64)]
491}
492
493fn sqrt_normal_lane(_sign: bool, exp: i32, sig: u64, f: Format) -> u64 {
494    let mut s = sig as u128;
495    let mut exp2 = exp - f.bias - f.frac as i32;
496    if exp2 & 1 != 0 {
497        s <<= 1;
498        exp2 -= 1;
499    }
500    let extra = 64u32;
501    s <<= extra;
502    let root = isqrt(s);
503    let rem = s - root * root;
504    let scale = exp2 / 2 - extra as i32 / 2;
505    pack_mag(false, scale, root, rem != 0, f)
506}
507
508fn debug_assert_unary32(got: &Bin32x4, a: Bin32x4, op: fn(u64, Format) -> u64) {
509    for i in 0..4 {
510        debug_assert_eq!(
511            got[i],
512            op(a[i] as u64, BIN32) as u32,
513            "SIMD lane bits must match the scalar integer kernel"
514        );
515    }
516}
517
518/// Software IEEE FMA \(a\cdot b + c\) on four binary32 lanes. Integer SIMD
519/// significand products when every lane is normal; then the scalar `fma_add`.
520pub fn fma_bin32_x4(a: Bin32x4, b: Bin32x4, c: Bin32x4) -> Bin32x4 {
521    let va = U32x4::load(a);
522    let vb = U32x4::load(b);
523    let vc = U32x4::load(c);
524    if all_normal_bin32(va) && all_normal_bin32(vb) && all_normal_bin32(vc) {
525        let sign = va.srli(31).xor(vb.srli(31));
526        let exp_a = va.srli(23).and(U32x4::splat(0xff));
527        let exp_b = vb.srli(23).and(U32x4::splat(0xff));
528        let frac = U32x4::splat(0x7f_ffff);
529        let hidden = U32x4::splat(1 << 23);
530        let sa = va.and(frac).or(hidden).to_array();
531        let sb = vb.and(frac).or(hidden).to_array();
532        let prod = mul_u32x4_integer(sa, sb);
533        let mut out = [0u32; 4];
534        for i in 0..4 {
535            let pe = exp_a.0[i] as i32 + exp_b.0[i] as i32 - BIN32.bias;
536            let uc = unpack(c[i] as u64, BIN32);
537            out[i] = fma_add(sign.0[i] != 0, pe, prod[i] as u128, uc, BIN32) as u32;
538        }
539        debug_assert_fma32(&out, a, b, c);
540        return out;
541    }
542    [
543        fma_bits(a[0] as u64, b[0] as u64, c[0] as u64, BIN32) as u32,
544        fma_bits(a[1] as u64, b[1] as u64, c[1] as u64, BIN32) as u32,
545        fma_bits(a[2] as u64, b[2] as u64, c[2] as u64, BIN32) as u32,
546        fma_bits(a[3] as u64, b[3] as u64, c[3] as u64, BIN32) as u32,
547    ]
548}
549
550/// Software IEEE FMA on two binary64 lanes.
551pub fn fma_bin64_x2(a: Bin64x2, b: Bin64x2, c: Bin64x2) -> Bin64x2 {
552    if let (Some(a0), Some(b0)) = (normal_sig64(a[0]), normal_sig64(b[0])) {
553        if let (Some(a1), Some(b1)) = (normal_sig64(a[1]), normal_sig64(b[1])) {
554            let u0 = unpack(c[0], BIN64);
555            let u1 = unpack(c[1], BIN64);
556            if u0.class == Class::Norm && u1.class == Class::Norm {
557                let p0 = mul_u64_integer(a0.sig, b0.sig);
558                let p1 = mul_u64_integer(a1.sig, b1.sig);
559                let out = [
560                    fma_add(
561                        a0.sign ^ b0.sign,
562                        a0.exp + b0.exp - BIN64.bias,
563                        p0,
564                        u0,
565                        BIN64,
566                    ),
567                    fma_add(
568                        a1.sign ^ b1.sign,
569                        a1.exp + b1.exp - BIN64.bias,
570                        p1,
571                        u1,
572                        BIN64,
573                    ),
574                ];
575                debug_assert_eq!(out[0], fma_bits(a[0], b[0], c[0], BIN64));
576                debug_assert_eq!(out[1], fma_bits(a[1], b[1], c[1], BIN64));
577                return out;
578            }
579        }
580    }
581    [
582        fma_bits(a[0], b[0], c[0], BIN64),
583        fma_bits(a[1], b[1], c[1], BIN64),
584    ]
585}
586
587fn debug_assert_fma32(got: &Bin32x4, a: Bin32x4, b: Bin32x4, c: Bin32x4) {
588    for i in 0..4 {
589        debug_assert_eq!(
590            got[i],
591            fma_bits(a[i] as u64, b[i] as u64, c[i] as u64, BIN32) as u32,
592            "SIMD lane bits must match the scalar integer kernel"
593        );
594    }
595}
596
597fn map_unary_u32(a: &[u32], chunk: fn(Bin32x4) -> Bin32x4, tail: fn(u64, Format) -> u64) -> alloc::vec::Vec<u32> {
598    let mut out = alloc::vec::Vec::with_capacity(a.len());
599    let mut i = 0;
600    while i + 4 <= a.len() {
601        out.extend_from_slice(&chunk([a[i], a[i + 1], a[i + 2], a[i + 3]]));
602        i += 4;
603    }
604    while i < a.len() {
605        out.push(tail(a[i] as u64, BIN32) as u32);
606        i += 1;
607    }
608    out
609}
610
611fn map_unary_u64(a: &[u64], chunk: fn(Bin64x2) -> Bin64x2, tail: fn(u64, Format) -> u64) -> alloc::vec::Vec<u64> {
612    let mut out = alloc::vec::Vec::with_capacity(a.len());
613    let mut i = 0;
614    while i + 2 <= a.len() {
615        out.extend_from_slice(&chunk([a[i], a[i + 1]]));
616        i += 2;
617    }
618    while i < a.len() {
619        out.push(tail(a[i], BIN64));
620        i += 1;
621    }
622    out
623}
624
625fn map_triples_u32(
626    a: &[u32],
627    b: &[u32],
628    c: &[u32],
629    chunk: fn(Bin32x4, Bin32x4, Bin32x4) -> Bin32x4,
630    tail: fn(u64, u64, u64, Format) -> u64,
631) -> alloc::vec::Vec<u32> {
632    debug_assert_eq!(a.len(), b.len());
633    debug_assert_eq!(a.len(), c.len());
634    let mut out = alloc::vec::Vec::with_capacity(a.len());
635    let mut i = 0;
636    while i + 4 <= a.len() {
637        let ca = [a[i], a[i + 1], a[i + 2], a[i + 3]];
638        let cb = [b[i], b[i + 1], b[i + 2], b[i + 3]];
639        let cc = [c[i], c[i + 1], c[i + 2], c[i + 3]];
640        out.extend_from_slice(&chunk(ca, cb, cc));
641        i += 4;
642    }
643    while i < a.len() {
644        out.push(tail(a[i] as u64, b[i] as u64, c[i] as u64, BIN32) as u32);
645        i += 1;
646    }
647    out
648}
649
650fn map_triples_u64(
651    a: &[u64],
652    b: &[u64],
653    c: &[u64],
654    chunk: fn(Bin64x2, Bin64x2, Bin64x2) -> Bin64x2,
655    tail: fn(u64, u64, u64, Format) -> u64,
656) -> alloc::vec::Vec<u64> {
657    debug_assert_eq!(a.len(), b.len());
658    debug_assert_eq!(a.len(), c.len());
659    let mut out = alloc::vec::Vec::with_capacity(a.len());
660    let mut i = 0;
661    while i + 2 <= a.len() {
662        let r = chunk(
663            [a[i], a[i + 1]],
664            [b[i], b[i + 1]],
665            [c[i], c[i + 1]],
666        );
667        out.extend_from_slice(&r);
668        i += 2;
669    }
670    while i < a.len() {
671        out.push(tail(a[i], b[i], c[i], BIN64));
672        i += 1;
673    }
674    out
675}
676
677/// Elementwise binary32 sub over slices of equal length.
678pub fn sub_u32_lanes(a: &[u32], b: &[u32]) -> alloc::vec::Vec<u32> {
679    map_pairs_u32(a, b, sub_bin32_x4, sub_bits)
680}
681
682/// Elementwise binary32 div over slices of equal length.
683pub fn div_u32_lanes(a: &[u32], b: &[u32]) -> alloc::vec::Vec<u32> {
684    map_pairs_u32(a, b, div_bin32_x4, div_bits)
685}
686
687/// Elementwise binary32 sqrt.
688pub fn sqrt_u32_lanes(a: &[u32]) -> alloc::vec::Vec<u32> {
689    map_unary_u32(a, sqrt_bin32_x4, sqrt_bits)
690}
691
692/// Elementwise binary32 FMA.
693pub fn fma_u32_lanes(a: &[u32], b: &[u32], c: &[u32]) -> alloc::vec::Vec<u32> {
694    map_triples_u32(a, b, c, fma_bin32_x4, fma_bits)
695}
696
697/// Elementwise binary64 sub over slices of equal length.
698pub fn sub_u64_lanes(a: &[u64], b: &[u64]) -> alloc::vec::Vec<u64> {
699    map_pairs_u64(a, b, sub_bin64_x2, sub_bits)
700}
701
702/// Elementwise binary64 div over slices of equal length.
703pub fn div_u64_lanes(a: &[u64], b: &[u64]) -> alloc::vec::Vec<u64> {
704    map_pairs_u64(a, b, div_bin64_x2, div_bits)
705}
706
707/// Elementwise binary64 sqrt.
708pub fn sqrt_u64_lanes(a: &[u64]) -> alloc::vec::Vec<u64> {
709    map_unary_u64(a, sqrt_bin64_x2, sqrt_bits)
710}
711
712/// Elementwise binary64 FMA.
713pub fn fma_u64_lanes(a: &[u64], b: &[u64], c: &[u64]) -> alloc::vec::Vec<u64> {
714    map_triples_u64(a, b, c, fma_bin64_x2, fma_bits)
715}
716
717#[cfg(test)]
718mod tests {
719    use super::*;
720    use crate::ieee_soft::Ieee32;
721
722    #[test]
723    fn bin32_x4_add_matches_scalar() {
724        let one = Ieee32::from_i32(1).to_bits();
725        let two = Ieee32::from_i32(2).to_bits();
726        let a = [one, two, one, two];
727        let b = [one, two, two, one];
728        let got = add_bin32_x4(a, b);
729        for i in 0..4 {
730            assert_eq!(got[i], add_bits(a[i] as u64, b[i] as u64, BIN32) as u32);
731        }
732        assert_eq!(IEEE_SIMD_LANE_WIDTH, 4);
733        assert_eq!(got[0], Ieee32::from_i32(2).to_bits());
734        assert_eq!(got[1], Ieee32::from_i32(4).to_bits());
735    }
736
737    #[test]
738    fn bin32_x4_mul_half() {
739        let two = Ieee32::from_i32(2).to_bits();
740        let half = 0x3F00_0000u32;
741        let a = [two, two, two, two];
742        let b = [half, half, half, half];
743        let got = mul_bin32_x4(a, b);
744        let one = Ieee32::from_i32(1).to_bits();
745        assert_eq!(got, [one, one, one, one]);
746    }
747
748    #[test]
749    fn bin32_x4_specials_match_scalar() {
750        let one = Ieee32::from_i32(1).to_bits();
751        let inf = 0x7F80_0000u32;
752        let a = [one, inf, 0, one];
753        let b = [one, one, one, 0];
754        let got = add_bin32_x4(a, b);
755        for i in 0..4 {
756            assert_eq!(got[i], add_bits(a[i] as u64, b[i] as u64, BIN32) as u32);
757        }
758    }
759
760    #[test]
761    fn simd_add_mul_bit_identical_to_scalar() {
762        let samples = [
763            0u32,
764            0x0000_0001,
765            0x0080_0000,
766            0x3F00_0000,
767            0x3F80_0000,
768            0x3F80_0001,
769            0x4000_0000,
770            0x7F7F_FFFF,
771            0x7F80_0000,
772            0x7FC0_0000,
773            0x8000_0000,
774            0xBF80_0000,
775            Ieee32::from_i32(-3).to_bits(),
776            Ieee32::from_i32(7).to_bits(),
777        ];
778        for (i, &ai) in samples.iter().enumerate() {
779            for (j, &bj) in samples.iter().enumerate() {
780                let a = [
781                    ai,
782                    samples[(i + 1) % samples.len()],
783                    samples[(i + 2) % samples.len()],
784                    samples[(i + 3) % samples.len()],
785                ];
786                let b = [
787                    bj,
788                    samples[(j + 1) % samples.len()],
789                    samples[(j + 2) % samples.len()],
790                    samples[(j + 3) % samples.len()],
791                ];
792                let add = add_bin32_x4(a, b);
793                let mul = mul_bin32_x4(a, b);
794                for k in 0..4 {
795                    assert_eq!(add[k], add_bits(a[k] as u64, b[k] as u64, BIN32) as u32);
796                    assert_eq!(mul[k], mul_bits(a[k] as u64, b[k] as u64, BIN32) as u32);
797                }
798            }
799        }
800        let one = crate::ieee_soft::Ieee64::from_i32(1).to_bits();
801        let two = crate::ieee_soft::Ieee64::from_i32(2).to_bits();
802        let add64 = add_bin64_x2([one, two], [two, one]);
803        assert_eq!(add64[0], add_bits(one, two, BIN64));
804        assert_eq!(add64[1], add_bits(two, one, BIN64));
805        let mul64 = mul_bin64_x2([two, two], [one, one]);
806        assert_eq!(mul64[0], mul_bits(two, one, BIN64));
807        assert_eq!(mul64[1], mul_bits(two, one, BIN64));
808        let four = crate::ieee_soft::Ieee64::from_i32(4).to_bits();
809        let div64 = div_bin64_x2([four, two], [two, one]);
810        assert_eq!(div64[0], div_bits(four, two, BIN64));
811        assert_eq!(div64[1], div_bits(two, one, BIN64));
812        let sq64 = sqrt_bin64_x2([four, four]);
813        assert_eq!(sq64[0], sqrt_bits(four, BIN64));
814        assert_eq!(sq64[1], two);
815        let fma64 = fma_bin64_x2([two, two], [two, one], [one, one]);
816        assert_eq!(fma64[0], fma_bits(two, two, one, BIN64));
817        assert_eq!(fma64[1], fma_bits(two, one, one, BIN64));
818    }
819
820    #[test]
821    fn simd_div_sqrt_fma_match_scalar_samples() {
822        let samples = [
823            0u32,
824            0x0000_0001,
825            0x0080_0000,
826            0x3F00_0000,
827            0x3F80_0000,
828            0x3F80_0001,
829            0x4000_0000,
830            0x7F7F_FFFF,
831            0x7F80_0000,
832            0x7FC0_0000,
833            0x8000_0000,
834            0xBF80_0000,
835            Ieee32::from_i32(-3).to_bits(),
836            Ieee32::from_i32(7).to_bits(),
837        ];
838        for (i, &ai) in samples.iter().enumerate() {
839            for (j, &bj) in samples.iter().enumerate() {
840                let a = [
841                    ai,
842                    samples[(i + 1) % samples.len()],
843                    samples[(i + 2) % samples.len()],
844                    samples[(i + 3) % samples.len()],
845                ];
846                let b = [
847                    bj,
848                    samples[(j + 1) % samples.len()],
849                    samples[(j + 2) % samples.len()],
850                    samples[(j + 3) % samples.len()],
851                ];
852                let div = div_bin32_x4(a, b);
853                let sub = sub_bin32_x4(a, b);
854                let sq = sqrt_bin32_x4(a);
855                let fma = fma_bin32_x4(a, b, a);
856                for k in 0..4 {
857                    assert_eq!(div[k], div_bits(a[k] as u64, b[k] as u64, BIN32) as u32);
858                    assert_eq!(sub[k], sub_bits(a[k] as u64, b[k] as u64, BIN32) as u32);
859                    assert_eq!(sq[k], sqrt_bits(a[k] as u64, BIN32) as u32);
860                    assert_eq!(
861                        fma[k],
862                        fma_bits(a[k] as u64, b[k] as u64, a[k] as u64, BIN32) as u32
863                    );
864                }
865            }
866        }
867    }
868}