11#ifndef EIGEN_BLASUTIL_H
12#define EIGEN_BLASUTIL_H
18#include "../InternalHeaderCheck.h"
25template <
typename LhsScalar,
typename RhsScalar,
typename Index,
typename DataMapper,
int mr,
int nr,
26 bool ConjugateLhs =
false,
bool ConjugateRhs =
false>
29template <
typename Scalar,
typename Index,
typename DataMapper,
int nr,
int StorageOrder,
bool Conjugate =
false,
30 bool PanelMode =
false>
33template <
typename Scalar,
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
int StorageOrder,
34 bool Conjugate =
false,
bool PanelMode =
false>
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;
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;
45template <
typename From,
typename To>
47 EIGEN_DEVICE_FUNC
constexpr static EIGEN_STRONG_INLINE To run(
const From& x) {
return To(x); }
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);
57template <
typename Scalar,
typename Index>
58class BlasVectorMapper {
60 EIGEN_DEVICE_FUNC
constexpr EIGEN_ALWAYS_INLINE BlasVectorMapper(Scalar* data) : m_data(data) {}
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);
68 template <
typename Packet>
69 EIGEN_DEVICE_FUNC
bool aligned(Index i)
const {
70 return (std::uintptr_t(m_data + i) %
sizeof(Packet)) == 0;
77template <
typename Scalar,
typename Index,
int AlignmentType,
int Incr = 1>
78class BlasLinearMapper;
80template <
typename Scalar,
typename Index,
int AlignmentType>
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);
88 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE
void prefetch(Index i)
const { internal::prefetch(&
operator()(i)); }
90 EIGEN_DEVICE_FUNC
constexpr EIGEN_ALWAYS_INLINE Scalar& operator()(Index i)
const {
return m_data[i]; }
92 template <
typename PacketType>
93 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE PacketType loadPacket(Index i)
const {
94 return ploadt<PacketType, AlignmentType>(m_data + i);
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);
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);
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);
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);
123template <
typename Scalar,
typename Index,
int StorageOrder,
int AlignmentType = Unaligned,
int Incr = 1>
124class blas_data_mapper;
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]);
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]);
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);
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);
175template <
typename Scalar,
typename Index,
int StorageOrder,
int AlignmentType>
176class blas_data_mapper<Scalar, Index, StorageOrder,
AlignmentType, 1> {
178 using LinearMapper = BlasLinearMapper<Scalar, Index, AlignmentType>;
179 using SubMapper = blas_data_mapper<Scalar, Index, StorageOrder, AlignmentType>;
180 using VectorMapper = BlasVectorMapper<Scalar, Index>;
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);
188 EIGEN_DEVICE_FUNC
constexpr EIGEN_ALWAYS_INLINE SubMapper getSubMapper(Index i, Index j)
const {
189 return SubMapper(&
operator()(i, j), m_stride);
192 EIGEN_DEVICE_FUNC
constexpr EIGEN_ALWAYS_INLINE LinearMapper getLinearMapper(Index i, Index j)
const {
193 return LinearMapper(&
operator()(i, j));
196 EIGEN_DEVICE_FUNC
constexpr EIGEN_ALWAYS_INLINE VectorMapper getVectorMapper(Index i, Index j)
const {
197 return VectorMapper(&
operator()(i, j));
200 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE
void prefetch(Index i, Index j)
const { internal::prefetch(&
operator()(i, j)); }
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];
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));
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);
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));
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);
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);
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);
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);
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; }
247 EIGEN_DEVICE_FUNC Index firstAligned(Index size)
const {
248 if (std::uintptr_t(m_data) %
sizeof(Scalar)) {
251 return internal::first_default_aligned(m_data, size);
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);
262 Scalar* EIGEN_RESTRICT m_data;
263 const Index m_stride;
269template <
typename Scalar,
typename Index,
int AlignmentType,
int Incr>
270class BlasLinearMapper {
272 EIGEN_DEVICE_FUNC
constexpr EIGEN_ALWAYS_INLINE BlasLinearMapper(Scalar* data, Index incr)
273 : m_data(data), m_incr(incr) {}
275 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE
void prefetch(
int i)
const { internal::prefetch(&
operator()(i)); }
277 EIGEN_DEVICE_FUNC
constexpr EIGEN_ALWAYS_INLINE Scalar& operator()(Index i)
const {
278 return m_data[i * m_incr.value()];
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());
286 template <
typename PacketType>
287 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE PacketType loadPacketPartial(Index i, Index n, Index = 0)
const {
288 return pgather_partial<Scalar, PacketType>(m_data + i * m_incr.value(), m_incr.value(), n);
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());
296 template <
typename PacketType>
297 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE
void storePacketPartial(Index i,
const PacketType& p, Index n,
299 pscatter_partial<Scalar, PacketType>(m_data + i * m_incr.value(), p, m_incr.value(), n);
304 const internal::variable_if_dynamic<Index, Incr> m_incr;
307template <
typename Scalar,
typename Index,
int StorageOrder,
int AlignmentType,
int Incr>
308class blas_data_mapper {
310 using LinearMapper = BlasLinearMapper<Scalar, Index, AlignmentType, Incr>;
311 using SubMapper = blas_data_mapper;
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) {}
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());
320 EIGEN_DEVICE_FUNC
constexpr EIGEN_ALWAYS_INLINE LinearMapper getLinearMapper(Index i, Index j)
const {
321 return LinearMapper(&
operator()(i, j), m_incr.value());
324 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE
void prefetch(Index i, Index j)
const { internal::prefetch(&
operator()(i, j)); }
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];
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());
335 template <
typename PacketType>
336 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE PacketType loadPacketPartial(Index i, Index j, Index n,
338 return pgather_partial<Scalar, PacketType>(&
operator()(i, j), m_incr.value(), n);
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());
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());
351 template <
typename PacketType>
352 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE
void storePacketPartial(Index i, Index j,
const PacketType& p, Index n,
354 pscatter_partial<Scalar, PacketType>(&
operator()(i, j), p, m_incr.value(), n);
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);
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);
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]);
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]);
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]);
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 {}
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 {}
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 {}
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);
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; }
440 Scalar* EIGEN_RESTRICT m_data;
441 const Index m_stride;
442 const internal::variable_if_dynamic<Index, Incr> m_incr;
446template <
typename Scalar,
typename Index,
int StorageOrder>
447class const_blas_data_mapper :
public blas_data_mapper<const Scalar, Index, StorageOrder> {
449 using SubMapper = const_blas_data_mapper<Scalar, Index, StorageOrder>;
451 EIGEN_ALWAYS_INLINE const_blas_data_mapper(
const Scalar* data, Index stride)
452 : blas_data_mapper<const Scalar, Index, StorageOrder>(data, stride) {}
454 EIGEN_ALWAYS_INLINE SubMapper getSubMapper(Index i, Index j)
const {
455 return SubMapper(&(this->
operator()(i, j)), this->m_stride);
462template <
typename XprType>
464 using Scalar =
typename traits<XprType>::Scalar;
465 using ExtractType =
const XprType&;
466 using ExtractType_ = XprType;
468 IsComplex = NumTraits<Scalar>::IsComplex,
469 IsTransposed =
false,
470 NeedToConjugate =
false,
471 HasUsableDirectAccess =
473 (bool(XprType::IsVectorAtCompileTime) || int(inner_stride_at_compile_time<XprType>::value) == 1))
476 HasScalarFactor = false
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&) {
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;
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()));
501template <
typename Scalar,
typename NestedXpr,
typename Plain>
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>;
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());
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());
517template <
typename Scalar,
typename NestedXpr,
typename Plain>
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>;
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;
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> > {};
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());
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_>;
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()));
564 EIGEN_DEVICE_FUNC
static inline Scalar extractScalarFactor(
const XprType& x) {
565 return Base::extractScalarFactor(x.nestedExpression());
570struct blas_traits<const T> : blas_traits<T> {};
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();
580struct extract_data_selector<T, false> {
581 EIGEN_DEVICE_FUNC
constexpr static typename T::Scalar* run(
const T&) {
return 0; }
585EIGEN_DEVICE_FUNC
constexpr EIGEN_ALWAYS_INLINE
const typename T::Scalar* extract_data(
const T& m) {
586 return extract_data_selector<T>::run(m);
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);
598 EIGEN_DEVICE_FUNC
constexpr EIGEN_ALWAYS_INLINE
static ResScalar run(
const ResScalar& alpha,
const Lhs& lhs,
600 return alpha * blas_traits<Lhs>::extractScalarFactor(lhs) * blas_traits<Rhs>::extractScalarFactor(rhs);
603template <
typename Lhs,
typename 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);
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);
613template <
typename ResScalar,
typename Lhs,
typename Rhs>
614EIGEN_DEVICE_FUNC
constexpr EIGEN_ALWAYS_INLINE ResScalar combine_scalar_factors(
const ResScalar& alpha,
const Lhs& lhs,
616 return combine_scalar_factors_impl<ResScalar, Lhs, Rhs>::run(alpha, lhs, rhs);
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);
AlignmentType
Definition Constants.h:235
@ RowMajor
Definition Constants.h:321
constexpr unsigned int DirectAccessBit
Definition Constants.h:160
Definition BlasUtil.h:594