11#ifndef EIGEN_SPARSETRIANGULARSOLVER_H
12#define EIGEN_SPARSETRIANGULARSOLVER_H
15#include "./InternalHeaderCheck.h"
21template <
typename Lhs,
typename Rhs,
int Mode,
25 int StorageOrder = int(traits<Lhs>::Flags) &
RowMajorBit>
26struct sparse_solve_triangular_selector;
29template <
typename Lhs,
typename Rhs,
int Mode>
30struct sparse_solve_triangular_selector<Lhs, Rhs, Mode, Lower, RowMajor> {
31 using Scalar =
typename Rhs::Scalar;
32 using LhsEval = evaluator<Lhs>;
33 using LhsIterator =
typename evaluator<Lhs>::InnerIterator;
34 static void run(
const Lhs& lhs, Rhs& other) {
36 for (Index col = 0; col < other.cols(); ++col) {
37 for (Index i = 0; i < lhs.rows(); ++i) {
38 Scalar tmp = other.coeff(i, col);
41 for (LhsIterator it(lhsEval, i); it; ++it) {
43 lastIndex = it.index();
44 if (lastIndex == i)
break;
45 tmp = numext::madd<Scalar>(-lastVal, other.coeff(lastIndex, col), tmp);
48 other.coeffRef(i, col) = tmp;
50 eigen_assert(lastIndex == i);
51 other.coeffRef(i, col) = tmp / lastVal;
59template <
typename Lhs,
typename Rhs,
int Mode>
60struct sparse_solve_triangular_selector<Lhs, Rhs, Mode, Upper, RowMajor> {
61 using Scalar =
typename Rhs::Scalar;
62 using LhsEval = evaluator<Lhs>;
63 using LhsIterator =
typename evaluator<Lhs>::InnerIterator;
64 static void run(
const Lhs& lhs, Rhs& other) {
66 for (Index col = 0; col < other.cols(); ++col) {
67 for (Index i = lhs.rows() - 1; i >= 0; --i) {
68 Scalar tmp = other.coeff(i, col);
70 LhsIterator it(lhsEval, i);
71 while (it && it.index() < i) ++it;
72 EIGEN_IF_CONSTEXPR (!(Mode &
UnitDiag)) {
73 eigen_assert(it && it.index() == i);
76 }
else if (it && it.index() == i)
79 tmp = numext::madd<Scalar>(-it.value(), other.coeff(it.index(), col), tmp);
83 other.coeffRef(i, col) = tmp;
85 other.coeffRef(i, col) = tmp / l_ii;
92template <
typename Lhs,
typename Rhs,
int Mode>
93struct sparse_solve_triangular_selector<Lhs, Rhs, Mode, Lower, ColMajor> {
94 using Scalar =
typename Rhs::Scalar;
95 using LhsEval = evaluator<Lhs>;
96 using LhsIterator =
typename evaluator<Lhs>::InnerIterator;
97 static void run(
const Lhs& lhs, Rhs& other) {
99 for (Index col = 0; col < other.cols(); ++col) {
100 for (Index i = 0; i < lhs.cols(); ++i) {
101 Scalar& tmp = other.coeffRef(i, col);
102 if (!numext::is_exactly_zero(tmp))
104 LhsIterator it(lhsEval, i);
105 while (it && it.index() < i) ++it;
106 EIGEN_IF_CONSTEXPR (!(Mode &
UnitDiag)) {
107 eigen_assert(it && it.index() == i);
110 if (it && it.index() == i) ++it;
112 other.coeffRef(it.index(), col) = numext::madd<Scalar>(-tmp, it.value(), other.coeffRef(it.index(), col));
121template <
typename Lhs,
typename Rhs,
int Mode>
122struct sparse_solve_triangular_selector<Lhs, Rhs, Mode, Upper, ColMajor> {
123 using Scalar =
typename Rhs::Scalar;
124 using LhsEval = evaluator<Lhs>;
125 using LhsIterator =
typename evaluator<Lhs>::InnerIterator;
126 static void run(
const Lhs& lhs, Rhs& other) {
127 LhsEval lhsEval(lhs);
128 for (Index col = 0; col < other.cols(); ++col) {
129 for (Index i = lhs.cols() - 1; i >= 0; --i) {
130 Scalar& tmp = other.coeffRef(i, col);
131 if (!numext::is_exactly_zero(tmp))
133 EIGEN_IF_CONSTEXPR (!(Mode &
UnitDiag)) {
136 LhsIterator it(lhsEval, i);
137 while (it && it.index() != i) ++it;
138 eigen_assert(it && it.index() == i);
139 other.coeffRef(i, col) /= it.value();
141 LhsIterator it(lhsEval, i);
142 for (; it && it.index() < i; ++it) {
143 other.coeffRef(it.index(), col) = numext::madd<Scalar>(-tmp, it.value(), other.coeffRef(it.index(), col));
153#ifndef EIGEN_PARSED_BY_DOXYGEN
155template <
typename ExpressionType,
unsigned int Mode>
156template <
typename OtherDerived>
158 eigen_assert(derived().cols() == derived().rows() && derived().cols() == other.rows());
161 enum { copy = internal::traits<OtherDerived>::Flags &
RowMajorBit };
164 std::conditional_t<copy, typename internal::plain_matrix_type_column_major<OtherDerived>::type, OtherDerived&>;
165 OtherCopy otherCopy(other.derived());
167 internal::sparse_solve_triangular_selector<ExpressionType, std::remove_reference_t<OtherCopy>, Mode>::run(
168 derived().nestedExpression(), otherCopy);
170 if (copy) other = otherCopy;
178template <
typename Lhs,
typename Rhs,
int Mode,
183struct sparse_solve_triangular_sparse_selector;
189template <
typename Lhs,
typename Rhs>
190using rhs_matching_slice = std::integral_constant<
191 bool, has_compressed_access<Rhs>::value &&
192 std::is_same<typename traits<Rhs>::StorageIndex,
typename traits<Lhs>::StorageIndex>::value>;
199template <
bool Ordered,
bool Upper,
typename StorageIndex>
200struct reach_reorder {
201 static void run(StorageIndex* first, StorageIndex* last) { std::sort(first, last); }
203template <
bool Upper,
typename StorageIndex>
204struct reach_reorder<true,
Upper, StorageIndex> {
205 static void run(StorageIndex* , StorageIndex* ) {}
214template <
bool Ordered,
bool Upper,
typename Res,
typename StorageIndex,
typename Scalar>
215void reach_insert_column(Res& res, Index col, StorageIndex* xi, Index top, Index n, Scalar* xwork, uint8_t* mark) {
216 reach_reorder<Ordered, Upper, StorageIndex>::run(xi + top, xi + n);
219 constexpr bool Descending = Ordered &&
Upper;
220 for (Index k = top; k < n; ++k) {
221 StorageIndex j = xi[Descending ? top + n - 1 - k : k];
222 if (!numext::is_exactly_zero(xwork[j])) res.insert(j, col) = xwork[j];
223 xwork[j] = Scalar(0);
231template <
bool Upper,
bool UnitDiag,
typename Lhs,
typename Rhs,
typename Res,
typename Scalar,
232 std::enable_if_t<rhs_matching_slice<Lhs, Rhs>::value,
int> = 0>
233void reach_solve_columns(
const Lhs& lhs,
const Rhs& other, Res& res, uint8_t* mark, Scalar* xwork, Index n) {
234 using StorageIndex =
typename traits<Lhs>::StorageIndex;
235 Matrix<StorageIndex, Dynamic, 1> iwork(2 * n);
236 StorageIndex* xi = iwork.data();
237 for (Index col = 0; col < other.cols(); ++col) {
238 const StorageIndex* outer = other.outerIndexPtr();
239 const StorageIndex* nnz = other.innerNonZeroPtr();
240 Index p = outer ? outer[col] : 0;
241 Index bCount = outer ? (nnz ? Index(nnz[col]) : Index(outer[col + 1]) - p) : other.nonZeros();
242 const StorageIndex* roots = other.innerIndexPtr() + p;
243 const Scalar* vals = other.valuePtr() + p;
247 for (Index r = 0; r < bCount; ++r) xwork[roots[r]] = vals[r];
248 Index top = reach_solve_dense<Upper, UnitDiag>(lhs, roots, bCount, xi, mark, xwork);
249 reach_insert_column<!has_compressed_access<Lhs>::value,
Upper>(res, col, xi, top, n, xwork, mark);
256template <
bool Upper,
bool UnitDiag,
typename Lhs,
typename Rhs,
typename Res,
typename Scalar,
257 std::enable_if_t<!rhs_matching_slice<Lhs, Rhs>::value,
int> = 0>
258void reach_solve_columns(
const Lhs& lhs,
const Rhs& other, Res& res, uint8_t* mark, Scalar* xwork, Index n) {
259 using StorageIndex =
typename traits<Lhs>::StorageIndex;
260 Matrix<StorageIndex, Dynamic, 1> iwork(3 * n);
261 StorageIndex* xi = iwork.data();
262 StorageIndex* bIdx = iwork.data() + 2 * n;
263 for (Index col = 0; col < other.cols(); ++col) {
265 for (
typename Rhs::InnerIterator it(other, col); it; ++it) {
266 if (numext::is_exactly_zero(it.value()))
continue;
267 bIdx[bCount] = StorageIndex(it.index());
268 xwork[it.index()] = it.value();
271 Index top = reach_solve_dense<Upper, UnitDiag>(lhs, bIdx, bCount, xi, mark, xwork);
272 reach_insert_column<!has_compressed_access<Lhs>::value,
Upper>(res, col, xi, top, n, xwork, mark);
286template <
bool Upper,
typename Lhs,
typename Rhs,
int Mode>
287void run_sparse_reach_triangular_solve(
const Lhs& lhs, Rhs& other) {
288 using Scalar =
typename Rhs::Scalar;
289 Index n = lhs.rows();
290 Matrix<uint8_t, Dynamic, 1> mark = Matrix<uint8_t, Dynamic, 1>::Zero(n);
291 Matrix<Scalar, Dynamic, 1> xwork = Matrix<Scalar, Dynamic, 1>::Zero(n);
292 Rhs res(other.rows(), other.cols());
293 res.reserve(other.nonZeros());
294 reach_solve_columns<Upper, bool(Mode & UnitDiag)>(lhs, other, res, mark.data(), xwork.data(), n);
296 other = res.markAsRValue();
300template <
typename Lhs,
typename Rhs,
int Mode,
int UpLo>
301struct sparse_solve_triangular_sparse_selector<Lhs, Rhs, Mode, UpLo,
ColMajor> {
302 static void run(
const Lhs& lhs, Rhs& other) {
303 run_sparse_reach_triangular_solve<UpLo == Upper, Lhs, Rhs, Mode>(lhs, other);
309#ifndef EIGEN_PARSED_BY_DOXYGEN
310template <
typename ExpressionType,
unsigned int Mode>
311template <
typename OtherDerived>
313 eigen_assert(derived().cols() == derived().rows() && derived().cols() == other.rows());
316 internal::sparse_solve_triangular_sparse_selector<ExpressionType, OtherDerived, Mode>::run(
317 derived().nestedExpression(), other.derived());
Base class for all dense matrices, vectors, and expressions.
Definition MatrixBase.h:53
Base class of any sparse matrices or sparse expressions.
Definition SparseMatrixBase.h:31
@ UnitDiag
Definition Constants.h:216
@ ZeroDiag
Definition Constants.h:218
@ Lower
Definition Constants.h:212
@ Upper
Definition Constants.h:214
@ ColMajor
Definition Constants.h:319
constexpr unsigned int RowMajorBit
Definition Constants.h:71