Skip to main content

tract_linalg/x86_64/
exp.rs

1//! f32 exp kernels: the [`crate::generic::exp`] reduction and fit, eight or sixteen lanes
2//! at a time.
3
4routine_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            // Clamped by selection rather than by min and max: a NaN is neither, and both
28            // instructions would answer with the bound instead of propagating it.
29            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}