Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
KroneckerTensorProduct.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2011 Kolja Brix <brix@igpm.rwth-aachen.de>
5// Copyright (C) 2011 Andreas Platen <andiplaten@gmx.de>
6// Copyright (C) 2012 Chen-Pang He <jdh8@ms63.hinet.net>
7//
8// This Source Code Form is subject to the terms of the Mozilla
9// Public License v. 2.0. If a copy of the MPL was not distributed
10// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
11// SPDX-License-Identifier: MPL-2.0
12
13#ifndef KRONECKER_TENSOR_PRODUCT_H
14#define KRONECKER_TENSOR_PRODUCT_H
15
16// IWYU pragma: private
17#include "./InternalHeaderCheck.h"
18
19namespace Eigen {
20
28template <typename Derived>
29class KroneckerProductBase : public ReturnByValue<Derived> {
30 private:
31 typedef typename internal::traits<Derived> Traits;
32 typedef typename Traits::Scalar Scalar;
33
34 protected:
35 typedef typename Traits::Lhs Lhs;
36 typedef typename Traits::Rhs Rhs;
37
38 public:
40 KroneckerProductBase(const Lhs& A, const Rhs& B) : m_A(A), m_B(B) {}
41
42 inline Index rows() const { return m_A.rows() * m_B.rows(); }
43 inline Index cols() const { return m_A.cols() * m_B.cols(); }
44
49 Scalar coeff(Index row, Index col) const {
50 return m_A.coeff(row / m_B.rows(), col / m_B.cols()) * m_B.coeff(row % m_B.rows(), col % m_B.cols());
51 }
52
57 Scalar coeff(Index i) const {
58 EIGEN_STATIC_ASSERT_VECTOR_ONLY(Derived);
59 return m_A.coeff(i / m_B.size()) * m_B.coeff(i % m_B.size());
60 }
61
62 protected:
63 typename Lhs::Nested m_A;
64 typename Rhs::Nested m_B;
65};
66
79template <typename Lhs, typename Rhs>
80class KroneckerProduct : public KroneckerProductBase<KroneckerProduct<Lhs, Rhs> > {
81 private:
83 using Base::m_A;
84 using Base::m_B;
85
86 public:
88 KroneckerProduct(const Lhs& A, const Rhs& B) : Base(A, B) {}
89
91 template <typename Dest>
92 void evalTo(Dest& dst) const;
93};
94
110template <typename Lhs, typename Rhs>
111class KroneckerProductSparse : public KroneckerProductBase<KroneckerProductSparse<Lhs, Rhs> > {
112 private:
114 using Base::m_A;
115 using Base::m_B;
116
117 public:
119 KroneckerProductSparse(const Lhs& A, const Rhs& B) : Base(A, B) {}
120
122 template <typename Dest>
123 void evalTo(Dest& dst) const;
124};
125
126template <typename Lhs, typename Rhs>
127template <typename Dest>
129 const int BlockRows = Rhs::RowsAtCompileTime, BlockCols = Rhs::ColsAtCompileTime;
130 const Index Br = m_B.rows(), Bc = m_B.cols();
131 for (Index i = 0; i < m_A.rows(); ++i)
132 for (Index j = 0; j < m_A.cols(); ++j)
133 Block<Dest, BlockRows, BlockCols>(dst, i * Br, j * Bc, Br, Bc) = m_A.coeff(i, j) * m_B;
134}
135
136template <typename Lhs, typename Rhs>
137template <typename Dest>
139 Index Br = m_B.rows(), Bc = m_B.cols();
140 dst.resize(this->rows(), this->cols());
141 dst.resizeNonZeros(0);
142
143 // 1 - evaluate the operands if needed:
144 typedef typename internal::nested_eval<Lhs, Dynamic>::type Lhs1;
145 typedef internal::remove_all_t<Lhs1> Lhs1Cleaned;
146 const Lhs1 lhs1(m_A);
147 typedef typename internal::nested_eval<Rhs, Dynamic>::type Rhs1;
148 typedef internal::remove_all_t<Rhs1> Rhs1Cleaned;
149 const Rhs1 rhs1(m_B);
150
151 // 2 - construct respective iterators
152 typedef Eigen::InnerIterator<Lhs1Cleaned> LhsInnerIterator;
153 typedef Eigen::InnerIterator<Rhs1Cleaned> RhsInnerIterator;
154
155 // compute number of non-zeros per innervectors of dst
156 {
157 // TODO: VectorXi is not necessarily big enough!
158 VectorXi nnzA = VectorXi::Zero(Dest::IsRowMajor ? m_A.rows() : m_A.cols());
159 for (Index kA = 0; kA < m_A.outerSize(); ++kA)
160 for (LhsInnerIterator itA(lhs1, kA); itA; ++itA) nnzA(Dest::IsRowMajor ? itA.row() : itA.col())++;
161
162 VectorXi nnzB = VectorXi::Zero(Dest::IsRowMajor ? m_B.rows() : m_B.cols());
163 for (Index kB = 0; kB < m_B.outerSize(); ++kB)
164 for (RhsInnerIterator itB(rhs1, kB); itB; ++itB) nnzB(Dest::IsRowMajor ? itB.row() : itB.col())++;
165
166 Matrix<int, Dynamic, Dynamic, ColMajor> nnzAB = nnzB * nnzA.transpose();
167 dst.reserve(nnzAB.reshaped());
168 }
169
170 for (Index kA = 0; kA < m_A.outerSize(); ++kA) {
171 for (Index kB = 0; kB < m_B.outerSize(); ++kB) {
172 for (LhsInnerIterator itA(lhs1, kA); itA; ++itA) {
173 for (RhsInnerIterator itB(rhs1, kB); itB; ++itB) {
174 Index i = itA.row() * Br + itB.row(), j = itA.col() * Bc + itB.col();
175 dst.insert(i, j) = itA.value() * itB.value();
176 }
177 }
178 }
179 }
180}
181
182namespace internal {
183
184template <typename Lhs_, typename Rhs_>
185struct traits<KroneckerProduct<Lhs_, Rhs_> > {
186 typedef remove_all_t<Lhs_> Lhs;
187 typedef remove_all_t<Rhs_> Rhs;
189 typedef typename promote_index_type<typename Lhs::StorageIndex, typename Rhs::StorageIndex>::type StorageIndex;
190
191 enum {
192 Rows = size_at_compile_time(traits<Lhs>::RowsAtCompileTime, traits<Rhs>::RowsAtCompileTime),
193 Cols = size_at_compile_time(traits<Lhs>::ColsAtCompileTime, traits<Rhs>::ColsAtCompileTime),
194 MaxRows = size_at_compile_time(traits<Lhs>::MaxRowsAtCompileTime, traits<Rhs>::MaxRowsAtCompileTime),
195 MaxCols = size_at_compile_time(traits<Lhs>::MaxColsAtCompileTime, traits<Rhs>::MaxColsAtCompileTime)
196 };
197
198 typedef Matrix<Scalar, Rows, Cols> ReturnType;
199};
200
201template <typename Lhs_, typename Rhs_>
202struct traits<KroneckerProductSparse<Lhs_, Rhs_> > {
203 typedef MatrixXpr XprKind;
204 typedef remove_all_t<Lhs_> Lhs;
205 typedef remove_all_t<Rhs_> Rhs;
206 typedef typename ScalarBinaryOpTraits<typename Lhs::Scalar, typename Rhs::Scalar>::ReturnType Scalar;
207 typedef typename cwise_promote_storage_type<typename traits<Lhs>::StorageKind, typename traits<Rhs>::StorageKind,
208 scalar_product_op<typename Lhs::Scalar, typename Rhs::Scalar> >::ret
209 StorageKind;
210 typedef typename promote_index_type<typename Lhs::StorageIndex, typename Rhs::StorageIndex>::type StorageIndex;
211
212 enum {
213 LhsFlags = Lhs::Flags,
214 RhsFlags = Rhs::Flags,
215
216 RowsAtCompileTime = size_at_compile_time(traits<Lhs>::RowsAtCompileTime, traits<Rhs>::RowsAtCompileTime),
217 ColsAtCompileTime = size_at_compile_time(traits<Lhs>::ColsAtCompileTime, traits<Rhs>::ColsAtCompileTime),
218 MaxRowsAtCompileTime = size_at_compile_time(traits<Lhs>::MaxRowsAtCompileTime, traits<Rhs>::MaxRowsAtCompileTime),
219 MaxColsAtCompileTime = size_at_compile_time(traits<Lhs>::MaxColsAtCompileTime, traits<Rhs>::MaxColsAtCompileTime),
220
221 EvalToRowMajor = (int(LhsFlags) & int(RhsFlags) & RowMajorBit),
222 RemovedBits = ~(EvalToRowMajor ? 0 : RowMajorBit),
223
224 Flags = ((int(LhsFlags) | int(RhsFlags)) & HereditaryBits & RemovedBits) | EvalBeforeNestingBit,
225 CoeffReadCost = HugeCost
226 };
227
228 typedef SparseMatrix<Scalar, 0, StorageIndex> ReturnType;
229};
230
231} // end namespace internal
232
252template <typename A, typename B>
254 return KroneckerProduct<A, B>(a.derived(), b.derived());
255}
256
278template <typename A, typename B>
282
283} // end namespace Eigen
284
285#endif // KRONECKER_TENSOR_PRODUCT_H
Scalar coeff(Index row, Index col) const
Definition KroneckerTensorProduct.h:49
KroneckerProductBase(const Lhs &A, const Rhs &B)
Constructor.
Definition KroneckerTensorProduct.h:40
Scalar coeff(Index i) const
Definition KroneckerTensorProduct.h:57
Kronecker tensor product helper class for sparse matrices.
Definition KroneckerTensorProduct.h:111
void evalTo(Dest &dst) const
Evaluate the Kronecker tensor product.
Definition KroneckerTensorProduct.h:138
KroneckerProductSparse(const Lhs &A, const Rhs &B)
Constructor.
Definition KroneckerTensorProduct.h:119
Kronecker tensor product helper class for dense matrices.
Definition KroneckerTensorProduct.h:80
KroneckerProduct(const Lhs &A, const Rhs &B)
Constructor.
Definition KroneckerTensorProduct.h:88
void evalTo(Dest &dst) const
Evaluate the Kronecker tensor product.
Definition KroneckerTensorProduct.h:128
KroneckerProduct< A, B > kroneckerProduct(const MatrixBase< A > &a, const MatrixBase< B > &b)
Definition KroneckerTensorProduct.h:253
constexpr unsigned int RowMajorBit
Matrix< int, Dynamic, 1 > VectorXi
Namespace containing all symbols from the Eigen library.
constexpr Derived & derived()