Eigen  5.0.1
 
Loading...
Searching...
No Matches
TriangularMatrixVector.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2009 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_TRIANGULARMATRIXVECTOR_H
12#define EIGEN_TRIANGULARMATRIXVECTOR_H
13
14// IWYU pragma: private
15#include "../InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20
21template <typename Index, int Mode, typename LhsScalar, bool ConjLhs, typename RhsScalar, bool ConjRhs,
22 int StorageOrder, int Version = Specialized>
23struct triangular_matrix_vector_product;
24
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;
29 static constexpr bool HasUnitDiag = (Mode & UnitDiag) == UnitDiag;
30 static constexpr bool HasZeroDiag = (Mode & ZeroDiag) == ZeroDiag;
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);
34};
35
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;
46
47 using LhsMapper = const_blas_data_mapper<LhsScalar, Index, ColMajor>;
48 using RhsMapper = const_blas_data_mapper<RhsScalar, Index, RowMajor>;
49
50 conj_if<ConjLhs> cjl;
51 conj_if<ConjRhs> cjr;
52
53 for (Index pi = 0; pi < size; pi += PanelWidth) {
54 Index actualPanelWidth = (std::min)(PanelWidth, size - pi);
55
56 // Process the triangular panel using raw pointer operations with 2-column batching
57 // to eliminate expression template overhead and share result loads/stores.
58 EIGEN_IF_CONSTEXPR (IsLower) {
59 Index k = 0;
60 for (; k + 1 < actualPanelWidth; k += 2) {
61 Index i0 = pi + k;
62 Index i1 = i0 + 1;
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;
67
68 // Diagonal of column 0
69 EIGEN_IF_CONSTEXPR (!(HasUnitDiag || HasZeroDiag)) res_[i0] += s0 * cjl(c0[i0]);
70 // Row i1: contribution from column 0 + diagonal of column 1
71 {
72 ResScalar r1 = s0 * cjl(c0[i1]);
73 EIGEN_IF_CONSTEXPR (!(HasUnitDiag || HasZeroDiag)) r1 += s1 * cjl(c1[i1]);
74 res_[i1] += r1;
75 }
76 // Shared rows where both columns contribute
77 Index panelEnd = pi + actualPanelWidth;
78 for (Index j = i1 + 1; j < panelEnd; ++j) res_[j] += s0 * cjl(c0[j]) + s1 * cjl(c1[j]);
79
80 EIGEN_IF_CONSTEXPR (HasUnitDiag) {
81 res_[i0] += s0;
82 res_[i1] += s1;
83 }
84 }
85 if (k < actualPanelWidth) {
86 Index i = pi + k;
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;
91 }
92 } else {
93 // Upper triangular: process 2 columns at a time
94 Index k = 0;
95 for (; k + 1 < actualPanelWidth; k += 2) {
96 Index i0 = pi + k;
97 Index i1 = i0 + 1;
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;
102
103 // Shared rows before the diagonal block
104 for (Index j = pi; j < i0; ++j) res_[j] += s0 * cjl(c0[j]) + s1 * cjl(c1[j]);
105
106 // Row i0: diagonal of col0 + contribution from col1
107 {
108 ResScalar r0 = s1 * cjl(c1[i0]);
109 EIGEN_IF_CONSTEXPR (!(HasUnitDiag || HasZeroDiag)) r0 += s0 * cjl(c0[i0]);
110 res_[i0] += r0;
111 }
112 // Diagonal of column 1
113 EIGEN_IF_CONSTEXPR (!(HasUnitDiag || HasZeroDiag)) res_[i1] += s1 * cjl(c1[i1]);
114
115 EIGEN_IF_CONSTEXPR (HasUnitDiag) {
116 res_[i0] += s0;
117 res_[i1] += s1;
118 }
119 }
120 if (k < actualPanelWidth) {
121 Index i = pi + k;
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;
127 }
128 }
129
130 // Rectangular part: delegate to optimized GEMV
131 Index r = IsLower ? rows - pi - actualPanelWidth : pi;
132 if (r > 0) {
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);
137 }
138 }
139 EIGEN_IF_CONSTEXPR (!IsLower) {
140 if (cols > size) {
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);
144 }
145 }
146}
147
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;
152 static constexpr bool HasUnitDiag = (Mode & UnitDiag) == UnitDiag;
153 static constexpr bool HasZeroDiag = (Mode & ZeroDiag) == ZeroDiag;
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);
157};
158
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;
169
170 using LhsMapper = const_blas_data_mapper<LhsScalar, Index, RowMajor>;
171 using RhsMapper = const_blas_data_mapper<RhsScalar, Index, RowMajor>;
172
173 conj_if<ConjLhs> cjl;
174 conj_if<ConjRhs> cjr;
175
176 for (Index pi = 0; pi < diagSize; pi += PanelWidth) {
177 Index actualPanelWidth = (std::min)(PanelWidth, diagSize - pi);
178
179 // Process the triangular panel using raw dot products to eliminate
180 // the cwiseProduct().sum() expression template overhead.
181 for (Index k = 0; k < actualPanelWidth; ++k) {
182 Index i = pi + k;
183 const LhsScalar* EIGEN_RESTRICT row_i = lhs_ + i * lhsStride;
184 ResScalar dot = ResScalar(0);
185
186 EIGEN_IF_CONSTEXPR (IsLower) {
187 Index s = pi;
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]);
190 } else {
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]);
194 }
195 res_[i * resIncr] += alpha * dot;
196 EIGEN_IF_CONSTEXPR (HasUnitDiag) res_[i * resIncr] += alpha * cjr(rhs_[i]);
197 }
198
199 // Rectangular part: delegate to optimized GEMV
200 Index r = IsLower ? pi : cols - pi - actualPanelWidth;
201 if (r > 0) {
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);
206 }
207 }
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);
213 }
214 }
215}
216
217/***************************************************************************
218 * Wrapper to product_triangular_vector
219 ***************************************************************************/
220
221template <int Mode, int StorageOrder>
222struct trmv_selector;
223
224} // end namespace internal
225
226namespace internal {
227
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());
233
234 internal::trmv_selector<Mode, (int(internal::traits<Lhs>::Flags) & RowMajorBit) ? RowMajor : ColMajor>::run(
235 lhs, rhs, dst, alpha);
236 }
237};
238
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());
244
245 Transpose<Dest> dstT(dst);
246 internal::trmv_selector<(Mode & (UnitDiag | ZeroDiag)) | ((Mode & Lower) ? Upper : Lower),
247 (int(internal::traits<Rhs>::Flags) & RowMajorBit) ? ColMajor
248 : RowMajor>::run(rhs.transpose(),
249 lhs.transpose(), dstT,
250 alpha);
251 }
252};
253
254} // end namespace internal
255
256namespace internal {
257
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) {
261 EIGEN_IF_CONSTEXPR ((Mode & UnitDiag) == UnitDiag) {
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);
265 }
266 }
267}
268
269template <int Mode>
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;
276
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);
283
284 LhsScalar lhs_alpha = LhsBlasTraits::extractScalarFactor(lhs);
285 RhsScalar rhs_alpha = RhsBlasTraits::extractScalarFactor(rhs);
286 ResScalar actualAlpha = alpha * lhs_alpha * rhs_alpha;
287
288 // FIXME find a way to allow an inner stride on the result if packet_traits<Scalar>::size==1
289 // On the other hand, it is good for the cache to pack the vector anyways...
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;
293
294 gemv_static_vector_if<ResScalar, Dest::SizeAtCompileTime, Dest::MaxSizeAtCompileTime, MightCannotUseDest>
295 static_dest;
296
297 gemv_destination_policy<RhsScalar, ResScalar, EvalToDestAtCompileTime, ComplexByReal> destPolicy(actualAlpha);
298
299 ei_declare_aligned_stack_constructed_variable(ResScalar, actualDestPtr, dest.size(),
300 destPolicy.eval_to_dest() ? dest.data() : static_dest.data());
301
302 destPolicy.prepare(dest, actualDestPtr);
303
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());
310
311 destPolicy.copy_back(dest, actualDestPtr);
312
313 trmv_correct_unit_diagonal<Mode>(lhs, rhs, dest, lhs_alpha);
314 }
315};
316
317template <int Mode>
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;
324
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>;
330
331 std::add_const_t<ActualLhsType> actualLhs = LhsBlasTraits::extract(lhs);
332 std::add_const_t<ActualRhsType> actualRhs = RhsBlasTraits::extract(rhs);
333
334 LhsScalar lhs_alpha = LhsBlasTraits::extractScalarFactor(lhs);
335 RhsScalar rhs_alpha = RhsBlasTraits::extractScalarFactor(rhs);
336 ResScalar actualAlpha = alpha * lhs_alpha * rhs_alpha;
337
338 constexpr bool DirectlyUseRhs = ActualRhsTypeCleaned::InnerStrideAtCompileTime == 1;
339
340 gemv_static_vector_if<RhsScalar, ActualRhsTypeCleaned::SizeAtCompileTime,
341 ActualRhsTypeCleaned::MaxSizeAtCompileTime, !DirectlyUseRhs>
342 static_rhs;
343
344 ei_declare_aligned_stack_constructed_variable(
345 RhsScalar, actualRhsPtr, actualRhs.size(),
346 DirectlyUseRhs ? const_cast<RhsScalar*>(actualRhs.data()) : static_rhs.data());
347
348 gemv_prepare_rhs<DirectlyUseRhs>(actualRhs, actualRhsPtr);
349
350 internal::triangular_matrix_vector_product<Index, Mode, LhsScalar, LhsBlasTraits::NeedToConjugate, RhsScalar,
351 RhsBlasTraits::NeedToConjugate, RowMajor>::run(actualLhs.rows(),
352 actualLhs.cols(),
353 actualLhs.data(),
354 actualLhs.outerStride(),
355 actualRhsPtr, 1,
356 dest.data(),
357 dest.innerStride(),
358 actualAlpha);
359
360 trmv_correct_unit_diagonal<Mode>(lhs, rhs, dest, lhs_alpha);
361 }
362};
363
364} // end namespace internal
365
366} // end namespace Eigen
367
368#endif // EIGEN_TRIANGULARMATRIXVECTOR_H
@ 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