Eigen  5.0.1
 
Loading...
Searching...
No Matches
TrsmKernel.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_TRSM_KERNEL_H
12#define EIGEN_CORE_ARCH_AVX512_TRSM_KERNEL_H
13
14// IWYU pragma: private
15#include "../../InternalHeaderCheck.h"
16#include <utility>
17
18#if !defined(EIGEN_USE_AVX512_TRSM_KERNELS)
19#define EIGEN_USE_AVX512_TRSM_KERNELS 1
20#endif
21
22// TRSM kernels currently unconditionally rely on malloc with AVX512.
23// Disable them if malloc is explicitly disabled at compile-time.
24#ifdef EIGEN_NO_MALLOC
25#undef EIGEN_USE_AVX512_TRSM_KERNELS
26#define EIGEN_USE_AVX512_TRSM_KERNELS 0
27#endif
28
29#if EIGEN_USE_AVX512_TRSM_KERNELS
30#if !defined(EIGEN_USE_AVX512_TRSM_R_KERNELS)
31#define EIGEN_USE_AVX512_TRSM_R_KERNELS 1
32#endif
33#if !defined(EIGEN_USE_AVX512_TRSM_L_KERNELS)
34#define EIGEN_USE_AVX512_TRSM_L_KERNELS 1
35#endif
36#else // EIGEN_USE_AVX512_TRSM_KERNELS == 0
37#define EIGEN_USE_AVX512_TRSM_R_KERNELS 0
38#define EIGEN_USE_AVX512_TRSM_L_KERNELS 0
39#endif
40
41// Need this for some std::min calls.
42#ifdef min
43#undef min
44#endif
45
46namespace Eigen {
47namespace internal {
48
49#if (EIGEN_USE_AVX512_TRSM_KERNELS)
50
51#define EIGEN_AVX_MAX_NUM_ACC (int64_t(24))
52#define EIGEN_AVX_MAX_NUM_ROW (int64_t(8)) // Denoted L in code.
53#define EIGEN_AVX_MAX_K_UNROL (int64_t(4))
54#define EIGEN_AVX_B_LOAD_SETS (int64_t(2))
55#define EIGEN_AVX_MAX_A_BCAST (int64_t(2))
56typedef Packet16f vecFullFloat;
57typedef Packet8d vecFullDouble;
58typedef Packet8f vecHalfFloat;
59typedef Packet4d vecHalfDouble;
60
61// Compile-time unrolls are implemented here.
62// Note: this depends on macros and typedefs above.
63#include "TrsmUnrolls.inc"
64
65#if (EIGEN_COMP_CLANG != 0)
66
83#if !defined(EIGEN_ENABLE_AVX512_NOCOPY_TRSM_CUTOFFS)
84#define EIGEN_ENABLE_AVX512_NOCOPY_TRSM_CUTOFFS 1
85#endif
86
87#if EIGEN_ENABLE_AVX512_NOCOPY_TRSM_CUTOFFS
88
89#if EIGEN_USE_AVX512_TRSM_R_KERNELS
90#if !defined(EIGEN_ENABLE_AVX512_NOCOPY_TRSM_R_CUTOFFS)
91#define EIGEN_ENABLE_AVX512_NOCOPY_TRSM_R_CUTOFFS 1
92#endif // !defined(EIGEN_ENABLE_AVX512_NOCOPY_TRSM_R_CUTOFFS)
93#endif
94
95#if EIGEN_USE_AVX512_TRSM_L_KERNELS
96#if !defined(EIGEN_ENABLE_AVX512_NOCOPY_TRSM_L_CUTOFFS)
97#define EIGEN_ENABLE_AVX512_NOCOPY_TRSM_L_CUTOFFS 1
98#endif
99#endif // EIGEN_USE_AVX512_TRSM_L_KERNELS
100
101#else // EIGEN_ENABLE_AVX512_NOCOPY_TRSM_CUTOFFS == 0
102#define EIGEN_ENABLE_AVX512_NOCOPY_TRSM_R_CUTOFFS 0
103#define EIGEN_ENABLE_AVX512_NOCOPY_TRSM_L_CUTOFFS 0
104#endif // EIGEN_ENABLE_AVX512_NOCOPY_TRSM_CUTOFFS
105
106template <typename Scalar>
107int64_t avx512_trsm_cutoff(int64_t L2Size, int64_t N, double L2Cap) {
108 const int64_t U3 = 3 * packet_traits<Scalar>::size;
109 const int64_t MaxNb = 5 * U3;
110 int64_t Nb = std::min(MaxNb, N);
111 double cutoff_d =
112 (((L2Size * L2Cap) / (sizeof(Scalar))) - (EIGEN_AVX_MAX_NUM_ROW)*Nb) / ((EIGEN_AVX_MAX_NUM_ROW) + Nb);
113 int64_t cutoff_l = static_cast<int64_t>(cutoff_d);
114 return (cutoff_l / EIGEN_AVX_MAX_NUM_ROW) * EIGEN_AVX_MAX_NUM_ROW;
115}
116#else // !(EIGEN_USE_AVX512_TRSM_KERNELS) || !(EIGEN_COMP_CLANG != 0)
117#define EIGEN_ENABLE_AVX512_NOCOPY_TRSM_CUTOFFS 0
118#define EIGEN_ENABLE_AVX512_NOCOPY_TRSM_R_CUTOFFS 0
119#define EIGEN_ENABLE_AVX512_NOCOPY_TRSM_L_CUTOFFS 0
120#endif
121
125template <typename Scalar, typename vec, int64_t unrollM, int64_t unrollN, bool remM, bool remN>
126EIGEN_ALWAYS_INLINE void transStoreC(PacketBlock<vec, EIGEN_ARCH_DEFAULT_NUMBER_OF_REGISTERS>& zmm, Scalar* C_arr,
127 int64_t LDC, int64_t remM_ = 0, int64_t remN_ = 0) {
128 EIGEN_UNUSED_VARIABLE(remN_);
129 EIGEN_UNUSED_VARIABLE(remM_);
130 using urolls = unrolls::trans<Scalar>;
131
132 constexpr int64_t U3 = urolls::PacketSize * 3;
133 constexpr int64_t U2 = urolls::PacketSize * 2;
134 constexpr int64_t U1 = urolls::PacketSize * 1;
135
136 static_assert(unrollN == U1 || unrollN == U2 || unrollN == U3, "unrollN should be a multiple of PacketSize");
137 static_assert(unrollM == EIGEN_AVX_MAX_NUM_ROW, "unrollM should be equal to EIGEN_AVX_MAX_NUM_ROW");
138
139 urolls::template transpose<unrollN, 0>(zmm);
140 EIGEN_IF_CONSTEXPR (unrollN > U2) urolls::template transpose<unrollN, 2>(zmm);
141 EIGEN_IF_CONSTEXPR (unrollN > U1) urolls::template transpose<unrollN, 1>(zmm);
142
143 static_assert((remN && unrollN == U1) || !remN, "When handling N remainder set unrollN=U1");
144 EIGEN_IF_CONSTEXPR (!remN) {
145 urolls::template storeC<std::min(unrollN, U1), unrollN, 0, remM>(C_arr, LDC, zmm, remM_);
146 EIGEN_IF_CONSTEXPR (unrollN > U1) {
147 constexpr int64_t unrollN_ = std::min(unrollN - U1, U1);
148 urolls::template storeC<unrollN_, unrollN, 1, remM>(C_arr + U1 * LDC, LDC, zmm, remM_);
149 }
150 EIGEN_IF_CONSTEXPR (unrollN > U2) {
151 constexpr int64_t unrollN_ = std::min(unrollN - U2, U1);
152 urolls::template storeC<unrollN_, unrollN, 2, remM>(C_arr + U2 * LDC, LDC, zmm, remM_);
153 }
154 } else {
155 EIGEN_IF_CONSTEXPR ((std::is_same<Scalar, float>::value)) {
156 // Note: without "if constexpr" this section of code will also be
157 // parsed by the compiler so each of the storeC will still be instantiated.
158 // We use enable_if in aux_storeC to set it to an empty function for
159 // these cases.
160 if (remN_ == 15)
161 urolls::template storeC<15, unrollN, 0, remM>(C_arr, LDC, zmm, remM_);
162 else if (remN_ == 14)
163 urolls::template storeC<14, unrollN, 0, remM>(C_arr, LDC, zmm, remM_);
164 else if (remN_ == 13)
165 urolls::template storeC<13, unrollN, 0, remM>(C_arr, LDC, zmm, remM_);
166 else if (remN_ == 12)
167 urolls::template storeC<12, unrollN, 0, remM>(C_arr, LDC, zmm, remM_);
168 else if (remN_ == 11)
169 urolls::template storeC<11, unrollN, 0, remM>(C_arr, LDC, zmm, remM_);
170 else if (remN_ == 10)
171 urolls::template storeC<10, unrollN, 0, remM>(C_arr, LDC, zmm, remM_);
172 else if (remN_ == 9)
173 urolls::template storeC<9, unrollN, 0, remM>(C_arr, LDC, zmm, remM_);
174 else if (remN_ == 8)
175 urolls::template storeC<8, unrollN, 0, remM>(C_arr, LDC, zmm, remM_);
176 else if (remN_ == 7)
177 urolls::template storeC<7, unrollN, 0, remM>(C_arr, LDC, zmm, remM_);
178 else if (remN_ == 6)
179 urolls::template storeC<6, unrollN, 0, remM>(C_arr, LDC, zmm, remM_);
180 else if (remN_ == 5)
181 urolls::template storeC<5, unrollN, 0, remM>(C_arr, LDC, zmm, remM_);
182 else if (remN_ == 4)
183 urolls::template storeC<4, unrollN, 0, remM>(C_arr, LDC, zmm, remM_);
184 else if (remN_ == 3)
185 urolls::template storeC<3, unrollN, 0, remM>(C_arr, LDC, zmm, remM_);
186 else if (remN_ == 2)
187 urolls::template storeC<2, unrollN, 0, remM>(C_arr, LDC, zmm, remM_);
188 else if (remN_ == 1)
189 urolls::template storeC<1, unrollN, 0, remM>(C_arr, LDC, zmm, remM_);
190 } else {
191 if (remN_ == 7)
192 urolls::template storeC<7, unrollN, 0, remM>(C_arr, LDC, zmm, remM_);
193 else if (remN_ == 6)
194 urolls::template storeC<6, unrollN, 0, remM>(C_arr, LDC, zmm, remM_);
195 else if (remN_ == 5)
196 urolls::template storeC<5, unrollN, 0, remM>(C_arr, LDC, zmm, remM_);
197 else if (remN_ == 4)
198 urolls::template storeC<4, unrollN, 0, remM>(C_arr, LDC, zmm, remM_);
199 else if (remN_ == 3)
200 urolls::template storeC<3, unrollN, 0, remM>(C_arr, LDC, zmm, remM_);
201 else if (remN_ == 2)
202 urolls::template storeC<2, unrollN, 0, remM>(C_arr, LDC, zmm, remM_);
203 else if (remN_ == 1)
204 urolls::template storeC<1, unrollN, 0, remM>(C_arr, LDC, zmm, remM_);
205 }
206 }
207}
208
223template <typename Scalar, bool isARowMajor, bool isCRowMajor, bool isAdd, bool handleKRem>
224void gemmKernel(Scalar* A_arr, Scalar* B_arr, Scalar* C_arr, int64_t M, int64_t N, int64_t K, int64_t LDA, int64_t LDB,
225 int64_t LDC) {
226 using urolls = unrolls::gemm<Scalar, isAdd>;
227 constexpr int64_t U3 = urolls::PacketSize * 3;
228 constexpr int64_t U2 = urolls::PacketSize * 2;
229 constexpr int64_t U1 = urolls::PacketSize * 1;
230 using vec = std::conditional_t<std::is_same<Scalar, float>::value, vecFullFloat, vecFullDouble>;
231 int64_t N_ = (N / U3) * U3;
232 int64_t M_ = (M / EIGEN_AVX_MAX_NUM_ROW) * EIGEN_AVX_MAX_NUM_ROW;
233 int64_t K_ = (K / EIGEN_AVX_MAX_K_UNROL) * EIGEN_AVX_MAX_K_UNROL;
234 int64_t j = 0;
235 for (; j < N_; j += U3) {
236 constexpr int64_t EIGEN_AVX_MAX_B_LOAD = EIGEN_AVX_B_LOAD_SETS * 3;
237 int64_t i = 0;
238 for (; i < M_; i += EIGEN_AVX_MAX_NUM_ROW) {
239 Scalar *A_t = &A_arr[idA<isARowMajor>(i, 0, LDA)], *B_t = &B_arr[0 * LDB + j];
240 PacketBlock<vec, EIGEN_ARCH_DEFAULT_NUMBER_OF_REGISTERS> zmm;
241 urolls::template setzero<3, EIGEN_AVX_MAX_NUM_ROW>(zmm);
242 for (int64_t k = 0; k < K_; k += EIGEN_AVX_MAX_K_UNROL) {
243 urolls::template microKernel<isARowMajor, 3, EIGEN_AVX_MAX_NUM_ROW, EIGEN_AVX_MAX_K_UNROL, EIGEN_AVX_MAX_B_LOAD,
244 EIGEN_AVX_MAX_A_BCAST>(B_t, A_t, LDB, LDA, zmm);
245 B_t += EIGEN_AVX_MAX_K_UNROL * LDB;
246 EIGEN_IF_CONSTEXPR (isARowMajor)
247 A_t += EIGEN_AVX_MAX_K_UNROL;
248 else
249 A_t += EIGEN_AVX_MAX_K_UNROL * LDA;
250 }
251 EIGEN_IF_CONSTEXPR (handleKRem) {
252 for (int64_t k = K_; k < K; k++) {
253 urolls::template microKernel<isARowMajor, 3, EIGEN_AVX_MAX_NUM_ROW, 1, EIGEN_AVX_B_LOAD_SETS * 3,
254 EIGEN_AVX_MAX_A_BCAST>(B_t, A_t, LDB, LDA, zmm);
255 B_t += LDB;
256 EIGEN_IF_CONSTEXPR (isARowMajor)
257 A_t++;
258 else
259 A_t += LDA;
260 }
261 }
262 EIGEN_IF_CONSTEXPR (isCRowMajor) {
263 urolls::template updateC<3, EIGEN_AVX_MAX_NUM_ROW>(&C_arr[i * LDC + j], LDC, zmm);
264 urolls::template storeC<3, EIGEN_AVX_MAX_NUM_ROW>(&C_arr[i * LDC + j], LDC, zmm);
265 } else {
266 transStoreC<Scalar, vec, EIGEN_AVX_MAX_NUM_ROW, U3, false, false>(zmm, &C_arr[i + j * LDC], LDC);
267 }
268 }
269 if (M - i >= 4) { // Note: this block assumes EIGEN_AVX_MAX_NUM_ROW = 8. Should be removed otherwise
270 Scalar* A_t = &A_arr[idA<isARowMajor>(i, 0, LDA)];
271 Scalar* B_t = &B_arr[0 * LDB + j];
272 PacketBlock<vec, EIGEN_ARCH_DEFAULT_NUMBER_OF_REGISTERS> zmm;
273 urolls::template setzero<3, 4>(zmm);
274 for (int64_t k = 0; k < K_; k += EIGEN_AVX_MAX_K_UNROL) {
275 urolls::template microKernel<isARowMajor, 3, 4, EIGEN_AVX_MAX_K_UNROL, EIGEN_AVX_B_LOAD_SETS * 3,
276 EIGEN_AVX_MAX_A_BCAST>(B_t, A_t, LDB, LDA, zmm);
277 B_t += EIGEN_AVX_MAX_K_UNROL * LDB;
278 EIGEN_IF_CONSTEXPR (isARowMajor)
279 A_t += EIGEN_AVX_MAX_K_UNROL;
280 else
281 A_t += EIGEN_AVX_MAX_K_UNROL * LDA;
282 }
283 EIGEN_IF_CONSTEXPR (handleKRem) {
284 for (int64_t k = K_; k < K; k++) {
285 urolls::template microKernel<isARowMajor, 3, 4, 1, EIGEN_AVX_B_LOAD_SETS * 3, EIGEN_AVX_MAX_A_BCAST>(
286 B_t, A_t, LDB, LDA, zmm);
287 B_t += LDB;
288 EIGEN_IF_CONSTEXPR (isARowMajor)
289 A_t++;
290 else
291 A_t += LDA;
292 }
293 }
294 EIGEN_IF_CONSTEXPR (isCRowMajor) {
295 urolls::template updateC<3, 4>(&C_arr[i * LDC + j], LDC, zmm);
296 urolls::template storeC<3, 4>(&C_arr[i * LDC + j], LDC, zmm);
297 } else {
298 transStoreC<Scalar, vec, EIGEN_AVX_MAX_NUM_ROW, U3, true, false>(zmm, &C_arr[i + j * LDC], LDC, 4);
299 }
300 i += 4;
301 }
302 if (M - i >= 2) {
303 Scalar* A_t = &A_arr[idA<isARowMajor>(i, 0, LDA)];
304 Scalar* B_t = &B_arr[0 * LDB + j];
305 PacketBlock<vec, EIGEN_ARCH_DEFAULT_NUMBER_OF_REGISTERS> zmm;
306 urolls::template setzero<3, 2>(zmm);
307 for (int64_t k = 0; k < K_; k += EIGEN_AVX_MAX_K_UNROL) {
308 urolls::template microKernel<isARowMajor, 3, 2, EIGEN_AVX_MAX_K_UNROL, EIGEN_AVX_B_LOAD_SETS * 3,
309 EIGEN_AVX_MAX_A_BCAST>(B_t, A_t, LDB, LDA, zmm);
310 B_t += EIGEN_AVX_MAX_K_UNROL * LDB;
311 EIGEN_IF_CONSTEXPR (isARowMajor)
312 A_t += EIGEN_AVX_MAX_K_UNROL;
313 else
314 A_t += EIGEN_AVX_MAX_K_UNROL * LDA;
315 }
316 EIGEN_IF_CONSTEXPR (handleKRem) {
317 for (int64_t k = K_; k < K; k++) {
318 urolls::template microKernel<isARowMajor, 3, 2, 1, EIGEN_AVX_B_LOAD_SETS * 3, EIGEN_AVX_MAX_A_BCAST>(
319 B_t, A_t, LDB, LDA, zmm);
320 B_t += LDB;
321 EIGEN_IF_CONSTEXPR (isARowMajor)
322 A_t++;
323 else
324 A_t += LDA;
325 }
326 }
327 EIGEN_IF_CONSTEXPR (isCRowMajor) {
328 urolls::template updateC<3, 2>(&C_arr[i * LDC + j], LDC, zmm);
329 urolls::template storeC<3, 2>(&C_arr[i * LDC + j], LDC, zmm);
330 } else {
331 transStoreC<Scalar, vec, EIGEN_AVX_MAX_NUM_ROW, U3, true, false>(zmm, &C_arr[i + j * LDC], LDC, 2);
332 }
333 i += 2;
334 }
335 if (M - i > 0) {
336 Scalar* A_t = &A_arr[idA<isARowMajor>(i, 0, LDA)];
337 Scalar* B_t = &B_arr[0 * LDB + j];
338 PacketBlock<vec, EIGEN_ARCH_DEFAULT_NUMBER_OF_REGISTERS> zmm;
339 urolls::template setzero<3, 1>(zmm);
340 {
341 for (int64_t k = 0; k < K_; k += EIGEN_AVX_MAX_K_UNROL) {
342 urolls::template microKernel<isARowMajor, 3, 1, EIGEN_AVX_MAX_K_UNROL, EIGEN_AVX_B_LOAD_SETS * 3, 1>(
343 B_t, A_t, LDB, LDA, zmm);
344 B_t += EIGEN_AVX_MAX_K_UNROL * LDB;
345 EIGEN_IF_CONSTEXPR (isARowMajor)
346 A_t += EIGEN_AVX_MAX_K_UNROL;
347 else
348 A_t += EIGEN_AVX_MAX_K_UNROL * LDA;
349 }
350 EIGEN_IF_CONSTEXPR (handleKRem) {
351 for (int64_t k = K_; k < K; k++) {
352 urolls::template microKernel<isARowMajor, 3, 1, 1, EIGEN_AVX_B_LOAD_SETS * 3, 1>(B_t, A_t, LDB, LDA, zmm);
353 B_t += LDB;
354 EIGEN_IF_CONSTEXPR (isARowMajor)
355 A_t++;
356 else
357 A_t += LDA;
358 }
359 }
360 EIGEN_IF_CONSTEXPR (isCRowMajor) {
361 urolls::template updateC<3, 1>(&C_arr[i * LDC + j], LDC, zmm);
362 urolls::template storeC<3, 1>(&C_arr[i * LDC + j], LDC, zmm);
363 } else {
364 transStoreC<Scalar, vec, EIGEN_AVX_MAX_NUM_ROW, U3, true, false>(zmm, &C_arr[i + j * LDC], LDC, 1);
365 }
366 }
367 }
368 }
369 if (N - j >= U2) {
370 constexpr int64_t EIGEN_AVX_MAX_B_LOAD = EIGEN_AVX_B_LOAD_SETS * 2;
371 int64_t i = 0;
372 for (; i < M_; i += EIGEN_AVX_MAX_NUM_ROW) {
373 Scalar *A_t = &A_arr[idA<isARowMajor>(i, 0, LDA)], *B_t = &B_arr[0 * LDB + j];
374 EIGEN_IF_CONSTEXPR (isCRowMajor) B_t = &B_arr[0 * LDB + j];
375 PacketBlock<vec, EIGEN_ARCH_DEFAULT_NUMBER_OF_REGISTERS> zmm;
376 urolls::template setzero<2, EIGEN_AVX_MAX_NUM_ROW>(zmm);
377 for (int64_t k = 0; k < K_; k += EIGEN_AVX_MAX_K_UNROL) {
378 urolls::template microKernel<isARowMajor, 2, EIGEN_AVX_MAX_NUM_ROW, EIGEN_AVX_MAX_K_UNROL, EIGEN_AVX_MAX_B_LOAD,
379 EIGEN_AVX_MAX_A_BCAST>(B_t, A_t, LDB, LDA, zmm);
380 B_t += EIGEN_AVX_MAX_K_UNROL * LDB;
381 EIGEN_IF_CONSTEXPR (isARowMajor)
382 A_t += EIGEN_AVX_MAX_K_UNROL;
383 else
384 A_t += EIGEN_AVX_MAX_K_UNROL * LDA;
385 }
386 EIGEN_IF_CONSTEXPR (handleKRem) {
387 for (int64_t k = K_; k < K; k++) {
388 urolls::template microKernel<isARowMajor, 2, EIGEN_AVX_MAX_NUM_ROW, 1, EIGEN_AVX_MAX_B_LOAD,
389 EIGEN_AVX_MAX_A_BCAST>(B_t, A_t, LDB, LDA, zmm);
390 B_t += LDB;
391 EIGEN_IF_CONSTEXPR (isARowMajor)
392 A_t++;
393 else
394 A_t += LDA;
395 }
396 }
397 EIGEN_IF_CONSTEXPR (isCRowMajor) {
398 urolls::template updateC<2, EIGEN_AVX_MAX_NUM_ROW>(&C_arr[i * LDC + j], LDC, zmm);
399 urolls::template storeC<2, EIGEN_AVX_MAX_NUM_ROW>(&C_arr[i * LDC + j], LDC, zmm);
400 } else {
401 transStoreC<Scalar, vec, EIGEN_AVX_MAX_NUM_ROW, U2, false, false>(zmm, &C_arr[i + j * LDC], LDC);
402 }
403 }
404 if (M - i >= 4) { // Note: this block assumes EIGEN_AVX_MAX_NUM_ROW = 8. Should be removed otherwise
405 Scalar* A_t = &A_arr[idA<isARowMajor>(i, 0, LDA)];
406 Scalar* B_t = &B_arr[0 * LDB + j];
407 PacketBlock<vec, EIGEN_ARCH_DEFAULT_NUMBER_OF_REGISTERS> zmm;
408 urolls::template setzero<2, 4>(zmm);
409 for (int64_t k = 0; k < K_; k += EIGEN_AVX_MAX_K_UNROL) {
410 urolls::template microKernel<isARowMajor, 2, 4, EIGEN_AVX_MAX_K_UNROL, EIGEN_AVX_MAX_B_LOAD,
411 EIGEN_AVX_MAX_A_BCAST>(B_t, A_t, LDB, LDA, zmm);
412 B_t += EIGEN_AVX_MAX_K_UNROL * LDB;
413 EIGEN_IF_CONSTEXPR (isARowMajor)
414 A_t += EIGEN_AVX_MAX_K_UNROL;
415 else
416 A_t += EIGEN_AVX_MAX_K_UNROL * LDA;
417 }
418 EIGEN_IF_CONSTEXPR (handleKRem) {
419 for (int64_t k = K_; k < K; k++) {
420 urolls::template microKernel<isARowMajor, 2, 4, 1, EIGEN_AVX_MAX_B_LOAD, EIGEN_AVX_MAX_A_BCAST>(B_t, A_t, LDB,
421 LDA, zmm);
422 B_t += LDB;
423 EIGEN_IF_CONSTEXPR (isARowMajor)
424 A_t++;
425 else
426 A_t += LDA;
427 }
428 }
429 EIGEN_IF_CONSTEXPR (isCRowMajor) {
430 urolls::template updateC<2, 4>(&C_arr[i * LDC + j], LDC, zmm);
431 urolls::template storeC<2, 4>(&C_arr[i * LDC + j], LDC, zmm);
432 } else {
433 transStoreC<Scalar, vec, EIGEN_AVX_MAX_NUM_ROW, U2, true, false>(zmm, &C_arr[i + j * LDC], LDC, 4);
434 }
435 i += 4;
436 }
437 if (M - i >= 2) {
438 Scalar* A_t = &A_arr[idA<isARowMajor>(i, 0, LDA)];
439 Scalar* B_t = &B_arr[0 * LDB + j];
440 PacketBlock<vec, EIGEN_ARCH_DEFAULT_NUMBER_OF_REGISTERS> zmm;
441 urolls::template setzero<2, 2>(zmm);
442 for (int64_t k = 0; k < K_; k += EIGEN_AVX_MAX_K_UNROL) {
443 urolls::template microKernel<isARowMajor, 2, 2, EIGEN_AVX_MAX_K_UNROL, EIGEN_AVX_MAX_B_LOAD,
444 EIGEN_AVX_MAX_A_BCAST>(B_t, A_t, LDB, LDA, zmm);
445 B_t += EIGEN_AVX_MAX_K_UNROL * LDB;
446 EIGEN_IF_CONSTEXPR (isARowMajor)
447 A_t += EIGEN_AVX_MAX_K_UNROL;
448 else
449 A_t += EIGEN_AVX_MAX_K_UNROL * LDA;
450 }
451 EIGEN_IF_CONSTEXPR (handleKRem) {
452 for (int64_t k = K_; k < K; k++) {
453 urolls::template microKernel<isARowMajor, 2, 2, 1, EIGEN_AVX_MAX_B_LOAD, EIGEN_AVX_MAX_A_BCAST>(B_t, A_t, LDB,
454 LDA, zmm);
455 B_t += LDB;
456 EIGEN_IF_CONSTEXPR (isARowMajor)
457 A_t++;
458 else
459 A_t += LDA;
460 }
461 }
462 EIGEN_IF_CONSTEXPR (isCRowMajor) {
463 urolls::template updateC<2, 2>(&C_arr[i * LDC + j], LDC, zmm);
464 urolls::template storeC<2, 2>(&C_arr[i * LDC + j], LDC, zmm);
465 } else {
466 transStoreC<Scalar, vec, EIGEN_AVX_MAX_NUM_ROW, U2, true, false>(zmm, &C_arr[i + j * LDC], LDC, 2);
467 }
468 i += 2;
469 }
470 if (M - i > 0) {
471 Scalar* A_t = &A_arr[idA<isARowMajor>(i, 0, LDA)];
472 Scalar* B_t = &B_arr[0 * LDB + j];
473 PacketBlock<vec, EIGEN_ARCH_DEFAULT_NUMBER_OF_REGISTERS> zmm;
474 urolls::template setzero<2, 1>(zmm);
475 for (int64_t k = 0; k < K_; k += EIGEN_AVX_MAX_K_UNROL) {
476 urolls::template microKernel<isARowMajor, 2, 1, EIGEN_AVX_MAX_K_UNROL, EIGEN_AVX_MAX_B_LOAD, 1>(B_t, A_t, LDB,
477 LDA, zmm);
478 B_t += EIGEN_AVX_MAX_K_UNROL * LDB;
479 EIGEN_IF_CONSTEXPR (isARowMajor)
480 A_t += EIGEN_AVX_MAX_K_UNROL;
481 else
482 A_t += EIGEN_AVX_MAX_K_UNROL * LDA;
483 }
484 EIGEN_IF_CONSTEXPR (handleKRem) {
485 for (int64_t k = K_; k < K; k++) {
486 urolls::template microKernel<isARowMajor, 2, 1, 1, EIGEN_AVX_MAX_B_LOAD, 1>(B_t, A_t, LDB, LDA, zmm);
487 B_t += LDB;
488 EIGEN_IF_CONSTEXPR (isARowMajor)
489 A_t++;
490 else
491 A_t += LDA;
492 }
493 }
494 EIGEN_IF_CONSTEXPR (isCRowMajor) {
495 urolls::template updateC<2, 1>(&C_arr[i * LDC + j], LDC, zmm);
496 urolls::template storeC<2, 1>(&C_arr[i * LDC + j], LDC, zmm);
497 } else {
498 transStoreC<Scalar, vec, EIGEN_AVX_MAX_NUM_ROW, U2, true, false>(zmm, &C_arr[i + j * LDC], LDC, 1);
499 }
500 }
501 j += U2;
502 }
503 if (N - j >= U1) {
504 constexpr int64_t EIGEN_AVX_MAX_B_LOAD = EIGEN_AVX_B_LOAD_SETS * 1;
505 int64_t i = 0;
506 for (; i < M_; i += EIGEN_AVX_MAX_NUM_ROW) {
507 Scalar *A_t = &A_arr[idA<isARowMajor>(i, 0, LDA)], *B_t = &B_arr[0 * LDB + j];
508 PacketBlock<vec, EIGEN_ARCH_DEFAULT_NUMBER_OF_REGISTERS> zmm;
509 urolls::template setzero<1, EIGEN_AVX_MAX_NUM_ROW>(zmm);
510 for (int64_t k = 0; k < K_; k += EIGEN_AVX_MAX_K_UNROL) {
511 urolls::template microKernel<isARowMajor, 1, EIGEN_AVX_MAX_NUM_ROW, EIGEN_AVX_MAX_K_UNROL, EIGEN_AVX_MAX_B_LOAD,
512 EIGEN_AVX_MAX_A_BCAST>(B_t, A_t, LDB, LDA, zmm);
513 B_t += EIGEN_AVX_MAX_K_UNROL * LDB;
514 EIGEN_IF_CONSTEXPR (isARowMajor)
515 A_t += EIGEN_AVX_MAX_K_UNROL;
516 else
517 A_t += EIGEN_AVX_MAX_K_UNROL * LDA;
518 }
519 EIGEN_IF_CONSTEXPR (handleKRem) {
520 for (int64_t k = K_; k < K; k++) {
521 urolls::template microKernel<isARowMajor, 1, EIGEN_AVX_MAX_NUM_ROW, 1, EIGEN_AVX_B_LOAD_SETS * 1,
522 EIGEN_AVX_MAX_A_BCAST>(B_t, A_t, LDB, LDA, zmm);
523 B_t += LDB;
524 EIGEN_IF_CONSTEXPR (isARowMajor)
525 A_t++;
526 else
527 A_t += LDA;
528 }
529 }
530 EIGEN_IF_CONSTEXPR (isCRowMajor) {
531 urolls::template updateC<1, EIGEN_AVX_MAX_NUM_ROW>(&C_arr[i * LDC + j], LDC, zmm);
532 urolls::template storeC<1, EIGEN_AVX_MAX_NUM_ROW>(&C_arr[i * LDC + j], LDC, zmm);
533 } else {
534 transStoreC<Scalar, vec, EIGEN_AVX_MAX_NUM_ROW, U1, false, false>(zmm, &C_arr[i + j * LDC], LDC);
535 }
536 }
537 if (M - i >= 4) { // Note: this block assumes EIGEN_AVX_MAX_NUM_ROW = 8. Should be removed otherwise
538 Scalar* A_t = &A_arr[idA<isARowMajor>(i, 0, LDA)];
539 Scalar* B_t = &B_arr[0 * LDB + j];
540 PacketBlock<vec, EIGEN_ARCH_DEFAULT_NUMBER_OF_REGISTERS> zmm;
541 urolls::template setzero<1, 4>(zmm);
542 for (int64_t k = 0; k < K_; k += EIGEN_AVX_MAX_K_UNROL) {
543 urolls::template microKernel<isARowMajor, 1, 4, EIGEN_AVX_MAX_K_UNROL, EIGEN_AVX_MAX_B_LOAD,
544 EIGEN_AVX_MAX_A_BCAST>(B_t, A_t, LDB, LDA, zmm);
545 B_t += EIGEN_AVX_MAX_K_UNROL * LDB;
546 EIGEN_IF_CONSTEXPR (isARowMajor)
547 A_t += EIGEN_AVX_MAX_K_UNROL;
548 else
549 A_t += EIGEN_AVX_MAX_K_UNROL * LDA;
550 }
551 EIGEN_IF_CONSTEXPR (handleKRem) {
552 for (int64_t k = K_; k < K; k++) {
553 urolls::template microKernel<isARowMajor, 1, 4, 1, EIGEN_AVX_MAX_B_LOAD, EIGEN_AVX_MAX_A_BCAST>(B_t, A_t, LDB,
554 LDA, zmm);
555 B_t += LDB;
556 EIGEN_IF_CONSTEXPR (isARowMajor)
557 A_t++;
558 else
559 A_t += LDA;
560 }
561 }
562 EIGEN_IF_CONSTEXPR (isCRowMajor) {
563 urolls::template updateC<1, 4>(&C_arr[i * LDC + j], LDC, zmm);
564 urolls::template storeC<1, 4>(&C_arr[i * LDC + j], LDC, zmm);
565 } else {
566 transStoreC<Scalar, vec, EIGEN_AVX_MAX_NUM_ROW, U1, true, false>(zmm, &C_arr[i + j * LDC], LDC, 4);
567 }
568 i += 4;
569 }
570 if (M - i >= 2) {
571 Scalar* A_t = &A_arr[idA<isARowMajor>(i, 0, LDA)];
572 Scalar* B_t = &B_arr[0 * LDB + j];
573 PacketBlock<vec, EIGEN_ARCH_DEFAULT_NUMBER_OF_REGISTERS> zmm;
574 urolls::template setzero<1, 2>(zmm);
575 for (int64_t k = 0; k < K_; k += EIGEN_AVX_MAX_K_UNROL) {
576 urolls::template microKernel<isARowMajor, 1, 2, EIGEN_AVX_MAX_K_UNROL, EIGEN_AVX_MAX_B_LOAD,
577 EIGEN_AVX_MAX_A_BCAST>(B_t, A_t, LDB, LDA, zmm);
578 B_t += EIGEN_AVX_MAX_K_UNROL * LDB;
579 EIGEN_IF_CONSTEXPR (isARowMajor)
580 A_t += EIGEN_AVX_MAX_K_UNROL;
581 else
582 A_t += EIGEN_AVX_MAX_K_UNROL * LDA;
583 }
584 EIGEN_IF_CONSTEXPR (handleKRem) {
585 for (int64_t k = K_; k < K; k++) {
586 urolls::template microKernel<isARowMajor, 1, 2, 1, EIGEN_AVX_MAX_B_LOAD, EIGEN_AVX_MAX_A_BCAST>(B_t, A_t, LDB,
587 LDA, zmm);
588 B_t += LDB;
589 EIGEN_IF_CONSTEXPR (isARowMajor)
590 A_t++;
591 else
592 A_t += LDA;
593 }
594 }
595 EIGEN_IF_CONSTEXPR (isCRowMajor) {
596 urolls::template updateC<1, 2>(&C_arr[i * LDC + j], LDC, zmm);
597 urolls::template storeC<1, 2>(&C_arr[i * LDC + j], LDC, zmm);
598 } else {
599 transStoreC<Scalar, vec, EIGEN_AVX_MAX_NUM_ROW, U1, true, false>(zmm, &C_arr[i + j * LDC], LDC, 2);
600 }
601 i += 2;
602 }
603 if (M - i > 0) {
604 Scalar* A_t = &A_arr[idA<isARowMajor>(i, 0, LDA)];
605 Scalar* B_t = &B_arr[0 * LDB + j];
606 PacketBlock<vec, EIGEN_ARCH_DEFAULT_NUMBER_OF_REGISTERS> zmm;
607 urolls::template setzero<1, 1>(zmm);
608 {
609 for (int64_t k = 0; k < K_; k += EIGEN_AVX_MAX_K_UNROL) {
610 urolls::template microKernel<isARowMajor, 1, 1, EIGEN_AVX_MAX_K_UNROL, EIGEN_AVX_MAX_B_LOAD, 1>(B_t, A_t, LDB,
611 LDA, zmm);
612 B_t += EIGEN_AVX_MAX_K_UNROL * LDB;
613 EIGEN_IF_CONSTEXPR (isARowMajor)
614 A_t += EIGEN_AVX_MAX_K_UNROL;
615 else
616 A_t += EIGEN_AVX_MAX_K_UNROL * LDA;
617 }
618 EIGEN_IF_CONSTEXPR (handleKRem) {
619 for (int64_t k = K_; k < K; k++) {
620 urolls::template microKernel<isARowMajor, 1, 1, 1, EIGEN_AVX_B_LOAD_SETS * 1, 1>(B_t, A_t, LDB, LDA, zmm);
621 B_t += LDB;
622 EIGEN_IF_CONSTEXPR (isARowMajor)
623 A_t++;
624 else
625 A_t += LDA;
626 }
627 }
628 EIGEN_IF_CONSTEXPR (isCRowMajor) {
629 urolls::template updateC<1, 1>(&C_arr[i * LDC + j], LDC, zmm);
630 urolls::template storeC<1, 1>(&C_arr[i * LDC + j], LDC, zmm);
631 } else {
632 transStoreC<Scalar, vec, EIGEN_AVX_MAX_NUM_ROW, U1, true, false>(zmm, &C_arr[i + j * LDC], LDC, 1);
633 }
634 }
635 }
636 j += U1;
637 }
638 if (N - j > 0) {
639 constexpr int64_t EIGEN_AVX_MAX_B_LOAD = EIGEN_AVX_B_LOAD_SETS * 1;
640 int64_t i = 0;
641 for (; i < M_; i += EIGEN_AVX_MAX_NUM_ROW) {
642 Scalar* A_t = &A_arr[idA<isARowMajor>(i, 0, LDA)];
643 Scalar* B_t = &B_arr[0 * LDB + j];
644 PacketBlock<vec, EIGEN_ARCH_DEFAULT_NUMBER_OF_REGISTERS> zmm;
645 urolls::template setzero<1, EIGEN_AVX_MAX_NUM_ROW>(zmm);
646 for (int64_t k = 0; k < K_; k += EIGEN_AVX_MAX_K_UNROL) {
647 urolls::template microKernel<isARowMajor, 1, EIGEN_AVX_MAX_NUM_ROW, EIGEN_AVX_MAX_K_UNROL, EIGEN_AVX_MAX_B_LOAD,
648 EIGEN_AVX_MAX_A_BCAST, true>(B_t, A_t, LDB, LDA, zmm, N - j);
649 B_t += EIGEN_AVX_MAX_K_UNROL * LDB;
650 EIGEN_IF_CONSTEXPR (isARowMajor)
651 A_t += EIGEN_AVX_MAX_K_UNROL;
652 else
653 A_t += EIGEN_AVX_MAX_K_UNROL * LDA;
654 }
655 EIGEN_IF_CONSTEXPR (handleKRem) {
656 for (int64_t k = K_; k < K; k++) {
657 urolls::template microKernel<isARowMajor, 1, EIGEN_AVX_MAX_NUM_ROW, 1, EIGEN_AVX_MAX_B_LOAD,
658 EIGEN_AVX_MAX_A_BCAST, true>(B_t, A_t, LDB, LDA, zmm, N - j);
659 B_t += LDB;
660 EIGEN_IF_CONSTEXPR (isARowMajor)
661 A_t++;
662 else
663 A_t += LDA;
664 }
665 }
666 EIGEN_IF_CONSTEXPR (isCRowMajor) {
667 urolls::template updateC<1, EIGEN_AVX_MAX_NUM_ROW, true>(&C_arr[i * LDC + j], LDC, zmm, N - j);
668 urolls::template storeC<1, EIGEN_AVX_MAX_NUM_ROW, true>(&C_arr[i * LDC + j], LDC, zmm, N - j);
669 } else {
670 transStoreC<Scalar, vec, EIGEN_AVX_MAX_NUM_ROW, U1, false, true>(zmm, &C_arr[i + j * LDC], LDC, 0, N - j);
671 }
672 }
673 if (M - i >= 4) { // Note: this block assumes EIGEN_AVX_MAX_NUM_ROW = 8. Should be removed otherwise
674 Scalar* A_t = &A_arr[idA<isARowMajor>(i, 0, LDA)];
675 Scalar* B_t = &B_arr[0 * LDB + j];
676 PacketBlock<vec, EIGEN_ARCH_DEFAULT_NUMBER_OF_REGISTERS> zmm;
677 urolls::template setzero<1, 4>(zmm);
678 for (int64_t k = 0; k < K_; k += EIGEN_AVX_MAX_K_UNROL) {
679 urolls::template microKernel<isARowMajor, 1, 4, EIGEN_AVX_MAX_K_UNROL, EIGEN_AVX_MAX_B_LOAD,
680 EIGEN_AVX_MAX_A_BCAST, true>(B_t, A_t, LDB, LDA, zmm, N - j);
681 B_t += EIGEN_AVX_MAX_K_UNROL * LDB;
682 EIGEN_IF_CONSTEXPR (isARowMajor)
683 A_t += EIGEN_AVX_MAX_K_UNROL;
684 else
685 A_t += EIGEN_AVX_MAX_K_UNROL * LDA;
686 }
687 EIGEN_IF_CONSTEXPR (handleKRem) {
688 for (int64_t k = K_; k < K; k++) {
689 urolls::template microKernel<isARowMajor, 1, 4, 1, EIGEN_AVX_MAX_B_LOAD, EIGEN_AVX_MAX_A_BCAST, true>(
690 B_t, A_t, LDB, LDA, zmm, N - j);
691 B_t += LDB;
692 EIGEN_IF_CONSTEXPR (isARowMajor)
693 A_t++;
694 else
695 A_t += LDA;
696 }
697 }
698 EIGEN_IF_CONSTEXPR (isCRowMajor) {
699 urolls::template updateC<1, 4, true>(&C_arr[i * LDC + j], LDC, zmm, N - j);
700 urolls::template storeC<1, 4, true>(&C_arr[i * LDC + j], LDC, zmm, N - j);
701 } else {
702 transStoreC<Scalar, vec, EIGEN_AVX_MAX_NUM_ROW, U1, true, true>(zmm, &C_arr[i + j * LDC], LDC, 4, N - j);
703 }
704 i += 4;
705 }
706 if (M - i >= 2) {
707 Scalar* A_t = &A_arr[idA<isARowMajor>(i, 0, LDA)];
708 Scalar* B_t = &B_arr[0 * LDB + j];
709 PacketBlock<vec, EIGEN_ARCH_DEFAULT_NUMBER_OF_REGISTERS> zmm;
710 urolls::template setzero<1, 2>(zmm);
711 for (int64_t k = 0; k < K_; k += EIGEN_AVX_MAX_K_UNROL) {
712 urolls::template microKernel<isARowMajor, 1, 2, EIGEN_AVX_MAX_K_UNROL, EIGEN_AVX_MAX_B_LOAD,
713 EIGEN_AVX_MAX_A_BCAST, true>(B_t, A_t, LDB, LDA, zmm, N - j);
714 B_t += EIGEN_AVX_MAX_K_UNROL * LDB;
715 EIGEN_IF_CONSTEXPR (isARowMajor)
716 A_t += EIGEN_AVX_MAX_K_UNROL;
717 else
718 A_t += EIGEN_AVX_MAX_K_UNROL * LDA;
719 }
720 EIGEN_IF_CONSTEXPR (handleKRem) {
721 for (int64_t k = K_; k < K; k++) {
722 urolls::template microKernel<isARowMajor, 1, 2, 1, EIGEN_AVX_MAX_B_LOAD, EIGEN_AVX_MAX_A_BCAST, true>(
723 B_t, A_t, LDB, LDA, zmm, N - j);
724 B_t += LDB;
725 EIGEN_IF_CONSTEXPR (isARowMajor)
726 A_t++;
727 else
728 A_t += LDA;
729 }
730 }
731 EIGEN_IF_CONSTEXPR (isCRowMajor) {
732 urolls::template updateC<1, 2, true>(&C_arr[i * LDC + j], LDC, zmm, N - j);
733 urolls::template storeC<1, 2, true>(&C_arr[i * LDC + j], LDC, zmm, N - j);
734 } else {
735 transStoreC<Scalar, vec, EIGEN_AVX_MAX_NUM_ROW, U1, true, true>(zmm, &C_arr[i + j * LDC], LDC, 2, N - j);
736 }
737 i += 2;
738 }
739 if (M - i > 0) {
740 Scalar* A_t = &A_arr[idA<isARowMajor>(i, 0, LDA)];
741 Scalar* B_t = &B_arr[0 * LDB + j];
742 PacketBlock<vec, EIGEN_ARCH_DEFAULT_NUMBER_OF_REGISTERS> zmm;
743 urolls::template setzero<1, 1>(zmm);
744 for (int64_t k = 0; k < K_; k += EIGEN_AVX_MAX_K_UNROL) {
745 urolls::template microKernel<isARowMajor, 1, 1, EIGEN_AVX_MAX_K_UNROL, EIGEN_AVX_MAX_B_LOAD, 1, true>(
746 B_t, A_t, LDB, LDA, zmm, N - j);
747 B_t += EIGEN_AVX_MAX_K_UNROL * LDB;
748 EIGEN_IF_CONSTEXPR (isARowMajor)
749 A_t += EIGEN_AVX_MAX_K_UNROL;
750 else
751 A_t += EIGEN_AVX_MAX_K_UNROL * LDA;
752 }
753 EIGEN_IF_CONSTEXPR (handleKRem) {
754 for (int64_t k = K_; k < K; k++) {
755 urolls::template microKernel<isARowMajor, 1, 1, 1, EIGEN_AVX_MAX_B_LOAD, 1, true>(B_t, A_t, LDB, LDA, zmm,
756 N - j);
757 B_t += LDB;
758 EIGEN_IF_CONSTEXPR (isARowMajor)
759 A_t++;
760 else
761 A_t += LDA;
762 }
763 }
764 EIGEN_IF_CONSTEXPR (isCRowMajor) {
765 urolls::template updateC<1, 1, true>(&C_arr[i * LDC + j], LDC, zmm, N - j);
766 urolls::template storeC<1, 1, true>(&C_arr[i * LDC + j], LDC, zmm, N - j);
767 } else {
768 transStoreC<Scalar, vec, EIGEN_AVX_MAX_NUM_ROW, U1, true, true>(zmm, &C_arr[i + j * LDC], LDC, 1, N - j);
769 }
770 }
771 }
772}
773
782template <typename Scalar, typename vec, int64_t unrollM, bool isARowMajor, bool isFWDSolve, bool isUnitDiag>
783EIGEN_ALWAYS_INLINE void triSolveKernel(Scalar* A_arr, Scalar* B_arr, int64_t K, int64_t LDA, int64_t LDB) {
784 static_assert(unrollM <= EIGEN_AVX_MAX_NUM_ROW, "unrollM should be equal to EIGEN_AVX_MAX_NUM_ROW");
785 using urolls = unrolls::trsm<Scalar>;
786 constexpr int64_t U3 = urolls::PacketSize * 3;
787 constexpr int64_t U2 = urolls::PacketSize * 2;
788 constexpr int64_t U1 = urolls::PacketSize * 1;
789
790 PacketBlock<vec, EIGEN_AVX_MAX_NUM_ACC> RHSInPacket;
791 PacketBlock<vec, EIGEN_AVX_MAX_NUM_ROW> AInPacket;
792
793 int64_t k = 0;
794 while (K - k >= U3) {
795 urolls::template loadRHS<isFWDSolve, unrollM, 3>(B_arr + k, LDB, RHSInPacket);
796 urolls::template triSolveMicroKernel<isARowMajor, isFWDSolve, isUnitDiag, unrollM, 3>(A_arr, LDA, RHSInPacket,
797 AInPacket);
798 urolls::template storeRHS<isFWDSolve, unrollM, 3>(B_arr + k, LDB, RHSInPacket);
799 k += U3;
800 }
801 if (K - k >= U2) {
802 urolls::template loadRHS<isFWDSolve, unrollM, 2>(B_arr + k, LDB, RHSInPacket);
803 urolls::template triSolveMicroKernel<isARowMajor, isFWDSolve, isUnitDiag, unrollM, 2>(A_arr, LDA, RHSInPacket,
804 AInPacket);
805 urolls::template storeRHS<isFWDSolve, unrollM, 2>(B_arr + k, LDB, RHSInPacket);
806 k += U2;
807 }
808 if (K - k >= U1) {
809 urolls::template loadRHS<isFWDSolve, unrollM, 1>(B_arr + k, LDB, RHSInPacket);
810 urolls::template triSolveMicroKernel<isARowMajor, isFWDSolve, isUnitDiag, unrollM, 1>(A_arr, LDA, RHSInPacket,
811 AInPacket);
812 urolls::template storeRHS<isFWDSolve, unrollM, 1>(B_arr + k, LDB, RHSInPacket);
813 k += U1;
814 }
815 if (K - k > 0) {
816 // Handle remaining number of RHS
817 urolls::template loadRHS<isFWDSolve, unrollM, 1, true>(B_arr + k, LDB, RHSInPacket, K - k);
818 urolls::template triSolveMicroKernel<isARowMajor, isFWDSolve, isUnitDiag, unrollM, 1>(A_arr, LDA, RHSInPacket,
819 AInPacket);
820 urolls::template storeRHS<isFWDSolve, unrollM, 1, true>(B_arr + k, LDB, RHSInPacket, K - k);
821 }
822}
823
832template <typename Scalar, bool isARowMajor, bool isFWDSolve, bool isUnitDiag>
833void triSolveKernelLxK(Scalar* A_arr, Scalar* B_arr, int64_t M, int64_t K, int64_t LDA, int64_t LDB) {
834 // Note: this assumes EIGEN_AVX_MAX_NUM_ROW = 8. Unrolls should be adjusted
835 // accordingly if EIGEN_AVX_MAX_NUM_ROW is smaller.
836 using vec = std::conditional_t<std::is_same<Scalar, float>::value, vecFullFloat, vecFullDouble>;
837 if (M == 8)
838 triSolveKernel<Scalar, vec, 8, isARowMajor, isFWDSolve, isUnitDiag>(A_arr, B_arr, K, LDA, LDB);
839 else if (M == 7)
840 triSolveKernel<Scalar, vec, 7, isARowMajor, isFWDSolve, isUnitDiag>(A_arr, B_arr, K, LDA, LDB);
841 else if (M == 6)
842 triSolveKernel<Scalar, vec, 6, isARowMajor, isFWDSolve, isUnitDiag>(A_arr, B_arr, K, LDA, LDB);
843 else if (M == 5)
844 triSolveKernel<Scalar, vec, 5, isARowMajor, isFWDSolve, isUnitDiag>(A_arr, B_arr, K, LDA, LDB);
845 else if (M == 4)
846 triSolveKernel<Scalar, vec, 4, isARowMajor, isFWDSolve, isUnitDiag>(A_arr, B_arr, K, LDA, LDB);
847 else if (M == 3)
848 triSolveKernel<Scalar, vec, 3, isARowMajor, isFWDSolve, isUnitDiag>(A_arr, B_arr, K, LDA, LDB);
849 else if (M == 2)
850 triSolveKernel<Scalar, vec, 2, isARowMajor, isFWDSolve, isUnitDiag>(A_arr, B_arr, K, LDA, LDB);
851 else if (M == 1)
852 triSolveKernel<Scalar, vec, 1, isARowMajor, isFWDSolve, isUnitDiag>(A_arr, B_arr, K, LDA, LDB);
853 return;
854}
855
863template <typename Scalar, bool toTemp = true, bool remM = false>
864EIGEN_ALWAYS_INLINE void copyBToRowMajor(Scalar* B_arr, int64_t LDB, int64_t K, Scalar* B_temp, int64_t LDB_,
865 int64_t remM_ = 0) {
866 EIGEN_UNUSED_VARIABLE(remM_);
867 using urolls = unrolls::transB<Scalar>;
868 using vecHalf = std::conditional_t<std::is_same<Scalar, float>::value, vecHalfFloat, vecFullDouble>;
869 PacketBlock<vecHalf, EIGEN_ARCH_DEFAULT_NUMBER_OF_REGISTERS> ymm;
870 constexpr int64_t U3 = urolls::PacketSize * 3;
871 constexpr int64_t U2 = urolls::PacketSize * 2;
872 constexpr int64_t U1 = urolls::PacketSize * 1;
873 int64_t K_ = K / U3 * U3;
874 int64_t k = 0;
875
876 for (; k < K_; k += U3) {
877 urolls::template transB_kernel<U3, toTemp, remM>(B_arr + k * LDB, LDB, B_temp, LDB_, ymm, remM_);
878 B_temp += U3;
879 }
880 if (K - k >= U2) {
881 urolls::template transB_kernel<U2, toTemp, remM>(B_arr + k * LDB, LDB, B_temp, LDB_, ymm, remM_);
882 B_temp += U2;
883 k += U2;
884 }
885 if (K - k >= U1) {
886 urolls::template transB_kernel<U1, toTemp, remM>(B_arr + k * LDB, LDB, B_temp, LDB_, ymm, remM_);
887 B_temp += U1;
888 k += U1;
889 }
890 EIGEN_IF_CONSTEXPR (U1 > 8) {
891 // Note: without "if constexpr" this section of code will also be
892 // parsed by the compiler so there is an additional check in {load/store}BBlock
893 // to make sure the counter is not non-negative.
894 if (K - k >= 8) {
895 urolls::template transB_kernel<8, toTemp, remM>(B_arr + k * LDB, LDB, B_temp, LDB_, ymm, remM_);
896 B_temp += 8;
897 k += 8;
898 }
899 }
900 EIGEN_IF_CONSTEXPR (U1 > 4) {
901 // Note: without "if constexpr" this section of code will also be
902 // parsed by the compiler so there is an additional check in {load/store}BBlock
903 // to make sure the counter is not non-negative.
904 if (K - k >= 4) {
905 urolls::template transB_kernel<4, toTemp, remM>(B_arr + k * LDB, LDB, B_temp, LDB_, ymm, remM_);
906 B_temp += 4;
907 k += 4;
908 }
909 }
910 if (K - k >= 2) {
911 urolls::template transB_kernel<2, toTemp, remM>(B_arr + k * LDB, LDB, B_temp, LDB_, ymm, remM_);
912 B_temp += 2;
913 k += 2;
914 }
915 if (K - k >= 1) {
916 urolls::template transB_kernel<1, toTemp, remM>(B_arr + k * LDB, LDB, B_temp, LDB_, ymm, remM_);
917 B_temp += 1;
918 k += 1;
919 }
920}
921
949template <typename Scalar, bool isARowMajor = true, bool isBRowMajor = true, bool isFWDSolve = true,
950 bool isUnitDiag = false>
951void triSolve(Scalar* A_arr, Scalar* B_arr, int64_t M, int64_t numRHS, int64_t LDA, int64_t LDB) {
952 constexpr int64_t psize = packet_traits<Scalar>::size;
966 constexpr int64_t kB = (3 * psize) * 5; // 5*U3
967 constexpr int64_t numM = 8 * EIGEN_AVX_MAX_NUM_ROW;
968
969 int64_t sizeBTemp = 0;
970 Scalar* B_temp = nullptr;
971 EIGEN_IF_CONSTEXPR (!isBRowMajor) {
977 sizeBTemp = (((std::min(kB, numRHS) + psize - 1) / psize + 4) * psize) * numM;
978 }
979
980 EIGEN_IF_CONSTEXPR (!isBRowMajor) B_temp = (Scalar*)handmade_aligned_malloc(sizeof(Scalar) * sizeBTemp, 64);
981
982 for (int64_t k = 0; k < numRHS; k += kB) {
983 int64_t bK = numRHS - k > kB ? kB : numRHS - k;
984 int64_t M_ = (M / EIGEN_AVX_MAX_NUM_ROW) * EIGEN_AVX_MAX_NUM_ROW, gemmOff = 0;
985
986 // bK rounded up to next multiple of L=EIGEN_AVX_MAX_NUM_ROW. When B_temp is used, we solve for bkL RHS
987 // instead of bK RHS in triSolveKernelLxK.
988 int64_t bkL = ((bK + (EIGEN_AVX_MAX_NUM_ROW - 1)) / EIGEN_AVX_MAX_NUM_ROW) * EIGEN_AVX_MAX_NUM_ROW;
989 const int64_t numScalarPerCache = 64 / sizeof(Scalar);
990 // Leading dimension of B_temp, will be a multiple of the cache line size.
991 int64_t LDT = ((bkL + (numScalarPerCache - 1)) / numScalarPerCache) * numScalarPerCache;
992 int64_t offsetBTemp = 0;
993 for (int64_t i = 0; i < M_; i += EIGEN_AVX_MAX_NUM_ROW) {
994 EIGEN_IF_CONSTEXPR (!isBRowMajor) {
995 int64_t indA_i = isFWDSolve ? i : M - 1 - i;
996 int64_t indB_i = isFWDSolve ? i : M - (i + EIGEN_AVX_MAX_NUM_ROW);
997 int64_t offB_1 = isFWDSolve ? offsetBTemp : sizeBTemp - EIGEN_AVX_MAX_NUM_ROW * LDT - offsetBTemp;
998 int64_t offB_2 = isFWDSolve ? offsetBTemp : sizeBTemp - LDT - offsetBTemp;
999 // Copy values from B to B_temp.
1000 copyBToRowMajor<Scalar, true, false>(B_arr + indB_i + k * LDB, LDB, bK, B_temp + offB_1, LDT);
1001 // Triangular solve with a small block of A and long horizontal blocks of B (or B_temp if B col-major)
1002 triSolveKernelLxK<Scalar, isARowMajor, isFWDSolve, isUnitDiag>(
1003 &A_arr[idA<isARowMajor>(indA_i, indA_i, LDA)], B_temp + offB_2, EIGEN_AVX_MAX_NUM_ROW, bkL, LDA, LDT);
1004 // Copy values from B_temp back to B. B_temp will be reused in gemm call below.
1005 copyBToRowMajor<Scalar, false, false>(B_arr + indB_i + k * LDB, LDB, bK, B_temp + offB_1, LDT);
1006
1007 offsetBTemp += EIGEN_AVX_MAX_NUM_ROW * LDT;
1008 } else {
1009 int64_t ind = isFWDSolve ? i : M - 1 - i;
1010 triSolveKernelLxK<Scalar, isARowMajor, isFWDSolve, isUnitDiag>(
1011 &A_arr[idA<isARowMajor>(ind, ind, LDA)], B_arr + k + ind * LDB, EIGEN_AVX_MAX_NUM_ROW, bK, LDA, LDB);
1012 }
1013 if (i + EIGEN_AVX_MAX_NUM_ROW < M_) {
1026 EIGEN_IF_CONSTEXPR (isBRowMajor) {
1027 int64_t indA_i = isFWDSolve ? i + EIGEN_AVX_MAX_NUM_ROW : M - (i + 2 * EIGEN_AVX_MAX_NUM_ROW);
1028 int64_t indA_j = isFWDSolve ? 0 : M - (i + EIGEN_AVX_MAX_NUM_ROW);
1029 int64_t indB_i = isFWDSolve ? 0 : M - (i + EIGEN_AVX_MAX_NUM_ROW);
1030 int64_t indB_i2 = isFWDSolve ? i + EIGEN_AVX_MAX_NUM_ROW : M - (i + 2 * EIGEN_AVX_MAX_NUM_ROW);
1031 gemmKernel<Scalar, isARowMajor, isBRowMajor, false, false>(
1032 &A_arr[idA<isARowMajor>(indA_i, indA_j, LDA)], B_arr + k + indB_i * LDB, B_arr + k + indB_i2 * LDB,
1033 EIGEN_AVX_MAX_NUM_ROW, bK, i + EIGEN_AVX_MAX_NUM_ROW, LDA, LDB, LDB);
1034 } else {
1035 if (offsetBTemp + EIGEN_AVX_MAX_NUM_ROW * LDT > sizeBTemp) {
1044 int64_t indA_i = isFWDSolve ? i + EIGEN_AVX_MAX_NUM_ROW : 0;
1045 int64_t indA_j = isFWDSolve ? gemmOff : M - (i + EIGEN_AVX_MAX_NUM_ROW);
1046 int64_t indB_i = isFWDSolve ? i + EIGEN_AVX_MAX_NUM_ROW : 0;
1047 int64_t offB_1 = isFWDSolve ? 0 : sizeBTemp - offsetBTemp;
1048 gemmKernel<Scalar, isARowMajor, isBRowMajor, false, false>(
1049 &A_arr[idA<isARowMajor>(indA_i, indA_j, LDA)], B_temp + offB_1, B_arr + indB_i + (k)*LDB,
1050 M - (i + EIGEN_AVX_MAX_NUM_ROW), bK, i + EIGEN_AVX_MAX_NUM_ROW - gemmOff, LDA, LDT, LDB);
1051 offsetBTemp = 0;
1052 gemmOff = i + EIGEN_AVX_MAX_NUM_ROW;
1053 } else {
1057 int64_t indA_i = isFWDSolve ? i + EIGEN_AVX_MAX_NUM_ROW : M - (i + 2 * EIGEN_AVX_MAX_NUM_ROW);
1058 int64_t indA_j = isFWDSolve ? gemmOff : M - (i + EIGEN_AVX_MAX_NUM_ROW);
1059 int64_t indB_i = isFWDSolve ? i + EIGEN_AVX_MAX_NUM_ROW : M - (i + 2 * EIGEN_AVX_MAX_NUM_ROW);
1060 int64_t offB_1 = isFWDSolve ? 0 : sizeBTemp - offsetBTemp;
1061 gemmKernel<Scalar, isARowMajor, isBRowMajor, false, false>(
1062 &A_arr[idA<isARowMajor>(indA_i, indA_j, LDA)], B_temp + offB_1, B_arr + indB_i + (k)*LDB,
1063 EIGEN_AVX_MAX_NUM_ROW, bK, i + EIGEN_AVX_MAX_NUM_ROW - gemmOff, LDA, LDT, LDB);
1064 }
1065 }
1066 }
1067 }
1068 // Handle M remainder..
1069 int64_t bM = M - M_;
1070 if (bM > 0) {
1071 if (M_ > 0) {
1072 EIGEN_IF_CONSTEXPR (isBRowMajor) {
1073 int64_t indA_i = isFWDSolve ? M_ : 0;
1074 int64_t indA_j = isFWDSolve ? 0 : bM;
1075 int64_t indB_i = isFWDSolve ? 0 : bM;
1076 int64_t indB_i2 = isFWDSolve ? M_ : 0;
1077 gemmKernel<Scalar, isARowMajor, isBRowMajor, false, false>(
1078 &A_arr[idA<isARowMajor>(indA_i, indA_j, LDA)], B_arr + k + indB_i * LDB, B_arr + k + indB_i2 * LDB, bM,
1079 bK, M_, LDA, LDB, LDB);
1080 } else {
1081 int64_t indA_i = isFWDSolve ? M_ : 0;
1082 int64_t indA_j = isFWDSolve ? gemmOff : bM;
1083 int64_t indB_i = isFWDSolve ? M_ : 0;
1084 int64_t offB_1 = isFWDSolve ? 0 : sizeBTemp - offsetBTemp;
1085 gemmKernel<Scalar, isARowMajor, isBRowMajor, false, false>(&A_arr[idA<isARowMajor>(indA_i, indA_j, LDA)],
1086 B_temp + offB_1, B_arr + indB_i + (k)*LDB, bM, bK,
1087 M_ - gemmOff, LDA, LDT, LDB);
1088 }
1089 }
1090 EIGEN_IF_CONSTEXPR (!isBRowMajor) {
1091 int64_t indA_i = isFWDSolve ? M_ : M - 1 - M_;
1092 int64_t indB_i = isFWDSolve ? M_ : 0;
1093 int64_t offB_1 = isFWDSolve ? 0 : (bM - 1) * bkL;
1094 copyBToRowMajor<Scalar, true, true>(B_arr + indB_i + k * LDB, LDB, bK, B_temp, bkL, bM);
1095 triSolveKernelLxK<Scalar, isARowMajor, isFWDSolve, isUnitDiag>(&A_arr[idA<isARowMajor>(indA_i, indA_i, LDA)],
1096 B_temp + offB_1, bM, bkL, LDA, bkL);
1097 copyBToRowMajor<Scalar, false, true>(B_arr + indB_i + k * LDB, LDB, bK, B_temp, bkL, bM);
1098 } else {
1099 int64_t ind = isFWDSolve ? M_ : M - 1 - M_;
1100 triSolveKernelLxK<Scalar, isARowMajor, isFWDSolve, isUnitDiag>(&A_arr[idA<isARowMajor>(ind, ind, LDA)],
1101 B_arr + k + ind * LDB, bM, bK, LDA, LDB);
1102 }
1103 }
1104 }
1105
1106 EIGEN_IF_CONSTEXPR (!isBRowMajor) handmade_aligned_free(B_temp);
1107}
1108
1109// Template specializations of trsmKernelL/R for float/double and inner strides of 1.
1110#if (EIGEN_USE_AVX512_TRSM_R_KERNELS)
1111template <typename Scalar, typename Index, int Mode, bool Conjugate, int TriStorageOrder, int OtherInnerStride,
1112 bool Specialized>
1113struct trsmKernelR;
1114
1115template <typename Index, int Mode, int TriStorageOrder>
1116struct trsmKernelR<float, Index, Mode, false, TriStorageOrder, 1, true> {
1117 static void kernel(Index size, Index otherSize, const float* _tri, Index triStride, float* _other, Index otherIncr,
1118 Index otherStride);
1119};
1120
1121template <typename Index, int Mode, int TriStorageOrder>
1122struct trsmKernelR<double, Index, Mode, false, TriStorageOrder, 1, true> {
1123 static void kernel(Index size, Index otherSize, const double* _tri, Index triStride, double* _other, Index otherIncr,
1124 Index otherStride);
1125};
1126
1127template <typename Index, int Mode, int TriStorageOrder>
1128EIGEN_DONT_INLINE void trsmKernelR<float, Index, Mode, false, TriStorageOrder, 1, true>::kernel(
1129 Index size, Index otherSize, const float* _tri, Index triStride, float* _other, Index otherIncr,
1130 Index otherStride) {
1131 EIGEN_UNUSED_VARIABLE(otherIncr);
1132#ifdef EIGEN_RUNTIME_NO_MALLOC
1133 if (!is_malloc_allowed()) {
1134 trsmKernelR<float, Index, Mode, false, TriStorageOrder, 1, /*Specialized=*/false>::kernel(
1135 size, otherSize, _tri, triStride, _other, otherIncr, otherStride);
1136 return;
1137 }
1138#endif
1139 triSolve<float, TriStorageOrder != RowMajor, true, (Mode & Lower) != Lower, (Mode & UnitDiag) != 0>(
1140 const_cast<float*>(_tri), _other, size, otherSize, triStride, otherStride);
1141}
1142
1143template <typename Index, int Mode, int TriStorageOrder>
1144EIGEN_DONT_INLINE void trsmKernelR<double, Index, Mode, false, TriStorageOrder, 1, true>::kernel(
1145 Index size, Index otherSize, const double* _tri, Index triStride, double* _other, Index otherIncr,
1146 Index otherStride) {
1147 EIGEN_UNUSED_VARIABLE(otherIncr);
1148#ifdef EIGEN_RUNTIME_NO_MALLOC
1149 if (!is_malloc_allowed()) {
1150 trsmKernelR<double, Index, Mode, false, TriStorageOrder, 1, /*Specialized=*/false>::kernel(
1151 size, otherSize, _tri, triStride, _other, otherIncr, otherStride);
1152 return;
1153 }
1154#endif
1155 triSolve<double, TriStorageOrder != RowMajor, true, (Mode & Lower) != Lower, (Mode & UnitDiag) != 0>(
1156 const_cast<double*>(_tri), _other, size, otherSize, triStride, otherStride);
1157}
1158#endif // (EIGEN_USE_AVX512_TRSM_R_KERNELS)
1159
1160// These trsm kernels require temporary memory allocation
1161#if (EIGEN_USE_AVX512_TRSM_L_KERNELS)
1162template <typename Scalar, typename Index, int Mode, bool Conjugate, int TriStorageOrder, int OtherInnerStride,
1163 bool Specialized = true>
1164struct trsmKernelL;
1165
1166template <typename Index, int Mode, int TriStorageOrder>
1167struct trsmKernelL<float, Index, Mode, false, TriStorageOrder, 1, true> {
1168 static void kernel(Index size, Index otherSize, const float* _tri, Index triStride, float* _other, Index otherIncr,
1169 Index otherStride);
1170};
1171
1172template <typename Index, int Mode, int TriStorageOrder>
1173struct trsmKernelL<double, Index, Mode, false, TriStorageOrder, 1, true> {
1174 static void kernel(Index size, Index otherSize, const double* _tri, Index triStride, double* _other, Index otherIncr,
1175 Index otherStride);
1176};
1177
1178template <typename Index, int Mode, int TriStorageOrder>
1179EIGEN_DONT_INLINE void trsmKernelL<float, Index, Mode, false, TriStorageOrder, 1, true>::kernel(
1180 Index size, Index otherSize, const float* _tri, Index triStride, float* _other, Index otherIncr,
1181 Index otherStride) {
1182 EIGEN_UNUSED_VARIABLE(otherIncr);
1183#ifdef EIGEN_RUNTIME_NO_MALLOC
1184 if (!is_malloc_allowed()) {
1185 // The unspecialized kernel takes an upper-triangular panel by its bottom-right element and
1186 // indexes it backwards, while triangular_solve_matrix() hands these specializations the
1187 // panel's top-left element. Move the origin before delegating.
1188 const Index shift = ((Mode & Lower) == Lower) ? Index(0) : size - 1;
1189 trsmKernelL<float, Index, Mode, false, TriStorageOrder, 1, /*Specialized=*/false>::kernel(
1190 size, otherSize, _tri + shift + shift * triStride, triStride, _other + shift, otherIncr, otherStride);
1191 return;
1192 }
1193#endif
1194 triSolve<float, TriStorageOrder == RowMajor, false, (Mode & Lower) == Lower, (Mode & UnitDiag) != 0>(
1195 const_cast<float*>(_tri), _other, size, otherSize, triStride, otherStride);
1196}
1197
1198template <typename Index, int Mode, int TriStorageOrder>
1199EIGEN_DONT_INLINE void trsmKernelL<double, Index, Mode, false, TriStorageOrder, 1, true>::kernel(
1200 Index size, Index otherSize, const double* _tri, Index triStride, double* _other, Index otherIncr,
1201 Index otherStride) {
1202 EIGEN_UNUSED_VARIABLE(otherIncr);
1203#ifdef EIGEN_RUNTIME_NO_MALLOC
1204 if (!is_malloc_allowed()) {
1205 // The unspecialized kernel takes an upper-triangular panel by its bottom-right element and
1206 // indexes it backwards, while triangular_solve_matrix() hands these specializations the
1207 // panel's top-left element. Move the origin before delegating.
1208 const Index shift = ((Mode & Lower) == Lower) ? Index(0) : size - 1;
1209 trsmKernelL<double, Index, Mode, false, TriStorageOrder, 1, /*Specialized=*/false>::kernel(
1210 size, otherSize, _tri + shift + shift * triStride, triStride, _other + shift, otherIncr, otherStride);
1211 return;
1212 }
1213#endif
1214 triSolve<double, TriStorageOrder == RowMajor, false, (Mode & Lower) == Lower, (Mode & UnitDiag) != 0>(
1215 const_cast<double*>(_tri), _other, size, otherSize, triStride, otherStride);
1216}
1217#endif // EIGEN_USE_AVX512_TRSM_L_KERNELS
1218
1219#endif // EIGEN_USE_AVX512_TRSM_KERNELS
1220
1221} // namespace internal
1222} // namespace Eigen
1223#endif // EIGEN_CORE_ARCH_AVX512_TRSM_KERNEL_H
@ Lower
Definition Constants.h:212