Eigen  5.0.1
 
Loading...
Searching...
No Matches
IterativeSolverBase.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2011-2014 Gael Guennebaud <gael.guennebaud@inria.fr>
5//
6// This Source Code Form is subject to the terms of the Mozilla
7// Public License v. 2.0. If a copy of the MPL was not distributed
8// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
9// SPDX-License-Identifier: MPL-2.0
10
11#ifndef EIGEN_ITERATIVE_SOLVER_BASE_H
12#define EIGEN_ITERATIVE_SOLVER_BASE_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20
21template <typename T>
22auto is_ref_compatible_test(T& matrix) -> decltype(Ref<const T>(matrix), std::true_type());
23std::false_type is_ref_compatible_test(...);
24
25template <typename MatrixType>
26using is_ref_compatible = decltype(is_ref_compatible_test(std::declval<remove_all_t<MatrixType>&>()));
27
28// Returns a \a rows x \a cols matrix whose columns are an orthonormal basis of a random subspace,
29// obtained by QR-orthonormalizing a random matrix. The IDR(s)-type solvers use this to build the
30// shadow space; the basis only has to be (almost surely) non-degenerate, so any random seed is fine.
31template <typename MatrixType>
32MatrixType random_orthonormal_basis(Index rows, Index cols) {
33 HouseholderQR<MatrixType> qr(MatrixType::Random(rows, cols));
34 return qr.householderQ() * MatrixType::Identity(rows, cols);
35}
36
37template <typename RealScalar>
38EIGEN_DEVICE_FUNC RealScalar iterative_solver_scaling_factor(RealScalar norm) {
39 // Scale only when forming a quadratic recurrence term from norm may underflow or overflow.
40 const RealScalar sqrtMin = numext::sqrt((std::numeric_limits<RealScalar>::min)());
41 const RealScalar sqrtMax = numext::sqrt((std::numeric_limits<RealScalar>::max)());
42 return norm < sqrtMin || norm > sqrtMax ? norm : RealScalar(1);
43}
44
45template <typename MatrixType, bool MatrixFree = !internal::is_ref_compatible<MatrixType>::value>
46class generic_matrix_wrapper;
47
48// We have an explicit matrix at hand, compatible with Ref<>
49template <typename MatrixType>
50class generic_matrix_wrapper<MatrixType, false> {
51 public:
52 using ActualMatrixType = Ref<const MatrixType>;
53 template <int UpLo>
54 struct ConstSelfAdjointViewReturnType {
55 using Type = typename ActualMatrixType::template ConstSelfAdjointViewReturnType<UpLo>::Type;
56 };
57
58 enum { MatrixFree = false };
59
60 // Default-construct: passing (0,0) trips the size assertion for fixed-size MatrixType (#1704).
61 generic_matrix_wrapper() : m_dummy(), m_matrix(m_dummy) {}
62
63 template <typename InputType>
64 generic_matrix_wrapper(const InputType& mat) : m_matrix(mat) {}
65
66 const ActualMatrixType& matrix() const { return m_matrix; }
67
68 template <typename MatrixDerived>
69 void grab(const EigenBase<MatrixDerived>& mat) {
70 internal::destroy_at(&m_matrix);
71 internal::construct_at(&m_matrix, mat.derived());
72 }
73
74 void grab(const Ref<const MatrixType>& mat) {
75 if (&(mat.derived()) != &m_matrix) {
76 internal::destroy_at(&m_matrix);
77 internal::construct_at(&m_matrix, mat);
78 }
79 }
80
81 protected:
82 MatrixType m_dummy; // used to default initialize the Ref<> object
83 ActualMatrixType m_matrix;
84};
85
86// MatrixType is not compatible with Ref<> -> matrix-free wrapper
87template <typename MatrixType>
88class generic_matrix_wrapper<MatrixType, true> {
89 public:
90 using ActualMatrixType = MatrixType;
91 template <int UpLo>
92 struct ConstSelfAdjointViewReturnType {
93 using Type = ActualMatrixType;
94 };
95
96 enum { MatrixFree = true };
97
98 generic_matrix_wrapper() = default;
99
100 generic_matrix_wrapper(const MatrixType& mat) : mp_matrix(&mat) {}
101
102 const ActualMatrixType& matrix() const { return *mp_matrix; }
103
104 void grab(const MatrixType& mat) { mp_matrix = &mat; }
105
106 protected:
107 const ActualMatrixType* mp_matrix = nullptr;
108};
109
110} // namespace internal
111
117template <typename Derived>
118class IterativeSolverBase : public SparseSolverBase<Derived> {
119 protected:
120 using Base = SparseSolverBase<Derived>;
121 using Base::m_isInitialized;
122
123 public:
124 using MatrixType = typename internal::traits<Derived>::MatrixType;
125 using Preconditioner = typename internal::traits<Derived>::Preconditioner;
126 using Scalar = typename MatrixType::Scalar;
127 using StorageIndex = typename MatrixType::StorageIndex;
128 using RealScalar = typename MatrixType::RealScalar;
129
130 enum { ColsAtCompileTime = MatrixType::ColsAtCompileTime, MaxColsAtCompileTime = MatrixType::MaxColsAtCompileTime };
131
132 public:
133 using Base::derived;
134
137
148 template <typename MatrixDerived>
149 explicit IterativeSolverBase(const EigenBase<MatrixDerived>& A) : m_matrixWrapper(A.derived()) {
150 init();
151 compute(matrix());
152 }
153
155
161 template <typename MatrixDerived>
163 grab(A.derived());
164 m_preconditioner.analyzePattern(matrix());
165 m_isInitialized = true;
166 m_analysisIsOk = true;
167 m_info = m_preconditioner.info();
168 return derived();
169 }
170
181 template <typename MatrixDerived>
183 eigen_assert(m_analysisIsOk && "You must first call analyzePattern()");
184 grab(A.derived());
185 m_preconditioner.factorize(matrix());
186 m_factorizationIsOk = true;
187 m_info = m_preconditioner.info();
188 return derived();
189 }
190
201 template <typename MatrixDerived>
203 grab(A.derived());
204 m_preconditioner.compute(matrix());
205 m_isInitialized = true;
206 m_analysisIsOk = true;
207 m_factorizationIsOk = true;
208 m_info = m_preconditioner.info();
209 return derived();
210 }
211
213 constexpr Index rows() const noexcept { return matrix().rows(); }
214
216 constexpr Index cols() const noexcept { return matrix().cols(); }
217
221 RealScalar tolerance() const { return m_tolerance; }
222
230 Derived& setTolerance(const RealScalar& tolerance) {
231 m_tolerance = tolerance;
232 return derived();
233 }
234
236 Preconditioner& preconditioner() { return m_preconditioner; }
237
239 const Preconditioner& preconditioner() const { return m_preconditioner; }
240
245 Index maxIterations() const { return (m_maxIterations < 0) ? 2 * matrix().cols() : m_maxIterations; }
246
250 Derived& setMaxIterations(Index maxIters) {
251 m_maxIterations = maxIters;
252 return derived();
253 }
254
256 Index iterations() const {
257 eigen_assert(m_isInitialized && "IterativeSolverBase is not initialized.");
258 return m_iterations;
259 }
260
266 RealScalar error() const {
267 eigen_assert(m_isInitialized && "IterativeSolverBase is not initialized.");
268 return m_error;
269 }
270
276 template <typename Rhs, typename Guess>
277 inline const SolveWithGuess<Derived, Rhs, Guess> solveWithGuess(const MatrixBase<Rhs>& b, const Guess& x0) const {
278 eigen_assert(m_isInitialized && "Solver is not initialized.");
279 eigen_assert(derived().rows() == b.rows() && "solve(): invalid number of rows of the right hand side matrix b");
280 return SolveWithGuess<Derived, Rhs, Guess>(derived(), b.derived(), x0);
281 }
282
293 template <typename Rhs, typename Dest>
294 void solveWithGuessInPlace(const Rhs& b, Dest& x) const {
295 eigen_assert(m_isInitialized && "Solver is not initialized.");
296 eigen_assert(derived().rows() == b.rows() &&
297 "solveWithGuessInPlace(): invalid number of rows of the right hand side b");
298 derived()._solve_vector_with_guess_impl(b, x);
299 }
300
303 eigen_assert(m_isInitialized && "IterativeSolverBase is not initialized.");
304 return m_info;
305 }
306
308 template <typename Rhs, typename DestDerived>
309 void _solve_with_guess_impl(const Rhs& b, SparseMatrixBase<DestDerived>& aDest) const {
310 eigen_assert(rows() == b.rows());
311
312 Index rhsCols = b.cols();
313 Index size = b.rows();
314 DestDerived& dest(aDest.derived());
315 using DestScalar = typename DestDerived::Scalar;
318 // We do not directly fill dest because sparse expressions have to be free of aliasing issue.
319 // For non square least-square problems, b and dest might not have the same size whereas they might alias
320 // each-other.
321 typename DestDerived::PlainObject tmp(cols(), rhsCols);
322 ComputationInfo global_info = Success;
323 for (Index k = 0; k < rhsCols; ++k) {
324 tb = b.col(k);
325 tx = dest.col(k);
326 derived()._solve_vector_with_guess_impl(tb, tx);
327 tmp.col(k) = tx.sparseView(0);
328
329 // The call to _solve_vector_with_guess_impl updates m_info, so if it failed for a previous column
330 // we need to restore it to the worst value.
331 if (m_info == NumericalIssue)
332 global_info = NumericalIssue;
333 else if (m_info == NoConvergence)
334 global_info = NoConvergence;
335 }
336 m_info = global_info;
337 dest.swap(tmp);
338 }
339
340 template <typename Rhs, typename DestDerived>
341 std::enable_if_t<Rhs::ColsAtCompileTime != 1 && DestDerived::ColsAtCompileTime != 1> _solve_with_guess_impl(
342 const Rhs& b, MatrixBase<DestDerived>& aDest) const {
343 eigen_assert(rows() == b.rows());
344
345 Index rhsCols = b.cols();
346 DestDerived& dest(aDest.derived());
347 ComputationInfo global_info = Success;
348 for (Index k = 0; k < rhsCols; ++k) {
349 typename DestDerived::ColXpr xk(dest, k);
350 typename Rhs::ConstColXpr bk(b, k);
351 derived()._solve_vector_with_guess_impl(bk, xk);
352
353 // The call to _solve_vector_with_guess_impl updates m_info, so if it failed for a previous column
354 // we need to restore it to the worst value.
355 if (m_info == NumericalIssue)
356 global_info = NumericalIssue;
357 else if (m_info == NoConvergence)
358 global_info = NoConvergence;
359 }
360 m_info = global_info;
361 }
362
363 template <typename Rhs, typename DestDerived>
364 std::enable_if_t<Rhs::ColsAtCompileTime == 1 || DestDerived::ColsAtCompileTime == 1> _solve_with_guess_impl(
365 const Rhs& b, MatrixBase<DestDerived>& dest) const {
366 derived()._solve_vector_with_guess_impl(b, dest.derived());
367 }
368
370 template <typename Rhs, typename Dest>
371 void _solve_impl(const Rhs& b, Dest& x) const {
372 x.setZero();
373 derived()._solve_with_guess_impl(b, x);
374 }
375
376 protected:
377 void init() {
378 m_isInitialized = false;
379 m_analysisIsOk = false;
380 m_factorizationIsOk = false;
381 m_maxIterations = -1;
382 m_tolerance = NumTraits<Scalar>::epsilon();
383 }
384
385 using MatrixWrapper = internal::generic_matrix_wrapper<MatrixType>;
386 using ActualMatrixType = typename MatrixWrapper::ActualMatrixType;
387
388 const ActualMatrixType& matrix() const { return m_matrixWrapper.matrix(); }
389
390 template <typename InputType>
391 void grab(const InputType& A) {
392 m_matrixWrapper.grab(A);
393 }
394
395 MatrixWrapper m_matrixWrapper;
396 Preconditioner m_preconditioner;
397
398 Index m_maxIterations;
399 RealScalar m_tolerance;
400
401 mutable RealScalar m_error;
402 mutable Index m_iterations;
403 mutable ComputationInfo m_info;
404 mutable bool m_analysisIsOk, m_factorizationIsOk;
405};
406
407} // end namespace Eigen
408
409#endif // EIGEN_ITERATIVE_SOLVER_BASE_H
Base class for linear iterative solvers.
Definition IterativeSolverBase.h:118
IterativeSolverBase()
Definition IterativeSolverBase.h:136
ComputationInfo info() const
Definition IterativeSolverBase.h:302
RealScalar error() const
Definition IterativeSolverBase.h:266
Index maxIterations() const
Definition IterativeSolverBase.h:245
Derived & setMaxIterations(Index maxIters)
Definition IterativeSolverBase.h:250
Derived & compute(const EigenBase< MatrixDerived > &A)
Definition IterativeSolverBase.h:202
IterativeSolverBase(const EigenBase< MatrixDerived > &A)
Definition IterativeSolverBase.h:149
void solveWithGuessInPlace(const Rhs &b, Dest &x) const
Definition IterativeSolverBase.h:294
const Preconditioner & preconditioner() const
Definition IterativeSolverBase.h:239
Derived & analyzePattern(const EigenBase< MatrixDerived > &A)
Definition IterativeSolverBase.h:162
Derived & factorize(const EigenBase< MatrixDerived > &A)
Definition IterativeSolverBase.h:182
Preconditioner & preconditioner()
Definition IterativeSolverBase.h:236
Derived & setTolerance(const RealScalar &tolerance)
Definition IterativeSolverBase.h:230
RealScalar tolerance() const
Definition IterativeSolverBase.h:221
const SolveWithGuess< Derived, Rhs, Guess > solveWithGuess(const MatrixBase< Rhs > &b, const Guess &x0) const
Definition IterativeSolverBase.h:277
Index iterations() const
Definition IterativeSolverBase.h:256
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
Pseudo expression representing a solving operation.
Definition SolveWithGuess.h:44
Base class of any sparse matrices or sparse expressions.
Definition SparseMatrixBase.h:31
ComputationInfo
Definition Constants.h:455
@ NumericalIssue
Definition Constants.h:459
@ Success
Definition Constants.h:457
@ NoConvergence
Definition Constants.h:461
Definition EigenBase.h:34
constexpr Derived & derived()
Definition EigenBase.h:50