Eigen  5.0.1
 
Loading...
Searching...
No Matches
PacketMath.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2014 Benoit Steiner <benoit.steiner.goog@gmail.com>
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// SPDX-License-Identifier: MPL-2.0
10
11#ifndef EIGEN_PACKET_MATH_GPU_H
12#define EIGEN_PACKET_MATH_GPU_H
13
14// IWYU pragma: private
15#include "../../InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20
21// Read-only data cached load (__ldg) and native FP16 arithmetic are available
22// on all supported GPU architectures (sm_60+ for CUDA, GFX906+ for HIP).
23
24// Make sure this is only available when targeting a GPU: we don't want to
25// introduce conflicts between these packet_traits definitions and the ones
26// we'll use on the host side (SSE, AVX, ...)
27#if defined(EIGEN_GPUCC) && defined(EIGEN_USE_GPU)
28
29template <>
30struct is_arithmetic<float4> : std::true_type {};
31template <>
32struct is_arithmetic<double2> : std::true_type {};
33
34template <>
35struct packet_traits<float> : default_packet_traits {
36 using type = float4;
37 using half = float4;
38 static constexpr int Vectorizable = 1;
39 static constexpr int AlignedOnScalar = 1;
40 static constexpr int size = 4;
41
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;
62
63 static constexpr int HasCmp = 1;
64};
65
66template <>
67struct packet_traits<double> : default_packet_traits {
68 using type = double2;
69 using half = double2;
70 static constexpr int Vectorizable = 1;
71 static constexpr int AlignedOnScalar = 1;
72 static constexpr int size = 2;
73
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;
92
93 static constexpr int HasCmp = 1;
94};
95
96template <>
97struct unpacket_traits<float4> {
98 using type = float;
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;
104 using half = float4;
105};
106template <>
107struct unpacket_traits<double2> {
108 using type = double;
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;
115};
116
117template <>
118EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 pset1<float4>(const float& from) {
119 return make_float4(from, from, from, from);
120}
121template <>
122EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 pset1<double2>(const double& from) {
123 return make_double2(from, from);
124}
125
126// Bit-level helpers on the scalar lanes. numext::bit_cast rather than the __int_as_float family, which are device
127// intrinsics: a .cu translation unit that defines EIGEN_USE_GPU instantiates these packet types in the host pass
128// too, and an operation that exists in only one of the two passes makes packet_traits differ between them.
129template <typename T>
130using lane_bits_t = typename numext::get_integer_by_size<sizeof(T)>::unsigned_type;
131
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))); \
137 }
138
139EIGEN_MAKE_BITWISE_BINOP(and, &)
140EIGEN_MAKE_BITWISE_BINOP(or, |)
141EIGEN_MAKE_BITWISE_BINOP(xor, ^)
142EIGEN_MAKE_BITWISE_BINOP(andnot, &~)
143
144#undef EIGEN_MAKE_BITWISE_BINOP
145
146// A comparison returns an all-ones lane where it holds and an all-zero lane elsewhere, so that the result can be
147// consumed bitwise by pselect, pand and pandnot.
148template <typename T>
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));
152}
153
154template <typename T>
155EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T eq_mask(const T& a, const T& b) {
156 return mask_from<T>(a == b);
157}
158template <typename T>
159EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T lt_mask(const T& a, const T& b) {
160 return mask_from<T>(a < b);
161}
162template <typename T>
163EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T le_mask(const T& a, const T& b) {
164 return mask_from<T>(a <= b);
165}
166// !(a >= b), so a NaN operand makes the lane true.
167template <typename T>
168EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T lt_or_nan_mask(const T& a, const T& b) {
169 return mask_from<T>(!(a >= b));
170}
171
172template <>
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));
175}
176template <>
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));
179}
180
181template <>
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));
184}
185template <>
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));
188}
189
190template <>
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));
193}
194template <>
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));
197}
198
199template <>
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));
203}
204template <>
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));
207}
208
209template <>
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));
212}
213template <>
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));
216}
217template <>
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));
220}
221template <>
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));
225}
226template <>
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));
229}
230template <>
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));
233}
234template <>
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));
237}
238template <>
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));
241}
242
243template <>
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;
248}
249
250template <>
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;
254}
255
256// numext::sign: 0 for either zero, NaN for NaN, +/-1 otherwise. default_packet_traits advertises HasSign, so
257// without these the flag was a promise the backend did not keep (the generic form does not compile for float4).
258template <>
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));
261}
262template <>
263EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 psign<double2>(const double2& a) {
264 return make_double2(numext::sign(a.x), numext::sign(a.y));
265}
266
267template <>
268EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 plset<float4>(const float& a) {
269 return make_float4(a, a + 1, a + 2, a + 3);
270}
271template <>
272EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 plset<double2>(const double& a) {
273 return make_double2(a, a + 1);
274}
275
276template <>
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);
279}
280template <>
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);
283}
284
285template <>
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);
288}
289template <>
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);
292}
293
294template <>
295EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 pnegate(const float4& a) {
296 return make_float4(-a.x, -a.y, -a.z, -a.w);
297}
298template <>
299EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 pnegate(const double2& a) {
300 return make_double2(-a.x, -a.y);
301}
302
303template <>
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);
306}
307template <>
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);
310}
311
312template <>
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);
315}
316template <>
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);
319}
320
321template <>
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));
324}
325template <>
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));
328}
329
330template <>
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));
333}
334template <>
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));
337}
338
339template <>
340EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 pload<float4>(const float* from) {
341 return *reinterpret_cast<const float4*>(from);
342}
343
344template <>
345EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 pload<double2>(const double* from) {
346 return *reinterpret_cast<const double2*>(from);
347}
348
349template <>
350EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 ploadu<float4>(const float* from) {
351 return make_float4(from[0], from[1], from[2], from[3]);
352}
353template <>
354EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 ploadu<double2>(const double* from) {
355 return make_double2(from[0], from[1]);
356}
357
358template <>
359EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float4 ploaddup<float4>(const float* from) {
360 return make_float4(from[0], from[0], from[1], from[1]);
361}
362template <>
363EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE double2 ploaddup<double2>(const double* from) {
364 return make_double2(from[0], from[0]);
365}
366
367template <>
368EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void pstore<float>(float* to, const float4& from) {
369 *reinterpret_cast<float4*>(to) = from;
370}
371
372template <>
373EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void pstore<double>(double* to, const double2& from) {
374 *reinterpret_cast<double2*>(to) = from;
375}
376
377template <>
378EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void pstoreu<float>(float* to, const float4& from) {
379 to[0] = from.x;
380 to[1] = from.y;
381 to[2] = from.z;
382 to[3] = from.w;
383}
384
385template <>
386EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void pstoreu<double>(double* to, const double2& from) {
387 to[0] = from.x;
388 to[1] = from.y;
389}
390
391template <>
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));
395#else
396 return make_float4(from[0], from[1], from[2], from[3]);
397#endif
398}
399template <>
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));
403#else
404 return make_double2(from[0], from[1]);
405#endif
406}
407
408template <>
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));
412#else
413 return make_float4(from[0], from[1], from[2], from[3]);
414#endif
415}
416template <>
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));
420#else
421 return make_double2(from[0], from[1]);
422#endif
423}
424
425template <>
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]);
428}
429
430template <>
431EIGEN_DEVICE_FUNC inline double2 pgather<double, double2>(const double* from, Index stride) {
432 return make_double2(from[0 * stride], from[1 * stride]);
433}
434
435template <>
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;
441}
442template <>
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;
446}
447
448template <>
449EIGEN_DEVICE_FUNC inline float4 preverse(const float4& a) {
450 return make_float4(a.w, a.z, a.y, a.x);
451}
452template <>
453EIGEN_DEVICE_FUNC inline double2 preverse(const double2& a) {
454 return make_double2(a.y, a.x);
455}
456
457template <>
458EIGEN_DEVICE_FUNC inline float pfirst<float4>(const float4& a) {
459 return a.x;
460}
461template <>
462EIGEN_DEVICE_FUNC inline double pfirst<double2>(const double2& a) {
463 return a.x;
464}
465
466template <>
467EIGEN_DEVICE_FUNC inline float predux<float4>(const float4& a) {
468 return a.x + a.y + a.z + a.w;
469}
470template <>
471EIGEN_DEVICE_FUNC inline double predux<double2>(const double2& a) {
472 return a.x + a.y;
473}
474
475template <>
476EIGEN_DEVICE_FUNC inline float predux_max<float4>(const float4& a) {
477 return fmaxf(fmaxf(a.x, a.y), fmaxf(a.z, a.w));
478}
479template <>
480EIGEN_DEVICE_FUNC inline double predux_max<double2>(const double2& a) {
481 return fmax(a.x, a.y);
482}
483
484template <>
485EIGEN_DEVICE_FUNC inline float predux_min<float4>(const float4& a) {
486 return fminf(fminf(a.x, a.y), fminf(a.z, a.w));
487}
488template <>
489EIGEN_DEVICE_FUNC inline double predux_min<double2>(const double2& a) {
490 return fmin(a.x, a.y);
491}
492
493template <>
494EIGEN_DEVICE_FUNC inline float predux_mul<float4>(const float4& a) {
495 return a.x * a.y * a.z * a.w;
496}
497template <>
498EIGEN_DEVICE_FUNC inline double predux_mul<double2>(const double2& a) {
499 return a.x * a.y;
500}
501
502template <>
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));
505}
506template <>
507EIGEN_DEVICE_FUNC inline double2 pabs<double2>(const double2& a) {
508 return make_double2(fabs(a.x), fabs(a.y));
509}
510
511template <>
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));
514}
515template <>
516EIGEN_DEVICE_FUNC inline double2 pfloor<double2>(const double2& a) {
517 return make_double2(floor(a.x), floor(a.y));
518}
519
520template <>
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));
523}
524template <>
525EIGEN_DEVICE_FUNC inline double2 pceil<double2>(const double2& a) {
526 return make_double2(ceil(a.x), ceil(a.y));
527}
528
529template <>
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));
532}
533template <>
534EIGEN_DEVICE_FUNC inline double2 print<double2>(const double2& a) {
535 return make_double2(rint(a.x), rint(a.y));
536}
537
538template <>
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));
541}
542template <>
543EIGEN_DEVICE_FUNC inline double2 ptrunc<double2>(const double2& a) {
544 return make_double2(trunc(a.x), trunc(a.y));
545}
546
547template <>
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));
550}
551template <>
552EIGEN_DEVICE_FUNC inline double2 pround<double2>(const double2& a) {
553 return make_double2(round(a.x), round(a.y));
554}
555
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;
560
561 tmp = kernel.packet[0].z;
562 kernel.packet[0].z = kernel.packet[2].x;
563 kernel.packet[2].x = tmp;
564
565 tmp = kernel.packet[0].w;
566 kernel.packet[0].w = kernel.packet[3].x;
567 kernel.packet[3].x = tmp;
568
569 tmp = kernel.packet[1].z;
570 kernel.packet[1].z = kernel.packet[2].y;
571 kernel.packet[2].y = tmp;
572
573 tmp = kernel.packet[1].w;
574 kernel.packet[1].w = kernel.packet[3].y;
575 kernel.packet[3].y = tmp;
576
577 tmp = kernel.packet[2].w;
578 kernel.packet[2].w = kernel.packet[3].z;
579 kernel.packet[3].z = tmp;
580}
581
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;
586}
587
588#endif // defined(EIGEN_GPUCC) && defined(EIGEN_USE_GPU)
589
590// Half-packet functions are only available in GPU device compilation — they use
591// intrinsics (__half2, etc.) that have no host-side benefit.
592#if defined(EIGEN_GPU_COMPILE_PHASE)
593
594// Two packets over Eigen::half: the native two-lane half2, and Packet4h2, eight halves in four half2 lanes.
595// Packet4h2 stays an alias of ulonglong2 because the Tensor GPU kernels and TensorFlow name it that way.
596using Packet4h2 = ulonglong2;
597template <>
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;
606};
607template <>
608struct is_arithmetic<Packet4h2> : std::true_type {};
609
610template <>
611struct unpacket_traits<half2> {
612 using type = Eigen::half;
613 static constexpr int size = 2;
614 // half2 needs 4-byte alignment; Aligned8 is the smallest value the enum offers that satisfies it.
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;
619 using half = half2;
620};
621template <>
622struct is_arithmetic<half2> : std::true_type {};
623
624template <>
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;
641 // default_packet_traits turns these on, but there is no pfloor/pceil/print/ptrunc/pround and no psign for the
642 // half packets; advertising them makes the generic fallback the evaluator picks fail to compile.
643 static constexpr int HasRound = 0;
644 static constexpr int HasSign = 0;
645};
646
647// ---------------------------------------------------------------------------------------------------------------
648// half2, the native two-lane packet.
649
650template <>
651EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pset1<half2>(const Eigen::half& from) {
652 return __half2half2(from);
653}
654
655template <>
656EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pload<half2>(const Eigen::half* from) {
657 return *reinterpret_cast<const half2*>(from);
658}
659
660template <>
661EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 ploadu<half2>(const Eigen::half* from) {
662 return __halves2half2(from[0], from[1]);
663}
664
665template <>
666EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 ploaddup<half2>(const Eigen::half* from) {
667 return __halves2half2(from[0], from[0]);
668}
669
670template <>
671EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void pstore<Eigen::half>(Eigen::half* to, const half2& from) {
672 *reinterpret_cast<half2*>(to) = from;
673}
674
675template <>
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);
679}
680
681template <>
682EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE half2 ploadt_ro<half2, Aligned>(const Eigen::half* from) {
683 // Input is guaranteed to be properly aligned.
684 return __ldg(reinterpret_cast<const half2*>(from));
685}
686
687template <>
688EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE half2 ploadt_ro<half2, Unaligned>(const Eigen::half* from) {
689 return __halves2half2(__ldg(from + 0), __ldg(from + 1));
690}
691
692template <>
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]);
695}
696
697template <>
698EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void pscatter<Eigen::half, half2>(Eigen::half* to, const half2& from,
699 Index stride) {
700 to[stride * 0] = __low2half(from);
701 to[stride * 1] = __high2half(from);
702}
703
704template <>
705EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 preverse(const half2& a) {
706 return __halves2half2(__high2half(a), __low2half(a));
707}
708
709template <>
710EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Eigen::half pfirst<half2>(const half2& a) {
711 return __low2half(a);
712}
713
714// __low2half returns the CUDA type; Eigen::half is what carries the accessible raw field.
715EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE numext::uint16_t low_bits(const half2& a) {
716 const Eigen::half low = __low2half(a);
717 return low.x;
718}
719EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE numext::uint16_t high_bits(const half2& a) {
720 const Eigen::half high = __high2half(a);
721 return high.x;
722}
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));
725}
726// An all-ones or all-zero lane: the mask form pselect and the bitwise operations consume.
727EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Eigen::half half_mask_from(bool condition) {
728 return half_impl::raw_uint16_to_half(condition ? 0xffffu : 0x0000u);
729}
730
731template <>
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);
734}
735
736template <>
737EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 ptrue<half2>(const half2& /*a*/) {
738 return pset1<half2>(half_mask_from(true));
739}
740
741template <>
742EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pzero<half2>(const half2& /*a*/) {
743 return pset1<half2>(half_mask_from(false));
744}
745
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);
753}
754
755template <>
756EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 plset<half2>(const Eigen::half& a) {
757 return __halves2half2(a, __hadd(a, __float2half(1.0f)));
758}
759
760template <>
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);
767}
768
769// Conversion to float is exact for both half operands.
770template <>
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)));
774}
775template <>
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)));
779}
780template <>
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)));
784}
785
786template <>
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));
789}
790template <>
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));
793}
794template <>
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));
797}
798template <>
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));
801}
802
803template <>
804EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 padd<half2>(const half2& a, const half2& b) {
805 return __hadd2(a, b);
806}
807
808template <>
809EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 psub<half2>(const half2& a, const half2& b) {
810 return __hsub2(a, b);
811}
812
813template <>
814EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pnegate(const half2& a) {
815 return __hneg2(a);
816}
817
818template <>
819EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pmul<half2>(const half2& a, const half2& b) {
820 return __hmul2(a, b);
821}
822
823template <>
824EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pmadd<half2>(const half2& a, const half2& b, const half2& c) {
825 return __hfma2(a, b, c);
826}
827
828template <>
829EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pdiv<half2>(const half2& a, const half2& b) {
830 return __h2div(a, b);
831}
832
833// Compared through float, and every comparison against a NaN is false, so the result is b. This differs from
834// the fminf/fmaxf of the float packets, which return the non-NaN operand: here a NaN in b propagates and a
835// NaN in a does not.
836template <>
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);
841}
842
843template <>
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);
848}
849
850template <>
851EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Eigen::half predux<half2>(const half2& a) {
852 return __hadd(__low2half(a), __high2half(a));
853}
854
855template <>
856EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool predux_any(const half2& a) {
857 return (low_bits(a) | high_bits(a)) != 0;
858}
859
860template <>
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;
865}
866
867template <>
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;
872}
873
874template <>
875EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Eigen::half predux_mul<half2>(const half2& a) {
876 return __hmul(__low2half(a), __high2half(a));
877}
878
879template <>
880EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 plog<half2>(const half2& a) {
881 return h2log(a);
882}
883
884template <>
885EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pexp<half2>(const half2& a) {
886 return h2exp(a);
887}
888
889template <>
890EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 psqrt<half2>(const half2& a) {
891 return h2sqrt(a);
892}
893
894template <>
895EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 prsqrt<half2>(const half2& a) {
896 return h2rsqrt(a);
897}
898
899// No native h2log1p/h2expm1; both go through float.
900template <>
901EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 plog1p<half2>(const half2& a) {
902 return __floats2half2_rn(log1pf(__low2float(a)), log1pf(__high2float(a)));
903}
904
905template <>
906EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 pexpm1<half2>(const half2& a) {
907 return __floats2half2_rn(expm1f(__low2float(a)), expm1f(__high2float(a)));
908}
909
910// ---------------------------------------------------------------------------------------------------------------
911// Packet4h2 is four half2 lanes in the two words of a ulonglong2, the alias the Tensor GPU kernels expect. Do not
912// replace the cast: reaching the lanes conformingly (shift and reassemble, or __half2_raw) stops the compiler
913// keeping them in registers, costing 168 SASS instructions against 152 for eight mixed half operations and 152
914// against 64 for an 8x8 ptranspose (sm_89, nvcc 13.3).
915EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE half2 lane_half2(const Packet4h2& p, int i) {
916 return reinterpret_cast<const half2*>(&p)[i];
917}
918
919EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet4h2 make_packet4h2(const half2& l0, const half2& l1, const half2& l2,
920 const half2& l3) {
921 Packet4h2 r;
922 half2* lanes = reinterpret_cast<half2*>(&r);
923 lanes[0] = l0;
924 lanes[1] = l1;
925 lanes[2] = l2;
926 lanes[3] = l3;
927 return r;
928}
929
930// Every lane-wise operation is its half2 form applied to the four lanes; the macros keep that unroll in one place.
931#define EIGEN_GPU_PACKET4H2_UNARY(NAME) \
932 template <> \
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))); \
936 }
937
938#define EIGEN_GPU_PACKET4H2_BINARY(NAME) \
939 template <> \
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))); \
943 }
944
945template <>
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);
949}
950
951template <>
952EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet4h2 pload<Packet4h2>(const Eigen::half* from) {
953 return *reinterpret_cast<const Packet4h2*>(from);
954}
955
956template <>
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));
960}
961
962template <>
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));
966}
967
968template <>
969EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void pstore<Eigen::half>(Eigen::half* to, const Packet4h2& from) {
970 *reinterpret_cast<Packet4h2*>(to) = from;
971}
972
973template <>
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));
979}
980
981template <>
982EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet4h2 ploadt_ro<Packet4h2, Aligned>(const Eigen::half* from) {
983 return __ldg(reinterpret_cast<const Packet4h2*>(from));
984}
985
986template <>
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));
990}
991
992template <>
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]));
997}
998
999template <>
1000EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void pscatter<Eigen::half, Packet4h2>(Eigen::half* to, const Packet4h2& from,
1001 Index stride) {
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);
1006}
1007
1008template <>
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)));
1012}
1013
1014template <>
1015EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Eigen::half pfirst<Packet4h2>(const Packet4h2& a) {
1016 return pfirst<half2>(lane_half2(a, 0));
1017}
1018
1019template <>
1020EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet4h2 ptrue<Packet4h2>(const Packet4h2& /*a*/) {
1021 return pset1<Packet4h2>(half_impl::raw_uint16_to_half(0xffffu));
1022}
1023
1024template <>
1025EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet4h2 pzero<Packet4h2>(const Packet4h2& /*a*/) {
1026 return pset1<Packet4h2>(half_impl::raw_uint16_to_half(0x0000u));
1027}
1028
1029// An 8x8 transpose of halves: packet r of the result is column r of the input.
1030EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void ptranspose(PacketBlock<Packet4h2, 8>& kernel) {
1031 Eigen::half elements[8][8];
1032 EIGEN_UNROLL_LOOP
1033 for (int row = 0; row < 8; ++row) {
1034 EIGEN_UNROLL_LOOP
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);
1039 }
1040 }
1041 EIGEN_UNROLL_LOOP
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]));
1046 }
1047}
1048
1049template <>
1050EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet4h2 plset<Packet4h2>(const Eigen::half& a) {
1051 // Add each offset once: (a + 2*k) + 1 can round twice. Preserve a's sign in lane zero.
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)));
1055}
1056
1057template <>
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)));
1064}
1065
1066template <>
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)));
1073}
1074
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)
1083
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)
1097
1098#undef EIGEN_GPU_PACKET4H2_UNARY
1099#undef EIGEN_GPU_PACKET4H2_BINARY
1100
1101template <>
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));
1104}
1105
1106template <>
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));
1110}
1111
1112template <>
1113EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Eigen::half predux_max<Packet4h2>(const Packet4h2& a) {
1114 const half2 m =
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);
1117}
1118
1119template <>
1120EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Eigen::half predux_min<Packet4h2>(const Packet4h2& a) {
1121 const half2 m =
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);
1124}
1125
1126// Likely to overflow or underflow: eight halves multiplied together.
1127template <>
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);
1132}
1133
1134#endif // defined(EIGEN_GPU_COMPILE_PHASE)
1135
1136} // end namespace internal
1137
1138} // end namespace Eigen
1139
1140#endif // EIGEN_PACKET_MATH_GPU_H
@ Aligned8
Definition Constants.h:237
@ Aligned16
Definition Constants.h:238