Eigen  5.0.1
 
Loading...
Searching...
No Matches
SelfadjointProduct.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_SELFADJOINT_PRODUCT_H
12#define EIGEN_SELFADJOINT_PRODUCT_H
13
14/**********************************************************************
15 * This file implements a self adjoint product: C += A A^T updating only
16 * half of the selfadjoint matrix C.
17 * It corresponds to the level 3 SYRK and level 2 SYR Blas routines.
18 **********************************************************************/
19
20// IWYU pragma: private
21#include "../InternalHeaderCheck.h"
22
23namespace Eigen {
24
25template <typename Scalar, typename Index, int UpLo, bool ConjLhs, bool ConjRhs>
26struct selfadjoint_rank1_update<Scalar, Index, ColMajor, UpLo, ConjLhs, ConjRhs> {
27 static void run(Index size, Scalar* mat, Index stride, const Scalar* vecX, const Scalar* vecY, const Scalar& alpha) {
28 using Packet = typename internal::packet_traits<Scalar>::type;
29 const Index PacketSize = internal::unpacket_traits<Packet>::size;
30
31 internal::conj_if<ConjRhs> cjy;
32 internal::conj_if<ConjLhs> cjx;
33 internal::conj_helper<Packet, Packet, ConjLhs, false> pcj;
34
35 // Process 2 columns at a time to share vecX loads and reduce loop overhead.
36 Index j = 0;
37 for (; j + 1 < size; j += 2) {
38 Scalar s0 = alpha * cjy(vecY[j]);
39 Scalar s1 = alpha * cjy(vecY[j + 1]);
40 Packet ps0 = internal::pset1<Packet>(s0);
41 Packet ps1 = internal::pset1<Packet>(s1);
42
43 EIGEN_IF_CONSTEXPR (UpLo == Lower) {
44 Scalar* EIGEN_RESTRICT col0 = mat + stride * j + j;
45 Scalar* EIGEN_RESTRICT col1 = mat + stride * (j + 1) + (j + 1);
46
47 // Diagonal and cross-diagonal scalar elements
48 col0[0] += s0 * cjx(vecX[j]);
49 col0[1] += s0 * cjx(vecX[j + 1]);
50 col1[0] += s1 * cjx(vecX[j + 1]);
51
52 // Shared vectorized loop for rows j+2..size-1
53 Index len = size - j - 2;
54 const Scalar* EIGEN_RESTRICT xp = vecX + j + 2;
55 Scalar* EIGEN_RESTRICT d0 = col0 + 2;
56 Scalar* EIGEN_RESTRICT d1 = col1 + 1;
57
58 Index k = 0;
59 Index vectorizedEnd = (len / PacketSize) * PacketSize;
60 for (; k < vectorizedEnd; k += PacketSize) {
61 Packet xi = internal::ploadu<Packet>(xp + k);
62 Packet m0 = internal::ploadu<Packet>(d0 + k);
63 m0 = pcj.pmadd(xi, ps0, m0);
64 internal::pstoreu(d0 + k, m0);
65 Packet m1 = internal::ploadu<Packet>(d1 + k);
66 m1 = pcj.pmadd(xi, ps1, m1);
67 internal::pstoreu(d1 + k, m1);
68 }
69 for (; k < len; ++k) {
70 Scalar cx = cjx(xp[k]);
71 d0[k] += s0 * cx;
72 d1[k] += s1 * cx;
73 }
74 } else {
75 // UpLo == Upper
76 Scalar* EIGEN_RESTRICT col0 = mat + stride * j;
77 Scalar* EIGEN_RESTRICT col1 = mat + stride * (j + 1);
78
79 // Shared vectorized loop for rows 0..j-1
80 const Scalar* EIGEN_RESTRICT xp = vecX;
81 Index len = j;
82 Index k = 0;
83 Index vectorizedEnd = (len / PacketSize) * PacketSize;
84 for (; k < vectorizedEnd; k += PacketSize) {
85 Packet xi = internal::ploadu<Packet>(xp + k);
86 Packet m0 = internal::ploadu<Packet>(col0 + k);
87 Packet m1 = internal::ploadu<Packet>(col1 + k);
88 m0 = pcj.pmadd(xi, ps0, m0);
89 m1 = pcj.pmadd(xi, ps1, m1);
90 internal::pstoreu(col0 + k, m0);
91 internal::pstoreu(col1 + k, m1);
92 }
93 for (; k < len; ++k) {
94 Scalar cx = cjx(xp[k]);
95 col0[k] += s0 * cx;
96 col1[k] += s1 * cx;
97 }
98
99 // Diagonal and cross-diagonal scalar elements
100 col0[j] += s0 * cjx(vecX[j]);
101 col1[j] += s1 * cjx(vecX[j]);
102 col1[j + 1] += s1 * cjx(vecX[j + 1]);
103 }
104 }
105
106 // Handle last column if size is odd
107 if (j < size) {
108 Scalar s = alpha * cjy(vecY[j]);
109 Packet ps = internal::pset1<Packet>(s);
110 Index start = UpLo == Lower ? j : 0;
111 Index len = UpLo == Lower ? size - j : j + 1;
112 Scalar* EIGEN_RESTRICT dst = mat + stride * j + start;
113 const Scalar* EIGEN_RESTRICT xp = vecX + start;
114
115 Index k = 0;
116 Index vectorizedEnd = (len / PacketSize) * PacketSize;
117 for (; k < vectorizedEnd; k += PacketSize) {
118 Packet xi = internal::ploadu<Packet>(xp + k);
119 Packet di = internal::ploadu<Packet>(dst + k);
120 di = pcj.pmadd(xi, ps, di);
121 internal::pstoreu(dst + k, di);
122 }
123 for (; k < len; ++k) {
124 dst[k] += s * cjx(xp[k]);
125 }
126 }
127 }
128};
129
130template <typename Scalar, typename Index, int UpLo, bool ConjLhs, bool ConjRhs>
131struct selfadjoint_rank1_update<Scalar, Index, RowMajor, UpLo, ConjLhs, ConjRhs> {
132 static void run(Index size, Scalar* mat, Index stride, const Scalar* vecX, const Scalar* vecY, const Scalar& alpha) {
133 selfadjoint_rank1_update<Scalar, Index, ColMajor, UpLo == Lower ? Upper : Lower, ConjRhs, ConjLhs>::run(
134 size, mat, stride, vecY, vecX, alpha);
135 }
136};
137
138template <typename MatrixType, typename OtherType, int UpLo, bool OtherIsVector = OtherType::IsVectorAtCompileTime>
139struct selfadjoint_product_selector;
140
141template <typename MatrixType, typename OtherType, int UpLo>
142struct selfadjoint_product_selector<MatrixType, OtherType, UpLo, true> {
143 static void run(MatrixType& mat, const OtherType& other, const typename MatrixType::Scalar& alpha) {
144 using Scalar = typename MatrixType::Scalar;
145 using OtherBlasTraits = internal::blas_traits<OtherType>;
146 using ActualOtherType = typename OtherBlasTraits::DirectLinearAccessType;
147 using ActualOtherType_ = internal::remove_all_t<ActualOtherType>;
148 internal::add_const_on_value_type_t<ActualOtherType> actualOther = OtherBlasTraits::extract(other.derived());
149
150 Scalar actualAlpha = alpha * OtherBlasTraits::extractScalarFactor(other.derived());
151
152 enum {
153 StorageOrder = (internal::traits<MatrixType>::Flags & RowMajorBit) ? RowMajor : ColMajor,
154 UseOtherDirectly = ActualOtherType_::InnerStrideAtCompileTime == 1
155 };
156 internal::gemv_static_vector_if<Scalar, OtherType::SizeAtCompileTime, OtherType::MaxSizeAtCompileTime,
157 !UseOtherDirectly>
158 static_other;
159
160 ei_declare_aligned_stack_constructed_variable(
161 Scalar, actualOtherPtr, other.size(),
162 (UseOtherDirectly ? const_cast<Scalar*>(actualOther.data()) : static_other.data()));
163
164 EIGEN_IF_CONSTEXPR (!UseOtherDirectly) {
165 Map<typename ActualOtherType_::PlainObject>(actualOtherPtr, actualOther.size()) = actualOther;
166 }
167
168 selfadjoint_rank1_update<
169 Scalar, Index, StorageOrder, UpLo, OtherBlasTraits::NeedToConjugate && NumTraits<Scalar>::IsComplex,
170 (!OtherBlasTraits::NeedToConjugate) && NumTraits<Scalar>::IsComplex>::run(other.size(), mat.data(),
171 mat.outerStride(), actualOtherPtr,
172 actualOtherPtr, actualAlpha);
173 }
174};
175
176template <typename MatrixType, typename OtherType, int UpLo>
177struct selfadjoint_product_selector<MatrixType, OtherType, UpLo, false> {
178 static void run(MatrixType& mat, const OtherType& other, const typename MatrixType::Scalar& alpha) {
179 using Scalar = typename MatrixType::Scalar;
180 using OtherBlasTraits = internal::blas_traits<OtherType>;
181 using ActualOtherType = typename OtherBlasTraits::DirectLinearAccessType;
182 using ActualOtherType_ = internal::remove_all_t<ActualOtherType>;
183 internal::add_const_on_value_type_t<ActualOtherType> actualOther = OtherBlasTraits::extract(other.derived());
184
185 Scalar actualAlpha = alpha * OtherBlasTraits::extractScalarFactor(other.derived());
186
187 enum {
188 IsRowMajor = (internal::traits<MatrixType>::Flags & RowMajorBit) ? 1 : 0,
189 OtherIsRowMajor = ActualOtherType_::Flags & RowMajorBit ? 1 : 0
190 };
191
192 Index size = mat.cols();
193 Index depth = actualOther.cols();
194 eigen_assert(actualOther.rows() == size);
195 if (size == 0 || depth == 0) return;
196
197 using BlockingType =
198 internal::gemm_blocking_space<IsRowMajor ? RowMajor : ColMajor, Scalar, Scalar,
199 MatrixType::MaxColsAtCompileTime, MatrixType::MaxColsAtCompileTime,
200 ActualOtherType_::MaxColsAtCompileTime>;
201
202 BlockingType blocking(size, size, depth, 1, false);
203
204 internal::general_matrix_matrix_triangular_product<
205 Index, Scalar, OtherIsRowMajor ? RowMajor : ColMajor,
206 OtherBlasTraits::NeedToConjugate && NumTraits<Scalar>::IsComplex, Scalar, OtherIsRowMajor ? ColMajor : RowMajor,
207 (!OtherBlasTraits::NeedToConjugate) && NumTraits<Scalar>::IsComplex, IsRowMajor ? RowMajor : ColMajor,
208 MatrixType::InnerStrideAtCompileTime, UpLo>::run(size, depth, actualOther.data(), actualOther.outerStride(),
209 actualOther.data(), actualOther.outerStride(), mat.data(),
210 mat.innerStride(), mat.outerStride(), actualAlpha, blocking);
211 }
212};
213
214// high level API
215
216template <typename MatrixType, unsigned int UpLo>
217template <typename DerivedU>
218EIGEN_DEVICE_FUNC SelfAdjointView<MatrixType, UpLo>& SelfAdjointView<MatrixType, UpLo>::rankUpdate(
219 const MatrixBase<DerivedU>& u, const Scalar& alpha) {
220 selfadjoint_product_selector<MatrixType, DerivedU, UpLo>::run(nestedExpression(), u.derived(), alpha);
221
222 return *this;
223}
224
225} // end namespace Eigen
226
227#endif // EIGEN_SELFADJOINT_PRODUCT_H
@ Lower
Definition Constants.h:212
@ ColMajor
Definition Constants.h:319
@ RowMajor
Definition Constants.h:321
constexpr unsigned int RowMajorBit
Definition Constants.h:71