Eigen  5.0.1
 
Loading...
Searching...
No Matches
BlasUtil.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2009-2010 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_BLASUTIL_H
12#define EIGEN_BLASUTIL_H
13
14// This file contains many lightweight helper classes used to
15// implement and control fast level 2 and level 3 BLAS-like routines.
16
17// IWYU pragma: private
18#include "../InternalHeaderCheck.h"
19
20namespace Eigen {
21
22namespace internal {
23
24// forward declarations
25template <typename LhsScalar, typename RhsScalar, typename Index, typename DataMapper, int mr, int nr,
26 bool ConjugateLhs = false, bool ConjugateRhs = false>
27struct gebp_kernel;
28
29template <typename Scalar, typename Index, typename DataMapper, int nr, int StorageOrder, bool Conjugate = false,
30 bool PanelMode = false>
31struct gemm_pack_rhs;
32
33template <typename Scalar, typename Index, typename DataMapper, int Pack1, int Pack2, typename Packet, int StorageOrder,
34 bool Conjugate = false, bool PanelMode = false>
35struct gemm_pack_lhs;
36
37template <typename Index, typename LhsScalar, int LhsStorageOrder, bool ConjugateLhs, typename RhsScalar,
38 int RhsStorageOrder, bool ConjugateRhs, int ResStorageOrder, int ResInnerStride>
39struct general_matrix_matrix_product;
40
41template <typename Index, typename LhsScalar, typename LhsMapper, int LhsStorageOrder, bool ConjugateLhs,
42 typename RhsScalar, typename RhsMapper, bool ConjugateRhs, int Version = Specialized>
43struct general_matrix_vector_product;
44
45template <typename From, typename To>
46struct get_factor {
47 EIGEN_DEVICE_FUNC constexpr static EIGEN_STRONG_INLINE To run(const From& x) { return To(x); }
48};
49
50template <typename Scalar>
51struct get_factor<Scalar, typename NumTraits<Scalar>::Real> {
52 EIGEN_DEVICE_FUNC constexpr static EIGEN_STRONG_INLINE typename NumTraits<Scalar>::Real run(const Scalar& x) {
53 return numext::real(x);
54 }
55};
56
57template <typename Scalar, typename Index>
58class BlasVectorMapper {
59 public:
60 EIGEN_DEVICE_FUNC constexpr EIGEN_ALWAYS_INLINE BlasVectorMapper(Scalar* data) : m_data(data) {}
61
62 EIGEN_DEVICE_FUNC constexpr EIGEN_ALWAYS_INLINE Scalar operator()(Index i) const { return m_data[i]; }
63 template <typename Packet, int AlignmentType>
64 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet load(Index i) const {
65 return ploadt<Packet, AlignmentType>(m_data + i);
66 }
67
68 template <typename Packet>
69 EIGEN_DEVICE_FUNC bool aligned(Index i) const {
70 return (std::uintptr_t(m_data + i) % sizeof(Packet)) == 0;
71 }
72
73 protected:
74 Scalar* m_data;
75};
76
77template <typename Scalar, typename Index, int AlignmentType, int Incr = 1>
78class BlasLinearMapper;
79
80template <typename Scalar, typename Index, int AlignmentType>
81class BlasLinearMapper<Scalar, Index, AlignmentType> {
82 public:
83 EIGEN_DEVICE_FUNC constexpr EIGEN_ALWAYS_INLINE BlasLinearMapper(Scalar* data, Index incr = 1) : m_data(data) {
84 EIGEN_ONLY_USED_FOR_DEBUG(incr);
85 eigen_assert(incr == 1);
86 }
87
88 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void prefetch(Index i) const { internal::prefetch(&operator()(i)); }
89
90 EIGEN_DEVICE_FUNC constexpr EIGEN_ALWAYS_INLINE Scalar& operator()(Index i) const { return m_data[i]; }
91
92 template <typename PacketType>
93 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE PacketType loadPacket(Index i) const {
94 return ploadt<PacketType, AlignmentType>(m_data + i);
95 }
96
97 template <typename PacketType>
98 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE PacketType loadPacketPartial(Index i, Index n, Index offset = 0) const {
99 return ploadt_partial<PacketType, AlignmentType>(m_data + i, n, offset);
100 }
101
102 template <typename PacketType, int AlignmentT>
103 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE PacketType load(Index i) const {
104 return ploadt<PacketType, AlignmentT>(m_data + i);
105 }
106
107 template <typename PacketType>
108 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void storePacket(Index i, const PacketType& p) const {
109 pstoret<Scalar, PacketType, AlignmentType>(m_data + i, p);
110 }
111
112 template <typename PacketType>
113 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void storePacketPartial(Index i, const PacketType& p, Index n,
114 Index offset = 0) const {
115 pstoret_partial<Scalar, PacketType, AlignmentType>(m_data + i, p, n, offset);
116 }
117
118 protected:
119 Scalar* m_data;
120};
121
122// Lightweight helper class to access matrix coefficients.
123template <typename Scalar, typename Index, int StorageOrder, int AlignmentType = Unaligned, int Incr = 1>
124class blas_data_mapper;
125
126// TMP to help PacketBlock store implementation.
127// There's currently no known use case for PacketBlock load.
128// The default implementation assumes ColMajor order.
129// It always store each packet sequentially one `stride` apart.
130template <typename Index, typename Scalar, typename Packet, int n, int idx, int StorageOrder>
131struct PacketBlockManagement {
132 PacketBlockManagement<Index, Scalar, Packet, n, idx - 1, StorageOrder> pbm;
133 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void store(Scalar* to, const Index stride, Index i, Index j,
134 const PacketBlock<Packet, n>& block) const {
135 pbm.store(to, stride, i, j, block);
136 pstoreu<Scalar>(to + i + (j + idx) * stride, block.packet[idx]);
137 }
138};
139
140// PacketBlockManagement specialization to take care of RowMajor order without ifs.
141template <typename Index, typename Scalar, typename Packet, int n, int idx>
142struct PacketBlockManagement<Index, Scalar, Packet, n, idx, RowMajor> {
143 PacketBlockManagement<Index, Scalar, Packet, n, idx - 1, RowMajor> pbm;
144 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void store(Scalar* to, const Index stride, Index i, Index j,
145 const PacketBlock<Packet, n>& block) const {
146 pbm.store(to, stride, i, j, block);
147 pstoreu<Scalar>(to + j + (i + idx) * stride, block.packet[idx]);
148 }
149};
150
151template <typename Index, typename Scalar, typename Packet, int n, int StorageOrder>
152struct PacketBlockManagement<Index, Scalar, Packet, n, -1, StorageOrder> {
153 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void store(Scalar* to, const Index stride, Index i, Index j,
154 const PacketBlock<Packet, n>& block) const {
155 EIGEN_UNUSED_VARIABLE(to);
156 EIGEN_UNUSED_VARIABLE(stride);
157 EIGEN_UNUSED_VARIABLE(i);
158 EIGEN_UNUSED_VARIABLE(j);
159 EIGEN_UNUSED_VARIABLE(block);
160 }
161};
162
163template <typename Index, typename Scalar, typename Packet, int n>
164struct PacketBlockManagement<Index, Scalar, Packet, n, -1, RowMajor> {
165 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void store(Scalar* to, const Index stride, Index i, Index j,
166 const PacketBlock<Packet, n>& block) const {
167 EIGEN_UNUSED_VARIABLE(to);
168 EIGEN_UNUSED_VARIABLE(stride);
169 EIGEN_UNUSED_VARIABLE(i);
170 EIGEN_UNUSED_VARIABLE(j);
171 EIGEN_UNUSED_VARIABLE(block);
172 }
173};
174
175template <typename Scalar, typename Index, int StorageOrder, int AlignmentType>
176class blas_data_mapper<Scalar, Index, StorageOrder, AlignmentType, 1> {
177 public:
178 using LinearMapper = BlasLinearMapper<Scalar, Index, AlignmentType>;
179 using SubMapper = blas_data_mapper<Scalar, Index, StorageOrder, AlignmentType>;
180 using VectorMapper = BlasVectorMapper<Scalar, Index>;
181
182 EIGEN_DEVICE_FUNC constexpr EIGEN_ALWAYS_INLINE blas_data_mapper(Scalar* data, Index stride, Index incr = 1)
183 : m_data(data), m_stride(stride) {
184 EIGEN_ONLY_USED_FOR_DEBUG(incr);
185 eigen_assert(incr == 1);
186 }
187
188 EIGEN_DEVICE_FUNC constexpr EIGEN_ALWAYS_INLINE SubMapper getSubMapper(Index i, Index j) const {
189 return SubMapper(&operator()(i, j), m_stride);
190 }
191
192 EIGEN_DEVICE_FUNC constexpr EIGEN_ALWAYS_INLINE LinearMapper getLinearMapper(Index i, Index j) const {
193 return LinearMapper(&operator()(i, j));
194 }
195
196 EIGEN_DEVICE_FUNC constexpr EIGEN_ALWAYS_INLINE VectorMapper getVectorMapper(Index i, Index j) const {
197 return VectorMapper(&operator()(i, j));
198 }
199
200 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void prefetch(Index i, Index j) const { internal::prefetch(&operator()(i, j)); }
201
202 EIGEN_DEVICE_FUNC constexpr EIGEN_ALWAYS_INLINE Scalar& operator()(Index i, Index j) const {
203 return m_data[StorageOrder == RowMajor ? j + i * m_stride : i + j * m_stride];
204 }
205
206 template <typename PacketType>
207 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE PacketType loadPacket(Index i, Index j) const {
208 return ploadt<PacketType, AlignmentType>(&operator()(i, j));
209 }
210
211 template <typename PacketType>
212 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE PacketType loadPacketPartial(Index i, Index j, Index n,
213 Index offset = 0) const {
214 return ploadt_partial<PacketType, AlignmentType>(&operator()(i, j), n, offset);
215 }
216
217 template <typename PacketT, int AlignmentT>
218 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE PacketT load(Index i, Index j) const {
219 return ploadt<PacketT, AlignmentT>(&operator()(i, j));
220 }
221
222 template <typename PacketType>
223 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void storePacket(Index i, Index j, const PacketType& p) const {
224 pstoret<Scalar, PacketType, AlignmentType>(&operator()(i, j), p);
225 }
226
227 template <typename PacketType>
228 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void storePacketPartial(Index i, Index j, const PacketType& p, Index n,
229 Index offset = 0) const {
230 pstoret_partial<Scalar, PacketType, AlignmentType>(&operator()(i, j), p, n, offset);
231 }
232
233 template <typename SubPacket>
234 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void scatterPacket(Index i, Index j, const SubPacket& p) const {
235 pscatter<Scalar, SubPacket>(&operator()(i, j), p, m_stride);
236 }
237
238 template <typename SubPacket>
239 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE SubPacket gatherPacket(Index i, Index j) const {
240 return pgather<Scalar, SubPacket>(&operator()(i, j), m_stride);
241 }
242
243 EIGEN_DEVICE_FUNC constexpr const Index stride() const { return m_stride; }
244 EIGEN_DEVICE_FUNC constexpr const Index incr() const { return 1; }
245 EIGEN_DEVICE_FUNC constexpr const Scalar* data() const { return m_data; }
246
247 EIGEN_DEVICE_FUNC Index firstAligned(Index size) const {
248 if (std::uintptr_t(m_data) % sizeof(Scalar)) {
249 return -1;
250 }
251 return internal::first_default_aligned(m_data, size);
252 }
253
254 template <typename SubPacket, int n>
255 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void storePacketBlock(Index i, Index j,
256 const PacketBlock<SubPacket, n>& block) const {
257 PacketBlockManagement<Index, Scalar, SubPacket, n, n - 1, StorageOrder> pbm;
258 pbm.store(m_data, m_stride, i, j, block);
259 }
260
261 protected:
262 Scalar* EIGEN_RESTRICT m_data;
263 const Index m_stride;
264};
265
266// Implementation of non-natural increment (i.e. inner-stride != 1)
267// The exposed API is not complete yet compared to the Incr==1 case
268// because some features makes less sense in this case.
269template <typename Scalar, typename Index, int AlignmentType, int Incr>
270class BlasLinearMapper {
271 public:
272 EIGEN_DEVICE_FUNC constexpr EIGEN_ALWAYS_INLINE BlasLinearMapper(Scalar* data, Index incr)
273 : m_data(data), m_incr(incr) {}
274
275 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void prefetch(int i) const { internal::prefetch(&operator()(i)); }
276
277 EIGEN_DEVICE_FUNC constexpr EIGEN_ALWAYS_INLINE Scalar& operator()(Index i) const {
278 return m_data[i * m_incr.value()];
279 }
280
281 template <typename PacketType>
282 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE PacketType loadPacket(Index i) const {
283 return pgather<Scalar, PacketType>(m_data + i * m_incr.value(), m_incr.value());
284 }
285
286 template <typename PacketType>
287 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE PacketType loadPacketPartial(Index i, Index n, Index /*offset*/ = 0) const {
288 return pgather_partial<Scalar, PacketType>(m_data + i * m_incr.value(), m_incr.value(), n);
289 }
290
291 template <typename PacketType>
292 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void storePacket(Index i, const PacketType& p) const {
293 pscatter<Scalar, PacketType>(m_data + i * m_incr.value(), p, m_incr.value());
294 }
295
296 template <typename PacketType>
297 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void storePacketPartial(Index i, const PacketType& p, Index n,
298 Index /*offset*/ = 0) const {
299 pscatter_partial<Scalar, PacketType>(m_data + i * m_incr.value(), p, m_incr.value(), n);
300 }
301
302 protected:
303 Scalar* m_data;
304 const internal::variable_if_dynamic<Index, Incr> m_incr;
305};
306
307template <typename Scalar, typename Index, int StorageOrder, int AlignmentType, int Incr>
308class blas_data_mapper {
309 public:
310 using LinearMapper = BlasLinearMapper<Scalar, Index, AlignmentType, Incr>;
311 using SubMapper = blas_data_mapper;
312
313 EIGEN_DEVICE_FUNC constexpr EIGEN_ALWAYS_INLINE blas_data_mapper(Scalar* data, Index stride, Index incr)
314 : m_data(data), m_stride(stride), m_incr(incr) {}
315
316 EIGEN_DEVICE_FUNC constexpr EIGEN_ALWAYS_INLINE SubMapper getSubMapper(Index i, Index j) const {
317 return SubMapper(&operator()(i, j), m_stride, m_incr.value());
318 }
319
320 EIGEN_DEVICE_FUNC constexpr EIGEN_ALWAYS_INLINE LinearMapper getLinearMapper(Index i, Index j) const {
321 return LinearMapper(&operator()(i, j), m_incr.value());
322 }
323
324 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void prefetch(Index i, Index j) const { internal::prefetch(&operator()(i, j)); }
325
326 EIGEN_DEVICE_FUNC constexpr EIGEN_ALWAYS_INLINE Scalar& operator()(Index i, Index j) const {
327 return m_data[StorageOrder == RowMajor ? j * m_incr.value() + i * m_stride : i * m_incr.value() + j * m_stride];
328 }
329
330 template <typename PacketType>
331 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE PacketType loadPacket(Index i, Index j) const {
332 return pgather<Scalar, PacketType>(&operator()(i, j), m_incr.value());
333 }
334
335 template <typename PacketType>
336 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE PacketType loadPacketPartial(Index i, Index j, Index n,
337 Index /*offset*/ = 0) const {
338 return pgather_partial<Scalar, PacketType>(&operator()(i, j), m_incr.value(), n);
339 }
340
341 template <typename PacketT, int AlignmentT>
342 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE PacketT load(Index i, Index j) const {
343 return pgather<Scalar, PacketT>(&operator()(i, j), m_incr.value());
344 }
345
346 template <typename PacketType>
347 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void storePacket(Index i, Index j, const PacketType& p) const {
348 pscatter<Scalar, PacketType>(&operator()(i, j), p, m_incr.value());
349 }
350
351 template <typename PacketType>
352 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void storePacketPartial(Index i, Index j, const PacketType& p, Index n,
353 Index /*offset*/ = 0) const {
354 pscatter_partial<Scalar, PacketType>(&operator()(i, j), p, m_incr.value(), n);
355 }
356
357 template <typename SubPacket>
358 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void scatterPacket(Index i, Index j, const SubPacket& p) const {
359 pscatter<Scalar, SubPacket>(&operator()(i, j), p, m_stride);
360 }
361
362 template <typename SubPacket>
363 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE SubPacket gatherPacket(Index i, Index j) const {
364 return pgather<Scalar, SubPacket>(&operator()(i, j), m_stride);
365 }
366
367 // storePacketBlock_helper defines a way to access values inside the PacketBlock, this is essentially required by the
368 // Complex types.
369 template <typename SubPacket, typename Scalar_, int n, int idx>
370 struct storePacketBlock_helper {
371 storePacketBlock_helper<SubPacket, Scalar_, n, idx - 1> spbh;
372 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void store(
373 const blas_data_mapper<Scalar, Index, StorageOrder, AlignmentType, Incr>* sup, Index i, Index j,
374 const PacketBlock<SubPacket, n>& block) const {
375 spbh.store(sup, i, j, block);
376 sup->template storePacket<SubPacket>(i, j + idx, block.packet[idx]);
377 }
378 };
379
380 template <typename SubPacket, int n, int idx>
381 struct storePacketBlock_helper<SubPacket, std::complex<float>, n, idx> {
382 storePacketBlock_helper<SubPacket, std::complex<float>, n, idx - 1> spbh;
383 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void store(
384 const blas_data_mapper<Scalar, Index, StorageOrder, AlignmentType, Incr>* sup, Index i, Index j,
385 const PacketBlock<SubPacket, n>& block) const {
386 spbh.store(sup, i, j, block);
387 sup->template storePacket<SubPacket>(i, j + idx, block.packet[idx]);
388 }
389 };
390
391 template <typename SubPacket, int n, int idx>
392 struct storePacketBlock_helper<SubPacket, std::complex<double>, n, idx> {
393 storePacketBlock_helper<SubPacket, std::complex<double>, n, idx - 1> spbh;
394 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void store(
395 const blas_data_mapper<Scalar, Index, StorageOrder, AlignmentType, Incr>* sup, Index i, Index j,
396 const PacketBlock<SubPacket, n>& block) const {
397 spbh.store(sup, i, j, block);
398 for (int l = 0; l < unpacket_traits<SubPacket>::size; l++) {
399 std::complex<double>* v = &sup->operator()(i + l, j + idx);
400 v->real(block.packet[idx].v[2 * l + 0]);
401 v->imag(block.packet[idx].v[2 * l + 1]);
402 }
403 }
404 };
405
406 template <typename SubPacket, typename Scalar_, int n>
407 struct storePacketBlock_helper<SubPacket, Scalar_, n, -1> {
408 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void store(
409 const blas_data_mapper<Scalar, Index, StorageOrder, AlignmentType, Incr>*, Index, Index,
410 const PacketBlock<SubPacket, n>&) const {}
411 };
412
413 template <typename SubPacket, int n>
414 struct storePacketBlock_helper<SubPacket, std::complex<float>, n, -1> {
415 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void store(
416 const blas_data_mapper<Scalar, Index, StorageOrder, AlignmentType, Incr>*, Index, Index,
417 const PacketBlock<SubPacket, n>&) const {}
418 };
419
420 template <typename SubPacket, int n>
421 struct storePacketBlock_helper<SubPacket, std::complex<double>, n, -1> {
422 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void store(
423 const blas_data_mapper<Scalar, Index, StorageOrder, AlignmentType, Incr>*, Index, Index,
424 const PacketBlock<SubPacket, n>&) const {}
425 };
426 // This function stores a PacketBlock on m_data, this approach is really quite slow compare to Incr=1 and should be
427 // avoided when possible.
428 template <typename SubPacket, int n>
429 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void storePacketBlock(Index i, Index j,
430 const PacketBlock<SubPacket, n>& block) const {
431 storePacketBlock_helper<SubPacket, Scalar, n, n - 1> spb;
432 spb.store(this, i, j, block);
433 }
434
435 EIGEN_DEVICE_FUNC constexpr const Index stride() const { return m_stride; }
436 EIGEN_DEVICE_FUNC constexpr const Index incr() const { return m_incr.value(); }
437 EIGEN_DEVICE_FUNC constexpr Scalar* data() const { return m_data; }
438
439 protected:
440 Scalar* EIGEN_RESTRICT m_data;
441 const Index m_stride;
442 const internal::variable_if_dynamic<Index, Incr> m_incr;
443};
444
445// lightweight helper class to access matrix coefficients (const version)
446template <typename Scalar, typename Index, int StorageOrder>
447class const_blas_data_mapper : public blas_data_mapper<const Scalar, Index, StorageOrder> {
448 public:
449 using SubMapper = const_blas_data_mapper<Scalar, Index, StorageOrder>;
450
451 EIGEN_ALWAYS_INLINE const_blas_data_mapper(const Scalar* data, Index stride)
452 : blas_data_mapper<const Scalar, Index, StorageOrder>(data, stride) {}
453
454 EIGEN_ALWAYS_INLINE SubMapper getSubMapper(Index i, Index j) const {
455 return SubMapper(&(this->operator()(i, j)), this->m_stride);
456 }
457};
458
459/* Helper class to analyze the factors of a Product expression.
460 * In particular it allows to pop out operator-, scalar multiples,
461 * and conjugate */
462template <typename XprType>
463struct blas_traits {
464 using Scalar = typename traits<XprType>::Scalar;
465 using ExtractType = const XprType&;
466 using ExtractType_ = XprType;
467 enum {
468 IsComplex = NumTraits<Scalar>::IsComplex,
469 IsTransposed = false,
470 NeedToConjugate = false,
471 HasUsableDirectAccess =
472 ((int(XprType::Flags) & DirectAccessBit) &&
473 (bool(XprType::IsVectorAtCompileTime) || int(inner_stride_at_compile_time<XprType>::value) == 1))
474 ? 1
475 : 0,
476 HasScalarFactor = false
477 };
478 using DirectLinearAccessType =
479 std::conditional_t<bool(HasUsableDirectAccess), ExtractType, typename ExtractType_::PlainObject>;
480 EIGEN_DEVICE_FUNC static inline EIGEN_DEVICE_FUNC ExtractType extract(const XprType& x) { return x; }
481 EIGEN_DEVICE_FUNC static inline EIGEN_DEVICE_FUNC const Scalar extractScalarFactor(const XprType&) {
482 return Scalar(1);
483 }
484};
485
486// pop conjugate
487template <typename Scalar, typename NestedXpr>
488struct blas_traits<CwiseUnaryOp<scalar_conjugate_op<Scalar>, NestedXpr> > : blas_traits<NestedXpr> {
489 using Base = blas_traits<NestedXpr>;
490 using XprType = CwiseUnaryOp<scalar_conjugate_op<Scalar>, NestedXpr>;
491 using ExtractType = typename Base::ExtractType;
492
493 enum { IsComplex = NumTraits<Scalar>::IsComplex, NeedToConjugate = Base::NeedToConjugate ? 0 : IsComplex };
494 EIGEN_DEVICE_FUNC static inline ExtractType extract(const XprType& x) { return Base::extract(x.nestedExpression()); }
495 EIGEN_DEVICE_FUNC static inline Scalar extractScalarFactor(const XprType& x) {
496 return conj(Base::extractScalarFactor(x.nestedExpression()));
497 }
498};
499
500// pop scalar multiple
501template <typename Scalar, typename NestedXpr, typename Plain>
502struct blas_traits<
503 CwiseBinaryOp<scalar_product_op<Scalar>, const CwiseNullaryOp<scalar_constant_op<Scalar>, Plain>, NestedXpr> >
504 : blas_traits<NestedXpr> {
505 enum { HasScalarFactor = true };
506 using Base = blas_traits<NestedXpr>;
507 using XprType =
508 CwiseBinaryOp<scalar_product_op<Scalar>, const CwiseNullaryOp<scalar_constant_op<Scalar>, Plain>, NestedXpr>;
509 using ExtractType = typename Base::ExtractType;
510 EIGEN_DEVICE_FUNC static inline EIGEN_DEVICE_FUNC ExtractType extract(const XprType& x) {
511 return Base::extract(x.rhs());
512 }
513 EIGEN_DEVICE_FUNC static inline EIGEN_DEVICE_FUNC Scalar extractScalarFactor(const XprType& x) {
514 return x.lhs().functor().m_other * Base::extractScalarFactor(x.rhs());
515 }
516};
517template <typename Scalar, typename NestedXpr, typename Plain>
518struct blas_traits<
519 CwiseBinaryOp<scalar_product_op<Scalar>, NestedXpr, const CwiseNullaryOp<scalar_constant_op<Scalar>, Plain> > >
520 : blas_traits<NestedXpr> {
521 enum { HasScalarFactor = true };
522 using Base = blas_traits<NestedXpr>;
523 using XprType =
524 CwiseBinaryOp<scalar_product_op<Scalar>, NestedXpr, const CwiseNullaryOp<scalar_constant_op<Scalar>, Plain>>;
525 using ExtractType = typename Base::ExtractType;
526 EIGEN_DEVICE_FUNC static inline ExtractType extract(const XprType& x) { return Base::extract(x.lhs()); }
527 EIGEN_DEVICE_FUNC static inline Scalar extractScalarFactor(const XprType& x) {
528 return Base::extractScalarFactor(x.lhs()) * x.rhs().functor().m_other;
529 }
530};
531template <typename Scalar, typename Plain1, typename Plain2>
532struct blas_traits<CwiseBinaryOp<scalar_product_op<Scalar>, const CwiseNullaryOp<scalar_constant_op<Scalar>, Plain1>,
533 const CwiseNullaryOp<scalar_constant_op<Scalar>, Plain2> > >
534 : blas_traits<CwiseNullaryOp<scalar_constant_op<Scalar>, Plain1> > {};
535
536// pop opposite
537template <typename Scalar, typename NestedXpr>
538struct blas_traits<CwiseUnaryOp<scalar_opposite_op<Scalar>, NestedXpr> > : blas_traits<NestedXpr> {
539 enum { HasScalarFactor = true };
540 using Base = blas_traits<NestedXpr>;
541 using XprType = CwiseUnaryOp<scalar_opposite_op<Scalar>, NestedXpr>;
542 using ExtractType = typename Base::ExtractType;
543 EIGEN_DEVICE_FUNC static inline ExtractType extract(const XprType& x) { return Base::extract(x.nestedExpression()); }
544 EIGEN_DEVICE_FUNC static inline Scalar extractScalarFactor(const XprType& x) {
545 return -Base::extractScalarFactor(x.nestedExpression());
546 }
547};
548
549// pop/push transpose
550template <typename NestedXpr>
551struct blas_traits<Transpose<NestedXpr> > : blas_traits<NestedXpr> {
552 using Scalar = typename NestedXpr::Scalar;
553 using Base = blas_traits<NestedXpr>;
554 using XprType = Transpose<NestedXpr>;
555 using ExtractType = Transpose<const typename Base::ExtractType_>; // const to get rid of a compile error; anyway blas
556 // traits are only used on the RHS
557 using ExtractType_ = Transpose<const typename Base::ExtractType_>;
558 using DirectLinearAccessType =
559 std::conditional_t<bool(Base::HasUsableDirectAccess), ExtractType, typename ExtractType::PlainObject>;
560 enum { IsTransposed = Base::IsTransposed ? 0 : 1 };
561 EIGEN_DEVICE_FUNC static inline ExtractType extract(const XprType& x) {
562 return ExtractType(Base::extract(x.nestedExpression()));
563 }
564 EIGEN_DEVICE_FUNC static inline Scalar extractScalarFactor(const XprType& x) {
565 return Base::extractScalarFactor(x.nestedExpression());
566 }
567};
568
569template <typename T>
570struct blas_traits<const T> : blas_traits<T> {};
571
572template <typename T, bool HasUsableDirectAccess = blas_traits<T>::HasUsableDirectAccess>
573struct extract_data_selector {
574 EIGEN_DEVICE_FUNC constexpr EIGEN_ALWAYS_INLINE static const typename T::Scalar* run(const T& m) {
575 return blas_traits<T>::extract(m).data();
576 }
577};
578
579template <typename T>
580struct extract_data_selector<T, false> {
581 EIGEN_DEVICE_FUNC constexpr static typename T::Scalar* run(const T&) { return 0; }
582};
583
584template <typename T>
585EIGEN_DEVICE_FUNC constexpr EIGEN_ALWAYS_INLINE const typename T::Scalar* extract_data(const T& m) {
586 return extract_data_selector<T>::run(m);
587}
588
593template <typename ResScalar, typename Lhs, typename Rhs>
595 EIGEN_DEVICE_FUNC constexpr EIGEN_ALWAYS_INLINE static ResScalar run(const Lhs& lhs, const Rhs& rhs) {
596 return blas_traits<Lhs>::extractScalarFactor(lhs) * blas_traits<Rhs>::extractScalarFactor(rhs);
597 }
598 EIGEN_DEVICE_FUNC constexpr EIGEN_ALWAYS_INLINE static ResScalar run(const ResScalar& alpha, const Lhs& lhs,
599 const Rhs& rhs) {
600 return alpha * blas_traits<Lhs>::extractScalarFactor(lhs) * blas_traits<Rhs>::extractScalarFactor(rhs);
601 }
602};
603template <typename Lhs, typename Rhs>
604struct combine_scalar_factors_impl<bool, Lhs, Rhs> {
605 EIGEN_DEVICE_FUNC constexpr EIGEN_ALWAYS_INLINE static bool run(const Lhs& lhs, const Rhs& rhs) {
606 return blas_traits<Lhs>::extractScalarFactor(lhs) && blas_traits<Rhs>::extractScalarFactor(rhs);
607 }
608 EIGEN_DEVICE_FUNC constexpr EIGEN_ALWAYS_INLINE static bool run(const bool& alpha, const Lhs& lhs, const Rhs& rhs) {
609 return alpha && blas_traits<Lhs>::extractScalarFactor(lhs) && blas_traits<Rhs>::extractScalarFactor(rhs);
610 }
611};
612
613template <typename ResScalar, typename Lhs, typename Rhs>
614EIGEN_DEVICE_FUNC constexpr EIGEN_ALWAYS_INLINE ResScalar combine_scalar_factors(const ResScalar& alpha, const Lhs& lhs,
615 const Rhs& rhs) {
616 return combine_scalar_factors_impl<ResScalar, Lhs, Rhs>::run(alpha, lhs, rhs);
617}
618template <typename ResScalar, typename Lhs, typename Rhs>
619EIGEN_DEVICE_FUNC constexpr EIGEN_ALWAYS_INLINE ResScalar combine_scalar_factors(const Lhs& lhs, const Rhs& rhs) {
620 return combine_scalar_factors_impl<ResScalar, Lhs, Rhs>::run(lhs, rhs);
621}
622
623} // end namespace internal
624
625} // end namespace Eigen
626
627#endif // EIGEN_BLASUTIL_H
AlignmentType
Definition Constants.h:235
@ RowMajor
Definition Constants.h:321
constexpr unsigned int DirectAccessBit
Definition Constants.h:160