Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
KroneckerOperator.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] C. F. Van Loan, "The ubiquitous Kronecker product", Journal of
12// Computational and Applied Mathematics 123 (2000), 85-100. The vec
13// identity, with vec stacking columns, driving every product and solve
14// below -- stated there as
15// Y = C X B^T <=> vec(Y) = (B (x) C) vec(X), i.e. with this file's
16// operand naming (A (x) B) vec(X) = vec(B X A^T) -- together with the
17// factor-wise identities for the inverse, pseudo-inverse,
18// eigendecomposition, SVD and determinant of a Kronecker product.
19// [2] N. J. Higham, "Accuracy and Stability of Numerical Algorithms", 2nd ed.,
20// SIAM, 2002, chapter 27. Avoiding spurious overflow by rescaling with
21// powers of two, the technique behind the exponent-balanced determinant.
22// [3] P. H. Sterbenz, "Floating-Point Computation", Prentice-Hall, 1974.
23// Scaling by a power of two is exact, the property the determinant's
24// rescaling relies on.
25
26#ifndef EIGEN_STRUCTURED_KRONECKER_OPERATOR_H
27#define EIGEN_STRUCTURED_KRONECKER_OPERATOR_H
28
29// IWYU pragma: private
30#include "./InternalHeaderCheck.h"
31
32namespace Eigen {
33
34template <typename LhsMatrix, typename RhsMatrix>
36template <typename LhsMatrix, typename RhsMatrix>
37class KroneckerSum;
38
39namespace internal {
40
41template <typename LhsMatrix, typename RhsMatrix>
42struct traits<KroneckerOperator<LhsMatrix, RhsMatrix>> {
43 using Scalar = typename LhsMatrix::Scalar;
44 using StorageKind = Dense;
45 using XprKind = MatrixXpr;
46 using StorageIndex = int;
47 // size_at_compile_time is the compile-time dimension product, Dynamic when a
48 // factor is unknown or the product would overflow int.
49 static constexpr int RowsAtCompileTime =
50 size_at_compile_time(traits<LhsMatrix>::RowsAtCompileTime, traits<RhsMatrix>::RowsAtCompileTime);
51 static constexpr int ColsAtCompileTime =
52 size_at_compile_time(traits<LhsMatrix>::ColsAtCompileTime, traits<RhsMatrix>::ColsAtCompileTime);
53 static constexpr int MaxRowsAtCompileTime = RowsAtCompileTime;
54 static constexpr int MaxColsAtCompileTime = ColsAtCompileTime;
55 // Deliberately no NestByRefBit: transpose(), conjugate(), adjoint(), inverse(),
56 // eigenvectors(), matrixU() and matrixV() return owning temporaries (the
57 // operator stores its factors by value), so Product must nest the operator by
58 // value for a delayed-evaluated product expression to keep its left factor
59 // alive. The copy is O(m1 n1 + m2 n2), negligible against the product
60 // evaluation.
61 static constexpr int Flags = 0;
62};
63
64template <typename LhsMatrix, typename RhsMatrix>
65struct evaluator_traits<KroneckerOperator<LhsMatrix, RhsMatrix>> {
66 using Kind = IndexBased;
67 using Shape = StructuredShape;
68};
69
70// Factor kinds, each with its operations in the kron_factor_* helpers below:
71// dense; diagonal (stored as its diagonal); sparse (compressed, solved by
72// SparseLU); identity (a kron_identity_factor, dimensions only, skipped in
73// products); Kronecker (a nested KroneckerOperator, for three or more factors);
74// Kronecker sum (a KroneckerSum, see KroneckerSum.h).
75
76template <typename Scalar_, int Rows_, int Cols_>
77class kron_identity_factor;
78
79template <typename Scalar_, int Rows_, int Cols_>
80struct traits<kron_identity_factor<Scalar_, Rows_, Cols_>> : traits<Matrix<Scalar_, Rows_, Cols_>> {};
81
86template <typename Scalar_, int Rows_, int Cols_>
87class kron_identity_factor : public EigenBase<kron_identity_factor<Scalar_, Rows_, Cols_>> {
88 public:
89 using Scalar = Scalar_;
90 using PlainObject = Matrix<Scalar, Rows_, Cols_>;
91 static constexpr int RowsAtCompileTime = Rows_;
92 static constexpr int ColsAtCompileTime = Cols_;
93
94 kron_identity_factor(Index rows, Index cols) : m_rows(rows), m_cols(cols) {}
95 template <typename PlainObjectType>
96 explicit kron_identity_factor(const CwiseNullaryOp<scalar_identity_op<Scalar>, PlainObjectType>& identity)
97 : m_rows(identity.rows()), m_cols(identity.cols()) {}
98
99 EIGEN_DEVICE_FUNC constexpr Index rows() const { return m_rows.value(); }
100 EIGEN_DEVICE_FUNC constexpr Index cols() const { return m_cols.value(); }
101
102 private:
103 variable_if_dynamic<Index, Rows_> m_rows;
104 variable_if_dynamic<Index, Cols_> m_cols;
105};
106
107template <typename Factor>
108struct kron_factor_is_diagonal : std::false_type {};
109template <typename Scalar, int Size, int MaxSize>
110struct kron_factor_is_diagonal<DiagonalMatrix<Scalar, Size, MaxSize>> : std::true_type {};
111
112template <typename Factor>
113struct kron_factor_is_dense_matrix : std::false_type {};
114template <typename Scalar, int Rows, int Cols, int Options, int MaxRows, int MaxCols>
115struct kron_factor_is_dense_matrix<Matrix<Scalar, Rows, Cols, Options, MaxRows, MaxCols>> : std::true_type {};
116
117template <typename Factor>
118struct kron_factor_is_sparse_matrix : std::false_type {};
119template <typename Scalar, int Options, typename StorageIndex>
120struct kron_factor_is_sparse_matrix<SparseMatrix<Scalar, Options, StorageIndex>> : std::true_type {};
121
122template <typename Factor>
123struct kron_factor_is_identity : std::false_type {};
124template <typename Scalar, int Rows, int Cols>
125struct kron_factor_is_identity<kron_identity_factor<Scalar, Rows, Cols>> : std::true_type {};
126
127template <typename Factor>
128struct kron_factor_is_kronecker : std::false_type {};
129template <typename LhsMatrix, typename RhsMatrix>
130struct kron_factor_is_kronecker<KroneckerOperator<LhsMatrix, RhsMatrix>> : std::true_type {};
131
132template <typename Factor>
133struct kron_factor_is_kronecker_sum : std::false_type {};
134template <typename LhsMatrix, typename RhsMatrix>
135struct kron_factor_is_kronecker_sum<KroneckerSum<LhsMatrix, RhsMatrix>> : std::true_type {};
136
137// The factor kind, the dispatch key of kron_factor_ops and kron_factor_solver.
138constexpr int kKronDenseFactor = 0;
139constexpr int kKronDiagonalFactor = 1;
140constexpr int kKronSparseFactor = 2;
141constexpr int kKronIdentityFactor = 3;
142constexpr int kKronKroneckerFactor = 4;
143constexpr int kKronSumFactor = 5;
144
145template <typename Factor>
146constexpr int kron_factor_kind() {
147 return kron_factor_is_diagonal<Factor>::value ? kKronDiagonalFactor
148 : kron_factor_is_sparse_matrix<Factor>::value ? kKronSparseFactor
149 : kron_factor_is_identity<Factor>::value ? kKronIdentityFactor
150 : kron_factor_is_kronecker<Factor>::value ? kKronKroneckerFactor
151 : kron_factor_is_kronecker_sum<Factor>::value ? kKronSumFactor
152 : kKronDenseFactor;
153}
154
159template <typename Derived, bool StoredAsIs = kron_factor_is_identity<Derived>::value ||
160 kron_factor_is_kronecker<Derived>::value ||
161 kron_factor_is_kronecker_sum<Derived>::value>
162struct kron_factor_storage {
163 using type = typename Derived::PlainObject;
164};
165template <typename Derived>
166struct kron_factor_storage<Derived, true> {
167 using type = Derived;
168};
169template <typename Scalar, typename PlainObjectType>
170struct kron_factor_storage<CwiseNullaryOp<scalar_identity_op<Scalar>, PlainObjectType>, false> {
171 using type = kron_identity_factor<Scalar, PlainObjectType::RowsAtCompileTime, PlainObjectType::ColsAtCompileTime>;
172};
173
178template <typename Factor, int Kind = kron_factor_kind<Factor>()>
179struct kron_factor_visitable {
180 using type = Factor;
181 static const Factor& get(const Factor& f) { return f; }
182};
183
187template <typename LhsMatrix, typename RhsMatrix,
188 bool Unchanged = std::is_same<typename kron_factor_visitable<LhsMatrix>::type, LhsMatrix>::value &&
189 std::is_same<typename kron_factor_visitable<RhsMatrix>::type, RhsMatrix>::value>
190struct kron_operator_visitable {
191 using Factor = KroneckerOperator<LhsMatrix, RhsMatrix>;
192 using type = Factor;
193 static const Factor& get(const Factor& f) { return f; }
194};
195template <typename LhsMatrix, typename RhsMatrix>
196struct kron_operator_visitable<LhsMatrix, RhsMatrix, false> {
197 using LhsVisitable = kron_factor_visitable<LhsMatrix>;
198 using RhsVisitable = kron_factor_visitable<RhsMatrix>;
199 using type = KroneckerOperator<typename LhsVisitable::type, typename RhsVisitable::type>;
200 static type get(const KroneckerOperator<LhsMatrix, RhsMatrix>& f) {
201 return type(LhsVisitable::get(f.lhs()), RhsVisitable::get(f.rhs()));
202 }
203};
204template <typename LhsMatrix, typename RhsMatrix>
205struct kron_factor_visitable<KroneckerOperator<LhsMatrix, RhsMatrix>, kKronKroneckerFactor>
206 : kron_operator_visitable<LhsMatrix, RhsMatrix> {};
207
214template <typename Stacked, typename Xpr>
215void kron_stack_columns(Stacked& stacked, const Xpr& x, Index rows, Index cols) {
216 stacked.resize(rows * x.cols(), cols);
217 for (Index k = 0; k < x.cols(); ++k)
218 stacked.middleRows(k * rows, rows) = x.col(k).reshaped(rows, cols).template cast<typename Stacked::Scalar>();
219}
220
226template <typename WorkScalar>
227Index kron_rhs_chunk(Index perColumn) {
228 const Index budget = Index(l2CacheSize() / 8) / Index(sizeof(WorkScalar));
229 return numext::maxi(Index(1), budget / numext::maxi(perColumn, Index(1)));
230}
231
232template <typename Factor, int Kind = kron_factor_kind<Factor>()>
233struct kron_factor_ops {
234 // Dense factor.
235 using Scalar = typename Factor::Scalar;
236 using RealScalar = typename NumTraits<Scalar>::Real;
237 using TransposedFactor = Matrix<Scalar, Factor::ColsAtCompileTime, Factor::RowsAtCompileTime>;
238 using InverseFactor = Matrix<Scalar, Dynamic, Dynamic, ColMajor>;
239 // Every entry is stored: a materialization has no structurally zero blocks
240 // to clear, and forEachNonZero visits all m*n entries in column-major order.
241 static constexpr bool StoresAllEntries = true;
242
243 static void prepare(Factor&) {}
244 static bool coeffIfStored(const Factor& f, Index row, Index col, Scalar& value) {
245 value = f.coeff(row, col);
246 return true;
247 }
248 template <typename Visitor>
249 static void forEachNonZero(const Factor& f, Visitor&& visit) {
250 for (Index j = 0; j < f.cols(); ++j)
251 for (Index i = 0; i < f.rows(); ++i) visit(i, j, f.coeff(i, j));
252 }
255 static Matrix<Index, Dynamic, 1> innerNonZeros(const Factor& f, bool rowMajor) {
256 return Matrix<Index, Dynamic, 1>::Constant(rowMajor ? f.rows() : f.cols(), rowMajor ? f.cols() : f.rows());
257 }
258 static auto transposed(const Factor& f) { return f.transpose(); }
259 static auto conjugated(const Factor& f) { return f.conjugate(); }
260 static auto adjointed(const Factor& f) { return f.adjoint(); }
261 static auto inversed(const Factor& f) { return f.inverse(); }
262 // The factor as a dense expression, for the decomposition family.
263 static const Factor& denseFactor(const Factor& f) { return f; }
264 // The factor as the operand of a dense block assignment dst = a * f.
265 static const Factor& blockOperand(const Factor& f) { return f; }
266 static bool isSquareIdentity(const Factor&) { return false; }
267 // dst += alpha F X and dst += alpha X F^T: the two sides of the vec-trick product B X A^T.
268 // work: scratch for a nested factor's right product, owned by the caller's loop.
269 template <typename Dst, typename Alpha, typename Xpr>
270 static void addLeftProduct(Dst& dst, const Alpha& alpha, const Factor& f, const Xpr& X) {
271 dst.noalias() += alpha * (f * X);
272 }
273 template <typename Dst, typename Alpha, typename Xpr, typename Work>
274 static void addRightProduct(Dst& dst, const Alpha& alpha, const Xpr& X, const Factor& f, Work&) {
275 dst.noalias() += alpha * (X * f.transpose());
276 }
277 static int exponentBound(const Factor& f) { return structured_exponent_bound(f); }
282 static Scalar balancedDet(const Factor& M, Index& exponent) {
283 // Scale small factors up, exactly, so that the elimination runs on normal
284 // numbers: a subnormal near 2^e carries only e - min_exponent + digits bits.
285 // Scaling down could erase small pivots in factors with a wide exponent range,
286 // and every entry is scaled exactly because a small entry can be a pivot.
287 const int bound = exponentBound(M);
288 const int scaleExponent = numext::mini(bound, 0);
289 Factor normalized = M;
290 structured_ldexp_entries_exact(normalized, -scaleExponent);
291 exponent += M.rows() * scaleExponent;
292 PartialPivLU<Factor> lu(normalized);
293 Scalar m = Scalar(RealScalar(lu.permutationP().determinant())); // +-1
294 for (Index i = 0; i < M.rows(); ++i)
295 m = structured_balance(m * structured_balance(lu.matrixLU().coeff(i, i), exponent), exponent);
296 return m;
297 }
298};
299
300template <typename Factor>
301struct kron_factor_ops<Factor, kKronDiagonalFactor> {
302 // Diagonal factor: everything runs on the stored diagonal vector.
303 using Scalar = typename Factor::Scalar;
304 using TransposedFactor = typename Factor::PlainObject; // a diagonal matrix is its own transpose
305 using InverseFactor = typename Factor::PlainObject; // entrywise reciprocals stay diagonal
306 static constexpr bool StoresAllEntries = false;
307
308 static void prepare(Factor&) {}
309 static bool coeffIfStored(const Factor& f, Index row, Index col, Scalar& value) {
310 if (row != col) return false;
311 value = f.diagonal().coeff(row);
312 return true;
313 }
314 template <typename Visitor>
315 static void forEachNonZero(const Factor& f, Visitor&& visit) {
316 for (Index k = 0; k < f.rows(); ++k) visit(k, k, f.diagonal().coeff(k));
317 }
318 static Matrix<Index, Dynamic, 1> innerNonZeros(const Factor& f, bool) {
319 return Matrix<Index, Dynamic, 1>::Ones(f.rows());
320 }
321 static const Factor& transposed(const Factor& f) { return f; }
322 static auto conjugated(const Factor& f) { return f.diagonal().conjugate().asDiagonal(); }
323 static auto adjointed(const Factor& f) { return f.diagonal().conjugate().asDiagonal(); }
324 static auto inversed(const Factor& f) { return f.inverse(); }
325 static typename Factor::DenseMatrixType denseFactor(const Factor& f) { return f.toDenseMatrix(); }
326 static const Factor& blockOperand(const Factor& f) { return f; }
327 static bool isSquareIdentity(const Factor&) { return false; }
328 template <typename Dst, typename Alpha, typename Xpr>
329 static void addLeftProduct(Dst& dst, const Alpha& alpha, const Factor& f, const Xpr& X) {
330 dst.noalias() += alpha * (f * X);
331 }
332 template <typename Dst, typename Alpha, typename Xpr, typename Work>
333 static void addRightProduct(Dst& dst, const Alpha& alpha, const Xpr& X, const Factor& f, Work&) {
334 dst.noalias() += alpha * (X * f);
335 }
336 static int exponentBound(const Factor& f) { return structured_exponent_bound(f.diagonal()); }
337 static Scalar balancedDet(const Factor& D, Index& exponent) {
338 // Scale a small diagonal up first: the balancing below compares and scales
339 // in floating point, which reads a subnormal entry as zero under flush-to-zero.
340 const int bound = exponentBound(D);
341 const int scaleExponent = numext::mini(bound, 0);
342 typename Factor::DiagonalVectorType d = D.diagonal();
343 structured_ldexp_entries_exact(d, -scaleExponent);
344 exponent += D.rows() * scaleExponent;
345 Scalar m(1);
346 for (Index i = 0; i < D.rows(); ++i) m = structured_balance(m * structured_balance(d.coeff(i), exponent), exponent);
347 return m;
348 }
349};
350
355template <typename SparseType>
356class kron_sparse_lu : public SparseLU<SparseType> {
357 public:
358 using Base = SparseLU<SparseType>;
359 using Scalar = typename SparseType::Scalar;
360 using RealScalar = typename NumTraits<Scalar>::Real;
361
362 Scalar balancedDet(Index& exponent) const {
363 eigen_assert(Base::info() == Success);
364 Scalar m = Scalar(RealScalar(Base::m_detPermR * Base::m_detPermC));
365 for (Index j = 0; j < Base::cols(); ++j) {
366 typename Base::SCMatrix::InnerIterator it(Base::m_Lstore, j);
367 while (it && it.index() != j) ++it;
368 if (it) m = structured_balance(m * structured_balance(it.value(), exponent), exponent);
369 }
370 return m;
371 }
372};
373
374template <typename Factor>
375struct kron_factor_ops<Factor, kKronSparseFactor> {
376 // Sparse factor: the products run through the sparse-dense kernels, the
377 // solve and the determinant through a SparseLU factorization.
378 using Scalar = typename Factor::Scalar;
379 using RealScalar = typename NumTraits<Scalar>::Real;
380 using ColMajorFactor = SparseMatrix<Scalar, ColMajor, typename Factor::StorageIndex>;
381 using TransposedFactor = Factor; // the transpose is stored sparse, in the same storage order
382 using InverseFactor = Matrix<Scalar, Dynamic, Dynamic, ColMajor>; // dense in general
383 static constexpr bool StoresAllEntries = false;
384
385 // coeffs() below reads the compressed value array.
386 static void prepare(Factor& f) { f.makeCompressed(); }
387 static bool coeffIfStored(const Factor& f, Index row, Index col, Scalar& value) {
388 const Index outer = Factor::IsRowMajor ? row : col;
389 const Index inner = Factor::IsRowMajor ? col : row;
390 const Index start = f.outerIndexPtr()[outer], end = f.outerIndexPtr()[outer + 1];
391 if (start == end) return false;
392 const Index position = f.data().searchLowerIndex(start, end, inner);
393 if (position == end || f.innerIndexPtr()[position] != inner) return false;
394 value = f.valuePtr()[position];
395 return true;
396 }
397 template <typename Visitor>
398 static void forEachNonZero(const Factor& f, Visitor&& visit) {
399 for (Index k = 0; k < f.outerSize(); ++k)
400 for (typename Factor::InnerIterator it(f, k); it; ++it) visit(it.row(), it.col(), it.value());
401 }
402 static Matrix<Index, Dynamic, 1> innerNonZeros(const Factor& f, bool rowMajor) {
403 Matrix<Index, Dynamic, 1> counts = Matrix<Index, Dynamic, 1>::Zero(rowMajor ? f.rows() : f.cols());
404 for (Index k = 0; k < f.outerSize(); ++k)
405 for (typename Factor::InnerIterator it(f, k); it; ++it) ++counts[rowMajor ? it.row() : it.col()];
406 return counts;
407 }
408 static auto transposed(const Factor& f) { return f.transpose(); }
409 static auto conjugated(const Factor& f) { return f.conjugate(); }
410 static auto adjointed(const Factor& f) { return f.adjoint(); }
414 static InverseFactor inversed(const Factor& f) {
415 const ColMajorFactor colMajor(f);
416 SparseLU<ColMajorFactor> lu(colMajor);
417 if (lu.info() != Success) return nanMatrix(f.rows(), f.cols());
418 return lu.solve(InverseFactor::Identity(f.rows(), f.cols()));
419 }
423 static InverseFactor nanMatrix(Index rows, Index cols) {
424 return InverseFactor::Constant(rows, cols, Scalar(NumTraits<RealScalar>::quiet_NaN()));
425 }
426 static InverseFactor denseFactor(const Factor& f) { return InverseFactor(f); }
427 static const Factor& blockOperand(const Factor& f) { return f; }
428 static bool isSquareIdentity(const Factor&) { return false; }
429 template <typename Dst, typename Alpha, typename Xpr>
430 static void addLeftProduct(Dst& dst, const Alpha& alpha, const Factor& f, const Xpr& X) {
431 dst.noalias() += alpha * (f * X);
432 }
433 template <typename Dst, typename Alpha, typename Xpr, typename Work>
434 static void addRightProduct(Dst& dst, const Alpha& alpha, const Xpr& X, const Factor& f, Work&) {
435 dst.noalias() += alpha * (X * f.transpose());
436 }
437 static int exponentBound(const Factor& f) { return structured_exponent_bound(f.coeffs()); }
438 static Scalar balancedDet(const Factor& M, Index& exponent) {
439 // Scale small factors up, exactly, so that the elimination runs on normal
440 // numbers: a subnormal near 2^e carries only e - min_exponent + digits bits.
441 // Scaling down could erase small pivots in factors with a wide exponent range.
442 const int bound = exponentBound(M);
443 const int scaleExponent = numext::mini(bound, 0);
444 ColMajorFactor normalized(M);
445 if (scaleExponent != 0) {
446 auto values = normalized.coeffs();
447 structured_ldexp_entries_exact(values, -scaleExponent);
448 }
449 kron_sparse_lu<ColMajorFactor> lu;
450 lu.compute(normalized);
451 if (lu.info() == Success) {
452 exponent += M.rows() * scaleExponent;
453 return lu.balancedDet(exponent);
454 }
455 // An aborted factorization met an exactly zero pivot column -- an exactly
456 // singular factor -- unless a non-finite entry defeated the pivot search.
457 return M.coeffs().allFinite() ? Scalar(0) : Scalar(NumTraits<RealScalar>::quiet_NaN());
458 }
459};
460
461template <typename Factor>
462struct kron_factor_ops<Factor, kKronIdentityFactor> {
463 // Identity factor, possibly rectangular: the m x n matrix I with I(i,i) = 1
464 // for i < min(m, n). I X keeps the leading min(m, n) rows of X, X I^T its
465 // leading min(m, n) columns.
466 using Scalar = typename Factor::Scalar;
467 using TransposedFactor = kron_identity_factor<Scalar, Factor::ColsAtCompileTime, Factor::RowsAtCompileTime>;
468 using InverseFactor = Factor;
469 static constexpr bool StoresAllEntries = false;
470
471 static void prepare(Factor&) {}
472 static bool coeffIfStored(const Factor&, Index row, Index col, Scalar& value) {
473 if (row != col) return false;
474 value = Scalar(1);
475 return true;
476 }
477 template <typename Visitor>
478 static void forEachNonZero(const Factor& f, Visitor&& visit) {
479 for (Index k = 0; k < numext::mini(f.rows(), f.cols()); ++k) visit(k, k, Scalar(1));
480 }
481 static Matrix<Index, Dynamic, 1> innerNonZeros(const Factor& f, bool rowMajor) {
482 Matrix<Index, Dynamic, 1> counts = Matrix<Index, Dynamic, 1>::Zero(rowMajor ? f.rows() : f.cols());
483 counts.head(numext::mini(f.rows(), f.cols())).setOnes();
484 return counts;
485 }
486 static TransposedFactor transposed(const Factor& f) { return TransposedFactor(f.cols(), f.rows()); }
487 static const Factor& conjugated(const Factor& f) { return f; }
488 static TransposedFactor adjointed(const Factor& f) { return transposed(f); }
489 static const Factor& inversed(const Factor& f) { return f; }
490 static typename Factor::PlainObject denseFactor(const Factor& f) { return blockOperand(f); }
491 static typename Factor::PlainObject::IdentityReturnType blockOperand(const Factor& f) {
492 return Factor::PlainObject::Identity(f.rows(), f.cols());
493 }
494 static bool isSquareIdentity(const Factor& f) { return f.rows() == f.cols(); }
495 template <typename Dst, typename Alpha, typename Xpr>
496 static void addLeftProduct(Dst& dst, const Alpha& alpha, const Factor& f, const Xpr& X) {
497 const Index k = numext::mini(f.rows(), f.cols());
498 dst.topRows(k) += alpha * X.topRows(k);
499 }
500 template <typename Dst, typename Alpha, typename Xpr, typename Work>
501 static void addRightProduct(Dst& dst, const Alpha& alpha, const Xpr& X, const Factor& f, Work&) {
502 const Index k = numext::mini(f.rows(), f.cols());
503 dst.leftCols(k) += alpha * X.leftCols(k);
504 }
505 static Scalar balancedDet(const Factor&, Index&) { return Scalar(1); }
506};
507
508template <typename LhsMatrix, typename RhsMatrix>
509struct kron_factor_ops<KroneckerOperator<LhsMatrix, RhsMatrix>, kKronKroneckerFactor> {
510 // Nested Kronecker factor K = L (x) R: every operation recurses into L and R.
511 using Factor = KroneckerOperator<LhsMatrix, RhsMatrix>;
512 using Scalar = typename Factor::Scalar;
513 using LhsOps = kron_factor_ops<LhsMatrix>;
514 using RhsOps = kron_factor_ops<RhsMatrix>;
515 using DenseMatrix = Matrix<Scalar, Dynamic, Dynamic, ColMajor>;
516 using TransposedFactor = KroneckerOperator<typename LhsOps::TransposedFactor, typename RhsOps::TransposedFactor>;
517 using InverseFactor = KroneckerOperator<typename LhsOps::InverseFactor, typename RhsOps::InverseFactor>;
518 static constexpr bool StoresAllEntries = LhsOps::StoresAllEntries && RhsOps::StoresAllEntries;
519
520 static void prepare(Factor&) {}
521 static bool coeffIfStored(const Factor& f, Index row, Index col, Scalar& value) {
522 const Index m2 = f.rhs().rows(), n2 = f.rhs().cols();
523 Scalar a, b;
524 if (!LhsOps::coeffIfStored(f.lhs(), row / m2, col / n2, a) ||
525 !RhsOps::coeffIfStored(f.rhs(), row % m2, col % n2, b))
526 return false;
527 value = a * b;
528 return true;
529 }
530 // Visits (iA m2 + iB, jA n2 + jB) with the loop over L outermost, so a fixed
531 // column is visited in increasing row and a fixed row in increasing column,
532 // the order the sparse materialization inserts in.
533 template <typename Visitor>
534 static void forEachNonZero(const Factor& f, Visitor&& visit) {
535 using RhsVisitable = kron_factor_visitable<RhsMatrix>;
536 const Index m2 = f.rhs().rows(), n2 = f.rhs().cols();
537 const auto& R = RhsVisitable::get(f.rhs());
538 LhsOps::forEachNonZero(f.lhs(), [&R, &visit, m2, n2](Index iA, Index jA, const Scalar& a) {
539 kron_factor_ops<typename RhsVisitable::type>::forEachNonZero(
540 R, [&visit, m2, n2, iA, jA, &a](Index iB, Index jB, const Scalar& b) {
541 visit(iA * m2 + iB, jA * n2 + jB, a * b);
542 });
543 });
544 }
545 static Matrix<Index, Dynamic, 1> innerNonZeros(const Factor& f, bool rowMajor) {
546 const Matrix<Index, Dynamic, Dynamic, ColMajor> counts =
547 RhsOps::innerNonZeros(f.rhs(), rowMajor) * LhsOps::innerNonZeros(f.lhs(), rowMajor).transpose();
548 return counts.reshaped();
549 }
550 static TransposedFactor transposed(const Factor& f) { return f.transpose(); }
551 static Factor conjugated(const Factor& f) { return f.conjugate(); }
552 static TransposedFactor adjointed(const Factor& f) { return f.adjoint(); }
553 static InverseFactor inversed(const Factor& f) { return f.inverse(); }
554 static DenseMatrix denseFactor(const Factor& f) { return DenseMatrix(f); }
555 static DenseMatrix blockOperand(const Factor& f) { return denseFactor(f); }
556 static bool isSquareIdentity(const Factor&) { return false; }
557 template <typename Dst, typename Alpha, typename Xpr>
558 static void addLeftProduct(Dst& dst, const Alpha& alpha, const Factor& f, const Xpr& X) {
559 f.addProduct(dst, X, alpha);
560 }
561 // X K^T without forming X^T: with X_j = X(:, j nR : (j+1) nR - 1), j < nL,
562 // X (L (x) R)^T = W L^T, W(:, j) = vec(X_j R^T), W (p mR) x nL in work,
563 // and W = X_{[p nR x nL]} when R is a square identity.
564 template <typename Dst, typename Alpha, typename Xpr, typename Work>
565 static void addRightProduct(Dst& dst, const Alpha& alpha, const Xpr& X, const Factor& f, Work& work) {
566 const Index p = X.rows(), mL = f.lhs().rows(), nL = f.lhs().cols(), mR = f.rhs().rows(), nR = f.rhs().cols();
567 auto dstL = dst.reshaped(p * mR, mL);
568 if (RhsOps::isSquareIdentity(f.rhs())) {
569 LhsOps::addRightProduct(dstL, alpha, X.reshaped(p * nR, nL), f.lhs(), work);
570 } else if (LhsOps::isSquareIdentity(f.lhs())) {
571 for (Index j = 0; j < nL; ++j) {
572 auto dstj = dst.middleCols(j * mR, mR);
573 RhsOps::addRightProduct(dstj, alpha, X.middleCols(j * nR, nR), f.rhs(), work);
574 }
575 } else {
576 Work inner; // scratch for L and R while work holds W; allocated only if one is nested
577 work.setZero(p * mR, nL);
578 for (Index j = 0; j < nL; ++j) {
579 auto Wj = work.col(j).reshaped(p, mR);
580 RhsOps::addRightProduct(Wj, Alpha(1), X.middleCols(j * nR, nR), f.rhs(), inner);
581 }
582 LhsOps::addRightProduct(dstL, alpha, work, f.lhs(), inner);
583 }
584 }
585 static Scalar balancedDet(const Factor& f, Index& exponent) { return f.balancedDeterminant(exponent); }
586};
587
592template <typename Factor, int Kind = kron_factor_kind<Factor>()>
593class kron_factor_solver {
594 public:
595 using Scalar = typename Factor::Scalar;
596 using DenseMatrix = Matrix<Scalar, Dynamic, Dynamic, ColMajor>;
597
598 explicit kron_factor_solver(const Factor& f) : m_lu(f) {}
600 template <typename Xpr>
601 DenseMatrix solveLeft(const Xpr& M) const {
602 return m_lu.solve(M);
603 }
605 template <typename Xpr>
606 DenseMatrix solveTransposedRight(const Xpr& M) const {
607 return m_lu.solve(M.transpose()).transpose();
608 }
609
610 private:
611 PartialPivLU<DenseMatrix> m_lu;
612};
613
614template <typename Factor>
615class kron_factor_solver<Factor, kKronDiagonalFactor> {
616 public:
617 using Scalar = typename Factor::Scalar;
618 using DenseMatrix = Matrix<Scalar, Dynamic, Dynamic, ColMajor>;
619
620 explicit kron_factor_solver(const Factor& f) : m_d(f.diagonal()) {}
621 template <typename Xpr>
622 DenseMatrix solveLeft(const Xpr& M) const {
623 return (M.array().colwise() / m_d.array()).matrix();
624 }
625 template <typename Xpr>
626 DenseMatrix solveTransposedRight(const Xpr& M) const {
627 return (M.array().rowwise() / m_d.transpose().array()).matrix();
628 }
629
630 private:
631 typename Factor::DiagonalVectorType m_d;
632};
633
634template <typename Factor>
635class kron_factor_solver<Factor, kKronSparseFactor> {
636 public:
637 using Scalar = typename Factor::Scalar;
638 using DenseMatrix = Matrix<Scalar, Dynamic, Dynamic, ColMajor>;
639 using ColMajorFactor = SparseMatrix<Scalar, ColMajor, typename Factor::StorageIndex>;
640
641 explicit kron_factor_solver(const Factor& f) {
642 m_lu.compute(ColMajorFactor(f));
643 m_factorized = m_lu.info() == Success; // else the solves fill with NaN, see kron_factor_ops::nanMatrix
644 }
645 template <typename Xpr>
646 DenseMatrix solveLeft(const Xpr& M) const {
647 if (!m_factorized) return kron_factor_ops<Factor>::nanMatrix(M.rows(), M.cols());
648 return m_lu.solve(M);
649 }
650 template <typename Xpr>
651 DenseMatrix solveTransposedRight(const Xpr& M) const {
652 if (!m_factorized) return kron_factor_ops<Factor>::nanMatrix(M.rows(), M.cols());
653 // SparseLU solves into column-major storage only, which a transposed Solve
654 // expression would not evaluate into; the right-hand side stays a view.
655 const DenseMatrix Xt = m_lu.solve(M.transpose());
656 return Xt.transpose();
657 }
658
659 private:
660 bool m_factorized;
661 SparseLU<ColMajorFactor> m_lu;
662};
663
664template <typename Factor>
665class kron_factor_solver<Factor, kKronIdentityFactor> {
666 public:
667 using Scalar = typename Factor::Scalar;
668 using DenseMatrix = Matrix<Scalar, Dynamic, Dynamic, ColMajor>;
669
670 explicit kron_factor_solver(const Factor&) {}
671 template <typename Xpr>
672 DenseMatrix solveLeft(const Xpr& M) const {
673 return M;
674 }
675 template <typename Xpr>
676 DenseMatrix solveTransposedRight(const Xpr& M) const {
677 return M;
678 }
679};
680
683template <typename LhsMatrix, typename RhsMatrix>
684class kron_factor_solver<KroneckerOperator<LhsMatrix, RhsMatrix>, kKronKroneckerFactor> {
685 public:
686 using Factor = KroneckerOperator<LhsMatrix, RhsMatrix>;
687 using Scalar = typename Factor::Scalar;
688 using DenseMatrix = Matrix<Scalar, Dynamic, Dynamic, ColMajor>;
689
690 explicit kron_factor_solver(const Factor& f)
691 : m_solverA(squareFactor(f.lhs())),
692 m_solverB(squareFactor(f.rhs())),
693 m_n1(f.lhs().cols()),
694 m_n2(f.rhs().cols()) {}
695 template <typename Xpr>
696 DenseMatrix solveLeft(const Xpr& M) const {
697 DenseMatrix x(M.rows(), M.cols());
698 Factor::solveWith(m_solverA, m_solverB, m_n1, m_n2, M, x);
699 return x;
700 }
701 // M K^{-T} = (K^{-1} M^T)^T.
702 template <typename Xpr>
703 DenseMatrix solveTransposedRight(const Xpr& M) const {
704 return solveLeft(M.transpose()).transpose();
705 }
706
707 private:
708 template <typename F>
709 static const F& squareFactor(const F& f) {
710 eigen_assert(f.rows() == f.cols() && "KroneckerOperator::solve requires square factors");
711 return f;
712 }
713
714 kron_factor_solver<LhsMatrix> m_solverA;
715 kron_factor_solver<RhsMatrix> m_solverB;
716 Index m_n1, m_n2;
717};
718
727template <typename Factor, int Kind = kron_factor_kind<Factor>()>
728struct kron_factor_spectrum {
729 using Ops = kron_factor_ops<Factor>;
730 using Scalar = typename Factor::Scalar;
731 using RealScalar = typename NumTraits<Scalar>::Real;
732 using ComplexScalar = std::complex<RealScalar>;
733 using DenseMatrix = Matrix<Scalar, Dynamic, Dynamic, ColMajor>;
734 using RealVector = Matrix<RealScalar, Dynamic, 1>;
735 using ComplexMatrix = Matrix<ComplexScalar, Dynamic, Dynamic, ColMajor>;
736 using ComplexVector = Matrix<ComplexScalar, Dynamic, 1>;
737 using Eigenvectors = ComplexMatrix;
738 using SingularVectors = DenseMatrix;
739
740 static ComplexMatrix complexFactor(const Factor& f) { return Ops::denseFactor(f).template cast<ComplexScalar>(); }
741 static ComplexVector eigenvalues(const Factor& f) {
742 ComplexEigenSolver<ComplexMatrix> es(complexFactor(f), /*computeEigenvectors=*/false);
743 eigen_assert(es.info() == Success);
744 return es.eigenvalues();
745 }
746 static Eigenvectors eigenvectors(const Factor& f) {
747 ComplexEigenSolver<ComplexMatrix> es(complexFactor(f));
748 eigen_assert(es.info() == Success);
749 return es.eigenvectors();
750 }
751 static RealVector singularValues(const Factor& f) {
752 return BDCSVD<DenseMatrix>(Ops::denseFactor(f)).singularValues();
753 }
754 static SingularVectors matrixU(const Factor& f) {
755 return BDCSVD<DenseMatrix, ComputeThinU>(Ops::denseFactor(f)).matrixU();
756 }
757 static SingularVectors matrixV(const Factor& f) {
758 return BDCSVD<DenseMatrix, ComputeThinV>(Ops::denseFactor(f)).matrixV();
759 }
760};
761
762template <typename LhsMatrix, typename RhsMatrix>
763struct kron_factor_spectrum<KroneckerOperator<LhsMatrix, RhsMatrix>, kKronKroneckerFactor> {
764 using Factor = KroneckerOperator<LhsMatrix, RhsMatrix>;
765 using Eigenvectors = KroneckerOperator<typename kron_factor_spectrum<LhsMatrix>::Eigenvectors,
766 typename kron_factor_spectrum<RhsMatrix>::Eigenvectors>;
767 using SingularVectors = KroneckerOperator<typename kron_factor_spectrum<LhsMatrix>::SingularVectors,
768 typename kron_factor_spectrum<RhsMatrix>::SingularVectors>;
769
770 static typename Factor::ComplexVector eigenvalues(const Factor& f) { return f.eigenvalues(); }
771 static Eigenvectors eigenvectors(const Factor& f) { return f.eigenvectors(); }
772 static typename Factor::RealVector singularValues(const Factor& f) { return f.singularValues(); }
773 static SingularVectors matrixU(const Factor& f) { return f.matrixU(); }
774 static SingularVectors matrixV(const Factor& f) { return f.matrixV(); }
775};
776
777} // namespace internal
778
907template <typename LhsMatrix, typename RhsMatrix>
908class KroneckerOperator : public EigenBase<KroneckerOperator<LhsMatrix, RhsMatrix>> {
909 public:
910 using Scalar = typename LhsMatrix::Scalar;
911 using RealScalar = typename NumTraits<Scalar>::Real;
912 using StorageIndex = int;
913
914 static_assert(std::is_same<Scalar, typename RhsMatrix::Scalar>::value,
915 "KroneckerOperator requires both factors to have the same scalar type");
916 static_assert((internal::kron_factor_is_dense_matrix<LhsMatrix>::value ||
917 internal::kron_factor_kind<LhsMatrix>() != internal::kKronDenseFactor) &&
918 (internal::kron_factor_is_dense_matrix<RhsMatrix>::value ||
919 internal::kron_factor_kind<RhsMatrix>() != internal::kKronDenseFactor),
920 "KroneckerOperator factors must be plain Matrix, DiagonalMatrix or SparseMatrix types, identity "
921 "factors (makeKroneckerOperator stores an Identity() expression as one), KroneckerOperators or "
922 "KroneckerSums (owning their storage: views and other expressions would dangle)");
923
924 private:
925 // Factor-kind dispatch, see kron_factor_ops.
926 using LhsOps = internal::kron_factor_ops<LhsMatrix>;
927 using RhsOps = internal::kron_factor_ops<RhsMatrix>;
928 using LhsSpectrum = internal::kron_factor_spectrum<LhsMatrix>;
929 using RhsSpectrum = internal::kron_factor_spectrum<RhsMatrix>;
930 template <typename, int>
931 friend struct internal::kron_factor_ops;
932 template <typename, int>
933 friend class internal::kron_factor_solver;
934
935 public:
936 using ComplexScalar = std::complex<RealScalar>;
937 // The vec-trick reshapes below identify a vector of length n1*n2 with an
938 // n2 x n1 matrix whose columns are stacked, so every workspace taking part
939 // in a reshape is pinned to ColMajor explicitly: the semantics must not
940 // change under EIGEN_DEFAULT_TO_ROW_MAJOR.
942 using RealVector = Matrix<RealScalar, Dynamic, 1>;
944 using ComplexVector = Matrix<ComplexScalar, Dynamic, 1>;
945
946 static constexpr int RowsAtCompileTime =
947 internal::size_at_compile_time(LhsMatrix::RowsAtCompileTime, RhsMatrix::RowsAtCompileTime);
948 static constexpr int ColsAtCompileTime =
949 internal::size_at_compile_time(LhsMatrix::ColsAtCompileTime, RhsMatrix::ColsAtCompileTime);
950 static constexpr int MaxRowsAtCompileTime = RowsAtCompileTime;
951 static constexpr int MaxColsAtCompileTime = ColsAtCompileTime;
952 static constexpr int SizeAtCompileTime = internal::size_at_compile_time(RowsAtCompileTime, ColsAtCompileTime);
953 static constexpr int MaxSizeAtCompileTime = SizeAtCompileTime;
954 static constexpr bool IsRowMajor = false;
955 // Deliberately no IsVectorAtCompileTime: Ref<const KroneckerOperator>'s default
956 // StrideType argument reads it, so its absence makes internal::is_ref_compatible
957 // SFINAE to false and keeps the iterative solvers on their matrix-free path.
958
964 template <typename LhsDerived, typename RhsDerived>
966 : m_A(a.derived()), m_B(b.derived()) {
967 eigen_assert(m_A.size() > 0 && m_B.size() > 0 && "KroneckerOperator factors must be non-empty");
968 LhsOps::prepare(m_A);
969 RhsOps::prepare(m_B);
970 }
971
972 EIGEN_DEVICE_FUNC Index rows() const { return m_A.rows() * m_B.rows(); }
973 EIGEN_DEVICE_FUNC Index cols() const { return m_A.cols() * m_B.cols(); }
974
976 const LhsMatrix& lhs() const { return m_A; }
978 const RhsMatrix& rhs() const { return m_B; }
979
981 Scalar coeff(Index row, Index col) const {
982 eigen_assert(row >= 0 && row < rows() && col >= 0 && col < cols());
983 const Index m2 = m_B.rows(), n2 = m_B.cols();
984 Scalar a, b;
985 if (!LhsOps::coeffIfStored(m_A, row / m2, col / n2, a) || !RhsOps::coeffIfStored(m_B, row % m2, col % n2, b))
986 return Scalar(0);
987 return a * b;
988 }
989
993 return {LhsOps::transposed(m_A), RhsOps::transposed(m_B)};
994 }
995
997 KroneckerOperator conjugate() const { return {LhsOps::conjugated(m_A), RhsOps::conjugated(m_B)}; }
998
1002 return {LhsOps::adjointed(m_A), RhsOps::adjointed(m_B)};
1003 }
1004
1021 template <typename Rhs>
1023 EIGEN_STATIC_ASSERT(RowsAtCompileTime == Dynamic || Rhs::RowsAtCompileTime == Dynamic ||
1024 int(RowsAtCompileTime) == int(Rhs::RowsAtCompileTime),
1025 YOU_MIXED_MATRICES_OF_DIFFERENT_SIZES)
1026 const Index n1 = m_A.cols(), n2 = m_B.cols();
1027 eigen_assert(m_A.rows() == n1 && m_B.rows() == n2 && "KroneckerOperator::solve requires square factors");
1028 eigen_assert(b.rows() == n1 * n2 && "right-hand side has the wrong number of rows");
1029 const internal::kron_factor_solver<LhsMatrix> solverA(m_A);
1030 const internal::kron_factor_solver<RhsMatrix> solverB(m_B);
1032 solveWith(solverA, solverB, n1, n2, b.derived(), x);
1033 return x;
1034 }
1035
1058 template <typename Rhs>
1060 EIGEN_STATIC_ASSERT(RowsAtCompileTime == Dynamic || Rhs::RowsAtCompileTime == Dynamic ||
1061 int(RowsAtCompileTime) == int(Rhs::RowsAtCompileTime),
1062 YOU_MIXED_MATRICES_OF_DIFFERENT_SIZES)
1063 const Index m1 = m_A.rows(), m2 = m_B.rows(), n1 = m_A.cols(), n2 = m_B.cols(), r = b.cols();
1064 eigen_assert(b.rows() == m1 * m2 && "right-hand side has the wrong number of rows");
1066 // The pivoted QR squares column norms, which over- or underflow on extreme
1067 // factor magnitudes, so decompose 2^-eA A and 2^-eB B and fold
1068 // 2^-(eA+eB) back into the solution.
1069 DenseMatrix An = LhsOps::denseFactor(m_A), Bn = RhsOps::denseFactor(m_B);
1070 if (!An.allFinite() || !Bn.allFinite()) {
1071 // The pivoted QR would rank a non-finite column out instead of propagating it.
1073 return x;
1074 }
1075 const int eA = internal::structured_exponent_bound(An), eB = internal::structured_exponent_bound(Bn);
1076 internal::structured_ldexp_entries_exact(An, -eA);
1077 internal::structured_ldexp_entries_exact(Bn, -eB);
1078 const CompleteOrthogonalDecomposition<DenseMatrix> codA(An), codB(Bn);
1079 typename internal::nested_eval<Rhs, 1>::type actualRhs(b.derived());
1080 DenseMatrix S, Z, X;
1081 const Index chunk = internal::kron_rhs_chunk<Scalar>(m1 * m2 + n2 * m1 + n1 * n2);
1082 for (Index k0 = 0; k0 < r; k0 += chunk) {
1083 const Index c = numext::mini(chunk, r - k0);
1084 internal::kron_stack_columns(S, actualRhs.middleCols(k0, c), m2, m1);
1085 Z = codB.solve(S.reshaped(m2, c * m1));
1086 X = codA.solve(Z.reshaped(n2 * c, m1).transpose()).transpose();
1087 internal::structured_ldexp_entries(X, -(eA + eB));
1088 for (Index k = 0; k < c; ++k) x.col(k0 + k).reshaped(n2, n1) = X.middleRows(k * n2, n2);
1089 }
1090 return x;
1091 }
1092
1108 Index rank() const {
1109 BDCSVD<DenseMatrix> svdA(LhsOps::denseFactor(m_A)), svdB(RhsOps::denseFactor(m_B));
1110 const RealVector sa = svdA.singularValues(), sb = svdB.singularValues();
1111 // An exactly zero factor zeroes the whole operator (and would make the
1112 // ratios below 0/0).
1113 if (sa[0] == RealScalar(0) || sb[0] == RealScalar(0)) return 0;
1114 const RealScalar tol = relativeRankThreshold();
1115 const SingularModes modesA(sa), modesB(sb);
1116 Index r = 0;
1117 for (Index i = 0; i < sa.size(); ++i)
1118 for (Index j = 0; j < sb.size(); ++j)
1119 // Negated ratio comparison so NaN ratios count as non-zero.
1120 if (modesA.retains(modesB, i, j, tol)) ++r;
1121 return r;
1122 }
1123
1130 eigen_assert(m_A.rows() == m_A.cols() && m_B.rows() == m_B.cols() &&
1131 "KroneckerOperator::inverse requires square factors");
1132 return {LhsOps::inversed(m_A), RhsOps::inversed(m_B)};
1133 }
1134
1147 Scalar determinant() const {
1148 eigen_assert(m_A.rows() == m_A.cols() && m_B.rows() == m_B.cols() &&
1149 "KroneckerOperator::determinant requires square factors");
1150 Index exponent = 0;
1151 const Scalar mant = balancedDeterminant(exponent);
1152 return internal::structured_ldexp_clamped(mant, exponent);
1153 }
1154
1163 ComplexVector eigenvalues() const {
1164 eigen_assert(m_A.rows() == m_A.cols() && m_B.rows() == m_B.cols() &&
1165 "KroneckerOperator::eigenvalues requires square factors");
1166 // Column-major stacking of the nB x nA outer product puts mu_j(B) lambda_i(A)
1167 // at index i*nB + j, the Kronecker order.
1168 return (RhsSpectrum::eigenvalues(m_B) * LhsSpectrum::eigenvalues(m_A).transpose()).reshaped();
1169 }
1170
1176 eigen_assert(m_A.rows() == m_A.cols() && m_B.rows() == m_B.cols() &&
1177 "KroneckerOperator::eigenvectors requires square factors");
1178 return {LhsSpectrum::eigenvectors(m_A), RhsSpectrum::eigenvectors(m_B)};
1179 }
1180
1188 RealVector singularValues() const {
1189 // Column-major stacking of the kB x kA outer product puts sigma_j(B) sigma_i(A)
1190 // at index i*kB + j, the Kronecker order.
1191 return (RhsSpectrum::singularValues(m_B) * LhsSpectrum::singularValues(m_A).transpose()).reshaped();
1192 }
1193
1198 return {LhsSpectrum::matrixU(m_A), RhsSpectrum::matrixU(m_B)};
1199 }
1200
1205 return {LhsSpectrum::matrixV(m_A), RhsSpectrum::matrixV(m_B)};
1206 }
1207
1211 template <typename Dest>
1212 void evalTo(Dest& dst) const {
1213 evalToImpl(dst, IsSparseDestination<Dest>());
1214 }
1215
1217 template <typename Dest>
1218 void addTo(Dest& dst) const {
1219 addToImpl(dst, IsSparseDestination<Dest>());
1220 }
1221
1223 template <typename Dest>
1224 void subTo(Dest& dst) const {
1225 subToImpl(dst, IsSparseDestination<Dest>());
1226 }
1227
1228 private:
1229 template <typename Dest>
1230 using IsSparseDestination = std::is_same<typename internal::traits<Dest>::StorageKind, Sparse>;
1231
1237 template <typename Dest>
1238 void evalToImpl(Dest& dst, std::false_type) const {
1239 const Index m2 = m_B.rows(), n2 = m_B.cols();
1240 if (!LhsOps::StoresAllEntries) dst.setZero();
1241 const auto& B = RhsOps::blockOperand(m_B);
1242 LhsOps::forEachNonZero(
1243 m_A, [&dst, &B, m2, n2](Index i, Index j, const Scalar& a) { dst.block(i * m2, j * n2, m2, n2) = a * B; });
1244 }
1245
1246 template <typename Dest>
1247 void addToImpl(Dest& dst, std::false_type) const {
1248 const Index m2 = m_B.rows(), n2 = m_B.cols();
1249 const auto& B = RhsOps::blockOperand(m_B);
1250 LhsOps::forEachNonZero(
1251 m_A, [&dst, &B, m2, n2](Index i, Index j, const Scalar& a) { dst.block(i * m2, j * n2, m2, n2) += a * B; });
1252 }
1253
1254 template <typename Dest>
1255 void subToImpl(Dest& dst, std::false_type) const {
1256 const Index m2 = m_B.rows(), n2 = m_B.cols();
1257 const auto& B = RhsOps::blockOperand(m_B);
1258 LhsOps::forEachNonZero(
1259 m_A, [&dst, &B, m2, n2](Index i, Index j, const Scalar& a) { dst.block(i * m2, j * n2, m2, n2) -= a * B; });
1260 }
1261
1271 template <typename Dest>
1272 void evalToImpl(Dest& S, std::true_type) const {
1273 using LhsVisitable = internal::kron_factor_visitable<LhsMatrix>;
1274 using RhsVisitable = internal::kron_factor_visitable<RhsMatrix>;
1275 using VisitedLhsOps = internal::kron_factor_ops<typename LhsVisitable::type>;
1276 using VisitedRhsOps = internal::kron_factor_ops<typename RhsVisitable::type>;
1277 const Index m2 = m_B.rows(), n2 = m_B.cols();
1278 const auto& A = LhsVisitable::get(m_A);
1279 const auto& B = RhsVisitable::get(m_B);
1280 S.resize(rows(), cols());
1281 using IndexVector = Matrix<Index, Dynamic, 1>;
1282 const IndexVector nnzA = VisitedLhsOps::innerNonZeros(A, Dest::IsRowMajor);
1283 const IndexVector nnzB = VisitedRhsOps::innerNonZeros(B, Dest::IsRowMajor);
1284 // Inner vectors kA of A and kB of B meet in inner vector kA * nnzB.size() + kB
1285 // of the product: the column-major stacking of the count outer product.
1286 const Matrix<Index, Dynamic, Dynamic, ColMajor> counts = nnzB * nnzA.transpose();
1287 S.reserve(counts.reshaped());
1288 VisitedLhsOps::forEachNonZero(A, [&S, &B, m2, n2](Index iA, Index jA, const Scalar& a) {
1289 VisitedRhsOps::forEachNonZero(B, [&S, m2, n2, iA, jA, &a](Index iB, Index jB, const Scalar& b) {
1290 S.insert(iA * m2 + iB, jA * n2 + jB) = a * b;
1291 });
1292 });
1293 S.makeCompressed();
1294 }
1295
1296 template <typename Dest>
1297 void addToImpl(Dest& dst, std::true_type) const {
1298 typename Dest::PlainObject product;
1299 evalTo(product);
1300 dst += product;
1301 }
1302
1303 template <typename Dest>
1304 void subToImpl(Dest& dst, std::true_type) const {
1305 typename Dest::PlainObject product;
1306 evalTo(product);
1307 dst -= product;
1308 }
1309
1310 public:
1316 template <typename Rhs>
1318 EIGEN_STATIC_ASSERT(ColsAtCompileTime == Dynamic || Rhs::RowsAtCompileTime == Dynamic ||
1319 int(ColsAtCompileTime) == int(Rhs::RowsAtCompileTime),
1320 INVALID_MATRIX_PRODUCT)
1321 eigen_assert(x.rows() == cols() && "invalid product: dimensions do not match");
1322 return Product<KroneckerOperator, Rhs>(*this, x.derived());
1323 }
1324
1343 template <typename Dest, typename Rhs, typename ProductScalar>
1344 void addProduct(Dest& dst, const Rhs& rhs, const ProductScalar& alpha) const {
1345 using ProductMatrix = Matrix<ProductScalar, Dynamic, Dynamic, ColMajor>; // ColMajor: see DenseMatrix
1346 const Index m1 = m_A.rows(), n1 = m_A.cols(), m2 = m_B.rows(), n2 = m_B.cols(), r = rhs.cols();
1347 eigen_assert(rhs.rows() == n1 * n2 && "invalid product: dimensions do not match");
1348 typename internal::nested_eval<Rhs, 1>::type actualRhs(rhs);
1349 const bool skipA = LhsOps::isSquareIdentity(m_A), skipB = RhsOps::isSquareIdentity(m_B);
1350 ProductMatrix X, BX, Y, work;
1351 const Index chunk = internal::kron_rhs_chunk<ProductScalar>(n1 * n2 + m2 * n1 + m2 * m1);
1352 for (Index k0 = 0; k0 < r; k0 += chunk) {
1353 const Index c = numext::mini(chunk, r - k0);
1354 if (c == 1) {
1355 const auto Xk = actualRhs.col(k0).reshaped(n2, n1);
1356 auto Yk = dst.col(k0).reshaped(m2, m1);
1357 if (skipA) {
1358 RhsOps::addLeftProduct(Yk, alpha, m_B, Xk);
1359 } else if (skipB) {
1360 LhsOps::addRightProduct(Yk, alpha, Xk, m_A, work);
1361 } else {
1362 BX.setZero(m2, n1);
1363 RhsOps::addLeftProduct(BX, ProductScalar(1), m_B, Xk);
1364 LhsOps::addRightProduct(Yk, alpha, BX, m_A, work);
1365 }
1366 continue;
1367 }
1368 internal::kron_stack_columns(X, actualRhs.middleCols(k0, c), n2, n1);
1369 if (!skipB) {
1370 BX.setZero(m2, c * n1);
1371 RhsOps::addLeftProduct(BX, ProductScalar(1), m_B, X.reshaped(n2, c * n1));
1372 }
1373 const auto BXhat = (skipB ? X : BX).reshaped(m2 * c, n1);
1374 if (!skipA) {
1375 Y.setZero(m2 * c, m1);
1376 LhsOps::addRightProduct(Y, ProductScalar(1), BXhat, m_A, work);
1377 }
1378 for (Index k = 0; k < c; ++k) {
1379 auto Yk = dst.col(k0 + k).reshaped(m2, m1);
1380 if (skipA)
1381 Yk += alpha * BXhat.middleRows(k * m2, m2);
1382 else
1383 Yk += alpha * Y.middleRows(k * m2, m2);
1384 }
1385 }
1386 }
1387
1388 private:
1400 RealScalar relativeRankThreshold() const {
1401 return RealScalar(numext::mini(rows(), cols())) * NumTraits<RealScalar>::epsilon();
1402 }
1403
1404 // Factor-level ratios and frexp decompositions, independent of the other factor.
1405 struct SingularModes {
1406 explicit SingularModes(const RealVector& s) : ratios(s.size()), mantissas(s.size()), exponents(s.size()) {
1407 for (Index i = 0; i < s.size(); ++i) {
1408 ratios[i] = s[i] / s[0];
1409 exponents[i] = 0;
1410 EIGEN_USING_STD(frexp);
1411 mantissas[i] = frexp(s[i], &exponents[i]);
1412 }
1413 }
1414
1415 bool retains(const SingularModes& other, Index i, Index j, RealScalar tol) const {
1416 // Negation counts NaN ratios as retained, matching SVDBase.
1417 if (ratios[i] * other.ratios[j] < tol) return false;
1418 const RealScalar m = mantissas[i] * other.mantissas[j];
1419 if (!(numext::isfinite)(m)) return true;
1420 if (m == RealScalar(0)) return false;
1421 // s_i*s_j = m*2^(e_i+e_j), with m in [0.25,1). Test the
1422 // smallest-normal clamp without forming an overflowing/underflowing product.
1423 const int e = exponents[i] + other.exponents[j] - (m < RealScalar(0.5) ? 1 : 0);
1424 return e >= std::numeric_limits<RealScalar>::min_exponent;
1425 }
1426
1427 RealVector ratios, mantissas;
1428 Matrix<int, Dynamic, 1> exponents;
1429 };
1430
1433 Scalar balancedDeterminant(Index& exponent) const {
1434 eigen_assert(m_A.rows() == m_A.cols() && m_B.rows() == m_B.cols() &&
1435 "KroneckerOperator::determinant requires square factors");
1436 const Scalar mant = balancedDetPow(m_A, m_B.cols(), exponent);
1437 return internal::structured_balance(mant * balancedDetPow(m_B, m_A.cols(), exponent), exponent);
1438 }
1439
1444 template <typename SolverA, typename SolverB, typename Rhs, typename Dest>
1445 static void solveWith(const SolverA& solverA, const SolverB& solverB, Index n1, Index n2, const Rhs& b, Dest& x) {
1446 typename internal::nested_eval<Rhs, 1>::type actualRhs(b);
1447 DenseMatrix S, Z, X;
1448 const Index r = b.cols();
1449 const Index chunk = internal::kron_rhs_chunk<Scalar>(3 * n1 * n2);
1450 for (Index k0 = 0; k0 < r; k0 += chunk) {
1451 const Index c = numext::mini(chunk, r - k0);
1452 internal::kron_stack_columns(S, actualRhs.middleCols(k0, c), n2, n1);
1453 Z = solverB.solveLeft(S.reshaped(n2, c * n1));
1454 X = solverA.solveTransposedRight(Z.reshaped(n2 * c, n1));
1455 for (Index k = 0; k < c; ++k) x.col(k0 + k).reshaped(n2, n1) = X.middleRows(k * n2, n2);
1456 }
1457 }
1458
1468 template <typename Factor>
1469 static Scalar balancedDetPow(const Factor& M, Index power, Index& exponent) {
1470 Index e = 0;
1471 const Scalar m = internal::kron_factor_ops<Factor>::balancedDet(M, e);
1472 Scalar r(1);
1473 Index er = 0;
1474 for (Index k = 0; k < power; ++k) r = internal::structured_balance(r * m, er);
1475 exponent += power * e + er;
1476 return r;
1477 }
1478
1479 LhsMatrix m_A;
1480 RhsMatrix m_B;
1481};
1482
1491template <typename LhsDerived, typename RhsDerived>
1492KroneckerOperator<typename internal::kron_factor_storage<LhsDerived>::type,
1493 typename internal::kron_factor_storage<RhsDerived>::type>
1495 return {a.derived(), b.derived()};
1496}
1497
1501template <typename D1, typename D2, typename D3, typename... Rest>
1503 const Rest&... rest) {
1504 return makeKroneckerOperator(a, makeKroneckerOperator(b, c, rest...));
1505}
1506
1507namespace internal {
1508
1509template <typename LhsMatrix, typename RhsMatrix, typename Rhs, int ProductTag>
1510struct generic_product_impl<KroneckerOperator<LhsMatrix, RhsMatrix>, Rhs, StructuredShape, DenseShape, ProductTag>
1511 : structured_product_impl<KroneckerOperator<LhsMatrix, RhsMatrix>, Rhs> {};
1512
1513} // namespace internal
1514
1515} // namespace Eigen
1516
1517#endif // EIGEN_STRUCTURED_KRONECKER_OPERATOR_H
The Kronecker product as an implicit operator that is never materialized.
Definition KroneckerOperator.h:908
KroneckerOperator< typename LhsOps::InverseFactor, typename RhsOps::InverseFactor > inverse() const
Definition KroneckerOperator.h:1129
KroneckerOperator< typename LhsOps::TransposedFactor, typename RhsOps::TransposedFactor > adjoint() const
Definition KroneckerOperator.h:1001
KroneckerOperator< typename LhsSpectrum::SingularVectors, typename RhsSpectrum::SingularVectors > matrixU() const
Definition KroneckerOperator.h:1197
Index rank() const
Definition KroneckerOperator.h:1108
Scalar determinant() const
Definition KroneckerOperator.h:1147
const LhsMatrix & lhs() const
Definition KroneckerOperator.h:976
KroneckerOperator< typename LhsOps::TransposedFactor, typename RhsOps::TransposedFactor > transpose() const
Definition KroneckerOperator.h:992
ComplexVector eigenvalues() const
Definition KroneckerOperator.h:1163
Product< KroneckerOperator, Rhs > operator*(const MatrixBase< Rhs > &x) const
Definition KroneckerOperator.h:1317
KroneckerOperator(const EigenBase< LhsDerived > &a, const EigenBase< RhsDerived > &b)
Definition KroneckerOperator.h:965
KroneckerOperator< typename LhsSpectrum::SingularVectors, typename RhsSpectrum::SingularVectors > matrixV() const
Definition KroneckerOperator.h:1204
Matrix< Scalar, ColsAtCompileTime, Rhs::ColsAtCompileTime > leastSquaresSolve(const MatrixBase< Rhs > &b) const
Definition KroneckerOperator.h:1059
const RhsMatrix & rhs() const
Definition KroneckerOperator.h:978
KroneckerOperator< typename LhsSpectrum::Eigenvectors, typename RhsSpectrum::Eigenvectors > eigenvectors() const
Definition KroneckerOperator.h:1175
Matrix< Scalar, ColsAtCompileTime, Rhs::ColsAtCompileTime > solve(const MatrixBase< Rhs > &b) const
Definition KroneckerOperator.h:1022
KroneckerOperator conjugate() const
Definition KroneckerOperator.h:997
RealVector singularValues() const
Definition KroneckerOperator.h:1188
Scalar coeff(Index row, Index col) const
Definition KroneckerOperator.h:981
The Kronecker sum of two square matrices as an implicit operator that is never materialized.
Definition KroneckerSum.h:233
Derived & setConstant(Index rows, Index cols, const Scalar &val)
const SingularValuesType & singularValues() const
ComputationInfo info() const
KroneckerOperator< typename internal::kron_factor_storage< LhsDerived >::type, typename internal::kron_factor_storage< RhsDerived >::type > makeKroneckerOperator(const EigenBase< LhsDerived > &a, const EigenBase< RhsDerived > &b)
Definition KroneckerOperator.h:1494
Namespace containing all symbols from the Eigen library.
constexpr Derived & derived()