13#ifndef KRONECKER_TENSOR_PRODUCT_H
14#define KRONECKER_TENSOR_PRODUCT_H
17#include "./InternalHeaderCheck.h"
28template <
typename Derived>
31 typedef typename internal::traits<Derived> Traits;
32 typedef typename Traits::Scalar Scalar;
35 typedef typename Traits::Lhs Lhs;
36 typedef typename Traits::Rhs Rhs;
42 inline Index rows()
const {
return m_A.rows() * m_B.rows(); }
43 inline Index cols()
const {
return m_A.cols() * m_B.cols(); }
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());
58 EIGEN_STATIC_ASSERT_VECTOR_ONLY(Derived);
59 return m_A.coeff(i / m_B.size()) * m_B.coeff(i % m_B.size());
63 typename Lhs::Nested m_A;
64 typename Rhs::Nested m_B;
79template <
typename Lhs,
typename Rhs>
91 template <
typename Dest>
92 void evalTo(Dest& dst)
const;
110template <
typename Lhs,
typename Rhs>
122 template <
typename Dest>
123 void evalTo(Dest& dst)
const;
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)
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);
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);
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())++;
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())++;
167 dst.reserve(nnzAB.reshaped());
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();
184template <
typename Lhs_,
typename 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;
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)
198 typedef Matrix<Scalar, Rows, Cols> ReturnType;
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
210 typedef typename promote_index_type<typename Lhs::StorageIndex, typename Rhs::StorageIndex>::type StorageIndex;
213 LhsFlags = Lhs::Flags,
214 RhsFlags = Rhs::Flags,
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),
221 EvalToRowMajor = (int(LhsFlags) & int(RhsFlags) &
RowMajorBit),
222 RemovedBits = ~(EvalToRowMajor ? 0 : RowMajorBit),
224 Flags = ((int(LhsFlags) | int(RhsFlags)) & HereditaryBits & RemovedBits) | EvalBeforeNestingBit,
225 CoeffReadCost = HugeCost
228 typedef SparseMatrix<Scalar, 0, StorageIndex> ReturnType;
252template <
typename A,
typename B>
278template <
typename A,
typename B>
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()