11#ifndef EIGEN_TENSOR_TENSOR_CONCATENATION_H
12#define EIGEN_TENSOR_TENSOR_CONCATENATION_H
15#include "./InternalHeaderCheck.h"
20template <
typename Axis,
typename LhsXprType,
typename RhsXprType>
21struct traits<TensorConcatenationOp<Axis, LhsXprType, RhsXprType>> {
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;
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;
31 typedef std::conditional_t<Pointer_type_promotion<typename LhsXprType::Scalar, Scalar>::val,
32 typename traits<LhsXprType>::PointerType,
typename traits<RhsXprType>::PointerType>
36template <
typename Axis,
typename LhsXprType,
typename RhsXprType>
37struct eval<TensorConcatenationOp<Axis, LhsXprType, RhsXprType>, Eigen::Dense> {
38 typedef const TensorConcatenationOp<Axis, LhsXprType, RhsXprType>& type;
48template <
typename Axis,
typename LhsXprType,
typename RhsXprType>
49class TensorConcatenationOp :
public TensorBase<TensorConcatenationOp<Axis, LhsXprType, RhsXprType>, WriteAccessors> {
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;
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) {}
63 EIGEN_DEVICE_FUNC
const internal::remove_all_t<typename LhsXprType::Nested>& lhsExpression()
const {
67 EIGEN_DEVICE_FUNC
const internal::remove_all_t<typename RhsXprType::Nested>& rhsExpression()
const {
71 EIGEN_DEVICE_FUNC
const Axis& axis()
const {
return m_axis; }
73 EIGEN_INHERIT_ASSIGNMENT_OPERATORS(TensorConcatenationOp)
75 typename LhsXprType::Nested m_lhs_xpr;
76 typename RhsXprType::Nested m_rhs_xpr;
81template <
typename Axis,
typename LeftArgType,
typename RightArgType,
typename 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;
89 typedef typename XprType::Scalar Scalar;
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;
98 TensorEvaluator<LeftArgType, Device>::PacketAccess && TensorEvaluator<RightArgType, Device>::PacketAccess,
108 PreferBlockAccess =
true,
112 typedef std::remove_const_t<Scalar> ScalarNoConst;
115 typedef internal::TensorBlockDescriptor<NumDims, Index> TensorBlockDesc;
116 typedef internal::TensorBlockScratchAllocator<Device> TensorBlockScratch;
118 typedef typename internal::TensorMaterializedBlock<ScalarNoConst, NumDims, Layout, Index> TensorBlock;
121 EIGEN_STRONG_INLINE TensorEvaluator(
const XprType& op,
const Device& device)
122 : m_leftImpl(op.lhsExpression(), device),
123 m_rightImpl(op.rhsExpression(), device),
126 EIGEN_STATIC_ASSERT((
static_cast<int>(TensorEvaluator<LeftArgType, Device>::Layout) ==
127 static_cast<int>(TensorEvaluator<RightArgType, Device>::Layout) ||
129 YOU_MADE_A_PROGRAMMING_MISTAKE);
133 EIGEN_STATIC_ASSERT((NumDims == RightNumDims), YOU_MADE_A_PROGRAMMING_MISTAKE);
134 EIGEN_STATIC_ASSERT((NumDims > 0), YOU_MADE_A_PROGRAMMING_MISTAKE);
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();
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];
147 eigen_assert(lhs_dims[i] > 0);
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];
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;
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];
168 m_leftStrides[NumDims - 1] = 1;
169 m_rightStrides[NumDims - 1] = 1;
170 m_outputStrides[NumDims - 1] = 1;
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];
180 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
const Dimensions& dimensions()
const {
return m_dimensions; }
183 EIGEN_STRONG_INLINE
bool evalSubExprsIfNeeded(EvaluatorPointerType) {
184 m_leftImpl.evalSubExprsIfNeeded(
nullptr);
185 m_rightImpl.evalSubExprsIfNeeded(
nullptr);
189 EIGEN_STRONG_INLINE
void cleanup() {
190 m_leftImpl.cleanup();
191 m_rightImpl.cleanup();
194 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE internal::TensorBlockResourceRequirements getResourceRequirements()
const {
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()));
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;
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;
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);
237 if (desc.size() == 0) {
238 return TensorBlock(internal::TensorBlockKind::kView,
nullptr, desc.dimensions());
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];
248 out_coords[0] = remaining;
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];
254 out_coords[NumDims - 1] = remaining;
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;
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];
276 return TensorBlock(internal::TensorBlockKind::kView, m_leftImpl.data() + left_src_offset, desc.dimensions());
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) {
283 right_src_offset += out_coords[i] * m_rightStrides[i];
286 return TensorBlock(internal::TensorBlockKind::kView, m_rightImpl.data() + right_src_offset, desc.dimensions());
290 typedef internal::TensorBlockIO<ScalarNoConst, Index, NumDims, Layout> TensorBlockIO;
291 typedef typename TensorBlockIO::Dst TensorBlockIODst;
292 typedef typename TensorBlockIO::Src TensorBlockIOSrc;
297 typename TensorBlock::Storage block_storage =
298 TensorBlock::prepareStorage(desc, scratch, root_of_expr_ast);
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;
305 Index left_src_offset = 0;
306 for (
int i = 0; i < NumDims; ++i) {
307 left_src_offset += out_coords[i] * m_leftStrides[i];
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(),
314 TensorBlockIO::Copy(dst, src);
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;
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) {
328 right_src_offset += out_coords[i] * m_rightStrides[i];
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];
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);
342 return block_storage.AsTensorMaterializedBlock();
350 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffReturnType coeff(Index index)
const {
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];
360 for (
int i = 0; i < NumDims - 1; ++i) {
361 subs[i] = index / m_outputStrides[i];
362 index -= subs[i] * m_outputStrides[i];
364 subs[NumDims - 1] = index;
367 const Dimensions& left_dims = m_leftImpl.dimensions();
368 if (subs[m_axis] < left_dims[m_axis]) {
370 EIGEN_IF_CONSTEXPR (
static_cast<int>(Layout) ==
static_cast<int>(
ColMajor)) {
371 left_index = subs[0];
373 for (
int i = 1; i < NumDims; ++i) {
374 left_index += (subs[i] % left_dims[i]) * m_leftStrides[i];
377 left_index = subs[NumDims - 1];
379 for (
int i = NumDims - 2; i >= 0; --i) {
380 left_index += (subs[i] % left_dims[i]) * m_leftStrides[i];
383 return m_leftImpl.coeff(left_index);
385 subs[m_axis] -= left_dims[m_axis];
386 const Dimensions& right_dims = m_rightImpl.dimensions();
388 EIGEN_IF_CONSTEXPR (
static_cast<int>(Layout) ==
static_cast<int>(
ColMajor)) {
389 right_index = subs[0];
391 for (
int i = 1; i < NumDims; ++i) {
392 right_index += (subs[i] % right_dims[i]) * m_rightStrides[i];
395 right_index = subs[NumDims - 1];
397 for (
int i = NumDims - 2; i >= 0; --i) {
398 right_index += (subs[i] % right_dims[i]) * m_rightStrides[i];
401 return m_rightImpl.coeff(right_index);
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());
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];
433 subs_end[0] = remaining_end;
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];
441 subs[NumDims - 1] = remaining;
442 subs_end[NumDims - 1] = remaining_end;
445 const Dimensions& left_dims = m_leftImpl.dimensions();
446 const Index left_axis_size = left_dims[m_axis];
448 constexpr int innermost = (
static_cast<int>(Layout) ==
static_cast<int>(
ColMajor)) ? 0 : NumDims - 1;
449 bool packet_in_single_inner_row =
true;
451 for (
int i = 0; i < NumDims; ++i) {
452 if (i != innermost && subs[i] != subs_end[i]) {
453 packet_in_single_inner_row =
false;
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;
464 EIGEN_IF_CONSTEXPR (
static_cast<int>(Layout) ==
static_cast<int>(
ColMajor)) {
465 left_index = subs[0];
467 for (
int i = 1; i < NumDims; ++i) {
468 left_index += subs[i] * m_leftStrides[i];
471 left_index = subs[NumDims - 1];
473 for (
int i = NumDims - 2; i >= 0; --i) {
474 left_index += subs[i] * m_leftStrides[i];
477 return m_leftImpl.template packet<LoadMode>(left_index);
480 subs[m_axis] -= left_axis_size;
482 EIGEN_IF_CONSTEXPR (
static_cast<int>(Layout) ==
static_cast<int>(
ColMajor)) {
483 right_index = subs[0];
485 for (
int i = 1; i < NumDims; ++i) {
486 right_index += subs[i] * m_rightStrides[i];
489 right_index = subs[NumDims - 1];
491 for (
int i = NumDims - 2; i >= 0; --i) {
492 right_index += subs[i] * m_rightStrides[i];
495 return m_rightImpl.template packet<LoadMode>(right_index);
500 EIGEN_ALIGN_TO_BOUNDARY(internal::unpacket_traits<PacketReturnType>::alignment) CoeffReturnType values[packetSize];
502 for (
int i = 0; i < packetSize; ++i) {
503 values[i] = coeff(index + i);
505 PacketReturnType rslt = internal::pload<PacketReturnType>(values);
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);
518 EIGEN_DEVICE_FUNC EvaluatorPointerType data()
const {
return nullptr; }
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;
529 Index m_leftAxisSize;
533template <
typename Axis,
typename LeftArgType,
typename RightArgType,
typename 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;
544 TensorEvaluator<LeftArgType, Device>::PacketAccess && TensorEvaluator<RightArgType, Device>::PacketAccess,
547 BlockAccess = TensorEvaluator<LeftArgType, Device>::RawAccess && TensorEvaluator<RightArgType, Device>::RawAccess,
551 PreferBlockAccess =
true,
555 typedef std::remove_const_t<typename XprType::Scalar> ScalarNoConst;
558 typedef internal::TensorBlockDescriptor<NumDims, typename XprType::Index> TensorBlockDesc;
566 EIGEN_STRONG_INLINE TensorEvaluator(
const XprType& op,
const Device& device) : Base(op, device) {}
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;
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);
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];
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];
589 return this->m_leftImpl.coeffRef(left_index);
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];
597 return this->m_rightImpl.coeffRef(right_index);
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());
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];
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);
622 const DSizes<Index, NumDims> block_strides = internal::strides<Layout>(desc.dimensions());
625 const ScalarNoConst* block_buffer = block.data();
627 if (block_buffer ==
nullptr) {
628 mem = this->m_device.allocate(desc.size() *
sizeof(Scalar));
629 ScalarNoConst* buf =
static_cast<ScalarNoConst*
>(mem);
631 typedef internal::TensorBlockAssignment<ScalarNoConst, NumDims, typename TensorBlock::XprType, Index>
632 TensorBlockAssignment;
633 TensorBlockAssignment::Run(TensorBlockAssignment::target(desc.dimensions(), block_strides, buf), block.expr());
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];
646 out_coords[0] = remaining;
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];
652 out_coords[NumDims - 1] = remaining;
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;
660 typedef internal::TensorBlockIO<ScalarNoConst, Index, NumDims, Layout> TensorBlockIO;
661 typedef typename TensorBlockIO::Dst TensorBlockIODst;
662 typedef typename TensorBlockIO::Src TensorBlockIOSrc;
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;
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];
673 TensorBlockIOSrc src(block_strides, block_buffer, 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);
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);
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];
692 const Index src_offset = numext::maxi(Index(0), left_axis_size - axis_start) * block_strides[this->m_axis];
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);
701 if (mem !=
nullptr) this->m_device.deallocate(mem);
The tensor base class.
Definition TensorForwardDeclarations.h:69
Tensor concatenation class.
Definition TensorConcatenation.h:49
Namespace containing all symbols from the Eigen library.
The tensor evaluator class.
Definition TensorEvaluator.h:47