11#ifndef EIGEN_TRIANGULAR_MATRIX_MATRIX_H
12#define EIGEN_TRIANGULAR_MATRIX_MATRIX_H
15#include "../InternalHeaderCheck.h"
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;
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) {
38 ColMajor, ResInnerStride>::run(cols, rows, depth, rhs, rhsStride, lhs, lhsStride,
39 res, resIncr, resStride, alpha, blocking);
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>;
50 SmallPanelWidth = 2 * plain_enum_max(Traits::mr, Traits::nr),
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);
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) {
68 Index diagSize = (std::min)(_rows, _depth);
69 Index rows = IsLower ? _rows : diagSize;
70 Index depth = IsLower ? diagSize : _depth;
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);
80 Index kc = blocking.kc();
81 Index mc = (std::min)(rows, blocking.mc());
85 Index panelWidth = (std::min)(Index(SmallPanelWidth), (std::min)(kc, mc));
87 std::size_t sizeA = kc * mc;
88 std::size_t sizeB = kc * cols;
90 ei_declare_aligned_stack_constructed_variable(Scalar, blockA, sizeA, blocking.blockA());
91 ei_declare_aligned_stack_constructed_variable(Scalar, blockB, sizeB, blocking.blockB());
93 Matrix<Scalar, SmallPanelWidth, SmallPanelWidth, LhsStorageOrder> triangularBuffer;
94 triangularBuffer.setZero();
96 triangularBuffer.diagonal().setZero();
98 triangularBuffer.diagonal().setOnes();
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,
104 gemm_pack_rhs<Scalar, Index, RhsMapper, Traits::nr, RhsStorageOrder> pack_rhs;
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;
111 EIGEN_IF_CONSTEXPR (!IsLower) {
112 if ((k2 < rows) && (k2 + actual_kc > rows)) {
113 actual_kc = rows - k2;
114 k2 = k2 + actual_kc - kc;
118 pack_rhs(blockB, rhs.getSubMapper(actual_k2, 0), actual_kc, cols);
126 if (IsLower || actual_k2 < rows) {
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;
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);
142 pack_lhs(blockA, LhsMapper(triangularBuffer.data(), triangularBuffer.outerStride()), actualPanelWidth,
145 gebp_kernel(res.getSubMapper(startBlock, 0), blockA, blockB, actualPanelWidth, actualPanelWidth, cols, alpha,
146 actualPanelWidth, actual_kc, 0, blockBOffset);
149 if (lengthTarget > 0) {
150 Index startTarget = IsLower ? actual_k2 + k1 + actualPanelWidth : actual_k2;
152 pack_lhs(blockA, lhs.getSubMapper(startTarget, startBlock), actualPanelWidth, lengthTarget);
154 gebp_kernel(res.getSubMapper(startTarget, 0), blockA, blockB, lengthTarget, actualPanelWidth, cols, alpha,
155 actualPanelWidth, actual_kc, 0, blockBOffset);
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);
168 gebp_kernel(res.getSubMapper(i2, 0), blockA, blockB, actual_mc, actual_kc, cols, alpha, -1, -1, 0, 0);
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>;
181 SmallPanelWidth = plain_enum_max(Traits::mr, Traits::nr),
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);
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);
200 Index diagSize = (std::min)(_cols, _depth);
202 Index depth = IsLower ? _depth : diagSize;
203 Index cols = IsLower ? diagSize : _cols;
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);
212 Index kc = blocking.kc();
213 Index mc = (std::min)(rows, blocking.mc());
215 std::size_t sizeA = kc * mc;
216 std::size_t sizeB = kc * cols + EIGEN_MAX_ALIGN_BYTES /
sizeof(Scalar);
218 ei_declare_aligned_stack_constructed_variable(Scalar, blockA, sizeA, blocking.blockA());
219 ei_declare_aligned_stack_constructed_variable(Scalar, blockB, sizeB, blocking.blockB());
221 Matrix<Scalar, SmallPanelWidth, SmallPanelWidth, RhsStorageOrder> triangularBuffer;
222 triangularBuffer.setZero();
224 triangularBuffer.diagonal().setZero();
226 triangularBuffer.diagonal().setOnes();
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,
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;
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;
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;
248 Index rs = IsLower ? (std::min)(cols, actual_k2) : cols - k2;
250 Index ts = (IsLower && actual_k2 >= cols) ? 0 : actual_kc;
252 Scalar* geb = blockB + ts * ts;
253 geb = geb + internal::first_aligned<PacketBytes>(geb, PacketBytes /
sizeof(Scalar));
255 pack_rhs(geb, rhs.getSubMapper(actual_k2, IsLower ? 0 : k2), actual_kc, rs);
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;
265 pack_rhs_panel(blockB + j2 * actual_kc, rhs.getSubMapper(actual_k2 + panelOffset, actual_j2), panelLength,
266 actualPanelWidth, actual_kc, panelOffset);
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);
275 pack_rhs_panel(blockB + j2 * actual_kc, RhsMapper(triangularBuffer.data(), triangularBuffer.outerStride()),
276 actualPanelWidth, actualPanelWidth, actual_kc, j2);
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);
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;
291 gebp_kernel(res.getSubMapper(i2, actual_k2 + j2), blockA, blockB + j2 * actual_kc, actual_mc, panelLength,
292 actualPanelWidth, alpha, actual_kc, actual_kc,
293 blockOffset, blockOffset);
296 gebp_kernel(res.getSubMapper(i2, IsLower ? 0 : k2), blockA, geb, actual_mc, actual_kc, rs, alpha, -1, -1, 0, 0);
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;
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>;
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);
328 if (lhs.size() == 0 || rhs.size() == 0) {
332 LhsScalar lhs_alpha = LhsBlasTraits::extractScalarFactor(a_lhs);
333 RhsScalar rhs_alpha = RhsBlasTraits::extractScalarFactor(a_rhs);
334 Scalar actualAlpha = alpha * lhs_alpha * rhs_alpha;
337 Scalar, Lhs::MaxRowsAtCompileTime, Rhs::MaxColsAtCompileTime,
338 Lhs::MaxColsAtCompileTime, 4>;
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()));
346 BlockingType blocking(stripedRows, stripedCols, stripedDepth, 1,
false);
348 internal::product_triangular_matrix_matrix<
349 Scalar, Index, Mode, LhsIsTriangular,
351 LhsBlasTraits::NeedToConjugate,
354 Dest::InnerStrideAtCompileTime>::run(stripedRows, stripedCols, stripedDepth,
355 &lhs.coeffRef(0, 0), lhs.outerStride(),
356 &rhs.coeffRef(0, 0), rhs.outerStride(),
357 &dst.coeffRef(0, 0), dst.innerStride(), dst.outerStride(),
358 actualAlpha, blocking);
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);
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);
@ 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