Eigen  5.0.1
 
Loading...
Searching...
No Matches
GeneralMatrixVector.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2008-2016 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_GENERAL_MATRIX_VECTOR_H
12#define EIGEN_GENERAL_MATRIX_VECTOR_H
13
14// IWYU pragma: private
15#include "../InternalHeaderCheck.h"
16
17// C4804: unsafe use of type 'bool' in operation. Unavoidable in generic code
18// instantiated with bool scalars (e.g. += and * on bool).
19#if EIGEN_COMP_MSVC
20#pragma warning(push)
21#pragma warning(disable : 4804)
22#endif
23
24namespace Eigen {
25
26namespace internal {
27
28template <typename LhsScalar, typename RhsScalar, int PacketSize_ = GEBPPacketFull>
29class gemv_traits {
30 using ResScalar = typename ScalarBinaryOpTraits<LhsScalar, RhsScalar>::ReturnType;
31
32#define PACKET_DECL_COND_POSTFIX(postfix, name, packet_size) \
33 typedef typename packet_conditional< \
34 packet_size, typename packet_traits<name##Scalar>::type, typename packet_traits<name##Scalar>::half, \
35 typename unpacket_traits<typename packet_traits<name##Scalar>::half>::half>::type name##Packet##postfix
36
37 PACKET_DECL_COND_POSTFIX(_, Lhs, PacketSize_);
38 PACKET_DECL_COND_POSTFIX(_, Rhs, PacketSize_);
39 PACKET_DECL_COND_POSTFIX(_, Res, PacketSize_);
40#undef PACKET_DECL_COND_POSTFIX
41
42 public:
43 enum {
44 Vectorizable = unpacket_traits<LhsPacket_>::vectorizable && unpacket_traits<RhsPacket_>::vectorizable &&
45 int(unpacket_traits<LhsPacket_>::size) == int(unpacket_traits<RhsPacket_>::size),
46 LhsPacketSize = Vectorizable ? unpacket_traits<LhsPacket_>::size : 1,
47 RhsPacketSize = Vectorizable ? unpacket_traits<RhsPacket_>::size : 1,
48 ResPacketSize = Vectorizable ? unpacket_traits<ResPacket_>::size : 1
49 };
50
51 using LhsPacket = std::conditional_t<Vectorizable, LhsPacket_, LhsScalar>;
52 using RhsPacket = std::conditional_t<Vectorizable, RhsPacket_, RhsScalar>;
53 using ResPacket = std::conditional_t<Vectorizable, ResPacket_, ResScalar>;
54};
55
56// Whether a mapper's coefficients are addressable and consecutive along the given storage order: the BLAS mappers
57// with unit inner stride. Other mappers (e.g. the Tensor contraction input mappers) return coefficients by value.
58template <typename Mapper, int Order>
59struct gemv_mapper_is_contiguous : std::false_type {};
60template <typename Scalar, typename Index, int Order, int Alignment>
61struct gemv_mapper_is_contiguous<blas_data_mapper<Scalar, Index, Order, Alignment, 1>, Order> : std::true_type {};
62template <typename Scalar, typename Index, int Order>
63struct gemv_mapper_is_contiguous<const_blas_data_mapper<Scalar, Index, Order>, Order> : std::true_type {};
64
65// Whether the GEMV kernels finish a partial packet with masked segment loads. Segment loads zero the lanes outside
66// the segment, which the row-major kernel relies on before its horizontal reduction.
67template <typename LhsPacket, typename RhsPacket, typename ResPacket>
68using gemv_use_packet_segment =
69 bool_constant<has_packet_segment<ResPacket>::value && std::is_same<LhsPacket, ResPacket>::value &&
70 std::is_same<RhsPacket, ResPacket>::value>;
71
72// Loads count < packet-size consecutive coefficients starting at (i, j), along Order. Only contiguous mappers take
73// the masked load; the kernels never select segments for other mappers, whose branch merely has to compile.
74template <typename Packet, int Order, typename Mapper,
75 bool Contiguous = gemv_mapper_is_contiguous<Mapper, Order>::value>
76struct gemv_segment_loader {
77 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet run(const Mapper&, Index, Index, Index) {
78 eigen_internal_assert(false && "segment load from a non-contiguous mapper");
79 return pzero(Packet{});
80 }
81};
82template <typename Packet, int Order, typename Mapper>
83struct gemv_segment_loader<Packet, Order, Mapper, true> {
84 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet run(const Mapper& m, Index i, Index j, Index count) {
85 return ploaduSegment<Packet>(&m(i, j), 0, count);
86 }
87};
88
89/* Optimized col-major matrix * vector product:
90 * This algorithm processes the matrix per vertical panels,
91 * which are then processed horizontally per chunk of 8*PacketSize x 1 vertical segments.
92 *
93 * Mixing type logic: C += alpha * A * B
94 * | A | B |alpha| comments
95 * |real |cplx |cplx | no vectorization
96 * |real |cplx |real | alpha is converted to a cplx when calling the run function, no vectorization
97 * |cplx |real |cplx | invalid, the caller has to do tmp: = A * B; C += alpha*tmp
98 * |cplx |real |real | optimal case, vectorization possible via real-cplx mul
99 *
100 * The same reasoning applies for the transposed case.
101 */
102template <typename Index, typename LhsScalar, typename LhsMapper, bool ConjugateLhs, typename RhsScalar,
103 typename RhsMapper, bool ConjugateRhs, int Version>
104struct general_matrix_vector_product<Index, LhsScalar, LhsMapper, ColMajor, ConjugateLhs, RhsScalar, RhsMapper,
105 ConjugateRhs, Version> {
106 using Traits = gemv_traits<LhsScalar, RhsScalar>;
107 using HalfTraits = gemv_traits<LhsScalar, RhsScalar, GEBPPacketHalf>;
108 using QuarterTraits = gemv_traits<LhsScalar, RhsScalar, GEBPPacketQuarter>;
109
110 using ResScalar = typename ScalarBinaryOpTraits<LhsScalar, RhsScalar>::ReturnType;
111
112 using LhsPacket = typename Traits::LhsPacket;
113 using RhsPacket = typename Traits::RhsPacket;
114 using ResPacket = typename Traits::ResPacket;
115
116 using LhsPacketHalf = typename HalfTraits::LhsPacket;
117 using RhsPacketHalf = typename HalfTraits::RhsPacket;
118 using ResPacketHalf = typename HalfTraits::ResPacket;
119
120 using LhsPacketQuarter = typename QuarterTraits::LhsPacket;
121 using RhsPacketQuarter = typename QuarterTraits::RhsPacket;
122 using ResPacketQuarter = typename QuarterTraits::ResPacket;
123
124 EIGEN_DEVICE_FUNC inline static void run(Index rows, Index cols, const LhsMapper& lhs, const RhsMapper& rhs,
125 ResScalar* res, Index resIncr, RhsScalar alpha);
126
127 // Processes N full packets of rows starting at row i and, if Segment, the next count < ResPacketSize rows as one
128 // masked packet, all in a single pass over the columns [j2, jend).
129 template <int N, bool Segment = false>
130 EIGEN_DEVICE_FUNC static EIGEN_ALWAYS_INLINE void process_rows(
131 Index i, Index j2, Index jend, const LhsMapper& lhs, const RhsMapper& rhs, ResScalar* res,
132 const ResPacket& palpha, conj_helper<LhsPacket, RhsPacket, ConjugateLhs, ConjugateRhs>& pcj, Index count = 0);
133
134 // Finishes the rows from i on, full_packets < 10 full packets followed by 0 < count < ResPacketSize rows, in one
135 // process_rows pass. Out of line: it runs at most once per column block, and only when rows is not a multiple of
136 // the packet size. lhs is taken by value: a reference would make run() keep its copy in memory, which clang fills
137 // with a load that cannot be forwarded from the caller's stores.
138 EIGEN_DEVICE_FUNC static EIGEN_DONT_INLINE void process_segment_tail(
139 std::true_type, Index full_packets, Index i, Index j2, Index jend, LhsMapper lhs, const RhsMapper& rhs,
140 ResScalar* res, const ResPacket& palpha, conj_helper<LhsPacket, RhsPacket, ConjugateLhs, ConjugateRhs>& pcj,
141 Index count);
142 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void process_segment_tail(
143 std::false_type, Index, Index, Index, Index, const LhsMapper&, const RhsMapper&, ResScalar*, const ResPacket&,
144 conj_helper<LhsPacket, RhsPacket, ConjugateLhs, ConjugateRhs>&, Index) {}
145};
146
147// Integer-sequence helper for col-major GEMV full-packet row blocks.
148template <int N>
149struct gemv_colmajor_unroller {
150 template <typename Packet, int... K>
151 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void init_zero_impl(std::integer_sequence<int, K...>, Packet* c) {
152 int unused[] = {0, ((c[K] = pzero(Packet{})), 0)...};
153 EIGEN_UNUSED_VARIABLE(unused);
154 }
155
156 template <typename Packet>
157 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void init_zero(Packet* c) {
158 init_zero_impl(std::make_integer_sequence<int, N>{}, c);
159 }
160
161 template <typename LhsPacket, int LhsStride, int Alignment, typename AccPacket, typename RhsPacket,
162 typename ConjHelper, typename LhsMapper, typename Index, int... K>
163 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void madd_impl(std::integer_sequence<int, K...>, AccPacket* c,
164 const LhsMapper& lhs, Index i, Index j,
165 const RhsPacket& b0, ConjHelper& pcj) {
166 int unused[] = {
167 0, ((c[K] = pcj.pmadd(lhs.template load<LhsPacket, Alignment>(i + LhsStride * K, j), b0, c[K])), 0)...};
168 EIGEN_UNUSED_VARIABLE(unused);
169 }
170
171 template <typename LhsPacket, int LhsStride, int Alignment, typename AccPacket, typename RhsPacket,
172 typename ConjHelper, typename LhsMapper, typename Index>
173 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void madd(AccPacket* c, const LhsMapper& lhs, Index i, Index j,
174 const RhsPacket& b0, ConjHelper& pcj) {
175 madd_impl<LhsPacket, LhsStride, Alignment>(std::make_integer_sequence<int, N>{}, c, lhs, i, j, b0, pcj);
176 }
177
178 template <int K, typename ResPacket, int ResStride, typename ResScalar, typename Index>
179 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void store_one(const ResPacket* c, ResScalar* res, Index i,
180 const ResPacket& palpha) {
181 ResScalar* r = res + i + ResStride * K;
182 pstoreu(r, pmadd(c[K], palpha, ploadu<ResPacket>(r)));
183 }
184
185 template <typename ResPacket, int ResStride, typename ResScalar, typename Index, int... K>
186 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void store_impl(std::integer_sequence<int, K...>, const ResPacket* c,
187 ResScalar* res, Index i, const ResPacket& palpha) {
188 int unused[] = {0, (store_one<K, ResPacket, ResStride>(c, res, i, palpha), 0)...};
189 EIGEN_UNUSED_VARIABLE(unused);
190 }
191
192 template <typename ResPacket, int ResStride, typename ResScalar, typename Index>
193 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void store(const ResPacket* c, ResScalar* res, Index i,
194 const ResPacket& palpha) {
195 store_impl<ResPacket, ResStride>(std::make_integer_sequence<int, N>{}, c, res, i, palpha);
196 }
197};
198
199// No full packets: the final pass may hold only the masked partial packet.
200template <>
201struct gemv_colmajor_unroller<0> {
202 template <typename Packet>
203 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void init_zero(Packet*) {}
204
205 template <typename LhsPacket, int LhsStride, int Alignment, typename AccPacket, typename RhsPacket,
206 typename ConjHelper, typename LhsMapper, typename Index>
207 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void madd(AccPacket*, const LhsMapper&, Index, Index, const RhsPacket&,
208 ConjHelper&) {}
209
210 template <typename ResPacket, int ResStride, typename ResScalar, typename Index>
211 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void store(const ResPacket*, ResScalar*, Index, const ResPacket&) {}
212};
213
214template <typename Index, typename LhsScalar, typename LhsMapper, bool ConjugateLhs, typename RhsScalar,
215 typename RhsMapper, bool ConjugateRhs, int Version>
216template <int N, bool Segment>
217EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void
218general_matrix_vector_product<Index, LhsScalar, LhsMapper, ColMajor, ConjugateLhs, RhsScalar, RhsMapper, ConjugateRhs,
219 Version>::process_rows(Index i, Index j2, Index jend, const LhsMapper& lhs,
220 const RhsMapper& rhs, ResScalar* res, const ResPacket& palpha,
221 conj_helper<LhsPacket, RhsPacket, ConjugateLhs, ConjugateRhs>& pcj,
222 Index count) {
223 enum { LhsAlignment = Unaligned, LhsPacketSize = Traits::LhsPacketSize, ResPacketSize = Traits::ResPacketSize };
224 using Unroller = gemv_colmajor_unroller<N>;
225 const Index iseg = i + N * ResPacketSize;
226
227 // c_seg accumulates the masked partial packet when Segment is set. It is not part of c: a wider array grows the
228 // stack frame GCC estimates for run() and stops it from being inlined.
229 ResPacket c[N > 0 ? N : 1];
230 Unroller::init_zero(c);
231 ResPacket c_seg = pzero(ResPacket{});
232 for (Index j = j2; j < jend; ++j) {
233 RhsPacket b0 = pset1<RhsPacket>(rhs(j, 0));
234 Unroller::template madd<LhsPacket, LhsPacketSize, LhsAlignment>(c, lhs, i, j, b0, pcj);
235 EIGEN_IF_CONSTEXPR (Segment) {
236 c_seg = pcj.pmadd(gemv_segment_loader<LhsPacket, ColMajor, LhsMapper>::run(lhs, iseg, j, count), b0, c_seg);
237 }
238 }
239 Unroller::template store<ResPacket, ResPacketSize>(c, res, i, palpha);
240 EIGEN_IF_CONSTEXPR (Segment) {
241 pstoreuSegment(res + iseg, pmadd(c_seg, palpha, ploaduSegment<ResPacket>(res + iseg, 0, count)), 0, count);
242 }
243}
244
245template <typename Index, typename LhsScalar, typename LhsMapper, bool ConjugateLhs, typename RhsScalar,
246 typename RhsMapper, bool ConjugateRhs, int Version>
247EIGEN_DEVICE_FUNC EIGEN_DONT_INLINE void general_matrix_vector_product<
248 Index, LhsScalar, LhsMapper, ColMajor, ConjugateLhs, RhsScalar, RhsMapper, ConjugateRhs,
249 Version>::process_segment_tail(std::true_type, Index full_packets, Index i, Index j2, Index jend, LhsMapper lhs,
250 const RhsMapper& rhs, ResScalar* res, const ResPacket& palpha,
251 conj_helper<LhsPacket, RhsPacket, ConjugateLhs, ConjugateRhs>& pcj, Index count) {
252#define EIGEN_GEMV_PROCESS_ROW(n) \
253 case n: \
254 process_rows<n, true>(i, j2, jend, lhs, rhs, res, palpha, pcj, count); \
255 break
256 switch (full_packets) {
257 EIGEN_GEMV_PROCESS_ROW(0);
258 EIGEN_GEMV_PROCESS_ROW(1);
259 EIGEN_GEMV_PROCESS_ROW(2);
260 EIGEN_GEMV_PROCESS_ROW(3);
261 EIGEN_GEMV_PROCESS_ROW(4);
262 EIGEN_GEMV_PROCESS_ROW(5);
263 EIGEN_GEMV_PROCESS_ROW(6);
264 EIGEN_GEMV_PROCESS_ROW(7);
265 EIGEN_GEMV_PROCESS_ROW(8);
266 EIGEN_GEMV_PROCESS_ROW(9);
267 default:
268 eigen_internal_assert(false);
269 break;
270 }
271#undef EIGEN_GEMV_PROCESS_ROW
272}
273
274template <typename Index, typename LhsScalar, typename LhsMapper, bool ConjugateLhs, typename RhsScalar,
275 typename RhsMapper, bool ConjugateRhs, int Version>
276EIGEN_DEVICE_FUNC inline void
277general_matrix_vector_product<Index, LhsScalar, LhsMapper, ColMajor, ConjugateLhs, RhsScalar, RhsMapper, ConjugateRhs,
278 Version>::run(Index rows, Index cols, const LhsMapper& alhs, const RhsMapper& rhs,
279 ResScalar* res, Index resIncr, RhsScalar alpha) {
280 EIGEN_UNUSED_VARIABLE(resIncr);
281 eigen_internal_assert(resIncr == 1);
282
283 // BLAS contract: if alpha == 0, the result is unchanged (and lhs/rhs need not be read).
284 if (numext::is_exactly_zero(alpha)) return;
285
286 // The following copy tells the compiler that lhs's attributes are not modified outside this function
287 // This helps GCC to generate proper code.
288 LhsMapper lhs(alhs);
289
290 conj_helper<LhsScalar, RhsScalar, ConjugateLhs, ConjugateRhs> cj;
291 conj_helper<LhsPacket, RhsPacket, ConjugateLhs, ConjugateRhs> pcj;
292 conj_helper<LhsPacketHalf, RhsPacketHalf, ConjugateLhs, ConjugateRhs> pcj_half;
293 conj_helper<LhsPacketQuarter, RhsPacketQuarter, ConjugateLhs, ConjugateRhs> pcj_quarter;
294
295 const Index lhsStride = lhs.stride();
296 // LhsAlignment stays Unaligned; enabling aligned reads would require
297 // propagating the Mapper's Alignment through the run() template, and on
298 // modern x86 aligned/unaligned packet loads are equivalent anyway.
299 enum {
300 LhsAlignment = Unaligned,
301 ResPacketSize = Traits::ResPacketSize,
302 ResPacketSizeHalf = HalfTraits::ResPacketSize,
303 ResPacketSizeQuarter = QuarterTraits::ResPacketSize,
304 LhsPacketSize = Traits::LhsPacketSize,
305 HasHalf = (int)ResPacketSizeHalf < (int)ResPacketSize,
306 HasQuarter = (int)ResPacketSizeQuarter < (int)ResPacketSizeHalf,
307 UseSegment = gemv_use_packet_segment<LhsPacket, RhsPacket, ResPacket>::value &&
308 gemv_mapper_is_contiguous<LhsMapper, ColMajor>::value
309 };
310
311 using UnsignedIndex = std::make_unsigned_t<Index>;
312 // With segments, the count trailing rows of a partial packet take one masked pass together with the full packets
313 // before them. The 8-packet blocks then stop while at least 2 full packets remain, so that pass is never a lone
314 // packet whose accumulator chain is bound by the FMA latency. Exact multiples of the packet size keep the 4, 3, 2
315 // and 1-packet passes.
316 const Index count = UseSegment ? Index(UnsignedIndex(rows) % ResPacketSize) : Index(0);
317 const Index n8 = rows - (count > 0 ? 10 : 8) * ResPacketSize + 1;
318 const Index n4 = rows - 4 * ResPacketSize + 1;
319 const Index n3 = rows - 3 * ResPacketSize + 1;
320 const Index n2 = rows - 2 * ResPacketSize + 1;
321 const Index n1 = rows - 1 * ResPacketSize + 1;
322 const Index n_half = rows - 1 * ResPacketSizeHalf + 1;
323 const Index n_quarter = rows - 1 * ResPacketSizeQuarter + 1;
324
325 // Choose block_cols so that one column slice of the LHS roughly fits in L1.
326 // When it does not, fall back to a smaller batch to keep cache pressure down.
327 std::ptrdiff_t l1, l2, l3;
328 manage_caching_sizes(GetAction, &l1, &l2, &l3);
329 const Index block_cols =
330 cols < 128 ? cols : (lhsStride * Index(sizeof(LhsScalar)) < Index(l1) ? Index(16) : Index(4));
331 ResPacket palpha = pset1<ResPacket>(alpha);
332 ResPacketHalf palpha_half = pset1<ResPacketHalf>(alpha);
333 ResPacketQuarter palpha_quarter = pset1<ResPacketQuarter>(alpha);
334
335 for (Index j2 = 0; j2 < cols; j2 += block_cols) {
336 Index jend = numext::mini(j2 + block_cols, cols);
337 Index i = 0;
338 for (; i < n8; i += ResPacketSize * 8) process_rows<8>(i, j2, jend, lhs, rhs, res, palpha, pcj);
339 if (count > 0) {
340 process_segment_tail(bool_constant<UseSegment>(), Index(UnsignedIndex(rows - i) / ResPacketSize), i, j2, jend,
341 lhs, rhs, res, palpha, pcj, count);
342 } else {
343#define EIGEN_GEMV_PROCESS_ROW(k) \
344 if (i < n##k) { \
345 process_rows<k>(i, j2, jend, lhs, rhs, res, palpha, pcj); \
346 i += ResPacketSize * (k); \
347 } \
348 static_assert(true, "Trailing semicolon required")
349 EIGEN_GEMV_PROCESS_ROW(4);
350 EIGEN_GEMV_PROCESS_ROW(3);
351 EIGEN_GEMV_PROCESS_ROW(2);
352 EIGEN_GEMV_PROCESS_ROW(1);
353#undef EIGEN_GEMV_PROCESS_ROW
354 EIGEN_IF_CONSTEXPR (HasHalf) {
355 if (i < n_half) {
356 ResPacketHalf c0 = pzero(ResPacketHalf{});
357 for (Index j = j2; j < jend; j += 1) {
358 RhsPacketHalf b0 = pset1<RhsPacketHalf>(rhs(j, 0));
359 c0 = pcj_half.pmadd(lhs.template load<LhsPacketHalf, LhsAlignment>(i + 0, j), b0, c0);
360 }
361 pstoreu(res + i + ResPacketSizeHalf * 0,
362 pmadd(c0, palpha_half, ploadu<ResPacketHalf>(res + i + ResPacketSizeHalf * 0)));
363 i += ResPacketSizeHalf;
364 }
365 }
366 EIGEN_IF_CONSTEXPR (HasQuarter) {
367 if (i < n_quarter) {
368 ResPacketQuarter c0 = pzero(ResPacketQuarter{});
369 for (Index j = j2; j < jend; j += 1) {
370 RhsPacketQuarter b0 = pset1<RhsPacketQuarter>(rhs(j, 0));
371 c0 = pcj_quarter.pmadd(lhs.template load<LhsPacketQuarter, LhsAlignment>(i + 0, j), b0, c0);
372 }
373 pstoreu(res + i + ResPacketSizeQuarter * 0,
374 pmadd(c0, palpha_quarter, ploadu<ResPacketQuarter>(res + i + ResPacketSizeQuarter * 0)));
375 i += ResPacketSizeQuarter;
376 }
377 }
378 for (; i < rows; ++i) {
379 ResScalar c0(0);
380 for (Index j = j2; j < jend; j += 1) c0 += cj.pmul(lhs(i, j), rhs(j, 0));
381 res[i] += alpha * c0;
382 }
383 }
384 }
385}
386
387/* Optimized row-major matrix * vector product:
388 * This algorithm processes 4 rows at once that allows to both reduce
389 * the number of load/stores of the result by a factor 4 and to reduce
390 * the instruction dependency. Moreover, we know that all bands have the
391 * same alignment pattern.
392 *
393 * Mixing type logic:
394 * - alpha is always a complex (or converted to a complex)
395 * - no vectorization
396 */
397template <typename Index, typename LhsScalar, typename LhsMapper, bool ConjugateLhs, typename RhsScalar,
398 typename RhsMapper, bool ConjugateRhs, int Version>
399struct general_matrix_vector_product<Index, LhsScalar, LhsMapper, RowMajor, ConjugateLhs, RhsScalar, RhsMapper,
400 ConjugateRhs, Version> {
401 using Traits = gemv_traits<LhsScalar, RhsScalar>;
402 using HalfTraits = gemv_traits<LhsScalar, RhsScalar, GEBPPacketHalf>;
403 using QuarterTraits = gemv_traits<LhsScalar, RhsScalar, GEBPPacketQuarter>;
404
405 using ResScalar = typename ScalarBinaryOpTraits<LhsScalar, RhsScalar>::ReturnType;
406
407 using LhsPacket = typename Traits::LhsPacket;
408 using RhsPacket = typename Traits::RhsPacket;
409 using ResPacket = typename Traits::ResPacket;
410
411 using LhsPacketHalf = typename HalfTraits::LhsPacket;
412 using RhsPacketHalf = typename HalfTraits::RhsPacket;
413 using ResPacketHalf = typename HalfTraits::ResPacket;
414
415 using LhsPacketQuarter = typename QuarterTraits::LhsPacket;
416 using RhsPacketQuarter = typename QuarterTraits::RhsPacket;
417 using ResPacketQuarter = typename QuarterTraits::ResPacket;
418
419 EIGEN_DEVICE_FUNC static inline void run(Index rows, Index cols, const LhsMapper& lhs, const RhsMapper& rhs,
420 ResScalar* res, Index resIncr, ResScalar alpha);
421
422 // Specialized path for when cols < full packet size.
423 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE static void run_small_cols(Index rows, Index cols, const LhsMapper& lhs,
424 const RhsMapper& rhs, ResScalar* res, Index resIncr,
425 ResScalar alpha);
426
427 // Templated helper that processes N rows in run_small_cols. N is a compile-time
428 // constant; row-dimension unrolling is done inside flat helper loops.
429 template <int N>
430 EIGEN_DEVICE_FUNC static EIGEN_ALWAYS_INLINE void process_rows_small_cols(Index i, Index cols, const LhsMapper& lhs,
431 const RhsMapper& rhs, ResScalar* res,
432 Index resIncr, ResScalar alpha,
433 Index halfColBlockEnd,
434 Index quarterColBlockEnd);
435};
436
437template <typename Index, typename LhsScalar, typename LhsMapper, bool ConjugateLhs, typename RhsScalar,
438 typename RhsMapper, bool ConjugateRhs, int Version>
439EIGEN_DEVICE_FUNC inline void
440general_matrix_vector_product<Index, LhsScalar, LhsMapper, RowMajor, ConjugateLhs, RhsScalar, RhsMapper, ConjugateRhs,
441 Version>::run(Index rows, Index cols, const LhsMapper& alhs, const RhsMapper& rhs,
442 ResScalar* res, Index resIncr, ResScalar alpha) {
443 // BLAS contract: if alpha == 0, the result is unchanged (and lhs/rhs need not be read).
444 if (numext::is_exactly_zero(alpha)) return;
445
446 // When cols < full packet size, the main vectorized loops are empty.
447 // Use the sub-packet helper only when half or quarter packets can do useful work;
448 // otherwise it would just duplicate the scalar cleanup.
449 enum {
450 LhsPacketSize_ = Traits::LhsPacketSize,
451 MinUsefulCols_ =
452 ((int)QuarterTraits::LhsPacketSize < (int)HalfTraits::LhsPacketSize)
453 ? (int)QuarterTraits::LhsPacketSize
454 : (((int)HalfTraits::LhsPacketSize < (int)Traits::LhsPacketSize) ? (int)HalfTraits::LhsPacketSize
455 : (int)Traits::LhsPacketSize),
456 HasSubPackets_ = (int)MinUsefulCols_ < (int)LhsPacketSize_,
457 UseSegment_ = gemv_use_packet_segment<LhsPacket, RhsPacket, ResPacket>::value &&
458 gemv_mapper_is_contiguous<LhsMapper, RowMajor>::value &&
459 gemv_mapper_is_contiguous<RhsMapper, ColMajor>::value,
460 // With segments, one masked full packet per row beats a half packet followed by two or more scalar columns.
461 SmallColsEnd_ = UseSegment_ ? (int)HalfTraits::LhsPacketSize + 2 : (int)LhsPacketSize_
462 };
463 EIGEN_IF_CONSTEXPR (HasSubPackets_) {
464 if (cols >= MinUsefulCols_) {
465 if (cols < SmallColsEnd_ && cols < LhsPacketSize_) {
466 run_small_cols(rows, cols, alhs, rhs, res, resIncr, alpha);
467 return;
468 }
469 }
470 }
471
472 // The following copy tells the compiler that lhs's attributes are not modified outside this function
473 // This helps GCC to generate proper code.
474 LhsMapper lhs(alhs);
475
476 eigen_internal_assert(rhs.stride() == 1);
477 conj_helper<LhsScalar, RhsScalar, ConjugateLhs, ConjugateRhs> cj;
478 conj_helper<LhsPacket, RhsPacket, ConjugateLhs, ConjugateRhs> pcj;
479 conj_helper<LhsPacketHalf, RhsPacketHalf, ConjugateLhs, ConjugateRhs> pcj_half;
480 conj_helper<LhsPacketQuarter, RhsPacketQuarter, ConjugateLhs, ConjugateRhs> pcj_quarter;
481
482 // Disable the 8-row inner unroll once a single column slice no longer fits in L1; with very
483 // large LHS strides each unrolled iteration evicts the previously-loaded rows from cache.
484 std::ptrdiff_t l1, l2, l3;
485 manage_caching_sizes(GetAction, &l1, &l2, &l3);
486 const Index n8 = lhs.stride() * Index(sizeof(LhsScalar)) > Index(l1) ? 0 : rows - 7;
487 const Index n4 = rows - 3;
488 const Index n2 = rows - 1;
489
490 // LhsAlignment stays Unaligned; enabling aligned reads would require
491 // propagating the Mapper's Alignment through the run() template, and on
492 // modern x86 aligned/unaligned packet loads are equivalent anyway.
493 enum {
494 LhsAlignment = Unaligned,
495 ResPacketSize = Traits::ResPacketSize,
496 ResPacketSizeHalf = HalfTraits::ResPacketSize,
497 ResPacketSizeQuarter = QuarterTraits::ResPacketSize,
498 LhsPacketSize = Traits::LhsPacketSize,
499 LhsPacketSizeHalf = HalfTraits::LhsPacketSize,
500 LhsPacketSizeQuarter = QuarterTraits::LhsPacketSize,
501 HasHalf = (int)ResPacketSizeHalf < (int)ResPacketSize,
502 HasQuarter = (int)ResPacketSizeQuarter < (int)ResPacketSizeHalf,
503 UseSegment = UseSegment_
504 };
505
506 using UnsignedIndex = std::make_unsigned_t<Index>;
507 const Index fullColBlockEnd = LhsPacketSize * (UnsignedIndex(cols) / LhsPacketSize);
508 const Index halfColBlockEnd = LhsPacketSizeHalf * (UnsignedIndex(cols) / LhsPacketSizeHalf);
509 const Index quarterColBlockEnd = LhsPacketSizeQuarter * (UnsignedIndex(cols) / LhsPacketSizeQuarter);
510 // With segments, the last cols - fullColBlockEnd < LhsPacketSize columns take one masked packet per row and the
511 // scalar column loops below are empty.
512 const Index segmentCount = cols - fullColBlockEnd;
513 const Index scalarColStart = UseSegment ? cols : fullColBlockEnd;
514 using LhsSegmentLoader = gemv_segment_loader<LhsPacket, RowMajor, LhsMapper>;
515 using RhsSegmentLoader = gemv_segment_loader<RhsPacket, ColMajor, RhsMapper>;
516
517 Index i = 0;
518 for (; i < n8; i += 8) {
519 ResPacket c0 = pzero(ResPacket{}), c1 = pzero(ResPacket{}), c2 = pzero(ResPacket{}), c3 = pzero(ResPacket{}),
520 c4 = pzero(ResPacket{}), c5 = pzero(ResPacket{}), c6 = pzero(ResPacket{}), c7 = pzero(ResPacket{});
521
522 for (Index j = 0; j < fullColBlockEnd; j += LhsPacketSize) {
523 RhsPacket b0 = rhs.template load<RhsPacket, Unaligned>(j, 0);
524
525 c0 = pcj.pmadd(lhs.template load<LhsPacket, LhsAlignment>(i + 0, j), b0, c0);
526 c1 = pcj.pmadd(lhs.template load<LhsPacket, LhsAlignment>(i + 1, j), b0, c1);
527 c2 = pcj.pmadd(lhs.template load<LhsPacket, LhsAlignment>(i + 2, j), b0, c2);
528 c3 = pcj.pmadd(lhs.template load<LhsPacket, LhsAlignment>(i + 3, j), b0, c3);
529 c4 = pcj.pmadd(lhs.template load<LhsPacket, LhsAlignment>(i + 4, j), b0, c4);
530 c5 = pcj.pmadd(lhs.template load<LhsPacket, LhsAlignment>(i + 5, j), b0, c5);
531 c6 = pcj.pmadd(lhs.template load<LhsPacket, LhsAlignment>(i + 6, j), b0, c6);
532 c7 = pcj.pmadd(lhs.template load<LhsPacket, LhsAlignment>(i + 7, j), b0, c7);
533 }
534 if (UseSegment && segmentCount > 0) {
535 RhsPacket b0 = RhsSegmentLoader::run(rhs, fullColBlockEnd, 0, segmentCount);
536 c0 = pcj.pmadd(LhsSegmentLoader::run(lhs, i + 0, fullColBlockEnd, segmentCount), b0, c0);
537 c1 = pcj.pmadd(LhsSegmentLoader::run(lhs, i + 1, fullColBlockEnd, segmentCount), b0, c1);
538 c2 = pcj.pmadd(LhsSegmentLoader::run(lhs, i + 2, fullColBlockEnd, segmentCount), b0, c2);
539 c3 = pcj.pmadd(LhsSegmentLoader::run(lhs, i + 3, fullColBlockEnd, segmentCount), b0, c3);
540 c4 = pcj.pmadd(LhsSegmentLoader::run(lhs, i + 4, fullColBlockEnd, segmentCount), b0, c4);
541 c5 = pcj.pmadd(LhsSegmentLoader::run(lhs, i + 5, fullColBlockEnd, segmentCount), b0, c5);
542 c6 = pcj.pmadd(LhsSegmentLoader::run(lhs, i + 6, fullColBlockEnd, segmentCount), b0, c6);
543 c7 = pcj.pmadd(LhsSegmentLoader::run(lhs, i + 7, fullColBlockEnd, segmentCount), b0, c7);
544 }
545 ResScalar cc0 = predux(c0);
546 ResScalar cc1 = predux(c1);
547 ResScalar cc2 = predux(c2);
548 ResScalar cc3 = predux(c3);
549 ResScalar cc4 = predux(c4);
550 ResScalar cc5 = predux(c5);
551 ResScalar cc6 = predux(c6);
552 ResScalar cc7 = predux(c7);
553
554 for (Index j = scalarColStart; j < cols; ++j) {
555 RhsScalar b0 = rhs(j, 0);
556
557 cc0 += cj.pmul(lhs(i + 0, j), b0);
558 cc1 += cj.pmul(lhs(i + 1, j), b0);
559 cc2 += cj.pmul(lhs(i + 2, j), b0);
560 cc3 += cj.pmul(lhs(i + 3, j), b0);
561 cc4 += cj.pmul(lhs(i + 4, j), b0);
562 cc5 += cj.pmul(lhs(i + 5, j), b0);
563 cc6 += cj.pmul(lhs(i + 6, j), b0);
564 cc7 += cj.pmul(lhs(i + 7, j), b0);
565 }
566 res[(i + 0) * resIncr] += alpha * cc0;
567 res[(i + 1) * resIncr] += alpha * cc1;
568 res[(i + 2) * resIncr] += alpha * cc2;
569 res[(i + 3) * resIncr] += alpha * cc3;
570 res[(i + 4) * resIncr] += alpha * cc4;
571 res[(i + 5) * resIncr] += alpha * cc5;
572 res[(i + 6) * resIncr] += alpha * cc6;
573 res[(i + 7) * resIncr] += alpha * cc7;
574 }
575 for (; i < n4; i += 4) {
576 ResPacket c0 = pzero(ResPacket{}), c1 = pzero(ResPacket{}), c2 = pzero(ResPacket{}), c3 = pzero(ResPacket{});
577
578 for (Index j = 0; j < fullColBlockEnd; j += LhsPacketSize) {
579 RhsPacket b0 = rhs.template load<RhsPacket, Unaligned>(j, 0);
580
581 c0 = pcj.pmadd(lhs.template load<LhsPacket, LhsAlignment>(i + 0, j), b0, c0);
582 c1 = pcj.pmadd(lhs.template load<LhsPacket, LhsAlignment>(i + 1, j), b0, c1);
583 c2 = pcj.pmadd(lhs.template load<LhsPacket, LhsAlignment>(i + 2, j), b0, c2);
584 c3 = pcj.pmadd(lhs.template load<LhsPacket, LhsAlignment>(i + 3, j), b0, c3);
585 }
586 if (UseSegment && segmentCount > 0) {
587 RhsPacket b0 = RhsSegmentLoader::run(rhs, fullColBlockEnd, 0, segmentCount);
588 c0 = pcj.pmadd(LhsSegmentLoader::run(lhs, i + 0, fullColBlockEnd, segmentCount), b0, c0);
589 c1 = pcj.pmadd(LhsSegmentLoader::run(lhs, i + 1, fullColBlockEnd, segmentCount), b0, c1);
590 c2 = pcj.pmadd(LhsSegmentLoader::run(lhs, i + 2, fullColBlockEnd, segmentCount), b0, c2);
591 c3 = pcj.pmadd(LhsSegmentLoader::run(lhs, i + 3, fullColBlockEnd, segmentCount), b0, c3);
592 }
593 ResScalar cc0 = predux(c0);
594 ResScalar cc1 = predux(c1);
595 ResScalar cc2 = predux(c2);
596 ResScalar cc3 = predux(c3);
597
598 for (Index j = scalarColStart; j < cols; ++j) {
599 RhsScalar b0 = rhs(j, 0);
600
601 cc0 += cj.pmul(lhs(i + 0, j), b0);
602 cc1 += cj.pmul(lhs(i + 1, j), b0);
603 cc2 += cj.pmul(lhs(i + 2, j), b0);
604 cc3 += cj.pmul(lhs(i + 3, j), b0);
605 }
606 res[(i + 0) * resIncr] += alpha * cc0;
607 res[(i + 1) * resIncr] += alpha * cc1;
608 res[(i + 2) * resIncr] += alpha * cc2;
609 res[(i + 3) * resIncr] += alpha * cc3;
610 }
611 for (; i < n2; i += 2) {
612 ResPacket c0 = pzero(ResPacket{}), c1 = pzero(ResPacket{});
613
614 for (Index j = 0; j < fullColBlockEnd; j += LhsPacketSize) {
615 RhsPacket b0 = rhs.template load<RhsPacket, Unaligned>(j, 0);
616
617 c0 = pcj.pmadd(lhs.template load<LhsPacket, LhsAlignment>(i + 0, j), b0, c0);
618 c1 = pcj.pmadd(lhs.template load<LhsPacket, LhsAlignment>(i + 1, j), b0, c1);
619 }
620 if (UseSegment && segmentCount > 0) {
621 RhsPacket b0 = RhsSegmentLoader::run(rhs, fullColBlockEnd, 0, segmentCount);
622 c0 = pcj.pmadd(LhsSegmentLoader::run(lhs, i + 0, fullColBlockEnd, segmentCount), b0, c0);
623 c1 = pcj.pmadd(LhsSegmentLoader::run(lhs, i + 1, fullColBlockEnd, segmentCount), b0, c1);
624 }
625 ResScalar cc0 = predux(c0);
626 ResScalar cc1 = predux(c1);
627
628 for (Index j = scalarColStart; j < cols; ++j) {
629 RhsScalar b0 = rhs(j, 0);
630
631 cc0 += cj.pmul(lhs(i + 0, j), b0);
632 cc1 += cj.pmul(lhs(i + 1, j), b0);
633 }
634 res[(i + 0) * resIncr] += alpha * cc0;
635 res[(i + 1) * resIncr] += alpha * cc1;
636 }
637 for (; i < rows; ++i) {
638 ResPacket c0 = pzero(ResPacket{});
639 ResPacketHalf c0_h = pzero(ResPacketHalf{});
640 ResPacketQuarter c0_q = pzero(ResPacketQuarter{});
641
642 for (Index j = 0; j < fullColBlockEnd; j += LhsPacketSize) {
643 RhsPacket b0 = rhs.template load<RhsPacket, Unaligned>(j, 0);
644 c0 = pcj.pmadd(lhs.template load<LhsPacket, LhsAlignment>(i, j), b0, c0);
645 }
646 if (UseSegment && segmentCount > 0) {
647 RhsPacket b0 = RhsSegmentLoader::run(rhs, fullColBlockEnd, 0, segmentCount);
648 c0 = pcj.pmadd(LhsSegmentLoader::run(lhs, i, fullColBlockEnd, segmentCount), b0, c0);
649 }
650 ResScalar cc0 = predux(c0);
651 EIGEN_IF_CONSTEXPR (HasHalf && !UseSegment) {
652 for (Index j = fullColBlockEnd; j < halfColBlockEnd; j += LhsPacketSizeHalf) {
653 RhsPacketHalf b0 = rhs.template load<RhsPacketHalf, Unaligned>(j, 0);
654 c0_h = pcj_half.pmadd(lhs.template load<LhsPacketHalf, LhsAlignment>(i, j), b0, c0_h);
655 }
656 cc0 += predux(c0_h);
657 }
658 EIGEN_IF_CONSTEXPR (HasQuarter && !UseSegment) {
659 for (Index j = halfColBlockEnd; j < quarterColBlockEnd; j += LhsPacketSizeQuarter) {
660 RhsPacketQuarter b0 = rhs.template load<RhsPacketQuarter, Unaligned>(j, 0);
661 c0_q = pcj_quarter.pmadd(lhs.template load<LhsPacketQuarter, LhsAlignment>(i, j), b0, c0_q);
662 }
663 cc0 += predux(c0_q);
664 }
665 for (Index j = UseSegment ? cols : quarterColBlockEnd; j < cols; ++j) {
666 cc0 += cj.pmul(lhs(i, j), rhs(j, 0));
667 }
668 res[i * resIncr] += alpha * cc0;
669 }
670}
671
672// Integer-sequence helper for process_rows_small_cols.
673template <int N>
674struct gemv_small_cols_unroller {
675 template <typename LhsPacket, typename AccPacket, int Alignment, typename RhsType, typename ConjHelper,
676 typename LhsMapper, typename Index, int... K>
677 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void madd_impl(std::integer_sequence<int, K...>, AccPacket* acc,
678 const LhsMapper& lhs, Index i, Index j, const RhsType& b0,
679 ConjHelper& pcj) {
680 int unused[] = {0, ((acc[K] = pcj.pmadd(lhs.template load<LhsPacket, Alignment>(i + K, j), b0, acc[K])), 0)...};
681 EIGEN_UNUSED_VARIABLE(unused);
682 }
683
684 template <typename LhsPacket, typename AccPacket, int Alignment, typename RhsType, typename ConjHelper,
685 typename LhsMapper, typename Index>
686 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void madd(AccPacket* acc, const LhsMapper& lhs, Index i, Index j,
687 const RhsType& b0, ConjHelper& pcj) {
688 madd_impl<LhsPacket, AccPacket, Alignment>(std::make_integer_sequence<int, N>{}, acc, lhs, i, j, b0, pcj);
689 }
690
691 template <typename ResScalar, typename RhsScalar, typename ConjHelper, typename LhsMapper, typename Index, int... K>
692 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void scalar_madd_impl(std::integer_sequence<int, K...>, ResScalar* cc,
693 const LhsMapper& lhs, Index i, Index j,
694 const RhsScalar& b0, ConjHelper& cj) {
695 int unused[] = {0, ((cc[K] += cj.pmul(lhs(i + K, j), b0)), 0)...};
696 EIGEN_UNUSED_VARIABLE(unused);
697 }
698
699 template <typename ResScalar, typename RhsScalar, typename ConjHelper, typename LhsMapper, typename Index>
700 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void scalar_madd(ResScalar* cc, const LhsMapper& lhs, Index i, Index j,
701 const RhsScalar& b0, ConjHelper& cj) {
702 scalar_madd_impl(std::make_integer_sequence<int, N>{}, cc, lhs, i, j, b0, cj);
703 }
704
705 template <typename Scalar, typename Packet, int... K>
706 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void predux_accum_impl(std::integer_sequence<int, K...>, Scalar* cc,
707 const Packet* acc) {
708 int unused[] = {0, ((cc[K] += predux(acc[K])), 0)...};
709 EIGEN_UNUSED_VARIABLE(unused);
710 }
711
712 template <typename Scalar, typename Packet>
713 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void predux_accum(Scalar* cc, const Packet* acc) {
714 predux_accum_impl(std::make_integer_sequence<int, N>{}, cc, acc);
715 }
716
717 template <typename Packet, int... K>
718 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void init_zero_impl(std::integer_sequence<int, K...>, Packet* acc) {
719 int unused[] = {0, ((acc[K] = pzero(Packet{})), 0)...};
720 EIGEN_UNUSED_VARIABLE(unused);
721 }
722
723 template <typename Packet>
724 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void init_zero(Packet* acc) {
725 init_zero_impl(std::make_integer_sequence<int, N>{}, acc);
726 }
727
728 template <typename Scalar, typename Index, int... K>
729 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void write_result_impl(std::integer_sequence<int, K...>, Scalar* res,
730 Index resIncr, Index i, Scalar alpha,
731 const Scalar* cc) {
732 int unused[] = {0, ((res[(i + K) * resIncr] += alpha * cc[K]), 0)...};
733 EIGEN_UNUSED_VARIABLE(unused);
734 }
735
736 template <typename Scalar, typename Index>
737 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void write_result(Scalar* res, Index resIncr, Index i, Scalar alpha,
738 const Scalar* cc) {
739 write_result_impl(std::make_integer_sequence<int, N>{}, res, resIncr, i, alpha, cc);
740 }
741};
742
743template <typename Index, typename LhsScalar, typename LhsMapper, bool ConjugateLhs, typename RhsScalar,
744 typename RhsMapper, bool ConjugateRhs, int Version>
745template <int N>
746EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void
747general_matrix_vector_product<Index, LhsScalar, LhsMapper, RowMajor, ConjugateLhs, RhsScalar, RhsMapper, ConjugateRhs,
748 Version>::process_rows_small_cols(Index i, Index cols, const LhsMapper& lhs,
749 const RhsMapper& rhs, ResScalar* res, Index resIncr,
750 ResScalar alpha, Index halfColBlockEnd,
751 Index quarterColBlockEnd) {
752 conj_helper<LhsScalar, RhsScalar, ConjugateLhs, ConjugateRhs> cj;
753 conj_helper<LhsPacketHalf, RhsPacketHalf, ConjugateLhs, ConjugateRhs> pcj_half;
754 conj_helper<LhsPacketQuarter, RhsPacketQuarter, ConjugateLhs, ConjugateRhs> pcj_quarter;
755
756 enum {
757 LhsAlignment = Unaligned,
758 ResPacketSizeHalf = HalfTraits::ResPacketSize,
759 ResPacketSizeQuarter = QuarterTraits::ResPacketSize,
760 LhsPacketSizeHalf = HalfTraits::LhsPacketSize,
761 LhsPacketSizeQuarter = QuarterTraits::LhsPacketSize,
762 HasHalf = (int)ResPacketSizeHalf < (int)Traits::ResPacketSize,
763 HasQuarter = (int)ResPacketSizeQuarter < (int)ResPacketSizeHalf
764 };
765
766 using Unroll = gemv_small_cols_unroller<N>;
767
768 ResScalar cc[N] = {};
769 EIGEN_IF_CONSTEXPR (HasHalf) {
770 ResPacketHalf h[N];
771 Unroll::init_zero(h);
772 for (Index j = 0; j < halfColBlockEnd; j += LhsPacketSizeHalf) {
773 RhsPacketHalf b0 = rhs.template load<RhsPacketHalf, Unaligned>(j, 0);
774 Unroll::template madd<LhsPacketHalf, ResPacketHalf, LhsAlignment>(h, lhs, i, j, b0, pcj_half);
775 }
776 Unroll::predux_accum(cc, h);
777 }
778 EIGEN_IF_CONSTEXPR (HasQuarter) {
779 ResPacketQuarter q[N];
780 Unroll::init_zero(q);
781 for (Index j = halfColBlockEnd; j < quarterColBlockEnd; j += LhsPacketSizeQuarter) {
782 RhsPacketQuarter b0 = rhs.template load<RhsPacketQuarter, Unaligned>(j, 0);
783 Unroll::template madd<LhsPacketQuarter, ResPacketQuarter, LhsAlignment>(q, lhs, i, j, b0, pcj_quarter);
784 }
785 Unroll::predux_accum(cc, q);
786 }
787 for (Index j = quarterColBlockEnd; j < cols; ++j) {
788 RhsScalar b0 = rhs(j, 0);
789 Unroll::scalar_madd(cc, lhs, i, j, b0, cj);
790 }
791 Unroll::write_result(res, resIncr, i, alpha, cc);
792}
793
794template <typename Index, typename LhsScalar, typename LhsMapper, bool ConjugateLhs, typename RhsScalar,
795 typename RhsMapper, bool ConjugateRhs, int Version>
796EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void
797general_matrix_vector_product<Index, LhsScalar, LhsMapper, RowMajor, ConjugateLhs, RhsScalar, RhsMapper, ConjugateRhs,
798 Version>::run_small_cols(Index rows, Index cols, const LhsMapper& alhs,
799 const RhsMapper& rhs, ResScalar* res, Index resIncr,
800 ResScalar alpha) {
801 LhsMapper lhs(alhs);
802 eigen_internal_assert(rhs.stride() == 1);
803
804 enum {
805 LhsPacketSizeHalf = HalfTraits::LhsPacketSize,
806 LhsPacketSizeQuarter = QuarterTraits::LhsPacketSize,
807 };
808
809 using UnsignedIndex = std::make_unsigned_t<Index>;
810 const Index halfColBlockEnd = LhsPacketSizeHalf * (UnsignedIndex(cols) / LhsPacketSizeHalf);
811 const Index quarterColBlockEnd = LhsPacketSizeQuarter * (UnsignedIndex(cols) / LhsPacketSizeQuarter);
812
813 // Disable the 8-row inner unroll once a single column slice no longer fits in L1; with very
814 // large LHS strides each unrolled iteration evicts the previously-loaded rows from cache.
815 std::ptrdiff_t l1, l2, l3;
816 manage_caching_sizes(GetAction, &l1, &l2, &l3);
817 const Index n8 = lhs.stride() * Index(sizeof(LhsScalar)) > Index(l1) ? 0 : rows - 7;
818 const Index n4 = rows - 3;
819 const Index n2 = rows - 1;
820
821 Index i = 0;
822 for (; i < n8; i += 8) {
823 process_rows_small_cols<8>(i, cols, lhs, rhs, res, resIncr, alpha, halfColBlockEnd, quarterColBlockEnd);
824 }
825 // Process remaining groups of 4 rows in case n8 was 0.
826 for (; i < n4; i += 4) {
827 process_rows_small_cols<4>(i, cols, lhs, rhs, res, resIncr, alpha, halfColBlockEnd, quarterColBlockEnd);
828 }
829 if (i < n2) {
830 process_rows_small_cols<2>(i, cols, lhs, rhs, res, resIncr, alpha, halfColBlockEnd, quarterColBlockEnd);
831 i += 2;
832 }
833 if (i < rows) {
834 process_rows_small_cols<1>(i, cols, lhs, rhs, res, resIncr, alpha, halfColBlockEnd, quarterColBlockEnd);
835 }
836}
837
838} // end namespace internal
839
840} // end namespace Eigen
841
842#if EIGEN_COMP_MSVC
843#pragma warning(pop)
844#endif
845
846#endif // EIGEN_GENERAL_MATRIX_VECTOR_H
@ Unaligned
Definition Constants.h:236
@ ColMajor
Definition Constants.h:319
@ RowMajor
Definition Constants.h:321