Eigen  5.0.1
 
Loading...
Searching...
No Matches
PacketMathFP16.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2025 The Eigen Authors.
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_FP16_AVX512_H
12#define EIGEN_PACKET_MATH_FP16_AVX512_H
13
14// IWYU pragma: private
15#include "../../InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20
21typedef __m512h Packet32h;
22typedef __m256h Packet16h;
23typedef __m128h Packet8h;
24
25template <>
26struct is_arithmetic<Packet8h> {
27 enum { value = true };
28};
29
30template <>
31struct packet_traits<half> : default_packet_traits {
32 typedef Packet32h type;
33 typedef Packet16h half;
34 enum {
35 Vectorizable = 1,
36 AlignedOnScalar = 1,
37 size = 32,
38
39 HasCmp = 1,
40 HasAdd = 1,
41 HasSub = 1,
42 HasMul = 1,
43 HasDiv = 1,
44 HasNegate = 1,
45 HasAbs = 1,
46 HasMin = 1,
47 HasMax = 1,
48 HasConj = 1,
49 HasSetLinear = 0,
50 HasLog = 1,
51 HasLog1p = 1,
52 HasExp = 1,
53 HasExpm1 = 1,
54 HasSqrt = 1,
55 HasRsqrt = 1,
56 // These ones should be implemented in future
57 HasBessel = 0,
58 HasNdtri = 0,
59 HasSin = EIGEN_FAST_MATH,
60 HasCos = EIGEN_FAST_MATH,
61 HasTanh = EIGEN_FAST_MATH,
62 HasErf = 0, // EIGEN_FAST_MATH,
63 };
64};
65
66template <>
67struct unpacket_traits<Packet32h> {
68 typedef Eigen::half type;
69 typedef Packet16h half;
70 typedef Packet32s integer_packet;
71 enum {
72 size = 32,
73 alignment = Aligned64,
74 vectorizable = true,
75 masked_load_available = false,
76 masked_store_available = false
77 };
78};
79
80template <>
81struct unpacket_traits<Packet16h> {
82 typedef Eigen::half type;
83 typedef Packet8h half;
84 typedef Packet16s integer_packet;
85 enum {
86 size = 16,
87 alignment = Aligned32,
88 vectorizable = true,
89 masked_load_available = false,
90 masked_store_available = false
91 };
92};
93
94template <>
95struct unpacket_traits<Packet8h> {
96 typedef Eigen::half type;
97 typedef Packet8h half;
98 typedef Packet8s integer_packet;
99 enum {
100 size = 8,
101 alignment = Aligned16,
102 vectorizable = true,
103 masked_load_available = false,
104 masked_store_available = false
105 };
106};
107
108// Conversions
109
110EIGEN_STRONG_INLINE Packet16f half2float(const Packet16h& a) { return _mm512_cvtxph_ps(a); }
111
112EIGEN_STRONG_INLINE Packet8f half2float(const Packet8h& a) { return _mm256_cvtxph_ps(a); }
113
114EIGEN_STRONG_INLINE Packet16h float2half(const Packet16f& a) { return _mm512_cvtxps_ph(a); }
115
116EIGEN_STRONG_INLINE Packet8h float2half(const Packet8f& a) { return _mm256_cvtxps_ph(a); }
117
118// Memory functions
119
120// pset1
121
122template <>
123EIGEN_STRONG_INLINE Packet32h pset1<Packet32h>(const Eigen::half& from) {
124 return _mm512_set1_ph(from.x);
125}
126
127template <>
128EIGEN_STRONG_INLINE Packet16h pset1<Packet16h>(const Eigen::half& from) {
129 return _mm256_set1_ph(from.x);
130}
131
132template <>
133EIGEN_STRONG_INLINE Packet8h pset1<Packet8h>(const Eigen::half& from) {
134 return _mm_set1_ph(from.x);
135}
136
137template <>
138EIGEN_STRONG_INLINE Packet32h pzero(const Packet32h& /*a*/) {
139 return _mm512_setzero_ph();
140}
141
142template <>
143EIGEN_STRONG_INLINE Packet16h pzero(const Packet16h& /*a*/) {
144 return _mm256_setzero_ph();
145}
146
147template <>
148EIGEN_STRONG_INLINE Packet8h pzero(const Packet8h& /*a*/) {
149 return _mm_setzero_ph();
150}
151
152// pset1frombits
153template <>
154EIGEN_STRONG_INLINE Packet32h pset1frombits<Packet32h>(unsigned short from) {
155 return _mm512_castsi512_ph(_mm512_set1_epi16(from));
156}
157
158template <>
159EIGEN_STRONG_INLINE Packet16h pset1frombits<Packet16h>(unsigned short from) {
160 return _mm256_castsi256_ph(_mm256_set1_epi16(from));
161}
162
163template <>
164EIGEN_STRONG_INLINE Packet8h pset1frombits<Packet8h>(unsigned short from) {
165 return _mm_castsi128_ph(_mm_set1_epi16(from));
166}
167
168// pfirst
169
170template <>
171EIGEN_STRONG_INLINE Eigen::half pfirst<Packet32h>(const Packet32h& from) {
172 return Eigen::half(_mm512_cvtsh_h(from));
173}
174
175template <>
176EIGEN_STRONG_INLINE Eigen::half pfirst<Packet16h>(const Packet16h& from) {
177 return Eigen::half(_mm256_cvtsh_h(from));
178}
179
180template <>
181EIGEN_STRONG_INLINE Eigen::half pfirst<Packet8h>(const Packet8h& from) {
182 return Eigen::half(_mm_cvtsh_h(from));
183}
184
185// pload
186
187template <>
188EIGEN_STRONG_INLINE Packet32h pload<Packet32h>(const Eigen::half* from) {
189 EIGEN_DEBUG_ALIGNED_LOAD return _mm512_load_ph(from);
190}
191
192template <>
193EIGEN_STRONG_INLINE Packet16h pload<Packet16h>(const Eigen::half* from) {
194 EIGEN_DEBUG_ALIGNED_LOAD return _mm256_load_ph(from);
195}
196
197template <>
198EIGEN_STRONG_INLINE Packet8h pload<Packet8h>(const Eigen::half* from) {
199 EIGEN_DEBUG_ALIGNED_LOAD return _mm_load_ph(from);
200}
201
202// ploadu
203
204template <>
205EIGEN_STRONG_INLINE Packet32h ploadu<Packet32h>(const Eigen::half* from) {
206 EIGEN_DEBUG_UNALIGNED_LOAD return _mm512_loadu_ph(from);
207}
208
209template <>
210EIGEN_STRONG_INLINE Packet16h ploadu<Packet16h>(const Eigen::half* from) {
211 EIGEN_DEBUG_UNALIGNED_LOAD return _mm256_loadu_ph(from);
212}
213
214template <>
215EIGEN_STRONG_INLINE Packet8h ploadu<Packet8h>(const Eigen::half* from) {
216 EIGEN_DEBUG_UNALIGNED_LOAD return _mm_loadu_ph(from);
217}
218
219// pstore
220
221template <>
222EIGEN_STRONG_INLINE void pstore<half>(Eigen::half* to, const Packet32h& from) {
223 EIGEN_DEBUG_ALIGNED_STORE _mm512_store_ph(to, from);
224}
225
226template <>
227EIGEN_STRONG_INLINE void pstore<half>(Eigen::half* to, const Packet16h& from) {
228 EIGEN_DEBUG_ALIGNED_STORE _mm256_store_ph(to, from);
229}
230
231template <>
232EIGEN_STRONG_INLINE void pstore<half>(Eigen::half* to, const Packet8h& from) {
233 EIGEN_DEBUG_ALIGNED_STORE _mm_store_ph(to, from);
234}
235
236// pstoreu
237
238template <>
239EIGEN_STRONG_INLINE void pstoreu<half>(Eigen::half* to, const Packet32h& from) {
240 EIGEN_DEBUG_UNALIGNED_STORE _mm512_storeu_ph(to, from);
241}
242
243template <>
244EIGEN_STRONG_INLINE void pstoreu<half>(Eigen::half* to, const Packet16h& from) {
245 EIGEN_DEBUG_UNALIGNED_STORE _mm256_storeu_ph(to, from);
246}
247
248template <>
249EIGEN_STRONG_INLINE void pstoreu<half>(Eigen::half* to, const Packet8h& from) {
250 EIGEN_DEBUG_UNALIGNED_STORE _mm_storeu_ph(to, from);
251}
252
253// ploaddup
254template <>
255EIGEN_STRONG_INLINE Packet32h ploaddup<Packet32h>(const Eigen::half* from) {
256 __m512h a = _mm512_castph256_ph512(_mm256_loadu_ph(from));
257 return _mm512_permutexvar_ph(_mm512_set_epi16(15, 15, 14, 14, 13, 13, 12, 12, 11, 11, 10, 10, 9, 9, 8, 8, 7, 7, 6, 6,
258 5, 5, 4, 4, 3, 3, 2, 2, 1, 1, 0, 0),
259 a);
260}
261
262template <>
263EIGEN_STRONG_INLINE Packet16h ploaddup<Packet16h>(const Eigen::half* from) {
264 __m256h a = _mm256_castph128_ph256(_mm_loadu_ph(from));
265 return _mm256_permutexvar_ph(_mm256_set_epi16(7, 7, 6, 6, 5, 5, 4, 4, 3, 3, 2, 2, 1, 1, 0, 0), a);
266}
267
268template <>
269EIGEN_STRONG_INLINE Packet8h ploaddup<Packet8h>(const Eigen::half* from) {
270 return _mm_set_ph(from[3].x, from[3].x, from[2].x, from[2].x, from[1].x, from[1].x, from[0].x, from[0].x);
271}
272
273// ploadquad
274template <>
275EIGEN_STRONG_INLINE Packet32h ploadquad<Packet32h>(const Eigen::half* from) {
276 __m512h a = _mm512_castph128_ph512(_mm_loadu_ph(from));
277 return _mm512_permutexvar_ph(
278 _mm512_set_epi16(7, 7, 7, 7, 6, 6, 6, 6, 5, 5, 5, 5, 4, 4, 4, 4, 3, 3, 3, 3, 2, 2, 2, 2, 1, 1, 1, 1, 0, 0, 0, 0),
279 a);
280}
281
282template <>
283EIGEN_STRONG_INLINE Packet16h ploadquad<Packet16h>(const Eigen::half* from) {
284 return _mm256_set_ph(from[3].x, from[3].x, from[3].x, from[3].x, from[2].x, from[2].x, from[2].x, from[2].x,
285 from[1].x, from[1].x, from[1].x, from[1].x, from[0].x, from[0].x, from[0].x, from[0].x);
286}
287
288template <>
289EIGEN_STRONG_INLINE Packet8h ploadquad<Packet8h>(const Eigen::half* from) {
290 return _mm_set_ph(from[1].x, from[1].x, from[1].x, from[1].x, from[0].x, from[0].x, from[0].x, from[0].x);
291}
292
293// pabs
294
295template <>
296EIGEN_STRONG_INLINE Packet32h pabs<Packet32h>(const Packet32h& a) {
297 return _mm512_abs_ph(a);
298}
299
300template <>
301EIGEN_STRONG_INLINE Packet16h pabs<Packet16h>(const Packet16h& a) {
302 return _mm256_abs_ph(a);
303}
304
305template <>
306EIGEN_STRONG_INLINE Packet8h pabs<Packet8h>(const Packet8h& a) {
307 return _mm_abs_ph(a);
308}
309
310// psignbit
311
312template <>
313EIGEN_STRONG_INLINE Packet32h psignbit<Packet32h>(const Packet32h& a) {
314 return _mm512_castsi512_ph(_mm512_srai_epi16(_mm512_castph_si512(a), 15));
315}
316
317template <>
318EIGEN_STRONG_INLINE Packet16h psignbit<Packet16h>(const Packet16h& a) {
319 return _mm256_castsi256_ph(_mm256_srai_epi16(_mm256_castph_si256(a), 15));
320}
321
322template <>
323EIGEN_STRONG_INLINE Packet8h psignbit<Packet8h>(const Packet8h& a) {
324 return _mm_castsi128_ph(_mm_srai_epi16(_mm_castph_si128(a), 15));
325}
326
327// pmin
328
329template <>
330EIGEN_STRONG_INLINE Packet32h pmin<Packet32h>(const Packet32h& a, const Packet32h& b) {
331 return _mm512_min_ph(a, b);
332}
333
334template <>
335EIGEN_STRONG_INLINE Packet16h pmin<Packet16h>(const Packet16h& a, const Packet16h& b) {
336 return _mm256_min_ph(a, b);
337}
338
339template <>
340EIGEN_STRONG_INLINE Packet8h pmin<Packet8h>(const Packet8h& a, const Packet8h& b) {
341 return _mm_min_ph(a, b);
342}
343
344// pmax
345
346template <>
347EIGEN_STRONG_INLINE Packet32h pmax<Packet32h>(const Packet32h& a, const Packet32h& b) {
348 return _mm512_max_ph(a, b);
349}
350
351template <>
352EIGEN_STRONG_INLINE Packet16h pmax<Packet16h>(const Packet16h& a, const Packet16h& b) {
353 return _mm256_max_ph(a, b);
354}
355
356template <>
357EIGEN_STRONG_INLINE Packet8h pmax<Packet8h>(const Packet8h& a, const Packet8h& b) {
358 return _mm_max_ph(a, b);
359}
360
361// plset
362template <>
363EIGEN_STRONG_INLINE Packet32h plset<Packet32h>(const half& a) {
364 return _mm512_add_ph(pset1<Packet32h>(a), _mm512_set_ph(31, 30, 29, 28, 27, 26, 25, 24, 23, 22, 21, 20, 19, 18, 17,
365 16, 15, 14, 13, 12, 11, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1, 0));
366}
367
368template <>
369EIGEN_STRONG_INLINE Packet16h plset<Packet16h>(const half& a) {
370 return _mm256_add_ph(pset1<Packet16h>(a), _mm256_set_ph(15, 14, 13, 12, 11, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1, 0));
371}
372
373template <>
374EIGEN_STRONG_INLINE Packet8h plset<Packet8h>(const half& a) {
375 return _mm_add_ph(pset1<Packet8h>(a), _mm_set_ph(7, 6, 5, 4, 3, 2, 1, 0));
376}
377
378// por
379
380template <>
381EIGEN_STRONG_INLINE Packet32h por(const Packet32h& a, const Packet32h& b) {
382 return _mm512_castsi512_ph(_mm512_or_si512(_mm512_castph_si512(a), _mm512_castph_si512(b)));
383}
384
385template <>
386EIGEN_STRONG_INLINE Packet16h por(const Packet16h& a, const Packet16h& b) {
387 return _mm256_castsi256_ph(_mm256_or_si256(_mm256_castph_si256(a), _mm256_castph_si256(b)));
388}
389
390template <>
391EIGEN_STRONG_INLINE Packet8h por(const Packet8h& a, const Packet8h& b) {
392 return _mm_castsi128_ph(_mm_or_si128(_mm_castph_si128(a), _mm_castph_si128(b)));
393}
394
395// pxor
396
397template <>
398EIGEN_STRONG_INLINE Packet32h pxor(const Packet32h& a, const Packet32h& b) {
399 return _mm512_castsi512_ph(_mm512_xor_si512(_mm512_castph_si512(a), _mm512_castph_si512(b)));
400}
401
402template <>
403EIGEN_STRONG_INLINE Packet16h pxor(const Packet16h& a, const Packet16h& b) {
404 return _mm256_castsi256_ph(_mm256_xor_si256(_mm256_castph_si256(a), _mm256_castph_si256(b)));
405}
406
407template <>
408EIGEN_STRONG_INLINE Packet8h pxor(const Packet8h& a, const Packet8h& b) {
409 return _mm_castsi128_ph(_mm_xor_si128(_mm_castph_si128(a), _mm_castph_si128(b)));
410}
411
412// pand
413
414template <>
415EIGEN_STRONG_INLINE Packet32h pand(const Packet32h& a, const Packet32h& b) {
416 return _mm512_castsi512_ph(_mm512_and_si512(_mm512_castph_si512(a), _mm512_castph_si512(b)));
417}
418
419template <>
420EIGEN_STRONG_INLINE Packet16h pand(const Packet16h& a, const Packet16h& b) {
421 return _mm256_castsi256_ph(_mm256_and_si256(_mm256_castph_si256(a), _mm256_castph_si256(b)));
422}
423
424template <>
425EIGEN_STRONG_INLINE Packet8h pand(const Packet8h& a, const Packet8h& b) {
426 return _mm_castsi128_ph(_mm_and_si128(_mm_castph_si128(a), _mm_castph_si128(b)));
427}
428
429// pandnot
430
431template <>
432EIGEN_STRONG_INLINE Packet32h pandnot(const Packet32h& a, const Packet32h& b) {
433 return _mm512_castsi512_ph(_mm512_andnot_si512(_mm512_castph_si512(b), _mm512_castph_si512(a)));
434}
435
436template <>
437EIGEN_STRONG_INLINE Packet16h pandnot(const Packet16h& a, const Packet16h& b) {
438 return _mm256_castsi256_ph(_mm256_andnot_si256(_mm256_castph_si256(b), _mm256_castph_si256(a)));
439}
440
441template <>
442EIGEN_STRONG_INLINE Packet8h pandnot(const Packet8h& a, const Packet8h& b) {
443 return _mm_castsi128_ph(_mm_andnot_si128(_mm_castph_si128(b), _mm_castph_si128(a)));
444}
445
446// pselect
447
448template <>
449EIGEN_DEVICE_FUNC inline Packet32h pselect(const Packet32h& mask, const Packet32h& a, const Packet32h& b) {
450 __mmask32 mask32 = _mm512_cmp_epi16_mask(_mm512_castph_si512(mask), _mm512_setzero_epi32(), _MM_CMPINT_EQ);
451 return _mm512_mask_blend_ph(mask32, a, b);
452}
453
454template <>
455EIGEN_DEVICE_FUNC inline Packet16h pselect(const Packet16h& mask, const Packet16h& a, const Packet16h& b) {
456 __mmask16 mask16 = _mm256_cmp_epi16_mask(_mm256_castph_si256(mask), _mm256_setzero_si256(), _MM_CMPINT_EQ);
457 return _mm256_mask_blend_ph(mask16, a, b);
458}
459
460template <>
461EIGEN_DEVICE_FUNC inline Packet8h pselect(const Packet8h& mask, const Packet8h& a, const Packet8h& b) {
462 __mmask8 mask8 = _mm_cmp_epi16_mask(_mm_castph_si128(mask), _mm_setzero_si128(), _MM_CMPINT_EQ);
463 return _mm_mask_blend_ph(mask8, a, b);
464}
465
466// pcmp_eq
467
468template <>
469EIGEN_STRONG_INLINE Packet32h pcmp_eq(const Packet32h& a, const Packet32h& b) {
470 __mmask32 mask = _mm512_cmp_ph_mask(a, b, _CMP_EQ_OQ);
471 return _mm512_castsi512_ph(_mm512_mask_set1_epi16(_mm512_set1_epi32(0), mask, static_cast<short>(0xffffu)));
472}
473
474template <>
475EIGEN_STRONG_INLINE Packet16h pcmp_eq(const Packet16h& a, const Packet16h& b) {
476 __mmask16 mask = _mm256_cmp_ph_mask(a, b, _CMP_EQ_OQ);
477 return _mm256_castsi256_ph(_mm256_mask_set1_epi16(_mm256_set1_epi32(0), mask, static_cast<short>(0xffffu)));
478}
479
480template <>
481EIGEN_STRONG_INLINE Packet8h pcmp_eq(const Packet8h& a, const Packet8h& b) {
482 __mmask8 mask = _mm_cmp_ph_mask(a, b, _CMP_EQ_OQ);
483 return _mm_castsi128_ph(_mm_mask_set1_epi16(_mm_set1_epi32(0), mask, static_cast<short>(0xffffu)));
484}
485
486// pcmp_le
487
488template <>
489EIGEN_STRONG_INLINE Packet32h pcmp_le(const Packet32h& a, const Packet32h& b) {
490 __mmask32 mask = _mm512_cmp_ph_mask(a, b, _CMP_LE_OQ);
491 return _mm512_castsi512_ph(_mm512_mask_set1_epi16(_mm512_set1_epi32(0), mask, static_cast<short>(0xffffu)));
492}
493
494template <>
495EIGEN_STRONG_INLINE Packet16h pcmp_le(const Packet16h& a, const Packet16h& b) {
496 __mmask16 mask = _mm256_cmp_ph_mask(a, b, _CMP_LE_OQ);
497 return _mm256_castsi256_ph(_mm256_mask_set1_epi16(_mm256_set1_epi32(0), mask, static_cast<short>(0xffffu)));
498}
499
500template <>
501EIGEN_STRONG_INLINE Packet8h pcmp_le(const Packet8h& a, const Packet8h& b) {
502 __mmask8 mask = _mm_cmp_ph_mask(a, b, _CMP_LE_OQ);
503 return _mm_castsi128_ph(_mm_mask_set1_epi16(_mm_set1_epi32(0), mask, static_cast<short>(0xffffu)));
504}
505
506// pcmp_lt
507
508template <>
509EIGEN_STRONG_INLINE Packet32h pcmp_lt(const Packet32h& a, const Packet32h& b) {
510 __mmask32 mask = _mm512_cmp_ph_mask(a, b, _CMP_LT_OQ);
511 return _mm512_castsi512_ph(_mm512_mask_set1_epi16(_mm512_set1_epi32(0), mask, static_cast<short>(0xffffu)));
512}
513
514template <>
515EIGEN_STRONG_INLINE Packet16h pcmp_lt(const Packet16h& a, const Packet16h& b) {
516 __mmask16 mask = _mm256_cmp_ph_mask(a, b, _CMP_LT_OQ);
517 return _mm256_castsi256_ph(_mm256_mask_set1_epi16(_mm256_set1_epi32(0), mask, static_cast<short>(0xffffu)));
518}
519
520template <>
521EIGEN_STRONG_INLINE Packet8h pcmp_lt(const Packet8h& a, const Packet8h& b) {
522 __mmask8 mask = _mm_cmp_ph_mask(a, b, _CMP_LT_OQ);
523 return _mm_castsi128_ph(_mm_mask_set1_epi16(_mm_set1_epi32(0), mask, static_cast<short>(0xffffu)));
524}
525
526// pcmp_lt_or_nan
527
528template <>
529EIGEN_STRONG_INLINE Packet32h pcmp_lt_or_nan(const Packet32h& a, const Packet32h& b) {
530 __mmask32 mask = _mm512_cmp_ph_mask(a, b, _CMP_NGE_UQ);
531 return _mm512_castsi512_ph(_mm512_mask_set1_epi16(_mm512_set1_epi16(0), mask, static_cast<short>(0xffffu)));
532}
533
534template <>
535EIGEN_STRONG_INLINE Packet16h pcmp_lt_or_nan(const Packet16h& a, const Packet16h& b) {
536 __mmask16 mask = _mm256_cmp_ph_mask(a, b, _CMP_NGE_UQ);
537 return _mm256_castsi256_ph(_mm256_mask_set1_epi16(_mm256_set1_epi32(0), mask, static_cast<short>(0xffffu)));
538}
539
540template <>
541EIGEN_STRONG_INLINE Packet8h pcmp_lt_or_nan(const Packet8h& a, const Packet8h& b) {
542 __mmask8 mask = _mm_cmp_ph_mask(a, b, _CMP_NGE_UQ);
543 return _mm_castsi128_ph(_mm_mask_set1_epi16(_mm_set1_epi32(0), mask, static_cast<short>(0xffffu)));
544}
545
546// padd
547
548template <>
549EIGEN_STRONG_INLINE Packet32h padd<Packet32h>(const Packet32h& a, const Packet32h& b) {
550 return _mm512_add_ph(a, b);
551}
552
553template <>
554EIGEN_STRONG_INLINE Packet16h padd<Packet16h>(const Packet16h& a, const Packet16h& b) {
555 return _mm256_add_ph(a, b);
556}
557
558template <>
559EIGEN_STRONG_INLINE Packet8h padd<Packet8h>(const Packet8h& a, const Packet8h& b) {
560 return _mm_add_ph(a, b);
561}
562
563// psub
564
565template <>
566EIGEN_STRONG_INLINE Packet32h psub<Packet32h>(const Packet32h& a, const Packet32h& b) {
567 return _mm512_sub_ph(a, b);
568}
569
570template <>
571EIGEN_STRONG_INLINE Packet16h psub<Packet16h>(const Packet16h& a, const Packet16h& b) {
572 return _mm256_sub_ph(a, b);
573}
574
575template <>
576EIGEN_STRONG_INLINE Packet8h psub<Packet8h>(const Packet8h& a, const Packet8h& b) {
577 return _mm_sub_ph(a, b);
578}
579
580// pmul
581
582template <>
583EIGEN_STRONG_INLINE Packet32h pmul<Packet32h>(const Packet32h& a, const Packet32h& b) {
584 return _mm512_mul_ph(a, b);
585}
586
587template <>
588EIGEN_STRONG_INLINE Packet16h pmul<Packet16h>(const Packet16h& a, const Packet16h& b) {
589 return _mm256_mul_ph(a, b);
590}
591
592template <>
593EIGEN_STRONG_INLINE Packet8h pmul<Packet8h>(const Packet8h& a, const Packet8h& b) {
594 return _mm_mul_ph(a, b);
595}
596
597// pdiv
598
599template <>
600EIGEN_STRONG_INLINE Packet32h pdiv<Packet32h>(const Packet32h& a, const Packet32h& b) {
601 return _mm512_div_ph(a, b);
602}
603
604template <>
605EIGEN_STRONG_INLINE Packet16h pdiv<Packet16h>(const Packet16h& a, const Packet16h& b) {
606 return _mm256_div_ph(a, b);
607}
608
609template <>
610EIGEN_STRONG_INLINE Packet8h pdiv<Packet8h>(const Packet8h& a, const Packet8h& b) {
611 return _mm_div_ph(a, b);
612 ;
613}
614
615// pround
616
617template <>
618EIGEN_STRONG_INLINE Packet32h pround<Packet32h>(const Packet32h& a) {
619 // Work-around for default std::round rounding mode.
620
621 // Mask for the sign bit.
622 const Packet32h signMask =
623 pset1frombits<Packet32h>(static_cast<numext::uint16_t>(static_cast<std::uint16_t>(0x8000u)));
624 // The largest half-precision float less than 0.5.
625 const Packet32h prev0dot5 = pset1frombits<Packet32h>(static_cast<numext::uint16_t>(0x37FFu));
626
627 return _mm512_roundscale_ph(padd(por(pand(a, signMask), prev0dot5), a), _MM_FROUND_TO_ZERO);
628}
629
630template <>
631EIGEN_STRONG_INLINE Packet16h pround<Packet16h>(const Packet16h& a) {
632 // Work-around for default std::round rounding mode.
633
634 // Mask for the sign bit.
635 const Packet16h signMask =
636 pset1frombits<Packet16h>(static_cast<numext::uint16_t>(static_cast<std::uint16_t>(0x8000u)));
637 // The largest half-precision float less than 0.5.
638 const Packet16h prev0dot5 = pset1frombits<Packet16h>(static_cast<numext::uint16_t>(0x37FFu));
639
640 return _mm256_roundscale_ph(padd(por(pand(a, signMask), prev0dot5), a), _MM_FROUND_TO_ZERO);
641}
642
643template <>
644EIGEN_STRONG_INLINE Packet8h pround<Packet8h>(const Packet8h& a) {
645 // Work-around for default std::round rounding mode.
646
647 // Mask for the sign bit.
648 const Packet8h signMask = pset1frombits<Packet8h>(static_cast<numext::uint16_t>(static_cast<std::uint16_t>(0x8000u)));
649 // The largest half-precision float less than 0.5.
650 const Packet8h prev0dot5 = pset1frombits<Packet8h>(static_cast<numext::uint16_t>(0x37FFu));
651
652 return _mm_roundscale_ph(padd(por(pand(a, signMask), prev0dot5), a), _MM_FROUND_TO_ZERO);
653}
654
655// print
656
657template <>
658EIGEN_STRONG_INLINE Packet32h print<Packet32h>(const Packet32h& a) {
659 return _mm512_roundscale_ph(a, _MM_FROUND_CUR_DIRECTION);
660}
661
662template <>
663EIGEN_STRONG_INLINE Packet16h print<Packet16h>(const Packet16h& a) {
664 return _mm256_roundscale_ph(a, _MM_FROUND_CUR_DIRECTION);
665}
666
667template <>
668EIGEN_STRONG_INLINE Packet8h print<Packet8h>(const Packet8h& a) {
669 return _mm_roundscale_ph(a, _MM_FROUND_CUR_DIRECTION);
670}
671
672// pceil
673
674template <>
675EIGEN_STRONG_INLINE Packet32h pceil<Packet32h>(const Packet32h& a) {
676 return _mm512_roundscale_ph(a, _MM_FROUND_TO_POS_INF);
677}
678
679template <>
680EIGEN_STRONG_INLINE Packet16h pceil<Packet16h>(const Packet16h& a) {
681 return _mm256_roundscale_ph(a, _MM_FROUND_TO_POS_INF);
682}
683
684template <>
685EIGEN_STRONG_INLINE Packet8h pceil<Packet8h>(const Packet8h& a) {
686 return _mm_roundscale_ph(a, _MM_FROUND_TO_POS_INF);
687}
688
689// pfloor
690
691template <>
692EIGEN_STRONG_INLINE Packet32h pfloor<Packet32h>(const Packet32h& a) {
693 return _mm512_roundscale_ph(a, _MM_FROUND_TO_NEG_INF);
694}
695
696template <>
697EIGEN_STRONG_INLINE Packet16h pfloor<Packet16h>(const Packet16h& a) {
698 return _mm256_roundscale_ph(a, _MM_FROUND_TO_NEG_INF);
699}
700
701template <>
702EIGEN_STRONG_INLINE Packet8h pfloor<Packet8h>(const Packet8h& a) {
703 return _mm_roundscale_ph(a, _MM_FROUND_TO_NEG_INF);
704}
705
706// ptrunc
707
708template <>
709EIGEN_STRONG_INLINE Packet32h ptrunc<Packet32h>(const Packet32h& a) {
710 return _mm512_roundscale_ph(a, _MM_FROUND_TO_ZERO);
711}
712
713template <>
714EIGEN_STRONG_INLINE Packet16h ptrunc<Packet16h>(const Packet16h& a) {
715 return _mm256_roundscale_ph(a, _MM_FROUND_TO_ZERO);
716}
717
718template <>
719EIGEN_STRONG_INLINE Packet8h ptrunc<Packet8h>(const Packet8h& a) {
720 return _mm_roundscale_ph(a, _MM_FROUND_TO_ZERO);
721}
722
723// predux
724template <>
725EIGEN_STRONG_INLINE half predux<Packet32h>(const Packet32h& a) {
726 return half(_mm512_reduce_add_ph(a));
727}
728
729template <>
730EIGEN_STRONG_INLINE half predux<Packet16h>(const Packet16h& a) {
731 return half(_mm256_reduce_add_ph(a));
732}
733
734template <>
735EIGEN_STRONG_INLINE half predux<Packet8h>(const Packet8h& a) {
736 return half(_mm_reduce_add_ph(a));
737}
738
739template <>
740EIGEN_STRONG_INLINE bool predux_any(const Packet32h& a) {
741 return avx512_predux_any(_mm512_castph_si512(a));
742}
743
744template <>
745EIGEN_STRONG_INLINE bool predux_any(const Packet16h& a) {
746 const __m256i bits = _mm256_castph_si256(a);
747 return _mm256_testz_si256(bits, bits) == 0;
748}
749
750template <>
751EIGEN_STRONG_INLINE bool predux_any(const Packet8h& a) {
752 const __m128i bits = _mm_castph_si128(a);
753 return _mm_testz_si128(bits, bits) == 0;
754}
755
756template <>
757EIGEN_STRONG_INLINE bool predux_all(const Packet32h& a) {
758 return _mm512_cmp_ph_mask(a, _mm512_setzero_ph(), _CMP_EQ_OQ) == 0;
759}
760
761template <>
762EIGEN_STRONG_INLINE bool predux_all(const Packet16h& a) {
763 return _mm256_cmp_ph_mask(a, _mm256_setzero_ph(), _CMP_EQ_OQ) == 0;
764}
765
766template <>
767EIGEN_STRONG_INLINE bool predux_all(const Packet8h& a) {
768 return _mm_cmp_ph_mask(a, _mm_setzero_ph(), _CMP_EQ_OQ) == 0;
769}
770
771// predux_half
772template <>
773EIGEN_STRONG_INLINE Packet16h predux_half<Packet32h>(const Packet32h& a) {
774 const __m512i bits = _mm512_castph_si512(a);
775 Packet16h lo = _mm256_castsi256_ph(_mm512_castsi512_si256(bits));
776 Packet16h hi = _mm256_castsi256_ph(_mm512_extracti64x4_epi64(bits, 1));
777 return padd(lo, hi);
778}
779
780template <>
781EIGEN_STRONG_INLINE Packet8h predux_half<Packet16h>(const Packet16h& a) {
782 Packet8h lo = _mm_castsi128_ph(_mm256_castsi256_si128(_mm256_castph_si256(a)));
783 Packet8h hi = _mm_castps_ph(_mm256_extractf128_ps(_mm256_castph_ps(a), 1));
784 return padd(lo, hi);
785}
786
787// predux_max
788
789template <>
790EIGEN_STRONG_INLINE half predux_max<Packet32h>(const Packet32h& a) {
791 return half(_mm512_reduce_max_ph(a));
792}
793
794template <>
795EIGEN_STRONG_INLINE half predux_max<Packet16h>(const Packet16h& a) {
796 return half(_mm256_reduce_max_ph(a));
797}
798
799template <>
800EIGEN_STRONG_INLINE half predux_max<Packet8h>(const Packet8h& a) {
801 return half(_mm_reduce_max_ph(a));
802}
803
804// predux_min
805
806template <>
807EIGEN_STRONG_INLINE half predux_min<Packet32h>(const Packet32h& a) {
808 return half(_mm512_reduce_min_ph(a));
809}
810
811template <>
812EIGEN_STRONG_INLINE half predux_min<Packet16h>(const Packet16h& a) {
813 return half(_mm256_reduce_min_ph(a));
814}
815
816template <>
817EIGEN_STRONG_INLINE half predux_min<Packet8h>(const Packet8h& a) {
818 return half(_mm_reduce_min_ph(a));
819}
820
821// predux_mul
822
823template <>
824EIGEN_STRONG_INLINE half predux_mul<Packet32h>(const Packet32h& a) {
825 return half(_mm512_reduce_mul_ph(a));
826}
827
828template <>
829EIGEN_STRONG_INLINE half predux_mul<Packet16h>(const Packet16h& a) {
830 return half(_mm256_reduce_mul_ph(a));
831}
832
833template <>
834EIGEN_STRONG_INLINE half predux_mul<Packet8h>(const Packet8h& a) {
835 return half(_mm_reduce_mul_ph(a));
836}
837
838#ifdef EIGEN_VECTORIZE_FMA
839
840// pmadd
841
842template <>
843EIGEN_STRONG_INLINE Packet32h pmadd(const Packet32h& a, const Packet32h& b, const Packet32h& c) {
844 return _mm512_fmadd_ph(a, b, c);
845}
846
847template <>
848EIGEN_STRONG_INLINE Packet16h pmadd(const Packet16h& a, const Packet16h& b, const Packet16h& c) {
849 return _mm256_fmadd_ph(a, b, c);
850}
851
852template <>
853EIGEN_STRONG_INLINE Packet8h pmadd(const Packet8h& a, const Packet8h& b, const Packet8h& c) {
854 return _mm_fmadd_ph(a, b, c);
855}
856
857// pmsub
858
859template <>
860EIGEN_STRONG_INLINE Packet32h pmsub(const Packet32h& a, const Packet32h& b, const Packet32h& c) {
861 return _mm512_fmsub_ph(a, b, c);
862}
863
864template <>
865EIGEN_STRONG_INLINE Packet16h pmsub(const Packet16h& a, const Packet16h& b, const Packet16h& c) {
866 return _mm256_fmsub_ph(a, b, c);
867}
868
869template <>
870EIGEN_STRONG_INLINE Packet8h pmsub(const Packet8h& a, const Packet8h& b, const Packet8h& c) {
871 return _mm_fmsub_ph(a, b, c);
872}
873
874// pnmadd
875
876template <>
877EIGEN_STRONG_INLINE Packet32h pnmadd(const Packet32h& a, const Packet32h& b, const Packet32h& c) {
878 return _mm512_fnmadd_ph(a, b, c);
879}
880
881template <>
882EIGEN_STRONG_INLINE Packet16h pnmadd(const Packet16h& a, const Packet16h& b, const Packet16h& c) {
883 return _mm256_fnmadd_ph(a, b, c);
884}
885
886template <>
887EIGEN_STRONG_INLINE Packet8h pnmadd(const Packet8h& a, const Packet8h& b, const Packet8h& c) {
888 return _mm_fnmadd_ph(a, b, c);
889}
890
891// pnmsub
892
893template <>
894EIGEN_STRONG_INLINE Packet32h pnmsub(const Packet32h& a, const Packet32h& b, const Packet32h& c) {
895 return _mm512_fnmsub_ph(a, b, c);
896}
897
898template <>
899EIGEN_STRONG_INLINE Packet16h pnmsub(const Packet16h& a, const Packet16h& b, const Packet16h& c) {
900 return _mm256_fnmsub_ph(a, b, c);
901}
902
903template <>
904EIGEN_STRONG_INLINE Packet8h pnmsub(const Packet8h& a, const Packet8h& b, const Packet8h& c) {
905 return _mm_fnmsub_ph(a, b, c);
906}
907
908#endif
909
910// pnegate
911
912template <>
913EIGEN_STRONG_INLINE Packet32h pnegate<Packet32h>(const Packet32h& a) {
914 return _mm512_castsi512_ph(_mm512_xor_si512(_mm512_castph_si512(a), _mm512_set1_epi16(static_cast<short>(0x8000u))));
915}
916
917template <>
918EIGEN_STRONG_INLINE Packet16h pnegate<Packet16h>(const Packet16h& a) {
919 return _mm256_castsi256_ph(_mm256_xor_si256(_mm256_castph_si256(a), _mm256_set1_epi16(static_cast<short>(0x8000u))));
920}
921
922template <>
923EIGEN_STRONG_INLINE Packet8h pnegate<Packet8h>(const Packet8h& a) {
924 return _mm_castsi128_ph(_mm_xor_si128(_mm_castph_si128(a), _mm_set1_epi16(static_cast<short>(0x8000u))));
925}
926
927// pconj
928
929// Nothing, packets are real.
930
931// psqrt
932
933template <>
934EIGEN_STRONG_INLINE Packet32h psqrt<Packet32h>(const Packet32h& a) {
935 return generic_sqrt_newton_step<Packet32h>::run(a, _mm512_rsqrt_ph(a));
936}
937
938template <>
939EIGEN_STRONG_INLINE Packet16h psqrt<Packet16h>(const Packet16h& a) {
940 return generic_sqrt_newton_step<Packet16h>::run(a, _mm256_rsqrt_ph(a));
941}
942
943template <>
944EIGEN_STRONG_INLINE Packet8h psqrt<Packet8h>(const Packet8h& a) {
945 return generic_sqrt_newton_step<Packet8h>::run(a, _mm_rsqrt_ph(a));
946}
947
948// prsqrt
949
950template <>
951EIGEN_STRONG_INLINE Packet32h prsqrt<Packet32h>(const Packet32h& a) {
952 return generic_rsqrt_newton_step<Packet32h, /*Steps=*/1>::run(a, _mm512_rsqrt_ph(a));
953}
954
955template <>
956EIGEN_STRONG_INLINE Packet16h prsqrt<Packet16h>(const Packet16h& a) {
957 return generic_rsqrt_newton_step<Packet16h, /*Steps=*/1>::run(a, _mm256_rsqrt_ph(a));
958}
959
960template <>
961EIGEN_STRONG_INLINE Packet8h prsqrt<Packet8h>(const Packet8h& a) {
962 return generic_rsqrt_newton_step<Packet8h, /*Steps=*/1>::run(a, _mm_rsqrt_ph(a));
963}
964
965// preciprocal
966
967template <>
968EIGEN_STRONG_INLINE Packet32h preciprocal<Packet32h>(const Packet32h& a) {
969 return generic_reciprocal_newton_step<Packet32h, /*Steps=*/1>::run(a, _mm512_rcp_ph(a));
970}
971
972template <>
973EIGEN_STRONG_INLINE Packet16h preciprocal<Packet16h>(const Packet16h& a) {
974 return generic_reciprocal_newton_step<Packet16h, /*Steps=*/1>::run(a, _mm256_rcp_ph(a));
975}
976
977template <>
978EIGEN_STRONG_INLINE Packet8h preciprocal<Packet8h>(const Packet8h& a) {
979 return generic_reciprocal_newton_step<Packet8h, /*Steps=*/1>::run(a, _mm_rcp_ph(a));
980}
981
982// ptranspose
983
984EIGEN_DEVICE_FUNC inline void ptranspose(PacketBlock<Packet32h, 32>& a) {
985 __m512i t[32];
986
987 EIGEN_UNROLL_LOOP
988 for (int i = 0; i < 16; i++) {
989 t[2 * i] = _mm512_unpacklo_epi16(_mm512_castph_si512(a.packet[2 * i]), _mm512_castph_si512(a.packet[2 * i + 1]));
990 t[2 * i + 1] =
991 _mm512_unpackhi_epi16(_mm512_castph_si512(a.packet[2 * i]), _mm512_castph_si512(a.packet[2 * i + 1]));
992 }
993
994 __m512i p[32];
995
996 EIGEN_UNROLL_LOOP
997 for (int i = 0; i < 8; i++) {
998 p[4 * i] = _mm512_unpacklo_epi32(t[4 * i], t[4 * i + 2]);
999 p[4 * i + 1] = _mm512_unpackhi_epi32(t[4 * i], t[4 * i + 2]);
1000 p[4 * i + 2] = _mm512_unpacklo_epi32(t[4 * i + 1], t[4 * i + 3]);
1001 p[4 * i + 3] = _mm512_unpackhi_epi32(t[4 * i + 1], t[4 * i + 3]);
1002 }
1003
1004 __m512i q[32];
1005
1006 EIGEN_UNROLL_LOOP
1007 for (int i = 0; i < 4; i++) {
1008 q[8 * i] = _mm512_unpacklo_epi64(p[8 * i], p[8 * i + 4]);
1009 q[8 * i + 1] = _mm512_unpackhi_epi64(p[8 * i], p[8 * i + 4]);
1010 q[8 * i + 2] = _mm512_unpacklo_epi64(p[8 * i + 1], p[8 * i + 5]);
1011 q[8 * i + 3] = _mm512_unpackhi_epi64(p[8 * i + 1], p[8 * i + 5]);
1012 q[8 * i + 4] = _mm512_unpacklo_epi64(p[8 * i + 2], p[8 * i + 6]);
1013 q[8 * i + 5] = _mm512_unpackhi_epi64(p[8 * i + 2], p[8 * i + 6]);
1014 q[8 * i + 6] = _mm512_unpacklo_epi64(p[8 * i + 3], p[8 * i + 7]);
1015 q[8 * i + 7] = _mm512_unpackhi_epi64(p[8 * i + 3], p[8 * i + 7]);
1016 }
1017
1018 __m512i f[32] = {};
1019
1020#define PACKET32H_TRANSPOSE_HELPER(X, Y) \
1021 do { \
1022 f[Y * 8] = _mm512_inserti32x4(f[Y * 8], _mm512_extracti32x4_epi32(q[X * 8], Y), X); \
1023 f[Y * 8 + 1] = _mm512_inserti32x4(f[Y * 8 + 1], _mm512_extracti32x4_epi32(q[X * 8 + 1], Y), X); \
1024 f[Y * 8 + 2] = _mm512_inserti32x4(f[Y * 8 + 2], _mm512_extracti32x4_epi32(q[X * 8 + 2], Y), X); \
1025 f[Y * 8 + 3] = _mm512_inserti32x4(f[Y * 8 + 3], _mm512_extracti32x4_epi32(q[X * 8 + 3], Y), X); \
1026 f[Y * 8 + 4] = _mm512_inserti32x4(f[Y * 8 + 4], _mm512_extracti32x4_epi32(q[X * 8 + 4], Y), X); \
1027 f[Y * 8 + 5] = _mm512_inserti32x4(f[Y * 8 + 5], _mm512_extracti32x4_epi32(q[X * 8 + 5], Y), X); \
1028 f[Y * 8 + 6] = _mm512_inserti32x4(f[Y * 8 + 6], _mm512_extracti32x4_epi32(q[X * 8 + 6], Y), X); \
1029 f[Y * 8 + 7] = _mm512_inserti32x4(f[Y * 8 + 7], _mm512_extracti32x4_epi32(q[X * 8 + 7], Y), X); \
1030 } while (false);
1031
1032 PACKET32H_TRANSPOSE_HELPER(0, 0);
1033 PACKET32H_TRANSPOSE_HELPER(1, 1);
1034 PACKET32H_TRANSPOSE_HELPER(2, 2);
1035 PACKET32H_TRANSPOSE_HELPER(3, 3);
1036
1037 PACKET32H_TRANSPOSE_HELPER(1, 0);
1038 PACKET32H_TRANSPOSE_HELPER(2, 0);
1039 PACKET32H_TRANSPOSE_HELPER(3, 0);
1040 PACKET32H_TRANSPOSE_HELPER(2, 1);
1041 PACKET32H_TRANSPOSE_HELPER(3, 1);
1042 PACKET32H_TRANSPOSE_HELPER(3, 2);
1043
1044 PACKET32H_TRANSPOSE_HELPER(0, 1);
1045 PACKET32H_TRANSPOSE_HELPER(0, 2);
1046 PACKET32H_TRANSPOSE_HELPER(0, 3);
1047 PACKET32H_TRANSPOSE_HELPER(1, 2);
1048 PACKET32H_TRANSPOSE_HELPER(1, 3);
1049 PACKET32H_TRANSPOSE_HELPER(2, 3);
1050
1051#undef PACKET32H_TRANSPOSE_HELPER
1052
1053 EIGEN_UNROLL_LOOP
1054 for (int i = 0; i < 32; i++) {
1055 a.packet[i] = _mm512_castsi512_ph(f[i]);
1056 }
1057}
1058
1059EIGEN_DEVICE_FUNC inline void ptranspose(PacketBlock<Packet32h, 4>& a) {
1060 __m512i p0, p1, p2, p3, t0, t1, t2, t3, a0, a1, a2, a3;
1061 t0 = _mm512_unpacklo_epi16(_mm512_castph_si512(a.packet[0]), _mm512_castph_si512(a.packet[1]));
1062 t1 = _mm512_unpackhi_epi16(_mm512_castph_si512(a.packet[0]), _mm512_castph_si512(a.packet[1]));
1063 t2 = _mm512_unpacklo_epi16(_mm512_castph_si512(a.packet[2]), _mm512_castph_si512(a.packet[3]));
1064 t3 = _mm512_unpackhi_epi16(_mm512_castph_si512(a.packet[2]), _mm512_castph_si512(a.packet[3]));
1065
1066 p0 = _mm512_unpacklo_epi32(t0, t2);
1067 p1 = _mm512_unpackhi_epi32(t0, t2);
1068 p2 = _mm512_unpacklo_epi32(t1, t3);
1069 p3 = _mm512_unpackhi_epi32(t1, t3);
1070
1071 a0 = p0;
1072 a1 = p1;
1073 a2 = p2;
1074 a3 = p3;
1075
1076 a0 = _mm512_inserti32x4(a0, _mm512_extracti32x4_epi32(p1, 0), 1);
1077 a1 = _mm512_inserti32x4(a1, _mm512_extracti32x4_epi32(p0, 1), 0);
1078
1079 a0 = _mm512_inserti32x4(a0, _mm512_extracti32x4_epi32(p2, 0), 2);
1080 a2 = _mm512_inserti32x4(a2, _mm512_extracti32x4_epi32(p0, 2), 0);
1081
1082 a0 = _mm512_inserti32x4(a0, _mm512_extracti32x4_epi32(p3, 0), 3);
1083 a3 = _mm512_inserti32x4(a3, _mm512_extracti32x4_epi32(p0, 3), 0);
1084
1085 a1 = _mm512_inserti32x4(a1, _mm512_extracti32x4_epi32(p2, 1), 2);
1086 a2 = _mm512_inserti32x4(a2, _mm512_extracti32x4_epi32(p1, 2), 1);
1087
1088 a2 = _mm512_inserti32x4(a2, _mm512_extracti32x4_epi32(p3, 2), 3);
1089 a3 = _mm512_inserti32x4(a3, _mm512_extracti32x4_epi32(p2, 3), 2);
1090
1091 a1 = _mm512_inserti32x4(a1, _mm512_extracti32x4_epi32(p3, 1), 3);
1092 a3 = _mm512_inserti32x4(a3, _mm512_extracti32x4_epi32(p1, 3), 1);
1093
1094 a.packet[0] = _mm512_castsi512_ph(a0);
1095 a.packet[1] = _mm512_castsi512_ph(a1);
1096 a.packet[2] = _mm512_castsi512_ph(a2);
1097 a.packet[3] = _mm512_castsi512_ph(a3);
1098}
1099
1100EIGEN_STRONG_INLINE void ptranspose(PacketBlock<Packet16h, 16>& kernel) {
1101 __m256i a = _mm256_castph_si256(kernel.packet[0]);
1102 __m256i b = _mm256_castph_si256(kernel.packet[1]);
1103 __m256i c = _mm256_castph_si256(kernel.packet[2]);
1104 __m256i d = _mm256_castph_si256(kernel.packet[3]);
1105 __m256i e = _mm256_castph_si256(kernel.packet[4]);
1106 __m256i f = _mm256_castph_si256(kernel.packet[5]);
1107 __m256i g = _mm256_castph_si256(kernel.packet[6]);
1108 __m256i h = _mm256_castph_si256(kernel.packet[7]);
1109 __m256i i = _mm256_castph_si256(kernel.packet[8]);
1110 __m256i j = _mm256_castph_si256(kernel.packet[9]);
1111 __m256i k = _mm256_castph_si256(kernel.packet[10]);
1112 __m256i l = _mm256_castph_si256(kernel.packet[11]);
1113 __m256i m = _mm256_castph_si256(kernel.packet[12]);
1114 __m256i n = _mm256_castph_si256(kernel.packet[13]);
1115 __m256i o = _mm256_castph_si256(kernel.packet[14]);
1116 __m256i p = _mm256_castph_si256(kernel.packet[15]);
1117
1118 __m256i ab_07 = _mm256_unpacklo_epi16(a, b);
1119 __m256i cd_07 = _mm256_unpacklo_epi16(c, d);
1120 __m256i ef_07 = _mm256_unpacklo_epi16(e, f);
1121 __m256i gh_07 = _mm256_unpacklo_epi16(g, h);
1122 __m256i ij_07 = _mm256_unpacklo_epi16(i, j);
1123 __m256i kl_07 = _mm256_unpacklo_epi16(k, l);
1124 __m256i mn_07 = _mm256_unpacklo_epi16(m, n);
1125 __m256i op_07 = _mm256_unpacklo_epi16(o, p);
1126
1127 __m256i ab_8f = _mm256_unpackhi_epi16(a, b);
1128 __m256i cd_8f = _mm256_unpackhi_epi16(c, d);
1129 __m256i ef_8f = _mm256_unpackhi_epi16(e, f);
1130 __m256i gh_8f = _mm256_unpackhi_epi16(g, h);
1131 __m256i ij_8f = _mm256_unpackhi_epi16(i, j);
1132 __m256i kl_8f = _mm256_unpackhi_epi16(k, l);
1133 __m256i mn_8f = _mm256_unpackhi_epi16(m, n);
1134 __m256i op_8f = _mm256_unpackhi_epi16(o, p);
1135
1136 __m256i abcd_03 = _mm256_unpacklo_epi32(ab_07, cd_07);
1137 __m256i abcd_47 = _mm256_unpackhi_epi32(ab_07, cd_07);
1138 __m256i efgh_03 = _mm256_unpacklo_epi32(ef_07, gh_07);
1139 __m256i efgh_47 = _mm256_unpackhi_epi32(ef_07, gh_07);
1140 __m256i ijkl_03 = _mm256_unpacklo_epi32(ij_07, kl_07);
1141 __m256i ijkl_47 = _mm256_unpackhi_epi32(ij_07, kl_07);
1142 __m256i mnop_03 = _mm256_unpacklo_epi32(mn_07, op_07);
1143 __m256i mnop_47 = _mm256_unpackhi_epi32(mn_07, op_07);
1144
1145 __m256i abcd_8b = _mm256_unpacklo_epi32(ab_8f, cd_8f);
1146 __m256i abcd_cf = _mm256_unpackhi_epi32(ab_8f, cd_8f);
1147 __m256i efgh_8b = _mm256_unpacklo_epi32(ef_8f, gh_8f);
1148 __m256i efgh_cf = _mm256_unpackhi_epi32(ef_8f, gh_8f);
1149 __m256i ijkl_8b = _mm256_unpacklo_epi32(ij_8f, kl_8f);
1150 __m256i ijkl_cf = _mm256_unpackhi_epi32(ij_8f, kl_8f);
1151 __m256i mnop_8b = _mm256_unpacklo_epi32(mn_8f, op_8f);
1152 __m256i mnop_cf = _mm256_unpackhi_epi32(mn_8f, op_8f);
1153
1154 __m256i abcdefgh_01 = _mm256_unpacklo_epi64(abcd_03, efgh_03);
1155 __m256i abcdefgh_23 = _mm256_unpackhi_epi64(abcd_03, efgh_03);
1156 __m256i ijklmnop_01 = _mm256_unpacklo_epi64(ijkl_03, mnop_03);
1157 __m256i ijklmnop_23 = _mm256_unpackhi_epi64(ijkl_03, mnop_03);
1158 __m256i abcdefgh_45 = _mm256_unpacklo_epi64(abcd_47, efgh_47);
1159 __m256i abcdefgh_67 = _mm256_unpackhi_epi64(abcd_47, efgh_47);
1160 __m256i ijklmnop_45 = _mm256_unpacklo_epi64(ijkl_47, mnop_47);
1161 __m256i ijklmnop_67 = _mm256_unpackhi_epi64(ijkl_47, mnop_47);
1162 __m256i abcdefgh_89 = _mm256_unpacklo_epi64(abcd_8b, efgh_8b);
1163 __m256i abcdefgh_ab = _mm256_unpackhi_epi64(abcd_8b, efgh_8b);
1164 __m256i ijklmnop_89 = _mm256_unpacklo_epi64(ijkl_8b, mnop_8b);
1165 __m256i ijklmnop_ab = _mm256_unpackhi_epi64(ijkl_8b, mnop_8b);
1166 __m256i abcdefgh_cd = _mm256_unpacklo_epi64(abcd_cf, efgh_cf);
1167 __m256i abcdefgh_ef = _mm256_unpackhi_epi64(abcd_cf, efgh_cf);
1168 __m256i ijklmnop_cd = _mm256_unpacklo_epi64(ijkl_cf, mnop_cf);
1169 __m256i ijklmnop_ef = _mm256_unpackhi_epi64(ijkl_cf, mnop_cf);
1170
1171 // NOTE: no unpacklo/hi instr in this case, so using permute instr.
1172 __m256i a_p_0 = _mm256_permute2x128_si256(abcdefgh_01, ijklmnop_01, 0x20);
1173 __m256i a_p_1 = _mm256_permute2x128_si256(abcdefgh_23, ijklmnop_23, 0x20);
1174 __m256i a_p_2 = _mm256_permute2x128_si256(abcdefgh_45, ijklmnop_45, 0x20);
1175 __m256i a_p_3 = _mm256_permute2x128_si256(abcdefgh_67, ijklmnop_67, 0x20);
1176 __m256i a_p_4 = _mm256_permute2x128_si256(abcdefgh_89, ijklmnop_89, 0x20);
1177 __m256i a_p_5 = _mm256_permute2x128_si256(abcdefgh_ab, ijklmnop_ab, 0x20);
1178 __m256i a_p_6 = _mm256_permute2x128_si256(abcdefgh_cd, ijklmnop_cd, 0x20);
1179 __m256i a_p_7 = _mm256_permute2x128_si256(abcdefgh_ef, ijklmnop_ef, 0x20);
1180 __m256i a_p_8 = _mm256_permute2x128_si256(abcdefgh_01, ijklmnop_01, 0x31);
1181 __m256i a_p_9 = _mm256_permute2x128_si256(abcdefgh_23, ijklmnop_23, 0x31);
1182 __m256i a_p_a = _mm256_permute2x128_si256(abcdefgh_45, ijklmnop_45, 0x31);
1183 __m256i a_p_b = _mm256_permute2x128_si256(abcdefgh_67, ijklmnop_67, 0x31);
1184 __m256i a_p_c = _mm256_permute2x128_si256(abcdefgh_89, ijklmnop_89, 0x31);
1185 __m256i a_p_d = _mm256_permute2x128_si256(abcdefgh_ab, ijklmnop_ab, 0x31);
1186 __m256i a_p_e = _mm256_permute2x128_si256(abcdefgh_cd, ijklmnop_cd, 0x31);
1187 __m256i a_p_f = _mm256_permute2x128_si256(abcdefgh_ef, ijklmnop_ef, 0x31);
1188
1189 kernel.packet[0] = _mm256_castsi256_ph(a_p_0);
1190 kernel.packet[1] = _mm256_castsi256_ph(a_p_1);
1191 kernel.packet[2] = _mm256_castsi256_ph(a_p_2);
1192 kernel.packet[3] = _mm256_castsi256_ph(a_p_3);
1193 kernel.packet[4] = _mm256_castsi256_ph(a_p_4);
1194 kernel.packet[5] = _mm256_castsi256_ph(a_p_5);
1195 kernel.packet[6] = _mm256_castsi256_ph(a_p_6);
1196 kernel.packet[7] = _mm256_castsi256_ph(a_p_7);
1197 kernel.packet[8] = _mm256_castsi256_ph(a_p_8);
1198 kernel.packet[9] = _mm256_castsi256_ph(a_p_9);
1199 kernel.packet[10] = _mm256_castsi256_ph(a_p_a);
1200 kernel.packet[11] = _mm256_castsi256_ph(a_p_b);
1201 kernel.packet[12] = _mm256_castsi256_ph(a_p_c);
1202 kernel.packet[13] = _mm256_castsi256_ph(a_p_d);
1203 kernel.packet[14] = _mm256_castsi256_ph(a_p_e);
1204 kernel.packet[15] = _mm256_castsi256_ph(a_p_f);
1205}
1206
1207EIGEN_STRONG_INLINE void ptranspose(PacketBlock<Packet16h, 8>& kernel) {
1208 EIGEN_ALIGN64 half in[8][16];
1209 pstore<half>(in[0], kernel.packet[0]);
1210 pstore<half>(in[1], kernel.packet[1]);
1211 pstore<half>(in[2], kernel.packet[2]);
1212 pstore<half>(in[3], kernel.packet[3]);
1213 pstore<half>(in[4], kernel.packet[4]);
1214 pstore<half>(in[5], kernel.packet[5]);
1215 pstore<half>(in[6], kernel.packet[6]);
1216 pstore<half>(in[7], kernel.packet[7]);
1217
1218 EIGEN_ALIGN64 half out[8][16];
1219
1220 for (int i = 0; i < 8; ++i) {
1221 for (int j = 0; j < 8; ++j) {
1222 out[i][j] = in[j][2 * i];
1223 }
1224 for (int j = 0; j < 8; ++j) {
1225 out[i][j + 8] = in[j][2 * i + 1];
1226 }
1227 }
1228
1229 kernel.packet[0] = pload<Packet16h>(out[0]);
1230 kernel.packet[1] = pload<Packet16h>(out[1]);
1231 kernel.packet[2] = pload<Packet16h>(out[2]);
1232 kernel.packet[3] = pload<Packet16h>(out[3]);
1233 kernel.packet[4] = pload<Packet16h>(out[4]);
1234 kernel.packet[5] = pload<Packet16h>(out[5]);
1235 kernel.packet[6] = pload<Packet16h>(out[6]);
1236 kernel.packet[7] = pload<Packet16h>(out[7]);
1237}
1238
1239EIGEN_STRONG_INLINE void ptranspose(PacketBlock<Packet16h, 4>& kernel) {
1240 EIGEN_ALIGN64 half in[4][16];
1241 pstore<half>(in[0], kernel.packet[0]);
1242 pstore<half>(in[1], kernel.packet[1]);
1243 pstore<half>(in[2], kernel.packet[2]);
1244 pstore<half>(in[3], kernel.packet[3]);
1245
1246 EIGEN_ALIGN64 half out[4][16];
1247
1248 for (int i = 0; i < 4; ++i) {
1249 for (int j = 0; j < 4; ++j) {
1250 out[i][j] = in[j][4 * i];
1251 }
1252 for (int j = 0; j < 4; ++j) {
1253 out[i][j + 4] = in[j][4 * i + 1];
1254 }
1255 for (int j = 0; j < 4; ++j) {
1256 out[i][j + 8] = in[j][4 * i + 2];
1257 }
1258 for (int j = 0; j < 4; ++j) {
1259 out[i][j + 12] = in[j][4 * i + 3];
1260 }
1261 }
1262
1263 kernel.packet[0] = pload<Packet16h>(out[0]);
1264 kernel.packet[1] = pload<Packet16h>(out[1]);
1265 kernel.packet[2] = pload<Packet16h>(out[2]);
1266 kernel.packet[3] = pload<Packet16h>(out[3]);
1267}
1268
1269EIGEN_STRONG_INLINE void ptranspose(PacketBlock<Packet8h, 8>& kernel) {
1270 __m128i a = _mm_castph_si128(kernel.packet[0]);
1271 __m128i b = _mm_castph_si128(kernel.packet[1]);
1272 __m128i c = _mm_castph_si128(kernel.packet[2]);
1273 __m128i d = _mm_castph_si128(kernel.packet[3]);
1274 __m128i e = _mm_castph_si128(kernel.packet[4]);
1275 __m128i f = _mm_castph_si128(kernel.packet[5]);
1276 __m128i g = _mm_castph_si128(kernel.packet[6]);
1277 __m128i h = _mm_castph_si128(kernel.packet[7]);
1278
1279 __m128i a03b03 = _mm_unpacklo_epi16(a, b);
1280 __m128i c03d03 = _mm_unpacklo_epi16(c, d);
1281 __m128i e03f03 = _mm_unpacklo_epi16(e, f);
1282 __m128i g03h03 = _mm_unpacklo_epi16(g, h);
1283 __m128i a47b47 = _mm_unpackhi_epi16(a, b);
1284 __m128i c47d47 = _mm_unpackhi_epi16(c, d);
1285 __m128i e47f47 = _mm_unpackhi_epi16(e, f);
1286 __m128i g47h47 = _mm_unpackhi_epi16(g, h);
1287
1288 __m128i a01b01c01d01 = _mm_unpacklo_epi32(a03b03, c03d03);
1289 __m128i a23b23c23d23 = _mm_unpackhi_epi32(a03b03, c03d03);
1290 __m128i e01f01g01h01 = _mm_unpacklo_epi32(e03f03, g03h03);
1291 __m128i e23f23g23h23 = _mm_unpackhi_epi32(e03f03, g03h03);
1292 __m128i a45b45c45d45 = _mm_unpacklo_epi32(a47b47, c47d47);
1293 __m128i a67b67c67d67 = _mm_unpackhi_epi32(a47b47, c47d47);
1294 __m128i e45f45g45h45 = _mm_unpacklo_epi32(e47f47, g47h47);
1295 __m128i e67f67g67h67 = _mm_unpackhi_epi32(e47f47, g47h47);
1296
1297 __m128i a0b0c0d0e0f0g0h0 = _mm_unpacklo_epi64(a01b01c01d01, e01f01g01h01);
1298 __m128i a1b1c1d1e1f1g1h1 = _mm_unpackhi_epi64(a01b01c01d01, e01f01g01h01);
1299 __m128i a2b2c2d2e2f2g2h2 = _mm_unpacklo_epi64(a23b23c23d23, e23f23g23h23);
1300 __m128i a3b3c3d3e3f3g3h3 = _mm_unpackhi_epi64(a23b23c23d23, e23f23g23h23);
1301 __m128i a4b4c4d4e4f4g4h4 = _mm_unpacklo_epi64(a45b45c45d45, e45f45g45h45);
1302 __m128i a5b5c5d5e5f5g5h5 = _mm_unpackhi_epi64(a45b45c45d45, e45f45g45h45);
1303 __m128i a6b6c6d6e6f6g6h6 = _mm_unpacklo_epi64(a67b67c67d67, e67f67g67h67);
1304 __m128i a7b7c7d7e7f7g7h7 = _mm_unpackhi_epi64(a67b67c67d67, e67f67g67h67);
1305
1306 kernel.packet[0] = _mm_castsi128_ph(a0b0c0d0e0f0g0h0);
1307 kernel.packet[1] = _mm_castsi128_ph(a1b1c1d1e1f1g1h1);
1308 kernel.packet[2] = _mm_castsi128_ph(a2b2c2d2e2f2g2h2);
1309 kernel.packet[3] = _mm_castsi128_ph(a3b3c3d3e3f3g3h3);
1310 kernel.packet[4] = _mm_castsi128_ph(a4b4c4d4e4f4g4h4);
1311 kernel.packet[5] = _mm_castsi128_ph(a5b5c5d5e5f5g5h5);
1312 kernel.packet[6] = _mm_castsi128_ph(a6b6c6d6e6f6g6h6);
1313 kernel.packet[7] = _mm_castsi128_ph(a7b7c7d7e7f7g7h7);
1314}
1315
1316EIGEN_STRONG_INLINE void ptranspose(PacketBlock<Packet8h, 4>& kernel) {
1317 EIGEN_ALIGN32 Eigen::half in[4][8];
1318 pstore<Eigen::half>(in[0], kernel.packet[0]);
1319 pstore<Eigen::half>(in[1], kernel.packet[1]);
1320 pstore<Eigen::half>(in[2], kernel.packet[2]);
1321 pstore<Eigen::half>(in[3], kernel.packet[3]);
1322
1323 EIGEN_ALIGN32 Eigen::half out[4][8];
1324
1325 for (int i = 0; i < 4; ++i) {
1326 for (int j = 0; j < 4; ++j) {
1327 out[i][j] = in[j][2 * i];
1328 }
1329 for (int j = 0; j < 4; ++j) {
1330 out[i][j + 4] = in[j][2 * i + 1];
1331 }
1332 }
1333
1334 kernel.packet[0] = pload<Packet8h>(out[0]);
1335 kernel.packet[1] = pload<Packet8h>(out[1]);
1336 kernel.packet[2] = pload<Packet8h>(out[2]);
1337 kernel.packet[3] = pload<Packet8h>(out[3]);
1338}
1339
1340// preverse
1341
1342template <>
1343EIGEN_STRONG_INLINE Packet32h preverse(const Packet32h& a) {
1344 return _mm512_permutexvar_ph(_mm512_set_epi16(0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19,
1345 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31),
1346 a);
1347}
1348
1349template <>
1350EIGEN_STRONG_INLINE Packet16h preverse(const Packet16h& a) {
1351 __m128i m = _mm_setr_epi8(14, 15, 12, 13, 10, 11, 8, 9, 6, 7, 4, 5, 2, 3, 0, 1);
1352 return _mm256_castsi256_ph(_mm256_insertf128_si256(
1353 _mm256_castsi128_si256(_mm_shuffle_epi8(_mm256_extractf128_si256(_mm256_castph_si256(a), 1), m)),
1354 _mm_shuffle_epi8(_mm256_extractf128_si256(_mm256_castph_si256(a), 0), m), 1));
1355}
1356
1357template <>
1358EIGEN_STRONG_INLINE Packet8h preverse(const Packet8h& a) {
1359 __m128i m = _mm_setr_epi8(14, 15, 12, 13, 10, 11, 8, 9, 6, 7, 4, 5, 2, 3, 0, 1);
1360 return _mm_castsi128_ph(_mm_shuffle_epi8(_mm_castph_si128(a), m));
1361}
1362
1363// pscatter
1364
1365template <>
1366EIGEN_STRONG_INLINE void pscatter<half, Packet32h>(half* to, const Packet32h& from, Index stride) {
1367 EIGEN_ALIGN64 half aux[32];
1368 pstore(aux, from);
1369
1370 EIGEN_UNROLL_LOOP
1371 for (int i = 0; i < 32; i++) {
1372 to[stride * i] = aux[i];
1373 }
1374}
1375template <>
1376EIGEN_STRONG_INLINE void pscatter<half, Packet16h>(half* to, const Packet16h& from, Index stride) {
1377 EIGEN_ALIGN64 half aux[16];
1378 pstore(aux, from);
1379 to[stride * 0] = aux[0];
1380 to[stride * 1] = aux[1];
1381 to[stride * 2] = aux[2];
1382 to[stride * 3] = aux[3];
1383 to[stride * 4] = aux[4];
1384 to[stride * 5] = aux[5];
1385 to[stride * 6] = aux[6];
1386 to[stride * 7] = aux[7];
1387 to[stride * 8] = aux[8];
1388 to[stride * 9] = aux[9];
1389 to[stride * 10] = aux[10];
1390 to[stride * 11] = aux[11];
1391 to[stride * 12] = aux[12];
1392 to[stride * 13] = aux[13];
1393 to[stride * 14] = aux[14];
1394 to[stride * 15] = aux[15];
1395}
1396
1397template <>
1398EIGEN_STRONG_INLINE void pscatter<Eigen::half, Packet8h>(Eigen::half* to, const Packet8h& from, Index stride) {
1399 EIGEN_ALIGN32 Eigen::half aux[8];
1400 pstore(aux, from);
1401 to[stride * 0] = aux[0];
1402 to[stride * 1] = aux[1];
1403 to[stride * 2] = aux[2];
1404 to[stride * 3] = aux[3];
1405 to[stride * 4] = aux[4];
1406 to[stride * 5] = aux[5];
1407 to[stride * 6] = aux[6];
1408 to[stride * 7] = aux[7];
1409}
1410
1411// pgather
1412
1413template <>
1414EIGEN_STRONG_INLINE Packet32h pgather<Eigen::half, Packet32h>(const Eigen::half* from, Index stride) {
1415 return _mm512_set_ph(from[31 * stride].x, from[30 * stride].x, from[29 * stride].x, from[28 * stride].x,
1416 from[27 * stride].x, from[26 * stride].x, from[25 * stride].x, from[24 * stride].x,
1417 from[23 * stride].x, from[22 * stride].x, from[21 * stride].x, from[20 * stride].x,
1418 from[19 * stride].x, from[18 * stride].x, from[17 * stride].x, from[16 * stride].x,
1419 from[15 * stride].x, from[14 * stride].x, from[13 * stride].x, from[12 * stride].x,
1420 from[11 * stride].x, from[10 * stride].x, from[9 * stride].x, from[8 * stride].x,
1421 from[7 * stride].x, from[6 * stride].x, from[5 * stride].x, from[4 * stride].x,
1422 from[3 * stride].x, from[2 * stride].x, from[1 * stride].x, from[0 * stride].x);
1423}
1424
1425template <>
1426EIGEN_STRONG_INLINE Packet16h pgather<Eigen::half, Packet16h>(const Eigen::half* from, Index stride) {
1427 return _mm256_set_ph(from[15 * stride].x, from[14 * stride].x, from[13 * stride].x, from[12 * stride].x,
1428 from[11 * stride].x, from[10 * stride].x, from[9 * stride].x, from[8 * stride].x,
1429 from[7 * stride].x, from[6 * stride].x, from[5 * stride].x, from[4 * stride].x,
1430 from[3 * stride].x, from[2 * stride].x, from[1 * stride].x, from[0 * stride].x);
1431}
1432
1433template <>
1434EIGEN_STRONG_INLINE Packet8h pgather<Eigen::half, Packet8h>(const Eigen::half* from, Index stride) {
1435 return _mm_set_ph(from[7 * stride].x, from[6 * stride].x, from[5 * stride].x, from[4 * stride].x, from[3 * stride].x,
1436 from[2 * stride].x, from[1 * stride].x, from[0 * stride].x);
1437}
1438
1439/*---------------- load/store segment support ----------------*/
1440
1441// There are no masked FP16 moves; the word-granular AVX-512BW ones move the same bits.
1442
1443template <>
1444struct has_packet_segment<Packet32h> : std::true_type {};
1445
1446template <>
1447struct has_packet_segment<Packet16h> : std::true_type {};
1448
1449template <>
1450struct has_packet_segment<Packet8h> : std::true_type {};
1451
1452template <>
1453inline Packet32h ploaduSegment<Packet32h>(const Eigen::half* from, Index begin, Index count) {
1454 return _mm512_castsi512_ph(_mm512_maskz_loadu_epi16(static_cast<__mmask32>(segment_kmask(begin, count)), from));
1455}
1456
1457template <>
1458inline void pstoreuSegment<Eigen::half, Packet32h>(Eigen::half* to, const Packet32h& from, Index begin, Index count) {
1459 _mm512_mask_storeu_epi16(to, static_cast<__mmask32>(segment_kmask(begin, count)), _mm512_castph_si512(from));
1460}
1461
1462template <>
1463inline Packet16h ploaduSegment<Packet16h>(const Eigen::half* from, Index begin, Index count) {
1464 return _mm256_castsi256_ph(_mm256_maskz_loadu_epi16(static_cast<__mmask16>(segment_kmask(begin, count)), from));
1465}
1466
1467template <>
1468inline void pstoreuSegment<Eigen::half, Packet16h>(Eigen::half* to, const Packet16h& from, Index begin, Index count) {
1469 _mm256_mask_storeu_epi16(to, static_cast<__mmask16>(segment_kmask(begin, count)), _mm256_castph_si256(from));
1470}
1471
1472template <>
1473inline Packet8h ploaduSegment<Packet8h>(const Eigen::half* from, Index begin, Index count) {
1474 return _mm_castsi128_ph(_mm_maskz_loadu_epi16(static_cast<__mmask8>(segment_kmask(begin, count)), from));
1475}
1476
1477template <>
1478inline void pstoreuSegment<Eigen::half, Packet8h>(Eigen::half* to, const Packet8h& from, Index begin, Index count) {
1479 _mm_mask_storeu_epi16(to, static_cast<__mmask8>(segment_kmask(begin, count)), _mm_castph_si128(from));
1480}
1481
1482/*---------------- end load/store segment support ----------------*/
1483
1484} // end namespace internal
1485} // end namespace Eigen
1486
1487#endif // EIGEN_PACKET_MATH_FP16_AVX512_H
@ Aligned64
Definition Constants.h:240
@ Aligned32
Definition Constants.h:239
@ Aligned16
Definition Constants.h:238