11#ifndef EIGEN_SPARSE_PERMUTATION_H
12#define EIGEN_SPARSE_PERMUTATION_H
17#include "./InternalHeaderCheck.h"
23template <
typename ExpressionType,
typename PlainObjectType,
24 bool NeedEval = !std::is_same<ExpressionType, PlainObjectType>::value>
26 XprHelper(
const ExpressionType& xpr) : m_xpr(xpr) {}
27 inline const PlainObjectType& xpr()
const {
return m_xpr; }
29 const PlainObjectType m_xpr;
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; }
36 const PlainObjectType& m_xpr;
39template <
typename PermDerived,
bool NeedInverseEval>
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; }
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; }
58template <
typename ExpressionType,
int S
ide,
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>;
63 using Scalar =
typename MatrixTypeCleaned::Scalar;
64 using StorageIndex =
typename MatrixTypeCleaned::StorageIndex;
67 using ReturnType = SparseMatrix<Scalar, MatrixTypeCleaned::IsRowMajor ? RowMajor : ColMajor, StorageIndex>;
68 using TmpHelper = XprHelper<ExpressionType, ReturnType>;
70 static constexpr bool NeedOuterPermutation = ExpressionType::IsRowMajor ? Side ==
OnTheLeft : Side ==
OnTheRight;
71 static constexpr bool NeedInversePermutation = Transposed ? Side ==
OnTheLeft : Side ==
OnTheRight;
73 template <
typename Dest,
typename PermutationType>
74 static inline void permute_outer(Dest& dst,
const PermutationType& perm,
const ExpressionType& xpr) {
78 const TmpHelper tmpHelper(xpr);
79 const ReturnType& tmp = tmpHelper.xpr();
81 ReturnType result(tmp.rows(), tmp.cols());
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;
92 std::partial_sum(result.outerIndexPtr(), result.outerIndexPtr() + result.outerSize() + 1, result.outerIndexPtr());
93 result.resizeNonZeros(result.nonZeros());
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);
105 dst = std::move(result);
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;
116 const TmpHelper tmpHelper(xpr);
117 const ReturnType& tmp = tmpHelper.xpr();
121 const InnerPermHelper permHelper(perm);
122 const InnerPermType& innerPerm = permHelper.perm();
124 ReturnType result(tmp.rows(), tmp.cols());
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;
132 std::partial_sum(result.outerIndexPtr(), result.outerIndexPtr() + result.outerSize() + 1, result.outerIndexPtr());
133 result.resizeNonZeros(result.nonZeros());
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);
144 result.sortInnerIndices();
145 dst = std::move(result);
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);
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);
165template <
int ProductTag>
166struct product_promote_storage_type<Sparse, PermutationStorage, ProductTag> {
169template <
int ProductTag>
170struct product_promote_storage_type<PermutationStorage, Sparse, ProductTag> {
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>;
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());
193 PlainObject m_result;
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>;
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());
211 PlainObject m_result;
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());
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());
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());
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());
@ OnTheLeft
Definition Constants.h:332
@ OnTheRight
Definition Constants.h:334
constexpr unsigned int EvalBeforeNestingBit
Definition Constants.h:75