Eigen  5.0.1
 
Loading...
Searching...
No Matches
SelfadjointMatrixVector.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2008-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_MATRIX_VECTOR_H
12#define EIGEN_SELFADJOINT_MATRIX_VECTOR_H
13
14// IWYU pragma: private
15#include "../InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20
21/* Optimized selfadjoint matrix * vector product:
22 * This algorithm processes 4 columns at once to reduce the number of
23 * load/stores of the result vector by a factor of 4 compared to the
24 * naive approach, and to increase instruction-level parallelism.
25 * A 2-column cleanup handles the remaining even columns, and a
26 * 1-column loop handles any final odd column.
27 */
28
29template <typename Scalar, typename Index, int StorageOrder, int UpLo, bool ConjugateLhs, bool ConjugateRhs,
30 int Version = Specialized>
31struct selfadjoint_matrix_vector_product;
32
33template <typename Scalar, typename Index, int StorageOrder, int UpLo, bool ConjugateLhs, bool ConjugateRhs,
34 int Version>
35struct selfadjoint_matrix_vector_product
36
37{
38 static EIGEN_DONT_INLINE EIGEN_DEVICE_FUNC void run(Index size, const Scalar* lhs, Index lhsStride, const Scalar* rhs,
39 Scalar* res, Scalar alpha);
40};
41
42template <typename Scalar, typename Index, int StorageOrder, int UpLo, bool ConjugateLhs, bool ConjugateRhs,
43 int Version>
44EIGEN_DONT_INLINE EIGEN_DEVICE_FUNC void
45selfadjoint_matrix_vector_product<Scalar, Index, StorageOrder, UpLo, ConjugateLhs, ConjugateRhs, Version>::run(
46 Index size, const Scalar* lhs, Index lhsStride, const Scalar* rhs, Scalar* res, Scalar alpha) {
47 using Packet = typename packet_traits<Scalar>::type;
48 using RealScalar = typename NumTraits<Scalar>::Real;
49 const Index PacketSize = sizeof(Packet) / sizeof(Scalar);
50
51 enum {
52 IsRowMajor = StorageOrder == RowMajor ? 1 : 0,
53 IsLower = UpLo == Lower ? 1 : 0,
54 FirstTriangular = IsRowMajor == IsLower
55 };
56
57 conj_helper<Scalar, Scalar, NumTraits<Scalar>::IsComplex && logical_xor(ConjugateLhs, IsRowMajor), ConjugateRhs> cj0;
58 conj_helper<Scalar, Scalar, NumTraits<Scalar>::IsComplex && logical_xor(ConjugateLhs, !IsRowMajor), ConjugateRhs> cj1;
59 conj_helper<RealScalar, Scalar, false, ConjugateRhs> cjd;
60
61 conj_helper<Packet, Packet, NumTraits<Scalar>::IsComplex && logical_xor(ConjugateLhs, IsRowMajor), ConjugateRhs> pcj0;
62 conj_helper<Packet, Packet, NumTraits<Scalar>::IsComplex && logical_xor(ConjugateLhs, !IsRowMajor), ConjugateRhs>
63 pcj1;
64
65 Scalar cjAlpha = ConjugateRhs ? numext::conj(alpha) : alpha;
66
67 // Compute column counts for 4-col, 2-col, and 1-col processing phases.
68 // We leave up to ~8 columns near the diagonal for cleanup (short off-diagonal ranges).
69 Index n4 = (numext::maxi(Index(0), size - 8) / 4) * 4;
70 Index n2 = ((size - n4) / 2) * 2;
71 // Remaining (size - n4 - n2) is 0 or 1 columns.
72
73 // For !FirstTriangular: 4-col [0, n4), 2-col [n4, n4+n2), 1-col [n4+n2, size)
74 // For FirstTriangular: 1-col [0, size-n4-n2), 2-col [size-n4-n2, size-n4), 4-col [size-n4, size)
75
76 // === Phase 1: 4 columns at a time ===
77 {
78 Index jStart = FirstTriangular ? (size - n4) : 0;
79 Index jEnd = FirstTriangular ? size : n4;
80
81 for (Index j = jStart; j < jEnd; j += 4) {
82 const Scalar* EIGEN_RESTRICT A0 = lhs + j * lhsStride;
83 const Scalar* EIGEN_RESTRICT A1 = lhs + (j + 1) * lhsStride;
84 const Scalar* EIGEN_RESTRICT A2 = lhs + (j + 2) * lhsStride;
85 const Scalar* EIGEN_RESTRICT A3 = lhs + (j + 3) * lhsStride;
86
87 Scalar t0 = cjAlpha * rhs[j];
88 Scalar t1 = cjAlpha * rhs[j + 1];
89 Scalar t2 = cjAlpha * rhs[j + 2];
90 Scalar t3 = cjAlpha * rhs[j + 3];
91 Packet ptmp0 = pset1<Packet>(t0);
92 Packet ptmp1 = pset1<Packet>(t1);
93 Packet ptmp2 = pset1<Packet>(t2);
94 Packet ptmp3 = pset1<Packet>(t3);
95
96 Scalar t4(0), t5(0), t6(0), t7(0);
97 Packet ptmp4 = pzero(Packet{});
98 Packet ptmp5 = pzero(Packet{});
99 Packet ptmp6 = pzero(Packet{});
100 Packet ptmp7 = pzero(Packet{});
101
102 Index starti = FirstTriangular ? 0 : j + 4;
103 Index endi = FirstTriangular ? j : size;
104 Index alignedStart = starti + internal::first_default_aligned(&res[starti], endi - starti);
105 Index alignedEnd = alignedStart + ((endi - alignedStart) / PacketSize) * PacketSize;
106
107 // Handle the 4x4 diagonal block: diagonal elements
108 res[j] += cjd.pmul(numext::real(A0[j]), t0);
109 res[j + 1] += cjd.pmul(numext::real(A1[j + 1]), t1);
110 res[j + 2] += cjd.pmul(numext::real(A2[j + 2]), t2);
111 res[j + 3] += cjd.pmul(numext::real(A3[j + 3]), t3);
112
113 // Handle the 4x4 diagonal block: off-diagonal cross terms
114 EIGEN_IF_CONSTEXPR (FirstTriangular) {
115 // Upper triangle stored (A_k[l] for l <= k)
116 res[j] += cj0.pmul(A1[j], t1) + cj0.pmul(A2[j], t2) + cj0.pmul(A3[j], t3);
117 res[j + 1] += cj0.pmul(A2[j + 1], t2) + cj0.pmul(A3[j + 1], t3);
118 res[j + 2] += cj0.pmul(A3[j + 2], t3);
119
120 t5 += cj1.pmul(A1[j], rhs[j]);
121 t6 += cj1.pmul(A2[j], rhs[j]) + cj1.pmul(A2[j + 1], rhs[j + 1]);
122 t7 += cj1.pmul(A3[j], rhs[j]) + cj1.pmul(A3[j + 1], rhs[j + 1]) + cj1.pmul(A3[j + 2], rhs[j + 2]);
123 } else {
124 // Lower triangle stored (A_k[l] for l >= k)
125 res[j + 1] += cj0.pmul(A0[j + 1], t0);
126 res[j + 2] += cj0.pmul(A0[j + 2], t0) + cj0.pmul(A1[j + 2], t1);
127 res[j + 3] += cj0.pmul(A0[j + 3], t0) + cj0.pmul(A1[j + 3], t1) + cj0.pmul(A2[j + 3], t2);
128
129 t4 += cj1.pmul(A0[j + 1], rhs[j + 1]) + cj1.pmul(A0[j + 2], rhs[j + 2]) + cj1.pmul(A0[j + 3], rhs[j + 3]);
130 t5 += cj1.pmul(A1[j + 2], rhs[j + 2]) + cj1.pmul(A1[j + 3], rhs[j + 3]);
131 t6 += cj1.pmul(A2[j + 3], rhs[j + 3]);
132 }
133
134 // Pre-alignment scalar loop
135 for (Index i = starti; i < alignedStart; ++i) {
136 res[i] += cj0.pmul(A0[i], t0) + cj0.pmul(A1[i], t1) + cj0.pmul(A2[i], t2) + cj0.pmul(A3[i], t3);
137 t4 += cj1.pmul(A0[i], rhs[i]);
138 t5 += cj1.pmul(A1[i], rhs[i]);
139 t6 += cj1.pmul(A2[i], rhs[i]);
140 t7 += cj1.pmul(A3[i], rhs[i]);
141 }
142
143 // Main vectorized loop: 4 matrix column loads, 1 rhs load, 1 result load/store
144 const Scalar* EIGEN_RESTRICT a0It = A0 + alignedStart;
145 const Scalar* EIGEN_RESTRICT a1It = A1 + alignedStart;
146 const Scalar* EIGEN_RESTRICT a2It = A2 + alignedStart;
147 const Scalar* EIGEN_RESTRICT a3It = A3 + alignedStart;
148 const Scalar* EIGEN_RESTRICT rhsIt = rhs + alignedStart;
149 Scalar* EIGEN_RESTRICT resIt = res + alignedStart;
150 for (Index i = alignedStart; i < alignedEnd; i += PacketSize) {
151 Packet A0i = ploadu<Packet>(a0It);
152 a0It += PacketSize;
153 Packet A1i = ploadu<Packet>(a1It);
154 a1It += PacketSize;
155 Packet A2i = ploadu<Packet>(a2It);
156 a2It += PacketSize;
157 Packet A3i = ploadu<Packet>(a3It);
158 a3It += PacketSize;
159 Packet Bi = ploadu<Packet>(rhsIt);
160 rhsIt += PacketSize;
161 Packet Xi = pload<Packet>(resIt);
162
163 Xi = pcj0.pmadd(A0i, ptmp0, Xi);
164 Xi = pcj0.pmadd(A1i, ptmp1, Xi);
165 Xi = pcj0.pmadd(A2i, ptmp2, Xi);
166 Xi = pcj0.pmadd(A3i, ptmp3, Xi);
167 pstore(resIt, Xi);
168 resIt += PacketSize;
169
170 ptmp4 = pcj1.pmadd(A0i, Bi, ptmp4);
171 ptmp5 = pcj1.pmadd(A1i, Bi, ptmp5);
172 ptmp6 = pcj1.pmadd(A2i, Bi, ptmp6);
173 ptmp7 = pcj1.pmadd(A3i, Bi, ptmp7);
174 }
175
176 // Post-alignment scalar loop
177 for (Index i = alignedEnd; i < endi; ++i) {
178 res[i] += cj0.pmul(A0[i], t0) + cj0.pmul(A1[i], t1) + cj0.pmul(A2[i], t2) + cj0.pmul(A3[i], t3);
179 t4 += cj1.pmul(A0[i], rhs[i]);
180 t5 += cj1.pmul(A1[i], rhs[i]);
181 t6 += cj1.pmul(A2[i], rhs[i]);
182 t7 += cj1.pmul(A3[i], rhs[i]);
183 }
184
185 res[j] += alpha * (t4 + predux(ptmp4));
186 res[j + 1] += alpha * (t5 + predux(ptmp5));
187 res[j + 2] += alpha * (t6 + predux(ptmp6));
188 res[j + 3] += alpha * (t7 + predux(ptmp7));
189 }
190 }
191
192 // === Phase 2: 2 columns at a time ===
193 {
194 Index jStart = FirstTriangular ? (size - n4 - n2) : n4;
195 Index jEnd = FirstTriangular ? (size - n4) : (n4 + n2);
196
197 for (Index j = jStart; j < jEnd; j += 2) {
198 const Scalar* EIGEN_RESTRICT A0 = lhs + j * lhsStride;
199 const Scalar* EIGEN_RESTRICT A1 = lhs + (j + 1) * lhsStride;
200
201 Scalar t0 = cjAlpha * rhs[j];
202 Packet ptmp0 = pset1<Packet>(t0);
203 Scalar t1 = cjAlpha * rhs[j + 1];
204 Packet ptmp1 = pset1<Packet>(t1);
205
206 Scalar t2(0);
207 Packet ptmp2 = pzero(Packet{});
208 Scalar t3(0);
209 Packet ptmp3 = pzero(Packet{});
210
211 Index starti = FirstTriangular ? 0 : j + 2;
212 Index endi = FirstTriangular ? j : size;
213 Index alignedStart = starti + internal::first_default_aligned(&res[starti], endi - starti);
214 Index alignedEnd = alignedStart + ((endi - alignedStart) / PacketSize) * PacketSize;
215
216 res[j] += cjd.pmul(numext::real(A0[j]), t0);
217 res[j + 1] += cjd.pmul(numext::real(A1[j + 1]), t1);
218 EIGEN_IF_CONSTEXPR (FirstTriangular) {
219 res[j] += cj0.pmul(A1[j], t1);
220 t3 += cj1.pmul(A1[j], rhs[j]);
221 } else {
222 res[j + 1] += cj0.pmul(A0[j + 1], t0);
223 t2 += cj1.pmul(A0[j + 1], rhs[j + 1]);
224 }
225
226 for (Index i = starti; i < alignedStart; ++i) {
227 res[i] += cj0.pmul(A0[i], t0) + cj0.pmul(A1[i], t1);
228 t2 += cj1.pmul(A0[i], rhs[i]);
229 t3 += cj1.pmul(A1[i], rhs[i]);
230 }
231 const Scalar* EIGEN_RESTRICT a0It = A0 + alignedStart;
232 const Scalar* EIGEN_RESTRICT a1It = A1 + alignedStart;
233 const Scalar* EIGEN_RESTRICT rhsIt = rhs + alignedStart;
234 Scalar* EIGEN_RESTRICT resIt = res + alignedStart;
235 for (Index i = alignedStart; i < alignedEnd; i += PacketSize) {
236 Packet A0i = ploadu<Packet>(a0It);
237 a0It += PacketSize;
238 Packet A1i = ploadu<Packet>(a1It);
239 a1It += PacketSize;
240 Packet Bi = ploadu<Packet>(rhsIt);
241 rhsIt += PacketSize;
242 Packet Xi = pload<Packet>(resIt);
243
244 Xi = pcj0.pmadd(A0i, ptmp0, pcj0.pmadd(A1i, ptmp1, Xi));
245 ptmp2 = pcj1.pmadd(A0i, Bi, ptmp2);
246 ptmp3 = pcj1.pmadd(A1i, Bi, ptmp3);
247 pstore(resIt, Xi);
248 resIt += PacketSize;
249 }
250 for (Index i = alignedEnd; i < endi; i++) {
251 res[i] += cj0.pmul(A0[i], t0) + cj0.pmul(A1[i], t1);
252 t2 += cj1.pmul(A0[i], rhs[i]);
253 t3 += cj1.pmul(A1[i], rhs[i]);
254 }
255
256 res[j] += alpha * (t2 + predux(ptmp2));
257 res[j + 1] += alpha * (t3 + predux(ptmp3));
258 }
259 }
260
261 // === Phase 3: 1 column at a time ===
262 {
263 Index jStart = FirstTriangular ? 0 : (n4 + n2);
264 Index jEnd = FirstTriangular ? (size - n4 - n2) : size;
265
266 for (Index j = jStart; j < jEnd; j++) {
267 const Scalar* EIGEN_RESTRICT A0 = lhs + j * lhsStride;
268
269 Scalar t1 = cjAlpha * rhs[j];
270 Scalar t2(0);
271 Packet ptmp1 = pset1<Packet>(t1);
272 Packet ptmp2 = pzero(Packet{});
273
274 res[j] += cjd.pmul(numext::real(A0[j]), t1);
275
276 Index starti = FirstTriangular ? 0 : j + 1;
277 Index endi = FirstTriangular ? j : size;
278 Index alignedStart = starti + internal::first_default_aligned(&res[starti], endi - starti);
279 Index alignedEnd = alignedStart + ((endi - alignedStart) / PacketSize) * PacketSize;
280
281 for (Index i = starti; i < alignedStart; ++i) {
282 res[i] += cj0.pmul(A0[i], t1);
283 t2 += cj1.pmul(A0[i], rhs[i]);
284 }
285 const Scalar* EIGEN_RESTRICT a0It = A0 + alignedStart;
286 const Scalar* EIGEN_RESTRICT rhsIt = rhs + alignedStart;
287 Scalar* EIGEN_RESTRICT resIt = res + alignedStart;
288 for (Index i = alignedStart; i < alignedEnd; i += PacketSize) {
289 Packet A0i = ploadu<Packet>(a0It);
290 a0It += PacketSize;
291 Packet Bi = ploadu<Packet>(rhsIt);
292 rhsIt += PacketSize;
293 Packet Xi = pload<Packet>(resIt);
294
295 Xi = pcj0.pmadd(A0i, ptmp1, Xi);
296 pstore(resIt, Xi);
297 resIt += PacketSize;
298
299 ptmp2 = pcj1.pmadd(A0i, Bi, ptmp2);
300 }
301 for (Index i = alignedEnd; i < endi; i++) {
302 res[i] += cj0.pmul(A0[i], t1);
303 t2 += cj1.pmul(A0[i], rhs[i]);
304 }
305 res[j] += alpha * (t2 + predux(ptmp2));
306 }
307 }
308}
309
310} // end namespace internal
311
312/***************************************************************************
313 * Wrapper to product_selfadjoint_vector
314 ***************************************************************************/
315
316namespace internal {
317
318template <typename Lhs, int LhsMode, typename Rhs>
319struct selfadjoint_product_impl<Lhs, LhsMode, false, Rhs, 0, true> {
320 using Scalar = typename Product<Lhs, Rhs>::Scalar;
321
322 using LhsBlasTraits = internal::blas_traits<Lhs>;
323 using ActualLhsType = typename LhsBlasTraits::DirectLinearAccessType;
324 using ActualLhsTypeCleaned = internal::remove_all_t<ActualLhsType>;
325
326 using RhsBlasTraits = internal::blas_traits<Rhs>;
327 using ActualRhsType = typename RhsBlasTraits::DirectLinearAccessType;
328 using ActualRhsTypeCleaned = internal::remove_all_t<ActualRhsType>;
329
330 enum { LhsUpLo = LhsMode & (Upper | Lower) };
331
332 // Verify that the Rhs is a vector in the correct orientation.
333 // Otherwise, we break the assumption that we are multiplying
334 // MxN * Nx1.
335 static_assert(Rhs::ColsAtCompileTime == 1, "The RHS must be a column vector.");
336
337 template <typename Dest>
338 static EIGEN_DEVICE_FUNC void run(Dest& dest, const Lhs& a_lhs, const Rhs& a_rhs, const Scalar& alpha) {
339 using ResScalar = typename Dest::Scalar;
340 using RhsScalar = typename Rhs::Scalar;
341
342 eigen_assert(dest.rows() == a_lhs.rows() && dest.cols() == a_rhs.cols());
343
344 add_const_on_value_type_t<ActualLhsType> lhs = LhsBlasTraits::extract(a_lhs);
345 add_const_on_value_type_t<ActualRhsType> rhs = RhsBlasTraits::extract(a_rhs);
346
347 // Empty product, return early. Otherwise, we get `nullptr` use errors below when we try to access
348 // coeffRef(0,0).
349 if (lhs.size() == 0) return;
350
351 Scalar actualAlpha = combine_scalar_factors(alpha, a_lhs, a_rhs);
352
353 enum {
354 EvalToDest = (Dest::InnerStrideAtCompileTime == 1),
355 UseRhs = (ActualRhsTypeCleaned::InnerStrideAtCompileTime == 1)
356 };
357
358 internal::gemv_static_vector_if<ResScalar, Dest::SizeAtCompileTime, Dest::MaxSizeAtCompileTime, !EvalToDest>
359 static_dest;
360 internal::gemv_static_vector_if<RhsScalar, ActualRhsTypeCleaned::SizeAtCompileTime,
361 ActualRhsTypeCleaned::MaxSizeAtCompileTime, !UseRhs>
362 static_rhs;
363
364 ei_declare_aligned_stack_constructed_variable(ResScalar, actualDestPtr, dest.size(),
365 EvalToDest ? dest.data() : static_dest.data());
366
367 ei_declare_aligned_stack_constructed_variable(RhsScalar, actualRhsPtr, rhs.size(),
368 UseRhs ? const_cast<RhsScalar*>(rhs.data()) : static_rhs.data());
369
370 internal::gemv_prepare_destination<EvalToDest>(dest, actualDestPtr);
371 internal::gemv_prepare_rhs<UseRhs>(rhs, actualRhsPtr);
372
373 internal::selfadjoint_matrix_vector_product<
374 Scalar, Index, (internal::traits<ActualLhsTypeCleaned>::Flags & RowMajorBit) ? RowMajor : ColMajor,
375 int(LhsUpLo), bool(LhsBlasTraits::NeedToConjugate),
376 bool(RhsBlasTraits::NeedToConjugate)>::run(lhs.rows(), // size
377 &lhs.coeffRef(0, 0), lhs.outerStride(), // lhs info
378 actualRhsPtr, // rhs info
379 actualDestPtr, // result info
380 actualAlpha // scale factor
381 );
382
383 internal::gemv_copy_destination<EvalToDest>(dest, actualDestPtr);
384 }
385};
386
387template <typename Lhs, typename Rhs, int RhsMode>
388struct selfadjoint_product_impl<Lhs, 0, true, Rhs, RhsMode, false> {
389 using Scalar = typename Product<Lhs, Rhs>::Scalar;
390 enum { RhsUpLo = RhsMode & (Upper | Lower) };
391
392 template <typename Dest>
393 static void run(Dest& dest, const Lhs& a_lhs, const Rhs& a_rhs, const Scalar& alpha) {
394 // let's simply transpose the product
395 Transpose<Dest> destT(dest);
396 selfadjoint_product_impl<Transpose<const Rhs>, int(RhsUpLo) == Upper ? Lower : Upper, false, Transpose<const Lhs>,
397 0, true>::run(destT, a_rhs.transpose(), a_lhs.transpose(), alpha);
398 }
399};
400
401} // end namespace internal
402
403} // end namespace Eigen
404
405#endif // EIGEN_SELFADJOINT_MATRIX_VECTOR_H
@ 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