Eigen  5.0.1
 
Loading...
Searching...
No Matches
PermutationMatrix.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2009 Benoit Jacob <jacob.benoit.1@gmail.com>
5// Copyright (C) 2009-2015 Gael Guennebaud <gael.guennebaud@inria.fr>
6//
7// This Source Code Form is subject to the terms of the Mozilla
8// Public License v. 2.0. If a copy of the MPL was not distributed
9// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
10// SPDX-License-Identifier: MPL-2.0
11
12#ifndef EIGEN_PERMUTATIONMATRIX_H
13#define EIGEN_PERMUTATIONMATRIX_H
14
15// IWYU pragma: private
16#include "./InternalHeaderCheck.h"
17
18namespace Eigen {
19
20namespace internal {
21
22enum PermPermProduct_t { PermPermProduct };
23
29template <typename Scalar, typename IndicesType, bool Transposed>
30struct permutation_dense_op {
31 EIGEN_DEVICE_FUNC explicit permutation_dense_op(const IndicesType& indices) : m_indices(indices) {}
32
33 template <typename IndexType>
34 EIGEN_DEVICE_FUNC Scalar operator()(IndexType row, IndexType col) const {
35 const Index k = Transposed ? Index(row) : Index(col);
36 const Index image = Transposed ? Index(col) : Index(row);
37 return Index(m_indices.coeff(k)) == image ? Scalar(1) : Scalar(0);
38 }
39
40 EIGEN_DEVICE_FUNC const remove_all_t<typename IndicesType::Nested>& indices() const { return m_indices; }
41
42 typename IndicesType::Nested m_indices;
43};
44
45template <typename Scalar, typename IndicesType, bool Transposed>
46struct functor_traits<permutation_dense_op<Scalar, IndicesType, Transposed>> {
47 static constexpr int Cost = int(NumTraits<typename IndicesType::Scalar>::ReadCost) + int(NumTraits<Scalar>::AddCost);
48 static constexpr bool PacketAccess = false;
49 static constexpr bool IsRepeatable = true;
50};
51
52template <typename Scalar, bool Transposed, typename PermutationType>
53struct permutation_dense_expression {
54 using IndicesType = remove_all_t<typename PermutationType::IndicesType>;
55 using PlainObject = Matrix<Scalar, PermutationType::RowsAtCompileTime, PermutationType::ColsAtCompileTime, 0,
56 PermutationType::MaxRowsAtCompileTime, PermutationType::MaxColsAtCompileTime>;
57 using type = CwiseNullaryOp<permutation_dense_op<Scalar, IndicesType, Transposed>, PlainObject>;
58
59 static EIGEN_DEVICE_FUNC type run(const PermutationType& permutation) {
60 return type(permutation.rows(), permutation.cols(),
61 permutation_dense_op<Scalar, IndicesType, Transposed>(permutation.indices()));
62 }
63};
64
65} // end namespace internal
66
91template <typename Derived>
92class PermutationBase : public EigenBase<Derived> {
93 using Traits = internal::traits<Derived>;
94 using Base = EigenBase<Derived>;
95
96 public:
97#ifndef EIGEN_PARSED_BY_DOXYGEN
98 using IndicesType = typename Traits::IndicesType;
99 enum {
100 Flags = Traits::Flags,
101 RowsAtCompileTime = Traits::RowsAtCompileTime,
102 ColsAtCompileTime = Traits::ColsAtCompileTime,
103 MaxRowsAtCompileTime = Traits::MaxRowsAtCompileTime,
104 MaxColsAtCompileTime = Traits::MaxColsAtCompileTime
105 };
106 using StorageIndex = typename Traits::StorageIndex;
107 using DenseMatrixType =
109 using PlainPermutationType =
111 using PlainObject = PlainPermutationType;
112 using Base::derived;
113 using InverseReturnType = Inverse<Derived>;
114 using Scalar = void;
115#endif
116
118 template <typename OtherDerived>
120 indices() = other.indices();
121 return derived();
122 }
123
125 template <typename OtherDerived>
126 Derived& operator=(const TranspositionsBase<OtherDerived>& tr) {
127 setIdentity(tr.size());
128 for (Index k = size() - 1; k >= 0; --k) applyTranspositionOnTheRight(k, tr.coeff(k));
129 return derived();
130 }
131
133 inline EIGEN_DEVICE_FUNC Index rows() const { return Index(indices().size()); }
134
136 inline EIGEN_DEVICE_FUNC Index cols() const { return Index(indices().size()); }
137
139 inline EIGEN_DEVICE_FUNC Index size() const { return Index(indices().size()); }
140
141#ifndef EIGEN_PARSED_BY_DOXYGEN
142 template <typename DenseDerived>
143 void evalTo(MatrixBase<DenseDerived>& other) const {
144 other.setZero();
145 for (Index i = 0; i < rows(); ++i) other.coeffRef(indices().coeff(i), i) = typename DenseDerived::Scalar(1);
146 }
147#endif
148
153 DenseMatrixType toDenseMatrix() const { return derived(); }
154
156 DenseMatrixType eval() const { return toDenseMatrix(); }
157
159 EIGEN_DEVICE_FUNC constexpr const IndicesType& indices() const { return derived().indices(); }
161 EIGEN_DEVICE_FUNC constexpr IndicesType& indices() { return derived().indices(); }
162
165 EIGEN_DEVICE_FUNC void resize(Index newSize) { indices().resize(newSize); }
166
168 EIGEN_DEVICE_FUNC void setIdentity() {
169 StorageIndex n = StorageIndex(size());
170 for (StorageIndex i = 0; i < n; ++i) indices().coeffRef(i) = i;
171 }
172
175 EIGEN_DEVICE_FUNC void setIdentity(Index newSize) {
176 resize(newSize);
177 setIdentity();
178 }
179
190 eigen_assert(i >= 0 && j >= 0 && i < size() && j < size());
191 if (i == j) return derived();
192 EIGEN_IF_CONSTEXPR ((internal::evaluator<IndicesType>::Flags & PacketAccessBit) &&
193 internal::packet_traits<StorageIndex>::HasCmp) {
194 // Amortize packet setup over at least two packets.
195 if (size() >= 2 * internal::packet_traits<StorageIndex>::size) {
196 const StorageIndex first = StorageIndex(i), second = StorageIndex(j);
197 indices() =
198 indices().cwiseTypedEqual(first).select(second, indices().cwiseTypedEqual(second).select(first, indices()));
199 return derived();
200 }
201 }
202 for (Index k = 0; k < size(); ++k) {
203 if (indices().coeff(k) == i)
204 indices().coeffRef(k) = StorageIndex(j);
205 else if (indices().coeff(k) == j)
206 indices().coeffRef(k) = StorageIndex(i);
207 }
208 return derived();
209 }
210
220 eigen_assert(i >= 0 && j >= 0 && i < size() && j < size());
221 std::swap(indices().coeffRef(i), indices().coeffRef(j));
222 return derived();
223 }
224
229 inline InverseReturnType inverse() const { return InverseReturnType(derived()); }
234 inline InverseReturnType transpose() const { return InverseReturnType(derived()); }
240 InverseReturnType adjoint() const { return InverseReturnType(derived()); }
241
242 /**** multiplication helpers to hopefully get RVO ****/
243
244#ifndef EIGEN_PARSED_BY_DOXYGEN
245 protected:
246 template <typename OtherDerived>
247 void assignTranspose(const PermutationBase<OtherDerived>& other) {
248 for (Index i = 0; i < rows(); ++i) indices().coeffRef(other.indices().coeff(i)) = StorageIndex(i);
249 }
250 template <typename Lhs, typename Rhs>
251 void assignProduct(const Lhs& lhs, const Rhs& rhs) {
252 eigen_assert(lhs.cols() == rhs.rows());
253 for (Index i = 0; i < rows(); ++i) indices().coeffRef(i) = lhs.indices().coeff(rhs.indices().coeff(i));
254 }
255#endif
256
257 public:
262 template <typename Other>
263 inline PlainPermutationType operator*(const PermutationBase<Other>& other) const {
264 return PlainPermutationType(internal::PermPermProduct, derived(), other.derived());
265 }
266
271 template <typename Other>
272 inline PlainPermutationType operator*(const InverseImpl<Other, PermutationStorage>& other) const {
273 const auto& rhs = other.derived().nestedExpression();
274 eigen_assert(size() == rhs.size());
275 PlainPermutationType result(size());
276 // (P * Q.inverse())(Q(i)) = P(i).
277 for (Index i = 0; i < size(); ++i) result.indices().coeffRef(rhs.indices().coeff(i)) = indices().coeff(i);
278 return result;
279 }
280
285 template <typename Other>
286 friend inline PlainPermutationType operator*(const InverseImpl<Other, PermutationStorage>& other,
287 const PermutationBase& perm) {
288 return PlainPermutationType(internal::PermPermProduct, other.eval(), perm);
289 }
290
297 Index res = 1;
298 Index n = size();
300 mask.fill(false);
301 Index r = 0;
302 while (r < n) {
303 // search for the next seed
304 while (r < n && mask[r]) r++;
305 if (r >= n) break;
306 // we got one, let's follow it until we are back to the seed
307 Index k0 = r++;
308 mask.coeffRef(k0) = true;
309 for (Index k = indices().coeff(k0); k != k0; k = indices().coeff(k)) {
310 mask.coeffRef(k) = true;
311 res = -res;
312 }
313 }
314 return res;
315 }
316};
317
318namespace internal {
319template <int SizeAtCompileTime, int MaxSizeAtCompileTime, typename StorageIndex_>
320struct traits<PermutationMatrix<SizeAtCompileTime, MaxSizeAtCompileTime, StorageIndex_> >
321 : traits<
322 Matrix<StorageIndex_, SizeAtCompileTime, SizeAtCompileTime, 0, MaxSizeAtCompileTime, MaxSizeAtCompileTime> > {
323 using StorageKind = PermutationStorage;
324 using IndicesType = Matrix<StorageIndex_, SizeAtCompileTime, 1, 0, MaxSizeAtCompileTime, 1>;
325 using StorageIndex = StorageIndex_;
326 using Scalar = void;
327};
328} // namespace internal
329
344template <int SizeAtCompileTime, int MaxSizeAtCompileTime, typename StorageIndex_>
345class PermutationMatrix
346 : public PermutationBase<PermutationMatrix<SizeAtCompileTime, MaxSizeAtCompileTime, StorageIndex_> > {
348 using Traits = internal::traits<PermutationMatrix>;
349
350 public:
351 using Nested = const PermutationMatrix&;
352
353#ifndef EIGEN_PARSED_BY_DOXYGEN
354 using IndicesType = typename Traits::IndicesType;
355 using StorageIndex = typename Traits::StorageIndex;
356#endif
357
358 EIGEN_DEVICE_FUNC PermutationMatrix() = default;
359
362 EIGEN_DEVICE_FUNC explicit PermutationMatrix(Index size) : m_indices(size) {
363 eigen_internal_assert(size <= NumTraits<StorageIndex>::highest());
364 }
365
367 template <typename OtherDerived>
368 EIGEN_DEVICE_FUNC PermutationMatrix(const PermutationBase<OtherDerived>& other) : m_indices(other.indices()) {}
369
377 template <typename Other>
378 EIGEN_DEVICE_FUNC explicit PermutationMatrix(const MatrixBase<Other>& indices) : m_indices(indices) {}
379
381 template <typename Other>
382 explicit PermutationMatrix(const TranspositionsBase<Other>& tr) : m_indices(tr.size()) {
383 *this = tr;
384 }
385
387 template <typename Other>
388 EIGEN_DEVICE_FUNC PermutationMatrix& operator=(const PermutationBase<Other>& other) {
389 m_indices = other.indices();
390 return *this;
391 }
392
394 template <typename Other>
395 PermutationMatrix& operator=(const TranspositionsBase<Other>& tr) {
396 return Base::operator=(tr.derived());
397 }
398
400 EIGEN_DEVICE_FUNC constexpr const IndicesType& indices() const { return m_indices; }
402 EIGEN_DEVICE_FUNC constexpr IndicesType& indices() { return m_indices; }
403
404 /**** multiplication helpers to hopefully get RVO ****/
405
406#ifndef EIGEN_PARSED_BY_DOXYGEN
407 template <typename Other>
408 PermutationMatrix(const InverseImpl<Other, PermutationStorage>& other)
409 : m_indices(other.derived().nestedExpression().size()) {
410 eigen_internal_assert(m_indices.size() <= NumTraits<StorageIndex>::highest());
411 Base::assignTranspose(other.derived().nestedExpression());
412 }
413 template <typename Lhs, typename Rhs>
414 PermutationMatrix(internal::PermPermProduct_t, const Lhs& lhs, const Rhs& rhs) : m_indices(lhs.indices().size()) {
415 Base::assignProduct(lhs, rhs);
416 }
417#endif
418
419 protected:
420 IndicesType m_indices;
421};
422
423namespace internal {
424template <int SizeAtCompileTime, int MaxSizeAtCompileTime, typename StorageIndex_, int PacketAccess_>
425struct traits<Map<PermutationMatrix<SizeAtCompileTime, MaxSizeAtCompileTime, StorageIndex_>, PacketAccess_> >
426 : traits<
427 Matrix<StorageIndex_, SizeAtCompileTime, SizeAtCompileTime, 0, MaxSizeAtCompileTime, MaxSizeAtCompileTime> > {
428 using StorageKind = PermutationStorage;
429 using IndicesType = Map<const Matrix<StorageIndex_, SizeAtCompileTime, 1, 0, MaxSizeAtCompileTime, 1>, PacketAccess_>;
430 using StorageIndex = StorageIndex_;
431 using Scalar = void;
432};
433} // namespace internal
434
435template <int SizeAtCompileTime, int MaxSizeAtCompileTime, typename StorageIndex_, int PacketAccess_>
436class Map<PermutationMatrix<SizeAtCompileTime, MaxSizeAtCompileTime, StorageIndex_>, PacketAccess_>
437 : public PermutationBase<
438 Map<PermutationMatrix<SizeAtCompileTime, MaxSizeAtCompileTime, StorageIndex_>, PacketAccess_> > {
439 using Base = PermutationBase<Map>;
440 using Traits = internal::traits<Map>;
441
442 public:
443#ifndef EIGEN_PARSED_BY_DOXYGEN
444 using IndicesType = typename Traits::IndicesType;
445 using StorageIndex = typename IndicesType::Scalar;
446#endif
447
448 inline Map(const StorageIndex* indicesPtr) : m_indices(indicesPtr) {}
449
450 inline Map(const StorageIndex* indicesPtr, Index size) : m_indices(indicesPtr, size) {}
451
453 template <typename Other>
454 Map& operator=(const PermutationBase<Other>& other) {
455 return Base::operator=(other.derived());
456 }
457
459 template <typename Other>
460 Map& operator=(const TranspositionsBase<Other>& tr) {
461 return Base::operator=(tr.derived());
462 }
463
464#ifndef EIGEN_PARSED_BY_DOXYGEN
468 Map& operator=(const Map& other) {
469 m_indices = other.m_indices;
470 return *this;
471 }
472#endif
473
475 const IndicesType& indices() const { return m_indices; }
477 IndicesType& indices() { return m_indices; }
478
479 protected:
480 IndicesType m_indices;
481};
482
483namespace internal {
484template <typename IndicesType_>
485struct traits<PermutationWrapper<IndicesType_> > {
486 using StorageKind = PermutationStorage;
487 using Scalar = void;
488 using StorageIndex = typename IndicesType_::Scalar;
489 using IndicesType = IndicesType_;
490 enum {
491 RowsAtCompileTime = IndicesType_::SizeAtCompileTime,
492 ColsAtCompileTime = IndicesType_::SizeAtCompileTime,
493 MaxRowsAtCompileTime = IndicesType::MaxSizeAtCompileTime,
494 MaxColsAtCompileTime = IndicesType::MaxSizeAtCompileTime,
495 Flags = 0
496 };
497};
498} // namespace internal
499
511template <typename IndicesType_>
512class PermutationWrapper : public PermutationBase<PermutationWrapper<IndicesType_> > {
514 using Traits = internal::traits<PermutationWrapper>;
515
516 public:
517#ifndef EIGEN_PARSED_BY_DOXYGEN
518 using IndicesType = typename Traits::IndicesType;
519#endif
520
521 inline PermutationWrapper(const IndicesType& indices) : m_indices(indices) {}
522
524 const internal::remove_all_t<typename IndicesType::Nested>& indices() const { return m_indices; }
525
526 protected:
527 typename IndicesType::Nested m_indices;
528};
529
532template <typename MatrixDerived, typename PermutationDerived>
533EIGEN_DEVICE_FUNC const Product<MatrixDerived, PermutationDerived, DefaultProduct> operator*(
534 const MatrixBase<MatrixDerived>& matrix, const PermutationBase<PermutationDerived>& permutation) {
535 return Product<MatrixDerived, PermutationDerived, DefaultProduct>(matrix.derived(), permutation.derived());
536}
537
540template <typename PermutationDerived, typename MatrixDerived>
541EIGEN_DEVICE_FUNC const Product<PermutationDerived, MatrixDerived, DefaultProduct> operator*(
542 const PermutationBase<PermutationDerived>& permutation, const MatrixBase<MatrixDerived>& matrix) {
543 return Product<PermutationDerived, MatrixDerived, DefaultProduct>(permutation.derived(), matrix.derived());
544}
545
546// Sums with a permutation are lazy dense expressions: the permutation is read as a 0/1 matrix with the scalar
547// type of the other operand, so no dense copy of it is formed.
548
550template <typename MatrixDerived, typename PermutationDerived>
551EIGEN_DEVICE_FUNC auto operator+(const MatrixBase<MatrixDerived>& matrix,
552 const PermutationBase<PermutationDerived>& permutation) {
553 return matrix.derived() +
554 internal::permutation_dense_expression<typename MatrixDerived::Scalar, false, PermutationDerived>::run(
555 permutation.derived());
556}
557
559template <typename PermutationDerived, typename MatrixDerived>
560EIGEN_DEVICE_FUNC auto operator+(const PermutationBase<PermutationDerived>& permutation,
561 const MatrixBase<MatrixDerived>& matrix) {
562 return internal::permutation_dense_expression<typename MatrixDerived::Scalar, false, PermutationDerived>::run(
563 permutation.derived()) +
564 matrix.derived();
565}
566
568template <typename MatrixDerived, typename PermutationDerived>
569EIGEN_DEVICE_FUNC auto operator-(const MatrixBase<MatrixDerived>& matrix,
570 const PermutationBase<PermutationDerived>& permutation) {
571 return matrix.derived() -
572 internal::permutation_dense_expression<typename MatrixDerived::Scalar, false, PermutationDerived>::run(
573 permutation.derived());
574}
575
577template <typename PermutationDerived, typename MatrixDerived>
578EIGEN_DEVICE_FUNC auto operator-(const PermutationBase<PermutationDerived>& permutation,
579 const MatrixBase<MatrixDerived>& matrix) {
580 return internal::permutation_dense_expression<typename MatrixDerived::Scalar, false, PermutationDerived>::run(
581 permutation.derived()) -
582 matrix.derived();
583}
584
586template <typename DiagonalDerived, typename PermutationDerived>
587EIGEN_DEVICE_FUNC auto operator+(const DiagonalBase<DiagonalDerived>& diagonal,
588 const PermutationBase<PermutationDerived>& permutation) {
589 return diagonal.derived() +
590 internal::permutation_dense_expression<typename DiagonalDerived::Scalar, false, PermutationDerived>::run(
591 permutation.derived());
592}
593
595template <typename PermutationDerived, typename DiagonalDerived>
596EIGEN_DEVICE_FUNC auto operator+(const PermutationBase<PermutationDerived>& permutation,
597 const DiagonalBase<DiagonalDerived>& diagonal) {
598 return internal::permutation_dense_expression<typename DiagonalDerived::Scalar, false, PermutationDerived>::run(
599 permutation.derived()) +
600 diagonal.derived();
601}
602
604template <typename DiagonalDerived, typename PermutationDerived>
605EIGEN_DEVICE_FUNC auto operator-(const DiagonalBase<DiagonalDerived>& diagonal,
606 const PermutationBase<PermutationDerived>& permutation) {
607 return diagonal.derived() -
608 internal::permutation_dense_expression<typename DiagonalDerived::Scalar, false, PermutationDerived>::run(
609 permutation.derived());
610}
611
613template <typename PermutationDerived, typename DiagonalDerived>
614EIGEN_DEVICE_FUNC auto operator-(const PermutationBase<PermutationDerived>& permutation,
615 const DiagonalBase<DiagonalDerived>& diagonal) {
616 return internal::permutation_dense_expression<typename DiagonalDerived::Scalar, false, PermutationDerived>::run(
617 permutation.derived()) -
618 diagonal.derived();
619}
620
621template <typename PermutationType>
622class InverseImpl<PermutationType, PermutationStorage> : public EigenBase<Inverse<PermutationType> > {
623 using PlainPermutationType = typename PermutationType::PlainPermutationType;
624 using PermTraits = internal::traits<PermutationType>;
625
626 protected:
627 InverseImpl() = default;
628
629 public:
630 using InverseType = Inverse<PermutationType>;
631 using EigenBase<Inverse<PermutationType> >::derived;
632
633#ifndef EIGEN_PARSED_BY_DOXYGEN
634 using DenseMatrixType = typename PermutationType::DenseMatrixType;
635 enum {
636 RowsAtCompileTime = PermTraits::RowsAtCompileTime,
637 ColsAtCompileTime = PermTraits::ColsAtCompileTime,
638 MaxRowsAtCompileTime = PermTraits::MaxRowsAtCompileTime,
639 MaxColsAtCompileTime = PermTraits::MaxColsAtCompileTime
640 };
641#endif
642
643#ifndef EIGEN_PARSED_BY_DOXYGEN
644 template <typename DenseDerived>
645 void evalTo(MatrixBase<DenseDerived>& other) const {
646 other.setZero();
647 for (Index i = 0; i < derived().rows(); ++i)
648 other.coeffRef(i, derived().nestedExpression().indices().coeff(i)) = typename DenseDerived::Scalar(1);
649 }
650#endif
651
653 PlainPermutationType eval() const { return derived(); }
654
655 DenseMatrixType toDenseMatrix() const { return derived(); }
656
659 template <typename OtherDerived>
660 friend const Product<OtherDerived, InverseType, DefaultProduct> operator*(const MatrixBase<OtherDerived>& matrix,
661 const InverseType& trPerm) {
662 return Product<OtherDerived, InverseType, DefaultProduct>(matrix.derived(), trPerm.derived());
663 }
664
667 template <typename OtherDerived>
668 const Product<InverseType, OtherDerived, DefaultProduct> operator*(const MatrixBase<OtherDerived>& matrix) const {
669 return Product<InverseType, OtherDerived, DefaultProduct>(derived(), matrix.derived());
670 }
671
672 // Lazy sums with a dense or diagonal matrix, as for PermutationBase; the inverse is read as the transposed 0/1
673 // matrix.
675 template <typename OtherDerived>
676 EIGEN_DEVICE_FUNC auto operator+(const MatrixBase<OtherDerived>& matrix) const {
677 return denseExpression<typename OtherDerived::Scalar>() + matrix.derived();
678 }
680 template <typename OtherDerived>
681 EIGEN_DEVICE_FUNC friend auto operator+(const MatrixBase<OtherDerived>& matrix, const InverseType& inverse) {
682 return matrix.derived() + inverse.template denseExpression<typename OtherDerived::Scalar>();
683 }
685 template <typename OtherDerived>
686 EIGEN_DEVICE_FUNC auto operator-(const MatrixBase<OtherDerived>& matrix) const {
687 return denseExpression<typename OtherDerived::Scalar>() - matrix.derived();
688 }
690 template <typename OtherDerived>
691 EIGEN_DEVICE_FUNC friend auto operator-(const MatrixBase<OtherDerived>& matrix, const InverseType& inverse) {
692 return matrix.derived() - inverse.template denseExpression<typename OtherDerived::Scalar>();
693 }
695 template <typename OtherDerived>
696 EIGEN_DEVICE_FUNC auto operator+(const DiagonalBase<OtherDerived>& diagonal) const {
697 return denseExpression<typename OtherDerived::Scalar>() + diagonal.derived();
698 }
700 template <typename OtherDerived>
701 EIGEN_DEVICE_FUNC friend auto operator+(const DiagonalBase<OtherDerived>& diagonal, const InverseType& inverse) {
702 return diagonal.derived() + inverse.template denseExpression<typename OtherDerived::Scalar>();
703 }
705 template <typename OtherDerived>
706 EIGEN_DEVICE_FUNC auto operator-(const DiagonalBase<OtherDerived>& diagonal) const {
707 return denseExpression<typename OtherDerived::Scalar>() - diagonal.derived();
708 }
710 template <typename OtherDerived>
711 EIGEN_DEVICE_FUNC friend auto operator-(const DiagonalBase<OtherDerived>& diagonal, const InverseType& inverse) {
712 return diagonal.derived() - inverse.template denseExpression<typename OtherDerived::Scalar>();
713 }
714
715 private:
716 template <typename Scalar>
717 EIGEN_DEVICE_FUNC auto denseExpression() const {
718 using Nested = internal::remove_all_t<typename InverseType::XprTypeNestedCleaned>;
719 return internal::permutation_dense_expression<Scalar, true, Nested>::run(derived().nestedExpression());
720 }
721};
722
723template <typename Derived>
724const PermutationWrapper<const Derived> MatrixBase<Derived>::asPermutation() const {
725 return derived();
726}
727
728namespace internal {
729
730template <>
731struct AssignmentKind<DenseShape, PermutationShape> {
732 using Kind = EigenBase2EigenBase;
733};
734
735// Dense ?= (dense or diagonal) +/- permutation, in either order. The permutation's one in column k is at row
736// indices(k), and at column indices(k) of row k for its transpose; the column kernels below need one nonzero per
737// destination column, so P pairs with a column-major and P^T with a row-major destination, and the other pairings
738// stay on the generic path.
739template <typename Scalar, typename IndicesType>
740struct permutation_column_nonzeros {
741 EIGEN_DEVICE_FUNC explicit permutation_column_nonzeros(const IndicesType& indices) : m_indices(indices) {}
742 EIGEN_DEVICE_FUNC Index row(Index k) const { return Index(m_indices.coeff(k)); }
743 EIGEN_DEVICE_FUNC Scalar value(Index) const { return Scalar(1); }
744 const IndicesType& m_indices;
745};
746
747template <typename Scalar, typename IndicesType, bool Transposed, typename Plain>
748using permutation_dense_xpr = CwiseNullaryOp<permutation_dense_op<Scalar, IndicesType, Transposed>, Plain>;
749
750template <typename Scalar, typename IndicesType, bool Transposed, typename Plain>
751struct is_permutation_dense_xpr<permutation_dense_xpr<Scalar, IndicesType, Transposed, Plain>> : std::true_type {};
752
753// A diagonal operand never goes through the block pass, so only a dense one is subject to
754// dense_block_pass_cannot_overflow.
755template <typename Dst, typename OtherXpr, bool Transposed, typename Functor, bool NegateOther>
756struct dense_permutation_sum_fast_path
757 : bool_constant<(is_dense_shape<OtherXpr>::value || is_diagonal_shape<OtherXpr>::value) &&
758 !is_permutation_dense_xpr<OtherXpr>::value && bool(Dst::IsRowMajor) == Transposed &&
759 additive_assign_sign<Functor>::value != 0 &&
760 std::is_same<typename Dst::Scalar, typename OtherXpr::Scalar>::value &&
761 (is_diagonal_shape<OtherXpr>::value ||
762 dense_block_pass_cannot_overflow<typename Dst::Scalar, Functor, NegateOther>::value)> {};
763
764// dst ?= op(lhs, rhs) for a diagonal d and a permutation whose one in column k is at row r: op(d(k), [r == k]) at
765// (k, k), op(0, 1) or op(1, 0) at (r, k) when r != k, and structural zeros, which only = writes. A block's d(k) are
766// read before the block is written.
767template <bool DiagonalOnLeft>
768struct diagonal_permutation_sum_assignment {
769 template <typename Dst, typename Diagonal, typename Permutation, typename BinaryOp, typename Functor>
770 EIGEN_DEVICE_FUNC static void run(Dst& dst, const Diagonal& diagonal, const Permutation& permutation,
771 const BinaryOp& op, const Functor& func) {
772 using Scalar = typename Dst::Scalar;
773 constexpr Index kBlockColumns = 32;
774 const Scalar offDiagonalOne = DiagonalOnLeft ? op(Scalar(0), Scalar(1)) : op(Scalar(1), Scalar(0));
775 Scalar atDiagonal[kBlockColumns];
776 for (Index j = 0; j < dst.cols(); j += kBlockColumns) {
777 const Index columns = numext::mini(Index(kBlockColumns), dst.cols() - j);
778 for (Index k = 0; k < columns; ++k) {
779 const Scalar d = diagonal.value(j + k);
780 const Scalar p = permutation.row(j + k) == j + k ? Scalar(1) : Scalar(0);
781 atDiagonal[k] = DiagonalOnLeft ? op(d, p) : op(p, d);
782 }
783 EIGEN_IF_CONSTEXPR (is_plain_assign<Functor>::value) {
784 dst.middleCols(j, columns).setZero();
785 }
786 for (Index k = 0; k < columns; ++k) {
787 const Index r = permutation.row(j + k);
788 func.assignCoeff(dst.coeffRef(j + k, j + k), atDiagonal[k]);
789 if (r != j + k) {
790 func.assignCoeff(dst.coeffRef(r, j + k), offDiagonalOne);
791 }
792 }
793 }
794 }
795};
796
797template <bool OtherIsDiagonal, bool OtherOnLeft>
798struct permutation_sum_assignment {
799 template <typename Dst, typename SrcXprType, typename OtherXpr, typename Nonzeros, typename Functor>
800 EIGEN_DEVICE_FUNC static void run(Dst& dst, const SrcXprType& src, const OtherXpr& other, const Nonzeros& ones,
801 const Functor& func) {
802 dense_structured_sum_assignment<OtherOnLeft>::assign(dst, src, other, ones, func);
803 }
804};
805
806template <bool OtherOnLeft>
807struct permutation_sum_assignment<true, OtherOnLeft> {
808 template <typename Dst, typename SrcXprType, typename OtherXpr, typename Nonzeros, typename Functor>
809 EIGEN_DEVICE_FUNC static void run(Dst& dst, const SrcXprType& src, const OtherXpr& other, const Nonzeros& ones,
810 const Functor& func) {
811 const diagonal_column_nonzeros<OtherXpr> diagonal(other);
812 resize_if_allowed(dst, src, func);
813 auto&& dstView = column_major_view<bool(Dst::IsRowMajor)>::run(dst);
814 diagonal_permutation_sum_assignment<OtherOnLeft>::run(dstView, diagonal, ones, src.functor(), func);
815 }
816};
817
818// other +/- permutation
819template <typename DstXprType, typename BinaryOp, typename OtherXpr, typename Scalar, typename IndicesType,
820 bool Transposed, typename Plain, typename Functor>
821struct Assignment<
822 DstXprType,
823 CwiseBinaryOp<BinaryOp, const OtherXpr, const permutation_dense_xpr<Scalar, IndicesType, Transposed, Plain>>,
824 Functor, Dense2Dense,
825 std::enable_if_t<is_additive_binary_op<BinaryOp>::value &&
826 dense_permutation_sum_fast_path<DstXprType, OtherXpr, Transposed, Functor, false>::value>> {
827 using SrcXprType =
828 CwiseBinaryOp<BinaryOp, const OtherXpr, const permutation_dense_xpr<Scalar, IndicesType, Transposed, Plain>>;
829 EIGEN_DEVICE_FUNC static void run(DstXprType& dst, const SrcXprType& src, const Functor& func) {
830 const auto& indices = src.rhs().functor().indices();
831 permutation_sum_assignment<is_diagonal_shape<OtherXpr>::value, true>::run(
832 dst, src, src.lhs(), permutation_column_nonzeros<Scalar, remove_all_t<decltype(indices)>>(indices), func);
833 }
834};
835
836// permutation +/- other; a dense product on the right is left to the "xpr + product" rule.
837template <typename DstXprType, typename BinaryOp, typename OtherXpr, typename Scalar, typename IndicesType,
838 bool Transposed, typename Plain, typename Functor>
839struct Assignment<
840 DstXprType,
841 CwiseBinaryOp<BinaryOp, const permutation_dense_xpr<Scalar, IndicesType, Transposed, Plain>, const OtherXpr>,
842 Functor, Dense2Dense,
843 std::enable_if_t<is_additive_binary_op<BinaryOp>::value &&
844 dense_permutation_sum_fast_path<DstXprType, OtherXpr, Transposed, Functor,
845 is_difference_op<BinaryOp>::value>::value &&
846 !is_default_product<OtherXpr>::value>> {
847 using SrcXprType =
848 CwiseBinaryOp<BinaryOp, const permutation_dense_xpr<Scalar, IndicesType, Transposed, Plain>, const OtherXpr>;
849 EIGEN_DEVICE_FUNC static void run(DstXprType& dst, const SrcXprType& src, const Functor& func) {
850 const auto& indices = src.lhs().functor().indices();
851 permutation_sum_assignment<is_diagonal_shape<OtherXpr>::value, false>::run(
852 dst, src, src.rhs(), permutation_column_nonzeros<Scalar, remove_all_t<decltype(indices)>>(indices), func);
853 }
854};
855
856} // end namespace internal
857
858} // end namespace Eigen
859
860#endif // EIGEN_PERMUTATIONMATRIX_H
Derived & setZero()
Definition CwiseNullaryOp.h:521
Base class for diagonal matrices and expressions.
Definition DiagonalMatrix.h:34
Expression of the inverse of another expression.
Definition Inverse.h:44
A matrix or vector expression mapping an existing array of data.
Definition Map.h:97
constexpr Map(PointerArgType dataPtr, const StrideType &stride=StrideType())
Definition Map.h:124
Base class for all dense matrices, vectors, and expressions.
Definition MatrixBase.h:53
The matrix class, also used for vectors and row-vectors.
Definition Matrix.h:188
constexpr Scalar & coeffRef(Index rowId, Index colId)
Definition PlainObjectBase.h:205
Base class for permutations.
Definition PermutationMatrix.h:92
InverseReturnType transpose() const
Definition PermutationMatrix.h:234
Derived & applyTranspositionOnTheLeft(Index i, Index j)
Definition PermutationMatrix.h:189
void resize(Index newSize)
Definition PermutationMatrix.h:165
Index determinant() const
Definition PermutationMatrix.h:296
Index size() const
Definition PermutationMatrix.h:139
Index cols() const
Definition PermutationMatrix.h:136
InverseReturnType adjoint() const
Definition PermutationMatrix.h:240
DenseMatrixType eval() const
Definition PermutationMatrix.h:156
Derived & applyTranspositionOnTheRight(Index i, Index j)
Definition PermutationMatrix.h:219
friend PlainPermutationType operator*(const InverseImpl< Other, PermutationStorage > &other, const PermutationBase &perm)
Definition PermutationMatrix.h:286
constexpr const IndicesType & indices() const
Definition PermutationMatrix.h:159
void setIdentity()
Definition PermutationMatrix.h:168
Derived & operator=(const PermutationBase< OtherDerived > &other)
Definition PermutationMatrix.h:119
void setIdentity(Index newSize)
Definition PermutationMatrix.h:175
PlainPermutationType operator*(const InverseImpl< Other, PermutationStorage > &other) const
Definition PermutationMatrix.h:272
constexpr IndicesType & indices()
Definition PermutationMatrix.h:161
Index rows() const
Definition PermutationMatrix.h:133
InverseReturnType inverse() const
Definition PermutationMatrix.h:229
DenseMatrixType toDenseMatrix() const
Definition PermutationMatrix.h:153
Derived & operator=(const TranspositionsBase< OtherDerived > &tr)
Definition PermutationMatrix.h:126
PlainPermutationType operator*(const PermutationBase< Other > &other) const
Definition PermutationMatrix.h:263
Permutation matrix.
Definition PermutationMatrix.h:346
PermutationMatrix(const PermutationBase< OtherDerived > &other)
Definition PermutationMatrix.h:368
PermutationMatrix(const TranspositionsBase< Other > &tr)
Definition PermutationMatrix.h:382
PermutationMatrix & operator=(const PermutationBase< Other > &other)
Definition PermutationMatrix.h:388
constexpr const IndicesType & indices() const
Definition PermutationMatrix.h:400
PermutationMatrix(const MatrixBase< Other > &indices)
Definition PermutationMatrix.h:378
PermutationMatrix(Index size)
Definition PermutationMatrix.h:362
constexpr IndicesType & indices()
Definition PermutationMatrix.h:402
PermutationMatrix & operator=(const TranspositionsBase< Other > &tr)
Definition PermutationMatrix.h:395
Class to view a vector of integers as a permutation matrix.
Definition PermutationMatrix.h:512
const internal::remove_all_t< typename IndicesType::Nested > & indices() const
Definition PermutationMatrix.h:524
Expression of the product of two arbitrary matrices or vectors.
Definition Product.h:203
constexpr unsigned int PacketAccessBit
Definition Constants.h:98
Definition EigenBase.h:34
constexpr Derived & derived()
Definition EigenBase.h:50
Eigen::Index Index
The interface type of indices.
Definition EigenBase.h:44
Definition Constants.h:551