1pub 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
14pub type Bin32x4 = [u32; 4];
16pub 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
135pub 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
180pub 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
227pub 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
256pub 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 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
360pub 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
365pub 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
370pub 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
375pub 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
380pub 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
387pub 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
392pub 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
427pub 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
453pub 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
477pub 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
518pub 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
550pub 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
677pub 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
682pub 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
687pub fn sqrt_u32_lanes(a: &[u32]) -> alloc::vec::Vec<u32> {
689 map_unary_u32(a, sqrt_bin32_x4, sqrt_bits)
690}
691
692pub 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
697pub 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
702pub 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
707pub fn sqrt_u64_lanes(a: &[u64]) -> alloc::vec::Vec<u64> {
709 map_unary_u64(a, sqrt_bin64_x2, sqrt_bits)
710}
711
712pub 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}