[Numerics] BandMatrix and SquareMatrix support solving multiple RHS

This commit is contained in:
Ray Speth 2014-03-24 04:02:06 +00:00
parent a641992960
commit 4e72cf0334
5 changed files with 19 additions and 10 deletions

View file

@ -191,7 +191,7 @@ public:
* 0 indicates a success * 0 indicates a success
* ~0 Some error occurred, see the LAPACK documentation * ~0 Some error occurred, see the LAPACK documentation
*/ */
int solve(doublereal* b); int solve(doublereal* b, size_t nrhs=1, size_t ldb=0);
//! Returns an iterator for the start of the band storage data //! Returns an iterator for the start of the band storage data
/*! /*!

View file

@ -127,9 +127,12 @@ public:
//! Solves the Ax = b system returning x in the b spot. //! Solves the Ax = b system returning x in the b spot.
/*! /*!
* @param b Vector for the rhs of the equation system * @param b Vector for the rhs of the equation system
* @param nrhs Number of right-hand sides to solve, default 1
* @param ldb Leading dimension of the right-hand side array.
* Defaults to nRows()
*/ */
virtual int solve(doublereal* b) = 0; virtual int solve(doublereal* b, size_t nrhs=1, size_t ldb=0) = 0;
//! true if the current factorization is up to date with the matrix //! true if the current factorization is up to date with the matrix
virtual bool factored() const = 0; virtual bool factored() const = 0;

View file

@ -44,7 +44,7 @@ public:
//! Assignment operator //! Assignment operator
SquareMatrix& operator=(const SquareMatrix& right); SquareMatrix& operator=(const SquareMatrix& right);
int solve(doublereal* b); int solve(doublereal* b, size_t nrhs=1, size_t ldb=0);
void resize(size_t n, size_t m, doublereal v = 0.0); void resize(size_t n, size_t m, doublereal v = 0.0);

View file

@ -270,16 +270,19 @@ int BandMatrix::solve(const doublereal* const b, doublereal* const x)
return solve(x); return solve(x);
} }
int BandMatrix::solve(doublereal* b) int BandMatrix::solve(doublereal* b, size_t nrhs, size_t ldb)
{ {
int info = 0; int info = 0;
if (!m_factored) { if (!m_factored) {
info = factor(); info = factor();
} }
if (ldb == 0) {
ldb = nColumns();
}
if (info == 0) if (info == 0)
ct_dgbtrs(ctlapack::NoTranspose, nColumns(), nSubDiagonals(), ct_dgbtrs(ctlapack::NoTranspose, nColumns(), nSubDiagonals(),
nSuperDiagonals(), 1, DATA_PTR(ludata), ldim(), nSuperDiagonals(), nrhs, DATA_PTR(ludata), ldim(),
DATA_PTR(ipiv()), b, nColumns(), info); DATA_PTR(ipiv()), b, ldb, info);
// error handling // error handling
if (info != 0) { if (info != 0) {

View file

@ -60,7 +60,7 @@ SquareMatrix& SquareMatrix::operator=(const SquareMatrix& y)
return *this; return *this;
} }
int SquareMatrix::solve(doublereal* b) int SquareMatrix::solve(doublereal* b, size_t nrhs, size_t ldb)
{ {
if (useQR_) { if (useQR_) {
return solveQR(b); return solveQR(b);
@ -75,12 +75,15 @@ int SquareMatrix::solve(doublereal* b)
return retn; return retn;
} }
} }
if (ldb == 0) {
ldb = nColumns();
}
/* /*
* Solve the factored system * Solve the factored system
*/ */
ct_dgetrs(ctlapack::NoTranspose, static_cast<int>(nRows()), ct_dgetrs(ctlapack::NoTranspose, static_cast<int>(nRows()),
1, &(*(begin())), static_cast<int>(nRows()), nrhs, &(*(begin())), static_cast<int>(nRows()),
DATA_PTR(ipiv()), b, static_cast<int>(nColumns()), info); DATA_PTR(ipiv()), b, ldb, info);
if (info != 0) { if (info != 0) {
if (m_printLevel) { if (m_printLevel) {
writelogf("SquareMatrix::solve(): DGETRS returned INFO = %d\n", info); writelogf("SquareMatrix::solve(): DGETRS returned INFO = %d\n", info);