11#ifndef EIGEN_PACKET_MATH_GPU_H
12#define EIGEN_PACKET_MATH_GPU_H
15#include "../../InternalHeaderCheck.h"
27#if defined(EIGEN_GPUCC) && defined(EIGEN_USE_GPU)
30struct is_arithmetic<float4> : std::true_type {};
32struct is_arithmetic<double2> : std::true_type {};
35struct packet_traits<float> : default_packet_traits {
38 static constexpr int Vectorizable = 1;
39 static constexpr int AlignedOnScalar = 1;
40 static constexpr int size = 4;
42 static constexpr int HasDiv = 1;
43 static constexpr int HasSin = 0;
44 static constexpr int HasCos = 0;
45 static constexpr int HasLog = 1;
46 static constexpr int HasExp = 1;
47 static constexpr int HasSqrt = 1;
48 static constexpr int HasRsqrt = 1;
49 static constexpr int HasLGamma = 1;
50 static constexpr int HasDiGamma = 1;
51 static constexpr int HasZeta = 1;
52 static constexpr int HasPolygamma = 1;
53 static constexpr int HasErf = 1;
54 static constexpr int HasErfc = 1;
55 static constexpr int HasNdtri = 1;
56 static constexpr int HasBessel = 1;
57 static constexpr int HasIGamma = 1;
58 static constexpr int HasIGammaDerA = 1;
59 static constexpr int HasGammaSampleDerAlpha = 1;
60 static constexpr int HasIGammac = 1;
61 static constexpr int HasBetaInc = 1;
63 static constexpr int HasCmp = 1;
67struct packet_traits<double> : default_packet_traits {
70 static constexpr int Vectorizable = 1;
71 static constexpr int AlignedOnScalar = 1;
72 static constexpr int size = 2;
74 static constexpr int HasDiv = 1;
75 static constexpr int HasLog = 1;
76 static constexpr int HasExp = 1;
77 static constexpr int HasSqrt = 1;
78 static constexpr int HasRsqrt = 1;
79 static constexpr int HasLGamma = 1;
80 static constexpr int HasDiGamma = 1;
81 static constexpr int HasZeta = 1;
82 static constexpr int HasPolygamma = 1;
83 static constexpr int HasErf = 1;
84 static constexpr int HasErfc = 1;
85 static constexpr int HasNdtri = 1;
86 static constexpr int HasBessel = 1;
87 static constexpr int HasIGamma = 1;
88 static constexpr int HasIGammaDerA = 1;
89 static constexpr int HasGammaSampleDerAlpha = 1;
90 static constexpr int HasIGammac = 1;
91 static constexpr int HasBetaInc = 1;
93 static constexpr int HasCmp = 1;
97struct unpacket_traits<float4> {
99 static constexpr int size = 4;
100 static constexpr int alignment =
Aligned16;
101 static constexpr bool vectorizable =
true;
102 static constexpr bool masked_load_available =
false;
103 static constexpr bool masked_store_available =
false;
107struct unpacket_traits<double2> {
109 static constexpr int size = 2;
110 static constexpr int alignment =
Aligned16;
111 static constexpr bool vectorizable =
true;
112 static constexpr bool masked_load_available =
false;
113 static constexpr bool masked_store_available =
false;
114 using half = double2;
118EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 pset1<float4>(
const float& from) {
119 return make_float4(from, from, from, from);
122EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 pset1<double2>(
const double& from) {
123 return make_double2(from, from);
130using lane_bits_t =
typename numext::get_integer_by_size<
sizeof(T)>::unsigned_type;
132#define EIGEN_MAKE_BITWISE_BINOP(name, op) \
133 template <typename T> \
134 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T bitwise_##name(const T& a, const T& b) { \
135 using Bits = lane_bits_t<T>; \
136 return numext::bit_cast<T>(static_cast<Bits>(numext::bit_cast<Bits>(a) op numext::bit_cast<Bits>(b))); \
139EIGEN_MAKE_BITWISE_BINOP(and, &)
140EIGEN_MAKE_BITWISE_BINOP(or, |)
141EIGEN_MAKE_BITWISE_BINOP(xor, ^)
142EIGEN_MAKE_BITWISE_BINOP(andnot, &~)
144#undef EIGEN_MAKE_BITWISE_BINOP
149EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T mask_from(
bool condition) {
150 using Bits = lane_bits_t<T>;
151 return numext::bit_cast<T>(condition ? ~Bits(0) : Bits(0));
155EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T eq_mask(
const T& a,
const T& b) {
156 return mask_from<T>(a == b);
159EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T lt_mask(
const T& a,
const T& b) {
160 return mask_from<T>(a < b);
163EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T le_mask(
const T& a,
const T& b) {
164 return mask_from<T>(a <= b);
168EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T lt_or_nan_mask(
const T& a,
const T& b) {
169 return mask_from<T>(!(a >= b));
173EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 pand<float4>(
const float4& a,
const float4& b) {
174 return make_float4(bitwise_and(a.x, b.x), bitwise_and(a.y, b.y), bitwise_and(a.z, b.z), bitwise_and(a.w, b.w));
177EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 pand<double2>(
const double2& a,
const double2& b) {
178 return make_double2(bitwise_and(a.x, b.x), bitwise_and(a.y, b.y));
182EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 por<float4>(
const float4& a,
const float4& b) {
183 return make_float4(bitwise_or(a.x, b.x), bitwise_or(a.y, b.y), bitwise_or(a.z, b.z), bitwise_or(a.w, b.w));
186EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 por<double2>(
const double2& a,
const double2& b) {
187 return make_double2(bitwise_or(a.x, b.x), bitwise_or(a.y, b.y));
191EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 pxor<float4>(
const float4& a,
const float4& b) {
192 return make_float4(bitwise_xor(a.x, b.x), bitwise_xor(a.y, b.y), bitwise_xor(a.z, b.z), bitwise_xor(a.w, b.w));
195EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 pxor<double2>(
const double2& a,
const double2& b) {
196 return make_double2(bitwise_xor(a.x, b.x), bitwise_xor(a.y, b.y));
200EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 pandnot<float4>(
const float4& a,
const float4& b) {
201 return make_float4(bitwise_andnot(a.x, b.x), bitwise_andnot(a.y, b.y), bitwise_andnot(a.z, b.z),
202 bitwise_andnot(a.w, b.w));
205EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 pandnot<double2>(
const double2& a,
const double2& b) {
206 return make_double2(bitwise_andnot(a.x, b.x), bitwise_andnot(a.y, b.y));
210EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 pcmp_eq<float4>(
const float4& a,
const float4& b) {
211 return make_float4(eq_mask(a.x, b.x), eq_mask(a.y, b.y), eq_mask(a.z, b.z), eq_mask(a.w, b.w));
214EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 pcmp_lt<float4>(
const float4& a,
const float4& b) {
215 return make_float4(lt_mask(a.x, b.x), lt_mask(a.y, b.y), lt_mask(a.z, b.z), lt_mask(a.w, b.w));
218EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 pcmp_le<float4>(
const float4& a,
const float4& b) {
219 return make_float4(le_mask(a.x, b.x), le_mask(a.y, b.y), le_mask(a.z, b.z), le_mask(a.w, b.w));
222EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 pcmp_lt_or_nan<float4>(
const float4& a,
const float4& b) {
223 return make_float4(lt_or_nan_mask(a.x, b.x), lt_or_nan_mask(a.y, b.y), lt_or_nan_mask(a.z, b.z),
224 lt_or_nan_mask(a.w, b.w));
227EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 pcmp_eq<double2>(
const double2& a,
const double2& b) {
228 return make_double2(eq_mask(a.x, b.x), eq_mask(a.y, b.y));
231EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 pcmp_lt<double2>(
const double2& a,
const double2& b) {
232 return make_double2(lt_mask(a.x, b.x), lt_mask(a.y, b.y));
235EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 pcmp_le<double2>(
const double2& a,
const double2& b) {
236 return make_double2(le_mask(a.x, b.x), le_mask(a.y, b.y));
239EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 pcmp_lt_or_nan<double2>(
const double2& a,
const double2& b) {
240 return make_double2(lt_or_nan_mask(a.x, b.x), lt_or_nan_mask(a.y, b.y));
244EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
bool predux_any(
const float4& a) {
245 using Bits = lane_bits_t<float>;
246 return (numext::bit_cast<Bits>(a.x) | numext::bit_cast<Bits>(a.y) | numext::bit_cast<Bits>(a.z) |
247 numext::bit_cast<Bits>(a.w)) != 0;
251EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
bool predux_any(
const double2& a) {
252 using Bits = lane_bits_t<double>;
253 return (numext::bit_cast<Bits>(a.x) | numext::bit_cast<Bits>(a.y)) != 0;
259EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 psign<float4>(
const float4& a) {
260 return make_float4(numext::sign(a.x), numext::sign(a.y), numext::sign(a.z), numext::sign(a.w));
263EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 psign<double2>(
const double2& a) {
264 return make_double2(numext::sign(a.x), numext::sign(a.y));
268EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 plset<float4>(
const float& a) {
269 return make_float4(a, a + 1, a + 2, a + 3);
272EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 plset<double2>(
const double& a) {
273 return make_double2(a, a + 1);
277EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 padd<float4>(
const float4& a,
const float4& b) {
278 return make_float4(a.x + b.x, a.y + b.y, a.z + b.z, a.w + b.w);
281EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 padd<double2>(
const double2& a,
const double2& b) {
282 return make_double2(a.x + b.x, a.y + b.y);
286EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 psub<float4>(
const float4& a,
const float4& b) {
287 return make_float4(a.x - b.x, a.y - b.y, a.z - b.z, a.w - b.w);
290EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 psub<double2>(
const double2& a,
const double2& b) {
291 return make_double2(a.x - b.x, a.y - b.y);
295EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 pnegate(
const float4& a) {
296 return make_float4(-a.x, -a.y, -a.z, -a.w);
299EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 pnegate(
const double2& a) {
300 return make_double2(-a.x, -a.y);
304EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 pmul<float4>(
const float4& a,
const float4& b) {
305 return make_float4(a.x * b.x, a.y * b.y, a.z * b.z, a.w * b.w);
308EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 pmul<double2>(
const double2& a,
const double2& b) {
309 return make_double2(a.x * b.x, a.y * b.y);
313EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 pdiv<float4>(
const float4& a,
const float4& b) {
314 return make_float4(a.x / b.x, a.y / b.y, a.z / b.z, a.w / b.w);
317EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 pdiv<double2>(
const double2& a,
const double2& b) {
318 return make_double2(a.x / b.x, a.y / b.y);
322EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 pmin<float4>(
const float4& a,
const float4& b) {
323 return make_float4(fminf(a.x, b.x), fminf(a.y, b.y), fminf(a.z, b.z), fminf(a.w, b.w));
326EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 pmin<double2>(
const double2& a,
const double2& b) {
327 return make_double2(fmin(a.x, b.x), fmin(a.y, b.y));
331EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 pmax<float4>(
const float4& a,
const float4& b) {
332 return make_float4(fmaxf(a.x, b.x), fmaxf(a.y, b.y), fmaxf(a.z, b.z), fmaxf(a.w, b.w));
335EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 pmax<double2>(
const double2& a,
const double2& b) {
336 return make_double2(fmax(a.x, b.x), fmax(a.y, b.y));
340EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 pload<float4>(
const float* from) {
341 return *
reinterpret_cast<const float4*
>(from);
345EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 pload<double2>(
const double* from) {
346 return *
reinterpret_cast<const double2*
>(from);
350EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 ploadu<float4>(
const float* from) {
351 return make_float4(from[0], from[1], from[2], from[3]);
354EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 ploadu<double2>(
const double* from) {
355 return make_double2(from[0], from[1]);
359EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 ploaddup<float4>(
const float* from) {
360 return make_float4(from[0], from[0], from[1], from[1]);
363EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 ploaddup<double2>(
const double* from) {
364 return make_double2(from[0], from[0]);
368EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void pstore<float>(
float* to,
const float4& from) {
369 *
reinterpret_cast<float4*
>(to) = from;
373EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void pstore<double>(
double* to,
const double2& from) {
374 *
reinterpret_cast<double2*
>(to) = from;
378EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void pstoreu<float>(
float* to,
const float4& from) {
386EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void pstoreu<double>(
double* to,
const double2& from) {
392EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE float4 ploadt_ro<float4, Aligned>(
const float* from) {
393#if defined(EIGEN_GPU_COMPILE_PHASE)
394 return __ldg(
reinterpret_cast<const float4*
>(from));
396 return make_float4(from[0], from[1], from[2], from[3]);
400EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE double2 ploadt_ro<double2, Aligned>(
const double* from) {
401#if defined(EIGEN_GPU_COMPILE_PHASE)
402 return __ldg(
reinterpret_cast<const double2*
>(from));
404 return make_double2(from[0], from[1]);
409EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE float4 ploadt_ro<float4, Unaligned>(
const float* from) {
410#if defined(EIGEN_GPU_COMPILE_PHASE)
411 return make_float4(__ldg(from + 0), __ldg(from + 1), __ldg(from + 2), __ldg(from + 3));
413 return make_float4(from[0], from[1], from[2], from[3]);
417EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE double2 ploadt_ro<double2, Unaligned>(
const double* from) {
418#if defined(EIGEN_GPU_COMPILE_PHASE)
419 return make_double2(__ldg(from + 0), __ldg(from + 1));
421 return make_double2(from[0], from[1]);
426EIGEN_DEVICE_FUNC
inline float4 pgather<float, float4>(
const float* from, Index stride) {
427 return make_float4(from[0 * stride], from[1 * stride], from[2 * stride], from[3 * stride]);
431EIGEN_DEVICE_FUNC
inline double2 pgather<double, double2>(
const double* from, Index stride) {
432 return make_double2(from[0 * stride], from[1 * stride]);
436EIGEN_DEVICE_FUNC
inline void pscatter<float, float4>(
float* to,
const float4& from, Index stride) {
437 to[stride * 0] = from.x;
438 to[stride * 1] = from.y;
439 to[stride * 2] = from.z;
440 to[stride * 3] = from.w;
443EIGEN_DEVICE_FUNC
inline void pscatter<double, double2>(
double* to,
const double2& from, Index stride) {
444 to[stride * 0] = from.x;
445 to[stride * 1] = from.y;
449EIGEN_DEVICE_FUNC
inline float4 preverse(
const float4& a) {
450 return make_float4(a.w, a.z, a.y, a.x);
453EIGEN_DEVICE_FUNC
inline double2 preverse(
const double2& a) {
454 return make_double2(a.y, a.x);
458EIGEN_DEVICE_FUNC
inline float pfirst<float4>(
const float4& a) {
462EIGEN_DEVICE_FUNC
inline double pfirst<double2>(
const double2& a) {
467EIGEN_DEVICE_FUNC
inline float predux<float4>(
const float4& a) {
468 return a.x + a.y + a.z + a.w;
471EIGEN_DEVICE_FUNC
inline double predux<double2>(
const double2& a) {
476EIGEN_DEVICE_FUNC
inline float predux_max<float4>(
const float4& a) {
477 return fmaxf(fmaxf(a.x, a.y), fmaxf(a.z, a.w));
480EIGEN_DEVICE_FUNC
inline double predux_max<double2>(
const double2& a) {
481 return fmax(a.x, a.y);
485EIGEN_DEVICE_FUNC
inline float predux_min<float4>(
const float4& a) {
486 return fminf(fminf(a.x, a.y), fminf(a.z, a.w));
489EIGEN_DEVICE_FUNC
inline double predux_min<double2>(
const double2& a) {
490 return fmin(a.x, a.y);
494EIGEN_DEVICE_FUNC
inline float predux_mul<float4>(
const float4& a) {
495 return a.x * a.y * a.z * a.w;
498EIGEN_DEVICE_FUNC
inline double predux_mul<double2>(
const double2& a) {
503EIGEN_DEVICE_FUNC
inline float4 pabs<float4>(
const float4& a) {
504 return make_float4(fabsf(a.x), fabsf(a.y), fabsf(a.z), fabsf(a.w));
507EIGEN_DEVICE_FUNC
inline double2 pabs<double2>(
const double2& a) {
508 return make_double2(fabs(a.x), fabs(a.y));
512EIGEN_DEVICE_FUNC
inline float4 pfloor<float4>(
const float4& a) {
513 return make_float4(floorf(a.x), floorf(a.y), floorf(a.z), floorf(a.w));
516EIGEN_DEVICE_FUNC
inline double2 pfloor<double2>(
const double2& a) {
517 return make_double2(floor(a.x), floor(a.y));
521EIGEN_DEVICE_FUNC
inline float4 pceil<float4>(
const float4& a) {
522 return make_float4(ceilf(a.x), ceilf(a.y), ceilf(a.z), ceilf(a.w));
525EIGEN_DEVICE_FUNC
inline double2 pceil<double2>(
const double2& a) {
526 return make_double2(ceil(a.x), ceil(a.y));
530EIGEN_DEVICE_FUNC
inline float4 print<float4>(
const float4& a) {
531 return make_float4(rintf(a.x), rintf(a.y), rintf(a.z), rintf(a.w));
534EIGEN_DEVICE_FUNC
inline double2 print<double2>(
const double2& a) {
535 return make_double2(rint(a.x), rint(a.y));
539EIGEN_DEVICE_FUNC
inline float4 ptrunc<float4>(
const float4& a) {
540 return make_float4(truncf(a.x), truncf(a.y), truncf(a.z), truncf(a.w));
543EIGEN_DEVICE_FUNC
inline double2 ptrunc<double2>(
const double2& a) {
544 return make_double2(trunc(a.x), trunc(a.y));
548EIGEN_DEVICE_FUNC
inline float4 pround<float4>(
const float4& a) {
549 return make_float4(roundf(a.x), roundf(a.y), roundf(a.z), roundf(a.w));
552EIGEN_DEVICE_FUNC
inline double2 pround<double2>(
const double2& a) {
553 return make_double2(round(a.x), round(a.y));
556EIGEN_DEVICE_FUNC
inline void ptranspose(PacketBlock<float4, 4>& kernel) {
557 float tmp = kernel.packet[0].y;
558 kernel.packet[0].y = kernel.packet[1].x;
559 kernel.packet[1].x = tmp;
561 tmp = kernel.packet[0].z;
562 kernel.packet[0].z = kernel.packet[2].x;
563 kernel.packet[2].x = tmp;
565 tmp = kernel.packet[0].w;
566 kernel.packet[0].w = kernel.packet[3].x;
567 kernel.packet[3].x = tmp;
569 tmp = kernel.packet[1].z;
570 kernel.packet[1].z = kernel.packet[2].y;
571 kernel.packet[2].y = tmp;
573 tmp = kernel.packet[1].w;
574 kernel.packet[1].w = kernel.packet[3].y;
575 kernel.packet[3].y = tmp;
577 tmp = kernel.packet[2].w;
578 kernel.packet[2].w = kernel.packet[3].z;
579 kernel.packet[3].z = tmp;
582EIGEN_DEVICE_FUNC
inline void ptranspose(PacketBlock<double2, 2>& kernel) {
583 double tmp = kernel.packet[0].y;
584 kernel.packet[0].y = kernel.packet[1].x;
585 kernel.packet[1].x = tmp;
592#if defined(EIGEN_GPU_COMPILE_PHASE)
596using Packet4h2 = ulonglong2;
598struct unpacket_traits<Packet4h2> {
599 using type = Eigen::half;
600 static constexpr int size = 8;
601 static constexpr int alignment =
Aligned16;
602 static constexpr bool vectorizable =
true;
603 static constexpr bool masked_load_available =
false;
604 static constexpr bool masked_store_available =
false;
605 using half = Packet4h2;
608struct is_arithmetic<Packet4h2> : std::true_type {};
611struct unpacket_traits<half2> {
612 using type = Eigen::half;
613 static constexpr int size = 2;
615 static constexpr int alignment =
Aligned8;
616 static constexpr bool vectorizable =
true;
617 static constexpr bool masked_load_available =
false;
618 static constexpr bool masked_store_available =
false;
622struct is_arithmetic<half2> : std::true_type {};
625struct packet_traits<Eigen::half> : default_packet_traits {
626 using type = Packet4h2;
627 using half = Packet4h2;
628 static constexpr int Vectorizable = 1;
629 static constexpr int AlignedOnScalar = 1;
630 static constexpr int size = 8;
631 static constexpr int HasAdd = 1;
632 static constexpr int HasSub = 1;
633 static constexpr int HasMul = 1;
634 static constexpr int HasDiv = 1;
635 static constexpr int HasSqrt = 1;
636 static constexpr int HasRsqrt = 1;
637 static constexpr int HasExp = 1;
638 static constexpr int HasExpm1 = 1;
639 static constexpr int HasLog = 1;
640 static constexpr int HasLog1p = 1;
643 static constexpr int HasRound = 0;
644 static constexpr int HasSign = 0;
651EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pset1<half2>(
const Eigen::half& from) {
652 return __half2half2(from);
656EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pload<half2>(
const Eigen::half* from) {
657 return *
reinterpret_cast<const half2*
>(from);
661EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 ploadu<half2>(
const Eigen::half* from) {
662 return __halves2half2(from[0], from[1]);
666EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 ploaddup<half2>(
const Eigen::half* from) {
667 return __halves2half2(from[0], from[0]);
671EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void pstore<Eigen::half>(Eigen::half* to,
const half2& from) {
672 *
reinterpret_cast<half2*
>(to) = from;
676EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void pstoreu<Eigen::half>(Eigen::half* to,
const half2& from) {
677 to[0] = __low2half(from);
678 to[1] = __high2half(from);
682EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE half2 ploadt_ro<half2, Aligned>(
const Eigen::half* from) {
684 return __ldg(
reinterpret_cast<const half2*
>(from));
688EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE half2 ploadt_ro<half2, Unaligned>(
const Eigen::half* from) {
689 return __halves2half2(__ldg(from + 0), __ldg(from + 1));
693EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pgather<Eigen::half, half2>(
const Eigen::half* from, Index stride) {
694 return __halves2half2(from[0 * stride], from[1 * stride]);
698EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void pscatter<Eigen::half, half2>(Eigen::half* to,
const half2& from,
700 to[stride * 0] = __low2half(from);
701 to[stride * 1] = __high2half(from);
705EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 preverse(
const half2& a) {
706 return __halves2half2(__high2half(a), __low2half(a));
710EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Eigen::half pfirst<half2>(
const half2& a) {
711 return __low2half(a);
715EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE numext::uint16_t low_bits(
const half2& a) {
716 const Eigen::half low = __low2half(a);
719EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE numext::uint16_t high_bits(
const half2& a) {
720 const Eigen::half high = __high2half(a);
723EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 half2_from_bits(numext::uint16_t low, numext::uint16_t high) {
724 return __halves2half2(half_impl::raw_uint16_to_half(low), half_impl::raw_uint16_to_half(high));
727EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Eigen::half half_mask_from(
bool condition) {
728 return half_impl::raw_uint16_to_half(condition ? 0xffffu : 0x0000u);
732EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pabs<half2>(
const half2& a) {
733 return half2_from_bits(low_bits(a) & 0x7FFF, high_bits(a) & 0x7FFF);
737EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 ptrue<half2>(
const half2& ) {
738 return pset1<half2>(half_mask_from(
true));
742EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pzero<half2>(
const half2& ) {
743 return pset1<half2>(half_mask_from(
false));
746EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void ptranspose(PacketBlock<half2, 2>& kernel) {
747 const Eigen::half a1 = __low2half(kernel.packet[0]);
748 const Eigen::half a2 = __high2half(kernel.packet[0]);
749 const Eigen::half b1 = __low2half(kernel.packet[1]);
750 const Eigen::half b2 = __high2half(kernel.packet[1]);
751 kernel.packet[0] = __halves2half2(a1, b1);
752 kernel.packet[1] = __halves2half2(a2, b2);
756EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 plset<half2>(
const Eigen::half& a) {
757 return __halves2half2(a, __hadd(a, __float2half(1.0f)));
761EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pselect<half2>(
const half2& mask,
const half2& a,
const half2& b) {
762 const Eigen::half mask_low = __low2half(mask);
763 const Eigen::half mask_high = __high2half(mask);
764 const Eigen::half low = mask_low == Eigen::half(0) ? __low2half(b) : __low2half(a);
765 const Eigen::half high = mask_high == Eigen::half(0) ? __high2half(b) : __high2half(a);
766 return __halves2half2(low, high);
771EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pcmp_eq<half2>(
const half2& a,
const half2& b) {
772 return __halves2half2(half_mask_from(__low2float(a) == __low2float(b)),
773 half_mask_from(__high2float(a) == __high2float(b)));
776EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pcmp_lt<half2>(
const half2& a,
const half2& b) {
777 return __halves2half2(half_mask_from(__low2float(a) < __low2float(b)),
778 half_mask_from(__high2float(a) < __high2float(b)));
781EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pcmp_le<half2>(
const half2& a,
const half2& b) {
782 return __halves2half2(half_mask_from(__low2float(a) <= __low2float(b)),
783 half_mask_from(__high2float(a) <= __high2float(b)));
787EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pand<half2>(
const half2& a,
const half2& b) {
788 return half2_from_bits(low_bits(a) & low_bits(b), high_bits(a) & high_bits(b));
791EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 por<half2>(
const half2& a,
const half2& b) {
792 return half2_from_bits(low_bits(a) | low_bits(b), high_bits(a) | high_bits(b));
795EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pxor<half2>(
const half2& a,
const half2& b) {
796 return half2_from_bits(low_bits(a) ^ low_bits(b), high_bits(a) ^ high_bits(b));
799EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pandnot<half2>(
const half2& a,
const half2& b) {
800 return half2_from_bits(low_bits(a) & ~low_bits(b), high_bits(a) & ~high_bits(b));
804EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 padd<half2>(
const half2& a,
const half2& b) {
805 return __hadd2(a, b);
809EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 psub<half2>(
const half2& a,
const half2& b) {
810 return __hsub2(a, b);
814EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pnegate(
const half2& a) {
819EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pmul<half2>(
const half2& a,
const half2& b) {
820 return __hmul2(a, b);
824EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pmadd<half2>(
const half2& a,
const half2& b,
const half2& c) {
825 return __hfma2(a, b, c);
829EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pdiv<half2>(
const half2& a,
const half2& b) {
830 return __h2div(a, b);
837EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pmin<half2>(
const half2& a,
const half2& b) {
838 const __half low = __low2float(a) < __low2float(b) ? __low2half(a) : __low2half(b);
839 const __half high = __high2float(a) < __high2float(b) ? __high2half(a) : __high2half(b);
840 return __halves2half2(low, high);
844EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pmax<half2>(
const half2& a,
const half2& b) {
845 const __half low = __low2float(a) > __low2float(b) ? __low2half(a) : __low2half(b);
846 const __half high = __high2float(a) > __high2float(b) ? __high2half(a) : __high2half(b);
847 return __halves2half2(low, high);
851EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Eigen::half predux<half2>(
const half2& a) {
852 return __hadd(__low2half(a), __high2half(a));
856EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
bool predux_any(
const half2& a) {
857 return (low_bits(a) | high_bits(a)) != 0;
861EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Eigen::half predux_max<half2>(
const half2& a) {
862 const __half first = __low2half(a);
863 const __half second = __high2half(a);
864 return __hgt(first, second) ? first : second;
868EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Eigen::half predux_min<half2>(
const half2& a) {
869 const __half first = __low2half(a);
870 const __half second = __high2half(a);
871 return __hlt(first, second) ? first : second;
875EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Eigen::half predux_mul<half2>(
const half2& a) {
876 return __hmul(__low2half(a), __high2half(a));
880EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 plog<half2>(
const half2& a) {
885EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pexp<half2>(
const half2& a) {
890EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 psqrt<half2>(
const half2& a) {
895EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 prsqrt<half2>(
const half2& a) {
901EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 plog1p<half2>(
const half2& a) {
902 return __floats2half2_rn(log1pf(__low2float(a)), log1pf(__high2float(a)));
906EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pexpm1<half2>(
const half2& a) {
907 return __floats2half2_rn(expm1f(__low2float(a)), expm1f(__high2float(a)));
915EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 lane_half2(
const Packet4h2& p,
int i) {
916 return reinterpret_cast<const half2*
>(&p)[i];
919EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet4h2 make_packet4h2(
const half2& l0,
const half2& l1,
const half2& l2,
922 half2* lanes =
reinterpret_cast<half2*
>(&r);
931#define EIGEN_GPU_PACKET4H2_UNARY(NAME) \
933 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet4h2 NAME<Packet4h2>(const Packet4h2& a) { \
934 return make_packet4h2(NAME(lane_half2(a, 0)), NAME(lane_half2(a, 1)), NAME(lane_half2(a, 2)), \
935 NAME(lane_half2(a, 3))); \
938#define EIGEN_GPU_PACKET4H2_BINARY(NAME) \
940 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet4h2 NAME<Packet4h2>(const Packet4h2& a, const Packet4h2& b) { \
941 return make_packet4h2(NAME(lane_half2(a, 0), lane_half2(b, 0)), NAME(lane_half2(a, 1), lane_half2(b, 1)), \
942 NAME(lane_half2(a, 2), lane_half2(b, 2)), NAME(lane_half2(a, 3), lane_half2(b, 3))); \
946EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet4h2 pset1<Packet4h2>(
const Eigen::half& from) {
947 const half2 lane = pset1<half2>(from);
948 return make_packet4h2(lane, lane, lane, lane);
952EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet4h2 pload<Packet4h2>(
const Eigen::half* from) {
953 return *
reinterpret_cast<const Packet4h2*
>(from);
957EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet4h2 ploadu<Packet4h2>(
const Eigen::half* from) {
958 return make_packet4h2(ploadu<half2>(from + 0), ploadu<half2>(from + 2), ploadu<half2>(from + 4),
959 ploadu<half2>(from + 6));
963EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet4h2 ploaddup<Packet4h2>(
const Eigen::half* from) {
964 return make_packet4h2(ploaddup<half2>(from + 0), ploaddup<half2>(from + 1), ploaddup<half2>(from + 2),
965 ploaddup<half2>(from + 3));
969EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void pstore<Eigen::half>(Eigen::half* to,
const Packet4h2& from) {
970 *
reinterpret_cast<Packet4h2*
>(to) = from;
974EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void pstoreu<Eigen::half>(Eigen::half* to,
const Packet4h2& from) {
975 pstoreu<Eigen::half>(to + 0, lane_half2(from, 0));
976 pstoreu<Eigen::half>(to + 2, lane_half2(from, 1));
977 pstoreu<Eigen::half>(to + 4, lane_half2(from, 2));
978 pstoreu<Eigen::half>(to + 6, lane_half2(from, 3));
982EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet4h2 ploadt_ro<Packet4h2, Aligned>(
const Eigen::half* from) {
983 return __ldg(
reinterpret_cast<const Packet4h2*
>(from));
987EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet4h2 ploadt_ro<Packet4h2, Unaligned>(
const Eigen::half* from) {
988 return make_packet4h2(ploadt_ro<half2, Unaligned>(from + 0), ploadt_ro<half2, Unaligned>(from + 2),
989 ploadt_ro<half2, Unaligned>(from + 4), ploadt_ro<half2, Unaligned>(from + 6));
993EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet4h2 pgather<Eigen::half, Packet4h2>(
const Eigen::half* from, Index stride) {
994 return make_packet4h2(
995 __halves2half2(from[0 * stride], from[1 * stride]), __halves2half2(from[2 * stride], from[3 * stride]),
996 __halves2half2(from[4 * stride], from[5 * stride]), __halves2half2(from[6 * stride], from[7 * stride]));
1000EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void pscatter<Eigen::half, Packet4h2>(Eigen::half* to,
const Packet4h2& from,
1002 pscatter<Eigen::half, half2>(to + stride * 0, lane_half2(from, 0), stride);
1003 pscatter<Eigen::half, half2>(to + stride * 2, lane_half2(from, 1), stride);
1004 pscatter<Eigen::half, half2>(to + stride * 4, lane_half2(from, 2), stride);
1005 pscatter<Eigen::half, half2>(to + stride * 6, lane_half2(from, 3), stride);
1009EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet4h2 preverse(
const Packet4h2& a) {
1010 return make_packet4h2(preverse(lane_half2(a, 3)), preverse(lane_half2(a, 2)), preverse(lane_half2(a, 1)),
1011 preverse(lane_half2(a, 0)));
1015EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Eigen::half pfirst<Packet4h2>(
const Packet4h2& a) {
1016 return pfirst<half2>(lane_half2(a, 0));
1020EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet4h2 ptrue<Packet4h2>(
const Packet4h2& ) {
1021 return pset1<Packet4h2>(half_impl::raw_uint16_to_half(0xffffu));
1025EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet4h2 pzero<Packet4h2>(
const Packet4h2& ) {
1026 return pset1<Packet4h2>(half_impl::raw_uint16_to_half(0x0000u));
1030EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void ptranspose(PacketBlock<Packet4h2, 8>& kernel) {
1031 Eigen::half elements[8][8];
1033 for (
int row = 0; row < 8; ++row) {
1035 for (
int lane = 0; lane < 4; ++lane) {
1036 const half2 pair = lane_half2(kernel.packet[row], lane);
1037 elements[row][2 * lane] = __low2half(pair);
1038 elements[row][2 * lane + 1] = __high2half(pair);
1042 for (
int row = 0; row < 8; ++row) {
1043 kernel.packet[row] = make_packet4h2(
1044 __halves2half2(elements[0][row], elements[1][row]), __halves2half2(elements[2][row], elements[3][row]),
1045 __halves2half2(elements[4][row], elements[5][row]), __halves2half2(elements[6][row], elements[7][row]));
1050EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet4h2 plset<Packet4h2>(
const Eigen::half& a) {
1052 const half2 base = pset1<half2>(a);
1053 return make_packet4h2(plset<half2>(a), __hadd2(base, __floats2half2_rn(2.0f, 3.0f)),
1054 __hadd2(base, __floats2half2_rn(4.0f, 5.0f)), __hadd2(base, __floats2half2_rn(6.0f, 7.0f)));
1058EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet4h2 pselect<Packet4h2>(
const Packet4h2& mask,
const Packet4h2& a,
1059 const Packet4h2& b) {
1060 return make_packet4h2(pselect<half2>(lane_half2(mask, 0), lane_half2(a, 0), lane_half2(b, 0)),
1061 pselect<half2>(lane_half2(mask, 1), lane_half2(a, 1), lane_half2(b, 1)),
1062 pselect<half2>(lane_half2(mask, 2), lane_half2(a, 2), lane_half2(b, 2)),
1063 pselect<half2>(lane_half2(mask, 3), lane_half2(a, 3), lane_half2(b, 3)));
1067EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet4h2 pmadd<Packet4h2>(
const Packet4h2& a,
const Packet4h2& b,
1068 const Packet4h2& c) {
1069 return make_packet4h2(pmadd<half2>(lane_half2(a, 0), lane_half2(b, 0), lane_half2(c, 0)),
1070 pmadd<half2>(lane_half2(a, 1), lane_half2(b, 1), lane_half2(c, 1)),
1071 pmadd<half2>(lane_half2(a, 2), lane_half2(b, 2), lane_half2(c, 2)),
1072 pmadd<half2>(lane_half2(a, 3), lane_half2(b, 3), lane_half2(c, 3)));
1075EIGEN_GPU_PACKET4H2_UNARY(pabs)
1076EIGEN_GPU_PACKET4H2_UNARY(pnegate)
1077EIGEN_GPU_PACKET4H2_UNARY(plog)
1078EIGEN_GPU_PACKET4H2_UNARY(pexp)
1079EIGEN_GPU_PACKET4H2_UNARY(psqrt)
1080EIGEN_GPU_PACKET4H2_UNARY(prsqrt)
1081EIGEN_GPU_PACKET4H2_UNARY(plog1p)
1082EIGEN_GPU_PACKET4H2_UNARY(pexpm1)
1084EIGEN_GPU_PACKET4H2_BINARY(padd)
1085EIGEN_GPU_PACKET4H2_BINARY(psub)
1086EIGEN_GPU_PACKET4H2_BINARY(pmul)
1087EIGEN_GPU_PACKET4H2_BINARY(pdiv)
1088EIGEN_GPU_PACKET4H2_BINARY(pmin)
1089EIGEN_GPU_PACKET4H2_BINARY(pmax)
1090EIGEN_GPU_PACKET4H2_BINARY(pand)
1091EIGEN_GPU_PACKET4H2_BINARY(por)
1092EIGEN_GPU_PACKET4H2_BINARY(pxor)
1093EIGEN_GPU_PACKET4H2_BINARY(pandnot)
1094EIGEN_GPU_PACKET4H2_BINARY(pcmp_eq)
1095EIGEN_GPU_PACKET4H2_BINARY(pcmp_lt)
1096EIGEN_GPU_PACKET4H2_BINARY(pcmp_le)
1098#undef EIGEN_GPU_PACKET4H2_UNARY
1099#undef EIGEN_GPU_PACKET4H2_BINARY
1102EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Eigen::half predux<Packet4h2>(
const Packet4h2& a) {
1103 return predux(lane_half2(a, 0)) + predux(lane_half2(a, 1)) + predux(lane_half2(a, 2)) + predux(lane_half2(a, 3));
1107EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
bool predux_any(
const Packet4h2& a) {
1108 return predux_any(lane_half2(a, 0)) | predux_any(lane_half2(a, 1)) | predux_any(lane_half2(a, 2)) |
1109 predux_any(lane_half2(a, 3));
1113EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Eigen::half predux_max<Packet4h2>(
const Packet4h2& a) {
1115 pmax<half2>(pmax<half2>(lane_half2(a, 0), lane_half2(a, 1)), pmax<half2>(lane_half2(a, 2), lane_half2(a, 3)));
1116 return predux_max<half2>(m);
1120EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Eigen::half predux_min<Packet4h2>(
const Packet4h2& a) {
1122 pmin<half2>(pmin<half2>(lane_half2(a, 0), lane_half2(a, 1)), pmin<half2>(lane_half2(a, 2), lane_half2(a, 3)));
1123 return predux_min<half2>(m);
1128EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Eigen::half predux_mul<Packet4h2>(
const Packet4h2& a) {
1129 const half2 product =
1130 pmul<half2>(pmul<half2>(lane_half2(a, 0), lane_half2(a, 1)), pmul<half2>(lane_half2(a, 2), lane_half2(a, 3)));
1131 return predux_mul<half2>(product);
@ Aligned8
Definition Constants.h:237
@ Aligned16
Definition Constants.h:238