11#ifndef EIGEN_GENERAL_MATRIX_VECTOR_H
12#define EIGEN_GENERAL_MATRIX_VECTOR_H
15#include "../InternalHeaderCheck.h"
21#pragma warning(disable : 4804)
28template <
typename LhsScalar,
typename RhsScalar,
int PacketSize_ = GEBPPacketFull>
30 using ResScalar =
typename ScalarBinaryOpTraits<LhsScalar, RhsScalar>::ReturnType;
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
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
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
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>;
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 {};
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>;
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{});
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);
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>;
110 using ResScalar =
typename ScalarBinaryOpTraits<LhsScalar, RhsScalar>::ReturnType;
112 using LhsPacket =
typename Traits::LhsPacket;
113 using RhsPacket =
typename Traits::RhsPacket;
114 using ResPacket =
typename Traits::ResPacket;
116 using LhsPacketHalf =
typename HalfTraits::LhsPacket;
117 using RhsPacketHalf =
typename HalfTraits::RhsPacket;
118 using ResPacketHalf =
typename HalfTraits::ResPacket;
120 using LhsPacketQuarter =
typename QuarterTraits::LhsPacket;
121 using RhsPacketQuarter =
typename QuarterTraits::RhsPacket;
122 using ResPacketQuarter =
typename QuarterTraits::ResPacket;
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);
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);
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,
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) {}
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);
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);
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) {
167 0, ((c[K] = pcj.pmadd(lhs.template load<LhsPacket, Alignment>(i + LhsStride * K, j), b0, c[K])), 0)...};
168 EIGEN_UNUSED_VARIABLE(unused);
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);
178 template <
int K,
typename ResPacket,
int ResStr
ide,
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)));
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);
192 template <
typename ResPacket,
int ResStr
ide,
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);
201struct gemv_colmajor_unroller<0> {
202 template <
typename Packet>
203 EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE
void init_zero(Packet*) {}
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&,
210 template <
typename ResPacket,
int ResStr
ide,
typename ResScalar,
typename Index>
211 EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE
void store(
const ResPacket*, ResScalar*, Index,
const ResPacket&) {}
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,
223 enum { LhsAlignment =
Unaligned, LhsPacketSize = Traits::LhsPacketSize, ResPacketSize = Traits::ResPacketSize };
224 using Unroller = gemv_colmajor_unroller<N>;
225 const Index iseg = i + N * ResPacketSize;
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);
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);
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) \
254 process_rows<n, true>(i, j2, jend, lhs, rhs, res, palpha, pcj, count); \
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);
268 eigen_internal_assert(
false);
271#undef EIGEN_GEMV_PROCESS_ROW
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);
284 if (numext::is_exactly_zero(alpha))
return;
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;
295 const Index lhsStride = lhs.stride();
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
311 using UnsignedIndex = std::make_unsigned_t<Index>;
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;
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);
335 for (Index j2 = 0; j2 < cols; j2 += block_cols) {
336 Index jend = numext::mini(j2 + block_cols, cols);
338 for (; i < n8; i += ResPacketSize * 8) process_rows<8>(i, j2, jend, lhs, rhs, res, palpha, pcj);
340 process_segment_tail(bool_constant<UseSegment>(), Index(UnsignedIndex(rows - i) / ResPacketSize), i, j2, jend,
341 lhs, rhs, res, palpha, pcj, count);
343#define EIGEN_GEMV_PROCESS_ROW(k) \
345 process_rows<k>(i, j2, jend, lhs, rhs, res, palpha, pcj); \
346 i += ResPacketSize * (k); \
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) {
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);
361 pstoreu(res + i + ResPacketSizeHalf * 0,
362 pmadd(c0, palpha_half, ploadu<ResPacketHalf>(res + i + ResPacketSizeHalf * 0)));
363 i += ResPacketSizeHalf;
366 EIGEN_IF_CONSTEXPR (HasQuarter) {
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);
373 pstoreu(res + i + ResPacketSizeQuarter * 0,
374 pmadd(c0, palpha_quarter, ploadu<ResPacketQuarter>(res + i + ResPacketSizeQuarter * 0)));
375 i += ResPacketSizeQuarter;
378 for (; i < rows; ++i) {
380 for (Index j = j2; j < jend; j += 1) c0 += cj.pmul(lhs(i, j), rhs(j, 0));
381 res[i] += alpha * c0;
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>;
405 using ResScalar =
typename ScalarBinaryOpTraits<LhsScalar, RhsScalar>::ReturnType;
407 using LhsPacket =
typename Traits::LhsPacket;
408 using RhsPacket =
typename Traits::RhsPacket;
409 using ResPacket =
typename Traits::ResPacket;
411 using LhsPacketHalf =
typename HalfTraits::LhsPacket;
412 using RhsPacketHalf =
typename HalfTraits::RhsPacket;
413 using ResPacketHalf =
typename HalfTraits::ResPacket;
415 using LhsPacketQuarter =
typename QuarterTraits::LhsPacket;
416 using RhsPacketQuarter =
typename QuarterTraits::RhsPacket;
417 using ResPacketQuarter =
typename QuarterTraits::ResPacket;
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);
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,
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);
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) {
444 if (numext::is_exactly_zero(alpha))
return;
450 LhsPacketSize_ = Traits::LhsPacketSize,
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,
461 SmallColsEnd_ = UseSegment_ ? (
int)HalfTraits::LhsPacketSize + 2 : (int)LhsPacketSize_
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);
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;
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;
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_
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);
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>;
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{});
522 for (Index j = 0; j < fullColBlockEnd; j += LhsPacketSize) {
523 RhsPacket b0 = rhs.template load<RhsPacket, Unaligned>(j, 0);
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);
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);
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);
554 for (Index j = scalarColStart; j < cols; ++j) {
555 RhsScalar b0 = rhs(j, 0);
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);
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;
575 for (; i < n4; i += 4) {
576 ResPacket c0 = pzero(ResPacket{}), c1 = pzero(ResPacket{}), c2 = pzero(ResPacket{}), c3 = pzero(ResPacket{});
578 for (Index j = 0; j < fullColBlockEnd; j += LhsPacketSize) {
579 RhsPacket b0 = rhs.template load<RhsPacket, Unaligned>(j, 0);
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);
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);
593 ResScalar cc0 = predux(c0);
594 ResScalar cc1 = predux(c1);
595 ResScalar cc2 = predux(c2);
596 ResScalar cc3 = predux(c3);
598 for (Index j = scalarColStart; j < cols; ++j) {
599 RhsScalar b0 = rhs(j, 0);
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);
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;
611 for (; i < n2; i += 2) {
612 ResPacket c0 = pzero(ResPacket{}), c1 = pzero(ResPacket{});
614 for (Index j = 0; j < fullColBlockEnd; j += LhsPacketSize) {
615 RhsPacket b0 = rhs.template load<RhsPacket, Unaligned>(j, 0);
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);
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);
625 ResScalar cc0 = predux(c0);
626 ResScalar cc1 = predux(c1);
628 for (Index j = scalarColStart; j < cols; ++j) {
629 RhsScalar b0 = rhs(j, 0);
631 cc0 += cj.pmul(lhs(i + 0, j), b0);
632 cc1 += cj.pmul(lhs(i + 1, j), b0);
634 res[(i + 0) * resIncr] += alpha * cc0;
635 res[(i + 1) * resIncr] += alpha * cc1;
637 for (; i < rows; ++i) {
638 ResPacket c0 = pzero(ResPacket{});
639 ResPacketHalf c0_h = pzero(ResPacketHalf{});
640 ResPacketQuarter c0_q = pzero(ResPacketQuarter{});
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);
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);
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);
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);
665 for (Index j = UseSegment ? cols : quarterColBlockEnd; j < cols; ++j) {
666 cc0 += cj.pmul(lhs(i, j), rhs(j, 0));
668 res[i * resIncr] += alpha * cc0;
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,
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);
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);
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);
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);
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,
708 int unused[] = {0, ((cc[K] += predux(acc[K])), 0)...};
709 EIGEN_UNUSED_VARIABLE(unused);
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);
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);
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);
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,
732 int unused[] = {0, ((res[(i + K) * resIncr] += alpha * cc[K]), 0)...};
733 EIGEN_UNUSED_VARIABLE(unused);
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,
739 write_result_impl(std::make_integer_sequence<int, N>{}, res, resIncr, i, alpha, cc);
743template <
typename Index,
typename LhsScalar,
typename LhsMapper,
bool ConjugateLhs,
typename RhsScalar,
744 typename RhsMapper,
bool ConjugateRhs,
int Version>
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;
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
766 using Unroll = gemv_small_cols_unroller<N>;
768 ResScalar cc[N] = {};
769 EIGEN_IF_CONSTEXPR (HasHalf) {
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);
776 Unroll::predux_accum(cc, h);
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);
785 Unroll::predux_accum(cc, q);
787 for (Index j = quarterColBlockEnd; j < cols; ++j) {
788 RhsScalar b0 = rhs(j, 0);
789 Unroll::scalar_madd(cc, lhs, i, j, b0, cj);
791 Unroll::write_result(res, resIncr, i, alpha, cc);
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,
802 eigen_internal_assert(rhs.stride() == 1);
805 LhsPacketSizeHalf = HalfTraits::LhsPacketSize,
806 LhsPacketSizeQuarter = QuarterTraits::LhsPacketSize,
809 using UnsignedIndex = std::make_unsigned_t<Index>;
810 const Index halfColBlockEnd = LhsPacketSizeHalf * (UnsignedIndex(cols) / LhsPacketSizeHalf);
811 const Index quarterColBlockEnd = LhsPacketSizeQuarter * (UnsignedIndex(cols) / LhsPacketSizeQuarter);
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;
822 for (; i < n8; i += 8) {
823 process_rows_small_cols<8>(i, cols, lhs, rhs, res, resIncr, alpha, halfColBlockEnd, quarterColBlockEnd);
826 for (; i < n4; i += 4) {
827 process_rows_small_cols<4>(i, cols, lhs, rhs, res, resIncr, alpha, halfColBlockEnd, quarterColBlockEnd);
830 process_rows_small_cols<2>(i, cols, lhs, rhs, res, resIncr, alpha, halfColBlockEnd, quarterColBlockEnd);
834 process_rows_small_cols<1>(i, cols, lhs, rhs, res, resIncr, alpha, halfColBlockEnd, quarterColBlockEnd);
@ Unaligned
Definition Constants.h:236
@ ColMajor
Definition Constants.h:319
@ RowMajor
Definition Constants.h:321