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}