11#ifndef EIGEN_SELFADJOINT_MATRIX_MATRIX_H
12#define EIGEN_SELFADJOINT_MATRIX_MATRIX_H
15#include "../InternalHeaderCheck.h"
29template <
typename Scalar,
typename Index,
int Pack1,
int Pack2,
int StorageOrder>
32 static constexpr int PacketSize = packet_traits<Scalar>::size;
33 using PacketType =
typename packet_traits<Scalar>::type;
34 using Mapper = const_blas_data_mapper<Scalar, Index, StorageOrder>;
35 using TransposedMapper = const_blas_data_mapper<Scalar, Index, TransposedStorageOrder>;
39 using DirectPacker = gemm_pack_lhs<Scalar, Index, Mapper, Pack1, Pack2, PacketType, StorageOrder, false, false>;
40 using MirroredPacker =
41 gemm_pack_lhs<Scalar, Index, TransposedMapper, Pack1, Pack2, PacketType, TransposedStorageOrder, true, false>;
43 void pack_panel(Scalar* blockA,
const Mapper& lhs,
const TransposedMapper& lhs_t, Index cols, Index i, Index bw,
45 Scalar* panel = blockA + count;
46 const Index band_end = i + bw;
47 if (i > 0) DirectPacker()(panel, lhs.getSubMapper(i, 0), i, bw);
50 for (Index k = i; k < band_end; k++) {
51 Scalar* row = panel + k * bw;
52 const Index h = k - i;
53 for (Index w = 0; w < h; w++) row[w] = numext::conj(lhs(k, i + w));
54 row[h] = numext::real(lhs(k, k));
55 for (Index w = h + 1; w < bw; w++) row[w] = lhs(i + w, k);
57 if (cols > band_end) MirroredPacker()(panel + band_end * bw, lhs_t.getSubMapper(i, band_end), cols - band_end, bw);
58 count += bw * numext::maxi(cols, band_end);
61 void operator()(Scalar* blockA,
const Scalar* lhs_, Index lhsStride, Index cols, Index rows)
const {
62 using HalfPacket =
typename unpacket_traits<PacketType>::half;
63 using QuarterPacket =
typename unpacket_traits<HalfPacket>::half;
64 constexpr int HalfPacketSize = unpacket_traits<HalfPacket>::size;
65 constexpr int QuarterPacketSize = unpacket_traits<QuarterPacket>::size;
66 constexpr bool HasHalf = HalfPacketSize < PacketSize;
67 constexpr bool HasQuarter = QuarterPacketSize < HalfPacketSize;
69 const Mapper lhs(lhs_, lhsStride);
70 const TransposedMapper lhs_t(lhs_, lhsStride);
73 const Index peeled_mc3 = Pack1 >= 3 * PacketSize ? (rows / (3 * PacketSize)) * (3 * PacketSize) : 0;
74 const Index peeled_mc2 =
75 Pack1 >= 2 * PacketSize ? peeled_mc3 + ((rows - peeled_mc3) / (2 * PacketSize)) * (2 * PacketSize) : 0;
76 const Index peeled_mc1 =
77 Pack1 >= 1 * PacketSize ? peeled_mc2 + ((rows - peeled_mc2) / (1 * PacketSize)) * (1 * PacketSize) : 0;
78 const Index peeled_mc_half =
79 Pack1 >= HalfPacketSize ? peeled_mc1 + ((rows - peeled_mc1) / (HalfPacketSize)) * (HalfPacketSize) : 0;
80 const Index peeled_mc_quarter =
81 Pack1 >= QuarterPacketSize
82 ? peeled_mc_half + ((rows - peeled_mc_half) / (QuarterPacketSize)) * (QuarterPacketSize)
85 auto pack_rows = [&](Index begin, Index end, Index bw) {
86 for (Index i = begin; i < end; i += bw) pack_panel(blockA, lhs, lhs_t, cols, i, bw, count);
89 EIGEN_IF_CONSTEXPR (Pack1 >= 3 * PacketSize) pack_rows(0, peeled_mc3, 3 * PacketSize);
90 EIGEN_IF_CONSTEXPR (Pack1 >= 2 * PacketSize) pack_rows(peeled_mc3, peeled_mc2, 2 * PacketSize);
91 EIGEN_IF_CONSTEXPR (Pack1 >= 1 * PacketSize) pack_rows(peeled_mc2, peeled_mc1, 1 * PacketSize);
92 EIGEN_IF_CONSTEXPR (HasHalf && Pack1 >= HalfPacketSize) pack_rows(peeled_mc1, peeled_mc_half, HalfPacketSize);
93 EIGEN_IF_CONSTEXPR (HasQuarter && Pack1 >= QuarterPacketSize)
94 pack_rows(peeled_mc_half, peeled_mc_quarter, QuarterPacketSize);
97 pack_rows(peeled_mc_quarter, rows, 1);
101template <
typename Scalar,
typename Index,
int nr,
int StorageOrder>
102struct symm_pack_rhs {
104 using Mapper = const_blas_data_mapper<Scalar, Index, StorageOrder>;
105 using TransposedMapper = const_blas_data_mapper<Scalar, Index, TransposedStorageOrder>;
106 using DirectPacker = gemm_pack_rhs<Scalar, Index, Mapper, nr, StorageOrder, false, false>;
107 using MirroredPacker = gemm_pack_rhs<Scalar, Index, TransposedMapper, nr, TransposedStorageOrder, true, false>;
112 void pack_diagonal_panel(Scalar* blockB,
const Mapper& rhs,
const TransposedMapper& rhs_t, Index k2, Index end_k,
113 Index j2, Index& count)
const {
114 const Index band_end = j2 + Width;
116 MirroredPacker()(blockB + count, rhs_t.getSubMapper(k2, j2), j2 - k2, Width);
117 count += Width * (j2 - k2);
119 for (Index k = j2; k < band_end; k++) {
120 Scalar* row = blockB + count;
121 const Index h = k - j2;
122 for (Index w = 0; w < h; w++) row[w] = rhs(k, j2 + w);
123 row[h] = numext::real(rhs(k, k));
124 for (Index w = h + 1; w < Width; w++) row[w] = numext::conj(rhs(j2 + w, k));
127 if (end_k > band_end) {
128 DirectPacker()(blockB + count, rhs.getSubMapper(band_end, j2), end_k - band_end, Width);
129 count += Width * (end_k - band_end);
133 void operator()(Scalar* blockB,
const Scalar* rhs_, Index rhsStride, Index rows, Index cols, Index k2)
const {
134 const Index end_k = k2 + rows;
135 const Mapper rhs(rhs_, rhsStride);
136 const TransposedMapper rhs_t(rhs_, rhsStride);
138 const Index packet_cols8 = nr >= 8 ? (cols / 8) * 8 : 0;
139 const Index packet_cols4 = nr >= 4 ? (cols / 4) * 4 : 0;
147 eigen_internal_assert(k2 % nr == 0);
148 eigen_internal_assert(end_k % nr == 0 || end_k == cols);
150 DirectPacker()(blockB, rhs.getSubMapper(k2, 0), rows, k2);
155 const Index end8 = nr >= 8 ? numext::mini(end_k, packet_cols8) : k2;
156 const Index end4 = numext::mini(end_k, packet_cols4);
157 EIGEN_IF_CONSTEXPR (nr >= 8) {
158 for (Index j2 = k2; j2 < end8; j2 += 8) pack_diagonal_panel<8>(blockB, rhs, rhs_t, k2, end_k, j2, count);
160 EIGEN_IF_CONSTEXPR (nr >= 4) {
161 for (Index j2 = end8; j2 < end4; j2 += 4) pack_diagonal_panel<4>(blockB, rhs, rhs_t, k2, end_k, j2, count);
165 EIGEN_IF_CONSTEXPR (nr >= 8) {
166 if (packet_cols8 > end_k) {
167 MirroredPacker()(blockB + count, rhs_t.getSubMapper(k2, end_k), rows, packet_cols8 - end_k);
168 count += rows * (packet_cols8 - end_k);
171 EIGEN_IF_CONSTEXPR (nr >= 4) {
172 const Index j3 = numext::maxi(packet_cols8, end_k);
173 if (packet_cols4 > j3) {
174 MirroredPacker()(blockB + count, rhs_t.getSubMapper(k2, j3), rows, packet_cols4 - j3);
175 count += rows * (packet_cols4 - j3);
180 for (Index j2 = packet_cols4; j2 < cols; ++j2) {
182 Index half = numext::mini(end_k, j2);
183 for (Index k = k2; k < half; k++) {
184 blockB[count] = numext::conj(rhs(j2, k));
188 if (half == j2 && half < k2 + rows) {
189 blockB[count] = numext::real(rhs(j2, j2));
195 for (Index k = half + 1; k < k2 + rows; k++) {
196 blockB[count] = rhs(k, j2);
206template <
typename Scalar,
typename Index,
int LhsStorageOrder,
bool LhsSelfAdjoint,
bool ConjugateLhs,
207 int RhsStorageOrder,
bool RhsSelfAdjoint,
bool ConjugateRhs,
int ResStorageOrder,
int ResInnerStride>
208struct product_selfadjoint_matrix;
210template <
typename Scalar,
typename Index,
int LhsStorageOrder,
bool LhsSelfAdjoint,
bool ConjugateLhs,
211 int RhsStorageOrder,
bool RhsSelfAdjoint,
bool ConjugateRhs,
int ResInnerStride>
212struct product_selfadjoint_matrix<Scalar, Index, LhsStorageOrder, LhsSelfAdjoint, ConjugateLhs, RhsStorageOrder,
213 RhsSelfAdjoint, ConjugateRhs,
RowMajor, ResInnerStride> {
214 static EIGEN_STRONG_INLINE
void run(Index rows, Index cols,
const Scalar* lhs, Index lhsStride,
const Scalar* rhs,
215 Index rhsStride, Scalar* res, Index resIncr, Index resStride,
const Scalar& alpha,
216 level3_blocking<Scalar, Scalar>& blocking) {
217 product_selfadjoint_matrix<
219 NumTraits<Scalar>::IsComplex && logical_xor(RhsSelfAdjoint, ConjugateRhs),
221 NumTraits<Scalar>::IsComplex && logical_xor(LhsSelfAdjoint, ConjugateLhs),
ColMajor,
222 ResInnerStride>::run(cols, rows, rhs, rhsStride, lhs, lhsStride, res, resIncr, resStride, alpha, blocking);
226template <
typename Scalar,
typename Index,
int LhsStorageOrder,
bool ConjugateLhs,
int RhsStorageOrder,
227 bool ConjugateRhs,
int ResInnerStride>
228struct product_selfadjoint_matrix<Scalar, Index, LhsStorageOrder, true, ConjugateLhs, RhsStorageOrder, false,
229 ConjugateRhs,
ColMajor, ResInnerStride> {
230 static EIGEN_DONT_INLINE
void run(Index rows, Index cols,
const Scalar* lhs_, Index lhsStride,
const Scalar* rhs_,
231 Index rhsStride, Scalar* res, Index resIncr, Index resStride,
const Scalar& alpha,
232 level3_blocking<Scalar, Scalar>& blocking);
235template <
typename Scalar,
typename Index,
int LhsStorageOrder,
bool ConjugateLhs,
int RhsStorageOrder,
236 bool ConjugateRhs,
int ResInnerStride>
237EIGEN_DONT_INLINE
void
238product_selfadjoint_matrix<Scalar, Index, LhsStorageOrder,
true, ConjugateLhs, RhsStorageOrder,
false, ConjugateRhs,
239 ColMajor, ResInnerStride>::run(Index rows, Index cols,
const Scalar* lhs_, Index lhsStride,
240 const Scalar* rhs_, Index rhsStride, Scalar* res_,
241 Index resIncr, Index resStride,
const Scalar& alpha,
242 level3_blocking<Scalar, Scalar>& blocking) {
245 using Traits = gebp_traits<Scalar, Scalar>;
247 using LhsMapper = const_blas_data_mapper<Scalar, Index, LhsStorageOrder>;
248 using LhsTransposeMapper = const_blas_data_mapper<Scalar, Index, (LhsStorageOrder ==
RowMajor) ?
ColMajor :
RowMajor>;
249 using RhsMapper = const_blas_data_mapper<Scalar, Index, RhsStorageOrder>;
250 using ResMapper = blas_data_mapper<typename Traits::ResScalar, Index, ColMajor, Unaligned, ResInnerStride>;
251 LhsMapper lhs(lhs_, lhsStride);
252 LhsTransposeMapper lhs_transpose(lhs_, lhsStride);
253 RhsMapper rhs(rhs_, rhsStride);
254 ResMapper res(res_, resStride, resIncr);
256 Index kc = blocking.kc();
257 Index mc = (std::min)(rows, blocking.mc());
259 kc = (std::min)(kc, mc);
260 std::size_t sizeA = kc * mc;
261 std::size_t sizeB = kc * cols;
262 ei_declare_aligned_stack_constructed_variable(Scalar, blockA, sizeA, blocking.blockA());
263 ei_declare_aligned_stack_constructed_variable(Scalar, blockB, sizeB, blocking.blockB());
265 gebp_kernel<Scalar, Scalar, Index, ResMapper, Traits::mr, Traits::nr, ConjugateLhs, ConjugateRhs> gebp_kernel;
266 symm_pack_lhs<Scalar, Index, Traits::mr, Traits::LhsProgress, LhsStorageOrder> pack_lhs;
267 gemm_pack_rhs<Scalar, Index, RhsMapper, Traits::nr, RhsStorageOrder> pack_rhs;
268 gemm_pack_lhs<Scalar, Index, LhsTransposeMapper, Traits::mr, Traits::LhsProgress,
typename Traits::LhsPacket4Packing,
272 for (Index k2 = 0; k2 < size; k2 += kc) {
273 const Index actual_kc = (std::min)(k2 + kc, size) - k2;
278 pack_rhs(blockB, rhs.getSubMapper(k2, 0), actual_kc, cols);
284 for (Index i2 = 0; i2 < k2; i2 += mc) {
285 const Index actual_mc = (std::min)(i2 + mc, k2) - i2;
287 pack_lhs_transposed(blockA, lhs_transpose.getSubMapper(i2, k2), actual_kc, actual_mc);
289 gebp_kernel(res.getSubMapper(i2, 0), blockA, blockB, actual_mc, actual_kc, cols, alpha);
293 const Index actual_mc = (std::min)(k2 + kc, size) - k2;
295 pack_lhs(blockA, &lhs(k2, k2), lhsStride, actual_kc, actual_mc);
297 gebp_kernel(res.getSubMapper(k2, 0), blockA, blockB, actual_mc, actual_kc, cols, alpha);
300 for (Index i2 = k2 + kc; i2 < size; i2 += mc) {
301 const Index actual_mc = (std::min)(i2 + mc, size) - i2;
302 gemm_pack_lhs<Scalar, Index, LhsMapper, Traits::mr, Traits::LhsProgress,
typename Traits::LhsPacket4Packing,
303 LhsStorageOrder,
false>()(blockA, lhs.getSubMapper(i2, k2), actual_kc, actual_mc);
305 gebp_kernel(res.getSubMapper(i2, 0), blockA, blockB, actual_mc, actual_kc, cols, alpha);
311template <
typename Scalar,
typename Index,
int LhsStorageOrder,
bool ConjugateLhs,
int RhsStorageOrder,
312 bool ConjugateRhs,
int ResInnerStride>
313struct product_selfadjoint_matrix<Scalar, Index, LhsStorageOrder, false, ConjugateLhs, RhsStorageOrder, true,
314 ConjugateRhs,
ColMajor, ResInnerStride> {
315 static EIGEN_DONT_INLINE
void run(Index rows, Index cols,
const Scalar* lhs_, Index lhsStride,
const Scalar* rhs_,
316 Index rhsStride, Scalar* res, Index resIncr, Index resStride,
const Scalar& alpha,
317 level3_blocking<Scalar, Scalar>& blocking);
320template <
typename Scalar,
typename Index,
int LhsStorageOrder,
bool ConjugateLhs,
int RhsStorageOrder,
321 bool ConjugateRhs,
int ResInnerStride>
322EIGEN_DONT_INLINE
void
323product_selfadjoint_matrix<Scalar, Index, LhsStorageOrder,
false, ConjugateLhs, RhsStorageOrder,
true, ConjugateRhs,
324 ColMajor, ResInnerStride>::run(Index rows, Index cols,
const Scalar* lhs_, Index lhsStride,
325 const Scalar* rhs_, Index rhsStride, Scalar* res_,
326 Index resIncr, Index resStride,
const Scalar& alpha,
327 level3_blocking<Scalar, Scalar>& blocking) {
330 using Traits = gebp_traits<Scalar, Scalar>;
332 using LhsMapper = const_blas_data_mapper<Scalar, Index, LhsStorageOrder>;
333 using ResMapper = blas_data_mapper<typename Traits::ResScalar, Index, ColMajor, Unaligned, ResInnerStride>;
334 LhsMapper lhs(lhs_, lhsStride);
335 ResMapper res(res_, resStride, resIncr);
337 Index kc = blocking.kc();
338 Index mc = (std::min)(rows, blocking.mc());
339 std::size_t sizeA = kc * mc;
340 std::size_t sizeB = kc * cols;
341 ei_declare_aligned_stack_constructed_variable(Scalar, blockA, sizeA, blocking.blockA());
342 ei_declare_aligned_stack_constructed_variable(Scalar, blockB, sizeB, blocking.blockB());
344 gebp_kernel<Scalar, Scalar, Index, ResMapper, Traits::mr, Traits::nr, ConjugateLhs, ConjugateRhs> gebp_kernel;
345 gemm_pack_lhs<Scalar, Index, LhsMapper, Traits::mr, Traits::LhsProgress,
typename Traits::LhsPacket4Packing,
348 symm_pack_rhs<Scalar, Index, Traits::nr, RhsStorageOrder> pack_rhs;
350 for (Index k2 = 0; k2 < size; k2 += kc) {
351 const Index actual_kc = (std::min)(k2 + kc, size) - k2;
353 pack_rhs(blockB, rhs_, rhsStride, actual_kc, cols, k2);
356 for (Index i2 = 0; i2 < rows; i2 += mc) {
357 const Index actual_mc = (std::min)(i2 + mc, rows) - i2;
358 pack_lhs(blockA, lhs.getSubMapper(i2, k2), actual_kc, actual_mc);
360 gebp_kernel(res.getSubMapper(i2, 0), blockA, blockB, actual_mc, actual_kc, cols, alpha);
373template <
typename Lhs,
int LhsMode,
typename Rhs,
int RhsMode>
374struct selfadjoint_product_impl<Lhs, LhsMode, false, Rhs, RhsMode, false> {
375 using Scalar =
typename Product<Lhs, Rhs>::Scalar;
377 using LhsBlasTraits = internal::blas_traits<Lhs>;
378 using ActualLhsType =
typename LhsBlasTraits::DirectLinearAccessType;
379 using RhsBlasTraits = internal::blas_traits<Rhs>;
380 using ActualRhsType =
typename RhsBlasTraits::DirectLinearAccessType;
389 template <
typename Dest>
390 static void run(Dest& dst,
const Lhs& a_lhs,
const Rhs& a_rhs,
const Scalar& alpha) {
391 eigen_assert(dst.rows() == a_lhs.rows() && dst.cols() == a_rhs.cols());
393 add_const_on_value_type_t<ActualLhsType> lhs = LhsBlasTraits::extract(a_lhs);
394 add_const_on_value_type_t<ActualRhsType> rhs = RhsBlasTraits::extract(a_rhs);
398 if (lhs.size() == 0 || rhs.size() == 0)
return;
400 Scalar actualAlpha = alpha * LhsBlasTraits::extractScalarFactor(a_lhs) * RhsBlasTraits::extractScalarFactor(a_rhs);
403 Scalar, Lhs::MaxRowsAtCompileTime, Rhs::MaxColsAtCompileTime,
404 Lhs::MaxColsAtCompileTime, 1>;
406 BlockingType blocking(lhs.rows(), rhs.cols(), lhs.cols(), 1,
false);
408 internal::product_selfadjoint_matrix<
412 NumTraits<Scalar>::IsComplex && internal::logical_xor(LhsIsUpper,
bool(LhsBlasTraits::NeedToConjugate)),
415 NumTraits<Scalar>::IsComplex && internal::logical_xor(RhsIsUpper,
bool(RhsBlasTraits::NeedToConjugate)),
417 Dest::InnerStrideAtCompileTime>::run(lhs.rows(), rhs.cols(),
418 &lhs.coeffRef(0, 0), lhs.outerStride(),
419 &rhs.coeffRef(0, 0), rhs.outerStride(),
420 &dst.coeffRef(0, 0), dst.innerStride(), dst.outerStride(),
421 actualAlpha, blocking
@ SelfAdjoint
Definition Constants.h:228
@ 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