Eigen  5.0.1
 
Loading...
Searching...
No Matches
GeneralMatrixMatrix.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2008-2009 Gael Guennebaud <gael.guennebaud@inria.fr>
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_GENERAL_MATRIX_MATRIX_H
12#define EIGEN_GENERAL_MATRIX_MATRIX_H
13
14// IWYU pragma: private
15#include "../InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20
21template <typename LhsScalar_, typename RhsScalar_>
22class level3_blocking;
23
24// LHS-first loop order: mc -> kc -> nc. This is Eigen's usual sequential
25// schedule, with a fast path that can pack RHS once for tall skinny blocking.
26struct gemm_pack_lhs_first_loop_policy {
27 template <typename Index, typename LhsScalar, typename RhsScalar, typename ResScalar, typename LhsMapper,
28 typename RhsMapper, typename ResMapper, typename PackLhs, typename PackRhs, typename Gebp>
29 static EIGEN_STRONG_INLINE void run(Index rows, Index cols, Index depth, Index kc, Index mc, Index nc,
30 const LhsMapper& lhs, const RhsMapper& rhs, ResMapper& res, PackLhs& pack_lhs,
31 PackRhs& pack_rhs, Gebp& gebp, LhsScalar* blockA, RhsScalar* blockB,
32 ResScalar alpha) {
33 const bool pack_rhs_once = mc != rows && kc == depth && nc == cols;
34
35 // For each horizontal panel of the rhs, and corresponding panel of the lhs...
36 for (Index i2 = 0; i2 < rows; i2 += mc) {
37 const Index actual_mc = (std::min)(i2 + mc, rows) - i2;
38
39 for (Index k2 = 0; k2 < depth; k2 += kc) {
40 const Index actual_kc = (std::min)(k2 + kc, depth) - k2;
41
42 // OK, here we have selected one horizontal panel of rhs and one vertical panel of lhs.
43 // => Pack lhs's panel into a sequential chunk of memory (L2/L3 caching)
44 // Note that this panel will be read as many times as the number of blocks in the rhs's
45 // horizontal panel which is, in practice, a very low number.
46 pack_lhs(blockA, lhs.getSubMapper(i2, k2), actual_kc, actual_mc);
47
48 // For each kc x nc block of the rhs's horizontal panel...
49 for (Index j2 = 0; j2 < cols; j2 += nc) {
50 const Index actual_nc = (std::min)(j2 + nc, cols) - j2;
51
52 // We pack the rhs's block into a sequential chunk of memory (L2 caching)
53 // Note that this block will be read a very high number of times, which is equal to the number of
54 // micro horizontal panel of the large rhs's panel (e.g., rows/12 times).
55 if ((!pack_rhs_once) || i2 == 0) pack_rhs(blockB, rhs.getSubMapper(k2, j2), actual_kc, actual_nc);
56
57 // Everything is packed, we can now call the panel * block kernel:
58 gebp(res.getSubMapper(i2, j2), blockA, blockB, actual_mc, actual_kc, actual_nc, alpha);
59 }
60 }
61 }
62 }
63};
64
65#ifdef EIGEN_VECTORIZE_SME
66// Defined in arch/SME/GeneralBlockPanelKernel.h, included after this header:
67// whether the SME kernel can read this ColMajor LHS block straight from its
68// source instead of a packed panel.
69template <typename Scalar, typename Index>
70bool sme_direct_lhs_ok(Index lhsStride, Index rows, Index depth, Index cols);
71// True for the unit-stride ColMajor mapper the GEMM driver hands the packers.
72template <typename Mapper>
73struct sme_direct_lhs_mapper : std::false_type {};
74template <typename Scalar, typename Index>
75struct sme_direct_lhs_mapper<const_blas_data_mapper<Scalar, Index, ColMajor>>
76 : bool_constant<!NumTraits<Scalar>::IsComplex> {};
77// Runs the block on the in-place LHS when the kernel can take it; the false
78// overload keeps the call out of every other instantiation.
79template <typename Gebp, typename ResMapper, typename LhsMapper, typename Scalar, typename ResScalar, typename Index>
80EIGEN_ALWAYS_INLINE bool sme_run_direct_lhs(std::true_type, Gebp& gebp, const ResMapper& res, const LhsMapper& lhs,
81 Index i2, Index k2, const Scalar* blockB, Index mc, Index kc, Index nc,
82 ResScalar alpha) {
83 if (!sme_direct_lhs_ok<Scalar>(lhs.stride(), mc, kc, nc)) return false;
84 gebp.run_direct_lhs(res, &lhs(i2, k2), lhs.stride(), blockB, mc, kc, nc, alpha);
85 return true;
86}
87template <typename Scalar, int LhsOrder, int RhsOrder, typename Index>
88bool sme_tiny_gemm(Index rows, Index cols, Index depth, const Scalar* lhs, Index lhsStride, const Scalar* rhs,
89 Index rhsStride, Scalar* res, Index resIncr, Index resStride, Scalar alpha);
90// Where the tiny-result kernel beats both the coeff-based product and the packed paths, as tuned on Apple M4: short
91// depths keep the coeff-based product, and a full 2 * ps x 8 result over a long depth the SME kernel. The NEON
92// small-block switches pin the paths they select, so they turn it off.
93template <typename Scalar, typename Index>
94EIGEN_ALWAYS_INLINE bool sme_tiny_gemm_wins(Index rows, Index cols, Index depth) {
95#if defined(EIGEN_SME_NO_NEON_SMALL_BLOCKS) || defined(EIGEN_SME_FORCE_NEON_SMALL_BLOCKS) || defined(EIGEN_USE_BLAS)
96 EIGEN_UNUSED_VARIABLE(rows);
97 EIGEN_UNUSED_VARIABLE(cols);
98 EIGEN_UNUSED_VARIABLE(depth);
99 return false;
100#else
101 const Index ps = Index(16 / sizeof(Scalar)); // scalars per NEON vector
102 if (rows < 2 || cols < 2 || rows > 2 * ps || cols > 8) return false;
103 if (depth < 16 && (depth < 6 || rows * cols < 24)) return false;
104 return !(cols == 8 && rows > ps && depth >= 2048);
105#endif
106}
107// Real float and double on the SME kernel take the tiny-result kernel; the false overload keeps it out of every
108// other pair, including double without FEAT_SME_F64F64.
109template <typename LhsScalar, typename RhsScalar>
110struct sme_tiny_gemm_pair
111 : bool_constant<std::is_same<LhsScalar, RhsScalar>::value &&
112 (std::is_same<LhsScalar, float>::value || std::is_same<LhsScalar, double>::value) &&
113 sme_has_gebp_kernel<LhsScalar, RhsScalar>::value> {};
114template <int LhsOrder, int RhsOrder, typename Scalar, typename Index>
115EIGEN_ALWAYS_INLINE bool sme_run_tiny_gemm(std::true_type, Index rows, Index cols, Index depth, const Scalar* lhs,
116 Index lhsStride, const Scalar* rhs, Index rhsStride, Scalar* res,
117 Index resIncr, Index resStride, Scalar alpha) {
118 return sme_tiny_gemm<Scalar, LhsOrder, RhsOrder>(rows, cols, depth, lhs, lhsStride, rhs, rhsStride, res, resIncr,
119 resStride, alpha);
120}
121template <int LhsOrder, int RhsOrder, typename LhsScalar, typename RhsScalar, typename ResScalar, typename Index>
122EIGEN_ALWAYS_INLINE bool sme_run_tiny_gemm(std::false_type, Index, Index, Index, const LhsScalar*, Index,
123 const RhsScalar*, Index, ResScalar*, Index, Index, ResScalar) {
124 return false;
125}
126template <typename Gebp, typename ResMapper, typename LhsMapper, typename Scalar, typename ResScalar, typename Index>
127EIGEN_ALWAYS_INLINE bool sme_run_direct_lhs(std::false_type, Gebp&, const ResMapper&, const LhsMapper&, Index, Index,
128 const Scalar*, Index, Index, Index, ResScalar) {
129 return false;
130}
131#endif
132
133// RHS-first loop order: nc -> kc -> mc. Used by SME to stream ColMajor result
134// stores through adjacent row panels; a ColMajor LHS whose MR-row slices are
135// already contiguous per depth step is read from its source (no LHS packing).
136struct gemm_pack_rhs_first_loop_policy {
137 template <typename Index, typename LhsScalar, typename RhsScalar, typename ResScalar, typename LhsMapper,
138 typename RhsMapper, typename ResMapper, typename PackLhs, typename PackRhs, typename Gebp>
139 static EIGEN_STRONG_INLINE void run(Index rows, Index cols, Index depth, Index kc, Index mc, Index nc,
140 const LhsMapper& lhs, const RhsMapper& rhs, ResMapper& res, PackLhs& pack_lhs,
141 PackRhs& pack_rhs, Gebp& gebp, LhsScalar* blockA, RhsScalar* blockB,
142 ResScalar alpha) {
143 // Mirror of pack_rhs_once: reuse one full LHS panel across column blocks. The in-place read is decided per column
144 // block, so a block that falls back to packing packs unless an earlier one already did.
145 const bool pack_lhs_once = nc != cols && kc == depth && mc == rows;
146 bool lhs_packed = false;
147
148 for (Index j2 = 0; j2 < cols; j2 += nc) {
149 const Index actual_nc = (std::min)(j2 + nc, cols) - j2;
150
151 for (Index k2 = 0; k2 < depth; k2 += kc) {
152 const Index actual_kc = (std::min)(k2 + kc, depth) - k2;
153
154 // Pack one nc-strip of RHS and reuse it across all row panels.
155 pack_rhs(blockB, rhs.getSubMapper(k2, j2), actual_kc, actual_nc);
156
157 for (Index i2 = 0; i2 < rows; i2 += mc) {
158 const Index actual_mc = (std::min)(i2 + mc, rows) - i2;
159#ifdef EIGEN_VECTORIZE_SME
160 if (sme_run_direct_lhs(bool_constant<sme_direct_lhs_mapper<LhsMapper>::value>(), gebp,
161 res.getSubMapper(i2, j2), lhs, i2, k2, blockB, actual_mc, actual_kc, actual_nc, alpha))
162 continue;
163#endif
164 if (!pack_lhs_once || !lhs_packed) {
165 pack_lhs(blockA, lhs.getSubMapper(i2, k2), actual_kc, actual_mc);
166 lhs_packed = true;
167 }
168 gebp(res.getSubMapper(i2, j2), blockA, blockB, actual_mc, actual_kc, actual_nc, alpha);
169 }
170 }
171 }
172 }
173};
174
175/* Specialization for a row-major destination matrix => simple transposition of the product */
176template <typename Index, typename LhsScalar, int LhsStorageOrder, bool ConjugateLhs, typename RhsScalar,
177 int RhsStorageOrder, bool ConjugateRhs, int ResInnerStride>
178struct general_matrix_matrix_product<Index, LhsScalar, LhsStorageOrder, ConjugateLhs, RhsScalar, RhsStorageOrder,
179 ConjugateRhs, RowMajor, ResInnerStride> {
180 using Traits = gebp_traits<RhsScalar, LhsScalar>;
181
182 using ResScalar = typename ScalarBinaryOpTraits<LhsScalar, RhsScalar>::ReturnType;
183 static EIGEN_STRONG_INLINE void run(Index rows, Index cols, Index depth, const LhsScalar* lhs, Index lhsStride,
184 const RhsScalar* rhs, Index rhsStride, ResScalar* res, Index resIncr,
185 Index resStride, ResScalar alpha, level3_blocking<RhsScalar, LhsScalar>& blocking,
186 GemmParallelInfo<Index>* info = 0) {
187 // transpose the product such that the result is column major
188 general_matrix_matrix_product<Index, RhsScalar, RhsStorageOrder == RowMajor ? ColMajor : RowMajor, ConjugateRhs,
189 LhsScalar, LhsStorageOrder == RowMajor ? ColMajor : RowMajor, ConjugateLhs, ColMajor,
190 ResInnerStride>::run(cols, rows, depth, rhs, rhsStride, lhs, lhsStride, res, resIncr,
191 resStride, alpha, blocking, info);
192 }
193};
194
195/* Specialization for a col-major destination matrix
196 * => Blocking algorithm following Goto's paper */
197template <typename Index, typename LhsScalar, int LhsStorageOrder, bool ConjugateLhs, typename RhsScalar,
198 int RhsStorageOrder, bool ConjugateRhs, int ResInnerStride>
199struct general_matrix_matrix_product<Index, LhsScalar, LhsStorageOrder, ConjugateLhs, RhsScalar, RhsStorageOrder,
200 ConjugateRhs, ColMajor, ResInnerStride> {
201 using Traits = gebp_traits<LhsScalar, RhsScalar>;
202
203 using ResScalar = typename ScalarBinaryOpTraits<LhsScalar, RhsScalar>::ReturnType;
204 static void run(Index rows, Index cols, Index depth, const LhsScalar* lhs_, Index lhsStride, const RhsScalar* rhs_,
205 Index rhsStride, ResScalar* res_, Index resIncr, Index resStride, ResScalar alpha,
206 level3_blocking<LhsScalar, RhsScalar>& blocking, GemmParallelInfo<Index>* info = 0) {
207 // BLAS contract: if alpha == 0, the result is unchanged (and lhs/rhs need not be read).
208 if (numext::is_exactly_zero(alpha)) return;
209#ifdef EIGEN_VECTORIZE_SME
210 // A parallel session needs every thread to take the packed path; scaleAndAddTo runs tiny results on one thread.
211 if (info == nullptr && sme_run_tiny_gemm<LhsStorageOrder, RhsStorageOrder>(
212 bool_constant<sme_tiny_gemm_pair<LhsScalar, RhsScalar>::value>(), rows, cols, depth,
213 lhs_, lhsStride, rhs_, rhsStride, res_, resIncr, resStride, alpha))
214 return;
215#endif
216
217 using LhsMapper = const_blas_data_mapper<LhsScalar, Index, LhsStorageOrder>;
218 using RhsMapper = const_blas_data_mapper<RhsScalar, Index, RhsStorageOrder>;
219 using ResMapper = blas_data_mapper<typename Traits::ResScalar, Index, ColMajor, Unaligned, ResInnerStride>;
220 LhsMapper lhs(lhs_, lhsStride);
221 RhsMapper rhs(rhs_, rhsStride);
222 ResMapper res(res_, resStride, resIncr);
223
224 Index kc = blocking.kc(); // cache block size along the K direction
225 Index mc = (std::min)(rows, blocking.mc()); // cache block size along the M direction
226 Index nc = (std::min)(cols, blocking.nc()); // cache block size along the N direction
227
228 gemm_pack_lhs<LhsScalar, Index, LhsMapper, Traits::mr, Traits::LhsProgress, typename Traits::LhsPacket4Packing,
229 LhsStorageOrder>
230 pack_lhs;
231 gemm_pack_rhs<RhsScalar, Index, RhsMapper, Traits::nr, RhsStorageOrder> pack_rhs;
232 gebp_kernel<LhsScalar, RhsScalar, Index, ResMapper, Traits::mr, Traits::nr, ConjugateLhs, ConjugateRhs> gebp;
233
234#if !defined(EIGEN_USE_BLAS) && (defined(EIGEN_HAS_OPENMP) || defined(EIGEN_GEMM_THREADPOOL))
235 if (info) {
236 // this is the parallel version!
237 int tid = info->logical_thread_id;
238 int threads = info->num_threads;
239
240 LhsScalar* blockA = blocking.blockA();
241 eigen_internal_assert(blockA != 0);
242
243 std::size_t sizeB = kc * nc;
244 ei_declare_aligned_stack_constructed_variable(RhsScalar, blockB, sizeB, 0);
245
246 // For each horizontal panel of the rhs, and corresponding vertical panel of the lhs...
247 for (Index k = 0; k < depth; k += kc) {
248 const Index actual_kc = (std::min)(k + kc, depth) - k; // => rows of B', and cols of the A'
249
250 // In order to reduce the chance that a thread has to wait for the other,
251 // let's start by packing B'.
252 pack_rhs(blockB, rhs.getSubMapper(k, 0), actual_kc, nc);
253
254 // Pack A_k to A' in a parallel fashion:
255 // each thread packs the sub block A_k,i to A'_i where i is the thread id.
256
257 // However, before copying to A'_i, we have to make sure that no other thread is still using it,
258 // i.e., we test that info->task_info[tid].users equals 0.
259 // Then, we set info->task_info[tid].users to the number of threads to mark that all other threads are going to
260 // use it.
261 while (info->task_info[tid].users != 0) {
262 std::this_thread::yield();
263 }
264 info->task_info[tid].users = threads;
265
266 pack_lhs(blockA + info->task_info[tid].lhs_start * actual_kc,
267 lhs.getSubMapper(info->task_info[tid].lhs_start, k), actual_kc, info->task_info[tid].lhs_length);
268
269 // Notify the other threads that the part A'_i is ready to go.
270 info->task_info[tid].sync = k;
271
272 // Computes C_i += A' * B' per A'_i
273 for (int shift = 0; shift < threads; ++shift) {
274 int i = (tid + shift) % threads;
275
276 // At this point we have to make sure that A'_i has been updated by the thread i,
277 // we use testAndSetOrdered to mimic a volatile access.
278 // However, no need to wait for the B' part which has been updated by the current thread!
279 if (shift > 0) {
280 while (info->task_info[i].sync != k) {
281 std::this_thread::yield();
282 }
283 }
284
285 gebp(res.getSubMapper(info->task_info[i].lhs_start, 0), blockA + info->task_info[i].lhs_start * actual_kc,
286 blockB, info->task_info[i].lhs_length, actual_kc, nc, alpha);
287 }
288
289 // Then keep going as usual with the remaining B'
290 for (Index j = nc; j < cols; j += nc) {
291 const Index actual_nc = (std::min)(j + nc, cols) - j;
292
293 // pack B_k,j to B'
294 pack_rhs(blockB, rhs.getSubMapper(k, j), actual_kc, actual_nc);
295
296 // C_j += A' * B'
297 gebp(res.getSubMapper(0, j), blockA, blockB, rows, actual_kc, actual_nc, alpha);
298 }
299
300 // Release all the sub blocks A'_i of A' for the current thread,
301 // i.e., we simply decrement the number of users by 1
302 for (Index i = 0; i < threads; ++i) info->task_info[i].users -= 1;
303 }
304 } else
305#endif // !defined(EIGEN_USE_BLAS) && (defined(EIGEN_HAS_OPENMP) || defined(EIGEN_GEMM_THREADPOOL))
306 {
307 EIGEN_UNUSED_VARIABLE(info);
308
309 // this is the sequential version!
310 std::size_t sizeA = kc * mc;
311 std::size_t sizeB = kc * nc;
312
313 ei_declare_aligned_stack_constructed_variable(LhsScalar, blockA, sizeA, blocking.blockA());
314 ei_declare_aligned_stack_constructed_variable(RhsScalar, blockB, sizeB, blocking.blockB());
315
316 // The SME kernel uses RHS-first order so consecutive gebp calls stream
317 // through adjacent row panels of a ColMajor result. Other kernels --
318 // including the scalar pairs SME does not specialize -- keep Eigen's
319 // default LHS-first order.
320#ifdef EIGEN_VECTORIZE_SME
321 using SequentialGemmLoop = std::conditional_t<sme_has_gebp_kernel<LhsScalar, RhsScalar>::value,
322 gemm_pack_rhs_first_loop_policy, gemm_pack_lhs_first_loop_policy>;
323#else
324 using SequentialGemmLoop = gemm_pack_lhs_first_loop_policy;
325#endif
326
327 SequentialGemmLoop::run(rows, cols, depth, kc, mc, nc, lhs, rhs, res, pack_lhs, pack_rhs, gebp, blockA, blockB,
328 alpha);
329 }
330 }
331};
332
333/*********************************************************************************
334 * Specialization of generic_product_impl for "large" GEMM, i.e.,
335 * implementation of the high level wrapper to general_matrix_matrix_product
336 **********************************************************************************/
337
338template <typename Scalar, typename Index, typename Gemm, typename Lhs, typename Rhs, typename Dest,
339 typename BlockingType>
340struct gemm_functor {
341 gemm_functor(const Lhs& lhs, const Rhs& rhs, Dest& dest, const Scalar& actualAlpha, BlockingType& blocking)
342 : m_lhs(lhs), m_rhs(rhs), m_dest(dest), m_actualAlpha(actualAlpha), m_blocking(blocking) {}
343
344 void initParallelSession(Index num_threads) const {
345 m_blocking.initParallel(m_lhs.rows(), m_rhs.cols(), m_lhs.cols(), num_threads);
346 m_blocking.allocateA();
347 }
348
349 // Whether the blocking holds no preallocated buffers, so that disjoint parts of the product can run concurrently.
350 bool ownsNoBuffers() const { return m_blocking.blockA() == nullptr && m_blocking.blockB() == nullptr; }
351
352 void operator()(Index row, Index rows, Index col = 0, Index cols = -1, GemmParallelInfo<Index>* info = 0) const {
353 if (cols == -1) cols = m_rhs.cols();
354
355 Gemm::run(rows, cols, m_lhs.cols(), &m_lhs.coeffRef(row, 0), m_lhs.outerStride(), &m_rhs.coeffRef(0, col),
356 m_rhs.outerStride(), (Scalar*)&(m_dest.coeffRef(row, col)), m_dest.innerStride(), m_dest.outerStride(),
357 m_actualAlpha, m_blocking, info);
358 }
359
360 using Traits = typename Gemm::Traits;
361
362 protected:
363 const Lhs& m_lhs;
364 const Rhs& m_rhs;
365 Dest& m_dest;
366 Scalar m_actualAlpha;
367 BlockingType& m_blocking;
368};
369
370template <int StorageOrder, typename LhsScalar, typename RhsScalar, int MaxRows, int MaxCols, int MaxDepth,
371 int KcFactor = 1, bool FiniteAtCompileTime = MaxRows != Dynamic && MaxCols != Dynamic && MaxDepth != Dynamic>
372class gemm_blocking_space;
373
374template <typename LhsScalar_, typename RhsScalar_>
375class level3_blocking {
376 using LhsScalar = LhsScalar_;
377 using RhsScalar = RhsScalar_;
378
379 protected:
380 LhsScalar* m_blockA = nullptr;
381 RhsScalar* m_blockB = nullptr;
382
383 Index m_mc = 0;
384 Index m_nc = 0;
385 Index m_kc = 0;
386
387 public:
388 level3_blocking() = default;
389
390 inline Index mc() const { return m_mc; }
391 inline Index nc() const { return m_nc; }
392 inline Index kc() const { return m_kc; }
393
394 inline LhsScalar* blockA() { return m_blockA; }
395 inline RhsScalar* blockB() { return m_blockB; }
396};
397
398template <int StorageOrder, typename LhsScalar_, typename RhsScalar_, int MaxRows, int MaxCols, int MaxDepth,
399 int KcFactor>
400class gemm_blocking_space<StorageOrder, LhsScalar_, RhsScalar_, MaxRows, MaxCols, MaxDepth, KcFactor,
401 true /* == FiniteAtCompileTime */>
402 : public level3_blocking<std::conditional_t<StorageOrder == RowMajor, RhsScalar_, LhsScalar_>,
403 std::conditional_t<StorageOrder == RowMajor, LhsScalar_, RhsScalar_>> {
404 enum {
405 Transpose = StorageOrder == RowMajor,
406 ActualRows = Transpose ? MaxCols : MaxRows,
407 ActualCols = Transpose ? MaxRows : MaxCols
408 };
409 using LhsScalar = std::conditional_t<Transpose, RhsScalar_, LhsScalar_>;
410 using RhsScalar = std::conditional_t<Transpose, LhsScalar_, RhsScalar_>;
411 enum { SizeA = ActualRows * MaxDepth, SizeB = ActualCols * MaxDepth };
412
413#if EIGEN_MAX_STATIC_ALIGN_BYTES >= EIGEN_DEFAULT_ALIGN_BYTES
414 EIGEN_ALIGN_MAX LhsScalar m_staticA[SizeA];
415 EIGEN_ALIGN_MAX RhsScalar m_staticB[SizeB];
416#else
417 EIGEN_ALIGN_MAX char m_staticA[SizeA * sizeof(LhsScalar) + EIGEN_DEFAULT_ALIGN_BYTES - 1];
418 EIGEN_ALIGN_MAX char m_staticB[SizeB * sizeof(RhsScalar) + EIGEN_DEFAULT_ALIGN_BYTES - 1];
419#endif
420
421 public:
422 gemm_blocking_space(Index /*rows*/, Index /*cols*/, Index /*depth*/, Index /*num_threads*/,
423 bool /*full_rows = false*/) {
424 this->m_mc = ActualRows;
425 this->m_nc = ActualCols;
426 this->m_kc = MaxDepth;
427#if EIGEN_MAX_STATIC_ALIGN_BYTES >= EIGEN_DEFAULT_ALIGN_BYTES
428 this->m_blockA = m_staticA;
429 this->m_blockB = m_staticB;
430#else
431 this->m_blockA = reinterpret_cast<LhsScalar*>((std::uintptr_t(m_staticA) + (EIGEN_DEFAULT_ALIGN_BYTES - 1)) &
432 ~std::size_t(EIGEN_DEFAULT_ALIGN_BYTES - 1));
433 this->m_blockB = reinterpret_cast<RhsScalar*>((std::uintptr_t(m_staticB) + (EIGEN_DEFAULT_ALIGN_BYTES - 1)) &
434 ~std::size_t(EIGEN_DEFAULT_ALIGN_BYTES - 1));
435#endif
436 }
437
438 void initParallel(Index, Index, Index, Index) {}
439
440 inline void allocateA() {}
441 inline void allocateB() {}
442 inline void allocateAll() {}
443};
444
445template <int StorageOrder, typename LhsScalar_, typename RhsScalar_, int MaxRows, int MaxCols, int MaxDepth,
446 int KcFactor>
447class gemm_blocking_space<StorageOrder, LhsScalar_, RhsScalar_, MaxRows, MaxCols, MaxDepth, KcFactor, false>
448 : public level3_blocking<std::conditional_t<StorageOrder == RowMajor, RhsScalar_, LhsScalar_>,
449 std::conditional_t<StorageOrder == RowMajor, LhsScalar_, RhsScalar_>> {
450 enum { Transpose = StorageOrder == RowMajor };
451 using LhsScalar = std::conditional_t<Transpose, RhsScalar_, LhsScalar_>;
452 using RhsScalar = std::conditional_t<Transpose, LhsScalar_, RhsScalar_>;
453
454 Index m_sizeA;
455 Index m_sizeB;
456
457 public:
458 gemm_blocking_space(Index rows, Index cols, Index depth, Index num_threads, bool l3_blocking) {
459 this->m_mc = Transpose ? cols : rows;
460 this->m_nc = Transpose ? rows : cols;
461 this->m_kc = depth;
462
463 if (l3_blocking) {
464 computeProductBlockingSizes<LhsScalar, RhsScalar, KcFactor>(this->m_kc, this->m_mc, this->m_nc, num_threads);
465 } else // no l3 blocking
466 {
467 Index n = this->m_nc;
468 computeProductBlockingSizes<LhsScalar, RhsScalar, KcFactor>(this->m_kc, this->m_mc, n, num_threads);
469 }
470
471 m_sizeA = this->m_mc * this->m_kc;
472 m_sizeB = this->m_kc * this->m_nc;
473 }
474
475 void initParallel(Index rows, Index cols, Index depth, Index num_threads) {
476 this->m_mc = Transpose ? cols : rows;
477 this->m_nc = Transpose ? rows : cols;
478 this->m_kc = depth;
479
480 eigen_internal_assert(this->m_blockA == 0 && this->m_blockB == 0);
481 Index m = this->m_mc;
482 computeProductBlockingSizes<LhsScalar, RhsScalar, KcFactor>(this->m_kc, m, this->m_nc, num_threads);
483 m_sizeA = this->m_mc * this->m_kc;
484 m_sizeB = this->m_kc * this->m_nc;
485 }
486
487 // The blocking buffers get the temporaries' alignment (64 bytes in SME builds, EIGEN_STACK_ALIGN_BYTES).
488 void allocateA() {
489 if (this->m_blockA == 0) this->m_blockA = scratch_new<LhsScalar>(m_sizeA);
490 }
491
492 void allocateB() {
493 if (this->m_blockB == 0) this->m_blockB = scratch_new<RhsScalar>(m_sizeB);
494 }
495
496 void allocateAll() {
497 allocateA();
498 allocateB();
499 }
500
501 ~gemm_blocking_space() {
502 scratch_delete(this->m_blockA, m_sizeA);
503 scratch_delete(this->m_blockB, m_sizeB);
504 }
505};
506
507} // end namespace internal
508
509namespace internal {
510
511template <typename Lhs, typename Rhs>
512struct generic_product_impl<Lhs, Rhs, DenseShape, DenseShape, GemmProduct>
513 : generic_product_impl_base<Lhs, Rhs, generic_product_impl<Lhs, Rhs, DenseShape, DenseShape, GemmProduct>> {
514 using Scalar = typename Product<Lhs, Rhs>::Scalar;
515 using LhsScalar = typename Lhs::Scalar;
516 using RhsScalar = typename Rhs::Scalar;
517
518 using LhsBlasTraits = internal::blas_traits<Lhs>;
519 using ActualLhsType = typename LhsBlasTraits::DirectLinearAccessType;
520 using ActualLhsTypeCleaned = internal::remove_all_t<ActualLhsType>;
521
522 using RhsBlasTraits = internal::blas_traits<Rhs>;
523 using ActualRhsType = typename RhsBlasTraits::DirectLinearAccessType;
524 using ActualRhsTypeCleaned = internal::remove_all_t<ActualRhsType>;
525
526 enum { MaxDepthAtCompileTime = min_size_prefer_fixed(Lhs::MaxColsAtCompileTime, Rhs::MaxRowsAtCompileTime) };
527
528 using lazyproduct = generic_product_impl<Lhs, Rhs, DenseShape, DenseShape, CoeffBasedProductMode>;
529
530 // The runtime-size heuristic comes from bug 404 and was tuned with a
531 // helper program on Haswell. The threshold belongs to the kernel the GEMM
532 // path would select, not to the path itself, and for the SME kernel to the
533 // scalar type as well: it needs a larger product before it beats the
534 // coeff-based one, by an amount that falls as the scalar widens (see
535 // GeneralProduct.h). The rhs.rows() > 0 guard preserves the historical
536 // empty-product path through scaleAndAddTo().
537 static constexpr int kCoeffBasedThreshold =
538#ifdef EIGEN_VECTORIZE_SME
539 sme_has_gebp_kernel<LhsScalar, RhsScalar>::value ? sme_gemm_to_coeffbased_threshold<Scalar>::value :
540#endif
541 EIGEN_GEMM_TO_COEFFBASED_THRESHOLD;
542
543#ifdef EIGEN_VECTORIZE_SME
544 // Second bound for the SME kernel only: a small output over a long depth.
545 // The sum above grows with the depth and so never catches it, while the ZA
546 // grid this kernel fills is sized by the output (see GeneralProduct.h).
547 // Vector shapes are excluded -- scaleAndAddTo() routes those to GEMV, which
548 // is not the path being compared here.
549 static constexpr Index kCoeffBasedOutputArea = sme_has_gebp_kernel<LhsScalar, RhsScalar>::value
550 ? Index(EIGEN_SME_GEMM_TO_COEFFBASED_OUTPUT_AREA_THRESHOLD(Scalar))
551 : Index(0);
552
553 template <typename Dst>
554 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool outputAreaBelowThreshold(const Dst& dst) {
555 // rows * cols is dst.size(), which fits Index by construction.
556 return dst.rows() > 1 && dst.cols() > 1 && dst.size() <= kCoeffBasedOutputArea;
557 }
558#endif
559
560#ifdef EIGEN_VECTORIZE_SME
561 // Whether the tiny-result kernel takes a product the coeff-based one would; the GEMM driver sees a RowMajor
562 // result transposed.
563 template <typename Dst>
564 static EIGEN_STRONG_INLINE bool tinyKernelWins(const Dst& dst, const Rhs& rhs) {
565 constexpr bool row_major = (Dst::Flags & RowMajorBit) != 0;
566 return sme_tiny_gemm_pair<LhsScalar, RhsScalar>::value &&
567 sme_tiny_gemm_wins<Scalar>(row_major ? dst.cols() : dst.rows(), row_major ? dst.rows() : dst.cols(),
568 rhs.rows());
569 }
570#endif
571
572 template <typename Dst>
573 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool useRuntimeCoeffBasedProduct(const Dst& dst, const Rhs& rhs) {
574 if (rhs.rows() <= 0) return false;
575#ifdef EIGEN_VECTORIZE_SME
576 if ((rhs.rows() + dst.rows() + dst.cols()) < kCoeffBasedThreshold || outputAreaBelowThreshold(dst))
577 return !tinyKernelWins(dst, rhs);
578#else
579 if ((rhs.rows() + dst.rows() + dst.cols()) < kCoeffBasedThreshold) return true;
580#endif
581 return false;
582 }
583
584 // BLAS contract: a zero scalar factor leaves the destination unchanged and
585 // neither operand need be read, so that a non-finite coefficient cannot taint
586 // the result through 0 * Inf. general_matrix_matrix_product::run enforces it
587 // for the GEMM path, but the coeff-based path below has no such exit and
588 // would evaluate the product, so the factor is tested before the dispatch
589 // rather than inside either kernel.
590 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool scalarFactorIsZero(const Lhs& lhs, const Rhs& rhs) {
591 return numext::is_exactly_zero(combine_scalar_factors<Scalar>(lhs, rhs));
592 }
593
594 template <typename Dst>
595 static void evalTo(Dst& dst, const Lhs& lhs, const Rhs& rhs) {
596 if (scalarFactorIsZero(lhs, rhs)) {
597 dst.setZero();
598 return;
599 }
600 if (useRuntimeCoeffBasedProduct(dst, rhs))
601 lazyproduct::eval_dynamic(dst, lhs, rhs, internal::assign_op<typename Dst::Scalar, Scalar>());
602 else {
603 dst.setZero();
604 scaleAndAddTo(dst, lhs, rhs, Scalar(1));
605 }
606 }
607
608 template <typename Dst>
609 static void addTo(Dst& dst, const Lhs& lhs, const Rhs& rhs) {
610 if (scalarFactorIsZero(lhs, rhs)) return;
611 if (useRuntimeCoeffBasedProduct(dst, rhs))
612 lazyproduct::eval_dynamic(dst, lhs, rhs, internal::add_assign_op<typename Dst::Scalar, Scalar>());
613 else
614 scaleAndAddTo(dst, lhs, rhs, Scalar(1));
615 }
616
617 template <typename Dst>
618 static void subTo(Dst& dst, const Lhs& lhs, const Rhs& rhs) {
619 if (scalarFactorIsZero(lhs, rhs)) return;
620 if (useRuntimeCoeffBasedProduct(dst, rhs))
621 lazyproduct::eval_dynamic(dst, lhs, rhs, internal::sub_assign_op<typename Dst::Scalar, Scalar>());
622 else
623 scaleAndAddTo(dst, lhs, rhs, Scalar(-1));
624 }
625
626 template <typename Dest>
627 static void scaleAndAddTo(Dest& dst, const Lhs& a_lhs, const Rhs& a_rhs, const Scalar& alpha) {
628 eigen_assert(dst.rows() == a_lhs.rows() && dst.cols() == a_rhs.cols());
629 if (a_lhs.cols() == 0 || a_lhs.rows() == 0 || a_rhs.cols() == 0) return;
630
631 if (dst.cols() == 1) {
632 // Fallback to GEMV if either the lhs or rhs is a runtime vector
633 typename Dest::ColXpr dst_vec(dst.col(0));
634 return internal::generic_product_impl<Lhs, typename Rhs::ConstColXpr, DenseShape, DenseShape,
635 GemvProduct>::scaleAndAddTo(dst_vec, a_lhs, a_rhs.col(0), alpha);
636 } else if (dst.rows() == 1) {
637 // Fallback to GEMV if either the lhs or rhs is a runtime vector
638 typename Dest::RowXpr dst_vec(dst.row(0));
639 return internal::generic_product_impl<typename Lhs::ConstRowXpr, Rhs, DenseShape, DenseShape,
640 GemvProduct>::scaleAndAddTo(dst_vec, a_lhs.row(0), a_rhs, alpha);
641 }
642
643 add_const_on_value_type_t<ActualLhsType> lhs = LhsBlasTraits::extract(a_lhs);
644 add_const_on_value_type_t<ActualRhsType> rhs = RhsBlasTraits::extract(a_rhs);
645
646 Scalar actualAlpha = combine_scalar_factors(alpha, a_lhs, a_rhs);
647
648 using BlockingType =
649 internal::gemm_blocking_space<(Dest::Flags & RowMajorBit) ? RowMajor : ColMajor, LhsScalar, RhsScalar,
650 Dest::MaxRowsAtCompileTime, Dest::MaxColsAtCompileTime, MaxDepthAtCompileTime>;
651
652 using GemmFunctor = internal::gemm_functor<
653 Scalar, Index,
654 internal::general_matrix_matrix_product<
655 Index, LhsScalar, (ActualLhsTypeCleaned::Flags & RowMajorBit) ? RowMajor : ColMajor,
656 bool(LhsBlasTraits::NeedToConjugate), RhsScalar,
657 (ActualRhsTypeCleaned::Flags & RowMajorBit) ? RowMajor : ColMajor, bool(RhsBlasTraits::NeedToConjugate),
658 (Dest::Flags & RowMajorBit) ? RowMajor : ColMajor, Dest::InnerStrideAtCompileTime>,
659 ActualLhsTypeCleaned, ActualRhsTypeCleaned, Dest, BlockingType>;
660
661 BlockingType blocking(dst.rows(), dst.cols(), lhs.cols(), 1, true);
662#ifdef EIGEN_VECTORIZE_SME
663 // A result the tiny-result kernel takes is never worth a parallel session.
664 if (tinyKernelWins(dst, a_rhs))
665 return internal::parallelize_gemm<false>(GemmFunctor(lhs, rhs, dst, actualAlpha, blocking), a_lhs.rows(),
666 a_rhs.cols(), a_lhs.cols(), Dest::Flags & RowMajorBit);
667#endif
668 internal::parallelize_gemm<(Dest::MaxRowsAtCompileTime > 32 || Dest::MaxRowsAtCompileTime == Dynamic)>(
669 GemmFunctor(lhs, rhs, dst, actualAlpha, blocking), a_lhs.rows(), a_rhs.cols(), a_lhs.cols(),
670 Dest::Flags & RowMajorBit);
671 }
672};
673
674} // end namespace internal
675
676} // end namespace Eigen
677
678#endif // EIGEN_GENERAL_MATRIX_MATRIX_H
@ ColMajor
Definition Constants.h:319
@ RowMajor
Definition Constants.h:321
constexpr unsigned int RowMajorBit
Definition Constants.h:71