11#ifndef EIGEN_SELFADJOINTRANK2UPDATE_H
12#define EIGEN_SELFADJOINTRANK2UPDATE_H
15#include "../InternalHeaderCheck.h"
25template <
typename Scalar,
typename Index,
int UpLo>
26struct selfadjoint_rank2_update_selector;
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);
38 for (; j + 1 < size; j += 2) {
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]);
45 Packet ps0u = pset1<Packet>(s0u);
46 Packet ps0v = pset1<Packet>(s0v);
47 Packet ps1u = pset1<Packet>(s1u);
48 Packet ps1v = pset1<Packet>(s1v);
50 Scalar* EIGEN_RESTRICT col0 = mat + stride * j + j;
51 Scalar* EIGEN_RESTRICT col1 = mat + stride * (j + 1) + (j + 1);
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];
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;
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);
74 Packet m1 = ploadu<Packet>(d1 + k);
75 m1 = pmadd(vi, ps1u, m1);
76 m1 = pmadd(ui, ps1v, m1);
79 for (; k < len; ++k) {
80 d0[k] += s0u * vp[k] + s0v * up[k];
81 d1[k] += s1u * vp[k] + s1v * up[k];
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);
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;
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);
107 for (; k < len; ++k) {
108 dst[k] += su * vp[k] + sv * up[k];
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);
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]);
130 Packet ps0u = pset1<Packet>(s0u);
131 Packet ps0v = pset1<Packet>(s0v);
132 Packet ps1u = pset1<Packet>(s1u);
133 Packet ps1v = pset1<Packet>(s1v);
135 Scalar* EIGEN_RESTRICT col0 = mat + stride * j;
136 Scalar* EIGEN_RESTRICT col1 = mat + stride * (j + 1);
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);
154 for (; k < len; ++k) {
155 col0[k] += s0u * v[k] + s0v * u[k];
156 col1[k] += s1u * v[k] + s1v * u[k];
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];
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);
172 Scalar* EIGEN_RESTRICT dst = mat + stride * j;
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);
185 for (; k < len; ++k) {
186 dst[k] += su * v[k] + sv * u[k];
194template <
typename MatrixType,
unsigned int UpLo>
195template <
typename DerivedU,
typename DerivedV>
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());
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());
211 IsRowMajor = (internal::traits<MatrixType>::Flags &
RowMajorBit) ? 1 : 0,
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
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);
225 const Index size = u.size();
228 internal::gemv_static_vector_if<Scalar, DerivedU::SizeAtCompileTime, DerivedU::MaxSizeAtCompileTime, !UseUDirectly>
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) {
241 internal::gemv_static_vector_if<Scalar, DerivedV::SizeAtCompileTime, DerivedV::MaxSizeAtCompileTime, !UseVDirectly>
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) {
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);
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