Eigen  5.0.1
 
Loading...
Searching...
No Matches
SparsePermutation.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2012 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_SPARSE_PERMUTATION_H
12#define EIGEN_SPARSE_PERMUTATION_H
13
14// This file implements sparse * permutation products
15
16// IWYU pragma: private
17#include "./InternalHeaderCheck.h"
18
19namespace Eigen {
20
21namespace internal {
22
23template <typename ExpressionType, typename PlainObjectType,
24 bool NeedEval = !std::is_same<ExpressionType, PlainObjectType>::value>
25struct XprHelper {
26 XprHelper(const ExpressionType& xpr) : m_xpr(xpr) {}
27 inline const PlainObjectType& xpr() const { return m_xpr; }
28 // this is a new PlainObjectType initialized by xpr
29 const PlainObjectType m_xpr;
30};
31template <typename ExpressionType, typename PlainObjectType>
32struct XprHelper<ExpressionType, PlainObjectType, false> {
33 XprHelper(const ExpressionType& xpr) : m_xpr(xpr) {}
34 inline const PlainObjectType& xpr() const { return m_xpr; }
35 // this is a reference to xpr
36 const PlainObjectType& m_xpr;
37};
38
39template <typename PermDerived, bool NeedInverseEval>
40struct PermHelper {
41 using IndicesType = typename PermDerived::IndicesType;
42 using PermutationIndex = typename IndicesType::Scalar;
43 using type = PermutationMatrix<IndicesType::SizeAtCompileTime, IndicesType::MaxSizeAtCompileTime, PermutationIndex>;
44 PermHelper(const PermDerived& perm) : m_perm(perm.inverse()) {}
45 inline const type& perm() const { return m_perm; }
46 // this is a new PermutationMatrix initialized by perm.inverse()
47 const type m_perm;
48};
49template <typename PermDerived>
50struct PermHelper<PermDerived, false> {
51 using type = PermDerived;
52 PermHelper(const PermDerived& perm) : m_perm(perm) {}
53 inline const type& perm() const { return m_perm; }
54 // this is a reference to perm
55 const type& m_perm;
56};
57
58template <typename ExpressionType, int Side, bool Transposed>
59struct permutation_matrix_product<ExpressionType, Side, Transposed, SparseShape> {
60 using MatrixType = typename nested_eval<ExpressionType, 1>::type;
61 using MatrixTypeCleaned = remove_all_t<MatrixType>;
62
63 using Scalar = typename MatrixTypeCleaned::Scalar;
64 using StorageIndex = typename MatrixTypeCleaned::StorageIndex;
65
66 // the actual "return type" is `Dest`. this is a temporary type
67 using ReturnType = SparseMatrix<Scalar, MatrixTypeCleaned::IsRowMajor ? RowMajor : ColMajor, StorageIndex>;
68 using TmpHelper = XprHelper<ExpressionType, ReturnType>;
69
70 static constexpr bool NeedOuterPermutation = ExpressionType::IsRowMajor ? Side == OnTheLeft : Side == OnTheRight;
71 static constexpr bool NeedInversePermutation = Transposed ? Side == OnTheLeft : Side == OnTheRight;
72
73 template <typename Dest, typename PermutationType>
74 static inline void permute_outer(Dest& dst, const PermutationType& perm, const ExpressionType& xpr) {
75 // if ExpressionType is not ReturnType, evaluate `xpr` (allocation)
76 // otherwise, just reference `xpr`
77 // TODO: handle trivial expressions such as CwiseBinaryOp without temporary
78 const TmpHelper tmpHelper(xpr);
79 const ReturnType& tmp = tmpHelper.xpr();
80
81 ReturnType result(tmp.rows(), tmp.cols());
82
83 for (Index j = 0; j < tmp.outerSize(); j++) {
84 Index jp = perm.indices().coeff(j);
85 Index jsrc = NeedInversePermutation ? jp : j;
86 Index jdst = NeedInversePermutation ? j : jp;
87 Index begin = tmp.outerIndexPtr()[jsrc];
88 Index end = tmp.isCompressed() ? tmp.outerIndexPtr()[jsrc + 1] : begin + tmp.innerNonZeroPtr()[jsrc];
89 result.outerIndexPtr()[jdst + 1] += end - begin;
90 }
91
92 std::partial_sum(result.outerIndexPtr(), result.outerIndexPtr() + result.outerSize() + 1, result.outerIndexPtr());
93 result.resizeNonZeros(result.nonZeros());
94
95 for (Index j = 0; j < tmp.outerSize(); j++) {
96 Index jp = perm.indices().coeff(j);
97 Index jsrc = NeedInversePermutation ? jp : j;
98 Index jdst = NeedInversePermutation ? j : jp;
99 Index begin = tmp.outerIndexPtr()[jsrc];
100 Index end = tmp.isCompressed() ? tmp.outerIndexPtr()[jsrc + 1] : begin + tmp.innerNonZeroPtr()[jsrc];
101 Index target = result.outerIndexPtr()[jdst];
102 smart_copy(tmp.innerIndexPtr() + begin, tmp.innerIndexPtr() + end, result.innerIndexPtr() + target);
103 smart_copy(tmp.valuePtr() + begin, tmp.valuePtr() + end, result.valuePtr() + target);
104 }
105 dst = std::move(result);
106 }
107
108 template <typename Dest, typename PermutationType>
109 static inline void permute_inner(Dest& dst, const PermutationType& perm, const ExpressionType& xpr) {
110 using InnerPermHelper = PermHelper<PermutationType, NeedInversePermutation>;
111 using InnerPermType = typename InnerPermHelper::type;
112
113 // if ExpressionType is not ReturnType, evaluate `xpr` (allocation)
114 // otherwise, just reference `xpr`
115 // TODO: handle trivial expressions such as CwiseBinaryOp without temporary
116 const TmpHelper tmpHelper(xpr);
117 const ReturnType& tmp = tmpHelper.xpr();
118
119 // if inverse permutation of inner indices is requested, calculate perm.inverse() (allocation)
120 // otherwise, just reference `perm`
121 const InnerPermHelper permHelper(perm);
122 const InnerPermType& innerPerm = permHelper.perm();
123
124 ReturnType result(tmp.rows(), tmp.cols());
125
126 for (Index j = 0; j < tmp.outerSize(); j++) {
127 Index begin = tmp.outerIndexPtr()[j];
128 Index end = tmp.isCompressed() ? tmp.outerIndexPtr()[j + 1] : begin + tmp.innerNonZeroPtr()[j];
129 result.outerIndexPtr()[j + 1] += end - begin;
130 }
131
132 std::partial_sum(result.outerIndexPtr(), result.outerIndexPtr() + result.outerSize() + 1, result.outerIndexPtr());
133 result.resizeNonZeros(result.nonZeros());
134
135 for (Index j = 0; j < tmp.outerSize(); j++) {
136 Index begin = tmp.outerIndexPtr()[j];
137 Index end = tmp.isCompressed() ? tmp.outerIndexPtr()[j + 1] : begin + tmp.innerNonZeroPtr()[j];
138 Index target = result.outerIndexPtr()[j];
139 std::transform(tmp.innerIndexPtr() + begin, tmp.innerIndexPtr() + end, result.innerIndexPtr() + target,
140 [&innerPerm](StorageIndex i) { return innerPerm.indices().coeff(i); });
141 smart_copy(tmp.valuePtr() + begin, tmp.valuePtr() + end, result.valuePtr() + target);
142 }
143 // the inner indices were permuted, and must be sorted
144 result.sortInnerIndices();
145 dst = std::move(result);
146 }
147
148 template <typename Dest, typename PermutationType, bool DoOuter = NeedOuterPermutation,
149 std::enable_if_t<DoOuter, int> = 0>
150 static inline void run(Dest& dst, const PermutationType& perm, const ExpressionType& xpr) {
151 permute_outer(dst, perm, xpr);
152 }
153
154 template <typename Dest, typename PermutationType, bool DoOuter = NeedOuterPermutation,
155 std::enable_if_t<!DoOuter, int> = 0>
156 static inline void run(Dest& dst, const PermutationType& perm, const ExpressionType& xpr) {
157 permute_inner(dst, perm, xpr);
158 }
159};
160
161} // namespace internal
162
163namespace internal {
164
165template <int ProductTag>
166struct product_promote_storage_type<Sparse, PermutationStorage, ProductTag> {
167 using ret = Sparse;
168};
169template <int ProductTag>
170struct product_promote_storage_type<PermutationStorage, Sparse, ProductTag> {
171 using ret = Sparse;
172};
173
174// TODO, the following two overloads are only needed to define the right temporary type through
175// typename traits<permutation_matrix_product<Rhs,Lhs,OnTheRight,false> >::ReturnType
176// whereas it should be correctly handled by traits<Product<> >::PlainObject
177
178template <typename Lhs, typename Rhs, int ProductTag>
179struct product_evaluator<Product<Lhs, Rhs, AliasFreeProduct>, ProductTag, PermutationShape, SparseShape>
180 : public evaluator<typename permutation_matrix_product<Rhs, OnTheLeft, false, SparseShape>::ReturnType> {
181 using XprType = Product<Lhs, Rhs, AliasFreeProduct>;
182 using PlainObject = typename permutation_matrix_product<Rhs, OnTheLeft, false, SparseShape>::ReturnType;
183 using Base = evaluator<PlainObject>;
184
185 enum { Flags = Base::Flags | EvalBeforeNestingBit };
186
187 explicit product_evaluator(const XprType& xpr) : m_result(xpr.rows(), xpr.cols()) {
188 internal::construct_at<Base>(this, m_result);
189 generic_product_impl<Lhs, Rhs, PermutationShape, SparseShape, ProductTag>::evalTo(m_result, xpr.lhs(), xpr.rhs());
190 }
191
192 protected:
193 PlainObject m_result;
194};
195
196template <typename Lhs, typename Rhs, int ProductTag>
197struct product_evaluator<Product<Lhs, Rhs, AliasFreeProduct>, ProductTag, SparseShape, PermutationShape>
198 : public evaluator<typename permutation_matrix_product<Lhs, OnTheRight, false, SparseShape>::ReturnType> {
199 using XprType = Product<Lhs, Rhs, AliasFreeProduct>;
200 using PlainObject = typename permutation_matrix_product<Lhs, OnTheRight, false, SparseShape>::ReturnType;
201 using Base = evaluator<PlainObject>;
202
203 enum { Flags = Base::Flags | EvalBeforeNestingBit };
204
205 explicit product_evaluator(const XprType& xpr) : m_result(xpr.rows(), xpr.cols()) {
206 ::new (static_cast<Base*>(this)) Base(m_result);
207 generic_product_impl<Lhs, Rhs, SparseShape, PermutationShape, ProductTag>::evalTo(m_result, xpr.lhs(), xpr.rhs());
208 }
209
210 protected:
211 PlainObject m_result;
212};
213
214} // end namespace internal
215
218template <typename SparseDerived, typename PermDerived>
219inline const Product<SparseDerived, PermDerived, AliasFreeProduct> operator*(
220 const SparseMatrixBase<SparseDerived>& matrix, const PermutationBase<PermDerived>& perm) {
221 return Product<SparseDerived, PermDerived, AliasFreeProduct>(matrix.derived(), perm.derived());
222}
223
226template <typename SparseDerived, typename PermDerived>
227inline const Product<PermDerived, SparseDerived, AliasFreeProduct> operator*(
228 const PermutationBase<PermDerived>& perm, const SparseMatrixBase<SparseDerived>& matrix) {
229 return Product<PermDerived, SparseDerived, AliasFreeProduct>(perm.derived(), matrix.derived());
230}
231
234template <typename SparseDerived, typename PermutationType>
235inline const Product<SparseDerived, Inverse<PermutationType>, AliasFreeProduct> operator*(
236 const SparseMatrixBase<SparseDerived>& matrix, const InverseImpl<PermutationType, PermutationStorage>& tperm) {
237 return Product<SparseDerived, Inverse<PermutationType>, AliasFreeProduct>(matrix.derived(), tperm.derived());
238}
239
242template <typename SparseDerived, typename PermutationType>
243inline const Product<Inverse<PermutationType>, SparseDerived, AliasFreeProduct> operator*(
244 const InverseImpl<PermutationType, PermutationStorage>& tperm, const SparseMatrixBase<SparseDerived>& matrix) {
245 return Product<Inverse<PermutationType>, SparseDerived, AliasFreeProduct>(tperm.derived(), matrix.derived());
246}
247
248} // end namespace Eigen
249
250#endif // EIGEN_SPARSE_PERMUTATION_H
@ OnTheLeft
Definition Constants.h:332
@ OnTheRight
Definition Constants.h:334
constexpr unsigned int EvalBeforeNestingBit
Definition Constants.h:75