Eigen  5.0.1
 
Loading...
Searching...
No Matches
GeneralMatrixMatrixTriangular.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2009-2010 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_GENERAL_MATRIX_MATRIX_TRIANGULAR_H
12#define EIGEN_GENERAL_MATRIX_MATRIX_TRIANGULAR_H
13
14// IWYU pragma: private
15#include "../InternalHeaderCheck.h"
16
17namespace Eigen {
18
19template <typename Scalar, typename Index, int StorageOrder, int UpLo, bool ConjLhs, bool ConjRhs>
20struct selfadjoint_rank1_update;
21
22namespace internal {
23
24/**********************************************************************
25 * This file implements a general A * B product while
26 * evaluating only one triangular part of the product.
27 * This is a more general version of self adjoint product (C += A A^T)
28 * as the level 3 SYRK Blas routine.
29 **********************************************************************/
30
31// forward declarations (defined at the end of this file)
32template <typename LhsScalar, typename RhsScalar, typename Index, int mr, int nr, bool ConjLhs, bool ConjRhs,
33 int ResInnerStride, int UpLo>
34struct tribb_kernel;
35
36/* Optimized matrix-matrix product evaluating only one triangular half */
37template <typename Index, typename LhsScalar, int LhsStorageOrder, bool ConjugateLhs, typename RhsScalar,
38 int RhsStorageOrder, bool ConjugateRhs, int ResStorageOrder, int ResInnerStride, int UpLo,
39 int Version = Specialized>
40struct general_matrix_matrix_triangular_product;
41
42// as usual if the result is row major => we transpose the product
43template <typename Index, typename LhsScalar, int LhsStorageOrder, bool ConjugateLhs, typename RhsScalar,
44 int RhsStorageOrder, bool ConjugateRhs, int ResInnerStride, int UpLo, int Version>
45struct general_matrix_matrix_triangular_product<Index, LhsScalar, LhsStorageOrder, ConjugateLhs, RhsScalar,
46 RhsStorageOrder, ConjugateRhs, RowMajor, ResInnerStride, UpLo,
47 Version> {
48 using ResScalar = typename ScalarBinaryOpTraits<LhsScalar, RhsScalar>::ReturnType;
49 static EIGEN_STRONG_INLINE void run(Index size, Index depth, const LhsScalar* lhs, Index lhsStride,
50 const RhsScalar* rhs, Index rhsStride, ResScalar* res, Index resIncr,
51 Index resStride, const ResScalar& alpha,
52 level3_blocking<RhsScalar, LhsScalar>& blocking) {
53 general_matrix_matrix_triangular_product<Index, RhsScalar, RhsStorageOrder == RowMajor ? ColMajor : RowMajor,
54 ConjugateRhs, LhsScalar, LhsStorageOrder == RowMajor ? ColMajor : RowMajor,
55 ConjugateLhs, ColMajor, ResInnerStride,
56 UpLo == Lower ? Upper : Lower>::run(size, depth, rhs, rhsStride, lhs,
57 lhsStride, res, resIncr, resStride,
58 alpha, blocking);
59 }
60};
61
62template <typename Index, typename LhsScalar, int LhsStorageOrder, bool ConjugateLhs, typename RhsScalar,
63 int RhsStorageOrder, bool ConjugateRhs, int ResInnerStride, int UpLo, int Version>
64struct general_matrix_matrix_triangular_product<Index, LhsScalar, LhsStorageOrder, ConjugateLhs, RhsScalar,
65 RhsStorageOrder, ConjugateRhs, ColMajor, ResInnerStride, UpLo,
66 Version> {
67 using ResScalar = typename ScalarBinaryOpTraits<LhsScalar, RhsScalar>::ReturnType;
68 static EIGEN_STRONG_INLINE void run(Index size, Index depth, const LhsScalar* lhs_, Index lhsStride,
69 const RhsScalar* rhs_, Index rhsStride, ResScalar* res_, Index resIncr,
70 Index resStride, const ResScalar& alpha,
71 level3_blocking<LhsScalar, RhsScalar>& blocking) {
72 if (size == 0) {
73 return;
74 }
75
76 using Traits = gebp_traits<LhsScalar, RhsScalar>;
77
78 using LhsMapper = const_blas_data_mapper<LhsScalar, Index, LhsStorageOrder>;
79 using RhsMapper = const_blas_data_mapper<RhsScalar, Index, RhsStorageOrder>;
80 using ResMapper = blas_data_mapper<typename Traits::ResScalar, Index, ColMajor, Unaligned, ResInnerStride>;
81 LhsMapper lhs(lhs_, lhsStride);
82 RhsMapper rhs(rhs_, rhsStride);
83 ResMapper res(res_, resStride, resIncr);
84
85 Index kc = blocking.kc();
86 // Ensure that mc >= nr and <= size
87 Index mc = (std::min)(size, (std::max)(static_cast<decltype(blocking.mc())>(Traits::nr), blocking.mc()));
88
89 // !!! mc must be a multiple of nr
90 if (mc > Traits::nr) {
91 using UnsignedIndex = std::make_unsigned_t<Index>;
92 mc = (UnsignedIndex(mc) / Traits::nr) * Traits::nr;
93 }
94
95 std::size_t sizeA = kc * mc;
96 std::size_t sizeB = kc * size;
97
98 ei_declare_aligned_stack_constructed_variable(LhsScalar, blockA, sizeA, blocking.blockA());
99 ei_declare_aligned_stack_constructed_variable(RhsScalar, blockB, sizeB, blocking.blockB());
100
101 gemm_pack_lhs<LhsScalar, Index, LhsMapper, Traits::mr, Traits::LhsProgress, typename Traits::LhsPacket4Packing,
102 LhsStorageOrder>
103 pack_lhs;
104 gemm_pack_rhs<RhsScalar, Index, RhsMapper, Traits::nr, RhsStorageOrder> pack_rhs;
105 gebp_kernel<LhsScalar, RhsScalar, Index, ResMapper, Traits::mr, Traits::nr, ConjugateLhs, ConjugateRhs> gebp;
106 tribb_kernel<LhsScalar, RhsScalar, Index, Traits::mr, Traits::nr, ConjugateLhs, ConjugateRhs, ResInnerStride, UpLo>
107 sybb;
108
109 for (Index k2 = 0; k2 < depth; k2 += kc) {
110 const Index actual_kc = (std::min)(k2 + kc, depth) - k2;
111
112 // note that the actual rhs is the transpose/adjoint of mat
113 pack_rhs(blockB, rhs.getSubMapper(k2, 0), actual_kc, size);
114
115 for (Index i2 = 0; i2 < size; i2 += mc) {
116 const Index actual_mc = (std::min)(i2 + mc, size) - i2;
117
118 pack_lhs(blockA, lhs.getSubMapper(i2, k2), actual_kc, actual_mc);
119
120 // the selected actual_mc * size panel of res is split into three different parts:
121 // 1 - before the diagonal => processed with gebp or skipped
122 // 2 - the actual_mc x actual_mc symmetric block => processed with a special kernel
123 // 3 - after the diagonal => processed with gebp or skipped
124 EIGEN_IF_CONSTEXPR (UpLo == Lower) {
125 gebp(res.getSubMapper(i2, 0), blockA, blockB, actual_mc, actual_kc, (std::min)(size, i2), alpha, -1, -1, 0,
126 0);
127 }
128
129 sybb(res_ + resStride * i2 + resIncr * i2, resIncr, resStride, blockA, blockB + actual_kc * i2, actual_mc,
130 actual_kc, alpha);
131
132 EIGEN_IF_CONSTEXPR (UpLo == Upper) {
133 Index j2 = i2 + actual_mc;
134 gebp(res.getSubMapper(i2, j2), blockA, blockB + actual_kc * j2, actual_mc, actual_kc,
135 (std::max)(Index(0), size - j2), alpha, -1, -1, 0, 0);
136 }
137 }
138 }
139 }
140};
141
142// Optimized packed Block * packed Block product kernel evaluating only one given triangular part
143// This kernel is built on top of the gebp kernel:
144// - the current destination block is processed per panel of actual_mc x BlockSize
145// where BlockSize is set to the minimal value allowing gebp to be as fast as possible
146// - then, as usual, each panel is split into three parts along the diagonal,
147// the sub blocks above and below the diagonal are processed as usual,
148// while the triangular block overlapping the diagonal is evaluated into a
149// small temporary buffer which is then accumulated into the result using a
150// triangular traversal.
151template <typename LhsScalar, typename RhsScalar, typename Index, int mr, int nr, bool ConjLhs, bool ConjRhs,
152 int ResInnerStride, int UpLo>
153struct tribb_kernel {
154 using Traits = gebp_traits<LhsScalar, RhsScalar, ConjLhs, ConjRhs>;
155 using ResScalar = typename Traits::ResScalar;
156
157 enum { BlockSize = meta_least_common_multiple<plain_enum_max(mr, nr), plain_enum_min(mr, nr)>::value };
158 void operator()(ResScalar* res_, Index resIncr, Index resStride, const LhsScalar* blockA, const RhsScalar* blockB,
159 Index size, Index depth, const ResScalar& alpha) const {
160 using ResMapper = blas_data_mapper<ResScalar, Index, ColMajor, Unaligned, ResInnerStride>;
161 using BufferMapper = blas_data_mapper<ResScalar, Index, ColMajor, Unaligned>;
162 ResMapper res(res_, resStride, resIncr);
163 gebp_kernel<LhsScalar, RhsScalar, Index, ResMapper, mr, nr, ConjLhs, ConjRhs> gebp_kernel1;
164 gebp_kernel<LhsScalar, RhsScalar, Index, BufferMapper, mr, nr, ConjLhs, ConjRhs> gebp_kernel2;
165
166 Matrix<ResScalar, BlockSize, BlockSize, ColMajor> buffer;
167
168 // let's process the block per panel of actual_mc x BlockSize,
169 // again, each is split into three parts, etc.
170 for (Index j = 0; j < size; j += BlockSize) {
171 Index actualBlockSize = std::min<Index>(BlockSize, size - j);
172 const RhsScalar* actual_b = blockB + j * depth;
173
174 EIGEN_IF_CONSTEXPR (UpLo == Upper) {
175 gebp_kernel1(res.getSubMapper(0, j), blockA, actual_b, j, depth, actualBlockSize, alpha, -1, -1, 0, 0);
176 }
177
178 // selfadjoint micro block
179 {
180 Index i = j;
181 buffer.setZero();
182 // 1 - apply the kernel on the temporary buffer
183 gebp_kernel2(BufferMapper(buffer.data(), BlockSize), blockA + depth * i, actual_b, actualBlockSize, depth,
184 actualBlockSize, alpha, -1, -1, 0, 0);
185
186 // 2 - triangular accumulation
187 for (Index j1 = 0; j1 < actualBlockSize; ++j1) {
188 typename ResMapper::LinearMapper r = res.getLinearMapper(i, j + j1);
189 for (Index i1 = UpLo == Lower ? j1 : 0; UpLo == Lower ? i1 < actualBlockSize : i1 <= j1; ++i1)
190 r(i1) += buffer(i1, j1);
191 }
192 }
193
194 EIGEN_IF_CONSTEXPR (UpLo == Lower) {
195 Index i = j + actualBlockSize;
196 gebp_kernel1(res.getSubMapper(i, j), blockA + depth * i, actual_b, size - i, depth, actualBlockSize, alpha, -1,
197 -1, 0, 0);
198 }
199 }
200 }
201};
202
203} // end namespace internal
204
205// high level API
206
207template <typename MatrixType, typename ProductType, int UpLo, bool IsOuterProduct>
208struct general_product_to_triangular_selector;
209
210template <typename MatrixType, typename ProductType, int UpLo>
211struct general_product_to_triangular_selector<MatrixType, ProductType, UpLo, true> {
212 static void run(MatrixType& mat, const ProductType& prod, const typename MatrixType::Scalar& alpha, bool beta) {
213 using Scalar = typename MatrixType::Scalar;
214
215 using Lhs = internal::remove_all_t<typename ProductType::LhsNested>;
216 using LhsBlasTraits = internal::blas_traits<Lhs>;
217 using ActualLhs = typename LhsBlasTraits::DirectLinearAccessType;
218 using ActualLhs_ = internal::remove_all_t<ActualLhs>;
219 internal::add_const_on_value_type_t<ActualLhs> actualLhs = LhsBlasTraits::extract(prod.lhs());
220
221 using Rhs = internal::remove_all_t<typename ProductType::RhsNested>;
222 using RhsBlasTraits = internal::blas_traits<Rhs>;
223 using ActualRhs = typename RhsBlasTraits::DirectLinearAccessType;
224 using ActualRhs_ = internal::remove_all_t<ActualRhs>;
225 internal::add_const_on_value_type_t<ActualRhs> actualRhs = RhsBlasTraits::extract(prod.rhs());
226
227 Scalar actualAlpha = alpha * LhsBlasTraits::extractScalarFactor(prod.lhs().derived()) *
228 RhsBlasTraits::extractScalarFactor(prod.rhs().derived());
229
230 if (!beta) mat.template triangularView<UpLo>().setZero();
231
232 enum {
233 StorageOrder = (internal::traits<MatrixType>::Flags & RowMajorBit) ? RowMajor : ColMajor,
234 UseLhsDirectly = ActualLhs_::InnerStrideAtCompileTime == 1,
235 UseRhsDirectly = ActualRhs_::InnerStrideAtCompileTime == 1
236 };
237
238 internal::gemv_static_vector_if<Scalar, Lhs::SizeAtCompileTime, Lhs::MaxSizeAtCompileTime, !UseLhsDirectly>
239 static_lhs;
240 ei_declare_aligned_stack_constructed_variable(
241 Scalar, actualLhsPtr, actualLhs.size(),
242 (UseLhsDirectly ? const_cast<Scalar*>(actualLhs.data()) : static_lhs.data()));
243 EIGEN_IF_CONSTEXPR (!UseLhsDirectly) {
244 Map<typename ActualLhs_::PlainObject>(actualLhsPtr, actualLhs.size()) = actualLhs;
245 }
246
247 internal::gemv_static_vector_if<Scalar, Rhs::SizeAtCompileTime, Rhs::MaxSizeAtCompileTime, !UseRhsDirectly>
248 static_rhs;
249 ei_declare_aligned_stack_constructed_variable(
250 Scalar, actualRhsPtr, actualRhs.size(),
251 (UseRhsDirectly ? const_cast<Scalar*>(actualRhs.data()) : static_rhs.data()));
252 EIGEN_IF_CONSTEXPR (!UseRhsDirectly) {
253 Map<typename ActualRhs_::PlainObject>(actualRhsPtr, actualRhs.size()) = actualRhs;
254 }
255
256 selfadjoint_rank1_update<
257 Scalar, Index, StorageOrder, UpLo, LhsBlasTraits::NeedToConjugate && NumTraits<Scalar>::IsComplex,
258 RhsBlasTraits::NeedToConjugate && NumTraits<Scalar>::IsComplex>::run(actualLhs.size(), mat.data(),
259 mat.outerStride(), actualLhsPtr,
260 actualRhsPtr, actualAlpha);
261 }
262};
263
264template <typename MatrixType, typename ProductType, int UpLo>
265struct general_product_to_triangular_selector<MatrixType, ProductType, UpLo, false> {
266 static void run(MatrixType& mat, const ProductType& prod, const typename MatrixType::Scalar& alpha, bool beta) {
267 using Lhs = internal::remove_all_t<typename ProductType::LhsNested>;
268 using LhsBlasTraits = internal::blas_traits<Lhs>;
269 using ActualLhs = typename LhsBlasTraits::DirectLinearAccessType;
270 using ActualLhs_ = internal::remove_all_t<ActualLhs>;
271 internal::add_const_on_value_type_t<ActualLhs> actualLhs = LhsBlasTraits::extract(prod.lhs());
272
273 using Rhs = internal::remove_all_t<typename ProductType::RhsNested>;
274 using RhsBlasTraits = internal::blas_traits<Rhs>;
275 using ActualRhs = typename RhsBlasTraits::DirectLinearAccessType;
276 using ActualRhs_ = internal::remove_all_t<ActualRhs>;
277 internal::add_const_on_value_type_t<ActualRhs> actualRhs = RhsBlasTraits::extract(prod.rhs());
278
279 typename ProductType::Scalar actualAlpha = alpha * LhsBlasTraits::extractScalarFactor(prod.lhs().derived()) *
280 RhsBlasTraits::extractScalarFactor(prod.rhs().derived());
281
282 if (!beta) mat.template triangularView<UpLo>().setZero();
283
284 enum {
285 IsRowMajor = (internal::traits<MatrixType>::Flags & RowMajorBit) ? 1 : 0,
286 LhsIsRowMajor = ActualLhs_::Flags & RowMajorBit ? 1 : 0,
287 RhsIsRowMajor = ActualRhs_::Flags & RowMajorBit ? 1 : 0,
288 SkipDiag = (UpLo & (UnitDiag | ZeroDiag)) != 0
289 };
290
291 Index size = mat.cols();
292 EIGEN_IF_CONSTEXPR (SkipDiag) size--;
293 Index depth = actualLhs.cols();
294 eigen_assert(actualLhs.rows() == mat.rows() && actualRhs.cols() == mat.cols() &&
295 actualLhs.cols() == actualRhs.rows());
296 if (size <= 0 || depth == 0) return;
297
298 using BlockingType =
299 internal::gemm_blocking_space<IsRowMajor ? RowMajor : ColMajor, typename Lhs::Scalar, typename Rhs::Scalar,
300 MatrixType::MaxColsAtCompileTime, MatrixType::MaxColsAtCompileTime,
301 ActualRhs_::MaxColsAtCompileTime>;
302
303 BlockingType blocking(size, size, depth, 1, false);
304
305 internal::general_matrix_matrix_triangular_product<
306 Index, typename Lhs::Scalar, LhsIsRowMajor ? RowMajor : ColMajor, LhsBlasTraits::NeedToConjugate,
307 typename Rhs::Scalar, RhsIsRowMajor ? RowMajor : ColMajor, RhsBlasTraits::NeedToConjugate,
308 IsRowMajor ? RowMajor : ColMajor, MatrixType::InnerStrideAtCompileTime,
309 UpLo&(Lower | Upper)>::run(size, depth, &actualLhs.coeffRef(SkipDiag && (UpLo & Lower) == Lower ? 1 : 0, 0),
310 actualLhs.outerStride(),
311 &actualRhs.coeffRef(0, SkipDiag && (UpLo & Upper) == Upper ? 1 : 0),
312 actualRhs.outerStride(),
313 mat.data() +
314 (SkipDiag ? (bool(IsRowMajor) != ((UpLo & Lower) == Lower) ? mat.innerStride()
315 : mat.outerStride())
316 : 0),
317 mat.innerStride(), mat.outerStride(), actualAlpha, blocking);
318 }
319};
320
321template <typename MatrixType_, unsigned int Mode_>
322template <typename ProductType>
323EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE typename TriangularViewImpl<MatrixType_, Mode_, Dense>::TriangularViewType&
324TriangularViewImpl<MatrixType_, Mode_, Dense>::_assignProduct(
325 const ProductType& prod, const typename TriangularViewImpl<MatrixType_, Mode_, Dense>::Scalar& alpha, bool beta) {
326 EIGEN_STATIC_ASSERT((Mode_ & UnitDiag) == 0, WRITING_TO_TRIANGULAR_PART_WITH_UNIT_DIAGONAL_IS_NOT_SUPPORTED);
327 eigen_assert(derived().nestedExpression().rows() == prod.rows() && derived().cols() == prod.cols());
328
329 general_product_to_triangular_selector<MatrixType_, ProductType, Mode_,
330 internal::traits<ProductType>::InnerSize == 1>::run(derived()
331 .nestedExpression()
332 .const_cast_derived(),
333 prod, alpha, beta);
334
335 return derived();
336}
337
338} // end namespace Eigen
339
340#endif // EIGEN_GENERAL_MATRIX_MATRIX_TRIANGULAR_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