Eigen  5.0.1
 
Loading...
Searching...
No Matches
ConcatOp.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2026 Pavel Guzenfeld
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_CONCAT_OP_H
12#define EIGEN_CONCAT_OP_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20
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>;
30 enum {
31 // For vertical concat (stacking rows): rows add up, cols must match
32 // For horizontal concat (stacking cols): cols add up, rows must match
33 LhsRows = int(LhsType::RowsAtCompileTime),
34 RhsRows = int(RhsType::RowsAtCompileTime),
35 LhsCols = int(LhsType::ColsAtCompileTime),
36 RhsCols = int(RhsType::ColsAtCompileTime),
37
38 RowsAtCompileTime = Direction == Vertical
39 ? (LhsRows == Dynamic || RhsRows == Dynamic ? int(Dynamic) : LhsRows + RhsRows)
40 : size_prefer_fixed(LhsRows, RhsRows),
41 ColsAtCompileTime = Direction == Horizontal
42 ? (LhsCols == Dynamic || RhsCols == Dynamic ? int(Dynamic) : LhsCols + RhsCols)
43 : size_prefer_fixed(LhsCols, RhsCols),
44
45 LhsMaxRows = int(LhsType::MaxRowsAtCompileTime),
46 RhsMaxRows = int(RhsType::MaxRowsAtCompileTime),
47 LhsMaxCols = int(LhsType::MaxColsAtCompileTime),
48 RhsMaxCols = int(RhsType::MaxColsAtCompileTime),
49
50 MaxRowsAtCompileTime =
51 Direction == Vertical
52 ? (LhsMaxRows == Dynamic || RhsMaxRows == Dynamic ? int(Dynamic) : LhsMaxRows + RhsMaxRows)
53 : max_size_prefer_dynamic(LhsMaxRows, RhsMaxRows),
54 MaxColsAtCompileTime =
55 Direction == Horizontal
56 ? (LhsMaxCols == Dynamic || RhsMaxCols == Dynamic ? int(Dynamic) : LhsMaxCols + RhsMaxCols)
57 : max_size_prefer_dynamic(LhsMaxCols, RhsMaxCols),
58
59 IsRowMajor = MaxRowsAtCompileTime == 1 && MaxColsAtCompileTime != 1 ? 1
60 : MaxColsAtCompileTime == 1 && MaxRowsAtCompileTime != 1 ? 0
61 : (int(LhsType::Flags) & RowMajorBit) ? 1
62 : 0,
63 Flags = IsRowMajor ? RowMajorBit : 0
64 };
65};
66
67} // namespace internal
68
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_;
92
93 public:
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>;
98
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)
106 EIGEN_STATIC_ASSERT(
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");
120 } else {
121 eigen_assert(lhs.rows() == rhs.rows() && "hcat: number of rows must match");
122 }
123 }
124
125 EIGEN_DEVICE_FUNC constexpr Index rows() const {
126 return Direction == Vertical ? m_lhs.rows() + m_rhs.rows() : m_lhs.rows();
127 }
128 EIGEN_DEVICE_FUNC constexpr Index cols() const {
129 return Direction == Horizontal ? m_lhs.cols() + m_rhs.cols() : m_lhs.cols();
130 }
131
132 EIGEN_DEVICE_FUNC constexpr const LhsTypeNested_& lhs() const { return m_lhs; }
133 EIGEN_DEVICE_FUNC constexpr const RhsTypeNested_& rhs() const { return m_rhs; }
134
135 protected:
136 LhsTypeNested m_lhs;
137 RhsTypeNested m_rhs;
138};
139
140// Evaluator for Concat
141namespace internal {
142
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;
147
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>;
152
153 enum {
154 CoeffReadCost = plain_enum_max(evaluator<LhsNestedCleaned>::CoeffReadCost,
155 evaluator<RhsNestedCleaned>::CoeffReadCost) +
156 NumTraits<typename XprType::Scalar>::AddCost, // cost of the branch
157 LhsFlags = evaluator<LhsNestedCleaned>::Flags,
158 RhsFlags = evaluator<RhsNestedCleaned>::Flags,
159 IsRowMajor = int(traits<XprType>::Flags) & RowMajorBit,
160 IsVectorAtCompileTime = XprType::IsVectorAtCompileTime,
161 // Packets are loaded along the Concat's inner direction, so each operand must store its
162 // coefficients in that same order. Compile-time vectors are exempt: a vector evaluator with
163 // packet access is contiguous along its single extent, and requests along the other
164 // direction can only be single-lane.
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),
169 Flags = (traits<XprType>::Flags & RowMajorBit) | (BothHavePacketAccess ? PacketAccessBit : 0) |
170 (IsVectorAtCompileTime && BothHaveLinearAccess ? LinearAccessBit : 0),
171 Alignment = 0 // conservative: no alignment guarantees across boundary
172 };
173
174 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE explicit evaluator(const XprType& xpr)
175 : m_lhs(xpr.lhs()),
176 m_rhs(xpr.rhs()),
177 m_lhsImpl(m_lhs),
178 m_rhsImpl(m_rhs),
179 m_lhsRows(xpr.lhs().rows()),
180 m_lhsCols(xpr.lhs().cols()) {}
181
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);
186 else
187 return m_rhsImpl.coeff(row - m_lhsRows.value(), col);
188 } else {
189 if (col < m_lhsCols.value())
190 return m_lhsImpl.coeff(row, col);
191 else
192 return m_rhsImpl.coeff(row, col - m_lhsCols.value());
193 }
194 }
195
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);
200 else
201 return m_rhsImpl.coeff(index - boundary);
202 }
203
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);
210 // Column-major: inner=rows, packet extends along rows and may straddle the row boundary.
211 // Row-major: inner=cols, packet extends along cols — never crosses the row boundary.
212 EIGEN_IF_CONSTEXPR (!IsRowMajor) {
213 if (row + packetSize > boundary) return packetBoundary<LoadMode, PacketType>(row, col);
214 }
215 return m_lhsImpl.template packet<LoadMode, PacketType>(row, col);
216 } else {
217 const Index boundary = m_lhsCols.value();
218 if (col >= boundary) return m_rhsImpl.template packet<LoadMode, PacketType>(row, col - boundary);
219 // Row-major: inner=cols, packet extends along cols and may straddle the col boundary.
220 // Column-major: inner=rows, packet extends along rows — never crosses the col boundary.
221 EIGEN_IF_CONSTEXPR (IsRowMajor) {
222 if (col + packetSize > boundary) return packetBoundary<LoadMode, PacketType>(row, col);
223 }
224 return m_lhsImpl.template packet<LoadMode, PacketType>(row, col);
225 }
226 }
227
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);
235 }
236
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();
241 if (row >= boundary)
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);
245 }
246 return m_lhsImpl.template packetSegment<LoadMode, PacketType>(row, col, begin, count);
247 } else {
248 const Index boundary = m_lhsCols.value();
249 if (col >= boundary)
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);
253 }
254 return m_lhsImpl.template packetSegment<LoadMode, PacketType>(row, col, begin, count);
255 }
256 }
257
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);
265 }
266
267 protected:
268 using Scalar = typename XprType::Scalar;
269
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);
277 }
278
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);
285 }
286
287 template <int LoadMode, typename PacketType>
288 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketType packetSegmentBoundary(Index row, Index col, Index begin,
289 Index count) const {
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);
295 }
296
297 template <int LoadMode, typename PacketType>
298 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketType packetSegmentBoundaryLinear(Index index, Index begin,
299 Index count) const {
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);
304 }
305
306 LhsNested m_lhs;
307 RhsNested m_rhs;
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;
312};
313
314} // namespace internal
315
330template <typename Lhs, typename Rhs>
331EIGEN_DEVICE_FUNC inline const Concat<Horizontal, Lhs, Rhs> hcat(const DenseBase<Lhs>& lhs, const DenseBase<Rhs>& rhs) {
332 return Concat<Horizontal, Lhs, Rhs>(lhs.derived(), rhs.derived());
333}
334
349template <typename Lhs, typename Rhs>
350EIGEN_DEVICE_FUNC inline const Concat<Vertical, Lhs, Rhs> vcat(const DenseBase<Lhs>& lhs, const DenseBase<Rhs>& rhs) {
351 return Concat<Vertical, Lhs, Rhs>(lhs.derived(), rhs.derived());
352}
353
354} // end namespace Eigen
355
356#endif // EIGEN_CONCAT_OP_H
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