Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
TensorConcatenation.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2014 Benoit Steiner <benoit.steiner.goog@gmail.com>
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_TENSOR_TENSOR_CONCATENATION_H
12#define EIGEN_TENSOR_TENSOR_CONCATENATION_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20template <typename Axis, typename LhsXprType, typename RhsXprType>
21struct traits<TensorConcatenationOp<Axis, LhsXprType, RhsXprType>> {
22 // Type promotion to handle the case where the types of the lhs and the rhs are different.
23 typedef typename promote_storage_type<typename LhsXprType::Scalar, typename RhsXprType::Scalar>::ret Scalar;
24 typedef typename promote_storage_type<typename traits<LhsXprType>::StorageKind,
25 typename traits<RhsXprType>::StorageKind>::ret StorageKind;
26 typedef
27 typename promote_index_type<typename traits<LhsXprType>::Index, typename traits<RhsXprType>::Index>::type Index;
28 static constexpr int NumDimensions = traits<LhsXprType>::NumDimensions;
29 static constexpr int Layout = traits<LhsXprType>::Layout;
30 enum { Flags = 0 };
31 typedef std::conditional_t<Pointer_type_promotion<typename LhsXprType::Scalar, Scalar>::val,
32 typename traits<LhsXprType>::PointerType, typename traits<RhsXprType>::PointerType>
33 PointerType;
34};
35
36template <typename Axis, typename LhsXprType, typename RhsXprType>
37struct eval<TensorConcatenationOp<Axis, LhsXprType, RhsXprType>, Eigen::Dense> {
38 typedef const TensorConcatenationOp<Axis, LhsXprType, RhsXprType>& type;
39};
40
41} // end namespace internal
42
48template <typename Axis, typename LhsXprType, typename RhsXprType>
49class TensorConcatenationOp : public TensorBase<TensorConcatenationOp<Axis, LhsXprType, RhsXprType>, WriteAccessors> {
50 public:
52 typedef typename internal::traits<TensorConcatenationOp>::Scalar Scalar;
53 typedef typename internal::traits<TensorConcatenationOp>::StorageKind StorageKind;
54 typedef typename internal::traits<TensorConcatenationOp>::Index Index;
55 typedef typename internal::ref_selector<TensorConcatenationOp>::type Nested;
56 typedef typename internal::promote_storage_type<typename LhsXprType::CoeffReturnType,
57 typename RhsXprType::CoeffReturnType>::ret CoeffReturnType;
58 typedef typename NumTraits<Scalar>::Real RealScalar;
59
60 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorConcatenationOp(const LhsXprType& lhs, const RhsXprType& rhs, Axis axis)
61 : m_lhs_xpr(lhs), m_rhs_xpr(rhs), m_axis(axis) {}
62
63 EIGEN_DEVICE_FUNC const internal::remove_all_t<typename LhsXprType::Nested>& lhsExpression() const {
64 return m_lhs_xpr;
65 }
66
67 EIGEN_DEVICE_FUNC const internal::remove_all_t<typename RhsXprType::Nested>& rhsExpression() const {
68 return m_rhs_xpr;
69 }
70
71 EIGEN_DEVICE_FUNC const Axis& axis() const { return m_axis; }
72
73 EIGEN_INHERIT_ASSIGNMENT_OPERATORS(TensorConcatenationOp)
74 protected:
75 typename LhsXprType::Nested m_lhs_xpr;
76 typename RhsXprType::Nested m_rhs_xpr;
77 const Axis m_axis;
78};
79
80// Eval as rvalue
81template <typename Axis, typename LeftArgType, typename RightArgType, typename Device>
82struct TensorEvaluator<const TensorConcatenationOp<Axis, LeftArgType, RightArgType>, Device> {
84 typedef typename XprType::Index Index;
85 static constexpr int NumDims = internal::array_size<typename TensorEvaluator<LeftArgType, Device>::Dimensions>::value;
86 static constexpr int RightNumDims =
87 internal::array_size<typename TensorEvaluator<RightArgType, Device>::Dimensions>::value;
88 typedef DSizes<Index, NumDims> Dimensions;
89 typedef typename XprType::Scalar Scalar;
90 typedef typename XprType::CoeffReturnType CoeffReturnType;
91 typedef typename PacketType<CoeffReturnType, Device>::type PacketReturnType;
92 typedef StorageMemory<CoeffReturnType, Device> Storage;
93 typedef typename Storage::Type EvaluatorPointerType;
94 static constexpr int Layout = TensorEvaluator<LeftArgType, Device>::Layout;
95 enum {
96 IsAligned = false,
97 PacketAccess =
98 TensorEvaluator<LeftArgType, Device>::PacketAccess && TensorEvaluator<RightArgType, Device>::PacketAccess,
99 // block() reads each operand's data() pointer directly, so both must
100 // expose raw storage. Scalar-changing block consumers
101 // (TensorCwiseUnaryOp, TensorConversionOp) drop the forwarded destination
102 // buffer before reaching us, so prepareStorage always sees either a
103 // matching-Scalar buffer or no buffer at all.
105 // Matches TensorShuffling / TensorBroadcasting / TensorPadding: bulk copy
106 // wins over the per-element coeff/packet path (which pays div/mod per
107 // access) at every size we benchmarked, so always prefer block.
108 PreferBlockAccess = true,
109 RawAccess = false
110 };
111
112 typedef std::remove_const_t<Scalar> ScalarNoConst;
113
114 //===- Tensor block evaluation strategy (see TensorBlock.h) -------------===//
115 typedef internal::TensorBlockDescriptor<NumDims, Index> TensorBlockDesc;
116 typedef internal::TensorBlockScratchAllocator<Device> TensorBlockScratch;
117
118 typedef typename internal::TensorMaterializedBlock<ScalarNoConst, NumDims, Layout, Index> TensorBlock;
119 //===--------------------------------------------------------------------===//
120
121 EIGEN_STRONG_INLINE TensorEvaluator(const XprType& op, const Device& device)
122 : m_leftImpl(op.lhsExpression(), device),
123 m_rightImpl(op.rhsExpression(), device),
124 m_device(device),
125 m_axis(op.axis()) {
126 EIGEN_STATIC_ASSERT((static_cast<int>(TensorEvaluator<LeftArgType, Device>::Layout) ==
127 static_cast<int>(TensorEvaluator<RightArgType, Device>::Layout) ||
128 NumDims == 1),
129 YOU_MADE_A_PROGRAMMING_MISTAKE);
130 // TensorConcatenationOp requires both operands to have the same static
131 // rank. Reshape the lower-rank operand explicitly if you need to mix
132 // ranks; see the concatenate() entry in contrib/Eigen/src/Tensor/README.md.
133 EIGEN_STATIC_ASSERT((NumDims == RightNumDims), YOU_MADE_A_PROGRAMMING_MISTAKE);
134 EIGEN_STATIC_ASSERT((NumDims > 0), YOU_MADE_A_PROGRAMMING_MISTAKE);
135
136 eigen_assert(0 <= m_axis && m_axis < NumDims);
137 m_leftAxisSize = m_leftImpl.dimensions()[m_axis];
138 const Dimensions& lhs_dims = m_leftImpl.dimensions();
139 const Dimensions& rhs_dims = m_rightImpl.dimensions();
140 {
141 int i = 0;
142 for (; i < m_axis; ++i) {
143 eigen_assert(lhs_dims[i] > 0);
144 eigen_assert(lhs_dims[i] == rhs_dims[i]);
145 m_dimensions[i] = lhs_dims[i];
146 }
147 eigen_assert(lhs_dims[i] > 0); // Now i == m_axis.
148 eigen_assert(rhs_dims[i] > 0);
149 m_dimensions[i] = lhs_dims[i] + rhs_dims[i];
150 for (++i; i < NumDims; ++i) {
151 eigen_assert(lhs_dims[i] > 0);
152 eigen_assert(lhs_dims[i] == rhs_dims[i]);
153 m_dimensions[i] = lhs_dims[i];
154 }
155 }
156
157 EIGEN_IF_CONSTEXPR (static_cast<int>(Layout) == static_cast<int>(ColMajor)) {
158 m_leftStrides[0] = 1;
159 m_rightStrides[0] = 1;
160 m_outputStrides[0] = 1;
161
162 for (int j = 1; j < NumDims; ++j) {
163 m_leftStrides[j] = m_leftStrides[j - 1] * lhs_dims[j - 1];
164 m_rightStrides[j] = m_rightStrides[j - 1] * rhs_dims[j - 1];
165 m_outputStrides[j] = m_outputStrides[j - 1] * m_dimensions[j - 1];
166 }
167 } else {
168 m_leftStrides[NumDims - 1] = 1;
169 m_rightStrides[NumDims - 1] = 1;
170 m_outputStrides[NumDims - 1] = 1;
171
172 for (int j = NumDims - 2; j >= 0; --j) {
173 m_leftStrides[j] = m_leftStrides[j + 1] * lhs_dims[j + 1];
174 m_rightStrides[j] = m_rightStrides[j + 1] * rhs_dims[j + 1];
175 m_outputStrides[j] = m_outputStrides[j + 1] * m_dimensions[j + 1];
176 }
177 }
178 }
179
180 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Dimensions& dimensions() const { return m_dimensions; }
181
182 // TODO(phli): Add short-circuit memcpy evaluation if underlying data are linear.
183 EIGEN_STRONG_INLINE bool evalSubExprsIfNeeded(EvaluatorPointerType) {
184 m_leftImpl.evalSubExprsIfNeeded(nullptr);
185 m_rightImpl.evalSubExprsIfNeeded(nullptr);
186 return true;
187 }
188
189 EIGEN_STRONG_INLINE void cleanup() {
190 m_leftImpl.cleanup();
191 m_rightImpl.cleanup();
192 }
193
194 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE internal::TensorBlockResourceRequirements getResourceRequirements() const {
195 // Target the L1 cache. A block straddling the concat axis is materialized
196 // into a merged scratch slab that the cwise consumer reads straight back;
197 // sizing the block to L1 keeps that round-trip out of the last-level
198 // cache. It also splits an otherwise cache-resident output into many
199 // blocks, so the off-axis ones qualify for the zero-copy view path in
200 // block(). Matches TensorBroadcasting, which likewise targets L1 for its
201 // materialized block path.
202 const size_t target_size = m_device.firstLevelCacheSize();
203 return internal::TensorBlockResourceRequirements::merge(
204 internal::TensorBlockResourceRequirements::skewed<Scalar>(target_size),
205 internal::TensorBlockResourceRequirements::merge(m_leftImpl.getResourceRequirements(),
206 m_rightImpl.getResourceRequirements()));
207 }
208
209 // True when a block of shape `block_dims` is a contiguous run of an operand
210 // whose shape is `operand_dims` -- i.e. it can be addressed as a plain
211 // pointer offset with no per-row stride walk. Mirrors the direct-access test
212 // in TensorMaterializedBlock::materialize().
213 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool isContiguousOperandSlab(const Dimensions& operand_dims,
214 const Dimensions& block_dims) const {
215 static constexpr bool IsColMajor = Layout == static_cast<int>(ColMajor);
216 int matching_inner_dims = 0;
217 for (int i = 0; i < NumDims; ++i) {
218 const int dim = IsColMajor ? i : NumDims - i - 1;
219 if (operand_dims[dim] != block_dims[dim]) break;
220 ++matching_inner_dims;
221 }
222 // Every dimension above the single partial dimension must be of size 1.
223 for (int i = matching_inner_dims + 1; i < NumDims; ++i) {
224 const int dim = IsColMajor ? i : NumDims - i - 1;
225 if (block_dims[dim] != 1) return false;
226 }
227 return true;
228 }
229
230 // Returns a zero-copy view when the block lies within a single operand and
231 // is contiguous there; otherwise materializes the slab(s) with
232 // TensorBlockIO::Copy, which collapses to a memcpy per contiguous slab.
233 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorBlock block(TensorBlockDesc& desc, TensorBlockScratch& scratch,
234 bool root_of_expr_ast = false) const {
235 static constexpr bool IsColMajor = Layout == static_cast<int>(ColMajor);
236
237 if (desc.size() == 0) {
238 return TensorBlock(internal::TensorBlockKind::kView, nullptr, desc.dimensions());
239 }
240
241 Index remaining = desc.offset();
242 DSizes<Index, NumDims> out_coords;
243 EIGEN_IF_CONSTEXPR (IsColMajor) {
244 for (int i = NumDims - 1; i > 0; --i) {
245 out_coords[i] = remaining / m_outputStrides[i];
246 remaining -= out_coords[i] * m_outputStrides[i];
247 }
248 out_coords[0] = remaining;
249 } else {
250 for (int i = 0; i < NumDims - 1; ++i) {
251 out_coords[i] = remaining / m_outputStrides[i];
252 remaining -= out_coords[i] * m_outputStrides[i];
253 }
254 out_coords[NumDims - 1] = remaining;
255 }
256
257 const Index axis_start = out_coords[m_axis];
258 const Index axis_size = desc.dimension(static_cast<int>(m_axis));
259 const Index axis_end = axis_start + axis_size;
260
261 // Fast path: a block that lies entirely within one operand and forms a
262 // contiguous run of its storage needs no copy at all -- return a view
263 // straight into A / B. This is what keeps a cwise expression that reads
264 // the concatenation (e.g. `(A.concatenate(B, axis) + C)`) streaming: a
265 // cwise consumer drops our destination buffer before calling block() (see
266 // TensorCwiseBinaryOp::block), so without the view the slab would be
267 // bounced through a scratch buffer and read straight back, doubling cache
268 // traffic. Straddling blocks still need the merged buffer materialized
269 // below.
270 if (axis_end <= m_leftAxisSize) {
271 if (isContiguousOperandSlab(m_leftImpl.dimensions(), desc.dimensions())) {
272 Index left_src_offset = 0;
273 for (int i = 0; i < NumDims; ++i) {
274 left_src_offset += out_coords[i] * m_leftStrides[i];
275 }
276 return TensorBlock(internal::TensorBlockKind::kView, m_leftImpl.data() + left_src_offset, desc.dimensions());
277 }
278 } else if (axis_start >= m_leftAxisSize) {
279 if (isContiguousOperandSlab(m_rightImpl.dimensions(), desc.dimensions())) {
280 Index right_src_offset = (axis_start - m_leftAxisSize) * m_rightStrides[m_axis];
281 for (int i = 0; i < NumDims; ++i) {
282 if (i != m_axis) {
283 right_src_offset += out_coords[i] * m_rightStrides[i];
284 }
285 }
286 return TensorBlock(internal::TensorBlockKind::kView, m_rightImpl.data() + right_src_offset, desc.dimensions());
287 }
288 }
289
290 typedef internal::TensorBlockIO<ScalarNoConst, Index, NumDims, Layout> TensorBlockIO;
291 typedef typename TensorBlockIO::Dst TensorBlockIODst;
292 typedef typename TensorBlockIO::Src TensorBlockIOSrc;
293
294 // Strided destination buffers are safe here because we only ever write
295 // dense slabs into them; allowing strided storage lets a root-of-AST
296 // assignment materialize directly into the output tensor.
297 typename TensorBlock::Storage block_storage =
298 TensorBlock::prepareStorage(desc, scratch, /*allow_strided_storage=*/root_of_expr_ast);
299
300 if (axis_start < m_leftAxisSize) {
301 const Index left_rows_in_block = numext::mini(m_leftAxisSize, axis_end) - axis_start;
302 DSizes<Index, NumDims> left_sub_dims = desc.dimensions();
303 left_sub_dims[m_axis] = left_rows_in_block;
304
305 Index left_src_offset = 0;
306 for (int i = 0; i < NumDims; ++i) {
307 left_src_offset += out_coords[i] * m_leftStrides[i];
308 }
309
310 typename TensorBlockIO::Dimensions left_strides(m_leftStrides);
311 TensorBlockIOSrc src(left_strides, m_leftImpl.data(), left_src_offset);
312 TensorBlockIODst dst(left_sub_dims, block_storage.strides(), block_storage.data(),
313 /*dst_offset=*/0);
314 TensorBlockIO::Copy(dst, src);
315 }
316
317 if (axis_end > m_leftAxisSize) {
318 const Index right_rows_in_block = axis_end - numext::maxi(m_leftAxisSize, axis_start);
319 DSizes<Index, NumDims> right_sub_dims = desc.dimensions();
320 right_sub_dims[m_axis] = right_rows_in_block;
321
322 // Right operand has the same per-dim coords as out_coords except along
323 // the concat axis, where it starts at max(0, axis_start - left_size).
324 const Index right_axis_start = numext::maxi(Index(0), axis_start - m_leftAxisSize);
325 Index right_src_offset = right_axis_start * m_rightStrides[m_axis];
326 for (int i = 0; i < NumDims; ++i) {
327 if (i != m_axis) {
328 right_src_offset += out_coords[i] * m_rightStrides[i];
329 }
330 }
331
332 // Offset within the materialized buffer where the right-side slab starts.
333 const Index dst_axis_offset = numext::maxi(Index(0), m_leftAxisSize - axis_start);
334 const Index dst_offset = dst_axis_offset * block_storage.strides()[m_axis];
335
336 typename TensorBlockIO::Dimensions right_strides(m_rightStrides);
337 TensorBlockIOSrc src(right_strides, m_rightImpl.data(), right_src_offset);
338 TensorBlockIODst dst(right_sub_dims, block_storage.strides(), block_storage.data(), dst_offset);
339 TensorBlockIO::Copy(dst, src);
340 }
341
342 return block_storage.AsTensorMaterializedBlock();
343 }
344
345 // The per-dim integer div/mod in this loop are the obvious "slow" candidates,
346 // but the same TensorIntDivisor substitution was measured net-negative on
347 // Intel Raptor Lake when prototyped for TensorBroadcasting (see the comment
348 // on TensorBroadcasting::indexColMajor). Modern x86 hardware div (~20
349 // cycles) amortizes well and the mul+shifts don't win back the added state.
350 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffReturnType coeff(Index index) const {
351 // Collect dimension-wise indices (subs).
352 array<Index, NumDims> subs;
353 EIGEN_IF_CONSTEXPR (static_cast<int>(Layout) == static_cast<int>(ColMajor)) {
354 for (int i = NumDims - 1; i > 0; --i) {
355 subs[i] = index / m_outputStrides[i];
356 index -= subs[i] * m_outputStrides[i];
357 }
358 subs[0] = index;
359 } else {
360 for (int i = 0; i < NumDims - 1; ++i) {
361 subs[i] = index / m_outputStrides[i];
362 index -= subs[i] * m_outputStrides[i];
363 }
364 subs[NumDims - 1] = index;
365 }
366
367 const Dimensions& left_dims = m_leftImpl.dimensions();
368 if (subs[m_axis] < left_dims[m_axis]) {
369 Index left_index;
370 EIGEN_IF_CONSTEXPR (static_cast<int>(Layout) == static_cast<int>(ColMajor)) {
371 left_index = subs[0];
372 EIGEN_UNROLL_LOOP
373 for (int i = 1; i < NumDims; ++i) {
374 left_index += (subs[i] % left_dims[i]) * m_leftStrides[i];
375 }
376 } else {
377 left_index = subs[NumDims - 1];
378 EIGEN_UNROLL_LOOP
379 for (int i = NumDims - 2; i >= 0; --i) {
380 left_index += (subs[i] % left_dims[i]) * m_leftStrides[i];
381 }
382 }
383 return m_leftImpl.coeff(left_index);
384 } else {
385 subs[m_axis] -= left_dims[m_axis];
386 const Dimensions& right_dims = m_rightImpl.dimensions();
387 Index right_index;
388 EIGEN_IF_CONSTEXPR (static_cast<int>(Layout) == static_cast<int>(ColMajor)) {
389 right_index = subs[0];
390 EIGEN_UNROLL_LOOP
391 for (int i = 1; i < NumDims; ++i) {
392 right_index += (subs[i] % right_dims[i]) * m_rightStrides[i];
393 }
394 } else {
395 right_index = subs[NumDims - 1];
396 EIGEN_UNROLL_LOOP
397 for (int i = NumDims - 2; i >= 0; --i) {
398 right_index += (subs[i] % right_dims[i]) * m_rightStrides[i];
399 }
400 }
401 return m_rightImpl.coeff(right_index);
402 }
403 }
404
405 // When the packet sits entirely on one side of the concat boundary, delegate
406 // to that operand's packet<>() rather than assembling PacketSize coeff()
407 // calls. The packet stays on one side iff only the innermost dim varies
408 // across the packet -- i.e. all other subs match between the first and last
409 // index. When that holds, subs[m_axis] is either constant (m_axis is not
410 // innermost) or monotonic non-decreasing (m_axis is innermost), so checking
411 // just the endpoints decides the side. Otherwise subs[m_axis] can wrap back
412 // through the boundary mid-packet (as when the inner dim has fewer than
413 // PacketSize elements and the packet spills past the concat axis), so fall
414 // back to scalars.
415 template <int LoadMode>
416 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketReturnType packet(Index index) const {
417 const int packetSize = PacketType<CoeffReturnType, Device>::size;
418 EIGEN_STATIC_ASSERT((packetSize > 1), YOU_MADE_A_PROGRAMMING_MISTAKE)
419 eigen_assert(index + packetSize - 1 < dimensions().TotalSize());
420
421 array<Index, NumDims> subs;
422 array<Index, NumDims> subs_end;
423 Index remaining = index;
424 Index remaining_end = index + packetSize - 1;
425 EIGEN_IF_CONSTEXPR (static_cast<int>(Layout) == static_cast<int>(ColMajor)) {
426 for (int i = NumDims - 1; i > 0; --i) {
427 subs[i] = remaining / m_outputStrides[i];
428 remaining -= subs[i] * m_outputStrides[i];
429 subs_end[i] = remaining_end / m_outputStrides[i];
430 remaining_end -= subs_end[i] * m_outputStrides[i];
431 }
432 subs[0] = remaining;
433 subs_end[0] = remaining_end;
434 } else {
435 for (int i = 0; i < NumDims - 1; ++i) {
436 subs[i] = remaining / m_outputStrides[i];
437 remaining -= subs[i] * m_outputStrides[i];
438 subs_end[i] = remaining_end / m_outputStrides[i];
439 remaining_end -= subs_end[i] * m_outputStrides[i];
440 }
441 subs[NumDims - 1] = remaining;
442 subs_end[NumDims - 1] = remaining_end;
443 }
444
445 const Dimensions& left_dims = m_leftImpl.dimensions();
446 const Index left_axis_size = left_dims[m_axis];
447
448 constexpr int innermost = (static_cast<int>(Layout) == static_cast<int>(ColMajor)) ? 0 : NumDims - 1;
449 bool packet_in_single_inner_row = true;
450 EIGEN_UNROLL_LOOP
451 for (int i = 0; i < NumDims; ++i) {
452 if (i != innermost && subs[i] != subs_end[i]) {
453 packet_in_single_inner_row = false;
454 }
455 }
456
457 const bool on_left =
458 packet_in_single_inner_row && subs[m_axis] < left_axis_size && subs_end[m_axis] < left_axis_size;
459 const bool on_right =
460 packet_in_single_inner_row && subs[m_axis] >= left_axis_size && subs_end[m_axis] >= left_axis_size;
461
462 if (on_left) {
463 Index left_index;
464 EIGEN_IF_CONSTEXPR (static_cast<int>(Layout) == static_cast<int>(ColMajor)) {
465 left_index = subs[0];
466 EIGEN_UNROLL_LOOP
467 for (int i = 1; i < NumDims; ++i) {
468 left_index += subs[i] * m_leftStrides[i];
469 }
470 } else {
471 left_index = subs[NumDims - 1];
472 EIGEN_UNROLL_LOOP
473 for (int i = NumDims - 2; i >= 0; --i) {
474 left_index += subs[i] * m_leftStrides[i];
475 }
476 }
477 return m_leftImpl.template packet<LoadMode>(left_index);
478 }
479 if (on_right) {
480 subs[m_axis] -= left_axis_size;
481 Index right_index;
482 EIGEN_IF_CONSTEXPR (static_cast<int>(Layout) == static_cast<int>(ColMajor)) {
483 right_index = subs[0];
484 EIGEN_UNROLL_LOOP
485 for (int i = 1; i < NumDims; ++i) {
486 right_index += subs[i] * m_rightStrides[i];
487 }
488 } else {
489 right_index = subs[NumDims - 1];
490 EIGEN_UNROLL_LOOP
491 for (int i = NumDims - 2; i >= 0; --i) {
492 right_index += subs[i] * m_rightStrides[i];
493 }
494 }
495 return m_rightImpl.template packet<LoadMode>(right_index);
496 }
497
498 // The packet straddles the boundary or spans multiple inner rows: fall
499 // back to assembling scalars.
500 EIGEN_ALIGN_TO_BOUNDARY(internal::unpacket_traits<PacketReturnType>::alignment) CoeffReturnType values[packetSize];
501 EIGEN_UNROLL_LOOP
502 for (int i = 0; i < packetSize; ++i) {
503 values[i] = coeff(index + i);
504 }
505 PacketReturnType rslt = internal::pload<PacketReturnType>(values);
506 return rslt;
507 }
508
509 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorOpCost costPerCoeff(bool vectorized) const {
510 const double compute_cost = NumDims * (2 * TensorOpCost::AddCost<Index>() + 2 * TensorOpCost::MulCost<Index>() +
511 TensorOpCost::DivCost<Index>() + TensorOpCost::ModCost<Index>());
512 const double lhs_size = m_leftImpl.dimensions().TotalSize();
513 const double rhs_size = m_rightImpl.dimensions().TotalSize();
514 return (lhs_size / (lhs_size + rhs_size)) * m_leftImpl.costPerCoeff(vectorized) +
515 (rhs_size / (lhs_size + rhs_size)) * m_rightImpl.costPerCoeff(vectorized) + TensorOpCost(0, 0, compute_cost);
516 }
517
518 EIGEN_DEVICE_FUNC EvaluatorPointerType data() const { return nullptr; }
519
520 protected:
521 Dimensions m_dimensions;
522 array<Index, NumDims> m_outputStrides;
523 array<Index, NumDims> m_leftStrides;
524 array<Index, NumDims> m_rightStrides;
525 TensorEvaluator<LeftArgType, Device> m_leftImpl;
526 TensorEvaluator<RightArgType, Device> m_rightImpl;
527 const Device EIGEN_DEVICE_REF m_device;
528 const Axis m_axis;
529 Index m_leftAxisSize;
530};
531
532// Eval as lvalue
533template <typename Axis, typename LeftArgType, typename RightArgType, typename Device>
534struct TensorEvaluator<TensorConcatenationOp<Axis, LeftArgType, RightArgType>, Device>
535 : public TensorEvaluator<const TensorConcatenationOp<Axis, LeftArgType, RightArgType>, Device> {
536 typedef TensorEvaluator<const TensorConcatenationOp<Axis, LeftArgType, RightArgType>, Device> Base;
537 typedef TensorConcatenationOp<Axis, LeftArgType, RightArgType> XprType;
538 typedef typename Base::Dimensions Dimensions;
539 static constexpr int NumDims = Base::NumDims;
540 static constexpr int Layout = TensorEvaluator<LeftArgType, Device>::Layout;
541 enum {
542 IsAligned = false,
543 PacketAccess =
544 TensorEvaluator<LeftArgType, Device>::PacketAccess && TensorEvaluator<RightArgType, Device>::PacketAccess,
545 // writeBlock() splits the block at the concat axis and copies each piece
546 // straight into the operand's buffer, so both must expose raw storage.
547 BlockAccess = TensorEvaluator<LeftArgType, Device>::RawAccess && TensorEvaluator<RightArgType, Device>::RawAccess,
548 // The coeff/packet write path pays a div/mod cascade plus a per-dim mod
549 // for every scalar; the block path is a pair of bulk copies. Mirrors the
550 // rvalue evaluator.
551 PreferBlockAccess = true,
552 RawAccess = false
553 };
554
555 typedef std::remove_const_t<typename XprType::Scalar> ScalarNoConst;
556
557 //===- Tensor block evaluation strategy (see TensorBlock.h) -------------===//
558 typedef internal::TensorBlockDescriptor<NumDims, typename XprType::Index> TensorBlockDesc;
559 //===--------------------------------------------------------------------===//
560
561 // The ColMajor-only static_assert lives in coeffRef/writePacket rather than
562 // here so that passthrough evaluators (e.g. TensorSlicingOp's) can
563 // instantiate this type for RowMajor concat operands without ever calling
564 // its lvalue methods. writeBlock() below is layout-generic, so tiled
565 // assignments through a RowMajor concatenation are supported.
566 EIGEN_STRONG_INLINE TensorEvaluator(const XprType& op, const Device& device) : Base(op, device) {}
567
568 typedef typename XprType::Index Index;
569 typedef typename XprType::Scalar Scalar;
570 typedef typename XprType::CoeffReturnType CoeffReturnType;
571 typedef typename PacketType<CoeffReturnType, Device>::type PacketReturnType;
572
573 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffReturnType& coeffRef(Index index) const {
574 EIGEN_STATIC_ASSERT((static_cast<int>(Layout) == static_cast<int>(ColMajor)), YOU_MADE_A_PROGRAMMING_MISTAKE);
575 // Collect dimension-wise indices (subs).
576 array<Index, Base::NumDims> subs;
577 for (int i = Base::NumDims - 1; i > 0; --i) {
578 subs[i] = index / this->m_outputStrides[i];
579 index -= subs[i] * this->m_outputStrides[i];
580 }
581 subs[0] = index;
582
583 const Dimensions& left_dims = this->m_leftImpl.dimensions();
584 if (subs[this->m_axis] < left_dims[this->m_axis]) {
585 Index left_index = subs[0];
586 for (int i = 1; i < Base::NumDims; ++i) {
587 left_index += (subs[i] % left_dims[i]) * this->m_leftStrides[i];
588 }
589 return this->m_leftImpl.coeffRef(left_index);
590 } else {
591 subs[this->m_axis] -= left_dims[this->m_axis];
592 const Dimensions& right_dims = this->m_rightImpl.dimensions();
593 Index right_index = subs[0];
594 for (int i = 1; i < Base::NumDims; ++i) {
595 right_index += (subs[i] % right_dims[i]) * this->m_rightStrides[i];
596 }
597 return this->m_rightImpl.coeffRef(right_index);
598 }
599 }
600
601 template <int StoreMode>
602 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void writePacket(Index index, const PacketReturnType& x) const {
603 EIGEN_STATIC_ASSERT((static_cast<int>(Layout) == static_cast<int>(ColMajor)), YOU_MADE_A_PROGRAMMING_MISTAKE);
604 const int packetSize = PacketType<CoeffReturnType, Device>::size;
605 EIGEN_STATIC_ASSERT((packetSize > 1), YOU_MADE_A_PROGRAMMING_MISTAKE)
606 eigen_assert(index + packetSize - 1 < this->dimensions().TotalSize());
607
608 EIGEN_ALIGN_TO_BOUNDARY(internal::unpacket_traits<PacketReturnType>::alignment) CoeffReturnType values[packetSize];
609 internal::pstore<CoeffReturnType, PacketReturnType>(values, x);
610 for (int i = 0; i < packetSize; ++i) {
611 coeffRef(index + i) = values[i];
612 }
613 }
614
615 // Mirror of the rvalue block(): split the block at the concat axis and copy
616 // each piece into the corresponding operand's buffer.
617 template <typename TensorBlock>
618 EIGEN_STRONG_INLINE void writeBlock(const TensorBlockDesc& desc, const TensorBlock& block) {
619 if (desc.size() == 0) return;
620 eigen_assert(this->m_leftImpl.data() != nullptr && this->m_rightImpl.data() != nullptr);
621
622 const DSizes<Index, NumDims> block_strides = internal::strides<Layout>(desc.dimensions());
623
624 // Materialize the block into a temporary buffer if it is lazy.
625 const ScalarNoConst* block_buffer = block.data();
626 void* mem = nullptr;
627 if (block_buffer == nullptr) {
628 mem = this->m_device.allocate(desc.size() * sizeof(Scalar));
629 ScalarNoConst* buf = static_cast<ScalarNoConst*>(mem);
630
631 typedef internal::TensorBlockAssignment<ScalarNoConst, NumDims, typename TensorBlock::XprType, Index>
632 TensorBlockAssignment;
633 TensorBlockAssignment::Run(TensorBlockAssignment::target(desc.dimensions(), block_strides, buf), block.expr());
634
635 block_buffer = buf;
636 }
637
638 // Decompose the block's offset into output coordinates.
639 Index remaining = desc.offset();
640 DSizes<Index, NumDims> out_coords;
641 EIGEN_IF_CONSTEXPR (static_cast<int>(Layout) == static_cast<int>(ColMajor)) {
642 for (int i = NumDims - 1; i > 0; --i) {
643 out_coords[i] = remaining / this->m_outputStrides[i];
644 remaining -= out_coords[i] * this->m_outputStrides[i];
645 }
646 out_coords[0] = remaining;
647 } else {
648 for (int i = 0; i < NumDims - 1; ++i) {
649 out_coords[i] = remaining / this->m_outputStrides[i];
650 remaining -= out_coords[i] * this->m_outputStrides[i];
651 }
652 out_coords[NumDims - 1] = remaining;
653 }
654
655 const Index axis_start = out_coords[this->m_axis];
656 const Index axis_size = desc.dimension(static_cast<int>(this->m_axis));
657 const Index axis_end = axis_start + axis_size;
658 const Index left_axis_size = this->m_leftAxisSize;
659
660 typedef internal::TensorBlockIO<ScalarNoConst, Index, NumDims, Layout> TensorBlockIO;
661 typedef typename TensorBlockIO::Dst TensorBlockIODst;
662 typedef typename TensorBlockIO::Src TensorBlockIOSrc;
663
664 if (axis_start < left_axis_size) {
665 DSizes<Index, NumDims> left_sub_dims = desc.dimensions();
666 left_sub_dims[this->m_axis] = numext::mini(left_axis_size, axis_end) - axis_start;
667
668 Index left_dst_offset = 0;
669 for (int i = 0; i < NumDims; ++i) {
670 left_dst_offset += out_coords[i] * this->m_leftStrides[i];
671 }
672
673 TensorBlockIOSrc src(block_strides, block_buffer, /*src_offset=*/0);
674 TensorBlockIODst dst(left_sub_dims, typename TensorBlockIO::Dimensions(this->m_leftStrides),
675 this->m_leftImpl.data(), left_dst_offset);
676 TensorBlockIO::Copy(dst, src);
677 }
678
679 if (axis_end > left_axis_size) {
680 DSizes<Index, NumDims> right_sub_dims = desc.dimensions();
681 right_sub_dims[this->m_axis] = axis_end - numext::maxi(left_axis_size, axis_start);
682
683 const Index right_axis_start = numext::maxi(Index(0), axis_start - left_axis_size);
684 Index right_dst_offset = right_axis_start * this->m_rightStrides[this->m_axis];
685 for (int i = 0; i < NumDims; ++i) {
686 if (i != this->m_axis) {
687 right_dst_offset += out_coords[i] * this->m_rightStrides[i];
688 }
689 }
690
691 // Offset within the block buffer where the right-side piece starts.
692 const Index src_offset = numext::maxi(Index(0), left_axis_size - axis_start) * block_strides[this->m_axis];
693
694 TensorBlockIOSrc src(block_strides, block_buffer, src_offset);
695 TensorBlockIODst dst(right_sub_dims, typename TensorBlockIO::Dimensions(this->m_rightStrides),
696 this->m_rightImpl.data(), right_dst_offset);
697 TensorBlockIO::Copy(dst, src);
698 }
699
700 // Deallocate temporary buffer used for the block materialization.
701 if (mem != nullptr) this->m_device.deallocate(mem);
702 }
703};
704
705} // end namespace Eigen
706
707#endif // EIGEN_TENSOR_TENSOR_CONCATENATION_H
The tensor base class.
Definition TensorForwardDeclarations.h:69
Tensor concatenation class.
Definition TensorConcatenation.h:49
WriteAccessors
Namespace containing all symbols from the Eigen library.
The tensor evaluator class.
Definition TensorEvaluator.h:47