1routine_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 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}