11#ifndef EIGEN_INNER_PRODUCT_EVAL_H
12#define EIGEN_INNER_PRODUCT_EVAL_H
15#include "./InternalHeaderCheck.h"
22template <
typename Lhs,
typename Rhs>
23struct inner_product_assert {
24 EIGEN_STATIC_ASSERT_VECTOR_ONLY(Lhs)
25 EIGEN_STATIC_ASSERT_VECTOR_ONLY(Rhs)
26 EIGEN_STATIC_ASSERT_SAME_VECTOR_SIZE(Lhs, Rhs)
28 static EIGEN_DEVICE_FUNC
void run(
const Lhs& lhs,
const Rhs& rhs) {
29 eigen_assert((lhs.size() == rhs.size()) &&
"Inner product: lhs and rhs vectors must have same size");
32 static EIGEN_DEVICE_FUNC
void run(
const Lhs&,
const Rhs&) {}
36template <
typename Func,
typename Lhs,
typename Rhs>
37struct inner_product_evaluator {
38 static constexpr int LhsFlags = evaluator<Lhs>::Flags;
39 static constexpr int RhsFlags = evaluator<Rhs>::Flags;
40 static constexpr int SizeAtCompileTime = size_prefer_fixed(Lhs::SizeAtCompileTime, Rhs::SizeAtCompileTime);
41 static constexpr int MaxSizeAtCompileTime =
42 min_size_prefer_fixed(Lhs::MaxSizeAtCompileTime, Rhs::MaxSizeAtCompileTime);
43 static constexpr int LhsAlignment = evaluator<Lhs>::Alignment;
44 static constexpr int RhsAlignment = evaluator<Rhs>::Alignment;
46 using Scalar =
typename Func::result_type;
47 using Packet =
typename find_largest_packet<Scalar, SizeAtCompileTime>::type;
49 static constexpr bool Vectorize =
51 ((MaxSizeAtCompileTime == Dynamic) || (unpacket_traits<Packet>::size <= MaxSizeAtCompileTime));
53 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
explicit inner_product_evaluator(
const Lhs& lhs,
const Rhs& rhs,
55 : m_func(func), m_lhs(lhs), m_rhs(rhs), m_size(lhs.size()) {
56 inner_product_assert<Lhs, Rhs>::run(lhs, rhs);
59 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index size()
const {
return m_size.value(); }
61 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar coeff(Index index)
const {
62 return m_func.coeff(m_lhs.coeff(index), m_rhs.coeff(index));
65 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar coeff(
const Scalar& value, Index index)
const {
66 return m_func.coeff(value, m_lhs.coeff(index), m_rhs.coeff(index));
69 template <
typename PacketType,
int LhsMode = LhsAlignment,
int RhsMode = RhsAlignment>
70 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketType packet(Index index)
const {
71 return m_func.packet(m_lhs.template packet<LhsMode, PacketType>(index),
72 m_rhs.template packet<RhsMode, PacketType>(index));
75 template <
typename PacketType,
int LhsMode = LhsAlignment,
int RhsMode = RhsAlignment>
76 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketType packet(
const PacketType& value, Index index)
const {
77 return m_func.packet(value, m_lhs.template packet<LhsMode, PacketType>(index),
78 m_rhs.template packet<RhsMode, PacketType>(index));
82 const evaluator<Lhs> m_lhs;
83 const evaluator<Rhs> m_rhs;
84 const variable_if_dynamic<Index, SizeAtCompileTime> m_size;
87template <
typename Evaluator,
bool Vectorize = Evaluator::Vectorize>
88struct inner_product_impl;
91template <
typename Evaluator>
92struct inner_product_impl<Evaluator, false> {
93 using Scalar =
typename Evaluator::Scalar;
94 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar run(
const Evaluator& eval) {
95 const Index size = eval.size();
96 if (size == 0)
return Scalar(0);
98 Scalar result = eval.coeff(0);
99 for (Index k = 1; k < size; k++) {
100 result = eval.coeff(result, k);
108template <
typename Evaluator>
109struct inner_product_impl<Evaluator, true> {
110 using UnsignedIndex = std::make_unsigned_t<Index>;
111 using Scalar =
typename Evaluator::Scalar;
112 using Packet =
typename Evaluator::Packet;
113 static constexpr int PacketSize = unpacket_traits<Packet>::size;
114 static constexpr int MaxSize = Evaluator::MaxSizeAtCompileTime;
115 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar run(
const Evaluator& eval) {
116 const UnsignedIndex size =
static_cast<UnsignedIndex
>(eval.size());
117 if (size < PacketSize)
return inner_product_impl<Evaluator, false>::run(eval);
119 const UnsignedIndex packetEnd = numext::round_down(size, PacketSize);
120 const UnsignedIndex quadEnd = numext::round_down(size, 4 * PacketSize);
121 const UnsignedIndex numPackets = size / PacketSize;
122 const UnsignedIndex numRemPackets = (packetEnd - quadEnd) / PacketSize;
124 Packet presult0 = eval.template packet<Packet>(0 * PacketSize);
126 EIGEN_IF_CONSTEXPR (MaxSize == Dynamic || MaxSize / PacketSize >= 2) {
127 if (numPackets >= 2) {
128 Packet presult1 = eval.template packet<Packet>(1 * PacketSize);
129 EIGEN_IF_CONSTEXPR (MaxSize == Dynamic || MaxSize / PacketSize >= 3) {
130 if (numPackets >= 3) {
131 Packet presult2 = eval.template packet<Packet>(2 * PacketSize);
132 EIGEN_IF_CONSTEXPR (MaxSize == Dynamic || MaxSize / PacketSize >= 4) {
133 if (numPackets >= 4) {
134 Packet presult3 = eval.template packet<Packet>(3 * PacketSize);
136 for (UnsignedIndex k = 4 * PacketSize; k < quadEnd; k += 4 * PacketSize) {
137 presult0 = eval.packet(presult0, k + 0 * PacketSize);
138 presult1 = eval.packet(presult1, k + 1 * PacketSize);
139 presult2 = eval.packet(presult2, k + 2 * PacketSize);
140 presult3 = eval.packet(presult3, k + 3 * PacketSize);
143 EIGEN_IF_CONSTEXPR (MaxSize == Dynamic || MaxSize / PacketSize >= 5) {
144 if (numRemPackets >= 1) {
145 presult0 = eval.packet(presult0, quadEnd + 0 * PacketSize);
146 EIGEN_IF_CONSTEXPR (MaxSize == Dynamic || MaxSize / PacketSize >= 6) {
147 if (numRemPackets >= 2) {
148 presult1 = eval.packet(presult1, quadEnd + 1 * PacketSize);
149 EIGEN_IF_CONSTEXPR (MaxSize == Dynamic || MaxSize / PacketSize >= 7) {
150 if (numRemPackets == 3) presult2 = eval.packet(presult2, quadEnd + 2 * PacketSize);
157 presult2 = padd(presult2, presult3);
160 presult1 = padd(presult1, presult2);
163 presult0 = padd(presult0, presult1);
167 Scalar result = predux(presult0);
168 for (UnsignedIndex k = packetEnd; k < size; k++) {
169 result = eval.coeff(result, k);
176template <
typename LhsScalar,
typename RhsScalar,
bool Conj>
177struct scalar_inner_product_op {
178 using result_type =
typename ScalarBinaryOpTraits<LhsScalar, RhsScalar>::ReturnType;
179 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE result_type coeff(
const LhsScalar& a,
const RhsScalar& b)
const {
180 return (conj_if<Conj>()(a) * b);
182 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE result_type coeff(
const result_type& accum,
const LhsScalar& a,
183 const RhsScalar& b)
const {
184 return (conj_if<Conj>()(a) * b) + accum;
186 static constexpr bool PacketAccess =
false;
191template <
typename Scalar,
bool Conj>
192struct scalar_inner_product_op<
194 std::enable_if_t<std::is_same<typename ScalarBinaryOpTraits<Scalar, Scalar>::ReturnType, Scalar>::value, Scalar>,
196 using result_type = Scalar;
197 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar coeff(
const Scalar& a,
const Scalar& b)
const {
198 return pmul(conj_if<Conj>()(a), b);
200 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar coeff(
const Scalar& accum,
const Scalar& a,
const Scalar& b)
const {
201 return pmadd(conj_if<Conj>()(a), b, accum);
203 template <
typename Packet>
204 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packet(
const Packet& a,
const Packet& b)
const {
205 return pmul(conj_if<Conj>().pconj(a), b);
207 template <
typename Packet>
208 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packet(
const Packet& accum,
const Packet& a,
const Packet& b)
const {
209 return pmadd(conj_if<Conj>().pconj(a), b, accum);
211 static constexpr bool PacketAccess = packet_traits<Scalar>::HasMul && packet_traits<Scalar>::HasAdd;
215template <
typename Lhs,
typename Rhs,
bool Conj,
typename Enable =
void>
216struct default_inner_product_impl {
217 using LhsScalar =
typename traits<Lhs>::Scalar;
218 using RhsScalar =
typename traits<Rhs>::Scalar;
219 using Op = scalar_inner_product_op<LhsScalar, RhsScalar, Conj>;
220 using Evaluator = inner_product_evaluator<Op, Lhs, Rhs>;
221 using result_type =
typename Evaluator::Scalar;
222 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE result_type run(
const MatrixBase<Lhs>& a,
const MatrixBase<Rhs>& b) {
223 Evaluator eval(a.derived(), b.derived(), Op());
224 return inner_product_impl<Evaluator>::run(eval);
231 static constexpr bool HasDirectAccess = bool(traits<type>::Flags &
DirectAccessBit);
233 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE
constexpr T
const& get(T
const& xpr) {
return xpr; }
237struct unwrap_unary<T const> : unwrap_unary<T> {};
239template <
typename Op,
typename Xpr>
240struct unwrap_unary<CwiseUnaryOp<Op, Xpr>> : unwrap_unary<Xpr> {
241 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE
constexpr typename unwrap_unary<Xpr>::type
const& get(
242 CwiseUnaryOp<Op, Xpr>
const& xpr) {
243 return unwrap_unary<Xpr>::get(xpr.nestedExpression());
247template <
typename T,
typename Target>
251 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE
constexpr Target
const& apply(T
const&, Target
const& target) {
256template <
typename T,
typename Target>
257struct rewrap_unary<T const, Target> : rewrap_unary<T, Target const> {};
259template <
typename Op,
typename Xpr,
typename Target>
260struct rewrap_unary<CwiseUnaryOp<Op, Xpr>, Target> {
262 using type = CwiseUnaryOp<Op, std::add_const_t<typename rewrap_unary<Xpr, Target>::type>>;
264 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE
constexpr type apply(CwiseUnaryOp<Op, Xpr>
const& base,
265 Target
const& target) {
266 return rewrap_unary<Xpr, Target>::apply(base.nestedExpression(), target).unaryExpr(base.functor());
272template <
typename Lhs,
typename Rhs,
bool Conj,
273 bool MayMap = unwrap_unary<Lhs>::HasDirectAccess && unwrap_unary<Rhs>::HasDirectAccess &&
274 (inner_stride_at_compile_time<typename unwrap_unary<Lhs>::type>::value != 1 ||
275 inner_stride_at_compile_time<typename unwrap_unary<Rhs>::type>::value != 1)>
276struct inner_product_dispatch : default_inner_product_impl<Lhs, Rhs, Conj> {};
278template <
typename Lhs,
typename Rhs,
bool Conj>
279struct inner_product_dispatch<Lhs, Rhs, Conj, true> {
280 using Impl = default_inner_product_impl<Lhs, Rhs, Conj>;
281 using result_type =
typename Impl::result_type;
283 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE result_type run(
const MatrixBase<Lhs>& a,
const MatrixBase<Rhs>& b) {
284 EIGEN_IF_CONSTEXPR (Conj) {
285 return run_general(a, b);
289 typename Impl::Evaluator eval(a.derived(), b.derived());
290 if (eval.size() == 0)
return result_type(0);
291 result_type result = eval.coeff(0);
292 if (eval.size() > 1) result = eval.coeff(result, 1);
293 if (eval.size() > 2) result = eval.coeff(result, 2);
294 if (eval.size() > 3) result = eval.coeff(result, 3);
297 return run_large_product(a, b);
301 static EIGEN_DEVICE_FUNC EIGEN_DONT_INLINE result_type run_large_product(
const MatrixBase<Lhs>& a,
302 const MatrixBase<Rhs>& b) {
303 return run_general(a, b);
306 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE result_type run_general(
const MatrixBase<Lhs>& a,
307 const MatrixBase<Rhs>& b) {
308 using LhsUnwrapper = unwrap_unary<Lhs>;
309 using RhsUnwrapper = unwrap_unary<Rhs>;
310 using LhsInner =
typename LhsUnwrapper::type;
311 using RhsInner =
typename RhsUnwrapper::type;
313 LhsInner
const& lhs_inner = LhsUnwrapper::get(a.derived());
314 RhsInner
const& rhs_inner = RhsUnwrapper::get(b.derived());
316 if (lhs_inner.innerStride() == 1 && rhs_inner.innerStride() == 1) {
317 using LhsMap = Map<Vector<typename LhsInner::Scalar, size_of_xpr_at_compile_time<LhsInner>::value>
const,
318 evaluator<LhsInner>::Alignment>;
319 using RhsMap = Map<Vector<typename RhsInner::Scalar, size_of_xpr_at_compile_time<RhsInner>::value>
const,
320 evaluator<RhsInner>::Alignment>;
322 LhsMap
const lhs_map(lhs_inner.data(), lhs_inner.size());
323 RhsMap
const rhs_map(rhs_inner.data(), rhs_inner.size());
325 using LhsRewrap = rewrap_unary<Lhs, LhsMap>;
326 using RhsRewrap = rewrap_unary<Rhs, RhsMap>;
328 return default_inner_product_impl<typename LhsRewrap::type, typename RhsRewrap::type, Conj>::run(
329 LhsRewrap::apply(a.derived(), lhs_map), RhsRewrap::apply(b.derived(), rhs_map));
332 return Impl::run(a, b);
constexpr unsigned int PacketAccessBit
Definition Constants.h:98
constexpr unsigned int DirectAccessBit
Definition Constants.h:160