Skip to main content

tract_linalg/x86_64/
softmax.rs

1// Accurate f32 softmax_l2: Cody-Waite argument reduction and a degree-5 minimax
2// e^r fit, identical arithmetic to the aarch64 NEON body in generic/reduce.rs
3// so every SIMD target agrees. `map_neutral` is NEG_INFINITY so a padded lane
4// underflows to zero even on a fully masked row, where the row max equals the
5// neutral.
6// nr=32 (4x ymm of 8), 32-byte aligned.
7routine_map_reduce_rust!(x86_64;
8    f32,
9    x86_64_fma_softmax2_f32_32n,
10    32,
11    8,
12    #[inline(never)]
13    fn run(buf: &mut [f32], max: f32) -> f32 {
14        assert!(buf.len() % 32 == 0);
15        unsafe { x86_64_fma_softmax2_f32_32n_run(buf, max) }
16    },
17    op(Softmax2),
18    isa(X86_64Avx2, X86_64Fma)
19);
20
21#[cfg(target_arch = "x86_64")]
22#[target_feature(enable = "avx2,fma")]
23unsafe fn x86_64_fma_softmax2_f32_32n_run(buf: &mut [f32], max: f32) -> f32 {
24    use std::arch::x86_64::*;
25    // exp(x) for x = row_value - max, i.e. never positive off the padding lanes.
26    #[inline(always)]
27    unsafe fn exp8(x: __m256) -> __m256 {
28        unsafe {
29            let k = _mm256_cvtps_epi32(_mm256_mul_ps(x, _mm256_set1_ps(1.442_695_04)));
30            let kf = _mm256_cvtepi32_ps(k);
31            let mut rr = _mm256_fnmadd_ps(kf, _mm256_set1_ps(0.693_145_75), x);
32            rr = _mm256_fnmadd_ps(kf, _mm256_set1_ps(1.428_606_8e-6), rr);
33            let mut q = _mm256_set1_ps(8.297653546e-03);
34            q = _mm256_fmadd_ps(q, rr, _mm256_set1_ps(4.191538191e-02));
35            q = _mm256_fmadd_ps(q, rr, _mm256_set1_ps(1.666757475e-01));
36            q = _mm256_fmadd_ps(q, rr, _mm256_set1_ps(4.999889485e-01));
37            q = _mm256_fmadd_ps(q, rr, _mm256_set1_ps(9.999996920e-01));
38            q = _mm256_fmadd_ps(q, rr, _mm256_set1_ps(1.000000072e+00));
39            // Biased exponent must stay a valid field: k + 127 goes non-positive
40            // around x = -88, and shifting that in would build -inf.
41            let biased = _mm256_max_epi32(
42                _mm256_min_epi32(
43                    _mm256_add_epi32(k, _mm256_set1_epi32(127)),
44                    _mm256_set1_epi32(254),
45                ),
46                _mm256_set1_epi32(1),
47            );
48            let scale = _mm256_castsi256_ps(_mm256_slli_epi32::<23>(biased));
49            let out = _mm256_or_ps(
50                _mm256_cmp_ps::<_CMP_LT_OQ>(x, _mm256_set1_ps(-103.0)),
51                _mm256_cmp_ps::<_CMP_GT_OQ>(x, _mm256_set1_ps(0.0)),
52            );
53            _mm256_blendv_ps(_mm256_mul_ps(q, scale), _mm256_setzero_ps(), out)
54        }
55    }
56    unsafe {
57        let vm = _mm256_set1_ps(max);
58        let mut a0 = _mm256_setzero_ps();
59        let mut a1 = _mm256_setzero_ps();
60        let mut a2 = _mm256_setzero_ps();
61        let mut a3 = _mm256_setzero_ps();
62        let p = buf.as_mut_ptr();
63        let n = buf.len();
64        let mut i = 0;
65        while i + 32 <= n {
66            let y0 = exp8(_mm256_sub_ps(_mm256_load_ps(p.add(i)), vm));
67            let y1 = exp8(_mm256_sub_ps(_mm256_load_ps(p.add(i + 8)), vm));
68            let y2 = exp8(_mm256_sub_ps(_mm256_load_ps(p.add(i + 16)), vm));
69            let y3 = exp8(_mm256_sub_ps(_mm256_load_ps(p.add(i + 24)), vm));
70            _mm256_store_ps(p.add(i), y0);
71            _mm256_store_ps(p.add(i + 8), y1);
72            _mm256_store_ps(p.add(i + 16), y2);
73            _mm256_store_ps(p.add(i + 24), y3);
74            a0 = _mm256_add_ps(a0, y0);
75            a1 = _mm256_add_ps(a1, y1);
76            a2 = _mm256_add_ps(a2, y2);
77            a3 = _mm256_add_ps(a3, y3);
78            i += 32;
79        }
80        let acc = _mm256_add_ps(_mm256_add_ps(a0, a1), _mm256_add_ps(a2, a3));
81        let mut s = _mm_add_ps(_mm256_castps256_ps128(acc), _mm256_extractf128_ps::<1>(acc));
82        s = _mm_add_ps(s, _mm_movehl_ps(s, s));
83        s = _mm_add_ss(s, _mm_shuffle_ps::<1>(s, s));
84        _mm_cvtss_f32(s)
85    }
86}
87
88// AVX-512 accurate f32 softmax_l2: same arithmetic as the FMA kernel, 64 f32
89// (4x zmm of 16) per iteration. Declares avx512f.
90// nr=64, 64-byte aligned.
91routine_map_reduce_rust!(x86_64;
92    f32,
93    x86_64_avx512_softmax2_f32_64n,
94    64,
95    16,
96    #[inline(never)]
97    fn run(buf: &mut [f32], max: f32) -> f32 {
98        assert!(buf.len() % 64 == 0);
99        unsafe { x86_64_avx512_softmax2_f32_64n_run(buf, max) }
100    },
101    op(Softmax2),
102    isa(X86_64Avx512f)
103);
104
105#[cfg(target_arch = "x86_64")]
106#[target_feature(enable = "avx512f")]
107unsafe fn x86_64_avx512_softmax2_f32_64n_run(buf: &mut [f32], max: f32) -> f32 {
108    use std::arch::x86_64::*;
109    #[inline(always)]
110    unsafe fn exp16(x: __m512) -> __m512 {
111        unsafe {
112            let k = _mm512_cvtps_epi32(_mm512_mul_ps(x, _mm512_set1_ps(1.442_695_04)));
113            let kf = _mm512_cvtepi32_ps(k);
114            let mut rr = _mm512_fnmadd_ps(kf, _mm512_set1_ps(0.693_145_75), x);
115            rr = _mm512_fnmadd_ps(kf, _mm512_set1_ps(1.428_606_8e-6), rr);
116            let mut q = _mm512_set1_ps(8.297653546e-03);
117            q = _mm512_fmadd_ps(q, rr, _mm512_set1_ps(4.191538191e-02));
118            q = _mm512_fmadd_ps(q, rr, _mm512_set1_ps(1.666757475e-01));
119            q = _mm512_fmadd_ps(q, rr, _mm512_set1_ps(4.999889485e-01));
120            q = _mm512_fmadd_ps(q, rr, _mm512_set1_ps(9.999996920e-01));
121            q = _mm512_fmadd_ps(q, rr, _mm512_set1_ps(1.000000072e+00));
122            let biased = _mm512_max_epi32(
123                _mm512_min_epi32(
124                    _mm512_add_epi32(k, _mm512_set1_epi32(127)),
125                    _mm512_set1_epi32(254),
126                ),
127                _mm512_set1_epi32(1),
128            );
129            let scale = _mm512_castsi512_ps(_mm512_slli_epi32::<23>(biased));
130            let out = _mm512_cmp_ps_mask::<_CMP_LT_OQ>(x, _mm512_set1_ps(-103.0))
131                | _mm512_cmp_ps_mask::<_CMP_GT_OQ>(x, _mm512_set1_ps(0.0));
132            _mm512_mask_blend_ps(out, _mm512_mul_ps(q, scale), _mm512_setzero_ps())
133        }
134    }
135    unsafe {
136        let vm = _mm512_set1_ps(max);
137        let mut a0 = _mm512_setzero_ps();
138        let mut a1 = _mm512_setzero_ps();
139        let mut a2 = _mm512_setzero_ps();
140        let mut a3 = _mm512_setzero_ps();
141        let p = buf.as_mut_ptr();
142        let n = buf.len();
143        let mut i = 0;
144        while i + 64 <= n {
145            let y0 = exp16(_mm512_sub_ps(_mm512_load_ps(p.add(i)), vm));
146            let y1 = exp16(_mm512_sub_ps(_mm512_load_ps(p.add(i + 16)), vm));
147            let y2 = exp16(_mm512_sub_ps(_mm512_load_ps(p.add(i + 32)), vm));
148            let y3 = exp16(_mm512_sub_ps(_mm512_load_ps(p.add(i + 48)), vm));
149            _mm512_store_ps(p.add(i), y0);
150            _mm512_store_ps(p.add(i + 16), y1);
151            _mm512_store_ps(p.add(i + 32), y2);
152            _mm512_store_ps(p.add(i + 48), y3);
153            a0 = _mm512_add_ps(a0, y0);
154            a1 = _mm512_add_ps(a1, y1);
155            a2 = _mm512_add_ps(a2, y2);
156            a3 = _mm512_add_ps(a3, y3);
157            i += 64;
158        }
159        _mm512_reduce_add_ps(_mm512_add_ps(_mm512_add_ps(a0, a1), _mm512_add_ps(a2, a3)))
160    }
161}