Eigen  5.0.1
 
Loading...
Searching...
No Matches
GemmKernel.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2022 Intel Corporation
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_CORE_ARCH_AVX512_GEMM_KERNEL_H
12#define EIGEN_CORE_ARCH_AVX512_GEMM_KERNEL_H
13
14#if EIGEN_COMP_MSVC
15#include <intrin.h>
16#else
17#include <x86intrin.h>
18#endif
19#include <immintrin.h>
20#include <type_traits>
21#include <utility>
22
23// IWYU pragma: private
24#include "../../InternalHeaderCheck.h"
25
26#if !defined(EIGEN_USE_AVX512_GEMM_KERNELS)
27#define EIGEN_USE_AVX512_GEMM_KERNELS 1
28#endif
29
30#define SECOND_FETCH (32)
31#if (EIGEN_COMP_GNUC_STRICT != 0) && !defined(EIGEN_ARCH_AVX512_GEMM_KERNEL_USE_LESS_A_REGS)
32// Use less registers to load A elements to workaround compiler spills. Lose a
33// bit of performance (less than ~2%).
34#define EIGEN_ARCH_AVX512_GEMM_KERNEL_USE_LESS_A_REGS
35#endif
36
37namespace Eigen {
38namespace internal {
39
40#if EIGEN_USE_AVX512_GEMM_KERNELS
41
42template <typename Scalar, bool is_unit_inc>
43class gemm_class {
44 using vec = typename packet_traits<Scalar>::type;
45 using vec_ymm = typename unpacket_traits<vec>::half;
46 using vec_xmm = typename unpacket_traits<vec_ymm>::half;
47 using umask_t = typename unpacket_traits<vec>::mask_t;
48
49 static constexpr bool is_f32 = sizeof(Scalar) == sizeof(float);
50 static constexpr bool is_f64 = sizeof(Scalar) == sizeof(double);
51
52#ifndef EIGEN_ARCH_AVX512_GEMM_KERNEL_USE_LESS_A_REGS
53 static constexpr bool use_less_a_regs = !is_unit_inc;
54#else
55 static constexpr bool use_less_a_regs = true;
56#endif
57#ifndef EIGEN_ARCH_AVX512_GEMM_KERNEL_USE_LESS_B_REGS
58 static constexpr bool use_less_b_regs = !is_unit_inc;
59#else
60 static constexpr bool use_less_b_regs = true;
61#endif
62
63 static constexpr int a_regs[] = {0, 1, 2, use_less_a_regs ? 0 : 3, use_less_a_regs ? 1 : 4, use_less_a_regs ? 2 : 5};
64 static constexpr int b_regs[] = {6, use_less_b_regs ? 6 : 7};
65 static constexpr int c_regs[] = {
66 8, 16, 24, 9, 17, 25, 10, 18, 26, 11, 19, 27, 12, 20, 28, 13, 21, 29, 14, 22, 30, 15, 23, 31,
67 };
68
69 static constexpr int alpha_load_reg = 0;
70 static constexpr int c_load_regs[] = {1, 2, 6};
71
72 static constexpr int a_shift = 128;
73 static constexpr int b_shift = 128;
74
75 static constexpr int nelems_in_cache_line = is_f32 ? 16 : 8;
76 static constexpr int a_prefetch_size = nelems_in_cache_line * 2;
77 static constexpr int b_prefetch_size = nelems_in_cache_line * 8;
78
79 vec zmm[32];
80 umask_t mask;
81
82 // gemm arguments.
83 Index m;
84 const Index n, k, ldc;
85 const Index inc;
86 const Scalar* alpha;
87
88 const Scalar *a, *b;
89 Scalar* c;
90
91 const bool is_alpha1;
92 const bool is_beta0;
93
94 const Index a_stride, b_stride;
95 const Index a_off, b_off;
96
97 EIGEN_ALWAYS_INLINE void prefetch_a(const Scalar* a_addr) {
98 _mm_prefetch((char*)(a_prefetch_size + a_addr - a_shift), _MM_HINT_T0);
99 }
100
101 EIGEN_ALWAYS_INLINE void prefetch_b(const Scalar* b_addr) {
102 _mm_prefetch((char*)(b_prefetch_size + b_addr - b_shift), _MM_HINT_T0);
103 }
104
105 EIGEN_ALWAYS_INLINE void prefetch_x(const Scalar* x_addr) { _mm_prefetch((char*)(x_addr - a_shift), _MM_HINT_T2); }
106
107 EIGEN_ALWAYS_INLINE void prefetch_c(const Scalar* c_addr) {
108#if defined(__PRFCHW__) && __PRFCHW__ == 1
109 _m_prefetchw((void*)c_addr);
110#else
111 _mm_prefetch((char*)c_addr, _MM_HINT_T0);
112#endif
113 }
114
115 template <int nelems>
116 EIGEN_ALWAYS_INLINE void a_load(vec& a_reg, const Scalar* a_addr) {
117 switch (nelems * sizeof(*a_addr) * 8) {
118 default:
119 case 512 * 3:
120 a_reg = ploadu<vec>(a_addr);
121 break;
122 case 512 * 2:
123 a_reg = ploadu<vec>(a_addr);
124 break;
125 case 512 * 1:
126 a_reg = ploadu<vec>(a_addr);
127 break;
128 case 256 * 1:
129 a_reg = preinterpret<vec>(_mm512_broadcast_f64x4(ploadu<Packet4d>(reinterpret_cast<const double*>(a_addr))));
130 break;
131 case 128 * 1:
132 a_reg = preinterpret<vec>(_mm512_broadcast_f32x4(ploadu<Packet4f>(reinterpret_cast<const float*>(a_addr))));
133 break;
134 case 64 * 1:
135 a_reg = preinterpret<vec>(pload1<Packet8d>(reinterpret_cast<const double*>(a_addr)));
136 break;
137 case 32 * 1:
138 a_reg = pload1<vec>(a_addr);
139 break;
140 }
141 }
142
143 EIGEN_ALWAYS_INLINE void b_load(vec& b_reg, const Scalar* b_addr) { b_reg = pload1<vec>(b_addr); }
144
145 template <int nelems>
146 EIGEN_ALWAYS_INLINE void c_store(Scalar* mem, vec& src) {
147 EIGEN_IF_CONSTEXPR (is_unit_inc) {
148 switch (nelems * sizeof(*mem) * 8) {
149 default:
150 case 512 * 3:
151 pstoreu(mem, src);
152 break;
153 case 512 * 2:
154 pstoreu(mem, src);
155 break;
156 case 512 * 1:
157 pstoreu(mem, src);
158 break;
159 case 256 * 1:
160 pstoreu(mem, preinterpret<vec_ymm>(src));
161 break;
162 case 128 * 1:
163 pstoreu(mem, preinterpret<vec_xmm>(src));
164 break;
165 case 64 * 1:
166 pstorel(mem, preinterpret<vec_xmm>(src));
167 break;
168 case 32 * 1:
169 pstores(mem, preinterpret<vec_xmm>(src));
170 break;
171 }
172 } else {
173 switch (nelems * sizeof(*mem) * 8) {
174 default:
175 case 512 * 3:
176 pscatter(mem, src, inc);
177 break;
178 case 512 * 2:
179 pscatter(mem, src, inc);
180 break;
181 case 512 * 1:
182 pscatter(mem, src, inc);
183 break;
184 case 256 * 1:
185 pscatter(mem, src, inc, mask);
186 break;
187 case 128 * 1:
188 pscatter(mem, src, inc, mask);
189 break;
190 case 64 * 1:
191 pscatter(mem, src, inc, mask);
192 break;
193 case 32 * 1:
194 pscatter(mem, src, inc, mask);
195 break;
196 }
197 }
198 }
199
200 template <int nelems>
201 EIGEN_ALWAYS_INLINE void vaddm(vec& dst, const Scalar* mem, vec& src, vec& reg) {
202 EIGEN_IF_CONSTEXPR (is_unit_inc) {
203 switch (nelems * sizeof(*mem) * 8) {
204 default:
205 case 512 * 3:
206 dst = padd(src, ploadu<vec>(mem));
207 break;
208 case 512 * 2:
209 dst = padd(src, ploadu<vec>(mem));
210 break;
211 case 512 * 1:
212 dst = padd(src, ploadu<vec>(mem));
213 break;
214 case 256 * 1:
215 dst = preinterpret<vec>(padd(preinterpret<vec_ymm>(src), ploadu<vec_ymm>(mem)));
216 break;
217 case 128 * 1:
218 dst = preinterpret<vec>(padd(preinterpret<vec_xmm>(src), ploadu<vec_xmm>(mem)));
219 break;
220 case 64 * 1:
221 dst = preinterpret<vec>(padd(preinterpret<vec_xmm>(src), ploadl<vec_xmm>(mem)));
222 break;
223 case 32 * 1:
224 dst = preinterpret<vec>(padds(preinterpret<vec_xmm>(src), ploads<vec_xmm>(mem)));
225 break;
226 }
227 } else {
228 // Zero out scratch register
229 reg = pzero(reg);
230
231 switch (nelems * sizeof(*mem) * 8) {
232 default:
233 case 512 * 3:
234 reg = pgather<Scalar, vec>(mem, inc);
235 dst = padd(src, reg);
236 break;
237 case 512 * 2:
238 reg = pgather<Scalar, vec>(mem, inc);
239 dst = padd(src, reg);
240 break;
241 case 512 * 1:
242 reg = pgather<Scalar, vec>(mem, inc);
243 dst = padd(src, reg);
244 break;
245 case 256 * 1:
246 reg = preinterpret<vec>(pgather<Scalar, vec_ymm>(mem, inc));
247 dst = preinterpret<vec>(padd(preinterpret<vec_ymm>(src), preinterpret<vec_ymm>(reg)));
248 break;
249 case 128 * 1:
250 reg = preinterpret<vec>(pgather<Scalar, vec_xmm>(mem, inc));
251 dst = preinterpret<vec>(padd(preinterpret<vec_xmm>(src), preinterpret<vec_xmm>(reg)));
252 break;
253 case 64 * 1:
254 EIGEN_IF_CONSTEXPR (is_f32) {
255 reg = pgather(reg, mem, inc, mask);
256 dst = preinterpret<vec>(padd(preinterpret<vec_xmm>(src), preinterpret<vec_xmm>(reg)));
257 } else {
258 dst = preinterpret<vec>(padd(preinterpret<vec_xmm>(src), ploadl<vec_xmm>(mem)));
259 }
260 break;
261 case 32 * 1:
262 dst = preinterpret<vec>(padds(preinterpret<vec_xmm>(src), ploads<vec_xmm>(mem)));
263 break;
264 }
265 }
266 }
267
268 EIGEN_STRONG_INLINE void vfmadd(vec& dst, const vec& src1, const vec& src2) {
269 dst = pmadd(src1, src2, dst);
270
271#if (EIGEN_COMP_GNUC != 0) || (EIGEN_COMP_CLANG != 0)
272 // Workaround register spills for gcc and clang
273 __asm__("#" : [dst] "+v"(dst) : [src1] "%v"(src1), [src2] "v"(src2));
274#endif
275 }
276
277 template <int nelems>
278 EIGEN_ALWAYS_INLINE void vfmaddm(vec& dst, const Scalar* mem, vec& src, vec& scale, vec& reg) {
279 EIGEN_IF_CONSTEXPR (is_unit_inc) {
280 switch (nelems * sizeof(*mem) * 8) {
281 default:
282 case 512 * 3:
283 dst = pmadd(scale, src, ploadu<vec>(mem));
284 break;
285 case 512 * 2:
286 dst = pmadd(scale, src, ploadu<vec>(mem));
287 break;
288 case 512 * 1:
289 dst = pmadd(scale, src, ploadu<vec>(mem));
290 break;
291 case 256 * 1:
292 dst =
293 preinterpret<vec>(pmadd(preinterpret<vec_ymm>(scale), preinterpret<vec_ymm>(src), ploadu<vec_ymm>(mem)));
294 break;
295 case 128 * 1:
296 dst =
297 preinterpret<vec>(pmadd(preinterpret<vec_xmm>(scale), preinterpret<vec_xmm>(src), ploadu<vec_xmm>(mem)));
298 break;
299 case 64 * 1:
300 dst =
301 preinterpret<vec>(pmadd(preinterpret<vec_xmm>(scale), preinterpret<vec_xmm>(src), ploadl<vec_xmm>(mem)));
302 break;
303 case 32 * 1:
304 dst =
305 preinterpret<vec>(pmadds(preinterpret<vec_xmm>(scale), preinterpret<vec_xmm>(src), ploads<vec_xmm>(mem)));
306 break;
307 }
308 } else {
309 // Zero out scratch register
310 reg = pzero(reg);
311
312 switch (nelems * sizeof(*mem) * 8) {
313 default:
314 case 512 * 3:
315 reg = pgather<Scalar, vec>(mem, inc);
316 dst = pmadd(scale, src, reg);
317 break;
318 case 512 * 2:
319 reg = pgather<Scalar, vec>(mem, inc);
320 dst = pmadd(scale, src, reg);
321 break;
322 case 512 * 1:
323 reg = pgather<Scalar, vec>(mem, inc);
324 dst = pmadd(scale, src, reg);
325 break;
326 case 256 * 1:
327 reg = preinterpret<vec>(pgather<Scalar, vec_ymm>(mem, inc));
328 dst = preinterpret<vec>(
329 pmadd(preinterpret<vec_ymm>(scale), preinterpret<vec_ymm>(src), preinterpret<vec_ymm>(reg)));
330 break;
331 case 128 * 1:
332 reg = preinterpret<vec>(pgather<Scalar, vec_xmm>(mem, inc));
333 dst = preinterpret<vec>(
334 pmadd(preinterpret<vec_xmm>(scale), preinterpret<vec_xmm>(src), preinterpret<vec_xmm>(reg)));
335 break;
336 case 64 * 1:
337 EIGEN_IF_CONSTEXPR (is_f32) {
338 reg = pgather(reg, mem, inc, mask);
339 dst = preinterpret<vec>(
340 pmadd(preinterpret<vec_xmm>(scale), preinterpret<vec_xmm>(src), preinterpret<vec_xmm>(reg)));
341 } else {
342 dst = preinterpret<vec>(
343 pmadd(preinterpret<vec_xmm>(scale), preinterpret<vec_xmm>(src), ploadl<vec_xmm>(mem)));
344 }
345 break;
346 case 32 * 1:
347 dst =
348 preinterpret<vec>(pmadds(preinterpret<vec_xmm>(scale), preinterpret<vec_xmm>(src), ploads<vec_xmm>(mem)));
349 break;
350 }
351 }
352 }
353
354 template <int index, int endY, int nelems>
355 EIGEN_ALWAYS_INLINE void a_load_one(const Scalar* ao) {
356 constexpr int j = index / endY;
357 constexpr int i = index % endY;
358 auto& a_reg = zmm[a_regs[i + (j % 2) * 3]];
359 const Scalar* a_addr = ao + nelems * j + nelems_in_cache_line * i - a_shift;
360 a_load<nelems>(a_reg, a_addr);
361 }
362
363 template <int endY, int nelems, int... indices>
364 EIGEN_ALWAYS_INLINE void a_loads_impl(std::integer_sequence<int, indices...>, const Scalar* ao) {
365 int unused[] = {0, (a_load_one<indices, endY, nelems>(ao), 0)...};
366 EIGEN_UNUSED_VARIABLE(unused);
367 }
368
369 template <int j, int endX, int i, int endY, int nelems>
370 EIGEN_ALWAYS_INLINE void a_loads(const Scalar* ao) {
371 static_assert(j == 0 && i == 0, "a_loads expects to start at zero");
372 a_loads_impl<endY, nelems>(std::make_integer_sequence<int, endX * endY>{}, ao);
373 }
374
375 /* C prefetch loop structure.
376 * for (int un = 0; un < 8; un++) {
377 * if (b_unroll >= un + 1) {
378 * if (un == 4) co2 = co1 + 4 * ldc;
379 *
380 * for (int i = 0; i < um_vecs; i++) {
381 * Scalar *co = (un + 1 <= 4) ? co1 : co2;
382 * auto co_off = (un % 4) * ldc + a_unroll - 1 + i * nelems_in_cache_line * sizeof *co;
383 * prefetch_c(co + co_off);
384 * }
385 * }
386 * }
387 */
388
389 template <int index, int um_vecs, int a_unroll, int b_unroll>
390 EIGEN_ALWAYS_INLINE void prefetch_c_one(Scalar*& co1, Scalar*& co2) {
391 constexpr int un = index / um_vecs;
392 constexpr int i = index % um_vecs;
393
394 EIGEN_IF_CONSTEXPR (b_unroll >= un + 1) {
395 EIGEN_IF_CONSTEXPR (un == 4 && i == 0) {
396 co2 = co1 + 4 * ldc;
397 }
398
399 Scalar* co = (un + 1 <= 4) ? co1 : co2;
400 auto co_off = (un % 4) * ldc + a_unroll - 1 + i * nelems_in_cache_line * sizeof *co;
401 prefetch_c(co + co_off);
402 }
403 }
404
405 template <int um_vecs, int a_unroll, int b_unroll, int... indices>
406 EIGEN_ALWAYS_INLINE void prefetch_cs_impl(std::integer_sequence<int, indices...>, Scalar*& co1, Scalar*& co2) {
407 int unused[] = {0, (prefetch_c_one<indices, um_vecs, a_unroll, b_unroll>(co1, co2), 0)...};
408 EIGEN_UNUSED_VARIABLE(unused);
409 }
410
411 template <int un, int max_b_unroll, int i, int um_vecs, int a_unroll, int b_unroll>
412 EIGEN_ALWAYS_INLINE void prefetch_cs(Scalar*& co1, Scalar*& co2) {
413 static_assert(un == 0 && i == 0, "prefetch_cs expects to start at zero");
414 prefetch_cs_impl<um_vecs, a_unroll, b_unroll>(std::make_integer_sequence<int, max_b_unroll * um_vecs>{}, co1, co2);
415 }
416
417 // load_c
418 template <int i, int idx, int nelems>
419 EIGEN_ALWAYS_INLINE void scale_load_c_one(const Scalar* cox, vec& alpha_reg) {
420 auto& c_reg = zmm[c_regs[i + idx * 3]];
421 auto& c_load_reg = zmm[c_load_regs[i % 3]];
422 auto c_mem = cox;
423 EIGEN_IF_CONSTEXPR (is_unit_inc)
424 c_mem += i * nelems_in_cache_line;
425 else
426 c_mem += i * nelems_in_cache_line * inc;
427
428 if (!is_beta0 && is_alpha1)
429 vaddm<nelems>(c_reg, c_mem, c_reg, c_load_reg);
430 else if (!is_beta0 && !is_alpha1)
431 vfmaddm<nelems>(c_reg, c_mem, c_reg, alpha_reg, c_load_reg);
432 else if (is_beta0 && !is_alpha1)
433 c_reg = pmul(alpha_reg, c_reg);
434 }
435
436 template <int start, int idx, int nelems, int... indices>
437 EIGEN_ALWAYS_INLINE void scale_load_c_impl(std::integer_sequence<int, indices...>, const Scalar* cox,
438 vec& alpha_reg) {
439 int unused[] = {0, (scale_load_c_one<start + indices, idx, nelems>(cox, alpha_reg), 0)...};
440 EIGEN_UNUSED_VARIABLE(unused);
441 }
442
443 template <int i, int um_vecs, int idx, int nelems>
444 EIGEN_ALWAYS_INLINE void scale_load_c(const Scalar* cox, vec& alpha_reg) {
445 static_assert(i <= um_vecs, "invalid C load range");
446 scale_load_c_impl<i, idx, nelems>(std::make_integer_sequence<int, um_vecs - i>{}, cox, alpha_reg);
447 }
448
449 // store_c
450 template <int i, int idx, int nelems>
451 EIGEN_ALWAYS_INLINE void write_c_one(Scalar* cox) {
452 auto& c_reg = zmm[c_regs[i + idx * 3]];
453 auto c_mem = cox;
454 EIGEN_IF_CONSTEXPR (is_unit_inc)
455 c_mem += i * nelems_in_cache_line;
456 else
457 c_mem += i * nelems_in_cache_line * inc;
458
459 c_store<nelems>(c_mem, c_reg);
460 c_reg = pzero(c_reg);
461 }
462
463 template <int start, int idx, int nelems, int... indices>
464 EIGEN_ALWAYS_INLINE void write_c_impl(std::integer_sequence<int, indices...>, Scalar* cox) {
465 int unused[] = {0, (write_c_one<start + indices, idx, nelems>(cox), 0)...};
466 EIGEN_UNUSED_VARIABLE(unused);
467 }
468
469 template <int i, int um_vecs, int idx, int nelems>
470 EIGEN_ALWAYS_INLINE void write_c(Scalar* cox) {
471 static_assert(i <= um_vecs, "invalid C store range");
472 write_c_impl<i, idx, nelems>(std::make_integer_sequence<int, um_vecs - i>{}, cox);
473 }
474
475 /* C update loop structure.
476 * co2 = co1 + ldc;
477 *
478 * auto &alpha_reg = zmm[alpha_load_reg];
479 * if (!is_alpha1) alpha_reg = pload1<vec>(alpha);
480 *
481 * int idx = 0;
482 * for (pow = 1; pow <= 8; pow <<= 1) {
483 *
484 * if (b_unroll >= pow) {
485 * for (count = 1; count < (pow + 1) / 2 + 1; count++) {
486 * if (pow >= 4) co2 += ldc;
487 *
488 * const Scalar *cox = (idx == 0) ? co1 : co2;
489 *
490 * const int um_vecs = numext::div_ceil(a_unroll, nelems_in_cache_line);
491 * scale_load_c<0, um_vecs, idx, a_unroll>(cox, alpha_reg);
492 * write_c<0, um_vecs, idx, a_unroll>(cox);
493 *
494 * idx++;
495 * }
496 * }
497 * }
498 *
499 * if (b_unroll == 1)
500 * co1 += ldc;
501 * else
502 * co1 = co2 + ldc;
503 */
504
505 template <int pow, int a_unroll, int idx>
506 EIGEN_ALWAYS_INLINE void c_update_1count(Scalar*& cox) {
507 EIGEN_IF_CONSTEXPR (pow >= 4) {
508 cox += ldc;
509 }
510
511 const int um_vecs = numext::div_ceil(a_unroll, nelems_in_cache_line);
512 auto& alpha_reg = zmm[alpha_load_reg];
513
514 scale_load_c<0, um_vecs, idx, a_unroll>(cox, alpha_reg);
515 write_c<0, um_vecs, idx, a_unroll>(cox);
516 }
517
518 template <int pow, int a_unroll>
519 EIGEN_ALWAYS_INLINE void c_update_1pow(Scalar*& co1, Scalar*& co2) {
520 constexpr int idx = pow / 2;
521 Scalar*& cox = idx == 0 ? co1 : co2;
522
523 constexpr int max_count = (pow + 1) / 2;
524 static_assert(max_count <= 4, "Unsupported max_count.");
525
526 EIGEN_IF_CONSTEXPR (1 <= max_count) {
527 c_update_1count<pow, a_unroll, idx + 0>(cox);
528 }
529 EIGEN_IF_CONSTEXPR (2 <= max_count) {
530 c_update_1count<pow, a_unroll, idx + 1>(cox);
531 }
532 EIGEN_IF_CONSTEXPR (3 <= max_count) {
533 c_update_1count<pow, a_unroll, idx + 2>(cox);
534 }
535 EIGEN_IF_CONSTEXPR (4 <= max_count) {
536 c_update_1count<pow, a_unroll, idx + 3>(cox);
537 }
538 }
539
540 template <int max_b_unroll, int a_unroll, int b_unroll>
541 EIGEN_ALWAYS_INLINE void c_update(Scalar*& co1, Scalar*& co2) {
542 auto& alpha_reg = zmm[alpha_load_reg];
543
544 co2 = co1 + ldc;
545 if (!is_alpha1) {
546 alpha_reg = pload1<vec>(alpha);
547 }
548 EIGEN_IF_CONSTEXPR (!is_unit_inc && a_unroll < nelems_in_cache_line) {
549 mask = static_cast<umask_t>((1ull << a_unroll) - 1);
550 }
551
552 static_assert(max_b_unroll <= 8, "Unsupported max_b_unroll");
553
554 EIGEN_IF_CONSTEXPR (1 <= max_b_unroll && 1 <= b_unroll) {
555 c_update_1pow<1, a_unroll>(co1, co2);
556 }
557 EIGEN_IF_CONSTEXPR (2 <= max_b_unroll && 2 <= b_unroll) {
558 c_update_1pow<2, a_unroll>(co1, co2);
559 }
560 EIGEN_IF_CONSTEXPR (4 <= max_b_unroll && 4 <= b_unroll) {
561 c_update_1pow<4, a_unroll>(co1, co2);
562 }
563 EIGEN_IF_CONSTEXPR (8 <= max_b_unroll && 8 <= b_unroll) {
564 c_update_1pow<8, a_unroll>(co1, co2);
565 }
566
567 EIGEN_IF_CONSTEXPR (b_unroll == 1)
568 co1 += ldc;
569 else
570 co1 = co2 + ldc;
571 }
572
573 // compute
574 template <int um, int idx, int uk, bool fetch_x, bool ktail>
575 EIGEN_ALWAYS_INLINE void compute_one(const Scalar* ao, const Scalar* bo, int& fetchA_idx, int& fetchB_idx,
576 vec& b_reg) {
577 auto& c_reg = zmm[c_regs[um + idx * 3]];
578 auto& a_reg = zmm[a_regs[um + (uk % 2) * 3]];
579
580 vfmadd(c_reg, a_reg, b_reg);
581
582 EIGEN_IF_CONSTEXPR (!fetch_x && um == 0 &&
583 (((idx == 0 || idx == 6) && (uk % 2 == 0 || is_f64 || ktail)) ||
584 (idx == 3 && (uk % 2 == 1 || is_f64 || ktail)))) {
585 prefetch_a(ao + nelems_in_cache_line * fetchA_idx);
586 fetchA_idx++;
587 }
588
589 EIGEN_IF_CONSTEXPR (um == 0 && idx == 1 && (uk % 2 == 0 || is_f64 || ktail)) {
590 prefetch_b(bo + nelems_in_cache_line * fetchB_idx);
591 fetchB_idx++;
592 }
593 }
594
595 template <int start, int idx, int uk, bool fetch_x, bool ktail, int... indices>
596 EIGEN_ALWAYS_INLINE void compute_impl(std::integer_sequence<int, indices...>, const Scalar* ao, const Scalar* bo,
597 int& fetchA_idx, int& fetchB_idx, vec& b_reg) {
598 int unused[] = {
599 0, (compute_one<start + indices, idx, uk, fetch_x, ktail>(ao, bo, fetchA_idx, fetchB_idx, b_reg), 0)...};
600 EIGEN_UNUSED_VARIABLE(unused);
601 }
602
603 template <int um, int uk, int nelems, bool ktail>
604 EIGEN_ALWAYS_INLINE void load_a_one(const Scalar* ao) {
605 auto& a_reg = zmm[a_regs[um + (uk % 2) * 3]];
606 const Scalar* a_addr = ao + nelems * (1 + !ktail * !use_less_a_regs + uk) + nelems_in_cache_line * um - a_shift;
607 a_load<nelems>(a_reg, a_addr);
608 }
609
610 template <int um, int um_vecs, int idx, int uk, bool fetch_x, bool ktail>
611 EIGEN_ALWAYS_INLINE void compute(const Scalar* ao, const Scalar* bo, int& fetchA_idx, int& fetchB_idx, vec& b_reg) {
612 static_assert(um <= um_vecs, "invalid compute range");
613 compute_impl<um, idx, uk, fetch_x, ktail>(std::make_integer_sequence<int, um_vecs - um>{}, ao, bo, fetchA_idx,
614 fetchB_idx, b_reg);
615 }
616
617 // load_a
618 template <int start, int uk, int nelems, bool ktail, int... indices>
619 EIGEN_ALWAYS_INLINE void load_a_impl(std::integer_sequence<int, indices...>, const Scalar* ao) {
620 int unused[] = {0, (load_a_one<start + indices, uk, nelems, ktail>(ao), 0)...};
621 EIGEN_UNUSED_VARIABLE(unused);
622 }
623
624 template <int um, int um_vecs, int uk, int nelems, bool ktail>
625 EIGEN_ALWAYS_INLINE void load_a(const Scalar* ao) {
626 static_assert(um <= um_vecs, "invalid A load range");
627 load_a_impl<um, uk, nelems, ktail>(std::make_integer_sequence<int, um_vecs - um>{}, ao);
628 }
629
630 template <int uk, int pow, int count, int um_vecs, int b_unroll, bool ktail, bool fetch_x, bool preload_next_k = true>
631 EIGEN_ALWAYS_INLINE void innerkernel_1pow_one(const Scalar*& aa, const Scalar* const& ao, const Scalar* const& bo,
632 int& fetchA_idx, int& fetchB_idx) {
633 const int idx = (pow / 2) + count;
634
635 auto& b_reg = zmm[b_regs[idx % 2]];
636
637 EIGEN_IF_CONSTEXPR (fetch_x && uk == 3 && idx == 0) {
638 prefetch_x(aa);
639 }
640 EIGEN_IF_CONSTEXPR (fetch_x && uk == 3 && idx == 4) {
641 aa += 8;
642 }
643
644 EIGEN_IF_CONSTEXPR (b_unroll >= pow) {
645 compute<0, um_vecs, idx, uk, fetch_x, ktail>(ao, bo, fetchA_idx, fetchB_idx, b_reg);
646
647 constexpr int b_load_offset = idx + 1 + (b_unroll > 1) * !use_less_b_regs;
648 EIGEN_IF_CONSTEXPR (preload_next_k || b_load_offset < b_unroll) {
649 const Scalar* b_addr = bo + b_unroll * uk + b_load_offset - b_shift;
650 b_load(b_reg, b_addr);
651 }
652 }
653 }
654
655 template <int uk, int pow, int count, int um_vecs, int b_unroll, bool ktail, bool fetch_x, bool preload_next_k,
656 int... indices>
657 EIGEN_ALWAYS_INLINE void innerkernel_1pow_impl(std::integer_sequence<int, indices...>, const Scalar*& aa,
658 const Scalar* const& ao, const Scalar* const& bo, int& fetchA_idx,
659 int& fetchB_idx) {
660 int unused[] = {0,
661 (innerkernel_1pow_one<uk, pow, count + indices, um_vecs, b_unroll, ktail, fetch_x, preload_next_k>(
662 aa, ao, bo, fetchA_idx, fetchB_idx),
663 0)...};
664 EIGEN_UNUSED_VARIABLE(unused);
665 }
666
667 template <int uk, int pow, int count, int um_vecs, int b_unroll, bool ktail, bool fetch_x, bool c_fetch,
668 bool preload_next_k = true>
669 EIGEN_ALWAYS_INLINE void innerkernel_1pow(const Scalar*& aa, const Scalar* const& ao, const Scalar* const& bo,
670 Scalar*& co2, int& fetchA_idx, int& fetchB_idx) {
671 constexpr int max_count = (pow + 1) / 2;
672 static_assert(count <= max_count, "invalid B load range");
673 innerkernel_1pow_impl<uk, pow, count, um_vecs, b_unroll, ktail, fetch_x, preload_next_k>(
674 std::make_integer_sequence<int, max_count - count>{}, aa, ao, bo, fetchA_idx, fetchB_idx);
675
676 // Maybe prefetch C data after count-loop.
677 EIGEN_IF_CONSTEXPR (pow == 2 && c_fetch) {
678 EIGEN_IF_CONSTEXPR (uk % 3 == 0 && uk > 0) {
679 co2 += ldc;
680 } else {
681 prefetch_c(co2 + (uk % 3) * nelems_in_cache_line);
682 }
683 }
684 }
685
686 template <int uk, int max_b_unroll, int a_unroll, int b_unroll, bool ktail, bool fetch_x, bool c_fetch,
687 bool preload_next_k = true>
688 EIGEN_ALWAYS_INLINE void innerkernel_1uk(const Scalar*& aa, const Scalar* const& ao, const Scalar* const& bo,
689 Scalar*& co2, int& fetchA_idx, int& fetchB_idx) {
690 const int um_vecs = numext::div_ceil(a_unroll, nelems_in_cache_line);
691
692 EIGEN_IF_CONSTEXPR (max_b_unroll >= 1)
693 innerkernel_1pow<uk, 1, 0, um_vecs, b_unroll, ktail, fetch_x, c_fetch, preload_next_k>(aa, ao, bo, co2,
694 fetchA_idx, fetchB_idx);
695 EIGEN_IF_CONSTEXPR (max_b_unroll >= 2)
696 innerkernel_1pow<uk, 2, 0, um_vecs, b_unroll, ktail, fetch_x, c_fetch, preload_next_k>(aa, ao, bo, co2,
697 fetchA_idx, fetchB_idx);
698 EIGEN_IF_CONSTEXPR (max_b_unroll >= 4)
699 innerkernel_1pow<uk, 4, 0, um_vecs, b_unroll, ktail, fetch_x, c_fetch, preload_next_k>(aa, ao, bo, co2,
700 fetchA_idx, fetchB_idx);
701 EIGEN_IF_CONSTEXPR (max_b_unroll >= 8)
702 innerkernel_1pow<uk, 8, 0, um_vecs, b_unroll, ktail, fetch_x, c_fetch, preload_next_k>(aa, ao, bo, co2,
703 fetchA_idx, fetchB_idx);
704
705 // Load A after pow-loop. Skip this at the end to prevent running over the buffer
706 if (preload_next_k) load_a<0, um_vecs, uk, a_unroll, ktail>(ao);
707 }
708
709 /* Inner kernel loop structure.
710 * for (int uk = 0; uk < kfactor; uk++) {
711 * int idx = 0;
712 *
713 * for (pow = 1; pow < max_b_unroll << 1; pow <<= 1) {
714 * for (int count = 0; count < (pow + 1) / 2; count++) {
715 * auto &b_reg = zmm[b_regs[idx % 2]];
716 *
717 * if (fetch_x && uk == 3 && idx == 0) prefetch_x(aa);
718 * if (fetch_x && uk == 3 && idx == 4) aa += 8;
719 *
720 * if (b_unroll >= pow) {
721 * compute<0, um_vecs, idx, uk, fetchx, ktail>(ao, bo, fetchA_idx, fetchB_idx, b_reg);
722 *
723 * constexpr int b_load_offset = idx + 1 + (b_unroll > 1) * !use_less_b_regs;
724 * if (preload_next_k || b_load_offset < b_unroll)
725 * b_load(b_reg, bo + b_unroll * uk + b_load_offset - b_shift);
726 * }
727 * idx++;
728 * }
729 *
730 * Maybe prefetch C data.
731 * if (pow == 2 && c_fetch) {
732 * if (uk % 3 == 0 && uk > 0) {
733 * co2 += ldc;
734 * } else {
735 * prefetch_c(co2 + (uk % 3) * nelems_in_cache_line);
736 * }
737 * }
738 * }
739 *
740 * Load A.
741 * if (preload_next_k) load_a<0, um_vecs, uk, a_unroll, ktail>(ao);
742 * }
743 *
744 * Advance A/B pointers after uk-loop.
745 * ao += a_unroll * kfactor;
746 * bo += b_unroll * kfactor;
747 */
748
749 template <int a_unroll, int b_unroll, int k_factor, int max_b_unroll, int max_k_factor, bool c_fetch,
750 bool preload_next_k = true>
751 EIGEN_ALWAYS_INLINE void innerkernel(const Scalar*& aa, const Scalar*& ao, const Scalar*& bo, Scalar*& co2) {
752 int fetchA_idx = 0;
753 int fetchB_idx = 0;
754
755 const bool fetch_x = k_factor == max_k_factor;
756 const bool ktail = k_factor == 1;
757
758 static_assert(k_factor <= 4 && k_factor > 0, "innerkernel maximum k_factor supported is 4");
759 static_assert(preload_next_k || k_factor == 1, "skipping next-k preload only allowed when k unroll is 1");
760
761 if (k_factor > 0)
762 innerkernel_1uk<0, max_b_unroll, a_unroll, b_unroll, ktail, fetch_x, c_fetch, preload_next_k>(
763 aa, ao, bo, co2, fetchA_idx, fetchB_idx);
764 if (k_factor > 1)
765 innerkernel_1uk<1, max_b_unroll, a_unroll, b_unroll, ktail, fetch_x, c_fetch, preload_next_k>(
766 aa, ao, bo, co2, fetchA_idx, fetchB_idx);
767 if (k_factor > 2)
768 innerkernel_1uk<2, max_b_unroll, a_unroll, b_unroll, ktail, fetch_x, c_fetch, preload_next_k>(
769 aa, ao, bo, co2, fetchA_idx, fetchB_idx);
770 if (k_factor > 3)
771 innerkernel_1uk<3, max_b_unroll, a_unroll, b_unroll, ktail, fetch_x, c_fetch, preload_next_k>(
772 aa, ao, bo, co2, fetchA_idx, fetchB_idx);
773
774 // Advance A/B pointers after uk-loop.
775 ao += a_unroll * k_factor;
776 bo += b_unroll * k_factor;
777 }
778
779 template <int a_unroll, int b_unroll, int max_b_unroll>
780 EIGEN_ALWAYS_INLINE void kloop(const Scalar*& aa, const Scalar*& ao, const Scalar*& bo, Scalar*& co1, Scalar*& co2) {
781 const int um_vecs = numext::div_ceil(a_unroll, nelems_in_cache_line);
782 EIGEN_IF_CONSTEXPR (!use_less_a_regs) {
783 if (k > 1)
784 a_loads<0, 2, 0, um_vecs, a_unroll>(ao);
785 else
786 a_loads<0, 1, 0, um_vecs, a_unroll>(ao);
787 } else {
788 a_loads<0, 1, 0, um_vecs, a_unroll>(ao);
789 }
790
791 b_load(zmm[b_regs[0]], bo - b_shift + 0);
792 EIGEN_IF_CONSTEXPR (b_unroll > 1 && !use_less_b_regs) {
793 b_load(zmm[b_regs[1]], bo - b_shift + 1);
794 }
795
796#ifndef SECOND_FETCH
797 prefetch_cs<0, max_b_unroll, 0, um_vecs, a_unroll, b_unroll>(co1, co2);
798#endif // SECOND_FETCH
799
800 // Unrolling k-loop by a factor of 4.
801 const int max_k_factor = 4;
802 Index kRem = k % max_k_factor;
803 Index k_ = k - kRem;
804 if (k_ >= max_k_factor) {
805 k_ -= max_k_factor;
806 kRem += max_k_factor;
807 }
808 Index loop_count = k_ / max_k_factor;
809
810 if (loop_count > 0) {
811#ifdef SECOND_FETCH
812 loop_count -= SECOND_FETCH;
813#endif
814 while (loop_count > 0) {
815 innerkernel<a_unroll, b_unroll, max_k_factor, max_b_unroll, max_k_factor, 0>(aa, ao, bo, co2);
816 loop_count--;
817 }
818#ifdef SECOND_FETCH
819 co2 = co1 + nelems_in_cache_line - 1;
820
821 loop_count += b_unroll;
822 while (loop_count > 0) {
823 innerkernel<a_unroll, b_unroll, max_k_factor, max_b_unroll, max_k_factor, 1>(aa, ao, bo, co2);
824 loop_count--;
825 }
826
827 loop_count += SECOND_FETCH - b_unroll;
828 while (loop_count > 0) {
829 innerkernel<a_unroll, b_unroll, max_k_factor, max_b_unroll, max_k_factor, 0>(aa, ao, bo, co2);
830 loop_count--;
831 }
832#endif
833 }
834
835 // k-loop remainder handling.
836 loop_count = kRem;
837 while (loop_count > 1) {
838 innerkernel<a_unroll, b_unroll, 1, max_b_unroll, max_k_factor, 0>(aa, ao, bo, co2);
839 loop_count--;
840 }
841 if (loop_count > 0) {
842 innerkernel<a_unroll, b_unroll, 1, max_b_unroll, max_k_factor, 0, false>(aa, ao, bo, co2);
843 }
844
845 // Update C matrix.
846 c_update<max_b_unroll, a_unroll, b_unroll>(co1, co2);
847 }
848
849 template <int a_unroll, int b_unroll, int max_b_unroll>
850 EIGEN_ALWAYS_INLINE void nloop(const Scalar*& aa, const Scalar*& ao, const Scalar*& bo, Scalar*& co1, Scalar*& co2) {
851 // Set A matrix pointer.
852 ao = a + a_off * a_unroll;
853
854 // Set B matrix pointer if needed.
855 bo += b_unroll * b_off;
856
857 kloop<a_unroll, b_unroll, max_b_unroll>(aa, ao, bo, co1, co2);
858
859 // Advance B matrix pointer if needed.
860 bo += b_unroll * (b_stride - k - b_off);
861
862 // Advance prefetch A pointer.
863 aa += 16;
864 }
865
866 template <int a_unroll, int max_a_unroll, int max_b_unroll>
867 EIGEN_ALWAYS_INLINE void mloop(const Scalar*& ao, const Scalar*& bo, Scalar*& co1, Scalar*& co2) {
868 // Set prefetch A pointers.
869 const Scalar* aa = a + a_unroll * a_stride;
870
871 // Set C matrix pointers.
872 co1 = c;
873 if (a_unroll >= max_a_unroll) co2 = c + 2 * ldc;
874 EIGEN_IF_CONSTEXPR (is_unit_inc)
875 c += a_unroll;
876 else
877 c += a_unroll * inc;
878
879 // Set B matrix pointer.
880 bo = b;
881
882 // Main n-loop.
883 for (Index i = n / max_b_unroll; i > 0; i--) nloop<a_unroll, max_b_unroll, max_b_unroll>(aa, ao, bo, co1, co2);
884
885 // n-remainders.
886 if (n & 4 && max_b_unroll > 4) nloop<a_unroll, 4, max_b_unroll>(aa, ao, bo, co1, co2);
887 // Copy kernels don't support tails of n = 2 for single/double precision.
888 // Loop over ones.
889 int n_rem = 2 * ((n & 2) != 0) + 1 * ((n & 1) != 0);
890 while (n_rem > 0) {
891 nloop<a_unroll, 1, max_b_unroll>(aa, ao, bo, co1, co2);
892 n_rem--;
893 }
894
895 // Advance A matrix pointer.
896 a = ao + a_unroll * (a_stride - k - a_off);
897 }
898
899 public:
900 // Compute kernel unrolling C matrix by max_a_unroll x max_b_unroll.
901 template <int max_a_unroll, int max_b_unroll>
902 EIGEN_ALWAYS_INLINE void compute_kern() {
903 a -= -a_shift;
904 b -= -b_shift;
905
906 const Scalar* ao = nullptr;
907 const Scalar* bo = nullptr;
908 Scalar* co1 = nullptr;
909 Scalar* co2 = nullptr;
910
911 // Main m-loop.
912 for (; m >= max_a_unroll; m -= max_a_unroll) mloop<max_a_unroll, max_a_unroll, max_b_unroll>(ao, bo, co1, co2);
913
914 // m-remainders.
915 EIGEN_IF_CONSTEXPR (max_a_unroll > 32 && is_f32) {
916 constexpr int a_unroll32 = is_f32 ? 32 : 24;
917 if (m & 32) mloop<a_unroll32, max_a_unroll, max_b_unroll>(ao, bo, co1, co2);
918 }
919 EIGEN_IF_CONSTEXPR (max_a_unroll > 16) {
920 if (m & 16) mloop<16, max_a_unroll, max_b_unroll>(ao, bo, co1, co2);
921 }
922 EIGEN_IF_CONSTEXPR (max_a_unroll > 8) {
923 if (m & 8) mloop<8, max_a_unroll, max_b_unroll>(ao, bo, co1, co2);
924 }
925 EIGEN_IF_CONSTEXPR (max_a_unroll > 4) {
926 if (m & 4) mloop<4, max_a_unroll, max_b_unroll>(ao, bo, co1, co2);
927 }
928 EIGEN_IF_CONSTEXPR (max_a_unroll > 2 && is_f64) {
929 if (m & 2) mloop<2, max_a_unroll, max_b_unroll>(ao, bo, co1, co2);
930 }
931 EIGEN_IF_CONSTEXPR (max_a_unroll > 1 && is_f64) {
932 if (m & 1) mloop<1, max_a_unroll, max_b_unroll>(ao, bo, co1, co2);
933 }
934
935 // Copy kernels don't support tails of m = 2 for single precision.
936 // Loop over ones.
937 EIGEN_IF_CONSTEXPR (is_f32) {
938 int m_rem = 2 * ((m & 2) != 0) + 1 * ((m & 1) != 0);
939 while (m_rem > 0) {
940 mloop<1, max_a_unroll, max_b_unroll>(ao, bo, co1, co2);
941 m_rem--;
942 }
943 }
944 }
945
946 gemm_class(Index m_, Index n_, Index k_, Index ldc_, Index inc_, const Scalar* alpha_, const Scalar* a_,
947 const Scalar* b_, Scalar* c_, bool is_alpha1_, bool is_beta0_, Index a_stride_, Index b_stride_,
948 Index a_off_, Index b_off_)
949 : m(m_),
950 n(n_),
951 k(k_),
952 ldc(ldc_),
953 inc(inc_),
954 alpha(alpha_),
955 a(a_),
956 b(b_),
957 c(c_),
958 is_alpha1(is_alpha1_),
959 is_beta0(is_beta0_),
960 a_stride(a_stride_),
961 b_stride(b_stride_),
962 a_off(a_off_),
963 b_off(b_off_) {
964 // Zero out all accumulation registers.
965 zmm[8] = pzero(zmm[8]);
966 zmm[9] = pzero(zmm[9]);
967 zmm[10] = pzero(zmm[10]);
968 zmm[11] = pzero(zmm[11]);
969 zmm[12] = pzero(zmm[12]);
970 zmm[13] = pzero(zmm[13]);
971 zmm[14] = pzero(zmm[14]);
972 zmm[15] = pzero(zmm[15]);
973 zmm[16] = pzero(zmm[16]);
974 zmm[17] = pzero(zmm[17]);
975 zmm[18] = pzero(zmm[18]);
976 zmm[19] = pzero(zmm[19]);
977 zmm[20] = pzero(zmm[20]);
978 zmm[21] = pzero(zmm[21]);
979 zmm[22] = pzero(zmm[22]);
980 zmm[23] = pzero(zmm[23]);
981 zmm[24] = pzero(zmm[24]);
982 zmm[25] = pzero(zmm[25]);
983 zmm[26] = pzero(zmm[26]);
984 zmm[27] = pzero(zmm[27]);
985 zmm[28] = pzero(zmm[28]);
986 zmm[29] = pzero(zmm[29]);
987 zmm[30] = pzero(zmm[30]);
988 zmm[31] = pzero(zmm[31]);
989 }
990};
991
992template <typename Scalar, bool is_unit_inc>
993const int gemm_class<Scalar, is_unit_inc>::a_regs[];
994
995template <typename Scalar, bool is_unit_inc>
996const int gemm_class<Scalar, is_unit_inc>::b_regs[];
997
998template <typename Scalar, bool is_unit_inc>
999const int gemm_class<Scalar, is_unit_inc>::c_regs[];
1000
1001// Compute kernel with max unroll support of:
1002// Single precision:
1003// max_a_unroll: 48, 32, 16, 8, 4, 2, 1
1004// max_b_unroll: 8, 4, 2, 1
1005// Double precision:
1006// max_a_unroll: 24, 16, 8, 4, 2, 1
1007// max_b_unroll: 8, 4, 2, 1
1008template <typename Scalar, int max_a_unroll, int max_b_unroll, bool is_alpha1, bool is_beta0, bool is_unit_inc>
1009EIGEN_DONT_INLINE void gemm_kern_avx512(Index m, Index n, Index k, Scalar* alpha, const Scalar* a, const Scalar* b,
1010 Scalar* c, Index ldc, Index inc = 1, Index a_stride = -1, Index b_stride = -1,
1011 Index a_off = 0, Index b_off = 0) {
1012 if (m <= 0 || n <= 0 || k <= 0) return;
1013 if (a_stride == -1) a_stride = k;
1014 if (b_stride == -1) b_stride = k;
1015
1016 gemm_class<Scalar, is_unit_inc> g(m, n, k, ldc, inc, alpha, a, b, c, is_alpha1, is_beta0, a_stride, b_stride, a_off,
1017 b_off);
1018 g.template compute_kern<max_a_unroll, max_b_unroll>();
1019}
1020
1021// Template specializations of GEBP kernels with nr = 8.
1022template <bool ConjLhs_, bool ConjRhs_, int PacketSize_>
1023class gebp_traits<float, float, ConjLhs_, ConjRhs_, Architecture::Target, PacketSize_>
1024 : public gebp_traits<float, float, ConjLhs_, ConjRhs_, Architecture::Generic, PacketSize_> {
1025 using Base = gebp_traits<float, float, ConjLhs_, ConjRhs_, Architecture::Generic, PacketSize_>;
1026
1027 public:
1028 enum { nr = Base::Vectorizable ? 8 : 4 };
1029};
1030
1031template <bool ConjLhs_, bool ConjRhs_, int PacketSize_>
1032class gebp_traits<double, double, ConjLhs_, ConjRhs_, Architecture::Target, PacketSize_>
1033 : public gebp_traits<double, double, ConjLhs_, ConjRhs_, Architecture::Generic, PacketSize_> {
1034 using Base = gebp_traits<double, double, ConjLhs_, ConjRhs_, Architecture::Generic, PacketSize_>;
1035
1036 public:
1037 enum { nr = Base::Vectorizable ? 8 : 4 };
1038};
1039
1040template <typename Scalar, typename Index, typename DataMapper, bool Conjugate, bool PanelMode>
1041struct gemm_pack_rhs<Scalar, Index, DataMapper, 8, ColMajor, Conjugate, PanelMode> {
1042 typedef typename packet_traits<Scalar>::type Packet;
1043 typedef typename DataMapper::LinearMapper LinearMapper;
1044 enum { PacketSize = packet_traits<Scalar>::size };
1045 EIGEN_DONT_INLINE void operator()(Scalar* blockB, const DataMapper& rhs, Index depth, Index cols, Index stride = 0,
1046 Index offset = 0) const;
1047};
1048
1049template <typename Scalar, typename Index, typename DataMapper, bool Conjugate, bool PanelMode>
1050EIGEN_DONT_INLINE void gemm_pack_rhs<Scalar, Index, DataMapper, 8, ColMajor, Conjugate, PanelMode>::operator()(
1051 Scalar* blockB, const DataMapper& rhs, Index depth, Index cols, Index stride, Index offset) const {
1052 constexpr int nr = 8;
1053 EIGEN_ASM_COMMENT("EIGEN PRODUCT PACK RHS COLMAJOR");
1054 EIGEN_UNUSED_VARIABLE(stride);
1055 EIGEN_UNUSED_VARIABLE(offset);
1056 eigen_assert(((!PanelMode) && stride == 0 && offset == 0) || (PanelMode && stride >= depth && offset <= stride));
1057 conj_if<NumTraits<Scalar>::IsComplex && Conjugate> cj;
1058 Index packet_cols8 = nr >= 8 ? (cols / 8) * 8 : 0;
1059 Index packet_cols4 = nr >= 4 ? (cols / 4) * 4 : 0;
1060 Index count = 0;
1061 const Index peeled_k = (depth / PacketSize) * PacketSize;
1062 EIGEN_IF_CONSTEXPR (nr >= 8) {
1063 for (Index j2 = 0; j2 < packet_cols8; j2 += 8) {
1064 // skip what we have before
1065 EIGEN_IF_CONSTEXPR (PanelMode) count += 8 * offset;
1066 const LinearMapper dm0 = rhs.getLinearMapper(0, j2 + 0);
1067 const LinearMapper dm1 = rhs.getLinearMapper(0, j2 + 1);
1068 const LinearMapper dm2 = rhs.getLinearMapper(0, j2 + 2);
1069 const LinearMapper dm3 = rhs.getLinearMapper(0, j2 + 3);
1070 const LinearMapper dm4 = rhs.getLinearMapper(0, j2 + 4);
1071 const LinearMapper dm5 = rhs.getLinearMapper(0, j2 + 5);
1072 const LinearMapper dm6 = rhs.getLinearMapper(0, j2 + 6);
1073 const LinearMapper dm7 = rhs.getLinearMapper(0, j2 + 7);
1074 Index k = 0;
1075 EIGEN_IF_CONSTEXPR ((PacketSize % 8) == 0 || PacketSize == 4) {
1076 for (; k < peeled_k; k += PacketSize) {
1077 PacketBlock<Packet, 8> kernel;
1078
1079 kernel.packet[0] = dm0.template loadPacket<Packet>(k);
1080 kernel.packet[1] = dm1.template loadPacket<Packet>(k);
1081 kernel.packet[2] = dm2.template loadPacket<Packet>(k);
1082 kernel.packet[3] = dm3.template loadPacket<Packet>(k);
1083 kernel.packet[4] = dm4.template loadPacket<Packet>(k);
1084 kernel.packet[5] = dm5.template loadPacket<Packet>(k);
1085 kernel.packet[6] = dm6.template loadPacket<Packet>(k);
1086 kernel.packet[7] = dm7.template loadPacket<Packet>(k);
1087
1088 EIGEN_IF_CONSTEXPR (PacketSize == 4) {
1089 // For PacketSize==4 we cannot ptranspose 8 packets directly; compose two
1090 // 4-packet transposes (cols 0-3 and 4-7) and interleave the halves so
1091 // the 8 stores produce 4 rows of 8 packed elements.
1092 PacketBlock<Packet, 4> tmp_lo;
1093 tmp_lo.packet[0] = kernel.packet[0];
1094 tmp_lo.packet[1] = kernel.packet[1];
1095 tmp_lo.packet[2] = kernel.packet[2];
1096 tmp_lo.packet[3] = kernel.packet[3];
1097 ptranspose(tmp_lo);
1098 PacketBlock<Packet, 4> tmp_hi;
1099 tmp_hi.packet[0] = kernel.packet[4];
1100 tmp_hi.packet[1] = kernel.packet[5];
1101 tmp_hi.packet[2] = kernel.packet[6];
1102 tmp_hi.packet[3] = kernel.packet[7];
1103 ptranspose(tmp_hi);
1104 kernel.packet[0] = tmp_lo.packet[0];
1105 kernel.packet[1] = tmp_hi.packet[0];
1106 kernel.packet[2] = tmp_lo.packet[1];
1107 kernel.packet[3] = tmp_hi.packet[1];
1108 kernel.packet[4] = tmp_lo.packet[2];
1109 kernel.packet[5] = tmp_hi.packet[2];
1110 kernel.packet[6] = tmp_lo.packet[3];
1111 kernel.packet[7] = tmp_hi.packet[3];
1112 } else {
1113 ptranspose(kernel);
1114 }
1115
1116 pstoreu(blockB + count + 0 * PacketSize, cj.pconj(kernel.packet[0]));
1117 pstoreu(blockB + count + 1 * PacketSize, cj.pconj(kernel.packet[1]));
1118 pstoreu(blockB + count + 2 * PacketSize, cj.pconj(kernel.packet[2]));
1119 pstoreu(blockB + count + 3 * PacketSize, cj.pconj(kernel.packet[3]));
1120 pstoreu(blockB + count + 4 * PacketSize, cj.pconj(kernel.packet[4]));
1121 pstoreu(blockB + count + 5 * PacketSize, cj.pconj(kernel.packet[5]));
1122 pstoreu(blockB + count + 6 * PacketSize, cj.pconj(kernel.packet[6]));
1123 pstoreu(blockB + count + 7 * PacketSize, cj.pconj(kernel.packet[7]));
1124 count += 8 * PacketSize;
1125 }
1126 }
1127 for (; k < depth; k++) {
1128 blockB[count + 0] = cj(dm0(k));
1129 blockB[count + 1] = cj(dm1(k));
1130 blockB[count + 2] = cj(dm2(k));
1131 blockB[count + 3] = cj(dm3(k));
1132 blockB[count + 4] = cj(dm4(k));
1133 blockB[count + 5] = cj(dm5(k));
1134 blockB[count + 6] = cj(dm6(k));
1135 blockB[count + 7] = cj(dm7(k));
1136 count += 8;
1137 }
1138 // skip what we have after
1139 EIGEN_IF_CONSTEXPR (PanelMode) count += 8 * (stride - offset - depth);
1140 }
1141 }
1142
1143 EIGEN_IF_CONSTEXPR (nr >= 4) {
1144 for (Index j2 = packet_cols8; j2 < packet_cols4; j2 += 4) {
1145 // skip what we have before
1146 EIGEN_IF_CONSTEXPR (PanelMode) count += 4 * offset;
1147 const LinearMapper dm0 = rhs.getLinearMapper(0, j2 + 0);
1148 const LinearMapper dm1 = rhs.getLinearMapper(0, j2 + 1);
1149 const LinearMapper dm2 = rhs.getLinearMapper(0, j2 + 2);
1150 const LinearMapper dm3 = rhs.getLinearMapper(0, j2 + 3);
1151
1152 Index k = 0;
1153 EIGEN_IF_CONSTEXPR ((PacketSize % 4) == 0 || PacketSize == 2) {
1154 for (; k < peeled_k; k += PacketSize) {
1155 PacketBlock<Packet, 4> kernel;
1156 kernel.packet[0] = dm0.template loadPacket<Packet>(k);
1157 kernel.packet[1] = dm1.template loadPacket<Packet>(k);
1158 kernel.packet[2] = dm2.template loadPacket<Packet>(k);
1159 kernel.packet[3] = dm3.template loadPacket<Packet>(k);
1160 EIGEN_IF_CONSTEXPR (PacketSize == 2) {
1161 // See the matching note in GeneralBlockPanelKernel.h.
1162 PacketBlock<Packet, 2> tmp01;
1163 tmp01.packet[0] = kernel.packet[0];
1164 tmp01.packet[1] = kernel.packet[1];
1165 ptranspose(tmp01);
1166 PacketBlock<Packet, 2> tmp23;
1167 tmp23.packet[0] = kernel.packet[2];
1168 tmp23.packet[1] = kernel.packet[3];
1169 ptranspose(tmp23);
1170 kernel.packet[0] = tmp01.packet[0];
1171 kernel.packet[1] = tmp23.packet[0];
1172 kernel.packet[2] = tmp01.packet[1];
1173 kernel.packet[3] = tmp23.packet[1];
1174 } else {
1175 ptranspose(kernel);
1176 }
1177 pstoreu(blockB + count + 0 * PacketSize, cj.pconj(kernel.packet[0]));
1178 pstoreu(blockB + count + 1 * PacketSize, cj.pconj(kernel.packet[1]));
1179 pstoreu(blockB + count + 2 * PacketSize, cj.pconj(kernel.packet[2]));
1180 pstoreu(blockB + count + 3 * PacketSize, cj.pconj(kernel.packet[3]));
1181 count += 4 * PacketSize;
1182 }
1183 }
1184 for (; k < depth; k++) {
1185 blockB[count + 0] = cj(dm0(k));
1186 blockB[count + 1] = cj(dm1(k));
1187 blockB[count + 2] = cj(dm2(k));
1188 blockB[count + 3] = cj(dm3(k));
1189 count += 4;
1190 }
1191 // skip what we have after
1192 EIGEN_IF_CONSTEXPR (PanelMode) count += 4 * (stride - offset - depth);
1193 }
1194 }
1195
1196 // copy the remaining columns one at a time (nr==1)
1197 for (Index j2 = packet_cols4; j2 < cols; ++j2) {
1198 EIGEN_IF_CONSTEXPR (PanelMode) count += offset;
1199 const LinearMapper dm0 = rhs.getLinearMapper(0, j2);
1200 for (Index k = 0; k < depth; k++) {
1201 blockB[count] = cj(dm0(k));
1202 count += 1;
1203 }
1204 EIGEN_IF_CONSTEXPR (PanelMode) count += (stride - offset - depth);
1205 }
1206}
1207
1208template <typename Scalar, typename Index, typename DataMapper, bool Conjugate, bool PanelMode>
1209struct gemm_pack_rhs<Scalar, Index, DataMapper, 8, RowMajor, Conjugate, PanelMode> {
1210 typedef typename packet_traits<Scalar>::type Packet;
1211 typedef typename unpacket_traits<Packet>::half HalfPacket;
1212 typedef typename unpacket_traits<typename unpacket_traits<Packet>::half>::half QuarterPacket;
1213 typedef typename DataMapper::LinearMapper LinearMapper;
1214 enum {
1215 PacketSize = packet_traits<Scalar>::size,
1216 HalfPacketSize = unpacket_traits<HalfPacket>::size,
1217 QuarterPacketSize = unpacket_traits<QuarterPacket>::size
1218 };
1219 EIGEN_DONT_INLINE void operator()(Scalar* blockB, const DataMapper& rhs, Index depth, Index cols, Index stride = 0,
1220 Index offset = 0) const {
1221 constexpr int nr = 8;
1222 EIGEN_ASM_COMMENT("EIGEN PRODUCT PACK RHS ROWMAJOR");
1223 EIGEN_UNUSED_VARIABLE(stride);
1224 EIGEN_UNUSED_VARIABLE(offset);
1225 eigen_assert(((!PanelMode) && stride == 0 && offset == 0) || (PanelMode && stride >= depth && offset <= stride));
1226 constexpr bool HasHalf = (int)HalfPacketSize < (int)PacketSize;
1227 constexpr bool HasQuarter = (int)QuarterPacketSize < (int)HalfPacketSize;
1228 conj_if<NumTraits<Scalar>::IsComplex && Conjugate> cj;
1229 Index packet_cols8 = nr >= 8 ? (cols / 8) * 8 : 0;
1230 Index packet_cols4 = nr >= 4 ? (cols / 4) * 4 : 0;
1231 Index count = 0;
1232
1233 EIGEN_IF_CONSTEXPR (nr >= 8) {
1234 for (Index j2 = 0; j2 < packet_cols8; j2 += 8) {
1235 // skip what we have before
1236 EIGEN_IF_CONSTEXPR (PanelMode) count += 8 * offset;
1237 for (Index k = 0; k < depth; k++) {
1238 EIGEN_IF_CONSTEXPR (PacketSize == 8) {
1239 // Packet A = ploadu<Packet>(&rhs.data()[k*rhs.stride() + j2]);
1240 Packet A = rhs.template loadPacket<Packet>(k, j2);
1241 pstoreu(blockB + count, cj.pconj(A));
1242 } else EIGEN_IF_CONSTEXPR (HasHalf && HalfPacketSize == 8) {
1243 HalfPacket A = rhs.template loadPacket<HalfPacket>(k, j2);
1244 pstoreu(blockB + count, cj.pconj(A));
1245 } else EIGEN_IF_CONSTEXPR (HasQuarter && QuarterPacketSize == 8) {
1246 QuarterPacket A = rhs.template loadPacket<QuarterPacket>(k, j2);
1247 pstoreu(blockB + count, cj.pconj(A));
1248 } else EIGEN_IF_CONSTEXPR (PacketSize == 4) {
1249 // Packet A = ploadu<Packet>(&rhs.data()[k*rhs.stride() + j2]);
1250 // Packet B = ploadu<Packet>(&rhs.data()[k*rhs.stride() + j2 + PacketSize]);
1251 Packet A = rhs.template loadPacket<Packet>(k, j2);
1252 Packet B = rhs.template loadPacket<Packet>(k, j2 + PacketSize);
1253 pstoreu(blockB + count, cj.pconj(A));
1254 pstoreu(blockB + count + PacketSize, cj.pconj(B));
1255 } else {
1256 // const Scalar* b0 = &rhs.data()[k*rhs.stride() + j2];
1257 const LinearMapper dm0 = rhs.getLinearMapper(k, j2);
1258 blockB[count + 0] = cj(dm0(0));
1259 blockB[count + 1] = cj(dm0(1));
1260 blockB[count + 2] = cj(dm0(2));
1261 blockB[count + 3] = cj(dm0(3));
1262 blockB[count + 4] = cj(dm0(4));
1263 blockB[count + 5] = cj(dm0(5));
1264 blockB[count + 6] = cj(dm0(6));
1265 blockB[count + 7] = cj(dm0(7));
1266 }
1267 count += 8;
1268 }
1269 // skip what we have after
1270 EIGEN_IF_CONSTEXPR (PanelMode) count += 8 * (stride - offset - depth);
1271 }
1272 }
1273
1274 EIGEN_IF_CONSTEXPR (nr >= 4) {
1275 for (Index j2 = packet_cols8; j2 < packet_cols4; j2 += 4) {
1276 // skip what we have before
1277 EIGEN_IF_CONSTEXPR (PanelMode) count += 4 * offset;
1278 for (Index k = 0; k < depth; k++) {
1279 EIGEN_IF_CONSTEXPR (PacketSize == 4) {
1280 Packet A = rhs.template loadPacket<Packet>(k, j2);
1281 pstoreu(blockB + count, cj.pconj(A));
1282 count += PacketSize;
1283 } else EIGEN_IF_CONSTEXPR (HasHalf && HalfPacketSize == 4) {
1284 HalfPacket A = rhs.template loadPacket<HalfPacket>(k, j2);
1285 pstoreu(blockB + count, cj.pconj(A));
1286 count += HalfPacketSize;
1287 } else EIGEN_IF_CONSTEXPR (HasQuarter && QuarterPacketSize == 4) {
1288 QuarterPacket A = rhs.template loadPacket<QuarterPacket>(k, j2);
1289 pstoreu(blockB + count, cj.pconj(A));
1290 count += QuarterPacketSize;
1291 } else {
1292 const LinearMapper dm0 = rhs.getLinearMapper(k, j2);
1293 blockB[count + 0] = cj(dm0(0));
1294 blockB[count + 1] = cj(dm0(1));
1295 blockB[count + 2] = cj(dm0(2));
1296 blockB[count + 3] = cj(dm0(3));
1297 count += 4;
1298 }
1299 }
1300 // skip what we have after
1301 EIGEN_IF_CONSTEXPR (PanelMode) count += 4 * (stride - offset - depth);
1302 }
1303 }
1304 // copy the remaining columns one at a time (nr==1)
1305 for (Index j2 = packet_cols4; j2 < cols; ++j2) {
1306 EIGEN_IF_CONSTEXPR (PanelMode) count += offset;
1307 for (Index k = 0; k < depth; k++) {
1308 blockB[count] = cj(rhs(k, j2));
1309 count += 1;
1310 }
1311 EIGEN_IF_CONSTEXPR (PanelMode) count += stride - offset - depth;
1312 }
1313 }
1314};
1315
1316template <typename Scalar, typename Index, typename DataMapper, int mr, bool ConjugateLhs, bool ConjugateRhs>
1317struct gebp_kernel<Scalar, Scalar, Index, DataMapper, mr, 8, ConjugateLhs, ConjugateRhs> {
1318 EIGEN_ALWAYS_INLINE void operator()(const DataMapper& res, const Scalar* blockA, const Scalar* blockB, Index rows,
1319 Index depth, Index cols, Scalar alpha, Index strideA = -1, Index strideB = -1,
1320 Index offsetA = 0, Index offsetB = 0) const;
1321};
1322
1323template <typename Scalar, typename Index, typename DataMapper, int mr, bool ConjugateLhs, bool ConjugateRhs>
1324EIGEN_ALWAYS_INLINE void gebp_kernel<Scalar, Scalar, Index, DataMapper, mr, 8, ConjugateLhs, ConjugateRhs>::operator()(
1325 const DataMapper& res, const Scalar* blockA, const Scalar* blockB, Index rows, Index depth, Index cols,
1326 Scalar alpha, Index strideA, Index strideB, Index offsetA, Index offsetB) const {
1327 if (res.incr() == 1) {
1328 if (alpha == 1) {
1329 gemm_kern_avx512<Scalar, mr, 8, true, false, true>(rows, cols, depth, &alpha, blockA, blockB, (Scalar*)res.data(),
1330 res.stride(), res.incr(), strideA, strideB, offsetA, offsetB);
1331 } else {
1332 gemm_kern_avx512<Scalar, mr, 8, false, false, true>(rows, cols, depth, &alpha, blockA, blockB,
1333 (Scalar*)res.data(), res.stride(), res.incr(), strideA,
1334 strideB, offsetA, offsetB);
1335 }
1336 } else {
1337 if (alpha == 1) {
1338 gemm_kern_avx512<Scalar, mr, 8, true, false, false>(rows, cols, depth, &alpha, blockA, blockB,
1339 (Scalar*)res.data(), res.stride(), res.incr(), strideA,
1340 strideB, offsetA, offsetB);
1341 } else {
1342 gemm_kern_avx512<Scalar, mr, 8, false, false, false>(rows, cols, depth, &alpha, blockA, blockB,
1343 (Scalar*)res.data(), res.stride(), res.incr(), strideA,
1344 strideB, offsetA, offsetB);
1345 }
1346 }
1347}
1348#endif // EIGEN_USE_AVX512_GEMM_KERNELS
1349
1350} // namespace internal
1351} // namespace Eigen
1352
1353#undef SECOND_FETCH
1354
1355#endif // EIGEN_CORE_ARCH_AVX512_GEMM_KERNEL_H
@ ColMajor
Definition Constants.h:319
@ RowMajor
Definition Constants.h:321