Eigen  5.0.1
 
Loading...
Searching...
No Matches
ConjugateGradient.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_CONJUGATE_GRADIENT_H
12#define EIGEN_CONJUGATE_GRADIENT_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20
30template <typename MatrixType, typename Rhs, typename Dest, typename Preconditioner>
31EIGEN_DONT_INLINE void conjugate_gradient(const MatrixType& mat, const Rhs& rhs, Dest& x, const Preconditioner& precond,
32 Index& iters, typename Dest::RealScalar& tol_error) {
33 using RealScalar = typename Dest::RealScalar;
34 using Scalar = typename Dest::Scalar;
35 // Use Dest's plain (owning) type as VectorType. For CPU Matrix/Map this
36 // resolves to Matrix<Scalar,Dynamic,1>. For GPU DeviceMatrix, PlainObject
37 // is DeviceMatrix itself (already owning).
38 using VectorType = typename Dest::PlainObject;
39
40 RealScalar tol = tol_error;
41 Index maxIters = iters;
42
43 Index n = mat.cols();
44
45 VectorType residual = rhs - mat * x; // initial residual
46
47 RealScalar rhsNorm = rhs.stableNorm();
48 if (rhsNorm == 0) {
49 x.setZero();
50 iters = 0;
51 tol_error = 0;
52 return;
53 }
54 const RealScalar considerAsZero = (std::numeric_limits<RealScalar>::min)();
55 RealScalar threshold = numext::maxi(RealScalar(tol * rhsNorm), considerAsZero);
56 RealScalar residualNorm = residual.stableNorm();
57 if (residualNorm < threshold) {
58 iters = 0;
59 tol_error = residualNorm / rhsNorm;
60 return;
61 }
62
63 // Keep the quadratic recurrence terms representable for very small or large residuals.
64 const RealScalar residualScale = internal::iterative_solver_scaling_factor(residualNorm);
65 residual /= residualScale;
66 threshold /= residualScale;
67
68 VectorType p(n);
69 p = precond.solve(residual); // initial search direction
70
71 VectorType z(n), tmp(n);
72 RealScalar absNew = numext::real(residual.dot(p)); // the square of the absolute value of r scaled by invM
73 Index i = 0;
74 while (i < maxIters) {
75 tmp.noalias() = mat * p; // the bottleneck of the algorithm
76
77 Scalar alpha = absNew / p.dot(tmp); // the amount we travel on dir
78 x += (residualScale * alpha) * p; // update solution
79 residual -= alpha * tmp; // update residual
80
81 residualNorm = residual.stableNorm();
82 if (residualNorm < threshold) break;
83
84 z = precond.solve(residual); // approximately solve for "A z = residual"
85
86 RealScalar absOld = absNew;
87 absNew = numext::real(residual.dot(z)); // update the absolute value of r
88 RealScalar beta = absNew / absOld; // calculate the Gram-Schmidt value used to create the new search direction
89 p = z + beta * p; // update search direction
90 i++;
91 }
92 tol_error = residualNorm / (rhsNorm / residualScale);
93 iters = i;
94}
95
96} // namespace internal
97
98template <typename MatrixType_, int UpLo_ = Lower,
101
102namespace internal {
103
104template <typename MatrixType_, int UpLo_, typename Preconditioner_>
105struct traits<ConjugateGradient<MatrixType_, UpLo_, Preconditioner_> > {
106 using MatrixType = MatrixType_;
107 using Preconditioner = Preconditioner_;
108};
109
110} // namespace internal
111
160template <typename MatrixType_, int UpLo_, typename Preconditioner_>
161class ConjugateGradient : public IterativeSolverBase<ConjugateGradient<MatrixType_, UpLo_, Preconditioner_> > {
162 protected:
164 using Base::m_error;
165 using Base::m_info;
166 using Base::m_isInitialized;
167 using Base::m_iterations;
168 using Base::matrix;
169
170 public:
171 using MatrixType = MatrixType_;
172 using Scalar = typename MatrixType::Scalar;
173 using RealScalar = typename MatrixType::RealScalar;
174 using Preconditioner = Preconditioner_;
175
176 enum { UpLo = UpLo_ };
177
178 public:
180 ConjugateGradient() : Base() {}
181
192 template <typename MatrixDerived>
193 explicit ConjugateGradient(const EigenBase<MatrixDerived>& A) : Base(A.derived()) {}
194
196 template <typename Rhs, typename Dest>
197 void _solve_vector_with_guess_impl(const Rhs& b, Dest& x) const {
198 using MatrixWrapper = typename Base::MatrixWrapper;
199 using ActualMatrixType = typename Base::ActualMatrixType;
200 enum {
201 TransposeInput = (!MatrixWrapper::MatrixFree) && (UpLo == (Lower | Upper)) && (!MatrixType::IsRowMajor) &&
202 (!NumTraits<Scalar>::IsComplex)
203 };
204 using RowMajorWrapper =
205 std::conditional_t<TransposeInput, Transpose<const ActualMatrixType>, ActualMatrixType const&>;
206 EIGEN_STATIC_ASSERT(internal::check_implication(MatrixWrapper::MatrixFree, UpLo == (Lower | Upper)),
207 MATRIX_FREE_CONJUGATE_GRADIENT_IS_COMPATIBLE_WITH_UPPER_UNION_LOWER_MODE_ONLY);
208 using SelfAdjointWrapper =
209 std::conditional_t<UpLo == (Lower | Upper), RowMajorWrapper,
210 typename MatrixWrapper::template ConstSelfAdjointViewReturnType<UpLo>::Type>;
211
212 m_iterations = Base::maxIterations();
213 m_error = Base::m_tolerance;
214
215 RowMajorWrapper row_mat(matrix());
216 internal::conjugate_gradient(SelfAdjointWrapper(row_mat), b, x, Base::m_preconditioner, m_iterations, m_error);
217 m_info = m_error <= Base::m_tolerance ? Success : NoConvergence;
218 }
219};
220
221} // end namespace Eigen
222
223#endif // EIGEN_CONJUGATE_GRADIENT_H
A conjugate gradient solver for sparse (or dense) self-adjoint problems.
Definition ConjugateGradient.h:161
ConjugateGradient(const EigenBase< MatrixDerived > &A)
Definition ConjugateGradient.h:193
ConjugateGradient()
Definition ConjugateGradient.h:180
A preconditioner based on the diagonal entries.
Definition BasicPreconditioners.h:40
Index maxIterations() const
Definition IterativeSolverBase.h:245
Expression of an array as a mathematical vector or matrix.
Definition ArrayWrapper.h:120
@ Lower
Definition Constants.h:212
@ Upper
Definition Constants.h:214
@ Success
Definition Constants.h:457
@ NoConvergence
Definition Constants.h:461
Definition EigenBase.h:34