Eigen  5.0.1
 
Loading...
Searching...
No Matches
SparseSparseProductWithPruning.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2008-2014 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_SPARSESPARSEPRODUCTWITHPRUNING_H
12#define EIGEN_SPARSESPARSEPRODUCTWITHPRUNING_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20
21// perform a pseudo in-place sparse * sparse product assuming all matrices are col major
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;
28
29 // make sure to call innerSize/outerSize since we fake the storage order.
30 Index rows = lhs.innerSize();
31 Index cols = rhs.outerSize();
32 eigen_assert(lhs.outerSize() == rhs.innerSize());
33
34 // allocate a temporary buffer
35 AmbiVector<ResScalar, StorageIndex> tempVector(rows);
36
37 // mimics a resizeByInnerOuter:
38 EIGEN_IF_CONSTEXPR (ResultType::IsRowMajor) {
39 res.resize(cols, rows);
40 } else {
41 res.resize(rows, cols);
42 }
43
44 evaluator<Lhs> lhsEval(lhs);
45 evaluator<Rhs> rhsEval(rhs);
46
47 // estimate the number of non zero entries
48 // given a rhs column containing Y non zeros, we assume that the respective Y columns
49 // of the lhs differs in average of one non zeros, thus the number of non zeros for
50 // the product of a rhs column with the lhs is X+Y where X is the average number of non zero
51 // per column of the lhs.
52 // Therefore, we have nnz(lhs*rhs) = nnz(lhs) + nnz(rhs)
53 Index estimated_nnz_prod = lhsEval.nonZerosEstimate() + rhsEval.nonZerosEstimate();
54
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) {
58 // FIXME: compute a more accurate per-column nnz ratio for res.
59 tempVector.init(ratioColRes);
60 tempVector.setZero();
61 for (typename evaluator<Rhs>::InnerIterator rhsIt(rhsEval, j); rhsIt; ++rhsIt) {
62 // FIXME: rewrite as tmp += rhsIt.value() * lhs.col(rhsIt.index()).
63 tempVector.restart();
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;
67 }
68 }
69 res.startVec(j);
70 for (typename AmbiVector<ResScalar, StorageIndex>::Iterator it(tempVector, tolerance); it; ++it)
71 res.insertBackByOuterInner(j, it.index()) = it.value();
72 }
73 res.finalize();
74}
75
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;
80
81template <typename Lhs, typename Rhs, typename ResultType>
82struct sparse_sparse_product_with_pruning_selector<Lhs, Rhs, ResultType, ColMajor, ColMajor, ColMajor> {
83 using RealScalar = typename ResultType::RealScalar;
84
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);
88 res.swap(res_);
89 }
90};
91
92template <typename Lhs, typename Rhs, typename ResultType>
93struct sparse_sparse_product_with_pruning_selector<Lhs, Rhs, ResultType, ColMajor, ColMajor, RowMajor> {
94 using RealScalar = typename ResultType::RealScalar;
95 static void run(const Lhs& lhs, const Rhs& rhs, ResultType& res, const RealScalar& tolerance) {
96 // we need a col-major matrix to hold the result
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);
100 res = res_;
101 }
102};
103
104template <typename Lhs, typename Rhs, typename ResultType>
105struct sparse_sparse_product_with_pruning_selector<Lhs, Rhs, ResultType, RowMajor, RowMajor, RowMajor> {
106 using RealScalar = typename ResultType::RealScalar;
107 static void run(const Lhs& lhs, const Rhs& rhs, ResultType& res, const RealScalar& tolerance) {
108 // let's transpose the product to get a column x column product
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);
111 res.swap(res_);
112 }
113};
114
115template <typename Lhs, typename Rhs, typename ResultType>
116struct sparse_sparse_product_with_pruning_selector<Lhs, Rhs, ResultType, RowMajor, RowMajor, ColMajor> {
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,
124 res, tolerance);
125 }
126};
127
128template <typename Lhs, typename Rhs, typename ResultType>
129struct sparse_sparse_product_with_pruning_selector<Lhs, Rhs, ResultType, ColMajor, RowMajor, RowMajor> {
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,
135 res, tolerance);
136 }
137};
138
139template <typename Lhs, typename Rhs, typename ResultType>
140struct sparse_sparse_product_with_pruning_selector<Lhs, Rhs, ResultType, RowMajor, ColMajor, RowMajor> {
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);
147 }
148};
149
150template <typename Lhs, typename Rhs, typename ResultType>
151struct sparse_sparse_product_with_pruning_selector<Lhs, Rhs, ResultType, ColMajor, RowMajor, ColMajor> {
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);
157 }
158};
159
160template <typename Lhs, typename Rhs, typename ResultType>
161struct sparse_sparse_product_with_pruning_selector<Lhs, Rhs, ResultType, RowMajor, ColMajor, ColMajor> {
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);
167 }
168};
169
170} // end namespace internal
171
172} // end namespace Eigen
173
174#endif // EIGEN_SPARSESPARSEPRODUCTWITHPRUNING_H
Definition AmbiVector.h:308
@ ColMajor
Definition Constants.h:319
@ RowMajor
Definition Constants.h:321
constexpr unsigned int RowMajorBit
Definition Constants.h:71