Eigen  5.0.1
 
Loading...
Searching...
No Matches
PaStiXSupport.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2012 Désiré Nuentsa-Wakam <desire.nuentsa_wakam@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_PASTIXSUPPORT_H
12#define EIGEN_PASTIXSUPPORT_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19#if defined(DCOMPLEX)
20#define PASTIX_COMPLEX COMPLEX
21#define PASTIX_DCOMPLEX DCOMPLEX
22#else
23#define PASTIX_COMPLEX std::complex<float>
24#define PASTIX_DCOMPLEX std::complex<double>
25#endif
26
35template <typename MatrixType_, bool IsStrSym = false>
36class PastixLU;
37template <typename MatrixType_, int Options>
38class PastixLLT;
39template <typename MatrixType_, int Options>
40class PastixLDLT;
41
42namespace internal {
43
44template <class Pastix>
45struct pastix_traits;
46
47template <typename MatrixType_>
48struct pastix_traits<PastixLU<MatrixType_> > {
49 typedef MatrixType_ MatrixType;
50 typedef typename MatrixType_::Scalar Scalar;
51 typedef typename MatrixType_::RealScalar RealScalar;
52 typedef typename MatrixType_::StorageIndex StorageIndex;
53};
54
55template <typename MatrixType_, int Options>
56struct pastix_traits<PastixLLT<MatrixType_, Options> > {
57 typedef MatrixType_ MatrixType;
58 typedef typename MatrixType_::Scalar Scalar;
59 typedef typename MatrixType_::RealScalar RealScalar;
60 typedef typename MatrixType_::StorageIndex StorageIndex;
61};
62
63template <typename MatrixType_, int Options>
64struct pastix_traits<PastixLDLT<MatrixType_, Options> > {
65 typedef MatrixType_ MatrixType;
66 typedef typename MatrixType_::Scalar Scalar;
67 typedef typename MatrixType_::RealScalar RealScalar;
68 typedef typename MatrixType_::StorageIndex StorageIndex;
69};
70
71inline void eigen_pastix(pastix_data_t **pastix_data, int pastix_comm, int n, int *ptr, int *idx, float *vals,
72 int *perm, int *invp, float *x, int nbrhs, int *iparm, double *dparm) {
73 if (n == 0) {
74 ptr = nullptr;
75 idx = nullptr;
76 vals = nullptr;
77 }
78 if (nbrhs == 0) {
79 x = nullptr;
80 nbrhs = 1;
81 }
82 s_pastix(pastix_data, pastix_comm, n, ptr, idx, vals, perm, invp, x, nbrhs, iparm, dparm);
83}
84
85inline void eigen_pastix(pastix_data_t **pastix_data, int pastix_comm, int n, int *ptr, int *idx, double *vals,
86 int *perm, int *invp, double *x, int nbrhs, int *iparm, double *dparm) {
87 if (n == 0) {
88 ptr = nullptr;
89 idx = nullptr;
90 vals = nullptr;
91 }
92 if (nbrhs == 0) {
93 x = nullptr;
94 nbrhs = 1;
95 }
96 d_pastix(pastix_data, pastix_comm, n, ptr, idx, vals, perm, invp, x, nbrhs, iparm, dparm);
97}
98
99inline void eigen_pastix(pastix_data_t **pastix_data, int pastix_comm, int n, int *ptr, int *idx,
100 std::complex<float> *vals, int *perm, int *invp, std::complex<float> *x, int nbrhs, int *iparm,
101 double *dparm) {
102 if (n == 0) {
103 ptr = nullptr;
104 idx = nullptr;
105 vals = nullptr;
106 }
107 if (nbrhs == 0) {
108 x = nullptr;
109 nbrhs = 1;
110 }
111 c_pastix(pastix_data, pastix_comm, n, ptr, idx, reinterpret_cast<PASTIX_COMPLEX *>(vals), perm, invp,
112 reinterpret_cast<PASTIX_COMPLEX *>(x), nbrhs, iparm, dparm);
113}
114
115inline void eigen_pastix(pastix_data_t **pastix_data, int pastix_comm, int n, int *ptr, int *idx,
116 std::complex<double> *vals, int *perm, int *invp, std::complex<double> *x, int nbrhs,
117 int *iparm, double *dparm) {
118 if (n == 0) {
119 ptr = nullptr;
120 idx = nullptr;
121 vals = nullptr;
122 }
123 if (nbrhs == 0) {
124 x = nullptr;
125 nbrhs = 1;
126 }
127 z_pastix(pastix_data, pastix_comm, n, ptr, idx, reinterpret_cast<PASTIX_DCOMPLEX *>(vals), perm, invp,
128 reinterpret_cast<PASTIX_DCOMPLEX *>(x), nbrhs, iparm, dparm);
129}
130
131// Convert the matrix to Fortran-style Numbering
132template <typename MatrixType>
133void c_to_fortran_numbering(MatrixType &mat) {
134 if (!(mat.outerIndexPtr()[0])) {
135 int i;
136 for (i = 0; i <= mat.rows(); ++i) ++mat.outerIndexPtr()[i];
137 for (i = 0; i < mat.nonZeros(); ++i) ++mat.innerIndexPtr()[i];
138 }
139}
140} // namespace internal
141
142// This is the base class to interface with PaStiX functions.
143// Users should not use this class directly.
144template <class Derived>
145class PastixBase : public SparseSolverBase<Derived> {
146 protected:
147 typedef SparseSolverBase<Derived> Base;
148 using Base::derived;
149 using Base::m_isInitialized;
150
151 public:
152 using Base::_solve_impl;
153
154 typedef typename internal::pastix_traits<Derived>::MatrixType MatrixType_;
155 typedef MatrixType_ MatrixType;
156 typedef typename MatrixType::Scalar Scalar;
157 typedef typename MatrixType::RealScalar RealScalar;
158 typedef typename MatrixType::StorageIndex StorageIndex;
159 typedef Matrix<Scalar, Dynamic, 1> Vector;
160 typedef SparseMatrix<Scalar, ColMajor> ColSpMatrix;
161 enum { ColsAtCompileTime = MatrixType::ColsAtCompileTime, MaxColsAtCompileTime = MatrixType::MaxColsAtCompileTime };
162
163 public:
164 PastixBase() : m_initisOk(false), m_analysisIsOk(false), m_factorizationIsOk(false), m_pastixdata(0), m_size(0) {
165 init();
166 }
167
168 ~PastixBase() { clean(); }
169
170 template <typename Rhs, typename Dest>
171 bool _solve_impl(const MatrixBase<Rhs> &b, MatrixBase<Dest> &x) const;
172
178 Array<StorageIndex, IPARM_SIZE, 1> &iparm() { return m_iparm; }
179
183
184 int &iparm(int idxparam) { return m_iparm(idxparam); }
185
190 Array<double, DPARM_SIZE, 1> &dparm() { return m_dparm; }
191
195 double &dparm(int idxparam) { return m_dparm(idxparam); }
196
197 inline Index cols() const { return m_size; }
198 inline Index rows() const { return m_size; }
199
208 ComputationInfo info() const {
209 eigen_assert(m_isInitialized && "Decomposition is not initialized.");
210 return m_info;
211 }
212
213 protected:
214 // Initialize the Pastix data structure, check the matrix
215 void init();
216
217 // Compute the ordering and the symbolic factorization
218 void analyzePattern(ColSpMatrix &mat);
219
220 // Compute the numerical factorization
221 void factorize(ColSpMatrix &mat);
222
223 // Free all the data allocated by Pastix
224 void clean() {
225 eigen_assert(m_initisOk && "The Pastix structure should be allocated first");
226 m_iparm(IPARM_START_TASK) = API_TASK_CLEAN;
227 m_iparm(IPARM_END_TASK) = API_TASK_CLEAN;
228 internal::eigen_pastix(&m_pastixdata, MPI_COMM_WORLD, 0, 0, 0, (Scalar *)0, m_perm.data(), m_invp.data(), 0, 0,
229 m_iparm.data(), m_dparm.data());
230 }
231
232 void compute(ColSpMatrix &mat);
233
234 int m_initisOk;
235 int m_analysisIsOk;
236 int m_factorizationIsOk;
237 mutable ComputationInfo m_info;
238 mutable pastix_data_t *m_pastixdata; // Data structure for pastix
239 mutable int m_comm; // The MPI communicator identifier
240 mutable Array<int, IPARM_SIZE, 1> m_iparm; // integer vector for the input parameters
241 mutable Array<double, DPARM_SIZE, 1> m_dparm; // Scalar vector for the input parameters
242 mutable Matrix<StorageIndex, Dynamic, 1> m_perm; // Permutation vector
243 mutable Matrix<StorageIndex, Dynamic, 1> m_invp; // Inverse permutation vector
244 mutable int m_size; // Size of the matrix
245};
246
251template <class Derived>
252void PastixBase<Derived>::init() {
253 m_size = 0;
254 m_iparm.setZero(IPARM_SIZE);
255 m_dparm.setZero(DPARM_SIZE);
256
257 m_iparm(IPARM_MODIFY_PARAMETER) = API_NO;
258 pastix(&m_pastixdata, MPI_COMM_WORLD, 0, 0, 0, 0, 0, 0, 0, 1, m_iparm.data(), m_dparm.data());
259
260 m_iparm[IPARM_MATRIX_VERIFICATION] = API_NO;
261 m_iparm[IPARM_VERBOSE] = API_VERBOSE_NOT;
262 m_iparm[IPARM_ORDERING] = API_ORDER_SCOTCH;
263 m_iparm[IPARM_INCOMPLETE] = API_NO;
264 m_iparm[IPARM_OOC_LIMIT] = 2000;
265 m_iparm[IPARM_RHS_MAKING] = API_RHS_B;
266 m_iparm(IPARM_MATRIX_VERIFICATION) = API_NO;
267
268 m_iparm(IPARM_START_TASK) = API_TASK_INIT;
269 m_iparm(IPARM_END_TASK) = API_TASK_INIT;
270 internal::eigen_pastix(&m_pastixdata, MPI_COMM_WORLD, 0, 0, 0, (Scalar *)0, 0, 0, 0, 0, m_iparm.data(),
271 m_dparm.data());
272
273 // Check the returned error
274 if (m_iparm(IPARM_ERROR_NUMBER)) {
275 m_info = InvalidInput;
276 m_initisOk = false;
277 } else {
278 m_info = Success;
279 m_initisOk = true;
280 }
281}
282
283template <class Derived>
284void PastixBase<Derived>::compute(ColSpMatrix &mat) {
285 eigen_assert(mat.rows() == mat.cols() && "The input matrix should be squared");
286
287 analyzePattern(mat);
288 factorize(mat);
289
290 m_iparm(IPARM_MATRIX_VERIFICATION) = API_NO;
291}
292
293template <class Derived>
294void PastixBase<Derived>::analyzePattern(ColSpMatrix &mat) {
295 eigen_assert(m_initisOk && "The initialization of PaSTiX failed");
296
297 // clean previous calls
298 if (m_size > 0) clean();
299
300 m_size = internal::convert_index<int>(mat.rows());
301 m_perm.resize(m_size);
302 m_invp.resize(m_size);
303
304 m_iparm(IPARM_START_TASK) = API_TASK_ORDERING;
305 m_iparm(IPARM_END_TASK) = API_TASK_ANALYSE;
306 internal::eigen_pastix(&m_pastixdata, MPI_COMM_WORLD, m_size, mat.outerIndexPtr(), mat.innerIndexPtr(),
307 mat.valuePtr(), m_perm.data(), m_invp.data(), 0, 0, m_iparm.data(), m_dparm.data());
308
309 // Check the returned error
310 if (m_iparm(IPARM_ERROR_NUMBER)) {
311 m_info = NumericalIssue;
312 m_analysisIsOk = false;
313 } else {
314 m_info = Success;
315 m_analysisIsOk = true;
316 }
317}
318
319template <class Derived>
320void PastixBase<Derived>::factorize(ColSpMatrix &mat) {
321 eigen_assert(m_analysisIsOk && "The analysis phase should be called before the factorization phase");
322 m_iparm(IPARM_START_TASK) = API_TASK_NUMFACT;
323 m_iparm(IPARM_END_TASK) = API_TASK_NUMFACT;
324 m_size = internal::convert_index<int>(mat.rows());
325
326 internal::eigen_pastix(&m_pastixdata, MPI_COMM_WORLD, m_size, mat.outerIndexPtr(), mat.innerIndexPtr(),
327 mat.valuePtr(), m_perm.data(), m_invp.data(), 0, 0, m_iparm.data(), m_dparm.data());
328
329 // Check the returned error
330 if (m_iparm(IPARM_ERROR_NUMBER)) {
331 m_info = NumericalIssue;
332 m_factorizationIsOk = false;
333 m_isInitialized = false;
334 } else {
335 m_info = Success;
336 m_factorizationIsOk = true;
337 m_isInitialized = true;
338 }
339}
340
341/* Solve the system */
342template <typename Base>
343template <typename Rhs, typename Dest>
344bool PastixBase<Base>::_solve_impl(const MatrixBase<Rhs> &b, MatrixBase<Dest> &x) const {
345 eigen_assert(m_isInitialized && "The matrix should be factorized first");
346 EIGEN_STATIC_ASSERT((Dest::Flags & RowMajorBit) == 0, THIS_METHOD_IS_ONLY_FOR_COLUMN_MAJOR_MATRICES);
347 int rhs = 1;
348
349 x = b; /* on return, x is overwritten by the computed solution */
350
351 for (int i = 0; i < b.cols(); i++) {
352 m_iparm[IPARM_START_TASK] = API_TASK_SOLVE;
353 m_iparm[IPARM_END_TASK] = API_TASK_REFINE;
354
355 internal::eigen_pastix(&m_pastixdata, MPI_COMM_WORLD, internal::convert_index<int>(x.rows()), 0, 0, 0,
356 m_perm.data(), m_invp.data(), &x(0, i), rhs, m_iparm.data(), m_dparm.data());
357 }
358
359 // Check the returned error
360 m_info = m_iparm(IPARM_ERROR_NUMBER) == 0 ? Success : NumericalIssue;
361
362 return m_iparm(IPARM_ERROR_NUMBER) == 0;
363}
364
386template <typename MatrixType_, bool IsStrSym>
387class PastixLU : public PastixBase<PastixLU<MatrixType_> > {
388 public:
389 typedef MatrixType_ MatrixType;
390 typedef PastixBase<PastixLU<MatrixType> > Base;
391 typedef typename Base::ColSpMatrix ColSpMatrix;
392 typedef typename MatrixType::StorageIndex StorageIndex;
393
394 public:
395 PastixLU() : Base() { init(); }
396
397 explicit PastixLU(const MatrixType &matrix) : Base() {
398 init();
399 compute(matrix);
400 }
406 void compute(const MatrixType &matrix) {
407 m_structureIsUptodate = false;
408 ColSpMatrix temp;
409 grabMatrix(matrix, temp);
410 Base::compute(temp);
411 }
412
417 void analyzePattern(const MatrixType &matrix) {
418 m_structureIsUptodate = false;
419 ColSpMatrix temp;
420 grabMatrix(matrix, temp);
421 Base::analyzePattern(temp);
422 }
423
429 void factorize(const MatrixType &matrix) {
430 ColSpMatrix temp;
431 grabMatrix(matrix, temp);
432 Base::factorize(temp);
433 }
434
435 protected:
436 void init() {
437 m_structureIsUptodate = false;
438 m_iparm(IPARM_SYM) = API_SYM_NO;
439 m_iparm(IPARM_FACTORIZATION) = API_FACT_LU;
440 }
441
442 void grabMatrix(const MatrixType &matrix, ColSpMatrix &out) {
443 EIGEN_IF_CONSTEXPR (IsStrSym)
444 out = matrix;
445 else {
446 if (!m_structureIsUptodate) {
447 // update the transposed structure
448 m_transposedStructure = matrix.transpose();
449
450 // Set the elements of the matrix to zero
451 for (Index j = 0; j < m_transposedStructure.outerSize(); ++j)
452 for (typename ColSpMatrix::InnerIterator it(m_transposedStructure, j); it; ++it) it.valueRef() = 0.0;
453
454 m_structureIsUptodate = true;
455 }
456
457 out = m_transposedStructure + matrix;
458 }
459 internal::c_to_fortran_numbering(out);
460 }
461
462 using Base::m_dparm;
463 using Base::m_iparm;
464
465 ColSpMatrix m_transposedStructure;
466 bool m_structureIsUptodate;
467};
468
486template <typename MatrixType_, int UpLo_>
487class PastixLLT : public PastixBase<PastixLLT<MatrixType_, UpLo_> > {
488 public:
489 typedef MatrixType_ MatrixType;
490 typedef PastixBase<PastixLLT<MatrixType, UpLo_> > Base;
491 typedef typename Base::ColSpMatrix ColSpMatrix;
492
493 public:
494 enum { UpLo = UpLo_ };
495 PastixLLT() : Base() { init(); }
496
497 explicit PastixLLT(const MatrixType &matrix) : Base() {
498 init();
499 compute(matrix);
500 }
501
505 void compute(const MatrixType &matrix) {
506 ColSpMatrix temp;
507 grabMatrix(matrix, temp);
508 Base::compute(temp);
509 }
510
515 void analyzePattern(const MatrixType &matrix) {
516 ColSpMatrix temp;
517 grabMatrix(matrix, temp);
518 Base::analyzePattern(temp);
519 }
520
523 void factorize(const MatrixType &matrix) {
524 ColSpMatrix temp;
525 grabMatrix(matrix, temp);
526 Base::factorize(temp);
527 }
528
529 protected:
530 using Base::m_iparm;
531
532 void init() {
533 m_iparm(IPARM_SYM) = API_SYM_YES;
534 m_iparm(IPARM_FACTORIZATION) = API_FACT_LLT;
535 }
536
537 void grabMatrix(const MatrixType &matrix, ColSpMatrix &out) {
538 out.resize(matrix.rows(), matrix.cols());
539 // Pastix supports only lower, column-major matrices
540 out.template selfadjointView<Lower>() = matrix.template selfadjointView<UpLo>();
541 internal::c_to_fortran_numbering(out);
542 }
543};
544
562template <typename MatrixType_, int UpLo_>
563class PastixLDLT : public PastixBase<PastixLDLT<MatrixType_, UpLo_> > {
564 public:
565 typedef MatrixType_ MatrixType;
566 typedef PastixBase<PastixLDLT<MatrixType, UpLo_> > Base;
567 typedef typename Base::ColSpMatrix ColSpMatrix;
568
569 public:
570 enum { UpLo = UpLo_ };
571 PastixLDLT() : Base() { init(); }
572
573 explicit PastixLDLT(const MatrixType &matrix) : Base() {
574 init();
575 compute(matrix);
576 }
577
581 void compute(const MatrixType &matrix) {
582 ColSpMatrix temp;
583 grabMatrix(matrix, temp);
584 Base::compute(temp);
585 }
586
591 void analyzePattern(const MatrixType &matrix) {
592 ColSpMatrix temp;
593 grabMatrix(matrix, temp);
594 Base::analyzePattern(temp);
595 }
596
599 void factorize(const MatrixType &matrix) {
600 ColSpMatrix temp;
601 grabMatrix(matrix, temp);
602 Base::factorize(temp);
603 }
604
605 protected:
606 using Base::m_iparm;
607
608 void init() {
609 m_iparm(IPARM_SYM) = API_SYM_YES;
610 m_iparm(IPARM_FACTORIZATION) = API_FACT_LDLT;
611 }
612
613 void grabMatrix(const MatrixType &matrix, ColSpMatrix &out) {
614 // Pastix supports only lower, column-major matrices
615 out.resize(matrix.rows(), matrix.cols());
616 out.template selfadjointView<Lower>() = matrix.template selfadjointView<UpLo>();
617 internal::c_to_fortran_numbering(out);
618 }
619};
620
621} // end namespace Eigen
622
623#undef PASTIX_COMPLEX
624#undef PASTIX_DCOMPLEX
625
626#endif
Base class for all dense matrices, vectors, and expressions.
Definition MatrixBase.h:53
A sparse direct supernodal Cholesky (LLT) factorization and solver based on the PaStiX library.
Definition PaStiXSupport.h:563
void compute(const MatrixType &matrix)
Definition PaStiXSupport.h:581
void analyzePattern(const MatrixType &matrix)
Definition PaStiXSupport.h:591
void factorize(const MatrixType &matrix)
Definition PaStiXSupport.h:599
A sparse direct supernodal Cholesky (LLT) factorization and solver based on the PaStiX library.
Definition PaStiXSupport.h:487
void compute(const MatrixType &matrix)
Definition PaStiXSupport.h:505
void factorize(const MatrixType &matrix)
Definition PaStiXSupport.h:523
void analyzePattern(const MatrixType &matrix)
Definition PaStiXSupport.h:515
Interface to the PaStix solver.
Definition PaStiXSupport.h:387
void analyzePattern(const MatrixType &matrix)
Definition PaStiXSupport.h:417
void factorize(const MatrixType &matrix)
Definition PaStiXSupport.h:429
void compute(const MatrixType &matrix)
Definition PaStiXSupport.h:406
Index outerSize() const
Definition SparseMatrix.h:167
A base class for sparse solvers.
Definition SparseSolverBase.h:68
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