11#ifndef EIGEN_TRIANGULARMATRIXVECTOR_H
12#define EIGEN_TRIANGULARMATRIXVECTOR_H
15#include "../InternalHeaderCheck.h"
21template <
typename Index,
int Mode,
typename LhsScalar,
bool ConjLhs,
typename RhsScalar,
bool ConjRhs,
22 int StorageOrder,
int Version = Specialized>
23struct triangular_matrix_vector_product;
25template <
typename Index,
int Mode,
typename LhsScalar,
bool ConjLhs,
typename RhsScalar,
bool ConjRhs,
int Version>
26struct triangular_matrix_vector_product<Index, Mode, LhsScalar, ConjLhs, RhsScalar, ConjRhs,
ColMajor, Version> {
27 using ResScalar =
typename ScalarBinaryOpTraits<LhsScalar, RhsScalar>::ReturnType;
28 static constexpr bool IsLower = (Mode &
Lower) ==
Lower;
31 static EIGEN_DONT_INLINE
void run(Index _rows, Index _cols,
const LhsScalar* lhs_, Index lhsStride,
32 const RhsScalar* rhs_, Index rhsIncr, ResScalar* res_, Index resIncr,
33 const RhsScalar& alpha);
36template <
typename Index,
int Mode,
typename LhsScalar,
bool ConjLhs,
typename RhsScalar,
bool ConjRhs,
int Version>
37EIGEN_DONT_INLINE
void triangular_matrix_vector_product<Index, Mode, LhsScalar, ConjLhs, RhsScalar, ConjRhs,
ColMajor,
38 Version>::run(Index _rows, Index _cols,
const LhsScalar* lhs_,
39 Index lhsStride,
const RhsScalar* rhs_,
40 Index rhsIncr, ResScalar* res_, Index resIncr,
41 const RhsScalar& alpha) {
42 static const Index PanelWidth = EIGEN_TUNE_TRIANGULAR_PANEL_WIDTH;
43 Index size = (std::min)(_rows, _cols);
44 Index rows = IsLower ? _rows : (std::min)(_rows, _cols);
45 Index cols = IsLower ? (std::min)(_rows, _cols) : _cols;
47 using LhsMapper = const_blas_data_mapper<LhsScalar, Index, ColMajor>;
48 using RhsMapper = const_blas_data_mapper<RhsScalar, Index, RowMajor>;
53 for (Index pi = 0; pi < size; pi += PanelWidth) {
54 Index actualPanelWidth = (std::min)(PanelWidth, size - pi);
58 EIGEN_IF_CONSTEXPR (IsLower) {
60 for (; k + 1 < actualPanelWidth; k += 2) {
63 ResScalar s0 = alpha * cjr(rhs_[i0 * rhsIncr]);
64 ResScalar s1 = alpha * cjr(rhs_[i1 * rhsIncr]);
65 const LhsScalar* EIGEN_RESTRICT c0 = lhs_ + i0 * lhsStride;
66 const LhsScalar* EIGEN_RESTRICT c1 = lhs_ + i1 * lhsStride;
69 EIGEN_IF_CONSTEXPR (!(HasUnitDiag || HasZeroDiag)) res_[i0] += s0 * cjl(c0[i0]);
72 ResScalar r1 = s0 * cjl(c0[i1]);
73 EIGEN_IF_CONSTEXPR (!(HasUnitDiag || HasZeroDiag)) r1 += s1 * cjl(c1[i1]);
77 Index panelEnd = pi + actualPanelWidth;
78 for (Index j = i1 + 1; j < panelEnd; ++j) res_[j] += s0 * cjl(c0[j]) + s1 * cjl(c1[j]);
80 EIGEN_IF_CONSTEXPR (HasUnitDiag) {
85 if (k < actualPanelWidth) {
87 ResScalar s = alpha * cjr(rhs_[i * rhsIncr]);
88 const LhsScalar* EIGEN_RESTRICT c = lhs_ + i * lhsStride;
89 EIGEN_IF_CONSTEXPR (!(HasUnitDiag || HasZeroDiag)) res_[i] += s * cjl(c[i]);
90 EIGEN_IF_CONSTEXPR (HasUnitDiag) res_[i] += s;
95 for (; k + 1 < actualPanelWidth; k += 2) {
98 ResScalar s0 = alpha * cjr(rhs_[i0 * rhsIncr]);
99 ResScalar s1 = alpha * cjr(rhs_[i1 * rhsIncr]);
100 const LhsScalar* EIGEN_RESTRICT c0 = lhs_ + i0 * lhsStride;
101 const LhsScalar* EIGEN_RESTRICT c1 = lhs_ + i1 * lhsStride;
104 for (Index j = pi; j < i0; ++j) res_[j] += s0 * cjl(c0[j]) + s1 * cjl(c1[j]);
108 ResScalar r0 = s1 * cjl(c1[i0]);
109 EIGEN_IF_CONSTEXPR (!(HasUnitDiag || HasZeroDiag)) r0 += s0 * cjl(c0[i0]);
113 EIGEN_IF_CONSTEXPR (!(HasUnitDiag || HasZeroDiag)) res_[i1] += s1 * cjl(c1[i1]);
115 EIGEN_IF_CONSTEXPR (HasUnitDiag) {
120 if (k < actualPanelWidth) {
122 ResScalar s = alpha * cjr(rhs_[i * rhsIncr]);
123 const LhsScalar* EIGEN_RESTRICT c = lhs_ + i * lhsStride;
124 for (Index j = pi; j < i; ++j) res_[j] += s * cjl(c[j]);
125 EIGEN_IF_CONSTEXPR (!(HasUnitDiag || HasZeroDiag)) res_[i] += s * cjl(c[i]);
126 EIGEN_IF_CONSTEXPR (HasUnitDiag) res_[i] += s;
131 Index r = IsLower ? rows - pi - actualPanelWidth : pi;
133 Index s = IsLower ? pi + actualPanelWidth : 0;
134 general_matrix_vector_product<Index, LhsScalar, LhsMapper,
ColMajor, ConjLhs, RhsScalar, RhsMapper, ConjRhs,
135 BuiltIn>::run(r, actualPanelWidth, LhsMapper(&lhs_[pi * lhsStride + s], lhsStride),
136 RhsMapper(&rhs_[pi * rhsIncr], rhsIncr), &res_[s], resIncr, alpha);
139 EIGEN_IF_CONSTEXPR (!IsLower) {
141 general_matrix_vector_product<Index, LhsScalar, LhsMapper, ColMajor, ConjLhs, RhsScalar, RhsMapper, ConjRhs>::run(
142 rows, cols - size, LhsMapper(&lhs_[size * lhsStride], lhsStride), RhsMapper(&rhs_[size * rhsIncr], rhsIncr),
143 res_, resIncr, alpha);
148template <
typename Index,
int Mode,
typename LhsScalar,
bool ConjLhs,
typename RhsScalar,
bool ConjRhs,
int Version>
149struct triangular_matrix_vector_product<Index, Mode, LhsScalar, ConjLhs, RhsScalar, ConjRhs,
RowMajor, Version> {
150 using ResScalar =
typename ScalarBinaryOpTraits<LhsScalar, RhsScalar>::ReturnType;
151 static constexpr bool IsLower = (Mode &
Lower) ==
Lower;
154 static EIGEN_DONT_INLINE
void run(Index _rows, Index _cols,
const LhsScalar* lhs_, Index lhsStride,
155 const RhsScalar* rhs_, Index rhsIncr, ResScalar* res_, Index resIncr,
156 const ResScalar& alpha);
159template <
typename Index,
int Mode,
typename LhsScalar,
bool ConjLhs,
typename RhsScalar,
bool ConjRhs,
int Version>
160EIGEN_DONT_INLINE
void triangular_matrix_vector_product<Index, Mode, LhsScalar, ConjLhs, RhsScalar, ConjRhs,
RowMajor,
161 Version>::run(Index _rows, Index _cols,
const LhsScalar* lhs_,
162 Index lhsStride,
const RhsScalar* rhs_,
163 Index rhsIncr, ResScalar* res_, Index resIncr,
164 const ResScalar& alpha) {
165 static const Index PanelWidth = EIGEN_TUNE_TRIANGULAR_PANEL_WIDTH;
166 Index diagSize = (std::min)(_rows, _cols);
167 Index rows = IsLower ? _rows : diagSize;
168 Index cols = IsLower ? diagSize : _cols;
170 using LhsMapper = const_blas_data_mapper<LhsScalar, Index, RowMajor>;
171 using RhsMapper = const_blas_data_mapper<RhsScalar, Index, RowMajor>;
173 conj_if<ConjLhs> cjl;
174 conj_if<ConjRhs> cjr;
176 for (Index pi = 0; pi < diagSize; pi += PanelWidth) {
177 Index actualPanelWidth = (std::min)(PanelWidth, diagSize - pi);
181 for (Index k = 0; k < actualPanelWidth; ++k) {
183 const LhsScalar* EIGEN_RESTRICT row_i = lhs_ + i * lhsStride;
184 ResScalar dot = ResScalar(0);
186 EIGEN_IF_CONSTEXPR (IsLower) {
188 Index len = (HasUnitDiag || HasZeroDiag) ? k : k + 1;
189 for (Index j = 0; j < len; ++j) dot += cjl(row_i[s + j]) * cjr(rhs_[s + j]);
191 Index s = (HasUnitDiag || HasZeroDiag) ? i + 1 : i;
192 Index len = pi + actualPanelWidth - s;
193 for (Index j = 0; j < len; ++j) dot += cjl(row_i[s + j]) * cjr(rhs_[s + j]);
195 res_[i * resIncr] += alpha * dot;
196 EIGEN_IF_CONSTEXPR (HasUnitDiag) res_[i * resIncr] += alpha * cjr(rhs_[i]);
200 Index r = IsLower ? pi : cols - pi - actualPanelWidth;
202 Index s = IsLower ? 0 : pi + actualPanelWidth;
203 general_matrix_vector_product<Index, LhsScalar, LhsMapper,
RowMajor, ConjLhs, RhsScalar, RhsMapper, ConjRhs,
204 BuiltIn>::run(actualPanelWidth, r, LhsMapper(&lhs_[pi * lhsStride + s], lhsStride),
205 RhsMapper(&rhs_[s], rhsIncr), &res_[pi * resIncr], resIncr, alpha);
208 EIGEN_IF_CONSTEXPR (IsLower) {
209 if (rows > diagSize) {
210 general_matrix_vector_product<Index, LhsScalar, LhsMapper, RowMajor, ConjLhs, RhsScalar, RhsMapper, ConjRhs>::run(
211 rows - diagSize, cols, LhsMapper(&lhs_[diagSize * lhsStride], lhsStride), RhsMapper(rhs_, rhsIncr),
212 &res_[diagSize * resIncr], resIncr, alpha);
221template <
int Mode,
int StorageOrder>
228template <
int Mode,
typename Lhs,
typename Rhs>
229struct triangular_product_impl<Mode, true, Lhs, false, Rhs, true> {
230 template <
typename Dest>
231 static void run(Dest& dst,
const Lhs& lhs,
const Rhs& rhs,
const typename Dest::Scalar& alpha) {
232 eigen_assert(dst.rows() == lhs.rows() && dst.cols() == rhs.cols());
235 lhs, rhs, dst, alpha);
239template <
int Mode,
typename Lhs,
typename Rhs>
240struct triangular_product_impl<Mode, false, Lhs, true, Rhs, false> {
241 template <
typename Dest>
242 static void run(Dest& dst,
const Lhs& lhs,
const Rhs& rhs,
const typename Dest::Scalar& alpha) {
243 eigen_assert(dst.rows() == lhs.rows() && dst.cols() == rhs.cols());
245 Transpose<Dest> dstT(dst);
249 lhs.transpose(), dstT,
258template <
int Mode,
typename Lhs,
typename Rhs,
typename Dest,
typename LhsScalar>
259EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void trmv_correct_unit_diagonal(
const Lhs& lhs,
const Rhs& rhs, Dest& dest,
260 const LhsScalar& lhs_alpha) {
262 if (!numext::is_exactly_one(lhs_alpha)) {
263 Index diagSize = (std::min)(lhs.rows(), lhs.cols());
264 dest.head(diagSize) -= (lhs_alpha - LhsScalar(1)) * rhs.head(diagSize);
270struct trmv_selector<Mode,
ColMajor> {
271 template <
typename Lhs,
typename Rhs,
typename Dest>
272 static void run(
const Lhs& lhs,
const Rhs& rhs, Dest& dest,
const typename Dest::Scalar& alpha) {
273 using LhsScalar =
typename Lhs::Scalar;
274 using RhsScalar =
typename Rhs::Scalar;
275 using ResScalar =
typename Dest::Scalar;
277 using LhsBlasTraits = internal::blas_traits<Lhs>;
278 using ActualLhsType =
typename LhsBlasTraits::DirectLinearAccessType;
279 using RhsBlasTraits = internal::blas_traits<Rhs>;
280 using ActualRhsType =
typename RhsBlasTraits::DirectLinearAccessType;
281 add_const_on_value_type_t<ActualLhsType> actualLhs = LhsBlasTraits::extract(lhs);
282 add_const_on_value_type_t<ActualRhsType> actualRhs = RhsBlasTraits::extract(rhs);
284 LhsScalar lhs_alpha = LhsBlasTraits::extractScalarFactor(lhs);
285 RhsScalar rhs_alpha = RhsBlasTraits::extractScalarFactor(rhs);
286 ResScalar actualAlpha = alpha * lhs_alpha * rhs_alpha;
290 constexpr bool EvalToDestAtCompileTime = Dest::InnerStrideAtCompileTime == 1;
291 constexpr bool ComplexByReal = (NumTraits<LhsScalar>::IsComplex) && (!NumTraits<RhsScalar>::IsComplex);
292 constexpr bool MightCannotUseDest = (Dest::InnerStrideAtCompileTime != 1) || ComplexByReal;
294 gemv_static_vector_if<ResScalar, Dest::SizeAtCompileTime, Dest::MaxSizeAtCompileTime, MightCannotUseDest>
297 gemv_destination_policy<RhsScalar, ResScalar, EvalToDestAtCompileTime, ComplexByReal> destPolicy(actualAlpha);
299 ei_declare_aligned_stack_constructed_variable(ResScalar, actualDestPtr, dest.size(),
300 destPolicy.eval_to_dest() ? dest.data() : static_dest.data());
302 destPolicy.prepare(dest, actualDestPtr);
304 internal::triangular_matrix_vector_product<Index, Mode, LhsScalar, LhsBlasTraits::NeedToConjugate, RhsScalar,
305 RhsBlasTraits::NeedToConjugate,
306 ColMajor>::run(actualLhs.rows(), actualLhs.cols(), actualLhs.data(),
307 actualLhs.outerStride(), actualRhs.data(),
308 actualRhs.innerStride(), actualDestPtr, 1,
309 destPolicy.compatible_alpha());
311 destPolicy.copy_back(dest, actualDestPtr);
313 trmv_correct_unit_diagonal<Mode>(lhs, rhs, dest, lhs_alpha);
318struct trmv_selector<Mode,
RowMajor> {
319 template <
typename Lhs,
typename Rhs,
typename Dest>
320 static void run(
const Lhs& lhs,
const Rhs& rhs, Dest& dest,
const typename Dest::Scalar& alpha) {
321 using LhsScalar =
typename Lhs::Scalar;
322 using RhsScalar =
typename Rhs::Scalar;
323 using ResScalar =
typename Dest::Scalar;
325 using LhsBlasTraits = internal::blas_traits<Lhs>;
326 using ActualLhsType =
typename LhsBlasTraits::DirectLinearAccessType;
327 using RhsBlasTraits = internal::blas_traits<Rhs>;
328 using ActualRhsType =
typename RhsBlasTraits::DirectLinearAccessType;
329 using ActualRhsTypeCleaned = internal::remove_all_t<ActualRhsType>;
331 std::add_const_t<ActualLhsType> actualLhs = LhsBlasTraits::extract(lhs);
332 std::add_const_t<ActualRhsType> actualRhs = RhsBlasTraits::extract(rhs);
334 LhsScalar lhs_alpha = LhsBlasTraits::extractScalarFactor(lhs);
335 RhsScalar rhs_alpha = RhsBlasTraits::extractScalarFactor(rhs);
336 ResScalar actualAlpha = alpha * lhs_alpha * rhs_alpha;
338 constexpr bool DirectlyUseRhs = ActualRhsTypeCleaned::InnerStrideAtCompileTime == 1;
340 gemv_static_vector_if<RhsScalar, ActualRhsTypeCleaned::SizeAtCompileTime,
341 ActualRhsTypeCleaned::MaxSizeAtCompileTime, !DirectlyUseRhs>
344 ei_declare_aligned_stack_constructed_variable(
345 RhsScalar, actualRhsPtr, actualRhs.size(),
346 DirectlyUseRhs ?
const_cast<RhsScalar*
>(actualRhs.data()) : static_rhs.data());
348 gemv_prepare_rhs<DirectlyUseRhs>(actualRhs, actualRhsPtr);
350 internal::triangular_matrix_vector_product<Index, Mode, LhsScalar, LhsBlasTraits::NeedToConjugate, RhsScalar,
351 RhsBlasTraits::NeedToConjugate,
RowMajor>::run(actualLhs.rows(),
354 actualLhs.outerStride(),
360 trmv_correct_unit_diagonal<Mode>(lhs, rhs, dest, lhs_alpha);
@ 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
@ RowMajor
Definition Constants.h:321
constexpr unsigned int RowMajorBit
Definition Constants.h:71