Eigen  5.0.1
 
Loading...
Searching...
No Matches
KLUSupport.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2017 Kyle Macfarlan <kyle.macfarlan@gmail.com>
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_KLUSUPPORT_H
12#define EIGEN_KLUSUPPORT_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19/* TODO extract L, extract U, compute det, etc... */
20
36
37inline int klu_solve(klu_symbolic *Symbolic, klu_numeric *Numeric, Index ldim, Index nrhs, double B[],
38 klu_common *Common, double) {
39 return klu_solve(Symbolic, Numeric, internal::convert_index<int>(ldim), internal::convert_index<int>(nrhs), B,
40 Common);
41}
42
43inline int klu_solve(klu_symbolic *Symbolic, klu_numeric *Numeric, Index ldim, Index nrhs, std::complex<double> B[],
44 klu_common *Common, std::complex<double>) {
45 return klu_z_solve(Symbolic, Numeric, internal::convert_index<int>(ldim), internal::convert_index<int>(nrhs),
46 &numext::real_ref(B[0]), Common);
47}
48
49inline int klu_tsolve(klu_symbolic *Symbolic, klu_numeric *Numeric, Index ldim, Index nrhs, double B[],
50 klu_common *Common, double) {
51 return klu_tsolve(Symbolic, Numeric, internal::convert_index<int>(ldim), internal::convert_index<int>(nrhs), B,
52 Common);
53}
54
55inline int klu_tsolve(klu_symbolic *Symbolic, klu_numeric *Numeric, Index ldim, Index nrhs, std::complex<double> B[],
56 klu_common *Common, std::complex<double>) {
57 return klu_z_tsolve(Symbolic, Numeric, internal::convert_index<int>(ldim), internal::convert_index<int>(nrhs),
58 &numext::real_ref(B[0]), 0, Common);
59}
60
61inline klu_numeric *klu_factor(int Ap[], int Ai[], double Ax[], klu_symbolic *Symbolic, klu_common *Common, double) {
62 return klu_factor(Ap, Ai, Ax, Symbolic, Common);
63}
64
65inline klu_numeric *klu_factor(int Ap[], int Ai[], std::complex<double> Ax[], klu_symbolic *Symbolic,
66 klu_common *Common, std::complex<double>) {
67 return klu_z_factor(Ap, Ai, &numext::real_ref(Ax[0]), Symbolic, Common);
68}
69
70template <typename MatrixType_>
71class KLU : public SparseSolverBase<KLU<MatrixType_> > {
72 protected:
74 using Base::m_isInitialized;
75
76 public:
77 using Base::_solve_impl;
78 typedef MatrixType_ MatrixType;
79 typedef typename MatrixType::Scalar Scalar;
80 typedef typename MatrixType::RealScalar RealScalar;
81 typedef typename MatrixType::StorageIndex StorageIndex;
82 typedef Matrix<Scalar, Dynamic, 1> Vector;
83 typedef Matrix<int, 1, MatrixType::ColsAtCompileTime> IntRowVectorType;
84 typedef Matrix<int, MatrixType::RowsAtCompileTime, 1> IntColVectorType;
85 typedef SparseMatrix<Scalar> LUMatrixType;
86 typedef SparseMatrix<Scalar, ColMajor, int> KLUMatrixType;
87 typedef Ref<const KLUMatrixType, StandardCompressedFormat> KLUMatrixRef;
88 enum { ColsAtCompileTime = MatrixType::ColsAtCompileTime, MaxColsAtCompileTime = MatrixType::MaxColsAtCompileTime };
89
90 public:
91 KLU() : m_dummy(0, 0), mp_matrix(m_dummy) { init(); }
92
93 template <typename InputMatrixType>
94 explicit KLU(const InputMatrixType &matrix) : mp_matrix(matrix) {
95 init();
96 compute(matrix);
97 }
98
99 ~KLU() {
100 if (m_symbolic) klu_free_symbolic(&m_symbolic, &m_common);
101 if (m_numeric) klu_free_numeric(&m_numeric, &m_common);
102 }
103
104 constexpr Index rows() const noexcept { return mp_matrix.rows(); }
105 constexpr Index cols() const noexcept { return mp_matrix.cols(); }
106
112 ComputationInfo info() const {
113 eigen_assert(m_isInitialized && "Decomposition is not initialized.");
114 return m_info;
115 }
120 template <typename InputMatrixType>
121 void compute(const InputMatrixType &matrix) {
122 if (m_symbolic) klu_free_symbolic(&m_symbolic, &m_common);
123 if (m_numeric) klu_free_numeric(&m_numeric, &m_common);
124 grab(matrix.derived());
125 analyzePattern_impl();
126 factorize_impl();
127 }
128
135 template <typename InputMatrixType>
136 void analyzePattern(const InputMatrixType &matrix) {
137 if (m_symbolic) klu_free_symbolic(&m_symbolic, &m_common);
138 if (m_numeric) klu_free_numeric(&m_numeric, &m_common);
139
140 grab(matrix.derived());
141
142 analyzePattern_impl();
143 }
144
149 inline const klu_common &kluCommon() const { return m_common; }
150
157 inline klu_common &kluCommon() { return m_common; }
158
165 template <typename InputMatrixType>
166 void factorize(const InputMatrixType &matrix) {
167 eigen_assert(m_analysisIsOk && "KLU: you must first call analyzePattern()");
168 if (m_numeric) klu_free_numeric(&m_numeric, &m_common);
169
170 grab(matrix.derived());
171
172 factorize_impl();
173 }
174
176 template <typename BDerived, typename XDerived>
177 bool _solve_impl(const MatrixBase<BDerived> &b, MatrixBase<XDerived> &x) const;
178
179 protected:
180 void init() {
181 m_info = InvalidInput;
182 m_isInitialized = false;
183 m_numeric = 0;
184 m_symbolic = 0;
185 m_extractedDataAreDirty = true;
186
187 klu_defaults(&m_common);
188 }
189
190 void analyzePattern_impl() {
191 m_info = InvalidInput;
192 m_analysisIsOk = false;
193 m_factorizationIsOk = false;
194 m_symbolic = klu_analyze(internal::convert_index<int>(mp_matrix.rows()),
195 const_cast<StorageIndex *>(mp_matrix.outerIndexPtr()),
196 const_cast<StorageIndex *>(mp_matrix.innerIndexPtr()), &m_common);
197 if (m_symbolic) {
198 m_isInitialized = true;
199 m_info = Success;
200 m_analysisIsOk = true;
201 m_extractedDataAreDirty = true;
202 }
203 }
204
205 void factorize_impl() {
206 m_numeric = klu_factor(const_cast<StorageIndex *>(mp_matrix.outerIndexPtr()),
207 const_cast<StorageIndex *>(mp_matrix.innerIndexPtr()),
208 const_cast<Scalar *>(mp_matrix.valuePtr()), m_symbolic, &m_common, Scalar());
209
210 m_info = m_numeric ? Success : NumericalIssue;
211 m_factorizationIsOk = m_numeric ? 1 : 0;
212 m_extractedDataAreDirty = true;
213 }
214
215 template <typename MatrixDerived>
216 void grab(const EigenBase<MatrixDerived> &A) {
217 internal::destroy_at(&mp_matrix);
218 internal::construct_at(&mp_matrix, A.derived());
219 }
220
221 void grab(const KLUMatrixRef &A) {
222 if (&(A.derived()) != &mp_matrix) {
223 internal::destroy_at(&mp_matrix);
224 internal::construct_at(&mp_matrix, A);
225 }
226 }
227
228 KLUMatrixType m_dummy;
229 KLUMatrixRef mp_matrix;
230
231 klu_numeric *m_numeric;
232 klu_symbolic *m_symbolic;
233 klu_common m_common;
234 mutable ComputationInfo m_info;
235 int m_factorizationIsOk;
236 int m_analysisIsOk;
237 mutable bool m_extractedDataAreDirty;
238
239 private:
240 KLU(const KLU &) {}
241};
242
243template <typename MatrixType>
244template <typename BDerived, typename XDerived>
245bool KLU<MatrixType>::_solve_impl(const MatrixBase<BDerived> &b, MatrixBase<XDerived> &x) const {
246 Index rhsCols = b.cols();
247 EIGEN_STATIC_ASSERT((XDerived::Flags & RowMajorBit) == 0, THIS_METHOD_IS_ONLY_FOR_COLUMN_MAJOR_MATRICES);
248 eigen_assert(m_factorizationIsOk &&
249 "The decomposition is not in a valid state for solving, you must first call either compute() or "
250 "analyzePattern()/factorize()");
251
252 x = b;
253 int info = klu_solve(m_symbolic, m_numeric, b.rows(), rhsCols, x.const_cast_derived().data(),
254 const_cast<klu_common *>(&m_common), Scalar());
255
256 m_info = info != 0 ? Success : NumericalIssue;
257 return true;
258}
259
260} // end namespace Eigen
261
262#endif // EIGEN_KLUSUPPORT_H
Base class for all dense matrices, vectors, and expressions.
Definition MatrixBase.h:53
A base class for sparse solvers.
Definition SparseSolverBase.h:68
int klu_solve(klu_symbolic *Symbolic, klu_numeric *Numeric, Index ldim, Index nrhs, double B[], klu_common *Common, double)
A sparse LU factorization and solver based on KLU.
Definition KLUSupport.h:37
ComputationInfo
Definition Constants.h:455
@ NumericalIssue
Definition Constants.h:459
@ InvalidInput
Definition Constants.h:464
@ Success
Definition Constants.h:457
constexpr unsigned int RowMajorBit
Definition Constants.h:71