11#ifndef EIGEN_CONCAT_OP_H
12#define EIGEN_CONCAT_OP_H
15#include "./InternalHeaderCheck.h"
21template <
int Direction,
typename LhsType,
typename RhsType>
22struct traits<Concat<Direction, LhsType, RhsType>> : traits<LhsType> {
23 using Scalar =
typename LhsType::Scalar;
24 using StorageKind =
typename traits<LhsType>::StorageKind;
25 using XprKind =
typename traits<LhsType>::XprKind;
26 using LhsTypeNested =
typename ref_selector<LhsType>::type;
27 using RhsTypeNested =
typename ref_selector<RhsType>::type;
28 using LhsTypeNested_ = std::remove_reference_t<LhsTypeNested>;
29 using RhsTypeNested_ = std::remove_reference_t<RhsTypeNested>;
33 LhsRows = int(LhsType::RowsAtCompileTime),
34 RhsRows = int(RhsType::RowsAtCompileTime),
35 LhsCols = int(LhsType::ColsAtCompileTime),
36 RhsCols = int(RhsType::ColsAtCompileTime),
38 RowsAtCompileTime = Direction ==
Vertical
39 ? (LhsRows == Dynamic || RhsRows == Dynamic ? int(Dynamic) : LhsRows + RhsRows)
40 : size_prefer_fixed(LhsRows, RhsRows),
42 ? (LhsCols == Dynamic || RhsCols == Dynamic ? int(Dynamic) : LhsCols + RhsCols)
43 : size_prefer_fixed(LhsCols, RhsCols),
45 LhsMaxRows = int(LhsType::MaxRowsAtCompileTime),
46 RhsMaxRows = int(RhsType::MaxRowsAtCompileTime),
47 LhsMaxCols = int(LhsType::MaxColsAtCompileTime),
48 RhsMaxCols = int(RhsType::MaxColsAtCompileTime),
50 MaxRowsAtCompileTime =
52 ? (LhsMaxRows == Dynamic || RhsMaxRows == Dynamic ? int(Dynamic) : LhsMaxRows + RhsMaxRows)
53 : max_size_prefer_dynamic(LhsMaxRows, RhsMaxRows),
54 MaxColsAtCompileTime =
56 ? (LhsMaxCols == Dynamic || RhsMaxCols == Dynamic ? int(Dynamic) : LhsMaxCols + RhsMaxCols)
57 : max_size_prefer_dynamic(LhsMaxCols, RhsMaxCols),
59 IsRowMajor = MaxRowsAtCompileTime == 1 && MaxColsAtCompileTime != 1 ? 1
60 : MaxColsAtCompileTime == 1 && MaxRowsAtCompileTime != 1 ? 0
86template <
int Direction,
typename LhsType,
typename RhsType>
87class Concat :
public internal::dense_xpr_base<Concat<Direction, LhsType, RhsType>>::type {
88 using LhsTypeNested =
typename internal::traits<Concat>::LhsTypeNested;
89 using RhsTypeNested =
typename internal::traits<Concat>::RhsTypeNested;
90 using LhsTypeNested_ =
typename internal::traits<Concat>::LhsTypeNested_;
91 using RhsTypeNested_ =
typename internal::traits<Concat>::RhsTypeNested_;
94 using Base =
typename internal::dense_xpr_base<Concat>::type;
95 EIGEN_DENSE_PUBLIC_INTERFACE(Concat)
96 using LhsNestedExpression = internal::remove_all_t<LhsType>;
97 using RhsNestedExpression = internal::remove_all_t<RhsType>;
99 template <
typename OriginalLhsType,
typename OriginalRhsType>
100 EIGEN_DEVICE_FUNC
constexpr inline Concat(
const OriginalLhsType& lhs,
const OriginalRhsType& rhs)
101 : m_lhs(lhs), m_rhs(rhs) {
102 EIGEN_STATIC_ASSERT((std::is_same<std::remove_const_t<LhsType>, OriginalLhsType>::value),
103 THE_MATRIX_OR_EXPRESSION_THAT_YOU_PASSED_DOES_NOT_HAVE_THE_EXPECTED_TYPE)
104 EIGEN_STATIC_ASSERT((std::is_same<std::remove_const_t<RhsType>, OriginalRhsType>::value),
105 THE_MATRIX_OR_EXPRESSION_THAT_YOU_PASSED_DOES_NOT_HAVE_THE_EXPECTED_TYPE)
107 (std::is_same<typename LhsType::Scalar, typename RhsType::Scalar>::value),
108 YOU_MIXED_DIFFERENT_NUMERIC_TYPES__YOU_NEED_TO_USE_THE_CAST_METHOD_OF_MATRIXBASE_TO_CAST_NUMERIC_TYPES_EXPLICITLY)
109 EIGEN_STATIC_ASSERT_SAME_XPR_KIND(LhsType, RhsType)
110 EIGEN_STATIC_ASSERT(Direction !=
Horizontal ||
int(LhsType::RowsAtCompileTime) == Dynamic ||
111 int(RhsType::RowsAtCompileTime) == Dynamic ||
112 int(LhsType::RowsAtCompileTime) ==
int(RhsType::RowsAtCompileTime),
113 YOU_MIXED_MATRICES_OF_DIFFERENT_SIZES)
114 EIGEN_STATIC_ASSERT(Direction !=
Vertical ||
int(LhsType::ColsAtCompileTime) == Dynamic ||
115 int(RhsType::ColsAtCompileTime) == Dynamic ||
116 int(LhsType::ColsAtCompileTime) ==
int(RhsType::ColsAtCompileTime),
117 YOU_MIXED_MATRICES_OF_DIFFERENT_SIZES)
118 EIGEN_IF_CONSTEXPR (Direction ==
Vertical) {
119 eigen_assert(lhs.cols() == rhs.cols() &&
"vcat: number of columns must match");
121 eigen_assert(lhs.rows() == rhs.rows() &&
"hcat: number of rows must match");
125 EIGEN_DEVICE_FUNC
constexpr Index rows()
const {
126 return Direction ==
Vertical ? m_lhs.rows() + m_rhs.rows() : m_lhs.rows();
128 EIGEN_DEVICE_FUNC
constexpr Index cols()
const {
129 return Direction ==
Horizontal ? m_lhs.cols() + m_rhs.cols() : m_lhs.cols();
132 EIGEN_DEVICE_FUNC
constexpr const LhsTypeNested_& lhs()
const {
return m_lhs; }
133 EIGEN_DEVICE_FUNC
constexpr const RhsTypeNested_& rhs()
const {
return m_rhs; }
143template <
int Direction,
typename LhsType,
typename RhsType>
144struct evaluator<
Concat<Direction, LhsType, RhsType>> : evaluator_base<Concat<Direction, LhsType, RhsType>> {
146 using CoeffReturnType =
typename XprType::CoeffReturnType;
148 using LhsNested =
typename nested_eval<LhsType, 1>::type;
149 using RhsNested =
typename nested_eval<RhsType, 1>::type;
150 using LhsNestedCleaned = remove_all_t<LhsNested>;
151 using RhsNestedCleaned = remove_all_t<RhsNested>;
154 CoeffReadCost = plain_enum_max(evaluator<LhsNestedCleaned>::CoeffReadCost,
155 evaluator<RhsNestedCleaned>::CoeffReadCost) +
156 NumTraits<typename XprType::Scalar>::AddCost,
157 LhsFlags = evaluator<LhsNestedCleaned>::Flags,
158 RhsFlags = evaluator<RhsNestedCleaned>::Flags,
159 IsRowMajor = int(traits<XprType>::Flags) &
RowMajorBit,
160 IsVectorAtCompileTime = XprType::IsVectorAtCompileTime,
165 LhsOrderAgrees = (bool(
int(LhsFlags) &
RowMajorBit) == bool(IsRowMajor)) || LhsNestedCleaned::IsVectorAtCompileTime,
166 RhsOrderAgrees = (
bool(
int(RhsFlags) &
RowMajorBit) ==
bool(IsRowMajor)) || RhsNestedCleaned::IsVectorAtCompileTime,
167 BothHavePacketAccess = (int(LhsFlags) & int(RhsFlags) &
PacketAccessBit) && LhsOrderAgrees && RhsOrderAgrees,
168 BothHaveLinearAccess =
bool(
int(LhsFlags) &
int(RhsFlags) &
LinearAccessBit),
174 EIGEN_DEVICE_FUNC
constexpr EIGEN_STRONG_INLINE
explicit evaluator(
const XprType& xpr)
179 m_lhsRows(xpr.lhs().rows()),
180 m_lhsCols(xpr.lhs().cols()) {}
182 EIGEN_DEVICE_FUNC
constexpr EIGEN_STRONG_INLINE CoeffReturnType coeff(Index row, Index col)
const {
183 EIGEN_IF_CONSTEXPR (Direction == Vertical) {
184 if (row < m_lhsRows.value())
185 return m_lhsImpl.coeff(row, col);
187 return m_rhsImpl.coeff(row - m_lhsRows.value(), col);
189 if (col < m_lhsCols.value())
190 return m_lhsImpl.coeff(row, col);
192 return m_rhsImpl.coeff(row, col - m_lhsCols.value());
196 EIGEN_DEVICE_FUNC
constexpr EIGEN_STRONG_INLINE CoeffReturnType coeff(Index index)
const {
197 const Index boundary = Direction ==
Vertical ? m_lhsRows.value() : m_lhsCols.value();
198 if (index < boundary)
199 return m_lhsImpl.coeff(index);
201 return m_rhsImpl.coeff(index - boundary);
204 template <
int LoadMode,
typename PacketType>
205 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketType packet(Index row, Index col)
const {
206 constexpr int packetSize = unpacket_traits<PacketType>::size;
207 EIGEN_IF_CONSTEXPR (Direction == Vertical) {
208 const Index boundary = m_lhsRows.value();
209 if (row >= boundary)
return m_rhsImpl.template packet<LoadMode, PacketType>(row - boundary, col);
212 EIGEN_IF_CONSTEXPR (!IsRowMajor) {
213 if (row + packetSize > boundary)
return packetBoundary<LoadMode, PacketType>(row, col);
215 return m_lhsImpl.template packet<LoadMode, PacketType>(row, col);
217 const Index boundary = m_lhsCols.value();
218 if (col >= boundary)
return m_rhsImpl.template packet<LoadMode, PacketType>(row, col - boundary);
221 EIGEN_IF_CONSTEXPR (IsRowMajor) {
222 if (col + packetSize > boundary)
return packetBoundary<LoadMode, PacketType>(row, col);
224 return m_lhsImpl.template packet<LoadMode, PacketType>(row, col);
228 template <
int LoadMode,
typename PacketType>
229 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketType packet(Index index)
const {
230 constexpr int packetSize = unpacket_traits<PacketType>::size;
231 const Index boundary = Direction ==
Vertical ? m_lhsRows.value() : m_lhsCols.value();
232 if (index >= boundary)
return m_rhsImpl.template packet<LoadMode, PacketType>(index - boundary);
233 if (index + packetSize > boundary)
return packetBoundaryLinear<LoadMode, PacketType>(index);
234 return m_lhsImpl.template packet<LoadMode, PacketType>(index);
237 template <
int LoadMode,
typename PacketType>
238 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketType packetSegment(Index row, Index col, Index begin, Index count)
const {
239 EIGEN_IF_CONSTEXPR (Direction == Vertical) {
240 const Index boundary = m_lhsRows.value();
242 return m_rhsImpl.template packetSegment<LoadMode, PacketType>(row - boundary, col, begin, count);
243 EIGEN_IF_CONSTEXPR (!IsRowMajor) {
244 if (row + begin + count > boundary)
return packetSegmentBoundary<LoadMode, PacketType>(row, col, begin, count);
246 return m_lhsImpl.template packetSegment<LoadMode, PacketType>(row, col, begin, count);
248 const Index boundary = m_lhsCols.value();
250 return m_rhsImpl.template packetSegment<LoadMode, PacketType>(row, col - boundary, begin, count);
251 EIGEN_IF_CONSTEXPR (IsRowMajor) {
252 if (col + begin + count > boundary)
return packetSegmentBoundary<LoadMode, PacketType>(row, col, begin, count);
254 return m_lhsImpl.template packetSegment<LoadMode, PacketType>(row, col, begin, count);
258 template <
int LoadMode,
typename PacketType>
259 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketType packetSegment(Index index, Index begin, Index count)
const {
260 const Index boundary = Direction ==
Vertical ? m_lhsRows.value() : m_lhsCols.value();
261 if (index >= boundary)
262 return m_rhsImpl.template packetSegment<LoadMode, PacketType>(index - boundary, begin, count);
263 if (index + begin + count > boundary)
return packetSegmentBoundaryLinear<LoadMode, PacketType>(index, begin, count);
264 return m_lhsImpl.template packetSegment<LoadMode, PacketType>(index, begin, count);
268 using Scalar =
typename XprType::Scalar;
270 template <
int LoadMode,
typename PacketType>
271 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketType packetBoundary(Index row, Index col)
const {
272 constexpr int packetSize = unpacket_traits<PacketType>::size;
273 EIGEN_ALIGN_TO_BOUNDARY(unpacket_traits<PacketType>::alignment) Scalar tmp[packetSize];
274 for (
int i = 0; i < packetSize; ++i)
275 tmp[i] = coeff(row + (Direction == Vertical ? i : 0), col + (Direction == Horizontal ? i : 0));
276 return pload<PacketType>(tmp);
279 template <
int LoadMode,
typename PacketType>
280 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketType packetBoundaryLinear(Index index)
const {
281 constexpr int packetSize = unpacket_traits<PacketType>::size;
282 EIGEN_ALIGN_TO_BOUNDARY(unpacket_traits<PacketType>::alignment) Scalar tmp[packetSize];
283 for (
int i = 0; i < packetSize; ++i) tmp[i] = coeff(index + i);
284 return pload<PacketType>(tmp);
287 template <
int LoadMode,
typename PacketType>
288 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketType packetSegmentBoundary(Index row, Index col, Index begin,
290 constexpr int packetSize = unpacket_traits<PacketType>::size;
291 EIGEN_ALIGN_MAX Scalar tmp[packetSize];
292 for (Index i = begin; i < begin + count; ++i)
293 tmp[i] = coeff(row + (Direction == Vertical ? i : 0), col + (Direction == Horizontal ? i : 0));
294 return ploadSegment<PacketType>(tmp, begin, count);
297 template <
int LoadMode,
typename PacketType>
298 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketType packetSegmentBoundaryLinear(Index index, Index begin,
300 constexpr int packetSize = unpacket_traits<PacketType>::size;
301 EIGEN_ALIGN_MAX Scalar tmp[packetSize];
302 for (Index i = begin; i < begin + count; ++i) tmp[i] = coeff(index + i);
303 return ploadSegment<PacketType>(tmp, begin, count);
308 evaluator<LhsNestedCleaned> m_lhsImpl;
309 evaluator<RhsNestedCleaned> m_rhsImpl;
310 const variable_if_dynamic<Index, LhsType::RowsAtCompileTime> m_lhsRows;
311 const variable_if_dynamic<Index, LhsType::ColsAtCompileTime> m_lhsCols;
330template <
typename Lhs,
typename Rhs>
332 return Concat<Horizontal, Lhs, Rhs>(lhs.derived(), rhs.derived());
349template <
typename Lhs,
typename Rhs>
351 return Concat<Vertical, Lhs, Rhs>(lhs.derived(), rhs.derived());
Expression of the concatenation of two dense expressions.
Definition ConcatOp.h:87
const Concat< Vertical, Lhs, Rhs > vcat(const DenseBase< Lhs > &lhs, const DenseBase< Rhs > &rhs)
Definition ConcatOp.h:350
const Concat< Horizontal, Lhs, Rhs > hcat(const DenseBase< Lhs > &lhs, const DenseBase< Rhs > &rhs)
Definition ConcatOp.h:331
Base class for all dense matrices, vectors, and arrays.
Definition DenseBase.h:45
@ Horizontal
Definition Constants.h:270
@ Vertical
Definition Constants.h:267
constexpr unsigned int PacketAccessBit
Definition Constants.h:98
constexpr unsigned int LinearAccessBit
Definition Constants.h:134
constexpr unsigned int RowMajorBit
Definition Constants.h:71