Medial Code Documentation
Loading...
Searching...
No Matches
Complex.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2018 Gael Guennebaud <gael.guennebaud@inria.fr>
5//
6// This Source Code Form is subject to the terms of the Mozilla
7// Public License v. 2.0. If a copy of the MPL was not distributed
8// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
9
10#ifndef EIGEN_COMPLEX_AVX512_H
11#define EIGEN_COMPLEX_AVX512_H
12
13namespace Eigen {
14
15namespace internal {
16
17//---------- float ----------
19{
20 EIGEN_STRONG_INLINE Packet8cf() {}
21 EIGEN_STRONG_INLINE explicit Packet8cf(const __m512& a) : v(a) {}
22 __m512 v;
23};
24
25template<> struct packet_traits<std::complex<float> > : default_packet_traits
26{
27 typedef Packet8cf type;
28 typedef Packet4cf half;
29 enum {
30 Vectorizable = 1,
31 AlignedOnScalar = 1,
32 size = 8,
33 HasHalfPacket = 1,
34
35 HasAdd = 1,
36 HasSub = 1,
37 HasMul = 1,
38 HasDiv = 1,
39 HasNegate = 1,
40 HasSqrt = EIGEN_HAS_AVX512_MATH,
41 HasAbs = 0,
42 HasAbs2 = 0,
43 HasMin = 0,
44 HasMax = 0,
45 HasSetLinear = 0
46 };
47};
48
49template<> struct unpacket_traits<Packet8cf> {
50 typedef std::complex<float> type;
51 typedef Packet4cf half;
52 typedef Packet16f as_real;
53 enum {
54 size = 8,
56 vectorizable=true,
57 masked_load_available=false,
58 masked_store_available=false
59 };
60};
61
62template<> EIGEN_STRONG_INLINE Packet8cf ptrue<Packet8cf>(const Packet8cf& a) { return Packet8cf(ptrue(Packet16f(a.v))); }
63template<> EIGEN_STRONG_INLINE Packet8cf padd<Packet8cf>(const Packet8cf& a, const Packet8cf& b) { return Packet8cf(_mm512_add_ps(a.v,b.v)); }
64template<> EIGEN_STRONG_INLINE Packet8cf psub<Packet8cf>(const Packet8cf& a, const Packet8cf& b) { return Packet8cf(_mm512_sub_ps(a.v,b.v)); }
65template<> EIGEN_STRONG_INLINE Packet8cf pnegate(const Packet8cf& a)
66{
67 return Packet8cf(pnegate(a.v));
68}
69template<> EIGEN_STRONG_INLINE Packet8cf pconj(const Packet8cf& a)
70{
71 const __m512 mask = _mm512_castsi512_ps(_mm512_setr_epi32(
72 0x00000000,0x80000000,0x00000000,0x80000000,0x00000000,0x80000000,0x00000000,0x80000000,
73 0x00000000,0x80000000,0x00000000,0x80000000,0x00000000,0x80000000,0x00000000,0x80000000));
74 return Packet8cf(pxor(a.v,mask));
75}
76
77template<> EIGEN_STRONG_INLINE Packet8cf pmul<Packet8cf>(const Packet8cf& a, const Packet8cf& b)
78{
79 __m512 tmp2 = _mm512_mul_ps(_mm512_movehdup_ps(a.v), _mm512_permute_ps(b.v, _MM_SHUFFLE(2,3,0,1)));
80 return Packet8cf(_mm512_fmaddsub_ps(_mm512_moveldup_ps(a.v), b.v, tmp2));
81}
82
83template<> EIGEN_STRONG_INLINE Packet8cf pand <Packet8cf>(const Packet8cf& a, const Packet8cf& b) { return Packet8cf(pand(a.v,b.v)); }
84template<> EIGEN_STRONG_INLINE Packet8cf por <Packet8cf>(const Packet8cf& a, const Packet8cf& b) { return Packet8cf(por(a.v,b.v)); }
85template<> EIGEN_STRONG_INLINE Packet8cf pxor <Packet8cf>(const Packet8cf& a, const Packet8cf& b) { return Packet8cf(pxor(a.v,b.v)); }
86template<> EIGEN_STRONG_INLINE Packet8cf pandnot<Packet8cf>(const Packet8cf& a, const Packet8cf& b) { return Packet8cf(pandnot(a.v,b.v)); }
87
88template <>
89EIGEN_STRONG_INLINE Packet8cf pcmp_eq(const Packet8cf& a, const Packet8cf& b) {
90 __m512 eq = pcmp_eq<Packet16f>(a.v, b.v);
91 return Packet8cf(pand(eq, _mm512_permute_ps(eq, 0xB1)));
92}
93
94template<> EIGEN_STRONG_INLINE Packet8cf pload <Packet8cf>(const std::complex<float>* from) { EIGEN_DEBUG_ALIGNED_LOAD return Packet8cf(pload<Packet16f>(&numext::real_ref(*from))); }
95template<> EIGEN_STRONG_INLINE Packet8cf ploadu<Packet8cf>(const std::complex<float>* from) { EIGEN_DEBUG_UNALIGNED_LOAD return Packet8cf(ploadu<Packet16f>(&numext::real_ref(*from))); }
96
97
98template<> EIGEN_STRONG_INLINE Packet8cf pset1<Packet8cf>(const std::complex<float>& from)
99{
100 const float re = std::real(from);
101 const float im = std::imag(from);
102 return Packet8cf(_mm512_set_ps(im, re, im, re, im, re, im, re, im, re, im, re, im, re, im, re));
103}
104
105template<> EIGEN_STRONG_INLINE Packet8cf ploaddup<Packet8cf>(const std::complex<float>* from)
106{
107 return Packet8cf( _mm512_castpd_ps( ploaddup<Packet8d>((const double*)(const void*)from )) );
108}
109template<> EIGEN_STRONG_INLINE Packet8cf ploadquad<Packet8cf>(const std::complex<float>* from)
110{
111 return Packet8cf( _mm512_castpd_ps( ploadquad<Packet8d>((const double*)(const void*)from )) );
112}
113
114template<> EIGEN_STRONG_INLINE void pstore <std::complex<float> >(std::complex<float>* to, const Packet8cf& from) { EIGEN_DEBUG_ALIGNED_STORE pstore(&numext::real_ref(*to), from.v); }
115template<> EIGEN_STRONG_INLINE void pstoreu<std::complex<float> >(std::complex<float>* to, const Packet8cf& from) { EIGEN_DEBUG_UNALIGNED_STORE pstoreu(&numext::real_ref(*to), from.v); }
116
117template<> EIGEN_DEVICE_FUNC inline Packet8cf pgather<std::complex<float>, Packet8cf>(const std::complex<float>* from, Index stride)
118{
119 return Packet8cf(_mm512_castpd_ps(pgather<double,Packet8d>((const double*)(const void*)from, stride)));
120}
121
122template<> EIGEN_DEVICE_FUNC inline void pscatter<std::complex<float>, Packet8cf>(std::complex<float>* to, const Packet8cf& from, Index stride)
123{
124 pscatter((double*)(void*)to, _mm512_castps_pd(from.v), stride);
125}
126
127template<> EIGEN_STRONG_INLINE std::complex<float> pfirst<Packet8cf>(const Packet8cf& a)
128{
129 return pfirst(Packet2cf(_mm512_castps512_ps128(a.v)));
130}
131
132template<> EIGEN_STRONG_INLINE Packet8cf preverse(const Packet8cf& a) {
133 return Packet8cf(_mm512_castsi512_ps(
134 _mm512_permutexvar_epi64( _mm512_set_epi32(0, 0, 0, 1, 0, 2, 0, 3, 0, 4, 0, 5, 0, 6, 0, 7),
135 _mm512_castps_si512(a.v))));
136}
137
138template<> EIGEN_STRONG_INLINE std::complex<float> predux<Packet8cf>(const Packet8cf& a)
139{
140 return predux(padd(Packet4cf(extract256<0>(a.v)),
141 Packet4cf(extract256<1>(a.v))));
142}
143
144template<> EIGEN_STRONG_INLINE std::complex<float> predux_mul<Packet8cf>(const Packet8cf& a)
145{
146 return predux_mul(pmul(Packet4cf(extract256<0>(a.v)),
147 Packet4cf(extract256<1>(a.v))));
148}
149
150template <>
151EIGEN_STRONG_INLINE Packet4cf predux_half_dowto4<Packet8cf>(const Packet8cf& a) {
152 __m256 lane0 = extract256<0>(a.v);
153 __m256 lane1 = extract256<1>(a.v);
154 __m256 res = _mm256_add_ps(lane0, lane1);
155 return Packet4cf(res);
156}
157
158EIGEN_MAKE_CONJ_HELPER_CPLX_REAL(Packet8cf,Packet16f)
159
160template<> EIGEN_STRONG_INLINE Packet8cf pdiv<Packet8cf>(const Packet8cf& a, const Packet8cf& b)
161{
162 return pdiv_complex(a, b);
163}
164
165template<> EIGEN_STRONG_INLINE Packet8cf pcplxflip<Packet8cf>(const Packet8cf& x)
166{
167 return Packet8cf(_mm512_shuffle_ps(x.v, x.v, _MM_SHUFFLE(2, 3, 0 ,1)));
168}
169
170//---------- double ----------
172{
173 EIGEN_STRONG_INLINE Packet4cd() {}
174 EIGEN_STRONG_INLINE explicit Packet4cd(const __m512d& a) : v(a) {}
175 __m512d v;
176};
177
178template<> struct packet_traits<std::complex<double> > : default_packet_traits
179{
180 typedef Packet4cd type;
181 typedef Packet2cd half;
182 enum {
183 Vectorizable = 1,
184 AlignedOnScalar = 0,
185 size = 4,
186 HasHalfPacket = 1,
187
188 HasAdd = 1,
189 HasSub = 1,
190 HasMul = 1,
191 HasDiv = 1,
192 HasNegate = 1,
193 HasSqrt = EIGEN_HAS_AVX512_MATH,
194 HasAbs = 0,
195 HasAbs2 = 0,
196 HasMin = 0,
197 HasMax = 0,
198 HasSetLinear = 0
199 };
200};
201
202template<> struct unpacket_traits<Packet4cd> {
203 typedef std::complex<double> type;
204 typedef Packet2cd half;
205 typedef Packet8d as_real;
206 enum {
207 size = 4,
209 vectorizable=true,
210 masked_load_available=false,
211 masked_store_available=false
212 };
213};
214
215template<> EIGEN_STRONG_INLINE Packet4cd padd<Packet4cd>(const Packet4cd& a, const Packet4cd& b) { return Packet4cd(_mm512_add_pd(a.v,b.v)); }
216template<> EIGEN_STRONG_INLINE Packet4cd psub<Packet4cd>(const Packet4cd& a, const Packet4cd& b) { return Packet4cd(_mm512_sub_pd(a.v,b.v)); }
217template<> EIGEN_STRONG_INLINE Packet4cd pnegate(const Packet4cd& a) { return Packet4cd(pnegate(a.v)); }
218template<> EIGEN_STRONG_INLINE Packet4cd pconj(const Packet4cd& a)
219{
220 const __m512d mask = _mm512_castsi512_pd(
221 _mm512_set_epi32(0x80000000,0x0,0x0,0x0,0x80000000,0x0,0x0,0x0,
222 0x80000000,0x0,0x0,0x0,0x80000000,0x0,0x0,0x0));
223 return Packet4cd(pxor(a.v,mask));
224}
225
226template<> EIGEN_STRONG_INLINE Packet4cd pmul<Packet4cd>(const Packet4cd& a, const Packet4cd& b)
227{
228 __m512d tmp1 = _mm512_shuffle_pd(a.v,a.v,0x0);
229 __m512d tmp2 = _mm512_shuffle_pd(a.v,a.v,0xFF);
230 __m512d tmp3 = _mm512_shuffle_pd(b.v,b.v,0x55);
231 __m512d odd = _mm512_mul_pd(tmp2, tmp3);
232 return Packet4cd(_mm512_fmaddsub_pd(tmp1, b.v, odd));
233}
234
235template<> EIGEN_STRONG_INLINE Packet4cd ptrue<Packet4cd>(const Packet4cd& a) { return Packet4cd(ptrue(Packet8d(a.v))); }
236template<> EIGEN_STRONG_INLINE Packet4cd pand <Packet4cd>(const Packet4cd& a, const Packet4cd& b) { return Packet4cd(pand(a.v,b.v)); }
237template<> EIGEN_STRONG_INLINE Packet4cd por <Packet4cd>(const Packet4cd& a, const Packet4cd& b) { return Packet4cd(por(a.v,b.v)); }
238template<> EIGEN_STRONG_INLINE Packet4cd pxor <Packet4cd>(const Packet4cd& a, const Packet4cd& b) { return Packet4cd(pxor(a.v,b.v)); }
239template<> EIGEN_STRONG_INLINE Packet4cd pandnot<Packet4cd>(const Packet4cd& a, const Packet4cd& b) { return Packet4cd(pandnot(a.v,b.v)); }
240
241template <>
242EIGEN_STRONG_INLINE Packet4cd pcmp_eq(const Packet4cd& a, const Packet4cd& b) {
243 __m512d eq = pcmp_eq<Packet8d>(a.v, b.v);
244 return Packet4cd(pand(eq, _mm512_permute_pd(eq, 0x55)));
245}
246
247template<> EIGEN_STRONG_INLINE Packet4cd pload <Packet4cd>(const std::complex<double>* from)
248{ EIGEN_DEBUG_ALIGNED_LOAD return Packet4cd(pload<Packet8d>((const double*)from)); }
249template<> EIGEN_STRONG_INLINE Packet4cd ploadu<Packet4cd>(const std::complex<double>* from)
250{ EIGEN_DEBUG_UNALIGNED_LOAD return Packet4cd(ploadu<Packet8d>((const double*)from)); }
251
252template<> EIGEN_STRONG_INLINE Packet4cd pset1<Packet4cd>(const std::complex<double>& from)
253{
254 return Packet4cd(_mm512_castps_pd(_mm512_broadcast_f32x4( _mm_castpd_ps(pset1<Packet1cd>(from).v))));
255}
256
257template<> EIGEN_STRONG_INLINE Packet4cd ploaddup<Packet4cd>(const std::complex<double>* from) {
258 return Packet4cd(_mm512_insertf64x4(
259 _mm512_castpd256_pd512(ploaddup<Packet2cd>(from).v), ploaddup<Packet2cd>(from+1).v, 1));
260}
261
262template<> EIGEN_STRONG_INLINE void pstore <std::complex<double> >(std::complex<double> * to, const Packet4cd& from) { EIGEN_DEBUG_ALIGNED_STORE pstore((double*)to, from.v); }
263template<> EIGEN_STRONG_INLINE void pstoreu<std::complex<double> >(std::complex<double> * to, const Packet4cd& from) { EIGEN_DEBUG_UNALIGNED_STORE pstoreu((double*)to, from.v); }
264
265template<> EIGEN_DEVICE_FUNC inline Packet4cd pgather<std::complex<double>, Packet4cd>(const std::complex<double>* from, Index stride)
266{
267 return Packet4cd(_mm512_insertf64x4(_mm512_castpd256_pd512(
268 _mm256_insertf128_pd(_mm256_castpd128_pd256(ploadu<Packet1cd>(from+0*stride).v), ploadu<Packet1cd>(from+1*stride).v,1)),
269 _mm256_insertf128_pd(_mm256_castpd128_pd256(ploadu<Packet1cd>(from+2*stride).v), ploadu<Packet1cd>(from+3*stride).v,1), 1));
270}
271
272template<> EIGEN_DEVICE_FUNC inline void pscatter<std::complex<double>, Packet4cd>(std::complex<double>* to, const Packet4cd& from, Index stride)
273{
274 __m512i fromi = _mm512_castpd_si512(from.v);
275 double* tod = (double*)(void*)to;
276 _mm_storeu_pd(tod+0*stride, _mm_castsi128_pd(_mm512_extracti32x4_epi32(fromi,0)) );
277 _mm_storeu_pd(tod+2*stride, _mm_castsi128_pd(_mm512_extracti32x4_epi32(fromi,1)) );
278 _mm_storeu_pd(tod+4*stride, _mm_castsi128_pd(_mm512_extracti32x4_epi32(fromi,2)) );
279 _mm_storeu_pd(tod+6*stride, _mm_castsi128_pd(_mm512_extracti32x4_epi32(fromi,3)) );
280}
281
282template<> EIGEN_STRONG_INLINE std::complex<double> pfirst<Packet4cd>(const Packet4cd& a)
283{
284 __m128d low = extract128<0>(a.v);
285 EIGEN_ALIGN16 double res[2];
286 _mm_store_pd(res, low);
287 return std::complex<double>(res[0],res[1]);
288}
289
290template<> EIGEN_STRONG_INLINE Packet4cd preverse(const Packet4cd& a) {
291 return Packet4cd(_mm512_shuffle_f64x2(a.v, a.v, (shuffle_mask<3,2,1,0>::mask)));
292}
293
294template<> EIGEN_STRONG_INLINE std::complex<double> predux<Packet4cd>(const Packet4cd& a)
295{
296 return predux(padd(Packet2cd(_mm512_extractf64x4_pd(a.v,0)),
297 Packet2cd(_mm512_extractf64x4_pd(a.v,1))));
298}
299
300template<> EIGEN_STRONG_INLINE std::complex<double> predux_mul<Packet4cd>(const Packet4cd& a)
301{
302 return predux_mul(pmul(Packet2cd(_mm512_extractf64x4_pd(a.v,0)),
303 Packet2cd(_mm512_extractf64x4_pd(a.v,1))));
304}
305
306EIGEN_MAKE_CONJ_HELPER_CPLX_REAL(Packet4cd,Packet8d)
307
308template<> EIGEN_STRONG_INLINE Packet4cd pdiv<Packet4cd>(const Packet4cd& a, const Packet4cd& b)
309{
310 return pdiv_complex(a, b);
311}
312
313template<> EIGEN_STRONG_INLINE Packet4cd pcplxflip<Packet4cd>(const Packet4cd& x)
314{
315 return Packet4cd(_mm512_permute_pd(x.v,0x55));
316}
317
318EIGEN_DEVICE_FUNC inline void
319ptranspose(PacketBlock<Packet8cf,4>& kernel) {
320 PacketBlock<Packet8d,4> pb;
321
322 pb.packet[0] = _mm512_castps_pd(kernel.packet[0].v);
323 pb.packet[1] = _mm512_castps_pd(kernel.packet[1].v);
324 pb.packet[2] = _mm512_castps_pd(kernel.packet[2].v);
325 pb.packet[3] = _mm512_castps_pd(kernel.packet[3].v);
326 ptranspose(pb);
327 kernel.packet[0].v = _mm512_castpd_ps(pb.packet[0]);
328 kernel.packet[1].v = _mm512_castpd_ps(pb.packet[1]);
329 kernel.packet[2].v = _mm512_castpd_ps(pb.packet[2]);
330 kernel.packet[3].v = _mm512_castpd_ps(pb.packet[3]);
331}
332
333EIGEN_DEVICE_FUNC inline void
334ptranspose(PacketBlock<Packet8cf,8>& kernel) {
335 PacketBlock<Packet8d,8> pb;
336
337 pb.packet[0] = _mm512_castps_pd(kernel.packet[0].v);
338 pb.packet[1] = _mm512_castps_pd(kernel.packet[1].v);
339 pb.packet[2] = _mm512_castps_pd(kernel.packet[2].v);
340 pb.packet[3] = _mm512_castps_pd(kernel.packet[3].v);
341 pb.packet[4] = _mm512_castps_pd(kernel.packet[4].v);
342 pb.packet[5] = _mm512_castps_pd(kernel.packet[5].v);
343 pb.packet[6] = _mm512_castps_pd(kernel.packet[6].v);
344 pb.packet[7] = _mm512_castps_pd(kernel.packet[7].v);
345 ptranspose(pb);
346 kernel.packet[0].v = _mm512_castpd_ps(pb.packet[0]);
347 kernel.packet[1].v = _mm512_castpd_ps(pb.packet[1]);
348 kernel.packet[2].v = _mm512_castpd_ps(pb.packet[2]);
349 kernel.packet[3].v = _mm512_castpd_ps(pb.packet[3]);
350 kernel.packet[4].v = _mm512_castpd_ps(pb.packet[4]);
351 kernel.packet[5].v = _mm512_castpd_ps(pb.packet[5]);
352 kernel.packet[6].v = _mm512_castpd_ps(pb.packet[6]);
353 kernel.packet[7].v = _mm512_castpd_ps(pb.packet[7]);
354}
355
356EIGEN_DEVICE_FUNC inline void
357ptranspose(PacketBlock<Packet4cd,4>& kernel) {
358 __m512d T0 = _mm512_shuffle_f64x2(kernel.packet[0].v, kernel.packet[1].v, (shuffle_mask<0,1,0,1>::mask)); // [a0 a1 b0 b1]
359 __m512d T1 = _mm512_shuffle_f64x2(kernel.packet[0].v, kernel.packet[1].v, (shuffle_mask<2,3,2,3>::mask)); // [a2 a3 b2 b3]
360 __m512d T2 = _mm512_shuffle_f64x2(kernel.packet[2].v, kernel.packet[3].v, (shuffle_mask<0,1,0,1>::mask)); // [c0 c1 d0 d1]
361 __m512d T3 = _mm512_shuffle_f64x2(kernel.packet[2].v, kernel.packet[3].v, (shuffle_mask<2,3,2,3>::mask)); // [c2 c3 d2 d3]
362
363 kernel.packet[3] = Packet4cd(_mm512_shuffle_f64x2(T1, T3, (shuffle_mask<1,3,1,3>::mask))); // [a3 b3 c3 d3]
364 kernel.packet[2] = Packet4cd(_mm512_shuffle_f64x2(T1, T3, (shuffle_mask<0,2,0,2>::mask))); // [a2 b2 c2 d2]
365 kernel.packet[1] = Packet4cd(_mm512_shuffle_f64x2(T0, T2, (shuffle_mask<1,3,1,3>::mask))); // [a1 b1 c1 d1]
366 kernel.packet[0] = Packet4cd(_mm512_shuffle_f64x2(T0, T2, (shuffle_mask<0,2,0,2>::mask))); // [a0 b0 c0 d0]
367}
368
369#if EIGEN_HAS_AVX512_MATH
370
371template<> EIGEN_STRONG_INLINE Packet4cd psqrt<Packet4cd>(const Packet4cd& a) {
372 return psqrt_complex<Packet4cd>(a);
373}
374
375template<> EIGEN_STRONG_INLINE Packet8cf psqrt<Packet8cf>(const Packet8cf& a) {
376 return psqrt_complex<Packet8cf>(a);
377}
378
379#endif
380
381} // end namespace internal
382} // end namespace Eigen
383
384#endif // EIGEN_COMPLEX_AVX512_H
Base class for all dense matrices, vectors, and expressions.
Definition MatrixBase.h:50
Namespace containing all symbols from the Eigen library.
Definition LDLT.h:16
EIGEN_DEFAULT_DENSE_INDEX_TYPE Index
The Index type as used for the API.
Definition Meta.h:74
Definition BFloat16.h:88
Definition Half.h:140
Definition Complex.h:187
Definition Complex.h:172
Definition Complex.h:19
Definition Complex.h:19
Definition GenericPacketMath.h:43
Definition GenericPacketMath.h:107
Definition GenericPacketMath.h:133