163 static constexpr int value = FixedSizeCoeffBasedProduct ? CoeffBasedProductMode : ProductType;
164#ifdef EIGEN_DEBUG_PRODUCT
165 static void debug() {
166 const int rows_select = product_size_category<Rows, MaxRows>::value;
167 const int cols_select = product_size_category<Cols, MaxCols>::value;
168 const int depth_select = product_size_category<Depth, MaxDepth>::value;
169 EIGEN_DEBUG_VAR(Rows);
170 EIGEN_DEBUG_VAR(Cols);
171 EIGEN_DEBUG_VAR(Depth);
172 EIGEN_DEBUG_VAR(rows_select);
173 EIGEN_DEBUG_VAR(cols_select);
174 EIGEN_DEBUG_VAR(depth_select);
175 EIGEN_DEBUG_VAR(ProductType);
176 EIGEN_DEBUG_VAR(FixedSizeCoeffBasedProduct);
177 EIGEN_DEBUG_VAR(
value);
186template <
int M,
int N>
187struct product_type_selector<M, N, 1> : std::integral_constant<int, OuterProduct> {};
189struct product_type_selector<M, 1, 1> : std::integral_constant<int, LazyCoeffBasedProductMode> {};
191struct product_type_selector<1, N, 1> : std::integral_constant<int, LazyCoeffBasedProductMode> {};
193struct product_type_selector<1, 1, Depth> : std::integral_constant<int, InnerProduct> {};
195struct product_type_selector<1, 1, 1> : std::integral_constant<int, InnerProduct> {};
197struct product_type_selector<Small, 1, Small> : std::integral_constant<int, CoeffBasedProductMode> {};
199struct product_type_selector<1, Small, Small> : std::integral_constant<int, CoeffBasedProductMode> {};
201struct product_type_selector<Small, Small, Small> : std::integral_constant<int, CoeffBasedProductMode> {};
203struct product_type_selector<Small, Small, 1> : std::integral_constant<int, LazyCoeffBasedProductMode> {};
205struct product_type_selector<Small, Large, 1> : std::integral_constant<int, LazyCoeffBasedProductMode> {};
207struct product_type_selector<Large, Small, 1> : std::integral_constant<int, LazyCoeffBasedProductMode> {};
209struct product_type_selector<1, Large, Small> : std::integral_constant<int, CoeffBasedProductMode> {};
211struct product_type_selector<1, Large, Large> : std::integral_constant<int, GemvProduct> {};
213struct product_type_selector<1, Small, Large> : std::integral_constant<int, CoeffBasedProductMode> {};
215struct product_type_selector<Large, 1, Small> : std::integral_constant<int, CoeffBasedProductMode> {};
217struct product_type_selector<Large, 1, Large> : std::integral_constant<int, GemvProduct> {};
219struct product_type_selector<Small, 1, Large> : std::integral_constant<int, CoeffBasedProductMode> {};
221struct product_type_selector<Small, Small, Large> : std::integral_constant<int, GemmProduct> {};
223struct product_type_selector<Large, Small, Large> : std::integral_constant<int, GemmProduct> {};
225struct product_type_selector<Small, Large, Large> : std::integral_constant<int, GemmProduct> {};
227struct product_type_selector<Large, Large, Large> : std::integral_constant<int, GemmProduct> {};
229struct product_type_selector<Large, Small, Small> : std::integral_constant<int, CoeffBasedProductMode> {};
231struct product_type_selector<Small, Large, Small> : std::integral_constant<int, CoeffBasedProductMode> {};
233struct product_type_selector<Large, Large, Small> : std::integral_constant<int, GemmProduct> {};
264template <
int S
ide,
int StorageOrder,
bool BlasCompatible>
265struct gemv_dense_selector;
271template <
typename Scalar,
int Size,
int MaxSize,
bool Cond>
272struct gemv_static_vector_if;
274template <
typename Scalar,
int Size,
int MaxSize>
275struct gemv_static_vector_if<Scalar, Size, MaxSize, false> {
276 EIGEN_DEVICE_FUNC
constexpr Scalar* data() {
277 eigen_internal_assert(
false &&
"should never be called");
282template <
typename Scalar,
int Size>
283struct gemv_static_vector_if<Scalar, Size, Dynamic, true> {
284 EIGEN_DEVICE_FUNC
constexpr Scalar* data() {
return 0; }
287template <
typename Scalar,
int Size,
int MaxSize>
288struct gemv_static_vector_if<Scalar, Size, MaxSize, true> {
289#if EIGEN_MAX_STATIC_ALIGN_BYTES != 0
290 internal::plain_array<Scalar, internal::min_size_prefer_fixed(Size, MaxSize), 0, AlignedMax> m_data;
291 constexpr Scalar* data() {
return m_data.array; }
295 internal::plain_array<Scalar, internal::min_size_prefer_fixed(Size, MaxSize) + EIGEN_MAX_ALIGN_BYTES, 0> m_data;
296 constexpr Scalar* data() {
297 return reinterpret_cast<Scalar*
>((std::uintptr_t(m_data.array) & ~(std::size_t(EIGEN_MAX_ALIGN_BYTES - 1))) +
298 EIGEN_MAX_ALIGN_BYTES);
303template <
typename ResScalar>
304using gemv_mapped_destination =
305 Map<Matrix<ResScalar, Dynamic, 1>, plain_enum_min(AlignedMax, internal::packet_traits<ResScalar>::size)>;
307template <
typename Dest>
308EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE gemv_mapped_destination<typename Dest::Scalar> gemv_construct_mapped_destination(
309 Dest& dest,
typename Dest::Scalar* actual_dest_ptr) {
310#ifdef EIGEN_DENSE_STORAGE_CTOR_PLUGIN
311 constexpr int Size = Dest::SizeAtCompileTime;
312 Index size = dest.size();
313 EIGEN_DENSE_STORAGE_CTOR_PLUGIN
315 return gemv_mapped_destination<typename Dest::Scalar>(actual_dest_ptr, dest.size());
319template <
bool EvalToDest,
typename Dest>
320EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void gemv_prepare_destination(Dest& dest,
321 typename Dest::Scalar* actual_dest_ptr) {
322 EIGEN_IF_CONSTEXPR (!EvalToDest) gemv_construct_mapped_destination(dest, actual_dest_ptr) = dest;
325template <
bool EvalToDest,
typename Dest>
326EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void gemv_prepare_destination(Dest& dest,
typename Dest::Scalar* actual_dest_ptr,
327 bool initialize_to_zero) {
328 if (initialize_to_zero) {
329 gemv_construct_mapped_destination(dest, actual_dest_ptr).setZero();
331 gemv_prepare_destination<EvalToDest>(dest, actual_dest_ptr);
335template <
bool EvalToDest,
typename Dest>
336EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void gemv_copy_destination(Dest& dest,
typename Dest::Scalar* actual_dest_ptr) {
337 EIGEN_IF_CONSTEXPR (!EvalToDest) {
338 dest = gemv_mapped_destination<typename Dest::Scalar>(actual_dest_ptr, dest.size());
343template <
typename RhsScalar,
typename ResScalar,
bool EvalToDestAtCompileTime,
bool ComplexByReal>
344class gemv_destination_policy {
346 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
explicit gemv_destination_policy(
const ResScalar& alpha) : alpha_(alpha) {}
348 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
bool eval_to_dest()
const {
349 return EvalToDestAtCompileTime && alpha_is_compatible();
352 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE RhsScalar compatible_alpha()
const {
353 return alpha_is_compatible() ? get_factor<ResScalar, RhsScalar>::run(alpha_) : RhsScalar(1);
356 template <
typename Dest>
357 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void prepare(Dest& dest, ResScalar* actual_dest_ptr)
const {
358 gemv_prepare_destination<EvalToDestAtCompileTime>(dest, actual_dest_ptr, !alpha_is_compatible());
361 template <
typename Dest>
362 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void copy_back(Dest& dest, ResScalar* actual_dest_ptr)
const {
363 if (!alpha_is_compatible()) {
364 dest.matrix() += alpha_ * gemv_mapped_destination<ResScalar>(actual_dest_ptr, dest.size());
366 gemv_copy_destination<EvalToDestAtCompileTime>(dest, actual_dest_ptr);
371 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
bool alpha_is_compatible()
const {
372 return !ComplexByReal || numext::is_exactly_zero(numext::imag(alpha_));
379template <
bool DirectlyUseRhs,
typename ActualRhsType>
380EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void gemv_prepare_rhs(
381 const ActualRhsType& rhs,
typename remove_all_t<ActualRhsType>::Scalar* actual_rhs_ptr) {
382 using ActualRhsTypeCleaned = remove_all_t<ActualRhsType>;
383 EIGEN_IF_CONSTEXPR (!DirectlyUseRhs) {
384#ifdef EIGEN_DENSE_STORAGE_CTOR_PLUGIN
385 constexpr int Size = ActualRhsTypeCleaned::SizeAtCompileTime;
386 Index size = rhs.size();
387 EIGEN_DENSE_STORAGE_CTOR_PLUGIN
390 Map<typename ActualRhsTypeCleaned::PlainObject, AlignedMax>(actual_rhs_ptr, rhs.size()) = rhs;
395template <
int StorageOrder,
bool BlasCompatible>
396struct gemv_dense_selector<
OnTheLeft, StorageOrder, BlasCompatible> {
397 template <
typename Lhs,
typename Rhs,
typename Dest>
398 static void run(
const Lhs& lhs,
const Rhs& rhs, Dest& dest,
const typename Dest::Scalar& alpha) {
399 Transpose<Dest> destT(dest);
401 gemv_dense_selector<OnTheRight, OtherStorageOrder, BlasCompatible>::run(rhs.transpose(), lhs.transpose(), destT,
408 template <
typename Lhs,
typename Rhs,
typename Dest>
409 static inline void run(
const Lhs& lhs,
const Rhs& rhs, Dest& dest,
const typename Dest::Scalar& alpha) {
410 using LhsScalar =
typename Lhs::Scalar;
411 using RhsScalar =
typename Rhs::Scalar;
412 using ResScalar =
typename Dest::Scalar;
414 using LhsBlasTraits = internal::blas_traits<Lhs>;
415 using ActualLhsType =
typename LhsBlasTraits::DirectLinearAccessType;
416 using RhsBlasTraits = internal::blas_traits<Rhs>;
417 using ActualRhsType =
typename RhsBlasTraits::DirectLinearAccessType;
419 ActualLhsType actualLhs = LhsBlasTraits::extract(lhs);
420 ActualRhsType actualRhs = RhsBlasTraits::extract(rhs);
422 ResScalar actualAlpha = combine_scalar_factors(alpha, lhs, rhs);
425 using ActualDest = std::conditional_t<Dest::IsVectorAtCompileTime, Dest, typename Dest::ColXpr>;
430 EvalToDestAtCompileTime = (ActualDest::InnerStrideAtCompileTime == 1),
431 ComplexByReal = (NumTraits<LhsScalar>::IsComplex) && (!NumTraits<RhsScalar>::IsComplex),
432 MightCannotUseDest = ((!EvalToDestAtCompileTime) || ComplexByReal) && (ActualDest::MaxSizeAtCompileTime != 0)
435 using LhsMapper = const_blas_data_mapper<LhsScalar, Index, ColMajor>;
436 using RhsMapper = const_blas_data_mapper<RhsScalar, Index, RowMajor>;
437 EIGEN_IF_CONSTEXPR (!MightCannotUseDest) {
440 general_matrix_vector_product<
441 Index, LhsScalar, LhsMapper,
ColMajor, LhsBlasTraits::NeedToConjugate, RhsScalar, RhsMapper,
442 RhsBlasTraits::NeedToConjugate>::run(actualLhs.rows(), actualLhs.cols(),
443 LhsMapper(actualLhs.data(), actualLhs.outerStride()),
444 RhsMapper(actualRhs.data(), actualRhs.innerStride()), dest.data(), 1,
445 get_factor<ResScalar, RhsScalar>::run(actualAlpha));
447 gemv_static_vector_if<ResScalar, ActualDest::SizeAtCompileTime, ActualDest::MaxSizeAtCompileTime,
451 gemv_destination_policy<RhsScalar, ResScalar, EvalToDestAtCompileTime, ComplexByReal> destPolicy(actualAlpha);
453 ei_declare_aligned_stack_constructed_variable(ResScalar, actualDestPtr, dest.size(),
454 destPolicy.eval_to_dest() ? dest.data() : static_dest.data());
456 destPolicy.prepare(dest, actualDestPtr);
458 general_matrix_vector_product<Index, LhsScalar, LhsMapper,
ColMajor, LhsBlasTraits::NeedToConjugate, RhsScalar,
459 RhsMapper, RhsBlasTraits::NeedToConjugate>::run(actualLhs.rows(), actualLhs.cols(),
460 LhsMapper(actualLhs.data(),
461 actualLhs.outerStride()),
462 RhsMapper(actualRhs.data(),
463 actualRhs.innerStride()),
465 destPolicy.compatible_alpha());
467 destPolicy.copy_back(dest, actualDestPtr);
474 template <
typename Lhs,
typename Rhs,
typename Dest>
475 static void run(
const Lhs& lhs,
const Rhs& rhs, Dest& dest,
const typename Dest::Scalar& alpha) {
476 using LhsScalar =
typename Lhs::Scalar;
477 using RhsScalar =
typename Rhs::Scalar;
478 using ResScalar =
typename Dest::Scalar;
480 using LhsBlasTraits = internal::blas_traits<Lhs>;
481 using ActualLhsType =
typename LhsBlasTraits::DirectLinearAccessType;
482 using RhsBlasTraits = internal::blas_traits<Rhs>;
483 using ActualRhsType =
typename RhsBlasTraits::DirectLinearAccessType;
484 using ActualRhsTypeCleaned = internal::remove_all_t<ActualRhsType>;
486 std::add_const_t<ActualLhsType> actualLhs = LhsBlasTraits::extract(lhs);
487 std::add_const_t<ActualRhsType> actualRhs = RhsBlasTraits::extract(rhs);
489 ResScalar actualAlpha = combine_scalar_factors(alpha, lhs, rhs);
495 ActualRhsTypeCleaned::InnerStrideAtCompileTime == 1 || ActualRhsTypeCleaned::MaxSizeAtCompileTime == 0
498 gemv_static_vector_if<RhsScalar, ActualRhsTypeCleaned::SizeAtCompileTime,
499 ActualRhsTypeCleaned::MaxSizeAtCompileTime, !DirectlyUseRhs>
502 ei_declare_aligned_stack_constructed_variable(
503 RhsScalar, actualRhsPtr, actualRhs.size(),
504 DirectlyUseRhs ?
const_cast<RhsScalar*
>(actualRhs.data()) : static_rhs.data());
506 gemv_prepare_rhs<DirectlyUseRhs>(actualRhs, actualRhsPtr);
508 using LhsMapper = const_blas_data_mapper<LhsScalar, Index, RowMajor>;
509 using RhsMapper = const_blas_data_mapper<RhsScalar, Index, ColMajor>;
510 general_matrix_vector_product<Index, LhsScalar, LhsMapper,
RowMajor, LhsBlasTraits::NeedToConjugate, RhsScalar,
511 RhsMapper, RhsBlasTraits::NeedToConjugate>::
512 run(actualLhs.rows(), actualLhs.cols(), LhsMapper(actualLhs.data(), actualLhs.outerStride()),
513 RhsMapper(actualRhsPtr, 1), dest.data(),
514 dest.col(0).innerStride(),
522 template <
typename Lhs,
typename Rhs,
typename Dest>
523 static void run(
const Lhs& lhs,
const Rhs& rhs, Dest& dest,
const typename Dest::Scalar& alpha) {
524 EIGEN_STATIC_ASSERT((!nested_eval<Lhs, 1>::Evaluate),
525 EIGEN_INTERNAL_COMPILATION_ERROR_OR_YOU_MADE_A_PROGRAMMING_MISTAKE);
528 typename nested_eval<Rhs, 1>::type actual_rhs(rhs);
529 const Index size = rhs.rows();
530 for (Index k = 0; k < size; ++k) dest += (alpha * actual_rhs.coeff(k)) * lhs.col(k);
536 template <
typename Lhs,
typename Rhs,
typename Dest>
537 static void run(
const Lhs& lhs,
const Rhs& rhs, Dest& dest,
const typename Dest::Scalar& alpha) {
538 EIGEN_STATIC_ASSERT((!nested_eval<Lhs, 1>::Evaluate),
539 EIGEN_INTERNAL_COMPILATION_ERROR_OR_YOU_MADE_A_PROGRAMMING_MISTAKE);
540 typename nested_eval<Rhs, Lhs::RowsAtCompileTime>::type actual_rhs(rhs);
541 const Index rows = dest.rows();
542 for (Index i = 0; i < rows; ++i)
543 dest.coeffRef(i) += alpha * (lhs.row(i).cwiseProduct(actual_rhs.transpose())).sum();
559template <
typename Derived>
560template <
typename OtherDerived>
562 const MatrixBase<OtherDerived>& other)
const {
568 ProductIsValid = Derived::ColsAtCompileTime == Dynamic || OtherDerived::RowsAtCompileTime == Dynamic ||
569 int(Derived::ColsAtCompileTime) == int(OtherDerived::RowsAtCompileTime),
570 AreVectors = Derived::IsVectorAtCompileTime && OtherDerived::IsVectorAtCompileTime,
571 SameSizes = EIGEN_PREDICATE_SAME_MATRIX_SIZE(Derived, OtherDerived)
577 ProductIsValid || !(AreVectors && SameSizes),
578 INVALID_VECTOR_VECTOR_PRODUCT__IF_YOU_WANTED_A_DOT_OR_COEFF_WISE_PRODUCT_YOU_MUST_USE_THE_EXPLICIT_FUNCTIONS)
579 EIGEN_STATIC_ASSERT(ProductIsValid || !(SameSizes && !AreVectors),
580 INVALID_MATRIX_PRODUCT__IF_YOU_WANTED_A_COEFF_WISE_PRODUCT_YOU_MUST_USE_THE_EXPLICIT_FUNCTION)
581 EIGEN_STATIC_ASSERT(ProductIsValid || SameSizes, INVALID_MATRIX_PRODUCT)
582#ifdef EIGEN_DEBUG_PRODUCT
583 internal::product_type<Derived, OtherDerived>::debug();
600template <
typename Derived>
601template <
typename OtherDerived>
605 ProductIsValid = Derived::ColsAtCompileTime == Dynamic || OtherDerived::RowsAtCompileTime == Dynamic ||
606 int(Derived::ColsAtCompileTime) == int(OtherDerived::RowsAtCompileTime),
607 AreVectors = Derived::IsVectorAtCompileTime && OtherDerived::IsVectorAtCompileTime,
608 SameSizes = EIGEN_PREDICATE_SAME_MATRIX_SIZE(Derived, OtherDerived)
614 ProductIsValid || !(AreVectors && SameSizes),
615 INVALID_VECTOR_VECTOR_PRODUCT__IF_YOU_WANTED_A_DOT_OR_COEFF_WISE_PRODUCT_YOU_MUST_USE_THE_EXPLICIT_FUNCTIONS)
616 EIGEN_STATIC_ASSERT(ProductIsValid || !(SameSizes && !AreVectors),
617 INVALID_MATRIX_PRODUCT__IF_YOU_WANTED_A_COEFF_WISE_PRODUCT_YOU_MUST_USE_THE_EXPLICIT_FUNCTION)
618 EIGEN_STATIC_ASSERT(ProductIsValid || SameSizes, INVALID_MATRIX_PRODUCT)