Skip to main content

tract_linalg/x86_64/
erf.rs

1// AVX-512 (zmm, 16-wide) error function kernel. Mirrors generic/erf.rs::serf
2// (Abramowitz & Stegun 7.1.26 six-coefficient polynomial) but runs the
3// polynomial via FMA chains over 4 zmm registers per iteration (64 lanes per
4// loop step). Validated against the generic scalar reference via
5// erf_frame_tests! at SuperApproximate tolerance.
6//
7// Algorithm (per lane):
8//   signum = sign(x);  abs = |x|
9//   y = a6
10//   y = y*abs + a5            (Horner FMA)
11//   y = y*abs + a4
12//   y = y*abs + a3
13//   y = y*abs + a2
14//   y = y*abs + a1
15//   y = y * abs               (final factor of abs)
16//   y = y + 1
17//   y = y^16                  (4 sequential squares)
18//   y = 1 / y                 (vdivps, full IEEE precision)
19//   y = 1 - y
20//   result = copysign(y, x)
21
22routine_ew_rust!(x86_64;
23    f32,
24    x86_64_avx512_erf_f32_64n,
25    64,
26    16,
27    #[inline(never)]
28    fn run(buf: &mut [f32], _: ()) {
29        debug_assert!(buf.len() % Self::nr() == 0);
30        debug_assert!(buf.as_ptr() as usize % Self::alignment_bytes() == 0);
31        if buf.is_empty() {
32            return;
33        }
34        unsafe { x86_64_avx512_erf_f32_64n_run(buf) }
35    },
36    func(Erf),
37    isa(X86_64Avx512f)
38);
39
40#[cfg(target_arch = "x86_64")]
41#[target_feature(enable = "avx512f")]
42unsafe fn x86_64_avx512_erf_f32_64n_run(buf: &mut [f32]) {
43    unsafe {
44        let len = buf.len();
45        let ptr = buf.as_ptr();
46        const A1: f32 = 0.0705230784;
47        const A2: f32 = 0.0422820123;
48        const A3: f32 = 0.0092705272;
49        const A4: f32 = 0.0001520143;
50        const A5: f32 = 0.0002765672;
51        const A6: f32 = 0.0000430638;
52        // 0x7fffffff: positive-finite mask (clears sign bit). As f32 bits, this
53        // is NaN; we never use it as a numeric value — only as a bit mask via vandps.
54        const ABS_MASK: f32 = f32::from_bits(0x7fffffff);
55        const SIGN_MASK: f32 = f32::from_bits(0x80000000);
56        std::arch::asm!("
57            // broadcast constants (xmmN -> zmmN, broadcast across all 16 lanes)
58            vbroadcastss zmm0, xmm0           // a1
59            vbroadcastss zmm1, xmm1           // a2
60            vbroadcastss zmm2, xmm2           // a3
61            vbroadcastss zmm3, xmm3           // a4
62            vbroadcastss zmm4, xmm4           // a5
63            vbroadcastss zmm5, xmm5           // a6
64            vbroadcastss zmm6, xmm6           // 1.0
65            vbroadcastss zmm7, xmm7           // abs mask (0x7fffffff)
66            vbroadcastss zmm8, xmm8           // sign mask (0x80000000)
67            2:
68                // load 4 zmm of input
69                vmovaps zmm9,  [{ptr}]
70                vmovaps zmm10, [{ptr} + 64]
71                vmovaps zmm11, [{ptr} + 128]
72                vmovaps zmm12, [{ptr} + 192]
73
74                // sign[i] = x[i] & SIGN_MASK   (keeps only the sign bit)
75                vandps zmm13, zmm9,  zmm8
76                vandps zmm14, zmm10, zmm8
77                vandps zmm15, zmm11, zmm8
78                vandps zmm16, zmm12, zmm8
79
80                // abs[i] = x[i] & ABS_MASK     (clears the sign bit)
81                vandps zmm9,  zmm9,  zmm7
82                vandps zmm10, zmm10, zmm7
83                vandps zmm11, zmm11, zmm7
84                vandps zmm12, zmm12, zmm7
85
86                // y = a6 (in zmm17..20, 4 independent channels)
87                vmovaps zmm17, zmm5
88                vmovaps zmm18, zmm5
89                vmovaps zmm19, zmm5
90                vmovaps zmm20, zmm5
91
92                // y = y*abs + a5
93                vfmadd213ps zmm17, zmm9,  zmm4
94                vfmadd213ps zmm18, zmm10, zmm4
95                vfmadd213ps zmm19, zmm11, zmm4
96                vfmadd213ps zmm20, zmm12, zmm4
97
98                // y = y*abs + a4
99                vfmadd213ps zmm17, zmm9,  zmm3
100                vfmadd213ps zmm18, zmm10, zmm3
101                vfmadd213ps zmm19, zmm11, zmm3
102                vfmadd213ps zmm20, zmm12, zmm3
103
104                // y = y*abs + a3
105                vfmadd213ps zmm17, zmm9,  zmm2
106                vfmadd213ps zmm18, zmm10, zmm2
107                vfmadd213ps zmm19, zmm11, zmm2
108                vfmadd213ps zmm20, zmm12, zmm2
109
110                // y = y*abs + a2
111                vfmadd213ps zmm17, zmm9,  zmm1
112                vfmadd213ps zmm18, zmm10, zmm1
113                vfmadd213ps zmm19, zmm11, zmm1
114                vfmadd213ps zmm20, zmm12, zmm1
115
116                // y = y*abs + a1
117                vfmadd213ps zmm17, zmm9,  zmm0
118                vfmadd213ps zmm18, zmm10, zmm0
119                vfmadd213ps zmm19, zmm11, zmm0
120                vfmadd213ps zmm20, zmm12, zmm0
121
122                // y = y * abs  (final factor)
123                vmulps zmm17, zmm17, zmm9
124                vmulps zmm18, zmm18, zmm10
125                vmulps zmm19, zmm19, zmm11
126                vmulps zmm20, zmm20, zmm12
127
128                // y = y + 1
129                vaddps zmm17, zmm17, zmm6
130                vaddps zmm18, zmm18, zmm6
131                vaddps zmm19, zmm19, zmm6
132                vaddps zmm20, zmm20, zmm6
133
134                // y^16: square 4 times
135                vmulps zmm17, zmm17, zmm17
136                vmulps zmm18, zmm18, zmm18
137                vmulps zmm19, zmm19, zmm19
138                vmulps zmm20, zmm20, zmm20
139
140                vmulps zmm17, zmm17, zmm17
141                vmulps zmm18, zmm18, zmm18
142                vmulps zmm19, zmm19, zmm19
143                vmulps zmm20, zmm20, zmm20
144
145                vmulps zmm17, zmm17, zmm17
146                vmulps zmm18, zmm18, zmm18
147                vmulps zmm19, zmm19, zmm19
148                vmulps zmm20, zmm20, zmm20
149
150                vmulps zmm17, zmm17, zmm17
151                vmulps zmm18, zmm18, zmm18
152                vmulps zmm19, zmm19, zmm19
153                vmulps zmm20, zmm20, zmm20
154
155                // y = 1 / y      (full-precision reciprocal, matches generic .recip())
156                vdivps zmm21, zmm6, zmm17
157                vdivps zmm22, zmm6, zmm18
158                vdivps zmm23, zmm6, zmm19
159                vdivps zmm24, zmm6, zmm20
160
161                // y = 1 - y
162                vsubps zmm21, zmm6, zmm21
163                vsubps zmm22, zmm6, zmm22
164                vsubps zmm23, zmm6, zmm23
165                vsubps zmm24, zmm6, zmm24
166
167                // copysign: stamp the original sign bit onto the (positive) result
168                vorps zmm21, zmm21, zmm13
169                vorps zmm22, zmm22, zmm14
170                vorps zmm23, zmm23, zmm15
171                vorps zmm24, zmm24, zmm16
172
173                // store
174                vmovaps [{ptr}],       zmm21
175                vmovaps [{ptr} + 64],  zmm22
176                vmovaps [{ptr} + 128], zmm23
177                vmovaps [{ptr} + 192], zmm24
178
179                add {ptr}, 256
180                sub {len}, 64
181                jnz 2b
182            ",
183            len = inout(reg) len => _,
184            ptr = inout(reg) ptr => _,
185            inout("xmm0") A1 => _,
186            inout("xmm1") A2 => _,
187            inout("xmm2") A3 => _,
188            inout("xmm3") A4 => _,
189            inout("xmm4") A5 => _,
190            inout("xmm5") A6 => _,
191            inout("xmm6") 1f32 => _,
192            inout("xmm7") ABS_MASK => _,
193            inout("xmm8") SIGN_MASK => _,
194            out("zmm9")  _, out("zmm10") _, out("zmm11") _, out("zmm12") _,
195            out("zmm13") _, out("zmm14") _, out("zmm15") _, out("zmm16") _,
196            out("zmm17") _, out("zmm18") _, out("zmm19") _, out("zmm20") _,
197            out("zmm21") _, out("zmm22") _, out("zmm23") _, out("zmm24") _,
198        );
199    }
200}
201
202// AVX2/FMA (ymm, 8-wide) error function kernel. Same polynomial and shape as
203// the AVX-512 kernel above, over 4 ymm registers per iteration (32 lanes per
204// loop step), for the x86_64 tier without AVX-512.
205routine_ew_rust!(x86_64;
206    f32,
207    x86_64_fma_erf_f32_32n,
208    32,
209    8,
210    #[inline(never)]
211    fn run(buf: &mut [f32], _: ()) {
212        debug_assert!(buf.len() % Self::nr() == 0);
213        debug_assert!(buf.as_ptr() as usize % Self::alignment_bytes() == 0);
214        if buf.is_empty() {
215            return;
216        }
217        unsafe { x86_64_fma_erf_f32_32n_run(buf) }
218    },
219    func(Erf),
220    isa(X86_64Fma)
221);
222
223#[cfg(target_arch = "x86_64")]
224#[target_feature(enable = "avx,fma")]
225unsafe fn x86_64_fma_erf_f32_32n_run(buf: &mut [f32]) {
226    unsafe {
227        use std::arch::x86_64::*;
228        const A1: f32 = 0.0705230784;
229        const A2: f32 = 0.0422820123;
230        const A3: f32 = 0.0092705272;
231        const A4: f32 = 0.0001520143;
232        const A5: f32 = 0.0002765672;
233        const A6: f32 = 0.0000430638;
234        // 0x7fffffff / 0x80000000 are bit masks, never numeric values: as f32
235        // the first is NaN. Used only through vandps/vorps.
236        let abs_mask = _mm256_set1_ps(f32::from_bits(0x7fffffff));
237        let sign_mask = _mm256_set1_ps(f32::from_bits(0x80000000));
238        let a1 = _mm256_set1_ps(A1);
239        let a2 = _mm256_set1_ps(A2);
240        let a3 = _mm256_set1_ps(A3);
241        let a4 = _mm256_set1_ps(A4);
242        let a5 = _mm256_set1_ps(A5);
243        let a6 = _mm256_set1_ps(A6);
244        let one = _mm256_set1_ps(1.0);
245
246        let erf8 = |x: __m256| -> __m256 {
247            let abs = _mm256_and_ps(x, abs_mask);
248            let sign = _mm256_and_ps(x, sign_mask);
249            let y = _mm256_fmadd_ps(a6, abs, a5);
250            let y = _mm256_fmadd_ps(y, abs, a4);
251            let y = _mm256_fmadd_ps(y, abs, a3);
252            let y = _mm256_fmadd_ps(y, abs, a2);
253            let y = _mm256_fmadd_ps(y, abs, a1);
254            let y = _mm256_fmadd_ps(y, abs, one);
255            let y = _mm256_mul_ps(y, y);
256            let y = _mm256_mul_ps(y, y);
257            let y = _mm256_mul_ps(y, y);
258            let y = _mm256_mul_ps(y, y);
259            let y = _mm256_sub_ps(one, _mm256_div_ps(one, y));
260            _mm256_or_ps(_mm256_and_ps(y, abs_mask), sign)
261        };
262
263        let ptr = buf.as_mut_ptr();
264        let mut i = 0;
265        while i < buf.len() {
266            let p = ptr.add(i);
267            let r0 = erf8(_mm256_loadu_ps(p));
268            let r1 = erf8(_mm256_loadu_ps(p.add(8)));
269            let r2 = erf8(_mm256_loadu_ps(p.add(16)));
270            let r3 = erf8(_mm256_loadu_ps(p.add(24)));
271            _mm256_storeu_ps(p, r0);
272            _mm256_storeu_ps(p.add(8), r1);
273            _mm256_storeu_ps(p.add(16), r2);
274            _mm256_storeu_ps(p.add(24), r3);
275            i += 32;
276        }
277    }
278}