Eigen  5.0.1
 
Loading...
Searching...
No Matches
Product.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2008-2011 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_PRODUCT_H
12#define EIGEN_PRODUCT_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19template <typename Lhs, typename Rhs, int Option, typename StorageKind>
20class ProductImpl;
21
22namespace internal {
23
24template <typename Lhs, typename Rhs, int Option>
25struct traits<Product<Lhs, Rhs, Option>> {
26 using LhsCleaned = remove_all_t<Lhs>;
27 using RhsCleaned = remove_all_t<Rhs>;
28 using LhsTraits = traits<LhsCleaned>;
29 using RhsTraits = traits<RhsCleaned>;
30
31 using XprKind = MatrixXpr;
32
33 using Scalar = typename ScalarBinaryOpTraits<typename LhsTraits::Scalar, typename RhsTraits::Scalar>::ReturnType;
34 using StorageKind =
35 typename product_promote_storage_type<typename LhsTraits::StorageKind, typename RhsTraits::StorageKind,
36 internal::product_type<Lhs, Rhs>::value>::ret;
37 using StorageIndex =
38 typename promote_index_type<typename LhsTraits::StorageIndex, typename RhsTraits::StorageIndex>::type;
39
40 enum {
41 RowsAtCompileTime = LhsTraits::RowsAtCompileTime,
42 ColsAtCompileTime = RhsTraits::ColsAtCompileTime,
43 MaxRowsAtCompileTime = LhsTraits::MaxRowsAtCompileTime,
44 MaxColsAtCompileTime = RhsTraits::MaxColsAtCompileTime,
45
46 // FIXME: only needed by GeneralMatrixMatrixTriangular
47 InnerSize = min_size_prefer_fixed(LhsTraits::ColsAtCompileTime, RhsTraits::RowsAtCompileTime),
48
49 // The storage order is somewhat arbitrary here. The correct one will be determined through the evaluator.
50 Flags = (MaxRowsAtCompileTime == 1 && MaxColsAtCompileTime != 1) ? RowMajorBit
51 : (MaxColsAtCompileTime == 1 && MaxRowsAtCompileTime != 1) ? 0
52 : (((LhsTraits::Flags & NoPreferredStorageOrderBit) && (RhsTraits::Flags & RowMajorBit)) ||
53 ((RhsTraits::Flags & NoPreferredStorageOrderBit) && (LhsTraits::Flags & RowMajorBit)))
56 };
57};
58
59struct TransposeProductEnum {
60 // convenience enumerations to specialize transposed products
61 enum : int {
62 Default = 0x00,
63 Matrix = 0x01,
64 Permutation = 0x02,
65 MatrixMatrix = (Matrix << 8) | Matrix,
66 MatrixPermutation = (Matrix << 8) | Permutation,
67 PermutationMatrix = (Permutation << 8) | Matrix
68 };
69};
70template <typename Xpr>
71struct TransposeKind {
72 static constexpr int Kind = is_matrix_base_xpr<Xpr>::value ? TransposeProductEnum::Matrix
73 : is_permutation_base_xpr<Xpr>::value ? TransposeProductEnum::Permutation
74 : TransposeProductEnum::Default;
75};
76
77template <typename Lhs, typename Rhs>
78struct TransposeProductKind {
79 static constexpr int Kind = (TransposeKind<Lhs>::Kind << 8) | TransposeKind<Rhs>::Kind;
80};
81
82template <typename Lhs, typename Rhs, int Option, int Kind = TransposeProductKind<Lhs, Rhs>::Kind>
83struct product_transpose_helper {
84 // by default, don't optimize the transposed product
85 using Derived = Product<Lhs, Rhs, Option>;
86 using Scalar = typename Derived::Scalar;
87 using TransposeType = Transpose<const Derived>;
88 using ConjugateTransposeType = CwiseUnaryOp<scalar_conjugate_op<Scalar>, TransposeType>;
89 using AdjointType = std::conditional_t<NumTraits<Scalar>::IsComplex, ConjugateTransposeType, TransposeType>;
90
91 // return (lhs * rhs)^T
92 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TransposeType run_transpose(const Derived& derived) {
93 return TransposeType(derived);
94 }
95 // return (lhs * rhs)^H
96 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE AdjointType run_adjoint(const Derived& derived) {
97 return AdjointType(TransposeType(derived));
98 }
99};
100
101template <typename Lhs, typename Rhs, int Option>
102struct product_transpose_helper<Lhs, Rhs, Option, TransposeProductEnum::MatrixMatrix> {
103 // expand the transposed matrix-matrix product
104 using Derived = Product<Lhs, Rhs, Option>;
105
106 using LhsScalar = typename traits<Lhs>::Scalar;
107 using LhsTransposeType = typename DenseBase<Lhs>::ConstTransposeReturnType;
108 using LhsConjugateTransposeType = CwiseUnaryOp<scalar_conjugate_op<LhsScalar>, LhsTransposeType>;
109 using LhsAdjointType =
110 std::conditional_t<NumTraits<LhsScalar>::IsComplex, LhsConjugateTransposeType, LhsTransposeType>;
111
112 using RhsScalar = typename traits<Rhs>::Scalar;
113 using RhsTransposeType = typename DenseBase<Rhs>::ConstTransposeReturnType;
114 using RhsConjugateTransposeType = CwiseUnaryOp<scalar_conjugate_op<RhsScalar>, RhsTransposeType>;
115 using RhsAdjointType =
116 std::conditional_t<NumTraits<RhsScalar>::IsComplex, RhsConjugateTransposeType, RhsTransposeType>;
117
118 using TransposeType = Product<RhsTransposeType, LhsTransposeType, Option>;
119 using AdjointType = Product<RhsAdjointType, LhsAdjointType, Option>;
120
121 // return rhs^T * lhs^T
122 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TransposeType run_transpose(const Derived& derived) {
123 return TransposeType(RhsTransposeType(derived.rhs()), LhsTransposeType(derived.lhs()));
124 }
125 // return rhs^H * lhs^H
126 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE AdjointType run_adjoint(const Derived& derived) {
127 return AdjointType(RhsAdjointType(RhsTransposeType(derived.rhs())),
128 LhsAdjointType(LhsTransposeType(derived.lhs())));
129 }
130};
131template <typename Lhs, typename Rhs, int Option>
132struct product_transpose_helper<Lhs, Rhs, Option, TransposeProductEnum::PermutationMatrix> {
133 // expand the transposed permutation-matrix product
134 using Derived = Product<Lhs, Rhs, Option>;
135
136 using LhsInverseType = typename PermutationBase<Lhs>::InverseReturnType;
137
138 using RhsScalar = typename traits<Rhs>::Scalar;
139 using RhsTransposeType = typename DenseBase<Rhs>::ConstTransposeReturnType;
140 using RhsConjugateTransposeType = CwiseUnaryOp<scalar_conjugate_op<RhsScalar>, RhsTransposeType>;
141 using RhsAdjointType =
142 std::conditional_t<NumTraits<RhsScalar>::IsComplex, RhsConjugateTransposeType, RhsTransposeType>;
143
144 using TransposeType = Product<RhsTransposeType, LhsInverseType, Option>;
145 using AdjointType = Product<RhsAdjointType, LhsInverseType, Option>;
146
147 // return rhs^T * lhs^-1
148 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TransposeType run_transpose(const Derived& derived) {
149 return TransposeType(RhsTransposeType(derived.rhs()), LhsInverseType(derived.lhs()));
150 }
151 // return rhs^H * lhs^-1
152 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE AdjointType run_adjoint(const Derived& derived) {
153 return AdjointType(RhsAdjointType(RhsTransposeType(derived.rhs())), LhsInverseType(derived.lhs()));
154 }
155};
156template <typename Lhs, typename Rhs, int Option>
157struct product_transpose_helper<Lhs, Rhs, Option, TransposeProductEnum::MatrixPermutation> {
158 // expand the transposed matrix-permutation product
159 using Derived = Product<Lhs, Rhs, Option>;
160
161 using LhsScalar = typename traits<Lhs>::Scalar;
162 using LhsTransposeType = typename DenseBase<Lhs>::ConstTransposeReturnType;
163 using LhsConjugateTransposeType = CwiseUnaryOp<scalar_conjugate_op<LhsScalar>, LhsTransposeType>;
164 using LhsAdjointType =
165 std::conditional_t<NumTraits<LhsScalar>::IsComplex, LhsConjugateTransposeType, LhsTransposeType>;
166
167 using RhsInverseType = typename PermutationBase<Rhs>::InverseReturnType;
168
169 using TransposeType = Product<RhsInverseType, LhsTransposeType, Option>;
170 using AdjointType = Product<RhsInverseType, LhsAdjointType, Option>;
171
172 // return rhs^-1 * lhs^T
173 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TransposeType run_transpose(const Derived& derived) {
174 return TransposeType(RhsInverseType(derived.rhs()), LhsTransposeType(derived.lhs()));
175 }
176 // return rhs^-1 * lhs^H
177 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE AdjointType run_adjoint(const Derived& derived) {
178 return AdjointType(RhsInverseType(derived.rhs()), LhsAdjointType(LhsTransposeType(derived.lhs())));
179 }
180};
181
182} // end namespace internal
183
198template <typename Lhs_, typename Rhs_, int Option>
199class Product
200 : public ProductImpl<Lhs_, Rhs_, Option,
201 typename internal::product_promote_storage_type<
202 typename internal::traits<Lhs_>::StorageKind, typename internal::traits<Rhs_>::StorageKind,
203 internal::product_type<Lhs_, Rhs_>::value>::ret> {
204 public:
205 using Lhs = Lhs_;
206 using Rhs = Rhs_;
207
208 using Base =
209 typename ProductImpl<Lhs, Rhs, Option,
210 typename internal::product_promote_storage_type<
211 typename internal::traits<Lhs>::StorageKind, typename internal::traits<Rhs>::StorageKind,
212 internal::product_type<Lhs, Rhs>::value>::ret>::Base;
213 EIGEN_GENERIC_PUBLIC_INTERFACE(Product)
214
215 using LhsNested = typename internal::ref_selector<Lhs>::type;
216 using RhsNested = typename internal::ref_selector<Rhs>::type;
217 using LhsNestedCleaned = internal::remove_all_t<LhsNested>;
218 using RhsNestedCleaned = internal::remove_all_t<RhsNested>;
219
220 using TransposeReturnType = typename internal::product_transpose_helper<Lhs, Rhs, Option>::TransposeType;
221 using AdjointReturnType = typename internal::product_transpose_helper<Lhs, Rhs, Option>::AdjointType;
222
223 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE Product(const Lhs& lhs, const Rhs& rhs) : m_lhs(lhs), m_rhs(rhs) {
224 eigen_assert(lhs.cols() == rhs.rows() && "invalid matrix product" &&
225 "if you wanted a coeff-wise or a dot product use the respective explicit functions");
226 }
227
228 EIGEN_DEVICE_FUNC constexpr Index rows() const noexcept { return m_lhs.rows(); }
229 EIGEN_DEVICE_FUNC constexpr Index cols() const noexcept { return m_rhs.cols(); }
230
231 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE const LhsNestedCleaned& lhs() const { return m_lhs; }
232 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE const RhsNestedCleaned& rhs() const { return m_rhs; }
233
234 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TransposeReturnType transpose() const {
235 return internal::product_transpose_helper<Lhs, Rhs, Option>::run_transpose(*this);
236 }
237 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE AdjointReturnType adjoint() const {
238 return internal::product_transpose_helper<Lhs, Rhs, Option>::run_adjoint(*this);
239 }
240
241 protected:
242 LhsNested m_lhs;
243 RhsNested m_rhs;
244};
245
246namespace internal {
247
248template <typename Lhs, typename Rhs, int Option, int ProductTag = internal::product_type<Lhs, Rhs>::value>
249class dense_product_base : public internal::dense_xpr_base<Product<Lhs, Rhs, Option>>::type {};
250
252template <typename Lhs, typename Rhs, int Option>
253class dense_product_base<Lhs, Rhs, Option, InnerProduct>
254 : public internal::dense_xpr_base<Product<Lhs, Rhs, Option>>::type {
255 using ProductXpr = Product<Lhs, Rhs, Option>;
256 using Base = typename internal::dense_xpr_base<ProductXpr>::type;
257
258 public:
259 using Base::derived;
260 using Scalar = typename Base::Scalar;
261
262 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE operator const Scalar() const {
263 return internal::evaluator<ProductXpr>(derived()).coeff(0, 0);
264 }
265};
266
267} // namespace internal
268
269// Generic API dispatcher
270template <typename Lhs, typename Rhs, int Option, typename StorageKind>
271class ProductImpl : public internal::generic_xpr_base<Product<Lhs, Rhs, Option>, MatrixXpr, StorageKind>::type {
272 public:
273 using Base = typename internal::generic_xpr_base<Product<Lhs, Rhs, Option>, MatrixXpr, StorageKind>::type;
274};
275
276template <typename Lhs, typename Rhs, int Option>
277class ProductImpl<Lhs, Rhs, Option, Dense> : public internal::dense_product_base<Lhs, Rhs, Option> {
278 using Derived = Product<Lhs, Rhs, Option>;
279
280 public:
281 using Base = typename internal::dense_product_base<Lhs, Rhs, Option>;
282 EIGEN_DENSE_PUBLIC_INTERFACE(Derived)
283 protected:
284 enum {
285 IsOneByOne = (RowsAtCompileTime == 1 || RowsAtCompileTime == Dynamic) &&
286 (ColsAtCompileTime == 1 || ColsAtCompileTime == Dynamic),
287 EnableCoeff = IsOneByOne || Option == LazyProduct
288 };
289
290 public:
291 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar coeff(Index row, Index col) const {
292 EIGEN_STATIC_ASSERT(EnableCoeff, THIS_METHOD_IS_ONLY_FOR_INNER_OR_LAZY_PRODUCTS);
293 eigen_assert((Option == LazyProduct) || (this->rows() == 1 && this->cols() == 1));
294
295 return internal::evaluator<Derived>(derived()).coeff(row, col);
296 }
297
298 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar coeff(Index i) const {
299 EIGEN_STATIC_ASSERT(EnableCoeff, THIS_METHOD_IS_ONLY_FOR_INNER_OR_LAZY_PRODUCTS);
300 eigen_assert((Option == LazyProduct) || (this->rows() == 1 && this->cols() == 1));
301
302 return internal::evaluator<Derived>(derived()).coeff(i);
303 }
304};
305
306} // end namespace Eigen
307
308#endif // EIGEN_PRODUCT_H
Expression of the product of two arbitrary matrices or vectors.
Definition Product.h:203
constexpr unsigned int NoPreferredStorageOrderBit
Definition Constants.h:183
constexpr unsigned int RowMajorBit
Definition Constants.h:71
Definition Constants.h:557