Optimize visitor traversal in case of RowMajor.

This commit is contained in:
Antonio Sánchez
2022-03-23 15:27:57 +00:00
parent f2a3e03e9b
commit 19a6a827c4
2 changed files with 171 additions and 7 deletions

View File

@@ -23,8 +23,10 @@ template<typename Visitor, typename Derived, int UnrollCount>
struct visitor_impl<Visitor, Derived, UnrollCount, false>
{
enum {
col = (UnrollCount-1) / Derived::RowsAtCompileTime,
row = (UnrollCount-1) % Derived::RowsAtCompileTime
col = Derived::IsRowMajor ? (UnrollCount-1) % Derived::ColsAtCompileTime
: (UnrollCount-1) / Derived::RowsAtCompileTime,
row = Derived::IsRowMajor ? (UnrollCount-1) / Derived::ColsAtCompileTime
: (UnrollCount-1) % Derived::RowsAtCompileTime
};
EIGEN_DEVICE_FUNC
@@ -60,11 +62,25 @@ struct visitor_impl<Visitor, Derived, Dynamic, /*Vectorize=*/false>
static inline void run(const Derived& mat, Visitor& visitor)
{
visitor.init(mat.coeff(0,0), 0, 0);
for(Index i = 1; i < mat.rows(); ++i)
visitor(mat.coeff(i, 0), i, 0);
for(Index j = 1; j < mat.cols(); ++j)
for(Index i = 0; i < mat.rows(); ++i)
visitor(mat.coeff(i, j), i, j);
if (Derived::IsRowMajor) {
for(Index i = 1; i < mat.cols(); ++i) {
visitor(mat.coeff(0, i), 0, i);
}
for(Index j = 1; j < mat.rows(); ++j) {
for(Index i = 0; i < mat.cols(); ++i) {
visitor(mat.coeff(j, i), j, i);
}
}
} else {
for(Index i = 1; i < mat.rows(); ++i) {
visitor(mat.coeff(i, 0), i, 0);
}
for(Index j = 1; j < mat.cols(); ++j) {
for(Index i = 0; i < mat.rows(); ++i) {
visitor(mat.coeff(i, j), i, j);
}
}
}
}
};
@@ -114,6 +130,7 @@ public:
PacketAccess = Evaluator::Flags & PacketAccessBit,
IsRowMajor = XprType::IsRowMajor,
RowsAtCompileTime = XprType::RowsAtCompileTime,
ColsAtCompileTime = XprType::ColsAtCompileTime,
CoeffReadCost = Evaluator::CoeffReadCost
};