11#ifndef EIGEN_SELFADJOINT_PRODUCT_H
12#define EIGEN_SELFADJOINT_PRODUCT_H
21#include "../InternalHeaderCheck.h"
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;
31 internal::conj_if<ConjRhs> cjy;
32 internal::conj_if<ConjLhs> cjx;
33 internal::conj_helper<Packet, Packet, ConjLhs, false> pcj;
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);
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);
48 col0[0] += s0 * cjx(vecX[j]);
49 col0[1] += s0 * cjx(vecX[j + 1]);
50 col1[0] += s1 * cjx(vecX[j + 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;
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);
69 for (; k < len; ++k) {
70 Scalar cx = cjx(xp[k]);
76 Scalar* EIGEN_RESTRICT col0 = mat + stride * j;
77 Scalar* EIGEN_RESTRICT col1 = mat + stride * (j + 1);
80 const Scalar* EIGEN_RESTRICT xp = vecX;
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);
93 for (; k < len; ++k) {
94 Scalar cx = cjx(xp[k]);
100 col0[j] += s0 * cjx(vecX[j]);
101 col1[j] += s1 * cjx(vecX[j]);
102 col1[j + 1] += s1 * cjx(vecX[j + 1]);
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;
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);
123 for (; k < len; ++k) {
124 dst[k] += s * cjx(xp[k]);
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);
138template <
typename MatrixType,
typename OtherType,
int UpLo,
bool OtherIsVector = OtherType::IsVectorAtCompileTime>
139struct selfadjoint_product_selector;
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());
150 Scalar actualAlpha = alpha * OtherBlasTraits::extractScalarFactor(other.derived());
154 UseOtherDirectly = ActualOtherType_::InnerStrideAtCompileTime == 1
156 internal::gemv_static_vector_if<Scalar, OtherType::SizeAtCompileTime, OtherType::MaxSizeAtCompileTime,
160 ei_declare_aligned_stack_constructed_variable(
161 Scalar, actualOtherPtr, other.size(),
162 (UseOtherDirectly ?
const_cast<Scalar*
>(actualOther.data()) : static_other.data()));
164 EIGEN_IF_CONSTEXPR (!UseOtherDirectly) {
165 Map<typename ActualOtherType_::PlainObject>(actualOtherPtr, actualOther.size()) = actualOther;
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);
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());
185 Scalar actualAlpha = alpha * OtherBlasTraits::extractScalarFactor(other.derived());
188 IsRowMajor = (internal::traits<MatrixType>::Flags &
RowMajorBit) ? 1 : 0,
189 OtherIsRowMajor = ActualOtherType_::Flags &
RowMajorBit ? 1 : 0
192 Index size = mat.cols();
193 Index depth = actualOther.cols();
194 eigen_assert(actualOther.rows() == size);
195 if (size == 0 || depth == 0)
return;
198 internal::gemm_blocking_space<IsRowMajor ?
RowMajor :
ColMajor, Scalar, Scalar,
199 MatrixType::MaxColsAtCompileTime, MatrixType::MaxColsAtCompileTime,
200 ActualOtherType_::MaxColsAtCompileTime>;
202 BlockingType blocking(size, size, depth, 1,
false);
204 internal::general_matrix_matrix_triangular_product<
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);
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);
@ Lower
Definition Constants.h:212
@ ColMajor
Definition Constants.h:319
@ RowMajor
Definition Constants.h:321
constexpr unsigned int RowMajorBit
Definition Constants.h:71