11#ifndef EIGEN_SOLVETRIANGULAR_H
12#define EIGEN_SOLVETRIANGULAR_H
15#include "./InternalHeaderCheck.h"
23template <
typename LhsScalar,
typename RhsScalar,
typename Index,
int S
ide,
int Mode,
bool Conjugate,
int StorageOrder>
24struct triangular_solve_vector;
26template <
typename Scalar,
typename Index,
int Side,
int Mode,
bool Conjugate,
int TriStorageOrder,
27 int OtherStorageOrder,
int OtherInnerStride>
28struct triangular_solve_matrix;
31template <
typename Lhs,
typename Rhs,
int S
ide>
34 enum { RhsIsVectorAtCompileTime = (Side ==
OnTheLeft ? Rhs::ColsAtCompileTime : Rhs::RowsAtCompileTime) == 1 };
38 Unrolling = (RhsIsVectorAtCompileTime && Rhs::SizeAtCompileTime != Dynamic && Rhs::SizeAtCompileTime <= 8)
41 RhsVectors = RhsIsVectorAtCompileTime ? 1 : Dynamic
45template <
typename Lhs,
typename Rhs,
48 int Unrolling = trsolve_traits<Lhs, Rhs, Side>::Unrolling,
49 int RhsVectors = trsolve_traits<Lhs, Rhs, Side>::RhsVectors>
50struct triangular_solver_selector;
52template <
typename Lhs,
typename Rhs,
int S
ide,
int Mode>
53struct triangular_solver_selector<Lhs, Rhs, Side, Mode, NoUnrolling, 1> {
54 using LhsScalar =
typename Lhs::Scalar;
55 using RhsScalar =
typename Rhs::Scalar;
56 using LhsProductTraits = blas_traits<Lhs>;
57 using ActualLhsType =
typename LhsProductTraits::DirectLinearAccessType;
58 using ActualLhsTypeCleaned = remove_all_t<ActualLhsType>;
59 using MappedRhs = Map<Matrix<RhsScalar, Dynamic, 1>,
Aligned>;
60 static EIGEN_DEVICE_FUNC
void run(
const Lhs& lhs, Rhs& rhs) {
61 add_const_on_value_type_t<ActualLhsType> actualLhs = LhsProductTraits::extract(lhs);
65 bool useRhsDirectly = Rhs::InnerStrideAtCompileTime == 1 || rhs.innerStride() == 1;
67 ei_declare_aligned_stack_constructed_variable(RhsScalar, actualRhs, rhs.size(), (useRhsDirectly ? rhs.data() : 0));
69 if (!useRhsDirectly) MappedRhs(actualRhs, rhs.size()) = rhs;
71 triangular_solve_vector<LhsScalar, RhsScalar, Index, Side, Mode, LhsProductTraits::NeedToConjugate,
72 (int(ActualLhsTypeCleaned::Flags) &
RowMajorBit) ? RowMajor
75 actualLhs.outerStride(),
78 if (!useRhsDirectly) rhs = MappedRhs(actualRhs, rhs.size());
83template <
typename Lhs,
typename Rhs,
int S
ide,
int Mode>
84struct triangular_solver_selector<Lhs, Rhs, Side, Mode, NoUnrolling, Dynamic> {
85 using Scalar =
typename Rhs::Scalar;
86 using LhsProductTraits = blas_traits<Lhs>;
87 using ActualLhsType =
typename LhsProductTraits::DirectLinearAccessType;
89 static EIGEN_DEVICE_FUNC
void run(
const Lhs& lhs, Rhs& rhs) {
90 add_const_on_value_type_t<ActualLhsType> actualLhs = LhsProductTraits::extract(lhs);
92 const Index size = lhs.rows();
93 const Index othersize = Side ==
OnTheLeft ? rhs.cols() : rhs.rows();
95 using BlockingType = internal::gemm_blocking_space<(Rhs::Flags &
RowMajorBit) ? RowMajor :
ColMajor, Scalar, Scalar,
96 Rhs::MaxRowsAtCompileTime, Rhs::MaxColsAtCompileTime,
97 Lhs::MaxRowsAtCompileTime, 4>;
100 if (actualLhs.size() == 0 || rhs.size() == 0) {
104 BlockingType blocking(rhs.rows(), rhs.cols(), size, 1,
false);
106 triangular_solve_matrix<Scalar, Index, Side, Mode, LhsProductTraits::NeedToConjugate,
109 Rhs::InnerStrideAtCompileTime>::run(size, othersize, &actualLhs.coeffRef(0, 0),
110 actualLhs.outerStride(), &rhs.coeffRef(0, 0),
111 rhs.innerStride(), rhs.outerStride(), blocking);
119template <
typename Lhs,
typename Rhs,
int Mode,
int LoopIndex,
int Size,
bool Stop = LoopIndex == Size>
120struct triangular_solver_unroller;
122template <
typename Lhs,
typename Rhs,
int Mode,
int LoopIndex,
int Size>
123struct triangular_solver_unroller<Lhs, Rhs, Mode, LoopIndex, Size, false> {
126 DiagIndex = IsLower ? LoopIndex : Size - LoopIndex - 1,
127 StartIndex = IsLower ? 0 : DiagIndex + 1
129 static EIGEN_DEVICE_FUNC
void run(
const Lhs& lhs, Rhs& rhs) {
131 rhs.coeffRef(DiagIndex) -= lhs.row(DiagIndex)
132 .template segment<LoopIndex>(StartIndex)
134 .cwiseProduct(rhs.template segment<LoopIndex>(StartIndex))
137 EIGEN_IF_CONSTEXPR (!(Mode & UnitDiag)) rhs.coeffRef(DiagIndex) /= lhs.coeff(DiagIndex, DiagIndex);
139 triangular_solver_unroller<Lhs, Rhs, Mode, LoopIndex + 1, Size>::run(lhs, rhs);
143template <
typename Lhs,
typename Rhs,
int Mode,
int LoopIndex,
int Size>
144struct triangular_solver_unroller<Lhs, Rhs, Mode, LoopIndex, Size, true> {
145 static EIGEN_DEVICE_FUNC
void run(
const Lhs&, Rhs&) {}
148template <
typename Lhs,
typename Rhs,
int Mode>
149struct triangular_solver_selector<Lhs, Rhs,
OnTheLeft, Mode, CompleteUnrolling, 1> {
150 static EIGEN_DEVICE_FUNC
void run(
const Lhs& lhs, Rhs& rhs) {
151 triangular_solver_unroller<Lhs, Rhs, Mode, 0, Rhs::SizeAtCompileTime>::run(lhs, rhs);
155template <
typename Lhs,
typename Rhs,
int Mode>
156struct triangular_solver_selector<Lhs, Rhs,
OnTheRight, Mode, CompleteUnrolling, 1> {
157 static EIGEN_DEVICE_FUNC
void run(
const Lhs& lhs, Rhs& rhs) {
158 Transpose<const Lhs> trLhs(lhs);
159 Transpose<Rhs> trRhs(rhs);
161 triangular_solver_unroller<Transpose<const Lhs>, Transpose<Rhs>,
163 Rhs::SizeAtCompileTime>::run(trLhs, trRhs);
173#ifndef EIGEN_PARSED_BY_DOXYGEN
174template <
typename MatrixType,
unsigned int Mode>
175template <
int S
ide,
typename OtherDerived>
176EIGEN_DEVICE_FUNC
void TriangularViewImpl<MatrixType, Mode, Dense>::solveInPlace(
177 const MatrixBase<OtherDerived>& _other)
const {
178 OtherDerived& other = _other.const_cast_derived();
179 eigen_assert(derived().cols() == derived().rows() && ((Side ==
OnTheLeft && derived().cols() == other.rows()) ||
180 (Side ==
OnTheRight && derived().cols() == other.cols())));
181 eigen_assert((!(
int(Mode) &
int(
ZeroDiag))) &&
bool(
int(Mode) & (
int(
Upper) |
int(
Lower))));
183 if (derived().cols() == 0)
return;
186 OtherFlags = internal::traits<OtherDerived>::Flags,
188 (OtherFlags &
RowMajorBit) && OtherDerived::IsVectorAtCompileTime && OtherDerived::SizeAtCompileTime != 1,
191 using OtherPlainObject =
192 std::conditional_t<IsRowMajorVector, typename internal::plain_matrix_type_column_major<OtherDerived>::type,
193 typename internal::plain_matrix_type<OtherDerived>::type>;
194 using OtherCopy = std::conditional_t<copy, OtherPlainObject, OtherDerived&>;
195 OtherCopy otherCopy(other);
197 internal::triangular_solver_selector<MatrixType, std::remove_reference_t<OtherCopy>, Side, Mode>::run(
198 derived().nestedExpression(), otherCopy);
200 if (copy) other = otherCopy;
203template <
typename Derived,
unsigned int Mode>
204template <
int S
ide,
typename Other>
205const internal::triangular_solve_retval<Side, TriangularView<Derived, Mode>, Other>
206TriangularViewImpl<Derived, Mode, Dense>::solve(
const MatrixBase<Other>& other)
const {
207 return internal::triangular_solve_retval<Side, TriangularViewType, Other>(derived(), other.derived());
213template <
int S
ide,
typename TriangularType,
typename Rhs>
214struct traits<triangular_solve_retval<Side, TriangularType, Rhs> > {
215 using ReturnType =
typename internal::plain_matrix_type_column_major<Rhs>::type;
218template <
int S
ide,
typename TriangularType,
typename Rhs>
219struct triangular_solve_retval :
public ReturnByValue<triangular_solve_retval<Side, TriangularType, Rhs> > {
220 using Base = ReturnByValue<triangular_solve_retval>;
222 triangular_solve_retval(
const TriangularType& tri,
const Rhs& rhs) : m_triangularMatrix(tri), m_rhs(rhs) {}
224 constexpr Index rows() const noexcept {
return m_rhs.rows(); }
225 constexpr Index cols() const noexcept {
return m_rhs.cols(); }
227 template <
typename Dest>
228 inline void evalTo(Dest& dst)
const {
229 if (!is_same_dense(dst, m_rhs)) dst = m_rhs;
230 m_triangularMatrix.template solveInPlace<Side>(dst);
234 const TriangularType& m_triangularMatrix;
235 typename Rhs::Nested m_rhs;
@ ZeroDiag
Definition Constants.h:218
@ Lower
Definition Constants.h:212
@ Upper
Definition Constants.h:214
@ Aligned
Definition Constants.h:243
@ ColMajor
Definition Constants.h:319
@ RowMajor
Definition Constants.h:321
@ OnTheLeft
Definition Constants.h:332
@ OnTheRight
Definition Constants.h:334
constexpr unsigned int DirectAccessBit
Definition Constants.h:160
constexpr unsigned int RowMajorBit
Definition Constants.h:71