11#ifndef EIGEN_SPARSESPARSEPRODUCTWITHPRUNING_H
12#define EIGEN_SPARSESPARSEPRODUCTWITHPRUNING_H
15#include "./InternalHeaderCheck.h"
22template <
typename Lhs,
typename Rhs,
typename ResultType>
23static void sparse_sparse_product_with_pruning_impl(
const Lhs& lhs,
const Rhs& rhs, ResultType& res,
24 const typename ResultType::RealScalar& tolerance) {
25 using RhsScalar =
typename remove_all_t<Rhs>::Scalar;
26 using ResScalar =
typename remove_all_t<ResultType>::Scalar;
27 using StorageIndex =
typename remove_all_t<Lhs>::StorageIndex;
30 Index rows = lhs.innerSize();
31 Index cols = rhs.outerSize();
32 eigen_assert(lhs.outerSize() == rhs.innerSize());
35 AmbiVector<ResScalar, StorageIndex> tempVector(rows);
38 EIGEN_IF_CONSTEXPR (ResultType::IsRowMajor) {
39 res.resize(cols, rows);
41 res.resize(rows, cols);
44 evaluator<Lhs> lhsEval(lhs);
45 evaluator<Rhs> rhsEval(rhs);
53 Index estimated_nnz_prod = lhsEval.nonZerosEstimate() + rhsEval.nonZerosEstimate();
55 res.reserve(estimated_nnz_prod);
56 double ratioColRes = double(estimated_nnz_prod) / (double(lhs.rows()) * double(rhs.cols()));
57 for (Index j = 0; j < cols; ++j) {
59 tempVector.init(ratioColRes);
61 for (
typename evaluator<Rhs>::InnerIterator rhsIt(rhsEval, j); rhsIt; ++rhsIt) {
64 RhsScalar x = rhsIt.value();
65 for (
typename evaluator<Lhs>::InnerIterator lhsIt(lhsEval, rhsIt.index()); lhsIt; ++lhsIt) {
66 tempVector.coeffRef(lhsIt.index()) += lhsIt.value() * x;
71 res.insertBackByOuterInner(j, it.index()) = it.value();
76template <typename Lhs, typename Rhs, typename ResultType, int LhsStorageOrder = traits<Lhs>::Flags &
RowMajorBit,
77 int RhsStorageOrder = traits<Rhs>::Flags &
RowMajorBit,
78 int ResStorageOrder = traits<ResultType>::Flags &
RowMajorBit>
79struct sparse_sparse_product_with_pruning_selector;
81template <
typename Lhs,
typename Rhs,
typename ResultType>
83 using RealScalar =
typename ResultType::RealScalar;
85 static void run(
const Lhs& lhs,
const Rhs& rhs, ResultType& res,
const RealScalar& tolerance) {
86 remove_all_t<ResultType> res_{res.rows(), res.cols()};
87 internal::sparse_sparse_product_with_pruning_impl<Lhs, Rhs, ResultType>(lhs, rhs, res_, tolerance);
92template <
typename Lhs,
typename Rhs,
typename ResultType>
94 using RealScalar =
typename ResultType::RealScalar;
95 static void run(
const Lhs& lhs,
const Rhs& rhs, ResultType& res,
const RealScalar& tolerance) {
97 using SparseTemporaryType = SparseMatrix<typename ResultType::Scalar, ColMajor, typename ResultType::StorageIndex>;
98 SparseTemporaryType res_{res.rows(), res.cols()};
99 internal::sparse_sparse_product_with_pruning_impl<Lhs, Rhs, SparseTemporaryType>(lhs, rhs, res_, tolerance);
104template <
typename Lhs,
typename Rhs,
typename ResultType>
106 using RealScalar =
typename ResultType::RealScalar;
107 static void run(
const Lhs& lhs,
const Rhs& rhs, ResultType& res,
const RealScalar& tolerance) {
109 remove_all_t<ResultType> res_{res.rows(), res.cols()};
110 internal::sparse_sparse_product_with_pruning_impl<Rhs, Lhs, ResultType>(rhs, lhs, res_, tolerance);
115template <
typename Lhs,
typename Rhs,
typename ResultType>
117 using RealScalar =
typename ResultType::RealScalar;
118 static void run(
const Lhs& lhs,
const Rhs& rhs, ResultType& res,
const RealScalar& tolerance) {
119 using ColMajorMatrixLhs = SparseMatrix<typename Lhs::Scalar, ColMajor, typename Lhs::StorageIndex>;
120 using ColMajorMatrixRhs = SparseMatrix<typename Rhs::Scalar, ColMajor, typename Lhs::StorageIndex>;
121 ColMajorMatrixLhs colLhs(lhs);
122 ColMajorMatrixRhs colRhs(rhs);
123 internal::sparse_sparse_product_with_pruning_impl<ColMajorMatrixLhs, ColMajorMatrixRhs, ResultType>(colLhs, colRhs,
128template <
typename Lhs,
typename Rhs,
typename ResultType>
130 using RealScalar =
typename ResultType::RealScalar;
131 static void run(
const Lhs& lhs,
const Rhs& rhs, ResultType& res,
const RealScalar& tolerance) {
132 using RowMajorMatrixLhs = SparseMatrix<typename Lhs::Scalar, RowMajor, typename Lhs::StorageIndex>;
133 RowMajorMatrixLhs rowLhs(lhs);
134 sparse_sparse_product_with_pruning_selector<RowMajorMatrixLhs, Rhs, ResultType, RowMajor, RowMajor>(rowLhs, rhs,
139template <
typename Lhs,
typename Rhs,
typename ResultType>
141 using RealScalar =
typename ResultType::RealScalar;
142 static void run(
const Lhs& lhs,
const Rhs& rhs, ResultType& res,
const RealScalar& tolerance) {
143 using RowMajorMatrixRhs = SparseMatrix<typename Rhs::Scalar, RowMajor, typename Lhs::StorageIndex>;
144 RowMajorMatrixRhs rowRhs(rhs);
145 sparse_sparse_product_with_pruning_selector<Lhs, RowMajorMatrixRhs, ResultType, RowMajor, RowMajor, RowMajor>(
146 lhs, rowRhs, res, tolerance);
150template <
typename Lhs,
typename Rhs,
typename ResultType>
152 using RealScalar =
typename ResultType::RealScalar;
153 static void run(
const Lhs& lhs,
const Rhs& rhs, ResultType& res,
const RealScalar& tolerance) {
154 using ColMajorMatrixRhs = SparseMatrix<typename Rhs::Scalar, ColMajor, typename Lhs::StorageIndex>;
155 ColMajorMatrixRhs colRhs(rhs);
156 internal::sparse_sparse_product_with_pruning_impl<Lhs, ColMajorMatrixRhs, ResultType>(lhs, colRhs, res, tolerance);
160template <
typename Lhs,
typename Rhs,
typename ResultType>
162 using RealScalar =
typename ResultType::RealScalar;
163 static void run(
const Lhs& lhs,
const Rhs& rhs, ResultType& res,
const RealScalar& tolerance) {
164 using ColMajorMatrixLhs = SparseMatrix<typename Lhs::Scalar, ColMajor, typename Lhs::StorageIndex>;
165 ColMajorMatrixLhs colLhs(lhs);
166 internal::sparse_sparse_product_with_pruning_impl<ColMajorMatrixLhs, Rhs, ResultType>(colLhs, rhs, res, tolerance);
Definition AmbiVector.h:308
@ ColMajor
Definition Constants.h:319
@ RowMajor
Definition Constants.h:321
constexpr unsigned int RowMajorBit
Definition Constants.h:71