Eigen  5.0.1
 
Loading...
Searching...
No Matches
SelfadjointRank2Update.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_SELFADJOINTRANK2UPDATE_H
12#define EIGEN_SELFADJOINTRANK2UPDATE_H
13
14// IWYU pragma: private
15#include "../InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20
21/* Optimized selfadjoint matrix += alpha * uv' + conj(alpha)*vu'
22 * It corresponds to the Level2 syr2 BLAS routine
23 */
24
25template <typename Scalar, typename Index, int UpLo>
26struct selfadjoint_rank2_update_selector;
27
28template <typename Scalar, typename Index>
29struct selfadjoint_rank2_update_selector<Scalar, Index, Lower> {
30 EIGEN_DEVICE_FUNC static void run(Index size, Scalar* mat, Index stride, const Scalar* u, const Scalar* v,
31 const Scalar& alpha) {
32 using Packet = typename packet_traits<Scalar>::type;
33 const Index PacketSize = unpacket_traits<Packet>::size;
34 const Scalar cAlpha = numext::conj(alpha);
35
36 // Process 2 columns at a time to share u/v loads and reduce loop overhead.
37 Index j = 0;
38 for (; j + 1 < size; j += 2) {
39 // Scale factors: col[j:] += s0u * v[j:] + s0v * u[j:]
40 Scalar s0u = cAlpha * numext::conj(u[j]);
41 Scalar s0v = alpha * numext::conj(v[j]);
42 Scalar s1u = cAlpha * numext::conj(u[j + 1]);
43 Scalar s1v = alpha * numext::conj(v[j + 1]);
44
45 Packet ps0u = pset1<Packet>(s0u);
46 Packet ps0v = pset1<Packet>(s0v);
47 Packet ps1u = pset1<Packet>(s1u);
48 Packet ps1v = pset1<Packet>(s1v);
49
50 Scalar* EIGEN_RESTRICT col0 = mat + stride * j + j;
51 Scalar* EIGEN_RESTRICT col1 = mat + stride * (j + 1) + (j + 1);
52
53 // Diagonal and cross-diagonal scalar elements
54 col0[0] += s0u * v[j] + s0v * u[j];
55 col0[1] += s0u * v[j + 1] + s0v * u[j + 1];
56 col1[0] += s1u * v[j + 1] + s1v * u[j + 1];
57
58 // Shared vectorized loop for rows j+2..size-1
59 Index len = size - j - 2;
60 const Scalar* EIGEN_RESTRICT up = u + j + 2;
61 const Scalar* EIGEN_RESTRICT vp = v + j + 2;
62 Scalar* EIGEN_RESTRICT d0 = col0 + 2;
63 Scalar* EIGEN_RESTRICT d1 = col1 + 1;
64
65 Index k = 0;
66 Index vectorizedEnd = (len / PacketSize) * PacketSize;
67 for (; k < vectorizedEnd; k += PacketSize) {
68 Packet ui = ploadu<Packet>(up + k);
69 Packet vi = ploadu<Packet>(vp + k);
70 Packet m0 = ploadu<Packet>(d0 + k);
71 m0 = pmadd(vi, ps0u, m0);
72 m0 = pmadd(ui, ps0v, m0);
73 pstoreu(d0 + k, m0);
74 Packet m1 = ploadu<Packet>(d1 + k);
75 m1 = pmadd(vi, ps1u, m1);
76 m1 = pmadd(ui, ps1v, m1);
77 pstoreu(d1 + k, m1);
78 }
79 for (; k < len; ++k) {
80 d0[k] += s0u * vp[k] + s0v * up[k];
81 d1[k] += s1u * vp[k] + s1v * up[k];
82 }
83 }
84
85 // Handle last column if size is odd
86 if (j < size) {
87 Scalar su = cAlpha * numext::conj(u[j]);
88 Scalar sv = alpha * numext::conj(v[j]);
89 Packet psu = pset1<Packet>(su);
90 Packet psv = pset1<Packet>(sv);
91
92 Scalar* EIGEN_RESTRICT dst = mat + stride * j + j;
93 const Scalar* EIGEN_RESTRICT up = u + j;
94 const Scalar* EIGEN_RESTRICT vp = v + j;
95 Index len = size - j;
96
97 Index k = 0;
98 Index vectorizedEnd = (len / PacketSize) * PacketSize;
99 for (; k < vectorizedEnd; k += PacketSize) {
100 Packet ui = ploadu<Packet>(up + k);
101 Packet vi = ploadu<Packet>(vp + k);
102 Packet di = ploadu<Packet>(dst + k);
103 di = pmadd(vi, psu, di);
104 di = pmadd(ui, psv, di);
105 pstoreu(dst + k, di);
106 }
107 for (; k < len; ++k) {
108 dst[k] += su * vp[k] + sv * up[k];
109 }
110 }
111 }
112};
113
114template <typename Scalar, typename Index>
115struct selfadjoint_rank2_update_selector<Scalar, Index, Upper> {
116 EIGEN_DEVICE_FUNC static void run(Index size, Scalar* mat, Index stride, const Scalar* u, const Scalar* v,
117 const Scalar& alpha) {
118 using Packet = typename packet_traits<Scalar>::type;
119 const Index PacketSize = unpacket_traits<Packet>::size;
120 const Scalar cAlpha = numext::conj(alpha);
121
122 // Process 2 columns at a time to share u/v loads and reduce loop overhead.
123 Index j = 0;
124 for (; j + 1 < size; j += 2) {
125 Scalar s0u = cAlpha * numext::conj(u[j]);
126 Scalar s0v = alpha * numext::conj(v[j]);
127 Scalar s1u = cAlpha * numext::conj(u[j + 1]);
128 Scalar s1v = alpha * numext::conj(v[j + 1]);
129
130 Packet ps0u = pset1<Packet>(s0u);
131 Packet ps0v = pset1<Packet>(s0v);
132 Packet ps1u = pset1<Packet>(s1u);
133 Packet ps1v = pset1<Packet>(s1v);
134
135 Scalar* EIGEN_RESTRICT col0 = mat + stride * j;
136 Scalar* EIGEN_RESTRICT col1 = mat + stride * (j + 1);
137
138 // Shared vectorized loop for rows 0..j-1
139 Index len = j;
140 Index k = 0;
141 Index vectorizedEnd = (len / PacketSize) * PacketSize;
142 for (; k < vectorizedEnd; k += PacketSize) {
143 Packet ui = ploadu<Packet>(u + k);
144 Packet vi = ploadu<Packet>(v + k);
145 Packet m0 = ploadu<Packet>(col0 + k);
146 m0 = pmadd(vi, ps0u, m0);
147 m0 = pmadd(ui, ps0v, m0);
148 pstoreu(col0 + k, m0);
149 Packet m1 = ploadu<Packet>(col1 + k);
150 m1 = pmadd(vi, ps1u, m1);
151 m1 = pmadd(ui, ps1v, m1);
152 pstoreu(col1 + k, m1);
153 }
154 for (; k < len; ++k) {
155 col0[k] += s0u * v[k] + s0v * u[k];
156 col1[k] += s1u * v[k] + s1v * u[k];
157 }
158
159 // Diagonal and cross-diagonal scalar elements
160 col0[j] += s0u * v[j] + s0v * u[j];
161 col1[j] += s1u * v[j] + s1v * u[j];
162 col1[j + 1] += s1u * v[j + 1] + s1v * u[j + 1];
163 }
164
165 // Handle last column if size is odd
166 if (j < size) {
167 Scalar su = cAlpha * numext::conj(u[j]);
168 Scalar sv = alpha * numext::conj(v[j]);
169 Packet psu = pset1<Packet>(su);
170 Packet psv = pset1<Packet>(sv);
171
172 Scalar* EIGEN_RESTRICT dst = mat + stride * j;
173 Index len = j + 1;
174
175 Index k = 0;
176 Index vectorizedEnd = (len / PacketSize) * PacketSize;
177 for (; k < vectorizedEnd; k += PacketSize) {
178 Packet ui = ploadu<Packet>(u + k);
179 Packet vi = ploadu<Packet>(v + k);
180 Packet di = ploadu<Packet>(dst + k);
181 di = pmadd(vi, psu, di);
182 di = pmadd(ui, psv, di);
183 pstoreu(dst + k, di);
184 }
185 for (; k < len; ++k) {
186 dst[k] += su * v[k] + sv * u[k];
187 }
188 }
189 }
190};
191
192} // end namespace internal
193
194template <typename MatrixType, unsigned int UpLo>
195template <typename DerivedU, typename DerivedV>
197 const MatrixBase<DerivedU>& u, const MatrixBase<DerivedV>& v, const Scalar& alpha) {
198 using UBlasTraits = internal::blas_traits<DerivedU>;
199 using ActualUType = typename UBlasTraits::DirectLinearAccessType;
200 using ActualUType_ = internal::remove_all_t<ActualUType>;
201 internal::add_const_on_value_type_t<ActualUType> actualU = UBlasTraits::extract(u.derived());
202
203 using VBlasTraits = internal::blas_traits<DerivedV>;
204 using ActualVType = typename VBlasTraits::DirectLinearAccessType;
205 using ActualVType_ = internal::remove_all_t<ActualVType>;
206 internal::add_const_on_value_type_t<ActualVType> actualV = VBlasTraits::extract(v.derived());
207
208 // If MatrixType is row major, then we use the routine for lower triangular in the upper triangular case and
209 // vice versa, and take the complex conjugate of all coefficients and vector entries.
210 enum {
211 IsRowMajor = (internal::traits<MatrixType>::Flags & RowMajorBit) ? 1 : 0,
212 // Only need to conjugate if complex and the condition triggers
213 NeedConjU = (int(IsRowMajor) ^ int(UBlasTraits::NeedToConjugate)) && NumTraits<Scalar>::IsComplex,
214 NeedConjV = (int(IsRowMajor) ^ int(VBlasTraits::NeedToConjugate)) && NumTraits<Scalar>::IsComplex,
215 UseUDirectly = ActualUType_::InnerStrideAtCompileTime == 1 && !NeedConjU,
216 UseVDirectly = ActualVType_::InnerStrideAtCompileTime == 1 && !NeedConjV
217 };
218
219 Scalar actualAlpha = alpha * UBlasTraits::extractScalarFactor(u.derived()) *
220 numext::conj(VBlasTraits::extractScalarFactor(v.derived()));
221 EIGEN_IF_CONSTEXPR (IsRowMajor) {
222 actualAlpha = numext::conj(actualAlpha);
223 }
224
225 const Index size = u.size();
226
227 // Copy u to contiguous buffer, applying conjugation if needed
228 internal::gemv_static_vector_if<Scalar, DerivedU::SizeAtCompileTime, DerivedU::MaxSizeAtCompileTime, !UseUDirectly>
229 static_u;
230 ei_declare_aligned_stack_constructed_variable(Scalar, uPtr, size,
231 (UseUDirectly ? const_cast<Scalar*>(actualU.data()) : static_u.data()));
232 EIGEN_IF_CONSTEXPR (!UseUDirectly) {
233 EIGEN_IF_CONSTEXPR (NeedConjU) {
234 Map<typename ActualUType_::PlainObject>(uPtr, size) = actualU.conjugate();
235 } else {
236 Map<typename ActualUType_::PlainObject>(uPtr, size) = actualU;
237 }
238 }
239
240 // Copy v to contiguous buffer, applying conjugation if needed
241 internal::gemv_static_vector_if<Scalar, DerivedV::SizeAtCompileTime, DerivedV::MaxSizeAtCompileTime, !UseVDirectly>
242 static_v;
243 ei_declare_aligned_stack_constructed_variable(Scalar, vPtr, size,
244 (UseVDirectly ? const_cast<Scalar*>(actualV.data()) : static_v.data()));
245 EIGEN_IF_CONSTEXPR (!UseVDirectly) {
246 EIGEN_IF_CONSTEXPR (NeedConjV) {
247 Map<typename ActualVType_::PlainObject>(vPtr, size) = actualV.conjugate();
248 } else {
249 Map<typename ActualVType_::PlainObject>(vPtr, size) = actualV;
250 }
251 }
252
253 internal::selfadjoint_rank2_update_selector<
254 Scalar, Index, (IsRowMajor ? int(UpLo == Upper ? Lower : Upper) : UpLo)>::run(size, nestedExpression().data(),
255 nestedExpression().outerStride(),
256 uPtr, vPtr, actualAlpha);
257
258 return *this;
259}
260
261} // end namespace Eigen
262
263#endif // EIGEN_SELFADJOINTRANK2UPDATE_H
A matrix or vector expression mapping an existing array of data.
Definition Map.h:97
Base class for all dense matrices, vectors, and expressions.
Definition MatrixBase.h:53
Expression of a selfadjoint matrix from a triangular part of a dense matrix.
Definition SelfAdjointView.h:54
SelfAdjointView & rankUpdate(const MatrixBase< DerivedU > &u, const MatrixBase< DerivedV > &v, const Scalar &alpha=Scalar(1))
@ Lower
Definition Constants.h:212
@ Upper
Definition Constants.h:214
constexpr unsigned int RowMajorBit
Definition Constants.h:71
Holds information about the various numeric (i.e. scalar) types allowed by Eigen.
Definition NumTraits.h:233