1routine_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 #[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 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
88routine_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}