10#ifndef EIGEN_COMPLEX_SVE_H
11#define EIGEN_COMPLEX_SVE_H
14#include "../../InternalHeaderCheck.h"
24 EIGEN_STRONG_INLINE PacketXcf() {}
25 EIGEN_STRONG_INLINE
explicit PacketXcf(
const PacketXf& a) : v(a) {}
30 EIGEN_STRONG_INLINE PacketXcd() {}
31 EIGEN_STRONG_INLINE
explicit PacketXcd(
const PacketXd& a) : v(a) {}
36struct packet_traits<std::complex<float>> : default_packet_traits {
37 typedef PacketXcf type;
38 typedef PacketXcf half;
42 size = sve_packet_size_selector<std::complex<float>, EIGEN_ARM64_SVE_VL>::size,
64struct packet_traits<std::complex<double>> : default_packet_traits {
65 typedef PacketXcd type;
66 typedef PacketXcd half;
70 size = sve_packet_size_selector<std::complex<double>, EIGEN_ARM64_SVE_VL>::size,
92struct unpacket_traits<PacketXcf> {
93 typedef std::complex<float> type;
94 typedef PacketXcf half;
95 typedef PacketXf as_real;
97 size = sve_packet_size_selector<std::complex<float>, EIGEN_ARM64_SVE_VL>::size,
98 alignment = sve_packet_alignment_selector<EIGEN_ARM64_SVE_VL>::alignment,
100 masked_load_available =
false,
101 masked_store_available =
false
106struct unpacket_traits<PacketXcd> {
107 typedef std::complex<double> type;
108 typedef PacketXcd half;
109 typedef PacketXd as_real;
111 size = sve_packet_size_selector<std::complex<double>, EIGEN_ARM64_SVE_VL>::size,
112 alignment = sve_packet_alignment_selector<EIGEN_ARM64_SVE_VL>::alignment,
114 masked_load_available =
false,
115 masked_store_available =
false
122EIGEN_STRONG_INLINE PacketXcf pset1<PacketXcf>(
const std::complex<float>& from) {
124 return PacketXcf(svreinterpret_f32_u64(svdup_n_u64(numext::bit_cast<numext::uint64_t>(from))));
128EIGEN_STRONG_INLINE PacketXcf pload<PacketXcf>(
const std::complex<float>* from) {
129 return PacketXcf(pload<PacketXf>(
reinterpret_cast<const float*
>(from)));
133EIGEN_STRONG_INLINE PacketXcf ploadu<PacketXcf>(
const std::complex<float>* from) {
134 return PacketXcf(ploadu<PacketXf>(
reinterpret_cast<const float*
>(from)));
138EIGEN_STRONG_INLINE
void pstore<std::complex<float>>(std::complex<float>* to,
const PacketXcf& from) {
139 pstore(
reinterpret_cast<float*
>(to), from.v);
143EIGEN_STRONG_INLINE
void pstoreu<std::complex<float>>(std::complex<float>* to,
const PacketXcf& from) {
144 pstoreu(
reinterpret_cast<float*
>(to), from.v);
148EIGEN_STRONG_INLINE PacketXcf ploaddup<PacketXcf>(
const std::complex<float>* from) {
153 constexpr uint64_t kHalf = uint64_t(packet_traits<std::complex<float>>::size) / 2;
154 const svuint64_t lo =
155 svreinterpret_u64_f32(svld1_f32(svwhilelt_b32(uint64_t(0), 2 * kHalf),
reinterpret_cast<const float*
>(from)));
156 return PacketXcf(svreinterpret_f32_u64(svzip1_u64(lo, lo)));
160EIGEN_STRONG_INLINE PacketXcf ploadquad<PacketXcf>(
const std::complex<float>* from) {
164 constexpr uint64_t kQuarter = numext::maxi(uint64_t(packet_traits<std::complex<float>>::size) / 4, uint64_t(1));
166 svreinterpret_u64_f32(svld1_f32(svwhilelt_b32(uint64_t(0), 2 * kQuarter),
reinterpret_cast<const float*
>(from)));
167 lo = svzip1_u64(lo, lo);
168 return PacketXcf(svreinterpret_f32_u64(svzip1_u64(lo, lo)));
172EIGEN_STRONG_INLINE PacketXcf pgather<std::complex<float>, PacketXcf>(
const std::complex<float>* from, Index stride) {
173 const svuint64_t idx = svindex_u64(0, numext::uint64_t(stride));
174 return PacketXcf(svreinterpret_f32_u64(
175 svld1_gather_u64index_u64(svptrue_b64(),
reinterpret_cast<const numext::uint64_t*
>(from), idx)));
179EIGEN_STRONG_INLINE
void pscatter<std::complex<float>, PacketXcf>(std::complex<float>* to,
const PacketXcf& from,
181 const svuint64_t idx = svindex_u64(0, numext::uint64_t(stride));
182 svst1_scatter_u64index_u64(svptrue_b64(),
reinterpret_cast<numext::uint64_t*
>(to), idx,
183 svreinterpret_u64_f32(from.v));
187EIGEN_STRONG_INLINE std::complex<float> pfirst<PacketXcf>(
const PacketXcf& a) {
189 return numext::bit_cast<std::complex<float>>(svlasta_u64(svpfalse_b(), svreinterpret_u64_f32(a.v)));
193EIGEN_STRONG_INLINE PacketXcf pconj(
const PacketXcf& a) {
197 svreinterpret_f32_u64(sveor_n_u64_x(svptrue_b64(), svreinterpret_u64_f32(a.v), numext::uint64_t(1) << 63)));
201EIGEN_STRONG_INLINE PacketXcf pcplxflip<PacketXcf>(
const PacketXcf& a) {
203 return PacketXcf(svreinterpret_f32_u64(svrevw_u64_x(svptrue_b64(), svreinterpret_u64_f32(a.v))));
207EIGEN_STRONG_INLINE PacketXcf pdupreal<PacketXcf>(
const PacketXcf& a) {
208 return PacketXcf(svtrn1_f32(a.v, a.v));
212EIGEN_STRONG_INLINE PacketXcf pdupimag<PacketXcf>(
const PacketXcf& a) {
213 return PacketXcf(svtrn2_f32(a.v, a.v));
217EIGEN_STRONG_INLINE PacketXcf preverse(
const PacketXcf& a) {
219 return PacketXcf(svreinterpret_f32_u64(svrev_u64(svreinterpret_u64_f32(a.v))));
223EIGEN_STRONG_INLINE std::complex<float> predux<PacketXcf>(
const PacketXcf& a) {
226 const svbool_t even = svptrue_b64();
227 const svbool_t odd = svrev_b32(even);
228 return {svaddv_f32(even, a.v), svaddv_f32(odd, a.v)};
236EIGEN_STRONG_INLINE PacketXcd pset1<PacketXcd>(
const std::complex<double>& from) {
238 return PacketXcd(svdupq_n_f64(numext::real(from), numext::imag(from)));
242EIGEN_STRONG_INLINE PacketXcd pload<PacketXcd>(
const std::complex<double>* from) {
243 return PacketXcd(pload<PacketXd>(
reinterpret_cast<const double*
>(from)));
247EIGEN_STRONG_INLINE PacketXcd ploadu<PacketXcd>(
const std::complex<double>* from) {
248 return PacketXcd(ploadu<PacketXd>(
reinterpret_cast<const double*
>(from)));
252EIGEN_STRONG_INLINE
void pstore<std::complex<double>>(std::complex<double>* to,
const PacketXcd& from) {
253 pstore(
reinterpret_cast<double*
>(to), from.v);
257EIGEN_STRONG_INLINE
void pstoreu<std::complex<double>>(std::complex<double>* to,
const PacketXcd& from) {
258 pstoreu(
reinterpret_cast<double*
>(to), from.v);
263EIGEN_STRONG_INLINE svuint64_t sve_cd_component_index(
const svuint64_t& value_index) {
264 const svuint64_t lane = svindex_u64(0, 1);
265 return svadd_u64_x(svptrue_b64(), svlsl_n_u64_x(svptrue_b64(), value_index, 1),
266 svand_n_u64_x(svptrue_b64(), lane, 1));
273template <
int kLog2Repeat>
274EIGEN_STRONG_INLINE PacketXcd sve_cd_loadrepeat(
const std::complex<double>* from) {
275 constexpr uint64_t kValues =
276 numext::maxi(uint64_t(packet_traits<std::complex<double>>::size) >> kLog2Repeat, uint64_t(1));
277 const svfloat64_t lo = svld1_f64(svwhilelt_b64(uint64_t(0), 2 * kValues),
reinterpret_cast<const double*
>(from));
278 const svuint64_t lane = svindex_u64(0, 1);
279 const svuint64_t idx =
280 svorr_u64_x(svptrue_b64(), svlsl_n_u64_x(svptrue_b64(), svlsr_n_u64_x(svptrue_b64(), lane, kLog2Repeat + 1), 1),
281 svand_n_u64_x(svptrue_b64(), lane, 1));
282 return PacketXcd(svtbl_f64(lo, idx));
286EIGEN_STRONG_INLINE PacketXcd ploaddup<PacketXcd>(
const std::complex<double>* from) {
287 return sve_cd_loadrepeat<1>(from);
291EIGEN_STRONG_INLINE PacketXcd ploadquad<PacketXcd>(
const std::complex<double>* from) {
292 return sve_cd_loadrepeat<2>(from);
296EIGEN_STRONG_INLINE PacketXcd pgather<std::complex<double>, PacketXcd>(
const std::complex<double>* from, Index stride) {
297 const svuint64_t value =
298 svmul_n_u64_x(svptrue_b64(), svlsr_n_u64_x(svptrue_b64(), svindex_u64(0, 1), 1), numext::uint64_t(stride));
300 svld1_gather_u64index_f64(svptrue_b64(),
reinterpret_cast<const double*
>(from), sve_cd_component_index(value)));
304EIGEN_STRONG_INLINE
void pscatter<std::complex<double>, PacketXcd>(std::complex<double>* to,
const PacketXcd& from,
306 const svuint64_t value =
307 svmul_n_u64_x(svptrue_b64(), svlsr_n_u64_x(svptrue_b64(), svindex_u64(0, 1), 1), numext::uint64_t(stride));
308 svst1_scatter_u64index_f64(svptrue_b64(),
reinterpret_cast<double*
>(to), sve_cd_component_index(value), from.v);
312EIGEN_STRONG_INLINE std::complex<double> pfirst<PacketXcd>(
const PacketXcd& a) {
314 return {svlastb_f64(svptrue_pat_b64(SV_VL1), a.v), svlastb_f64(svptrue_pat_b64(SV_VL2), a.v)};
318EIGEN_STRONG_INLINE PacketXcd pconj(
const PacketXcd& a) {
320 const svuint64_t mask = svdupq_n_u64(0, numext::uint64_t(1) << 63);
321 return PacketXcd(svreinterpret_f64_u64(sveor_u64_x(svptrue_b64(), svreinterpret_u64_f64(a.v), mask)));
325EIGEN_STRONG_INLINE PacketXcd pcplxflip<PacketXcd>(
const PacketXcd& a) {
327 return PacketXcd(svtbl_f64(a.v, sveor_n_u64_x(svptrue_b64(), svindex_u64(0, 1), 1)));
331EIGEN_STRONG_INLINE PacketXcd pdupreal<PacketXcd>(
const PacketXcd& a) {
332 return PacketXcd(svtrn1_f64(a.v, a.v));
336EIGEN_STRONG_INLINE PacketXcd pdupimag<PacketXcd>(
const PacketXcd& a) {
337 return PacketXcd(svtrn2_f64(a.v, a.v));
341EIGEN_STRONG_INLINE PacketXcd preverse(
const PacketXcd& a) {
343 return pcplxflip<PacketXcd>(PacketXcd(svrev_f64(a.v)));
347EIGEN_STRONG_INLINE std::complex<double> predux<PacketXcd>(
const PacketXcd& a) {
348 const svbool_t even = svdupq_n_b64(
true,
false);
349 return {svaddv_f64(even, a.v), svaddv_f64(svrev_b64(even), a.v)};
358EIGEN_STRONG_INLINE svbool_t sve_cd_quadword_even() {
359 const svuint64_t quadword = svlsr_n_u64_x(svptrue_b64(), svindex_u64(0, 1), 1);
360 return svcmpeq_n_u64(svptrue_b64(), svand_n_u64_x(svptrue_b64(), quadword, 1), 0);
366EIGEN_STRONG_INLINE svuint64_t sve_cd_zip_index(uint64_t offset) {
367 const svuint64_t lane = svindex_u64(0, 1);
368 const svuint64_t pair = svlsl_n_u64_x(svptrue_b64(), svlsr_n_u64_x(svptrue_b64(), lane, 2), 1);
369 return svadd_n_u64_x(svptrue_b64(), svorr_u64_x(svptrue_b64(), pair, svand_n_u64_x(svptrue_b64(), lane, 1)), offset);
373EIGEN_DEVICE_FUNC
inline void ptranspose(PacketBlock<PacketXcd, N>& kernel) {
374 EIGEN_STATIC_ASSERT((N & (N - 1)) == 0, EIGEN_INTERNAL_ERROR_PLEASE_FILE_A_BUG_REPORT);
375 constexpr uint64_t kLanes = 2 * uint64_t(unpacket_traits<PacketXcd>::size);
376 const svbool_t even = sve_cd_quadword_even();
377 const svuint64_t lo_index = sve_cd_zip_index(0);
378 const svuint64_t hi_index = sve_cd_zip_index(kLanes / 2);
379 for (
int stride = N / 2; stride > 0; stride >>= 1) {
380 for (
int block = 0; block < N; block += 2 * stride) {
381 for (
int k = 0; k < stride; ++k) {
382 const svfloat64_t a = kernel.packet[block + k].v;
383 const svfloat64_t b = kernel.packet[block + k + stride].v;
384 const svfloat64_t lo = svsel_f64(even, svtbl_f64(a, lo_index), svtbl_f64(b, lo_index));
385 const svfloat64_t hi = svsel_f64(even, svtbl_f64(a, hi_index), svtbl_f64(b, hi_index));
386 kernel.packet[block + k] = PacketXcd(lo);
387 kernel.packet[block + k + stride] = PacketXcd(hi);
396#define EIGEN_SVE_COMPLEX_DELEGATE(PACKET_CPLX) \
398 EIGEN_STRONG_INLINE PACKET_CPLX padd<PACKET_CPLX>(const PACKET_CPLX& a, const PACKET_CPLX& b) { \
399 return PACKET_CPLX(padd(a.v, b.v)); \
402 EIGEN_STRONG_INLINE PACKET_CPLX psub<PACKET_CPLX>(const PACKET_CPLX& a, const PACKET_CPLX& b) { \
403 return PACKET_CPLX(psub(a.v, b.v)); \
406 EIGEN_STRONG_INLINE PACKET_CPLX pnegate(const PACKET_CPLX& a) { \
407 return PACKET_CPLX(pnegate(a.v)); \
410 EIGEN_STRONG_INLINE PACKET_CPLX pzero<PACKET_CPLX>(const PACKET_CPLX& a) { \
411 return PACKET_CPLX(pzero(a.v)); \
414 EIGEN_STRONG_INLINE PACKET_CPLX pand<PACKET_CPLX>(const PACKET_CPLX& a, const PACKET_CPLX& b) { \
415 return PACKET_CPLX(pand(a.v, b.v)); \
418 EIGEN_STRONG_INLINE PACKET_CPLX por<PACKET_CPLX>(const PACKET_CPLX& a, const PACKET_CPLX& b) { \
419 return PACKET_CPLX(por(a.v, b.v)); \
422 EIGEN_STRONG_INLINE PACKET_CPLX pxor<PACKET_CPLX>(const PACKET_CPLX& a, const PACKET_CPLX& b) { \
423 return PACKET_CPLX(pxor(a.v, b.v)); \
426 EIGEN_STRONG_INLINE PACKET_CPLX pandnot<PACKET_CPLX>(const PACKET_CPLX& a, const PACKET_CPLX& b) { \
427 return PACKET_CPLX(pandnot(a.v, b.v)); \
430 EIGEN_STRONG_INLINE PACKET_CPLX pselect<PACKET_CPLX>(const PACKET_CPLX& mask, const PACKET_CPLX& a, \
431 const PACKET_CPLX& b) { \
432 return PACKET_CPLX(pselect(mask.v, a.v, b.v)); \
435 EIGEN_STRONG_INLINE PACKET_CPLX pmul<PACKET_CPLX>(const PACKET_CPLX& a, const PACKET_CPLX& b) { \
436 return pmul_complex(a, b); \
439 EIGEN_STRONG_INLINE PACKET_CPLX pdiv<PACKET_CPLX>(const PACKET_CPLX& a, const PACKET_CPLX& b) { \
440 return pdiv_complex(a, b); \
445 EIGEN_STRONG_INLINE PACKET_CPLX pcmp_eq<PACKET_CPLX>(const PACKET_CPLX& a, const PACKET_CPLX& b) { \
446 const PACKET_CPLX t = PACKET_CPLX(pcmp_eq(a.v, b.v)); \
447 return PACKET_CPLX(pand(pdupreal(t).v, pdupimag(t).v)); \
450EIGEN_SVE_COMPLEX_DELEGATE(PacketXcf)
451EIGEN_SVE_COMPLEX_DELEGATE(PacketXcd)
452#undef EIGEN_SVE_COMPLEX_DELEGATE
454EIGEN_INSTANTIATE_COMPLEX_MATH_FUNCS(PacketXcf)
455EIGEN_INSTANTIATE_COMPLEX_MATH_FUNCS(PacketXcd)
460EIGEN_DEVICE_FUNC
inline void ptranspose(PacketBlock<PacketXcf, N>& kernel) {
461 EIGEN_STATIC_ASSERT((N & (N - 1)) == 0, EIGEN_INTERNAL_ERROR_PLEASE_FILE_A_BUG_REPORT);
462 for (
int stride = N / 2; stride > 0; stride >>= 1) {
463 for (
int block = 0; block < N; block += 2 * stride) {
464 for (
int k = 0; k < stride; ++k) {
465 const svuint64_t a = svreinterpret_u64_f32(kernel.packet[block + k].v);
466 const svuint64_t b = svreinterpret_u64_f32(kernel.packet[block + k + stride].v);
467 kernel.packet[block + k] = PacketXcf(svreinterpret_f32_u64(svzip1_u64(a, b)));
468 kernel.packet[block + k + stride] = PacketXcf(svreinterpret_f32_u64(svzip2_u64(a, b)));
474EIGEN_MAKE_CONJ_HELPER_CPLX_REAL(PacketXcf, PacketXf)
475EIGEN_MAKE_CONJ_HELPER_CPLX_REAL(PacketXcd, PacketXd)
482struct has_packet_segment<PacketXcf> : std::true_type {};
485inline PacketXcf ploaduSegment<PacketXcf>(
const std::complex<float>* from, Index begin, Index count) {
486 return PacketXcf(svld1_f32(sve_segment_predicate_b32(2 * begin, 2 * count),
reinterpret_cast<const float*
>(from)));
490inline void pstoreuSegment<std::complex<float>, PacketXcf>(std::complex<float>* to,
const PacketXcf& from, Index begin,
492 svst1_f32(sve_segment_predicate_b32(2 * begin, 2 * count),
reinterpret_cast<float*
>(to), from.v);
496struct has_packet_segment<PacketXcd> : std::true_type {};
499inline PacketXcd ploaduSegment<PacketXcd>(
const std::complex<double>* from, Index begin, Index count) {
500 return PacketXcd(svld1_f64(sve_segment_predicate_b64(2 * begin, 2 * count),
reinterpret_cast<const double*
>(from)));
504inline void pstoreuSegment<std::complex<double>, PacketXcd>(std::complex<double>* to,
const PacketXcd& from,
505 Index begin, Index count) {
506 svst1_f64(sve_segment_predicate_b64(2 * begin, 2 * count),
reinterpret_cast<double*
>(to), from.v);