Eigen  5.0.1
 
Loading...
Searching...
No Matches
TriangularSolver.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2008 Gael Guennebaud <gael.guennebaud@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_SPARSETRIANGULARSOLVER_H
12#define EIGEN_SPARSETRIANGULARSOLVER_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20
21template <typename Lhs, typename Rhs, int Mode,
22 int UpLo = (Mode & Lower) ? Lower
23 : (Mode & Upper) ? Upper
24 : -1,
25 int StorageOrder = int(traits<Lhs>::Flags) & RowMajorBit>
26struct sparse_solve_triangular_selector;
27
28// forward substitution, row-major
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) {
35 LhsEval lhsEval(lhs);
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);
39 Scalar lastVal(0);
40 Index lastIndex = 0;
41 for (LhsIterator it(lhsEval, i); it; ++it) {
42 lastVal = it.value();
43 lastIndex = it.index();
44 if (lastIndex == i) break;
45 tmp = numext::madd<Scalar>(-lastVal, other.coeff(lastIndex, col), tmp);
46 }
47 EIGEN_IF_CONSTEXPR (Mode & UnitDiag)
48 other.coeffRef(i, col) = tmp;
49 else {
50 eigen_assert(lastIndex == i);
51 other.coeffRef(i, col) = tmp / lastVal;
52 }
53 }
54 }
55 }
56};
57
58// backward substitution, row-major
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) {
65 LhsEval lhsEval(lhs);
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);
69 Scalar l_ii(0);
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);
74 l_ii = it.value();
75 ++it;
76 } else if (it && it.index() == i)
77 ++it;
78 for (; it; ++it) {
79 tmp = numext::madd<Scalar>(-it.value(), other.coeff(it.index(), col), tmp);
80 }
81
82 EIGEN_IF_CONSTEXPR (Mode & UnitDiag)
83 other.coeffRef(i, col) = tmp;
84 else
85 other.coeffRef(i, col) = tmp / l_ii;
86 }
87 }
88 }
89};
90
91// forward substitution, col-major
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) {
98 LhsEval lhsEval(lhs);
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)) // optimization when other is actually sparse
103 {
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);
108 tmp /= it.value();
109 }
110 if (it && it.index() == i) ++it;
111 for (; it; ++it) {
112 other.coeffRef(it.index(), col) = numext::madd<Scalar>(-tmp, it.value(), other.coeffRef(it.index(), col));
113 }
114 }
115 }
116 }
117 }
118};
119
120// backward substitution, col-major
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)) // optimization when other is actually sparse
132 {
133 EIGEN_IF_CONSTEXPR (!(Mode & UnitDiag)) {
134 // TODO: replace this with a binary search. make sure the binary search is safe for partially sorted
135 // elements
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();
140 }
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));
144 }
145 }
146 }
147 }
148 }
149};
150
151} // end namespace internal
152
153#ifndef EIGEN_PARSED_BY_DOXYGEN
154
155template <typename ExpressionType, unsigned int Mode>
156template <typename OtherDerived>
157void TriangularViewImpl<ExpressionType, Mode, Sparse>::solveInPlace(MatrixBase<OtherDerived>& other) const {
158 eigen_assert(derived().cols() == derived().rows() && derived().cols() == other.rows());
159 eigen_assert((!(Mode & ZeroDiag)) && bool(Mode & (Upper | Lower)));
160
161 enum { copy = internal::traits<OtherDerived>::Flags & RowMajorBit };
162
163 using OtherCopy =
164 std::conditional_t<copy, typename internal::plain_matrix_type_column_major<OtherDerived>::type, OtherDerived&>;
165 OtherCopy otherCopy(other.derived());
166
167 internal::sparse_solve_triangular_selector<ExpressionType, std::remove_reference_t<OtherCopy>, Mode>::run(
168 derived().nestedExpression(), otherCopy);
169
170 if (copy) other = otherCopy;
171}
172#endif
173
174// pure sparse path
175
176namespace internal {
177
178template <typename Lhs, typename Rhs, int Mode,
179 int UpLo = (Mode & Lower) ? Lower
180 : (Mode & Upper) ? Upper
181 : -1,
182 int StorageOrder = int(Lhs::Flags) & RowMajorBit>
183struct sparse_solve_triangular_sparse_selector;
184
185// True when the rhs exposes raw CSC storage with a StorageIndex matching the lhs, so a
186// column's stored index slice can serve as the reach roots directly (no bIdx copy). A
187// SparseVector qualifies too -- it is a single compressed column, handled below via the
188// null-outerIndexPtr guard (its outerIndexPtr() is null since it has no outer array).
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>;
193
194// The reach for a column arrives in one of three compile-time-known orders, but the
195// output column must store ascending inner index. An unordered reach (Ordered == false,
196// the pointer/DFS path -- Upper irrelevant) is sorted ascending here; the iterator reach
197// already arrives in solve order, ascending for lower and descending for upper, so it
198// needs no pass of its own -- the descending one is read back to front at insertion.
199template <bool Ordered, bool Upper, typename StorageIndex>
200struct reach_reorder { // Ordered == false: unordered pointer/DFS reach
201 static void run(StorageIndex* first, StorageIndex* last) { std::sort(first, last); }
202};
203template <bool Upper, typename StorageIndex>
204struct reach_reorder<true, Upper, StorageIndex> { // iterator reach: already in solve order
205 static void run(StorageIndex* /*first*/, StorageIndex* /*last*/) {}
206};
207
208// Common per-column finish: read the reach xi[top..n) in ascending inner index order,
209// insert reading values from xwork, and clear xwork and mark for the next column.
210// The reach is a structural bound, not a numeric one: a reached coefficient can be
211// exactly zero (a zero rhs entry, or numerical cancellation), so skip exact zeros at
212// insertion. This matches the AmbiVector path, which pruned zeros, and keeps a zero rhs
213// from materializing O(|reach|) stored zeros. xwork and mark are cleared regardless.
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);
217 // Only the iterator upper reach is descending; it is read back to front, so k walks
218 // up while the load walks down. Every other order is ascending by this point.
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);
224 mark[j] = 0;
225 }
226}
227
228// Column loop, fast path: the rhs matches (see rhs_matching_slice), so each column's
229// stored index slice is the reach root list and the value slice is scattered directly
230// -- no bIdx copy, so iwork is just 2n (xi | pstack).
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); // xi | pstack
236 StorageIndex* xi = iwork.data();
237 for (Index col = 0; col < other.cols(); ++col) {
238 const StorageIndex* outer = other.outerIndexPtr(); // null for a SparseVector (single column)
239 const StorageIndex* nnz = other.innerNonZeroPtr(); // null when compressed
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; // the column's stored indices
243 const Scalar* vals = other.valuePtr() + p;
244 // Roots are the rhs slice directly (no bIdx copy, so iwork stays 2n). An exact-zero
245 // stored rhs entry is seeded harmlessly -- it propagates zeros and is dropped at
246 // insertion; filtering it here would cost a compacted root buffer (the 3n path).
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);
250 }
251}
252
253// Column loop, fallback: read each column through the InnerIterator, copying indices
254// into the bIdx third of a 3n iwork. For a rhs without raw storage or with a
255// mismatched index type.
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); // xi | pstack | bIdx
261 StorageIndex* xi = iwork.data();
262 StorageIndex* bIdx = iwork.data() + 2 * n;
263 for (Index col = 0; col < other.cols(); ++col) {
264 Index bCount = 0;
265 for (typename Rhs::InnerIterator it(other, col); it; ++it) {
266 if (numext::is_exactly_zero(it.value())) continue; // a zero root seeds nothing; xwork stays clear there
267 bIdx[bCount] = StorageIndex(it.index());
268 xwork[it.index()] = it.value();
269 ++bCount;
270 }
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);
273 }
274}
275
276// Reach-based (Gilbert-Peierls) sparse triangular solve, col-major, for lower OR
277// upper. Only the columns reachable from each rhs column's pattern are touched, so
278// the cost is O(|reach| + flops) per column instead of a dense O(n)-per-column sweep
279// (which also pays a coeff(i,i) binary search per row in the upper case). It is
280// the sole col-major sparse-sparse selector, dispatching lower/upper via the UpLo
281// template argument. reach_solve_dense leaves the solution values in
282// xwork and the reached indices in iwork[top..n); reach_solve_columns (slice or
283// fallback, selected on the rhs storage) scatters each column and solves, and
284// reach_insert_column reads the values out and restores mark/xwork. Only mark and
285// xwork need zeroing -- iwork is entirely written before read.
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);
295 res.finalize();
296 other = res.markAsRValue();
297}
298
299// forward and backward substitution, col-major
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);
304 }
305};
306
307} // end namespace internal
308
309#ifndef EIGEN_PARSED_BY_DOXYGEN
310template <typename ExpressionType, unsigned int Mode>
311template <typename OtherDerived>
312void TriangularViewImpl<ExpressionType, Mode, Sparse>::solveInPlace(SparseMatrixBase<OtherDerived>& other) const {
313 eigen_assert(derived().cols() == derived().rows() && derived().cols() == other.rows());
314 eigen_assert((!(Mode & ZeroDiag)) && bool(Mode & (Upper | Lower)));
315
316 internal::sparse_solve_triangular_sparse_selector<ExpressionType, OtherDerived, Mode>::run(
317 derived().nestedExpression(), other.derived());
318}
319#endif
320
321} // end namespace Eigen
322
323#endif // EIGEN_SPARSETRIANGULARSOLVER_H
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