Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
KroneckerSum.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// This Source Code Form is subject to the terms of the Mozilla
5// Public License v. 2.0. If a copy of the MPL was not distributed
6// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
7// SPDX-FileCopyrightText: The Eigen Authors
8// SPDX-License-Identifier: MPL-2.0
9
10// References:
11// [1] R. H. Bartels and G. W. Stewart, "Solution of the matrix equation
12// AX + XB = C", Communications of the ACM 15 (1972), 820-826. The Schur
13// form solve of the Sylvester equation B X + X A^T = mat(b) that
14// BartelsStewart generalizes to Kronecker sums of any number of factors.
15// [2] R. E. Lynch, J. R. Rice and D. H. Thomas, "Direct solution of partial
16// difference equations by tensor product methods", Numerische Mathematik 6
17// (1964), 185-199. The fast diagonalization method: with Hermitian factors
18// the Schur forms are diagonal and the solve is a division by the
19// eigenvalue sums.
20
21#ifndef EIGEN_STRUCTURED_KRONECKER_SUM_H
22#define EIGEN_STRUCTURED_KRONECKER_SUM_H
23
24// IWYU pragma: private
25#include "./InternalHeaderCheck.h"
26
27namespace Eigen {
28
29template <typename KroneckerSumType>
30class BartelsStewart;
31
32namespace internal {
33
34template <typename LhsMatrix, typename RhsMatrix>
35struct traits<KroneckerSum<LhsMatrix, RhsMatrix>> {
36 using Scalar = typename LhsMatrix::Scalar;
37 using StorageKind = Dense;
38 using XprKind = MatrixXpr;
39 using StorageIndex = int;
40 static constexpr int RowsAtCompileTime =
41 size_at_compile_time(traits<LhsMatrix>::RowsAtCompileTime, traits<RhsMatrix>::RowsAtCompileTime);
42 static constexpr int ColsAtCompileTime = RowsAtCompileTime;
43 static constexpr int MaxRowsAtCompileTime = RowsAtCompileTime;
44 static constexpr int MaxColsAtCompileTime = RowsAtCompileTime;
45 // No NestByRefBit, for the reason given at traits<KroneckerOperator>.
46 static constexpr unsigned int Flags = 0;
47};
48
49template <typename LhsMatrix, typename RhsMatrix>
50struct evaluator_traits<KroneckerSum<LhsMatrix, RhsMatrix>> {
51 using Kind = IndexBased;
52 using Shape = StructuredShape;
53};
54
55template <typename KroneckerSumType>
56struct traits<BartelsStewart<KroneckerSumType>> : traits<Matrix<typename KroneckerSumType::Scalar, Dynamic, Dynamic>> {
57 using XprKind = MatrixXpr;
58 using StorageKind = SolverStorage;
59 using StorageIndex = int;
60 using BaseTraits = traits<Matrix<typename KroneckerSumType::Scalar, Dynamic, Dynamic>>;
61 static constexpr unsigned int Flags = BaseTraits::Flags & RowMajorBit;
62 static constexpr int CoeffReadCost = Dynamic;
63};
64
70template <typename LhsMatrix, typename RhsMatrix>
71struct kron_factor_ops<KroneckerSum<LhsMatrix, RhsMatrix>, kKronSumFactor> {
72 using Factor = KroneckerSum<LhsMatrix, RhsMatrix>;
73 using Scalar = typename Factor::Scalar;
74 using LhsOps = kron_factor_ops<LhsMatrix>;
75 using RhsOps = kron_factor_ops<RhsMatrix>;
76 using DenseMatrix = Matrix<Scalar, Dynamic, Dynamic, ColMajor>;
77 using SparseFactor = SparseMatrix<Scalar>;
78 using TransposedFactor = KroneckerSum<typename LhsOps::TransposedFactor, typename RhsOps::TransposedFactor>;
79 using InverseFactor = DenseMatrix;
80 static constexpr bool StoresAllEntries = false;
81
82 static void prepare(Factor&) {}
83 static bool coeffIfStored(const Factor& f, Index row, Index col, Scalar& value) {
84 const Index n2 = f.rhs().rows();
85 Scalar a(0), b(0);
86 const bool storedA = row % n2 == col % n2 && LhsOps::coeffIfStored(f.lhs(), row / n2, col / n2, a);
87 const bool storedB = row / n2 == col / n2 && RhsOps::coeffIfStored(f.rhs(), row % n2, col % n2, b);
88 if (!storedA && !storedB) return false;
89 value = a + b;
90 return true;
91 }
92 template <typename Visitor>
93 static void forEachNonZero(const Factor& f, Visitor&& visit) {
94 kron_factor_ops<SparseFactor>::forEachNonZero(kron_factor_visitable<Factor>::get(f), visit);
95 }
96 static Matrix<Index, Dynamic, 1> innerNonZeros(const Factor& f, bool rowMajor) {
97 return kron_factor_ops<SparseFactor>::innerNonZeros(kron_factor_visitable<Factor>::get(f), rowMajor);
98 }
99 static TransposedFactor transposed(const Factor& f) { return f.transpose(); }
100 static Factor conjugated(const Factor& f) { return f.conjugate(); }
101 static TransposedFactor adjointed(const Factor& f) { return f.adjoint(); }
102 static InverseFactor inversed(const Factor& f) { return f.solve(DenseMatrix::Identity(f.rows(), f.cols())); }
103 static DenseMatrix denseFactor(const Factor& f) { return DenseMatrix(f); }
104 static DenseMatrix blockOperand(const Factor& f) { return denseFactor(f); }
105 static bool isSquareIdentity(const Factor&) { return false; }
106 template <typename Dst, typename Alpha, typename Xpr>
107 static void addLeftProduct(Dst& dst, const Alpha& alpha, const Factor& f, const Xpr& X) {
108 f.addProduct(dst, X, alpha);
109 }
110 // With X_j = X(:, j n2 : (j+1) n2 - 1), j < n1, the p x n2 slices of X,
111 // X (L (+) R)^T = X (L (x) I)^T + X (I (x) R)^T = X_{[p n2 x n1]} L^T + [X_j R^T]_j.
112 template <typename Dst, typename Alpha, typename Xpr, typename Work>
113 static void addRightProduct(Dst& dst, const Alpha& alpha, const Xpr& X, const Factor& f, Work& work) {
114 const Index p = X.rows(), n1 = f.lhs().rows(), n2 = f.rhs().rows();
115 auto dstL = dst.reshaped(p * n2, n1);
116 LhsOps::addRightProduct(dstL, alpha, X.reshaped(p * n2, n1), f.lhs(), work);
117 for (Index j = 0; j < n1; ++j) {
118 auto dstj = dst.middleCols(j * n2, n2);
119 RhsOps::addRightProduct(dstj, alpha, X.middleCols(j * n2, n2), f.rhs(), work);
120 }
121 }
122 // det = prod_{i,j} (lambda_i + mu_j) would cost only the factor Schur forms,
123 // but every delta lambda_i recurs in n_R factors of the product, and a shift
124 // split across L and R makes |lambda_i + mu_j| << |lambda_i| + |mu_j|. Against
125 // a 256-bit reference that product lost 1 to 6 digits to the LU of the
126 // materialized sum.
127 static Scalar balancedDet(const Factor& f, Index& exponent) {
128 return kron_factor_ops<DenseMatrix>::balancedDet(denseFactor(f), exponent);
129 }
130};
131
132template <typename LhsMatrix, typename RhsMatrix>
133struct kron_factor_visitable<KroneckerSum<LhsMatrix, RhsMatrix>, kKronSumFactor> {
134 using Factor = KroneckerSum<LhsMatrix, RhsMatrix>;
135 using type = SparseMatrix<typename Factor::Scalar>;
136 static type get(const Factor& f) {
137 type S;
138 S = f;
139 return S;
140 }
141};
142
147template <typename LhsMatrix, typename RhsMatrix>
148struct kron_factor_spectrum<KroneckerSum<LhsMatrix, RhsMatrix>, kKronSumFactor>
149 : kron_factor_spectrum<KroneckerSum<LhsMatrix, RhsMatrix>, kKronDenseFactor> {
150 using Factor = KroneckerSum<LhsMatrix, RhsMatrix>;
151 using Eigenvectors = KroneckerOperator<typename kron_factor_spectrum<LhsMatrix>::Eigenvectors,
152 typename kron_factor_spectrum<RhsMatrix>::Eigenvectors>;
153
154 static typename Factor::ComplexVector eigenvalues(const Factor& f) { return f.eigenvalues(); }
155 static Eigenvectors eigenvectors(const Factor& f) { return f.eigenvectors(); }
156};
157
158template <typename LhsMatrix, typename RhsMatrix>
159class kron_factor_solver<KroneckerSum<LhsMatrix, RhsMatrix>, kKronSumFactor> {
160 public:
161 using Factor = KroneckerSum<LhsMatrix, RhsMatrix>;
162 using Scalar = typename Factor::Scalar;
163 using DenseMatrix = Matrix<Scalar, Dynamic, Dynamic, ColMajor>;
164
165 explicit kron_factor_solver(const Factor& f) : m_solver(f) {}
166 template <typename Xpr>
167 DenseMatrix solveLeft(const Xpr& M) const {
168 return m_solver.solve(M);
169 }
170 // X = M S^{-T} solves S X^T = M^T, written through a transposed view of X.
171 template <typename Xpr>
172 DenseMatrix solveTransposedRight(const Xpr& M) const {
173 DenseMatrix X(M.rows(), M.cols());
174 X.transpose() = m_solver.solve(M.transpose());
175 return X;
176 }
177
178 private:
179 BartelsStewart<Factor> m_solver;
180};
181
182} // namespace internal
183
232template <typename LhsMatrix, typename RhsMatrix>
233class KroneckerSum : public EigenBase<KroneckerSum<LhsMatrix, RhsMatrix>> {
234 public:
235 using Scalar = typename LhsMatrix::Scalar;
236 using RealScalar = typename NumTraits<Scalar>::Real;
237 using StorageIndex = int;
238
239 static_assert(std::is_same<Scalar, typename RhsMatrix::Scalar>::value,
240 "KroneckerSum requires both factors to have the same scalar type");
241 static_assert((internal::kron_factor_is_dense_matrix<LhsMatrix>::value ||
242 internal::kron_factor_kind<LhsMatrix>() != internal::kKronDenseFactor) &&
243 (internal::kron_factor_is_dense_matrix<RhsMatrix>::value ||
244 internal::kron_factor_kind<RhsMatrix>() != internal::kKronDenseFactor),
245 "KroneckerSum factors must be plain Matrix, DiagonalMatrix or SparseMatrix types, identity factors "
246 "(makeKroneckerSum stores an Identity() expression as one), KroneckerOperators or KroneckerSums "
247 "(owning their storage)");
248
249 private:
250 using LhsOps = internal::kron_factor_ops<LhsMatrix>;
251 using RhsOps = internal::kron_factor_ops<RhsMatrix>;
252 using LhsSpectrum = internal::kron_factor_spectrum<LhsMatrix>;
253 using RhsSpectrum = internal::kron_factor_spectrum<RhsMatrix>;
254
255 public:
256 using ComplexScalar = std::complex<RealScalar>;
258 using ComplexVector = Matrix<ComplexScalar, Dynamic, 1>;
259
260 static constexpr int RowsAtCompileTime =
261 internal::size_at_compile_time(LhsMatrix::RowsAtCompileTime, RhsMatrix::RowsAtCompileTime);
262 static constexpr int ColsAtCompileTime = RowsAtCompileTime;
263 static constexpr int MaxRowsAtCompileTime = RowsAtCompileTime;
264 static constexpr int MaxColsAtCompileTime = RowsAtCompileTime;
265 static constexpr int SizeAtCompileTime = internal::size_at_compile_time(RowsAtCompileTime, ColsAtCompileTime);
266 static constexpr int MaxSizeAtCompileTime = SizeAtCompileTime;
267 static constexpr bool IsRowMajor = false;
268 // Deliberately no IsVectorAtCompileTime, for the reason given in KroneckerOperator.
269
272 template <typename LhsDerived, typename RhsDerived>
273 KroneckerSum(const EigenBase<LhsDerived>& a, const EigenBase<RhsDerived>& b) : m_A(a.derived()), m_B(b.derived()) {
274 eigen_assert(m_A.size() > 0 && m_B.size() > 0 && "KroneckerSum factors must be non-empty");
275 eigen_assert(m_A.rows() == m_A.cols() && m_B.rows() == m_B.cols() && "KroneckerSum factors must be square");
276 LhsOps::prepare(m_A);
277 RhsOps::prepare(m_B);
278 }
279
280 EIGEN_DEVICE_FUNC Index rows() const { return m_A.rows() * m_B.rows(); }
281 EIGEN_DEVICE_FUNC Index cols() const { return rows(); }
282
284 const LhsMatrix& lhs() const { return m_A; }
286 const RhsMatrix& rhs() const { return m_B; }
287
291 Scalar coeff(Index row, Index col) const {
292 eigen_assert(row >= 0 && row < rows() && col >= 0 && col < cols());
293 Scalar value;
294 return internal::kron_factor_ops<KroneckerSum>::coeffIfStored(*this, row, col, value) ? value : Scalar(0);
295 }
296
299 return {LhsOps::transposed(m_A), RhsOps::transposed(m_B)};
300 }
301
303 KroneckerSum conjugate() const { return {LhsOps::conjugated(m_A), RhsOps::conjugated(m_B)}; }
304
307 return {LhsOps::adjointed(m_A), RhsOps::adjointed(m_B)};
308 }
309
313 template <typename Rhs>
315 const BartelsStewart<KroneckerSum> solver(*this);
316 return solver.solve(b);
317 }
318
333 ComplexVector eigenvalues() const {
334 const ComplexVector lambda = LhsSpectrum::eigenvalues(m_A), mu = RhsSpectrum::eigenvalues(m_B);
335 return (mu.replicate(fix<1>, lambda.size()) + lambda.transpose().replicate(mu.size(), fix<1>)).reshaped();
336 }
337
344 return {LhsSpectrum::eigenvectors(m_A), RhsSpectrum::eigenvectors(m_B)};
345 }
346
349 template <typename Rhs>
351 EIGEN_STATIC_ASSERT(ColsAtCompileTime == Dynamic || Rhs::RowsAtCompileTime == Dynamic ||
352 int(ColsAtCompileTime) == int(Rhs::RowsAtCompileTime),
353 INVALID_MATRIX_PRODUCT)
354 eigen_assert(x.rows() == cols() && "invalid product: dimensions do not match");
355 return Product<KroneckerSum, Rhs>(*this, x.derived());
356 }
357
362 template <typename Dest, typename Rhs, typename ProductScalar>
363 void addProduct(Dest& dst, const Rhs& rhs, const ProductScalar& alpha) const {
365 const Index n1 = m_A.rows(), n2 = m_B.rows(), r = rhs.cols();
366 eigen_assert(rhs.rows() == n1 * n2 && "invalid product: dimensions do not match");
367 typename internal::nested_eval<Rhs, 1>::type actualRhs(rhs);
368 ProductMatrix X, Y, work;
369 const Index chunk = internal::kron_rhs_chunk<ProductScalar>(2 * n1 * n2);
370 for (Index k0 = 0; k0 < r; k0 += chunk) {
371 const Index c = numext::mini(chunk, r - k0);
372 if (c == 1) {
373 const auto Xk = actualRhs.col(k0).reshaped(n2, n1);
374 auto Yk = dst.col(k0).reshaped(n2, n1);
375 RhsOps::addLeftProduct(Yk, alpha, m_B, Xk);
376 LhsOps::addRightProduct(Yk, alpha, Xk, m_A, work);
377 continue;
378 }
379 internal::kron_stack_columns(X, actualRhs.middleCols(k0, c), n2, n1);
380 Y.setZero(n2 * c, n1);
381 auto Yflat = Y.reshaped(n2, c * n1);
382 RhsOps::addLeftProduct(Yflat, ProductScalar(1), m_B, X.reshaped(n2, c * n1));
383 LhsOps::addRightProduct(Y, ProductScalar(1), X, m_A, work);
384 for (Index k = 0; k < c; ++k) dst.col(k0 + k).reshaped(n2, n1) += alpha * Y.middleRows(k * n2, n2);
385 }
386 }
387
391 template <typename Dest>
392 void evalTo(Dest& dst) const {
393 evalToImpl(dst, IsSparseDestination<Dest>());
394 }
395
397 template <typename Dest>
398 void addTo(Dest& dst) const {
399 addToImpl(dst, Scalar(1), IsSparseDestination<Dest>());
400 }
401
403 template <typename Dest>
404 void subTo(Dest& dst) const {
405 addToImpl(dst, Scalar(-1), IsSparseDestination<Dest>());
406 }
407
408 private:
409 template <typename Dest>
410 using IsSparseDestination = std::is_same<typename internal::traits<Dest>::StorageKind, Sparse>;
411
412 template <typename Dest>
413 void evalToImpl(Dest& dst, std::false_type) const {
414 dst.setZero();
415 addToImpl(dst, Scalar(1), std::false_type());
416 }
417
420 template <typename Dest>
421 void addToImpl(Dest& dst, const Scalar& s, std::false_type) const {
422 const Index n1 = m_A.rows(), n2 = m_B.rows();
423 LhsOps::forEachNonZero(m_A, [&dst, &s, n2](Index i, Index j, const Scalar& a) {
424 dst.block(i * n2, j * n2, n2, n2).diagonal().array() += s * a;
425 });
426 const auto& B = RhsOps::blockOperand(m_B);
427 for (Index k = 0; k < n1; ++k) dst.block(k * n2, k * n2, n2, n2) += s * B;
428 }
429
435 template <typename Dest>
436 void evalToImpl(Dest& S, std::true_type) const {
437 using IndexVector = Matrix<Index, Dynamic, 1>;
438 using LhsVisitable = internal::kron_factor_visitable<LhsMatrix>;
439 using RhsVisitable = internal::kron_factor_visitable<RhsMatrix>;
440 using VisitedLhsOps = internal::kron_factor_ops<typename LhsVisitable::type>;
441 using VisitedRhsOps = internal::kron_factor_ops<typename RhsVisitable::type>;
442 const Index n1 = m_A.rows(), n2 = m_B.rows();
443 const auto& A = LhsVisitable::get(m_A);
444 const auto& B = RhsVisitable::get(m_B);
445 const IndexVector nnzA = VisitedLhsOps::innerNonZeros(A, Dest::IsRowMajor);
446 const IndexVector nnzB = VisitedRhsOps::innerNonZeros(B, Dest::IsRowMajor);
447 Dest SA(rows(), cols()), SB(rows(), cols());
448 SA.reserve(IndexVector(nnzA.transpose().replicate(n2, fix<1>).reshaped()));
449 VisitedLhsOps::forEachNonZero(A, [&SA, n2](Index i, Index j, const Scalar& a) {
450 for (Index k = 0; k < n2; ++k) SA.insert(i * n2 + k, j * n2 + k) = a;
451 });
452 SB.reserve(IndexVector(nnzB.replicate(n1, fix<1>)));
453 for (Index k = 0; k < n1; ++k)
454 VisitedRhsOps::forEachNonZero(
455 B, [&SB, n2, k](Index i, Index j, const Scalar& b) { SB.insert(k * n2 + i, k * n2 + j) = b; });
456 S = SA + SB;
457 }
458
459 template <typename Dest>
460 void addToImpl(Dest& dst, const Scalar& s, std::true_type) const {
461 typename Dest::PlainObject sum;
462 evalTo(sum);
463 if (s == Scalar(1))
464 dst += sum;
465 else
466 dst -= sum;
467 }
468
469 LhsMatrix m_A;
470 RhsMatrix m_B;
471};
472
476template <typename LhsDerived, typename RhsDerived>
478 typename internal::kron_factor_storage<RhsDerived>::type>
480 return {a.derived(), b.derived()};
481}
482
486template <typename D1, typename D2, typename D3, typename... Rest>
487auto makeKroneckerSum(const EigenBase<D1>& a, const EigenBase<D2>& b, const EigenBase<D3>& c, const Rest&... rest) {
488 return makeKroneckerSum(a, makeKroneckerSum(b, c, rest...));
489}
490
553template <typename KroneckerSumType>
554class BartelsStewart : public SolverBase<BartelsStewart<KroneckerSumType>> {
555 public:
556 using Base = SolverBase<BartelsStewart>;
557 friend class SolverBase<BartelsStewart>;
558 EIGEN_GENERIC_PUBLIC_INTERFACE(BartelsStewart)
559 EIGEN_STATIC_ASSERT_NON_INTEGER(RealScalar)
560 using MatrixType = KroneckerSumType;
561 using ComplexScalar = std::complex<RealScalar>;
564 using RealVector = Matrix<RealScalar, Dynamic, 1>;
565
567 BartelsStewart() = default;
568
570 explicit BartelsStewart(const KroneckerSumType& op) { compute(op); }
571
574 BartelsStewart& compute(const KroneckerSumType& op) {
575 std::vector<DenseMatrix> leaves;
576 collectLeaves(op, leaves);
577 const std::size_t d = leaves.size();
578 m_sizes.resize(d);
579 m_inner.resize(d);
580 m_unitary.clear();
581 m_triangular.clear();
582 m_triangularAdjoint.clear();
583 m_basis.clear();
584 m_quasiTriangular.clear();
585 m_info = Success;
586 bool hermitian = true;
587 for (std::size_t k = 0; k < d; ++k) {
588 m_sizes[k] = leaves[k].rows();
589 if (!leaves[k].allFinite()) m_info = InvalidInput;
590 hermitian = hermitian && leaves[k] == leaves[k].adjoint();
591 }
592 if (hermitian)
593 m_path = Path::Diagonal;
594 else if (!NumTraits<Scalar>::IsComplex && d <= std::size_t(kMaxRealFactors))
595 m_path = Path::RealSchur;
596 else
597 m_path = Path::ComplexSchur;
598 m_size = 1;
599 for (std::size_t k = d; k-- > 0;) {
600 m_inner[k] = m_size;
601 m_size *= m_sizes[k];
602 }
603 m_isInitialized = true;
604 if (m_info != Success) return *this;
605 switch (m_path) {
606 case Path::Diagonal: {
607 m_spectrum = RealVector::Zero(1);
609 for (std::size_t k = 0; k < d; ++k) {
610 es.compute(leaves[k]);
611 if (es.info() != Success) m_info = NoConvergence;
612 m_basis.push_back(es.eigenvectors());
613 // Kronecker order: the new factor's index runs fastest.
614 const RealVector& lambda = es.eigenvalues();
615 const RealVector previous = m_spectrum;
616 m_spectrum =
617 (lambda.replicate(fix<1>, previous.size()) + previous.transpose().replicate(lambda.size(), fix<1>))
618 .reshaped();
619 }
620 break;
621 }
622 case Path::RealSchur:
623 computeRealSchur(leaves, internal::bool_constant<!NumTraits<Scalar>::IsComplex>());
624 break;
625 case Path::ComplexSchur: {
626 ComplexSchur<ComplexMatrix> schur;
627 for (std::size_t k = 0; k < d; ++k) {
628 schur.compute(leaves[k].template cast<ComplexScalar>());
629 if (schur.info() != Success) m_info = NoConvergence;
630 m_unitary.push_back(schur.matrixU());
631 m_triangular.push_back(schur.matrixT());
632 m_triangularAdjoint.emplace_back(m_triangular.back().adjoint());
633 }
634 break;
635 }
636 }
637 return *this;
638 }
639
640 Index rows() const noexcept { return m_size; }
641 Index cols() const noexcept { return m_size; }
642
646 eigen_assert(m_isInitialized && "BartelsStewart is not initialized.");
647 return m_info;
648 }
649
652 bool isHermitian() const {
653 eigen_assert(m_isInitialized && "BartelsStewart is not initialized.");
654 return m_path == Path::Diagonal;
655 }
656
657#ifdef EIGEN_PARSED_BY_DOXYGEN
661 template <typename Rhs>
663#endif
664
665#ifndef EIGEN_PARSED_BY_DOXYGEN
666 template <typename RhsType, typename DstType>
667 void _solve_impl(const RhsType& rhs, DstType& dst) const {
668 solveImpl(rhs, dst, /*adjoint=*/false);
669 }
670
671 // M^T = conj(M^H): M^{-T} b = conj(M^{-H} conj(b)).
672 template <bool Conjugate, typename RhsType, typename DstType>
673 void _solve_impl_transposed(const RhsType& rhs, DstType& dst) const {
674 constexpr bool ConjugateRhs = !Conjugate && NumTraits<Scalar>::IsComplex;
675 solveImpl(rhs.template conjugateIf<ConjugateRhs>(), dst, /*adjoint=*/true);
676 if (ConjugateRhs) dst = dst.conjugate();
677 }
678#endif
679
680 private:
683 template <typename RhsType, typename DstType>
684 void solveImpl(const RhsType& rhs, DstType& dst, bool adjoint) const {
685 if (m_info != Success) {
686 // No usable decompositions; the nested solves never see info().
687 dst.setConstant(Scalar(NumTraits<RealScalar>::quiet_NaN()));
688 return;
689 }
690 switch (m_path) {
691 case Path::Diagonal: {
692 DenseMatrix W = rhs;
693 applyBasis(W, m_basis, /*adjoint=*/true);
694 W.array().colwise() /= m_spectrum.array();
695 applyBasis(W, m_basis, /*adjoint=*/false);
696 dst = W;
697 break;
698 }
699 case Path::RealSchur: {
700 DenseMatrix W = rhs;
701 solveRealSchur(W, adjoint, internal::bool_constant<!NumTraits<Scalar>::IsComplex>());
702 dst = W;
703 break;
704 }
705 case Path::ComplexSchur: {
706 ComplexMatrix W = rhs.template cast<ComplexScalar>();
707 applyBasis(W, m_unitary, /*adjoint=*/true);
708 for (Index j = 0; j < W.cols(); ++j) {
709 if (adjoint)
710 adjointTriangularSolve(0, ComplexScalar(0), W.col(j).data());
711 else
712 triangularSolve(0, ComplexScalar(0), W.col(j).data());
713 }
714 applyBasis(W, m_unitary, /*adjoint=*/false);
715 dst = internal::structured_scalar_part_impl<Scalar>::run(W);
716 break;
717 }
718 }
719 }
720
721 template <typename Factor>
722 static void collectLeaves(const Factor& f, std::vector<DenseMatrix>& leaves) {
723 collectLeaves(f, leaves, internal::kron_factor_is_kronecker_sum<Factor>());
724 }
725 template <typename Factor>
726 static void collectLeaves(const Factor& f, std::vector<DenseMatrix>& leaves, std::true_type) {
727 collectLeaves(f.lhs(), leaves);
728 collectLeaves(f.rhs(), leaves);
729 }
730 template <typename Factor>
731 static void collectLeaves(const Factor& f, std::vector<DenseMatrix>& leaves, std::false_type) {
732 leaves.push_back(DenseMatrix(internal::kron_factor_ops<Factor>::denseFactor(f)));
733 }
734
740 template <typename Work, typename Basis>
741 void applyBasis(Work& W, const std::vector<Basis>& Q, bool adjoint) const {
742 using WorkMatrix = Matrix<typename Work::Scalar, Dynamic, Dynamic, ColMajor>;
743 WorkMatrix T;
744 for (std::size_t k = 0; k < Q.size(); ++k) {
745 const Index n = m_sizes[k], s = m_inner[k], blocks = W.size() / (n * s);
746 if (s == 1) {
747 Map<WorkMatrix> U(W.data(), n, blocks);
748 if (adjoint)
749 T.noalias() = Q[k].adjoint() * U;
750 else
751 T.noalias() = Q[k] * U;
752 U = T;
753 continue;
754 }
755 for (Index b = 0; b < blocks; ++b) {
756 Map<WorkMatrix> V(W.data() + b * n * s, s, n);
757 if (adjoint)
758 T.noalias() = V * Q[k].conjugate();
759 else
760 T.noalias() = V * Q[k].transpose();
761 V = T;
762 }
763 }
764 }
765
770 void triangularSolve(std::size_t k, const ComplexScalar& sigma, ComplexScalar* y) const {
771 const Index n = m_sizes[k], s = m_inner[k];
772 const TriangularMatrix& T = m_triangular[k];
773 if (k + 1 == m_sizes.size()) {
774 // Last factor, s = 1: back substitution on the shifted T_d, by rows.
775 Map<Matrix<ComplexScalar, 1, Dynamic>> Y(y, n);
776 for (Index i = n - 1; i >= 0; --i) {
777 const Index tail = n - 1 - i;
778 Y(i) = (Y(i) - Y.tail(tail).cwiseProduct(T.row(i).tail(tail)).sum()) / (sigma + T(i, i));
779 }
780 return;
781 }
782 Map<ComplexMatrix> Y(y, s, n);
783 for (Index i = n - 1; i >= 0; --i) {
784 const Index tail = n - 1 - i;
785 if (tail > 0) Y.col(i).noalias() -= Y.rightCols(tail) * T.row(i).tail(tail).transpose();
786 triangularSolve(k + 1, sigma + T(i, i), Y.col(i).data());
787 }
788 }
789
793 void adjointTriangularSolve(std::size_t k, const ComplexScalar& sigma, ComplexScalar* y) const {
794 const Index n = m_sizes[k], s = m_inner[k];
795 const TriangularMatrix& L = m_triangularAdjoint[k];
796 if (k + 1 == m_sizes.size()) {
797 // Last factor, s = 1: forward substitution on the shifted L_d, by rows.
798 Map<Matrix<ComplexScalar, 1, Dynamic>> Y(y, n);
799 for (Index i = 0; i < n; ++i) Y(i) = (Y(i) - Y.head(i).cwiseProduct(L.row(i).head(i)).sum()) / (sigma + L(i, i));
800 return;
801 }
802 Map<ComplexMatrix> Y(y, s, n);
803 for (Index i = 0; i < n; ++i) {
804 if (i > 0) Y.col(i).noalias() -= Y.leftCols(i) * L.row(i).head(i).transpose();
805 adjointTriangularSolve(k + 1, sigma + L(i, i), Y.col(i).data());
806 }
807 }
808
809 // The real path's coupled systems are at most 2^d wide, see the class documentation.
810 static constexpr int kMaxRealFactors = 4;
811 static constexpr int kMaxCoupling = 1 << kMaxRealFactors;
812 using CouplingMatrix = Matrix<RealScalar, Dynamic, Dynamic, ColMajor, kMaxCoupling, kMaxCoupling>;
813 using CouplingVector = Matrix<RealScalar, Dynamic, 1, ColMajor, kMaxCoupling, 1>;
814 using RealMatrix = Matrix<RealScalar, Dynamic, Dynamic, ColMajor>;
815
816 void computeRealSchur(const std::vector<DenseMatrix>& leaves, std::true_type) {
817 RealSchur<DenseMatrix> schur;
818 for (const DenseMatrix& leaf : leaves) {
819 schur.compute(leaf);
820 if (schur.info() != Success) m_info = NoConvergence;
821 m_basis.push_back(schur.matrixU());
822 m_quasiTriangular.push_back(schur.matrixT());
823 }
824 }
825 void computeRealSchur(const std::vector<DenseMatrix>&, std::false_type) {}
826
827 void solveRealSchur(DenseMatrix& W, bool adjoint, std::true_type) const {
828 applyBasis(W, m_basis, /*adjoint=*/true);
829 std::vector<RealMatrix> work(m_sizes.size());
830 const CouplingMatrix outer = CouplingMatrix::Zero(1, 1);
831 for (Index j = 0; j < W.cols(); ++j)
832 quasiTriangularSolve(0, outer, Map<RealMatrix>(W.col(j).data(), W.rows(), 1), work, adjoint);
833 applyBasis(W, m_basis, /*adjoint=*/false);
834 }
835 void solveRealSchur(DenseMatrix&, bool, std::false_type) const {}
836
843 void quasiTriangularSolve(std::size_t k, const CouplingMatrix& Sigma, Map<RealMatrix> X,
844 std::vector<RealMatrix>& work, bool adjoint) const {
845 const Index n = m_sizes[k], s = m_inner[k], m = Sigma.rows();
846 const QuasiTriangularMatrix& T = m_quasiTriangular[k];
847 const bool last = k + 1 == m_sizes.size();
848 for (Index done = 0; done < n;) {
849 // Block I = [i, i + p): the next one up from the bottom, or down from the top when adjoint.
850 Index i, p;
851 if (adjoint) {
852 i = done;
853 p = i + 1 < n && T(i + 1, i) != RealScalar(0) ? 2 : 1;
854 } else {
855 const Index end = n - done;
856 p = end > 1 && T(end - 1, end - 2) != RealScalar(0) ? 2 : 1;
857 i = end - p;
858 }
859 const Index tail = n - i - p;
860 // T(I,I), T(J,I) for J < I, and T(I,J) for J > I.
861 const auto TII = T.block(i, i, p, p);
862 const auto colHead = T.block(0, i, i, p);
863 const auto rowTail = T.block(i, i + p, p, tail);
864 const CouplingMatrix coupled = adjoint ? kroneckerSum(Sigma, TII.transpose()) : kroneckerSum(Sigma, TII);
865 if (last) {
866 // s = 1: row i of X belongs to index i of the last factor.
867 auto XI = X.middleRows(i, p);
868 if (adjoint) {
869 if (i > 0) XI.noalias() -= colHead.transpose().lazyProduct(X.topRows(i));
870 } else if (tail > 0) {
871 XI.noalias() -= rowTail.lazyProduct(X.bottomRows(tail));
872 }
873 solveCoupled(coupled, XI);
874 } else {
875 for (Index t = 0; t < m; ++t) {
876 auto Xt = X.col(t).reshaped(s, n);
877 auto XtI = Xt.middleCols(i, p);
878 if (adjoint) {
879 if (i > 0) XtI.noalias() -= Xt.leftCols(i) * colHead;
880 } else if (tail > 0) {
881 XtI.noalias() -= Xt.rightCols(tail) * rowTail.transpose();
882 }
883 }
884 if (m == 1) {
885 quasiTriangularSolve(k + 1, coupled, Map<RealMatrix>(X.col(0).data() + i * s, s, p), work, adjoint);
886 } else {
887 // Columns (t, a) of Z, a fastest: column a of block I for coupling index t.
888 RealMatrix& Z = work[k];
889 Z.resize(s, m * p);
890 Z.reshaped(p * s, m) = X.middleRows(i * s, p * s);
891 quasiTriangularSolve(k + 1, coupled, Map<RealMatrix>(Z.data(), s, m * p), work, adjoint);
892 X.middleRows(i * s, p * s) = Z.reshaped(p * s, m);
893 }
894 }
895 done += p;
896 }
897 }
898
900 template <typename Block>
901 static CouplingMatrix kroneckerSum(const CouplingMatrix& Sigma, const Block& S) {
902 const Index m = Sigma.rows(), p = S.rows();
903 eigen_internal_assert(m * p <= kMaxCoupling);
904 CouplingMatrix R = CouplingMatrix::Zero(m * p, m * p);
905 for (Index t = 0; t < m; ++t) {
906 for (Index u = 0; u < m; ++u) R.block(t * p, u * p, p, p).diagonal().setConstant(Sigma(t, u));
907 R.block(t * p, t * p, p, p) += S;
908 }
909 return R;
910 }
911
913 template <typename Block>
914 static void solveCoupled(const CouplingMatrix& coupled, Block& XI) {
915 if (coupled.rows() == 1) {
916 XI(0, 0) /= coupled(0, 0);
917 return;
918 }
919 CouplingVector z = XI.reshaped();
920 z = coupled.partialPivLu().solve(z);
921 XI.reshaped() = z;
922 }
923
924 std::vector<Index> m_sizes, m_inner;
925 Index m_size = 0;
926 std::vector<DenseMatrix> m_basis;
927 RealVector m_spectrum;
928 // Row-major: the back substitution reads T_k by rows, and the forward
929 // substitution of the adjoint solve reads T_k^H, kept as its own copy, the
930 // same way.
931 using TriangularMatrix = Matrix<ComplexScalar, Dynamic, Dynamic, RowMajor>;
932 using QuasiTriangularMatrix = Matrix<RealScalar, Dynamic, Dynamic, RowMajor>;
933 std::vector<ComplexMatrix> m_unitary;
934 std::vector<TriangularMatrix> m_triangular, m_triangularAdjoint;
935 std::vector<QuasiTriangularMatrix> m_quasiTriangular;
936 enum class Path { Diagonal, RealSchur, ComplexSchur };
937 Path m_path = Path::Diagonal;
938 bool m_isInitialized = false;
940};
941
942namespace internal {
943
944template <typename LhsMatrix, typename RhsMatrix, typename Rhs, int ProductTag>
945struct generic_product_impl<KroneckerSum<LhsMatrix, RhsMatrix>, Rhs, StructuredShape, DenseShape, ProductTag>
946 : structured_product_impl<KroneckerSum<LhsMatrix, RhsMatrix>, Rhs> {};
947
948} // namespace internal
949
950} // namespace Eigen
951
952#endif // EIGEN_STRUCTURED_KRONECKER_SUM_H
Direct solver for Kronecker-sum systems .
Definition KroneckerSum.h:554
bool isHermitian() const
Definition KroneckerSum.h:652
ComputationInfo info() const
Definition KroneckerSum.h:645
BartelsStewart(const KroneckerSumType &op)
Definition KroneckerSum.h:570
BartelsStewart & compute(const KroneckerSumType &op)
Definition KroneckerSum.h:574
const Solve< BartelsStewart, Rhs > solve(const MatrixBase< Rhs > &b) const
ComputationInfo info() const
const MatrixTType & matrixT() const
ComplexSchur & compute(const EigenBase< InputType > &matrix, bool computeU=true)
const ComplexMatrixType & matrixU() const
The Kronecker product as an implicit operator that is never materialized.
Definition KroneckerOperator.h:908
The Kronecker sum of two square matrices as an implicit operator that is never materialized.
Definition KroneckerSum.h:233
KroneckerSum conjugate() const
Definition KroneckerSum.h:303
KroneckerSum< typename LhsOps::TransposedFactor, typename RhsOps::TransposedFactor > adjoint() const
Definition KroneckerSum.h:306
KroneckerSum(const EigenBase< LhsDerived > &a, const EigenBase< RhsDerived > &b)
Definition KroneckerSum.h:273
const RhsMatrix & rhs() const
Definition KroneckerSum.h:286
Matrix< Scalar, ColsAtCompileTime, Rhs::ColsAtCompileTime > solve(const MatrixBase< Rhs > &b) const
Definition KroneckerSum.h:314
KroneckerSum< typename LhsOps::TransposedFactor, typename RhsOps::TransposedFactor > transpose() const
Definition KroneckerSum.h:298
Scalar coeff(Index row, Index col) const
Definition KroneckerSum.h:291
ComplexVector eigenvalues() const
Definition KroneckerSum.h:333
const LhsMatrix & lhs() const
Definition KroneckerSum.h:284
KroneckerOperator< typename LhsSpectrum::Eigenvectors, typename RhsSpectrum::Eigenvectors > eigenvectors() const
Definition KroneckerSum.h:343
Product< KroneckerSum, Rhs > operator*(const MatrixBase< Rhs > &x) const
Definition KroneckerSum.h:350
SelfAdjointEigenSolver & compute(const EigenBase< InputType > &matrix, int options=ComputeEigenvectors)
ComputationInfo info() const
const RealVectorType & eigenvalues() const
const EigenvectorsType & eigenvectors() const
static const auto fix()
KroneckerSum< typename internal::kron_factor_storage< LhsDerived >::type, typename internal::kron_factor_storage< RhsDerived >::type > makeKroneckerSum(const EigenBase< LhsDerived > &a, const EigenBase< RhsDerived > &b)
Definition KroneckerSum.h:479
ComputationInfo
constexpr unsigned int RowMajorBit
Namespace containing all symbols from the Eigen library.
constexpr Derived & derived()
Eigen::Index Index