Eigen  5.0.1
 
Loading...
Searching...
No Matches
TriangularMatrixMatrix.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_TRIANGULAR_MATRIX_MATRIX_H
12#define EIGEN_TRIANGULAR_MATRIX_MATRIX_H
13
14// IWYU pragma: private
15#include "../InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20
21/* Optimized triangular matrix * matrix (_TRMM++) product built on top of
22 * the general matrix matrix product.
23 */
24template <typename Scalar, typename Index, int Mode, bool LhsIsTriangular, int LhsStorageOrder, bool ConjugateLhs,
25 int RhsStorageOrder, bool ConjugateRhs, int ResStorageOrder, int ResInnerStride, int Version = Specialized>
26struct product_triangular_matrix_matrix;
27
28template <typename Scalar, typename Index, int Mode, bool LhsIsTriangular, int LhsStorageOrder, bool ConjugateLhs,
29 int RhsStorageOrder, bool ConjugateRhs, int ResInnerStride, int Version>
30struct product_triangular_matrix_matrix<Scalar, Index, Mode, LhsIsTriangular, LhsStorageOrder, ConjugateLhs,
31 RhsStorageOrder, ConjugateRhs, RowMajor, ResInnerStride, Version> {
32 static EIGEN_STRONG_INLINE void run(Index rows, Index cols, Index depth, const Scalar* lhs, Index lhsStride,
33 const Scalar* rhs, Index rhsStride, Scalar* res, Index resIncr, Index resStride,
34 const Scalar& alpha, level3_blocking<Scalar, Scalar>& blocking) {
35 product_triangular_matrix_matrix<Scalar, Index, (Mode & (UnitDiag | ZeroDiag)) | ((Mode & Upper) ? Lower : Upper),
36 (!LhsIsTriangular), RhsStorageOrder == RowMajor ? ColMajor : RowMajor,
37 ConjugateRhs, LhsStorageOrder == RowMajor ? ColMajor : RowMajor, ConjugateLhs,
38 ColMajor, ResInnerStride>::run(cols, rows, depth, rhs, rhsStride, lhs, lhsStride,
39 res, resIncr, resStride, alpha, blocking);
40 }
41};
42
43// implements col-major += alpha * op(triangular) * op(general)
44template <typename Scalar, typename Index, int Mode, int LhsStorageOrder, bool ConjugateLhs, int RhsStorageOrder,
45 bool ConjugateRhs, int ResInnerStride, int Version>
46struct product_triangular_matrix_matrix<Scalar, Index, Mode, true, LhsStorageOrder, ConjugateLhs, RhsStorageOrder,
47 ConjugateRhs, ColMajor, ResInnerStride, Version> {
48 using Traits = gebp_traits<Scalar, Scalar>;
49 enum {
50 SmallPanelWidth = 2 * plain_enum_max(Traits::mr, Traits::nr),
51 IsLower = (Mode & Lower) == Lower,
52 SetDiag = (Mode & (ZeroDiag | UnitDiag)) ? 0 : 1
53 };
54
55 static EIGEN_DONT_INLINE void run(Index _rows, Index _cols, Index _depth, const Scalar* lhs_, Index lhsStride,
56 const Scalar* rhs_, Index rhsStride, Scalar* res, Index resIncr, Index resStride,
57 const Scalar& alpha, level3_blocking<Scalar, Scalar>& blocking);
58};
59
60template <typename Scalar, typename Index, int Mode, int LhsStorageOrder, bool ConjugateLhs, int RhsStorageOrder,
61 bool ConjugateRhs, int ResInnerStride, int Version>
62EIGEN_DONT_INLINE void product_triangular_matrix_matrix<
63 Scalar, Index, Mode, true, LhsStorageOrder, ConjugateLhs, RhsStorageOrder, ConjugateRhs, ColMajor, ResInnerStride,
64 Version>::run(Index _rows, Index _cols, Index _depth, const Scalar* lhs_, Index lhsStride, const Scalar* rhs_,
65 Index rhsStride, Scalar* res_, Index resIncr, Index resStride, const Scalar& alpha,
66 level3_blocking<Scalar, Scalar>& blocking) {
67 // strip zeros
68 Index diagSize = (std::min)(_rows, _depth);
69 Index rows = IsLower ? _rows : diagSize;
70 Index depth = IsLower ? diagSize : _depth;
71 Index cols = _cols;
72
73 using LhsMapper = const_blas_data_mapper<Scalar, Index, LhsStorageOrder>;
74 using RhsMapper = const_blas_data_mapper<Scalar, Index, RhsStorageOrder>;
75 using ResMapper = blas_data_mapper<typename Traits::ResScalar, Index, ColMajor, Unaligned, ResInnerStride>;
76 LhsMapper lhs(lhs_, lhsStride);
77 RhsMapper rhs(rhs_, rhsStride);
78 ResMapper res(res_, resStride, resIncr);
79
80 Index kc = blocking.kc(); // cache block size along the K direction
81 Index mc = (std::min)(rows, blocking.mc()); // cache block size along the M direction
82 // The small panel size must not be larger than blocking size.
83 // Usually this should never be the case because SmallPanelWidth^2 is very small
84 // compared to L2 cache size, but let's be safe:
85 Index panelWidth = (std::min)(Index(SmallPanelWidth), (std::min)(kc, mc));
86
87 std::size_t sizeA = kc * mc;
88 std::size_t sizeB = kc * cols;
89
90 ei_declare_aligned_stack_constructed_variable(Scalar, blockA, sizeA, blocking.blockA());
91 ei_declare_aligned_stack_constructed_variable(Scalar, blockB, sizeB, blocking.blockB());
92
93 Matrix<Scalar, SmallPanelWidth, SmallPanelWidth, LhsStorageOrder> triangularBuffer;
94 triangularBuffer.setZero();
95 EIGEN_IF_CONSTEXPR ((Mode & ZeroDiag) == ZeroDiag)
96 triangularBuffer.diagonal().setZero();
97 else
98 triangularBuffer.diagonal().setOnes();
99
100 gebp_kernel<Scalar, Scalar, Index, ResMapper, Traits::mr, Traits::nr, ConjugateLhs, ConjugateRhs> gebp_kernel;
101 gemm_pack_lhs<Scalar, Index, LhsMapper, Traits::mr, Traits::LhsProgress, typename Traits::LhsPacket4Packing,
102 LhsStorageOrder>
103 pack_lhs;
104 gemm_pack_rhs<Scalar, Index, RhsMapper, Traits::nr, RhsStorageOrder> pack_rhs;
105
106 for (Index k2 = IsLower ? depth : 0; IsLower ? k2 > 0 : k2 < depth; IsLower ? k2 -= kc : k2 += kc) {
107 Index actual_kc = (std::min)(IsLower ? k2 : depth - k2, kc);
108 Index actual_k2 = IsLower ? k2 - actual_kc : k2;
109
110 // align blocks with the end of the triangular part for trapezoidal lhs
111 EIGEN_IF_CONSTEXPR (!IsLower) {
112 if ((k2 < rows) && (k2 + actual_kc > rows)) {
113 actual_kc = rows - k2;
114 k2 = k2 + actual_kc - kc;
115 }
116 }
117
118 pack_rhs(blockB, rhs.getSubMapper(actual_k2, 0), actual_kc, cols);
119
120 // the selected lhs's panel has to be split in three different parts:
121 // 1 - the part which is zero => skip it
122 // 2 - the diagonal block => special kernel
123 // 3 - the dense panel below (lower case) or above (upper case) the diagonal block => GEPP
124
125 // the block diagonal, if any:
126 if (IsLower || actual_k2 < rows) {
127 // for each small vertical panels of lhs
128 for (Index k1 = 0; k1 < actual_kc; k1 += panelWidth) {
129 Index actualPanelWidth = std::min<Index>(actual_kc - k1, panelWidth);
130 Index lengthTarget = IsLower ? actual_kc - k1 - actualPanelWidth : k1;
131 Index startBlock = actual_k2 + k1;
132 Index blockBOffset = k1;
133
134 // => GEBP with the micro triangular block
135 // The trick is to pack this micro block while filling the opposite triangular part with zeros.
136 // To this end we do an extra triangular copy to a small temporary buffer
137 for (Index k = 0; k < actualPanelWidth; ++k) {
138 EIGEN_IF_CONSTEXPR (SetDiag) triangularBuffer.coeffRef(k, k) = lhs(startBlock + k, startBlock + k);
139 for (Index i = IsLower ? k + 1 : 0; IsLower ? i < actualPanelWidth : i < k; ++i)
140 triangularBuffer.coeffRef(i, k) = lhs(startBlock + i, startBlock + k);
141 }
142 pack_lhs(blockA, LhsMapper(triangularBuffer.data(), triangularBuffer.outerStride()), actualPanelWidth,
143 actualPanelWidth);
144
145 gebp_kernel(res.getSubMapper(startBlock, 0), blockA, blockB, actualPanelWidth, actualPanelWidth, cols, alpha,
146 actualPanelWidth, actual_kc, 0, blockBOffset);
147
148 // GEBP with remaining micro panel
149 if (lengthTarget > 0) {
150 Index startTarget = IsLower ? actual_k2 + k1 + actualPanelWidth : actual_k2;
151
152 pack_lhs(blockA, lhs.getSubMapper(startTarget, startBlock), actualPanelWidth, lengthTarget);
153
154 gebp_kernel(res.getSubMapper(startTarget, 0), blockA, blockB, lengthTarget, actualPanelWidth, cols, alpha,
155 actualPanelWidth, actual_kc, 0, blockBOffset);
156 }
157 }
158 }
159 // the part below (lower case) or above (upper case) the diagonal => GEPP
160 {
161 Index start = IsLower ? k2 : 0;
162 Index end = IsLower ? rows : (std::min)(actual_k2, rows);
163 for (Index i2 = start; i2 < end; i2 += mc) {
164 const Index actual_mc = (std::min)(i2 + mc, end) - i2;
165 gemm_pack_lhs<Scalar, Index, LhsMapper, Traits::mr, Traits::LhsProgress, typename Traits::LhsPacket4Packing,
166 LhsStorageOrder, false>()(blockA, lhs.getSubMapper(i2, actual_k2), actual_kc, actual_mc);
167
168 gebp_kernel(res.getSubMapper(i2, 0), blockA, blockB, actual_mc, actual_kc, cols, alpha, -1, -1, 0, 0);
169 }
170 }
171 }
172}
173
174// implements col-major += alpha * op(general) * op(triangular)
175template <typename Scalar, typename Index, int Mode, int LhsStorageOrder, bool ConjugateLhs, int RhsStorageOrder,
176 bool ConjugateRhs, int ResInnerStride, int Version>
177struct product_triangular_matrix_matrix<Scalar, Index, Mode, false, LhsStorageOrder, ConjugateLhs, RhsStorageOrder,
178 ConjugateRhs, ColMajor, ResInnerStride, Version> {
179 using Traits = gebp_traits<Scalar, Scalar>;
180 enum {
181 SmallPanelWidth = plain_enum_max(Traits::mr, Traits::nr),
182 IsLower = (Mode & Lower) == Lower,
183 SetDiag = (Mode & (ZeroDiag | UnitDiag)) ? 0 : 1
184 };
185
186 static EIGEN_DONT_INLINE void run(Index _rows, Index _cols, Index _depth, const Scalar* lhs_, Index lhsStride,
187 const Scalar* rhs_, Index rhsStride, Scalar* res, Index resIncr, Index resStride,
188 const Scalar& alpha, level3_blocking<Scalar, Scalar>& blocking);
189};
190
191template <typename Scalar, typename Index, int Mode, int LhsStorageOrder, bool ConjugateLhs, int RhsStorageOrder,
192 bool ConjugateRhs, int ResInnerStride, int Version>
193EIGEN_DONT_INLINE void product_triangular_matrix_matrix<
194 Scalar, Index, Mode, false, LhsStorageOrder, ConjugateLhs, RhsStorageOrder, ConjugateRhs, ColMajor, ResInnerStride,
195 Version>::run(Index _rows, Index _cols, Index _depth, const Scalar* lhs_, Index lhsStride, const Scalar* rhs_,
196 Index rhsStride, Scalar* res_, Index resIncr, Index resStride, const Scalar& alpha,
197 level3_blocking<Scalar, Scalar>& blocking) {
198 const Index PacketBytes = packet_traits<Scalar>::size * sizeof(Scalar);
199 // strip zeros
200 Index diagSize = (std::min)(_cols, _depth);
201 Index rows = _rows;
202 Index depth = IsLower ? _depth : diagSize;
203 Index cols = IsLower ? diagSize : _cols;
204
205 using LhsMapper = const_blas_data_mapper<Scalar, Index, LhsStorageOrder>;
206 using RhsMapper = const_blas_data_mapper<Scalar, Index, RhsStorageOrder>;
207 using ResMapper = blas_data_mapper<typename Traits::ResScalar, Index, ColMajor, Unaligned, ResInnerStride>;
208 LhsMapper lhs(lhs_, lhsStride);
209 RhsMapper rhs(rhs_, rhsStride);
210 ResMapper res(res_, resStride, resIncr);
211
212 Index kc = blocking.kc(); // cache block size along the K direction
213 Index mc = (std::min)(rows, blocking.mc()); // cache block size along the M direction
214
215 std::size_t sizeA = kc * mc;
216 std::size_t sizeB = kc * cols + EIGEN_MAX_ALIGN_BYTES / sizeof(Scalar);
217
218 ei_declare_aligned_stack_constructed_variable(Scalar, blockA, sizeA, blocking.blockA());
219 ei_declare_aligned_stack_constructed_variable(Scalar, blockB, sizeB, blocking.blockB());
220
221 Matrix<Scalar, SmallPanelWidth, SmallPanelWidth, RhsStorageOrder> triangularBuffer;
222 triangularBuffer.setZero();
223 EIGEN_IF_CONSTEXPR ((Mode & ZeroDiag) == ZeroDiag)
224 triangularBuffer.diagonal().setZero();
225 else
226 triangularBuffer.diagonal().setOnes();
227
228 gebp_kernel<Scalar, Scalar, Index, ResMapper, Traits::mr, Traits::nr, ConjugateLhs, ConjugateRhs> gebp_kernel;
229 gemm_pack_lhs<Scalar, Index, LhsMapper, Traits::mr, Traits::LhsProgress, typename Traits::LhsPacket4Packing,
230 LhsStorageOrder>
231 pack_lhs;
232 gemm_pack_rhs<Scalar, Index, RhsMapper, Traits::nr, RhsStorageOrder> pack_rhs;
233 gemm_pack_rhs<Scalar, Index, RhsMapper, Traits::nr, RhsStorageOrder, false, true> pack_rhs_panel;
234
235 for (Index k2 = IsLower ? 0 : depth; IsLower ? k2 < depth : k2 > 0; IsLower ? k2 += kc : k2 -= kc) {
236 Index actual_kc = (std::min)(IsLower ? depth - k2 : k2, kc);
237 Index actual_k2 = IsLower ? k2 : k2 - actual_kc;
238
239 // align blocks with the end of the triangular part for trapezoidal rhs
240 EIGEN_IF_CONSTEXPR (IsLower) {
241 if ((k2 < cols) && (actual_k2 + actual_kc > cols)) {
242 actual_kc = cols - k2;
243 k2 = actual_k2 + actual_kc - kc;
244 }
245 }
246
247 // remaining size
248 Index rs = IsLower ? (std::min)(cols, actual_k2) : cols - k2;
249 // size of the triangular part
250 Index ts = (IsLower && actual_k2 >= cols) ? 0 : actual_kc;
251
252 Scalar* geb = blockB + ts * ts;
253 geb = geb + internal::first_aligned<PacketBytes>(geb, PacketBytes / sizeof(Scalar));
254
255 pack_rhs(geb, rhs.getSubMapper(actual_k2, IsLower ? 0 : k2), actual_kc, rs);
256
257 // pack the triangular part of the rhs padding the unrolled blocks with zeros
258 if (ts > 0) {
259 for (Index j2 = 0; j2 < actual_kc; j2 += SmallPanelWidth) {
260 Index actualPanelWidth = std::min<Index>(actual_kc - j2, SmallPanelWidth);
261 Index actual_j2 = actual_k2 + j2;
262 Index panelOffset = IsLower ? j2 + actualPanelWidth : 0;
263 Index panelLength = IsLower ? actual_kc - j2 - actualPanelWidth : j2;
264 // general part
265 pack_rhs_panel(blockB + j2 * actual_kc, rhs.getSubMapper(actual_k2 + panelOffset, actual_j2), panelLength,
266 actualPanelWidth, actual_kc, panelOffset);
267
268 // append the triangular part via a temporary buffer
269 for (Index j = 0; j < actualPanelWidth; ++j) {
270 EIGEN_IF_CONSTEXPR (SetDiag) triangularBuffer.coeffRef(j, j) = rhs(actual_j2 + j, actual_j2 + j);
271 for (Index k = IsLower ? j + 1 : 0; IsLower ? k < actualPanelWidth : k < j; ++k)
272 triangularBuffer.coeffRef(k, j) = rhs(actual_j2 + k, actual_j2 + j);
273 }
274
275 pack_rhs_panel(blockB + j2 * actual_kc, RhsMapper(triangularBuffer.data(), triangularBuffer.outerStride()),
276 actualPanelWidth, actualPanelWidth, actual_kc, j2);
277 }
278 }
279
280 for (Index i2 = 0; i2 < rows; i2 += mc) {
281 const Index actual_mc = (std::min)(mc, rows - i2);
282 pack_lhs(blockA, lhs.getSubMapper(i2, actual_k2), actual_kc, actual_mc);
283
284 // triangular kernel
285 if (ts > 0) {
286 for (Index j2 = 0; j2 < actual_kc; j2 += SmallPanelWidth) {
287 Index actualPanelWidth = std::min<Index>(actual_kc - j2, SmallPanelWidth);
288 Index panelLength = IsLower ? actual_kc - j2 : j2 + actualPanelWidth;
289 Index blockOffset = IsLower ? j2 : 0;
290
291 gebp_kernel(res.getSubMapper(i2, actual_k2 + j2), blockA, blockB + j2 * actual_kc, actual_mc, panelLength,
292 actualPanelWidth, alpha, actual_kc, actual_kc, // strides
293 blockOffset, blockOffset); // offsets
294 }
295 }
296 gebp_kernel(res.getSubMapper(i2, IsLower ? 0 : k2), blockA, geb, actual_mc, actual_kc, rs, alpha, -1, -1, 0, 0);
297 }
298 }
299}
300
301/***************************************************************************
302 * Wrapper to product_triangular_matrix_matrix
303 ***************************************************************************/
304
305} // end namespace internal
306
307namespace internal {
308template <int Mode, bool LhsIsTriangular, typename Lhs, typename Rhs>
309struct triangular_product_impl<Mode, LhsIsTriangular, Lhs, false, Rhs, false> {
310 template <typename Dest>
311 static void run(Dest& dst, const Lhs& a_lhs, const Rhs& a_rhs, const typename Dest::Scalar& alpha) {
312 using LhsScalar = typename Lhs::Scalar;
313 using RhsScalar = typename Rhs::Scalar;
314 using Scalar = typename Dest::Scalar;
315
316 using LhsBlasTraits = internal::blas_traits<Lhs>;
317 using ActualLhsType = typename LhsBlasTraits::DirectLinearAccessType;
318 using ActualLhsTypeCleaned = internal::remove_all_t<ActualLhsType>;
319 using RhsBlasTraits = internal::blas_traits<Rhs>;
320 using ActualRhsType = typename RhsBlasTraits::DirectLinearAccessType;
321 using ActualRhsTypeCleaned = internal::remove_all_t<ActualRhsType>;
322
323 internal::add_const_on_value_type_t<ActualLhsType> lhs = LhsBlasTraits::extract(a_lhs);
324 internal::add_const_on_value_type_t<ActualRhsType> rhs = RhsBlasTraits::extract(a_rhs);
325
326 // Empty product, return early. Otherwise, we get `nullptr` use errors below when we try to access
327 // coeffRef(0,0).
328 if (lhs.size() == 0 || rhs.size() == 0) {
329 return;
330 }
331
332 LhsScalar lhs_alpha = LhsBlasTraits::extractScalarFactor(a_lhs);
333 RhsScalar rhs_alpha = RhsBlasTraits::extractScalarFactor(a_rhs);
334 Scalar actualAlpha = alpha * lhs_alpha * rhs_alpha;
335
336 using BlockingType = internal::gemm_blocking_space<(Dest::Flags & RowMajorBit) ? RowMajor : ColMajor, Scalar,
337 Scalar, Lhs::MaxRowsAtCompileTime, Rhs::MaxColsAtCompileTime,
338 Lhs::MaxColsAtCompileTime, 4>;
339
340 enum { IsLower = (Mode & Lower) == Lower };
341 Index stripedRows = ((!LhsIsTriangular) || (IsLower)) ? lhs.rows() : (std::min)(lhs.rows(), lhs.cols());
342 Index stripedCols = ((LhsIsTriangular) || (!IsLower)) ? rhs.cols() : (std::min)(rhs.cols(), rhs.rows());
343 Index stripedDepth = LhsIsTriangular ? ((!IsLower) ? lhs.cols() : (std::min)(lhs.cols(), lhs.rows()))
344 : ((IsLower) ? rhs.rows() : (std::min)(rhs.rows(), rhs.cols()));
345
346 BlockingType blocking(stripedRows, stripedCols, stripedDepth, 1, false);
347
348 internal::product_triangular_matrix_matrix<
349 Scalar, Index, Mode, LhsIsTriangular,
350 (internal::traits<ActualLhsTypeCleaned>::Flags & RowMajorBit) ? RowMajor : ColMajor,
351 LhsBlasTraits::NeedToConjugate,
352 (internal::traits<ActualRhsTypeCleaned>::Flags & RowMajorBit) ? RowMajor : ColMajor,
353 RhsBlasTraits::NeedToConjugate, (internal::traits<Dest>::Flags & RowMajorBit) ? RowMajor : ColMajor,
354 Dest::InnerStrideAtCompileTime>::run(stripedRows, stripedCols, stripedDepth, // sizes
355 &lhs.coeffRef(0, 0), lhs.outerStride(), // lhs info
356 &rhs.coeffRef(0, 0), rhs.outerStride(), // rhs info
357 &dst.coeffRef(0, 0), dst.innerStride(), dst.outerStride(), // result info
358 actualAlpha, blocking);
359
360 // Apply correction if the diagonal is unit and a scalar factor was nested:
361 EIGEN_IF_CONSTEXPR ((Mode & UnitDiag) == UnitDiag) {
362 EIGEN_IF_CONSTEXPR (LhsIsTriangular) {
363 if (!numext::is_exactly_one(lhs_alpha)) {
364 Index diagSize = (std::min)(lhs.rows(), lhs.cols());
365 dst.topRows(diagSize) -= ((lhs_alpha - LhsScalar(1)) * a_rhs).topRows(diagSize);
366 }
367 } else {
368 if (!numext::is_exactly_one(rhs_alpha)) {
369 Index diagSize = (std::min)(rhs.rows(), rhs.cols());
370 dst.leftCols(diagSize) -= (rhs_alpha - RhsScalar(1)) * a_lhs.leftCols(diagSize);
371 }
372 }
373 }
374 }
375};
376
377} // end namespace internal
378
379} // end namespace Eigen
380
381#endif // EIGEN_TRIANGULAR_MATRIX_MATRIX_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