11#ifndef EIGEN_GENERAL_MATRIX_MATRIX_H
12#define EIGEN_GENERAL_MATRIX_MATRIX_H
15#include "../InternalHeaderCheck.h"
21template <
typename LhsScalar_,
typename RhsScalar_>
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,
33 const bool pack_rhs_once = mc != rows && kc == depth && nc == cols;
36 for (Index i2 = 0; i2 < rows; i2 += mc) {
37 const Index actual_mc = (std::min)(i2 + mc, rows) - i2;
39 for (Index k2 = 0; k2 < depth; k2 += kc) {
40 const Index actual_kc = (std::min)(k2 + kc, depth) - k2;
46 pack_lhs(blockA, lhs.getSubMapper(i2, k2), actual_kc, actual_mc);
49 for (Index j2 = 0; j2 < cols; j2 += nc) {
50 const Index actual_nc = (std::min)(j2 + nc, cols) - j2;
55 if ((!pack_rhs_once) || i2 == 0) pack_rhs(blockB, rhs.getSubMapper(k2, j2), actual_kc, actual_nc);
58 gebp(res.getSubMapper(i2, j2), blockA, blockB, actual_mc, actual_kc, actual_nc, alpha);
65#ifdef EIGEN_VECTORIZE_SME
69template <
typename Scalar,
typename Index>
70bool sme_direct_lhs_ok(Index lhsStride, Index rows, Index depth, Index cols);
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> {};
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,
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);
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);
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);
101 const Index ps = Index(16 /
sizeof(Scalar));
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);
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,
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) {
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) {
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,
145 const bool pack_lhs_once = nc != cols && kc == depth && mc == rows;
146 bool lhs_packed =
false;
148 for (Index j2 = 0; j2 < cols; j2 += nc) {
149 const Index actual_nc = (std::min)(j2 + nc, cols) - j2;
151 for (Index k2 = 0; k2 < depth; k2 += kc) {
152 const Index actual_kc = (std::min)(k2 + kc, depth) - k2;
155 pack_rhs(blockB, rhs.getSubMapper(k2, j2), actual_kc, actual_nc);
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))
164 if (!pack_lhs_once || !lhs_packed) {
165 pack_lhs(blockA, lhs.getSubMapper(i2, k2), actual_kc, actual_mc);
168 gebp(res.getSubMapper(i2, j2), blockA, blockB, actual_mc, actual_kc, actual_nc, alpha);
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>;
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) {
190 ResInnerStride>::run(cols, rows, depth, rhs, rhsStride, lhs, lhsStride, res, resIncr,
191 resStride, alpha, blocking, info);
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>;
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) {
208 if (numext::is_exactly_zero(alpha))
return;
209#ifdef EIGEN_VECTORIZE_SME
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))
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);
224 Index kc = blocking.kc();
225 Index mc = (std::min)(rows, blocking.mc());
226 Index nc = (std::min)(cols, blocking.nc());
228 gemm_pack_lhs<LhsScalar, Index, LhsMapper, Traits::mr, Traits::LhsProgress,
typename Traits::LhsPacket4Packing,
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;
234#if !defined(EIGEN_USE_BLAS) && (defined(EIGEN_HAS_OPENMP) || defined(EIGEN_GEMM_THREADPOOL))
237 int tid = info->logical_thread_id;
238 int threads = info->num_threads;
240 LhsScalar* blockA = blocking.blockA();
241 eigen_internal_assert(blockA != 0);
243 std::size_t sizeB = kc * nc;
244 ei_declare_aligned_stack_constructed_variable(RhsScalar, blockB, sizeB, 0);
247 for (Index k = 0; k < depth; k += kc) {
248 const Index actual_kc = (std::min)(k + kc, depth) - k;
252 pack_rhs(blockB, rhs.getSubMapper(k, 0), actual_kc, nc);
261 while (info->task_info[tid].users != 0) {
262 std::this_thread::yield();
264 info->task_info[tid].users = threads;
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);
270 info->task_info[tid].sync = k;
273 for (
int shift = 0; shift < threads; ++shift) {
274 int i = (tid + shift) % threads;
280 while (info->task_info[i].sync != k) {
281 std::this_thread::yield();
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);
290 for (Index j = nc; j < cols; j += nc) {
291 const Index actual_nc = (std::min)(j + nc, cols) - j;
294 pack_rhs(blockB, rhs.getSubMapper(k, j), actual_kc, actual_nc);
297 gebp(res.getSubMapper(0, j), blockA, blockB, rows, actual_kc, actual_nc, alpha);
302 for (Index i = 0; i < threads; ++i) info->task_info[i].users -= 1;
307 EIGEN_UNUSED_VARIABLE(info);
310 std::size_t sizeA = kc * mc;
311 std::size_t sizeB = kc * nc;
313 ei_declare_aligned_stack_constructed_variable(LhsScalar, blockA, sizeA, blocking.blockA());
314 ei_declare_aligned_stack_constructed_variable(RhsScalar, blockB, sizeB, blocking.blockB());
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>;
324 using SequentialGemmLoop = gemm_pack_lhs_first_loop_policy;
327 SequentialGemmLoop::run(rows, cols, depth, kc, mc, nc, lhs, rhs, res, pack_lhs, pack_rhs, gebp, blockA, blockB,
338template <
typename Scalar,
typename Index,
typename Gemm,
typename Lhs,
typename Rhs,
typename Dest,
339 typename BlockingType>
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) {}
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();
350 bool ownsNoBuffers()
const {
return m_blocking.blockA() ==
nullptr && m_blocking.blockB() ==
nullptr; }
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();
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);
360 using Traits =
typename Gemm::Traits;
366 Scalar m_actualAlpha;
367 BlockingType& m_blocking;
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;
374template <
typename LhsScalar_,
typename RhsScalar_>
375class level3_blocking {
376 using LhsScalar = LhsScalar_;
377 using RhsScalar = RhsScalar_;
380 LhsScalar* m_blockA =
nullptr;
381 RhsScalar* m_blockB =
nullptr;
388 level3_blocking() =
default;
390 inline Index mc()
const {
return m_mc; }
391 inline Index nc()
const {
return m_nc; }
392 inline Index kc()
const {
return m_kc; }
394 inline LhsScalar* blockA() {
return m_blockA; }
395 inline RhsScalar* blockB() {
return m_blockB; }
398template <
int StorageOrder,
typename LhsScalar_,
typename RhsScalar_,
int MaxRows,
int MaxCols,
int MaxDepth,
400class gemm_blocking_space<StorageOrder, LhsScalar_, RhsScalar_, MaxRows, MaxCols, MaxDepth, KcFactor,
402 :
public level3_blocking<std::conditional_t<StorageOrder == RowMajor, RhsScalar_, LhsScalar_>,
403 std::conditional_t<StorageOrder == RowMajor, LhsScalar_, RhsScalar_>> {
405 Transpose = StorageOrder ==
RowMajor,
406 ActualRows = Transpose ? MaxCols : MaxRows,
407 ActualCols = Transpose ? MaxRows : MaxCols
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 };
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];
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];
422 gemm_blocking_space(Index , Index , Index , Index ,
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;
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));
438 void initParallel(Index, Index, Index, Index) {}
440 inline void allocateA() {}
441 inline void allocateB() {}
442 inline void allocateAll() {}
445template <
int StorageOrder,
typename LhsScalar_,
typename RhsScalar_,
int MaxRows,
int MaxCols,
int MaxDepth,
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_>;
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;
464 computeProductBlockingSizes<LhsScalar, RhsScalar, KcFactor>(this->m_kc, this->m_mc, this->m_nc, num_threads);
467 Index n = this->m_nc;
468 computeProductBlockingSizes<LhsScalar, RhsScalar, KcFactor>(this->m_kc, this->m_mc, n, num_threads);
471 m_sizeA = this->m_mc * this->m_kc;
472 m_sizeB = this->m_kc * this->m_nc;
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;
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;
489 if (this->m_blockA == 0) this->m_blockA = scratch_new<LhsScalar>(m_sizeA);
493 if (this->m_blockB == 0) this->m_blockB = scratch_new<RhsScalar>(m_sizeB);
501 ~gemm_blocking_space() {
502 scratch_delete(this->m_blockA, m_sizeA);
503 scratch_delete(this->m_blockB, m_sizeB);
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;
518 using LhsBlasTraits = internal::blas_traits<Lhs>;
519 using ActualLhsType =
typename LhsBlasTraits::DirectLinearAccessType;
520 using ActualLhsTypeCleaned = internal::remove_all_t<ActualLhsType>;
522 using RhsBlasTraits = internal::blas_traits<Rhs>;
523 using ActualRhsType =
typename RhsBlasTraits::DirectLinearAccessType;
524 using ActualRhsTypeCleaned = internal::remove_all_t<ActualRhsType>;
526 enum { MaxDepthAtCompileTime = min_size_prefer_fixed(Lhs::MaxColsAtCompileTime, Rhs::MaxRowsAtCompileTime) };
528 using lazyproduct = generic_product_impl<Lhs, Rhs, DenseShape, DenseShape, CoeffBasedProductMode>;
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 :
541 EIGEN_GEMM_TO_COEFFBASED_THRESHOLD;
543#ifdef EIGEN_VECTORIZE_SME
549 static constexpr Index kCoeffBasedOutputArea = sme_has_gebp_kernel<LhsScalar, RhsScalar>::value
550 ? Index(EIGEN_SME_GEMM_TO_COEFFBASED_OUTPUT_AREA_THRESHOLD(Scalar))
553 template <
typename Dst>
554 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
bool outputAreaBelowThreshold(
const Dst& dst) {
556 return dst.rows() > 1 && dst.cols() > 1 && dst.size() <= kCoeffBasedOutputArea;
560#ifdef EIGEN_VECTORIZE_SME
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(),
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);
579 if ((rhs.rows() + dst.rows() + dst.cols()) < kCoeffBasedThreshold)
return true;
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));
594 template <
typename Dst>
595 static void evalTo(Dst& dst,
const Lhs& lhs,
const Rhs& rhs) {
596 if (scalarFactorIsZero(lhs, rhs)) {
600 if (useRuntimeCoeffBasedProduct(dst, rhs))
601 lazyproduct::eval_dynamic(dst, lhs, rhs, internal::assign_op<typename Dst::Scalar, Scalar>());
604 scaleAndAddTo(dst, lhs, rhs, Scalar(1));
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>());
614 scaleAndAddTo(dst, lhs, rhs, Scalar(1));
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>());
623 scaleAndAddTo(dst, lhs, rhs, Scalar(-1));
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;
631 if (dst.cols() == 1) {
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) {
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);
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);
646 Scalar actualAlpha = combine_scalar_factors(alpha, a_lhs, a_rhs);
650 Dest::MaxRowsAtCompileTime, Dest::MaxColsAtCompileTime, MaxDepthAtCompileTime>;
652 using GemmFunctor = internal::gemm_functor<
654 internal::general_matrix_matrix_product<
656 bool(LhsBlasTraits::NeedToConjugate), RhsScalar,
659 ActualLhsTypeCleaned, ActualRhsTypeCleaned, Dest, BlockingType>;
661 BlockingType blocking(dst.rows(), dst.cols(), lhs.cols(), 1,
true);
662#ifdef EIGEN_VECTORIZE_SME
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);
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(),
@ ColMajor
Definition Constants.h:319
@ RowMajor
Definition Constants.h:321
constexpr unsigned int RowMajorBit
Definition Constants.h:71