tract_linalg/x86_64/
exp.rs1routine_ew_rust!(x86_64;
5 f32,
6 x86_64_fma_exp_f32_32n,
7 32,
8 8,
9 #[inline(never)]
10 fn run(buf: &mut [f32], _: ()) {
11 debug_assert!(buf.len() % Self::nr() == 0);
12 debug_assert!(buf.as_ptr() as usize % Self::alignment_bytes() == 0);
13 unsafe { x86_64_fma_exp_f32_32n_run(buf) }
14 },
15 func(Exp),
16 isa(X86_64Avx2, X86_64Fma)
17);
18
19#[cfg(target_arch = "x86_64")]
20#[target_feature(enable = "avx2,fma")]
21unsafe fn x86_64_fma_exp_f32_32n_run(buf: &mut [f32]) {
22 use crate::generic::exp::{HIGH, LN2_HI, LN2_LO, LOG2E, LOW, POLY, SCALE_BIAS};
23 use std::arch::x86_64::*;
24 #[inline(always)]
25 unsafe fn exp8(x: __m256) -> __m256 {
26 unsafe {
27 let high = _mm256_set1_ps(HIGH);
30 let low = _mm256_set1_ps(LOW);
31 let x = _mm256_blendv_ps(x, high, _mm256_cmp_ps::<_CMP_GT_OQ>(x, high));
32 let x = _mm256_blendv_ps(x, low, _mm256_cmp_ps::<_CMP_LT_OQ>(x, low));
33 let k = _mm256_cvtps_epi32(_mm256_mul_ps(x, _mm256_set1_ps(LOG2E)));
34 let kf = _mm256_cvtepi32_ps(k);
35 let mut r = _mm256_fnmadd_ps(kf, _mm256_set1_ps(LN2_HI), x);
36 r = _mm256_fnmadd_ps(kf, _mm256_set1_ps(LN2_LO), r);
37 let mut q = _mm256_set1_ps(POLY[0]);
38 for c in &POLY[1..] {
39 q = _mm256_fmadd_ps(q, r, _mm256_set1_ps(*c));
40 }
41 let bias = _mm256_set1_epi32(SCALE_BIAS);
42 let half = _mm256_srai_epi32::<1>(k);
43 let rest = _mm256_sub_epi32(k, half);
44 let scale = |k| _mm256_castsi256_ps(_mm256_slli_epi32::<23>(_mm256_add_epi32(k, bias)));
45 _mm256_mul_ps(_mm256_mul_ps(q, scale(half)), scale(rest))
46 }
47 }
48 unsafe {
49 let p = buf.as_mut_ptr();
50 for i in (0..buf.len()).step_by(8) {
51 _mm256_store_ps(p.add(i), exp8(_mm256_load_ps(p.add(i))));
52 }
53 }
54}
55
56routine_ew_rust!(x86_64;
57 f32,
58 x86_64_avx512_exp_f32_64n,
59 64,
60 16,
61 #[inline(never)]
62 fn run(buf: &mut [f32], _: ()) {
63 debug_assert!(buf.len() % Self::nr() == 0);
64 debug_assert!(buf.as_ptr() as usize % Self::alignment_bytes() == 0);
65 unsafe { x86_64_avx512_exp_f32_64n_run(buf) }
66 },
67 func(Exp),
68 isa(X86_64Avx512f)
69);
70
71#[cfg(target_arch = "x86_64")]
72#[target_feature(enable = "avx512f")]
73unsafe fn x86_64_avx512_exp_f32_64n_run(buf: &mut [f32]) {
74 use crate::generic::exp::{HIGH, LN2_HI, LN2_LO, LOG2E, LOW, POLY, SCALE_BIAS};
75 use std::arch::x86_64::*;
76 #[inline(always)]
77 unsafe fn exp16(x: __m512) -> __m512 {
78 unsafe {
79 let high = _mm512_set1_ps(HIGH);
80 let low = _mm512_set1_ps(LOW);
81 let x = _mm512_mask_blend_ps(_mm512_cmp_ps_mask::<_CMP_GT_OQ>(x, high), x, high);
82 let x = _mm512_mask_blend_ps(_mm512_cmp_ps_mask::<_CMP_LT_OQ>(x, low), x, low);
83 let k = _mm512_cvtps_epi32(_mm512_mul_ps(x, _mm512_set1_ps(LOG2E)));
84 let kf = _mm512_cvtepi32_ps(k);
85 let mut r = _mm512_fnmadd_ps(kf, _mm512_set1_ps(LN2_HI), x);
86 r = _mm512_fnmadd_ps(kf, _mm512_set1_ps(LN2_LO), r);
87 let mut q = _mm512_set1_ps(POLY[0]);
88 for c in &POLY[1..] {
89 q = _mm512_fmadd_ps(q, r, _mm512_set1_ps(*c));
90 }
91 let bias = _mm512_set1_epi32(SCALE_BIAS);
92 let half = _mm512_srai_epi32::<1>(k);
93 let rest = _mm512_sub_epi32(k, half);
94 let scale = |k| _mm512_castsi512_ps(_mm512_slli_epi32::<23>(_mm512_add_epi32(k, bias)));
95 _mm512_mul_ps(_mm512_mul_ps(q, scale(half)), scale(rest))
96 }
97 }
98 unsafe {
99 let p = buf.as_mut_ptr();
100 for i in (0..buf.len()).step_by(16) {
101 _mm512_store_ps(p.add(i), exp16(_mm512_load_ps(p.add(i))));
102 }
103 }
104}