Eigen  5.0.1
 
Loading...
Searching...
No Matches
InnerProduct.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2024 Charlie Schlosser <cs.schlosser@gmail.com>
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_INNER_PRODUCT_EVAL_H
12#define EIGEN_INNER_PRODUCT_EVAL_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20
21// Shared accumulation kernel for dot() (Conj = true) and vector products (Conj = false).
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)
27#ifndef EIGEN_NO_DEBUG
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");
30 }
31#else
32 static EIGEN_DEVICE_FUNC void run(const Lhs&, const Rhs&) {}
33#endif
34};
35
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;
45
46 using Scalar = typename Func::result_type;
47 using Packet = typename find_largest_packet<Scalar, SizeAtCompileTime>::type;
48
49 static constexpr bool Vectorize =
50 bool(LhsFlags & RhsFlags & PacketAccessBit) && Func::PacketAccess &&
51 ((MaxSizeAtCompileTime == Dynamic) || (unpacket_traits<Packet>::size <= MaxSizeAtCompileTime));
52
53 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE explicit inner_product_evaluator(const Lhs& lhs, const Rhs& rhs,
54 Func func = Func())
55 : m_func(func), m_lhs(lhs), m_rhs(rhs), m_size(lhs.size()) {
56 inner_product_assert<Lhs, Rhs>::run(lhs, rhs);
57 }
58
59 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index size() const { return m_size.value(); }
60
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));
63 }
64
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));
67 }
68
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));
73 }
74
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));
79 }
80
81 const Func m_func;
82 const evaluator<Lhs> m_lhs;
83 const evaluator<Rhs> m_rhs;
84 const variable_if_dynamic<Index, SizeAtCompileTime> m_size;
85};
86
87template <typename Evaluator, bool Vectorize = Evaluator::Vectorize>
88struct inner_product_impl;
89
90// scalar loop
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);
97
98 Scalar result = eval.coeff(0);
99 for (Index k = 1; k < size; k++) {
100 result = eval.coeff(result, k);
101 }
102
103 return result;
104 }
105};
106
107// vector loop
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);
118
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;
123
124 Packet presult0 = eval.template packet<Packet>(0 * PacketSize);
125 // Exclude unreachable packet loads for bounded vectors (also avoids GCC bounds warnings).
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);
135
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);
141 }
142
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);
151 }
152 }
153 }
154 }
155 }
156
157 presult2 = padd(presult2, presult3);
158 }
159 }
160 presult1 = padd(presult1, presult2);
161 }
162 }
163 presult0 = padd(presult0, presult1);
164 }
165 }
166
167 Scalar result = predux(presult0);
168 for (UnsignedIndex k = packetEnd; k < size; k++) {
169 result = eval.coeff(result, k);
170 }
171
172 return result;
173 }
174};
175
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);
181 }
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;
185 }
186 static constexpr bool PacketAccess = false;
187};
188
189// Partial specialization for packet access if and only if
190// LhsScalar == RhsScalar == ScalarBinaryOpTraits<LhsScalar, RhsScalar>::ReturnType.
191template <typename Scalar, bool Conj>
192struct scalar_inner_product_op<
193 Scalar,
194 std::enable_if_t<std::is_same<typename ScalarBinaryOpTraits<Scalar, Scalar>::ReturnType, Scalar>::value, Scalar>,
195 Conj> {
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);
199 }
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);
202 }
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);
206 }
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);
210 }
211 static constexpr bool PacketAccess = packet_traits<Scalar>::HasMul && packet_traits<Scalar>::HasAdd;
212};
213
214// Backends specialize Enable = enable_if_t<...> (void); Enable = false_type always names this generic definition.
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);
225 }
226};
227
228template <typename T>
229struct unwrap_unary {
230 using type = T;
231 static constexpr bool HasDirectAccess = bool(traits<type>::Flags & DirectAccessBit);
232
233 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE constexpr T const& get(T const& xpr) { return xpr; }
234};
235
236template <typename T>
237struct unwrap_unary<T const> : unwrap_unary<T> {};
238
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());
244 }
245};
246
247template <typename T, typename Target>
248struct rewrap_unary {
249 using type = Target;
250
251 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE constexpr Target const& apply(T const&, Target const& target) {
252 return target;
253 }
254};
255
256template <typename T, typename Target>
257struct rewrap_unary<T const, Target> : rewrap_unary<T, Target const> {};
258
259template <typename Op, typename Xpr, typename Target>
260struct rewrap_unary<CwiseUnaryOp<Op, Xpr>, Target> {
261 // unaryExpr() nests its operand as const; Xpr itself need not be (adjoint() nests a non-const Transpose).
262 using type = CwiseUnaryOp<Op, std::add_const_t<typename rewrap_unary<Xpr, Target>::type>>;
263
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());
267 }
268};
269
270// Rewrap coefficient-wise unary operations around contiguous maps only when both
271// operands expose storage and at least one lacks a compile-time unit inner stride.
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> {};
277
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;
282
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);
286 }
287 // Keep tiny products inlined without the remapping and packet-loop setup.
288 if (a.size() <= 4) {
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);
295 return result;
296 }
297 return run_large_product(a, b);
298 }
299
300 // Keep the larger product kernel out of tiny callers, while dot() retains its existing inlining.
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);
304 }
305
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;
312
313 LhsInner const& lhs_inner = LhsUnwrapper::get(a.derived());
314 RhsInner const& rhs_inner = RhsUnwrapper::get(b.derived());
315
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>;
321
322 LhsMap const lhs_map(lhs_inner.data(), lhs_inner.size());
323 RhsMap const rhs_map(rhs_inner.data(), rhs_inner.size());
324
325 using LhsRewrap = rewrap_unary<Lhs, LhsMap>;
326 using RhsRewrap = rewrap_unary<Rhs, RhsMap>;
327
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));
330 }
331
332 return Impl::run(a, b);
333 }
334};
335
336} // namespace internal
337} // namespace Eigen
338
339#endif // EIGEN_INNER_PRODUCT_EVAL_H
constexpr unsigned int PacketAccessBit
Definition Constants.h:98
constexpr unsigned int DirectAccessBit
Definition Constants.h:160