10#ifndef EIGEN_TRIANGULAR_REACH_SOLVER_H
11#define EIGEN_TRIANGULAR_REACH_SOLVER_H
14#include "./InternalHeaderCheck.h"
54template <
bool Upper,
typename StorageIndex>
55Index triangular_reach(
const StorageIndex* outerIndexPtr,
const StorageIndex* innerIndexPtr,
56 const StorageIndex* innerNonZeroPtr,
const StorageIndex* bIdx, Index bCount, StorageIndex* xi,
57 StorageIndex* pstack, uint8_t* mark, Index n) {
59 for (Index r = 0; r < bCount; ++r) {
60 StorageIndex root = bIdx[r];
61 if (mark[root])
continue;
66 StorageIndex j = xi[head];
67 Index colBeg = outerIndexPtr[j];
68 Index colEnd = innerNonZeroPtr ? outerIndexPtr[j] + innerNonZeroPtr[j] : outerIndexPtr[j + 1];
71 pstack[head] = StorageIndex(colBeg);
74 for (Index p = pstack[head]; p < colEnd; ++p) {
75 StorageIndex i = innerIndexPtr[p];
76 if (
Upper ? (i >= j) : (i <= j))
continue;
77 if (mark[i])
continue;
78 pstack[head] = StorageIndex(p + 1);
110template <
bool Upper,
bool UnitDiag,
typename StorageIndex,
typename LhsScalar,
typename RhsScalar>
111void triangular_solve_over_reach(
const StorageIndex* outerIndexPtr,
const StorageIndex* innerIndexPtr,
112 const LhsScalar* valuePtr,
const StorageIndex* innerNonZeroPtr,
const StorageIndex* xi,
113 Index top, Index n, RhsScalar* x) {
114 for (Index k = top; k < n; ++k) {
115 StorageIndex j = xi[k];
116 Index colBeg = outerIndexPtr[j];
117 Index colEnd = innerNonZeroPtr ? outerIndexPtr[j] + innerNonZeroPtr[j] : outerIndexPtr[j + 1];
119 Index offBeg, offEnd;
120 EIGEN_IF_CONSTEXPR (
Upper) {
123 if (e > colBeg && innerIndexPtr[e - 1] > j)
124 e = std::upper_bound(innerIndexPtr + colBeg, innerIndexPtr + colEnd, j) - innerIndexPtr;
125 bool hasDiag = e > colBeg && innerIndexPtr[e - 1] == j;
127 offEnd = hasDiag ? e - 1 : e;
129 eigen_assert(hasDiag &&
"sparse triangular solve: missing diagonal");
132 x[j] /= hasDiag ? valuePtr[e - 1] : LhsScalar(0);
137 if (s < colEnd && innerIndexPtr[s] < j)
138 s = std::lower_bound(innerIndexPtr + colBeg, innerIndexPtr + colEnd, j) - innerIndexPtr;
139 bool hasDiag = s < colEnd && innerIndexPtr[s] == j;
140 offBeg = hasDiag ? s + 1 : s;
143 eigen_assert(hasDiag &&
"sparse triangular solve: missing diagonal");
144 x[j] /= hasDiag ? valuePtr[s] : LhsScalar(0);
148 for (Index p = offBeg; p < offEnd; ++p) {
149 StorageIndex i = innerIndexPtr[p];
150 x[i] = numext::madd<RhsScalar>(-xj, RhsScalar(valuePtr[p]), x[i]);
176template <
bool Upper,
bool UnitDiag,
typename StorageIndex,
typename LhsScalar,
typename RhsScalar>
177Index reach_solve_dense(
const StorageIndex* outerIndexPtr,
const StorageIndex* innerIndexPtr,
const LhsScalar* valuePtr,
178 const StorageIndex* innerNonZeroPtr, Index n,
const StorageIndex* bIdx, Index bCount,
179 StorageIndex* iwork, uint8_t* mark, RhsScalar* xwork) {
181 triangular_reach<Upper>(outerIndexPtr, innerIndexPtr, innerNonZeroPtr, bIdx, bCount, iwork, iwork + n, mark, n);
182 triangular_solve_over_reach<Upper, UnitDiag>(outerIndexPtr, innerIndexPtr, valuePtr, innerNonZeroPtr, iwork, top, n,
201template <
bool Upper,
typename Eval,
typename StorageIndex>
202Index triangular_reach_iter(
const Eval& mat,
const StorageIndex* bIdx, Index bCount, StorageIndex* xi, uint8_t* mark,
206 for (Index r = 0; r < bCount; ++r) {
207 StorageIndex root = bIdx[r];
214 StorageIndex j = xi[--sp];
216 for (
typename Eval::InnerIterator it(mat, j); it; ++it) {
217 StorageIndex i = StorageIndex(it.index());
218 if (
Upper ? (i >= j) : (i <= j))
continue;
226 using Comp = std::conditional_t<Upper, std::greater<StorageIndex>, std::less<StorageIndex>>;
227 std::sort(xi + top, xi + n, Comp{});
235template <
bool Upper,
bool UnitDiag,
typename Eval,
typename StorageIndex,
typename Scalar>
236void triangular_solve_over_reach_iter(
const Eval& mat,
const StorageIndex* xi, Index top, Index n, Scalar* x) {
237 for (Index k = top; k < n; ++k) {
238 StorageIndex j = xi[k];
239 EIGEN_IF_CONSTEXPR (
Upper) {
242 bool hasDiag =
false;
243 for (
typename Eval::InnerIterator dt(mat, j); dt; ++dt)
244 if (StorageIndex(dt.index()) == j) {
248 eigen_assert(hasDiag &&
"sparse triangular solve: missing diagonal");
249 EIGEN_UNUSED_VARIABLE(hasDiag);
253 for (
typename Eval::InnerIterator it(mat, j); it && StorageIndex(it.index()) < j; ++it)
254 x[it.index()] = numext::madd<Scalar>(-xj, it.value(), x[it.index()]);
256 typename Eval::InnerIterator it(mat, j);
257 while (it && StorageIndex(it.index()) < j) ++it;
258 bool hasDiag = it && StorageIndex(it.index()) == j;
260 eigen_assert(hasDiag &&
"sparse triangular solve: missing diagonal");
262 x[j] /= hasDiag ? it.value() : Scalar(0);
266 for (; it; ++it) x[it.index()] = numext::madd<Scalar>(-xj, it.value(), x[it.index()]);
275template <
bool Upper,
bool UnitDiag,
typename LhsType,
typename StorageIndex,
typename Scalar>
276Index reach_solve_dense_iter(
const LhsType& lhs, Index n,
const StorageIndex* bIdx, Index bCount, StorageIndex* iwork,
277 uint8_t* mark, Scalar* xwork) {
278 evaluator<LhsType> mat(lhs);
279 Index top = triangular_reach_iter<Upper>(mat, bIdx, bCount, iwork, mark, n);
280 triangular_solve_over_reach_iter<Upper, UnitDiag>(mat, iwork, top, n, xwork);
295template <
bool Upper,
bool UnitDiag,
typename LhsType,
typename StorageIndex,
typename Scalar>
296Index reach_solve_dense_dispatch(std::true_type ,
const LhsType& lhs, Index n,
const StorageIndex* bIdx,
297 Index bCount, StorageIndex* iwork, uint8_t* mark, Scalar* xwork) {
298 return reach_solve_dense<Upper, UnitDiag>(lhs.outerIndexPtr(), lhs.innerIndexPtr(), lhs.valuePtr(),
299 lhs.innerNonZeroPtr(), n, bIdx, bCount, iwork, mark, xwork);
301template <
bool Upper,
bool UnitDiag,
typename LhsType,
typename StorageIndex,
typename Scalar>
302Index reach_solve_dense_dispatch(std::false_type ,
const LhsType& lhs, Index n,
const StorageIndex* bIdx,
303 Index bCount, StorageIndex* iwork, uint8_t* mark, Scalar* xwork) {
304 return reach_solve_dense_iter<Upper, UnitDiag>(lhs, n, bIdx, bCount, iwork, mark, xwork);
312template <
bool Upper,
bool UnitDiag,
typename LhsDerived,
typename StorageIndex,
typename Scalar>
313Index reach_solve_dense(
const SparseMatrixBase<LhsDerived>& lhs,
const StorageIndex* bIdx, Index bCount,
314 StorageIndex* iwork, uint8_t* mark, Scalar* xwork) {
315 return reach_solve_dense_dispatch<Upper, UnitDiag>(bool_constant<has_compressed_access<LhsDerived>::value>{},
316 lhs.derived(), lhs.rows(), bIdx, bCount, iwork, mark, xwork);
@ UnitDiag
Definition Constants.h:216
@ Upper
Definition Constants.h:214