10#ifndef EIGEN_TENSOR_TENSOR_BLOCK_H
11#define EIGEN_TENSOR_TENSOR_BLOCK_H
14#include "./InternalHeaderCheck.h"
21template <
typename Scalar,
typename IndexType,
int NumDims,
int Layout>
30template <
int Layout,
typename IndexType,
int NumDims>
31EIGEN_ALWAYS_INLINE std::enable_if_t<NumDims == 0, DSizes<IndexType, NumDims> > strides_impl(
32 const DSizes<IndexType, NumDims>& ) {
33 DSizes<IndexType, NumDims> strides;
37template <
int Layout,
typename IndexType,
int NumDims>
38EIGEN_ALWAYS_INLINE std::enable_if_t<(NumDims > 0), DSizes<IndexType, NumDims> > strides_impl(
39 const DSizes<IndexType, NumDims>& dimensions) {
40 DSizes<IndexType, NumDims> strides;
42 EIGEN_IF_CONSTEXPR (
static_cast<int>(Layout) ==
static_cast<int>(
ColMajor)) {
44 for (
int i = 1; i < NumDims; ++i) {
45 strides[i] = strides[i - 1] * dimensions[i - 1];
48 strides[NumDims - 1] = 1;
49 for (
int i = NumDims - 2; i >= 0; --i) {
50 strides[i] = strides[i + 1] * dimensions[i + 1];
57template <
int Layout,
typename IndexType,
int NumDims>
58EIGEN_ALWAYS_INLINE DSizes<IndexType, NumDims> strides(
const DSizes<IndexType, NumDims>& dimensions) {
59 return strides_impl<Layout>(dimensions);
62template <
int Layout,
typename IndexType,
size_t NumDims>
63EIGEN_ALWAYS_INLINE DSizes<IndexType, NumDims> strides(
const Eigen::array<IndexType, NumDims>& dimensions) {
64 return strides<Layout>(DSizes<IndexType, NumDims>(dimensions));
67template <
int Layout, std::ptrdiff_t... Indices>
68EIGEN_STRONG_INLINE DSizes<std::ptrdiff_t,
sizeof...(Indices)> strides(
const Sizes<Indices...>& sizes) {
69 return strides<Layout>(DSizes<std::ptrdiff_t,
sizeof...(Indices)>(sizes));
85enum class TensorBlockShapeType { kUniformAllDims, kSkewedInnerDims };
87struct TensorBlockResourceRequirements {
88 TensorBlockShapeType shape_type;
90 TensorOpCost cost_per_coeff;
96 EIGEN_DEVICE_FUNC TensorBlockResourceRequirements(TensorBlockShapeType shape_type_,
size_t size_, TensorOpCost cost_)
97 : shape_type(shape_type_), size(size_), cost_per_coeff(cost_) {}
100 template <
typename Scalar>
101 EIGEN_DEVICE_FUNC
static TensorBlockResourceRequirements withShapeAndSize(TensorBlockShapeType shape_type,
102 size_t size_in_bytes, TensorOpCost cost) {
103 const size_t size = numext::maxi(
size_t(1), size_in_bytes /
sizeof(Scalar));
104 return {shape_type, size, cost};
107 template <
typename Scalar>
108 EIGEN_DEVICE_FUNC
static TensorBlockResourceRequirements withShapeAndSize(TensorBlockShapeType shape_type,
109 size_t size_in_bytes) {
124 return withShapeAndSize<Scalar>(shape_type, size_in_bytes,
130 template <
typename Scalar>
131 EIGEN_DEVICE_FUNC
static TensorBlockResourceRequirements skewed(
size_t size_in_bytes) {
132 return withShapeAndSize<Scalar>(TensorBlockShapeType::kSkewedInnerDims, size_in_bytes);
135 template <
typename Scalar>
136 EIGEN_DEVICE_FUNC
static TensorBlockResourceRequirements uniform(
size_t size_in_bytes) {
137 return withShapeAndSize<Scalar>(TensorBlockShapeType::kUniformAllDims, size_in_bytes);
140 EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE TensorBlockResourceRequirements
141 merge(
const TensorBlockResourceRequirements& lhs,
const TensorBlockResourceRequirements& rhs) {
142 return {merge(lhs.shape_type, rhs.shape_type),
143 merge(lhs.size, rhs.size),
144 merge(lhs.cost_per_coeff, rhs.cost_per_coeff)};
147 EIGEN_DEVICE_FUNC TensorBlockResourceRequirements& addCostPerCoeff(TensorOpCost cost) {
148 cost_per_coeff += cost;
155 EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE TensorBlockResourceRequirements any() {
156 return {TensorBlockShapeType::kUniformAllDims, 1, {0, 0, 0}};
160 EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE
size_t merge(
size_t lhs_size,
size_t rhs_size) {
161 return numext::maxi(lhs_size, rhs_size);
164 EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE TensorBlockShapeType merge(TensorBlockShapeType lhs,
165 TensorBlockShapeType rhs) {
166 return (lhs == TensorBlockShapeType::kSkewedInnerDims || rhs == TensorBlockShapeType::kSkewedInnerDims)
167 ? TensorBlockShapeType::kSkewedInnerDims
168 : TensorBlockShapeType::kUniformAllDims;
171 EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE TensorOpCost merge(TensorOpCost lhs_cost, TensorOpCost rhs_cost) {
172 return lhs_cost + rhs_cost;
180template <
int NumDims,
typename IndexType = Eigen::Index>
181class TensorBlockDescriptor {
183 typedef DSizes<IndexType, NumDims> Dimensions;
195 class DestinationBuffer {
197 enum DestinationBufferKind :
int {
227 template <
typename Scalar>
228 Scalar* data()
const {
229 eigen_assert(m_data_type_size ==
sizeof(Scalar));
230 return static_cast<Scalar*
>(m_data);
233 const Dimensions& strides()
const {
return m_strides; }
234 const DestinationBufferKind& kind()
const {
return m_kind; }
237 friend class TensorBlockDescriptor<NumDims, IndexType>;
239 DestinationBuffer() =
default;
241 template <
typename Scalar>
242 DestinationBuffer(Scalar* data,
const Dimensions& strides, DestinationBufferKind kind)
243 : m_data(static_cast<void*>(data)), m_data_type_size(sizeof(Scalar)), m_strides(strides), m_kind(kind) {}
245 template <
int Layout,
typename Scalar>
246 static DestinationBuffer make(
const TensorBlockDescriptor& desc, Scalar* data,
const Dimensions& strides) {
247 return DestinationBuffer(data, strides, kind<Layout>(desc, strides));
250 template <
int Layout>
251 static DestinationBufferKind kind(
const TensorBlockDescriptor& desc,
const Dimensions& strides) {
252 const Dimensions& desc_dims = desc.dimensions();
253 const Dimensions& desc_strides = internal::strides<Layout>(desc_dims);
254 for (
int i = 0; i < NumDims; ++i) {
255 if (desc_dims[i] == 1)
continue;
256 if (desc_strides[i] != strides[i])
return kStrided;
263 void* m_data =
nullptr;
264 size_t m_data_type_size = 0;
268 Dimensions m_strides;
270 DestinationBufferKind m_kind = kEmpty;
273 TensorBlockDescriptor(
const IndexType offset,
const Dimensions& dimensions,
const DestinationBuffer& destination)
274 : m_offset(offset), m_dimensions(dimensions), m_destination(destination) {}
276 TensorBlockDescriptor(
const IndexType offset,
const Dimensions& dimensions)
277 : m_offset(offset), m_dimensions(dimensions), m_destination(DestinationBuffer()) {}
279 IndexType offset()
const {
return m_offset; }
280 const Dimensions& dimensions()
const {
return m_dimensions; }
281 IndexType dimension(
int index)
const {
return m_dimensions[index]; }
282 IndexType size()
const {
return array_prod<IndexType>(m_dimensions); }
284 const DestinationBuffer& destination()
const {
return m_destination; }
286 template <
int Layout,
typename Scalar>
287 void AddDestinationBuffer(Scalar* dst_base,
const Dimensions& dst_strides) {
288 eigen_assert(dst_base !=
nullptr);
289 m_destination = DestinationBuffer::template make<Layout>(*
this, dst_base, dst_strides);
292 template <
int Layout,
typename Scalar,
typename DstStr
idesIndexType>
293 void AddDestinationBuffer(Scalar* dst_base,
const DSizes<DstStridesIndexType, NumDims>& dst_strides) {
295 AddDestinationBuffer<Layout>(dst_base, Dimensions(dst_strides));
298 TensorBlockDescriptor& DropDestinationBuffer() {
299 m_destination.m_data =
nullptr;
300 m_destination.m_kind = DestinationBuffer::kEmpty;
304 bool HasDestinationBuffer()
const {
return m_destination.kind() != DestinationBuffer::kEmpty; }
307 TensorBlockDescriptor WithOffset(IndexType offset)
const {
308 return TensorBlockDescriptor(offset, m_dimensions, m_destination);
314 const IndexType m_offset;
315 const Dimensions m_dimensions;
316 DestinationBuffer m_destination;
322template <
int NumDims,
int Layout,
typename IndexType = Eigen::Index>
323class TensorBlockMapper {
324 typedef TensorBlockDescriptor<NumDims, IndexType> BlockDescriptor;
327 typedef DSizes<IndexType, NumDims> Dimensions;
329 TensorBlockMapper() =
default;
330 TensorBlockMapper(
const DSizes<IndexType, NumDims>& dimensions,
const TensorBlockResourceRequirements& requirements)
331 : m_tensor_dimensions(dimensions), m_requirements(requirements) {
333 InitializeBlockDimensions();
336 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE IndexType blockCount()
const {
return m_total_block_count; }
338 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE IndexType blockTotalSize()
const {
return m_block_dimensions.TotalSize(); }
340 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
const DSizes<IndexType, NumDims>& blockDimensions()
const {
341 return m_block_dimensions;
344 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE BlockDescriptor blockDescriptor(IndexType block_index)
const {
345 static constexpr bool isColMajor = Layout ==
static_cast<int>(
ColMajor);
347 IndexType offset = 0;
348 DSizes<IndexType, NumDims> dimensions;
350 EIGEN_IF_CONSTEXPR (NumDims == 0) return BlockDescriptor(offset, dimensions);
353 for (
int i = NumDims - 1; i >= 0; --i) {
354 const int dim = isColMajor ? i : NumDims - i - 1;
356 const IndexType idx = block_index / m_block_strides[dim];
357 block_index -= idx * m_block_strides[dim];
359 const IndexType coord = idx * m_block_dimensions[dim];
360 dimensions[dim] = numext::mini(m_tensor_dimensions[dim] - coord, m_block_dimensions[dim]);
361 offset += coord * m_tensor_strides[dim];
364 return {offset, dimensions};
368 void InitializeBlockDimensions() {
370 const TensorBlockShapeType shape_type = m_requirements.shape_type;
371 IndexType target_block_size = numext::maxi<IndexType>(1,
static_cast<IndexType
>(m_requirements.size));
373 IndexType tensor_size = m_tensor_dimensions.TotalSize();
379 if (tensor_size == 0) {
380 for (
int i = 0; i < NumDims; ++i) {
381 m_block_dimensions[i] = 1;
383 m_total_block_count = 0;
388 if (tensor_size <= target_block_size) {
389 m_block_dimensions = m_tensor_dimensions;
390 m_total_block_count = 1;
393 for (
int i = 0; i < NumDims; ++i) {
394 m_tensor_strides[i] = 0;
395 m_block_strides[i] = 1;
400 static constexpr bool isColMajor = Layout ==
static_cast<int>(
ColMajor);
403 if (shape_type == TensorBlockShapeType::kSkewedInnerDims) {
404 IndexType coeff_to_allocate = target_block_size;
406 for (
int i = 0; i < NumDims; ++i) {
407 const int dim = isColMajor ? i : NumDims - i - 1;
408 m_block_dimensions[dim] = numext::mini(coeff_to_allocate, m_tensor_dimensions[dim]);
410 numext::div_ceil(coeff_to_allocate, numext::maxi(
static_cast<IndexType
>(1), m_block_dimensions[dim]));
412 eigen_assert(coeff_to_allocate == 1);
414 }
else if (shape_type == TensorBlockShapeType::kUniformAllDims) {
417 const IndexType dim_size_target = convert_index<IndexType>(
418 numext::pow(
static_cast<float>(target_block_size), 1.0f /
static_cast<float>(m_block_dimensions.rank())));
420 for (
int i = 0; i < NumDims; ++i) {
425 m_block_dimensions[i] = numext::mini(dim_size_target, m_tensor_dimensions[i]);
429 IndexType total_size = m_block_dimensions.TotalSize();
430 for (
int i = 0; i < NumDims; ++i) {
431 const int dim = isColMajor ? i : NumDims - i - 1;
433 if (m_block_dimensions[dim] < m_tensor_dimensions[dim]) {
434 const IndexType total_size_other_dims = total_size / m_block_dimensions[dim];
435 const IndexType alloc_avail = numext::div_ceil<IndexType>(target_block_size, total_size_other_dims);
436 if (alloc_avail == m_block_dimensions[dim]) {
440 m_block_dimensions[dim] = numext::mini(m_tensor_dimensions[dim], alloc_avail);
441 total_size = total_size_other_dims * m_block_dimensions[dim];
449 eigen_assert(m_block_dimensions.TotalSize() >=
450 numext::mini<IndexType>(target_block_size, m_tensor_dimensions.TotalSize()));
453 DSizes<IndexType, NumDims> block_count;
454 for (
int i = 0; i < NumDims; ++i) {
455 block_count[i] = numext::div_ceil(m_tensor_dimensions[i], m_block_dimensions[i]);
457 m_total_block_count = array_prod(block_count);
460 m_tensor_strides = strides<Layout>(m_tensor_dimensions);
461 m_block_strides = strides<Layout>(block_count);
464 DSizes<IndexType, NumDims> m_tensor_dimensions;
465 TensorBlockResourceRequirements m_requirements;
467 DSizes<IndexType, NumDims> m_block_dimensions;
468 IndexType m_total_block_count;
470 DSizes<IndexType, NumDims> m_tensor_strides;
471 DSizes<IndexType, NumDims> m_block_strides;
483template <
typename Device>
484class TensorBlockScratchAllocator {
486 explicit TensorBlockScratchAllocator(
const Device& device) : m_device(device), m_allocation_index(0) {}
488 ~TensorBlockScratchAllocator() {
489 for (
size_t i = 0; i < m_allocations.size(); ++i) {
490 m_device.deallocate(m_allocations[i].ptr);
494 void* allocate(
size_t size) {
496 if (m_allocations.capacity() == 0) m_allocations.reserve(8);
499 const int num_allocations =
static_cast<int>(m_allocations.size());
500 const bool has_allocation = m_allocation_index < num_allocations;
503 eigen_assert(m_allocation_index <= num_allocations);
510 if (has_allocation && m_allocations[m_allocation_index].size < size) {
511 m_device.deallocate(m_allocations[m_allocation_index].ptr);
512 m_allocations[m_allocation_index].ptr = m_device.allocate(size);
513 m_allocations[m_allocation_index].size = size;
517 if (!has_allocation) {
518 Allocation allocation;
519 allocation.ptr = m_device.allocate(size);
520 allocation.size = size;
521 m_allocations.push_back(allocation);
524 eigen_assert(m_allocations[m_allocation_index].ptr !=
nullptr);
525 eigen_assert(m_allocations[m_allocation_index].size >= size);
527 return m_allocations[m_allocation_index++].ptr;
530 void reset() { m_allocation_index = 0; }
538 const Device& m_device;
539 int m_allocation_index;
541 std::vector<Allocation> m_allocations;
547enum TensorBlockKind {
559 kMaterializedInScratch,
568 kMaterializedInOutput
575class TensorBlockNotImplemented {
577 typedef void XprType;
580template <
typename Scalar,
int NumDims,
int Layout,
typename IndexType>
581class TensorBlockView;
583template <
typename Scalar,
int NumDims,
int Layout,
typename IndexType>
584struct traits<TensorBlockView<Scalar, NumDims, Layout, IndexType>>
585 : traits<Tensor<Scalar, NumDims, Layout, IndexType>> {
586 static constexpr unsigned int Flags = 0;
591template <
typename Scalar_,
int NumDims,
int Layout,
typename IndexType>
592class TensorBlockView :
public TensorBase<TensorBlockView<Scalar_, NumDims, Layout, IndexType>> {
594 using Scalar = Scalar_;
595 using Index = IndexType;
596 using Dimensions = DSizes<Index, NumDims>;
597 using Nested = TensorBlockView;
598 using StorageKind = Dense;
599 using CoeffReturnType = Scalar;
601 TensorBlockView(
const Scalar* data,
const Dimensions& dimensions)
602 : TensorBlockView(data, dimensions, internal::strides<Layout>(dimensions)) {}
604 TensorBlockView(
const Scalar* data,
const Dimensions& dimensions,
const Dimensions& strides)
605 : m_data(data), m_dimensions(dimensions), m_strides(strides), m_contiguous(true) {
606 eigen_assert(NumDims == 0 || strides[Layout ==
ColMajor ? 0 : NumDims - 1] == 1);
608 for (
int i = 0; i < NumDims; ++i) {
609 const int dim = Layout ==
ColMajor ? i : NumDims - 1 - i;
610 if (dimensions[dim] > 1 && strides[dim] != stride) m_contiguous =
false;
611 stride *= dimensions[dim];
615 EIGEN_DEVICE_FUNC
const Dimensions& dimensions()
const {
return m_dimensions; }
616 EIGEN_DEVICE_FUNC
const Dimensions& strides()
const {
return m_strides; }
617 EIGEN_DEVICE_FUNC
const Scalar* data()
const {
return m_contiguous ? m_data :
nullptr; }
618 EIGEN_DEVICE_FUNC
const Scalar* rawData()
const {
return m_data; }
621 const Scalar* m_data;
622 Dimensions m_dimensions;
623 Dimensions m_strides;
629template <
typename Scalar_,
int NumDims,
int Layout_,
typename IndexType,
typename Device>
630struct TensorEvaluator<const internal::TensorBlockView<Scalar_, NumDims, Layout_, IndexType>, Device> {
631 using XprType = internal::TensorBlockView<Scalar_, NumDims, Layout_, IndexType>;
632 using Scalar = Scalar_;
633 using Index = IndexType;
634 using Dimensions = DSizes<Index, NumDims>;
635 using CoeffReturnType = Scalar;
636 using PacketReturnType =
typename PacketType<Scalar, Device>::type;
637 using EvaluatorPointerType =
const Scalar*;
638 using TensorBlock = internal::TensorBlockNotImplemented;
639 static constexpr int Layout = Layout_;
640 static constexpr bool IsAligned =
false;
641 static constexpr bool PacketAccess = internal::packet_traits<Scalar>::Vectorizable;
642 static constexpr bool BlockAccess =
false;
643 static constexpr bool PreferBlockAccess =
false;
644 static constexpr bool CoordAccess =
false;
645 static constexpr bool RawAccess =
false;
647 TensorEvaluator(
const XprType& expression,
const Device&)
648 : m_expression(expression), m_output_strides(internal::strides<Layout>(expression.dimensions())) {
649 if (!expression.data()) {
650 for (int i = 0; i < NumDims; ++i) {
651 m_divisors[i] = internal::TensorIntDivisor<Index>(numext::maxi(Index(1), m_output_strides[i]));
656 EIGEN_DEVICE_FUNC
const Dimensions& dimensions()
const {
return m_expression.dimensions(); }
657 EIGEN_DEVICE_FUNC
bool evalSubExprsIfNeeded(EvaluatorPointerType) {
return true; }
658 EIGEN_DEVICE_FUNC
void cleanup() {}
659 EIGEN_DEVICE_FUNC
const Scalar* data()
const {
return m_expression.data(); }
661 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index srcCoeff(Index index)
const {
662 if (m_expression.data())
return index;
664 for (
int i = NumDims - 1; i > 0; --i) {
665 const int dim = Layout ==
ColMajor ? i : NumDims - 1 - i;
666 const Index coordinate = index / m_divisors[dim];
667 offset += coordinate * m_expression.strides()[dim];
668 index -= coordinate * m_output_strides[dim];
670 return offset + index;
673 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
CoeffReturnType coeff(Index index)
const {
674 return m_expression.rawData()[srcCoeff(index)];
677 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
const Scalar* coeffAddress(Index index)
const {
678 return m_expression.rawData() + srcCoeff(index);
681 template <
int LoadMode>
682 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketReturnType packet(Index index)
const {
683 constexpr int PacketSize = PacketType<Scalar, Device>::size;
684 const Index first = srcCoeff(index);
685 if (m_expression.data() || srcCoeff(index + PacketSize - 1) == first + PacketSize - 1) {
686 return internal::ploadu<PacketReturnType>(m_expression.rawData() + first);
688 EIGEN_ALIGN_MAX Scalar values[PacketSize];
689 for (
int i = 0; i < PacketSize; ++i) values[i] = coeff(index + i);
690 return internal::ploadu<PacketReturnType>(values);
693 EIGEN_DEVICE_FUNC TensorOpCost costPerCoeff(
bool vectorized)
const {
694 return TensorOpCost(
sizeof(Scalar), 0, m_expression.data() ? 0 : NumDims, vectorized,
695 PacketType<Scalar, Device>::size);
699 XprType m_expression;
701 array<internal::TensorIntDivisor<Index>, NumDims> m_divisors;
704template <
typename Scalar,
int Layout,
typename Index,
typename Device>
705struct TensorEvaluator<const internal::TensorBlockView<Scalar, 1, Layout, Index>, Device>
706 : TensorEvaluator<const TensorMap<const Tensor<Scalar, 1, Layout, Index>>, Device> {
707 using XprType = internal::TensorBlockView<Scalar, 1, Layout, Index>;
708 using MapType = TensorMap<const Tensor<Scalar, 1, Layout, Index>>;
709 using Base = TensorEvaluator<const MapType, Device>;
710 TensorEvaluator(
const XprType& expression,
const Device& device)
711 : Base(MapType(expression.rawData(), expression.dimensions()), device) {}
712 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
const Scalar* coeffAddress(Index index)
const {
return this->data() + index; }
722template <
typename XprType>
724 typedef typename XprType::Scalar type;
727struct XprScalar<void> {
749template <
typename Scalar,
int NumDims,
int Layout,
typename IndexType = Eigen::Index>
750class TensorMaterializedBlock {
752 typedef DSizes<IndexType, NumDims> Dimensions;
753 using XprType = TensorBlockView<Scalar, NumDims, Layout, IndexType>;
755 TensorMaterializedBlock(TensorBlockKind kind,
const Scalar* data,
const Dimensions& dimensions,
756 bool valid_expr =
true)
757 : m_kind(kind), m_data(data), m_dimensions(dimensions), m_expr(m_data, m_dimensions), m_valid_expr(valid_expr) {
758 eigen_assert(m_kind == internal::TensorBlockKind::kView ||
759 m_kind == internal::TensorBlockKind::kMaterializedInScratch ||
760 m_kind == internal::TensorBlockKind::kMaterializedInOutput);
763 TensorBlockKind kind()
const {
return m_kind; }
764 const XprType& expr()
const {
765 eigen_assert(m_valid_expr);
769 const Scalar* data()
const {
return m_valid_expr ? m_expr.data() : m_data; }
772 typedef internal::TensorBlockDescriptor<NumDims, IndexType> TensorBlockDesc;
784 Scalar* data()
const {
return m_data; }
785 const Dimensions& dimensions()
const {
return m_dimensions; }
786 const Dimensions& strides()
const {
return m_strides; }
788 TensorMaterializedBlock AsTensorMaterializedBlock()
const {
789 return TensorMaterializedBlock(m_materialized_in_output ? internal::TensorBlockKind::kMaterializedInOutput
790 : internal::TensorBlockKind::kMaterializedInScratch,
791 m_data, m_dimensions, !m_strided_storage);
795 friend class TensorMaterializedBlock<Scalar, NumDims, Layout, IndexType>;
797 Storage(Scalar* data,
const Dimensions& dimensions,
const Dimensions& strides,
bool materialized_in_output,
798 bool strided_storage)
800 m_dimensions(dimensions),
802 m_materialized_in_output(materialized_in_output),
803 m_strided_storage(strided_storage) {}
806 Dimensions m_dimensions;
807 Dimensions m_strides;
808 bool m_materialized_in_output;
809 bool m_strided_storage;
814 template <
typename TensorBlockScratch>
815 EIGEN_STRONG_INLINE
static Storage prepareStorage(TensorBlockDesc& desc, TensorBlockScratch& scratch,
816 bool allow_strided_storage =
false) {
818 typedef typename TensorBlockDesc::DestinationBuffer DestinationBuffer;
820 if (desc.destination().kind() == DestinationBuffer::kContiguous) {
821 Scalar* buffer = desc.destination().template data<Scalar>();
822 desc.DropDestinationBuffer();
823 return Storage(buffer, desc.dimensions(), internal::strides<Layout>(desc.dimensions()),
827 }
else if (desc.destination().kind() == DestinationBuffer::kStrided && allow_strided_storage) {
828 Scalar* buffer = desc.destination().template data<Scalar>();
829 desc.DropDestinationBuffer();
830 return Storage(buffer, desc.dimensions(), desc.destination().strides(),
834 void* mem = scratch.allocate(desc.size() *
sizeof(Scalar));
835 return Storage(
static_cast<Scalar*
>(mem), desc.dimensions(), internal::strides<Layout>(desc.dimensions()),
842 template <
typename DataDimensions,
typename TensorBlockScratch>
843 EIGEN_STRONG_INLINE
static TensorMaterializedBlock materialize(
const Scalar* data,
const DataDimensions& data_dims,
844 TensorBlockDesc& desc,
845 TensorBlockScratch& ) {
846 eigen_assert(array_size<DataDimensions>::value == desc.dimensions().size());
848 TensorMaterializedBlock block(internal::TensorBlockKind::kView, data + desc.offset(), desc.dimensions());
849 block.m_expr = XprType(data + desc.offset(), desc.dimensions(), internal::strides<Layout>(Dimensions(data_dims)));
854 TensorBlockKind m_kind;
855 const Scalar* m_data;
856 Dimensions m_dimensions;
865template <
typename UnaryOp,
typename ArgTensorBlock>
866class TensorCwiseUnaryBlock {
867 static constexpr bool NoArgBlockAccess = std::is_void<typename ArgTensorBlock::XprType>::value;
870 typedef std::conditional_t<NoArgBlockAccess, void,
871 TensorCwiseUnaryOp<UnaryOp, const typename ArgTensorBlock::XprType> >
874 typedef typename XprScalar<XprType>::type Scalar;
876 TensorCwiseUnaryBlock(
const ArgTensorBlock& arg_block,
const UnaryOp& functor)
877 : m_arg_block(arg_block), m_functor(functor) {}
879 TensorBlockKind kind()
const {
return internal::TensorBlockKind::kExpr; }
881 XprType expr()
const {
return XprType(m_arg_block.expr(), m_functor); }
882 const Scalar* data()
const {
return nullptr; }
883 void cleanup() { m_arg_block.cleanup(); }
886 ArgTensorBlock m_arg_block;
894template <
typename BinaryOp,
typename LhsTensorBlock,
typename RhsTensorBlock>
895class TensorCwiseBinaryBlock {
896 static constexpr bool NoArgBlockAccess =
897 std::is_void<typename LhsTensorBlock::XprType>::value || std::is_void<typename RhsTensorBlock::XprType>::value;
900 typedef std::conditional_t<
901 NoArgBlockAccess, void,
902 TensorCwiseBinaryOp<BinaryOp, const typename LhsTensorBlock::XprType, const typename RhsTensorBlock::XprType> >
905 typedef typename XprScalar<XprType>::type Scalar;
907 TensorCwiseBinaryBlock(
const LhsTensorBlock& left_block,
const RhsTensorBlock& right_block,
const BinaryOp& functor)
908 : m_left_block(left_block), m_right_block(right_block), m_functor(functor) {}
910 TensorBlockKind kind()
const {
return internal::TensorBlockKind::kExpr; }
912 XprType expr()
const {
return XprType(m_left_block.expr(), m_right_block.expr(), m_functor); }
914 const Scalar* data()
const {
return nullptr; }
917 m_left_block.cleanup();
918 m_right_block.cleanup();
922 LhsTensorBlock m_left_block;
923 RhsTensorBlock m_right_block;
932template <
typename BlockFactory,
typename ArgTensorBlock>
933class TensorUnaryExprBlock {
934 typedef typename ArgTensorBlock::XprType ArgXprType;
935 static constexpr bool NoArgBlockAccess = std::is_void<ArgXprType>::value;
938 typedef std::conditional_t<NoArgBlockAccess, void, typename BlockFactory::template XprType<ArgXprType>::type> XprType;
940 typedef typename XprScalar<XprType>::type Scalar;
942 TensorUnaryExprBlock(
const ArgTensorBlock& arg_block,
const BlockFactory& factory)
943 : m_arg_block(arg_block), m_factory(factory) {}
945 TensorBlockKind kind()
const {
return internal::TensorBlockKind::kExpr; }
946 XprType expr()
const {
return m_factory.expr(m_arg_block.expr()); }
947 const Scalar* data()
const {
return nullptr; }
948 void cleanup() { m_arg_block.cleanup(); }
951 ArgTensorBlock m_arg_block;
952 BlockFactory m_factory;
959template <
typename BlockFactory,
typename Arg1TensorBlock,
typename Arg2TensorBlock,
typename Arg3TensorBlock>
960class TensorTernaryExprBlock {
961 typedef typename Arg1TensorBlock::XprType Arg1XprType;
962 typedef typename Arg2TensorBlock::XprType Arg2XprType;
963 typedef typename Arg3TensorBlock::XprType Arg3XprType;
965 static constexpr bool NoArgBlockAccess =
966 std::is_void<Arg1XprType>::value || std::is_void<Arg2XprType>::value || std::is_void<Arg3XprType>::value;
969 typedef std::conditional_t<NoArgBlockAccess, void,
970 typename BlockFactory::template XprType<Arg1XprType, Arg2XprType, Arg3XprType>::type>
973 typedef typename XprScalar<XprType>::type Scalar;
975 TensorTernaryExprBlock(
const Arg1TensorBlock& arg1_block,
const Arg2TensorBlock& arg2_block,
976 const Arg3TensorBlock& arg3_block,
const BlockFactory& factory)
977 : m_arg1_block(arg1_block), m_arg2_block(arg2_block), m_arg3_block(arg3_block), m_factory(factory) {}
979 TensorBlockKind kind()
const {
return internal::TensorBlockKind::kExpr; }
980 XprType expr()
const {
return m_factory.expr(m_arg1_block.expr(), m_arg2_block.expr(), m_arg3_block.expr()); }
981 const Scalar* data()
const {
return nullptr; }
983 m_arg1_block.cleanup();
984 m_arg2_block.cleanup();
985 m_arg3_block.cleanup();
989 Arg1TensorBlock m_arg1_block;
990 Arg2TensorBlock m_arg2_block;
991 Arg3TensorBlock m_arg3_block;
992 BlockFactory m_factory;
999template <
typename Scalar,
typename IndexType>
1000class StridedLinearBufferCopy {
1001 typedef typename packet_traits<Scalar>::type Packet;
1002 typedef typename unpacket_traits<Packet>::half HalfPacket;
1004 Vectorizable = packet_traits<Scalar>::Vectorizable,
1005 PacketSize = packet_traits<Scalar>::size,
1006 HalfPacketSize = unpacket_traits<HalfPacket>::size,
1007 HasHalfPacket =
static_cast<int>(HalfPacketSize) <
static_cast<int>(PacketSize)
1025 Dst(IndexType o, IndexType s, Scalar* d) : offset(o), stride(s), data(d) {}
1033 Src(IndexType o, IndexType s,
const Scalar* d) : offset(o), stride(s), data(d) {}
1040 template <
typename Str
idedLinearBufferCopy::Kind kind>
1041 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void Run(
const Dst& dst,
const Src& src,
const size_t count) {
1042 Run<kind>(count, dst.offset, dst.stride, dst.data, src.offset, src.stride, src.data);
1046 template <
typename Str
idedLinearBufferCopy::Kind kind>
1047 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void Run(
const IndexType count,
const IndexType dst_offset,
1048 const IndexType dst_stride, Scalar* EIGEN_RESTRICT dst_data,
1049 const IndexType src_offset,
const IndexType src_stride,
1050 const Scalar* EIGEN_RESTRICT src_data) {
1051 const Scalar* src = &src_data[src_offset];
1052 Scalar* dst = &dst_data[dst_offset];
1054 EIGEN_IF_CONSTEXPR (!Vectorizable) {
1055 for (Index i = 0; i < count; ++i) {
1056 dst[i * dst_stride] = src[i * src_stride];
1061 const IndexType vectorized_size = PacketSize * (count / PacketSize);
1064 EIGEN_IF_CONSTEXPR (kind == StridedLinearBufferCopy::Kind::Linear ||
1065 kind == StridedLinearBufferCopy::Kind::ReverseBoth) {
1073 constexpr IndexType run_stride = kind == StridedLinearBufferCopy::Kind::ReverseBoth ? -1 : 1;
1074 eigen_assert(src_stride == run_stride && dst_stride == run_stride);
1075 const IndexType run_offset = run_stride == 1 ? 0 : count - 1;
1076 const Scalar* run_src = src - run_offset;
1077 Scalar* run_dst = dst - run_offset;
1078 const IndexType unrolled_size = (4 * PacketSize) * (count / (4 * PacketSize));
1079 for (; i < unrolled_size; i += 4 * PacketSize) {
1080 for (
int j = 0; j < 4; ++j) {
1081 Packet p = ploadu<Packet>(run_src + i + j * PacketSize);
1082 pstoreu<Scalar, Packet>(run_dst + i + j * PacketSize, p);
1085 for (; i < vectorized_size; i += PacketSize) {
1086 Packet p = ploadu<Packet>(run_src + i);
1087 pstoreu<Scalar, Packet>(run_dst + i, p);
1089 EIGEN_IF_CONSTEXPR (HasHalfPacket) {
1090 const IndexType vectorized_half_size = HalfPacketSize * (count / HalfPacketSize);
1091 if (i < vectorized_half_size) {
1092 HalfPacket p = ploadu<HalfPacket>(run_src + i);
1093 pstoreu<Scalar, HalfPacket>(run_dst + i, p);
1094 i += HalfPacketSize;
1097 for (; i < count; ++i) {
1098 run_dst[i] = run_src[i];
1101 }
else EIGEN_IF_CONSTEXPR (kind == StridedLinearBufferCopy::Kind::Scatter) {
1103 eigen_assert(src_stride == 1 && dst_stride != 1);
1104 for (; i < vectorized_size; i += PacketSize) {
1105 Packet p = ploadu<Packet>(src + i);
1106 pscatter<Scalar, Packet>(dst + i * dst_stride, p, dst_stride);
1108 EIGEN_IF_CONSTEXPR (HasHalfPacket) {
1109 const IndexType vectorized_half_size = HalfPacketSize * (count / HalfPacketSize);
1110 if (i < vectorized_half_size) {
1111 HalfPacket p = ploadu<HalfPacket>(src + i);
1112 pscatter<Scalar, HalfPacket>(dst + i * dst_stride, p, dst_stride);
1113 i += HalfPacketSize;
1116 for (; i < count; ++i) {
1117 dst[i * dst_stride] = src[i];
1120 }
else EIGEN_IF_CONSTEXPR (kind == StridedLinearBufferCopy::Kind::FillLinear) {
1122 eigen_assert(src_stride == 0 && dst_stride == 1);
1124 const IndexType unrolled_size = (4 * PacketSize) * (count / (4 * PacketSize));
1126 Packet p = pset1<Packet>(s);
1127 for (; i < unrolled_size; i += 4 * PacketSize) {
1128 for (
int j = 0; j < 4; ++j) {
1129 pstoreu<Scalar, Packet>(dst + i + j * PacketSize, p);
1132 for (; i < vectorized_size; i += PacketSize) {
1133 pstoreu<Scalar, Packet>(dst + i, p);
1135 EIGEN_IF_CONSTEXPR (HasHalfPacket) {
1136 const IndexType vectorized_half_size = HalfPacketSize * (count / HalfPacketSize);
1137 if (i < vectorized_half_size) {
1138 HalfPacket hp = pset1<HalfPacket>(s);
1139 pstoreu<Scalar, HalfPacket>(dst + i, hp);
1140 i += HalfPacketSize;
1143 for (; i < count; ++i) {
1147 }
else EIGEN_IF_CONSTEXPR (kind == StridedLinearBufferCopy::Kind::FillScatter) {
1149 eigen_assert(src_stride == 0 && dst_stride != 1);
1151 Packet p = pset1<Packet>(s);
1152 for (; i < vectorized_size; i += PacketSize) {
1153 pscatter<Scalar, Packet>(dst + i * dst_stride, p, dst_stride);
1155 EIGEN_IF_CONSTEXPR (HasHalfPacket) {
1156 const IndexType vectorized_half_size = HalfPacketSize * (count / HalfPacketSize);
1157 if (i < vectorized_half_size) {
1158 HalfPacket hp = pset1<HalfPacket>(s);
1159 pscatter<Scalar, HalfPacket>(dst + i * dst_stride, hp, dst_stride);
1160 i += HalfPacketSize;
1163 for (; i < count; ++i) {
1164 dst[i * dst_stride] = s;
1167 }
else EIGEN_IF_CONSTEXPR (kind == StridedLinearBufferCopy::Kind::Gather) {
1169 eigen_assert(dst_stride == 1);
1170 for (; i < vectorized_size; i += PacketSize) {
1171 Packet p = pgather<Scalar, Packet>(src + i * src_stride, src_stride);
1172 pstoreu<Scalar, Packet>(dst + i, p);
1174 EIGEN_IF_CONSTEXPR (HasHalfPacket) {
1175 const IndexType vectorized_half_size = HalfPacketSize * (count / HalfPacketSize);
1176 if (i < vectorized_half_size) {
1177 HalfPacket p = pgather<Scalar, HalfPacket>(src + i * src_stride, src_stride);
1178 pstoreu<Scalar, HalfPacket>(dst + i, p);
1179 i += HalfPacketSize;
1182 for (; i < count; ++i) {
1183 dst[i] = src[i * src_stride];
1186 }
else EIGEN_IF_CONSTEXPR (kind == StridedLinearBufferCopy::Kind::ReverseStore) {
1192 eigen_assert(src_stride == 1 && dst_stride == -1);
1193 for (; i < vectorized_size; i += PacketSize) {
1194 Packet p = ploadu<Packet>(src + i);
1195 pstoreu<Scalar, Packet>(dst - i - (PacketSize - 1), preverse(p));
1197 EIGEN_IF_CONSTEXPR (HasHalfPacket) {
1198 const IndexType vectorized_half_size = HalfPacketSize * (count / HalfPacketSize);
1199 if (i < vectorized_half_size) {
1200 HalfPacket p = ploadu<HalfPacket>(src + i);
1201 pstoreu<Scalar, HalfPacket>(dst - i - (HalfPacketSize - 1), preverse(p));
1202 i += HalfPacketSize;
1205 for (; i < count; ++i) {
1209 }
else EIGEN_IF_CONSTEXPR (kind == StridedLinearBufferCopy::Kind::ReverseLoad) {
1211 eigen_assert(src_stride == -1 && dst_stride == 1);
1212 for (; i < vectorized_size; i += PacketSize) {
1213 Packet p = ploadu<Packet>(src - i - (PacketSize - 1));
1214 pstoreu<Scalar, Packet>(dst + i, preverse(p));
1216 EIGEN_IF_CONSTEXPR (HasHalfPacket) {
1217 const IndexType vectorized_half_size = HalfPacketSize * (count / HalfPacketSize);
1218 if (i < vectorized_half_size) {
1219 HalfPacket p = ploadu<HalfPacket>(src - i - (HalfPacketSize - 1));
1220 pstoreu<Scalar, HalfPacket>(dst + i, preverse(p));
1221 i += HalfPacketSize;
1224 for (; i < count; ++i) {
1228 }
else EIGEN_IF_CONSTEXPR (kind == StridedLinearBufferCopy::Kind::Random) {
1230 for (; i < count; ++i) {
1231 dst[i * dst_stride] = src[i * src_stride];
1234 eigen_assert(
false);
1250template <
typename Scalar,
typename IndexType,
int NumDims,
int Layout>
1251class TensorBlockIO {
1252 static constexpr bool IsColMajor = Layout ==
ColMajor;
1254 typedef StridedLinearBufferCopy<Scalar, IndexType> LinCopy;
1257 typedef DSizes<IndexType, NumDims> Dimensions;
1258 typedef DSizes<int, NumDims> DimensionsMap;
1261 Dst(
const Dimensions& dst_dims,
const Dimensions& dst_strides, Scalar* dst, IndexType dst_offset = 0)
1262 : dims(dst_dims), strides(dst_strides), data(dst), offset(dst_offset) {}
1271 Src(
const Dimensions& src_strides,
const Scalar* src, IndexType src_offset = 0)
1272 : strides(src_strides), data(src), offset(src_offset) {}
1284 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE IndexType Copy(
const Dst& dst,
const Src& src,
1285 const DimensionsMap& dst_to_src_dim_map) {
1287 EIGEN_IF_CONSTEXPR (NumDims == 0) {
1288 *(dst.data + dst.offset) = *(src.data + src.offset);
1293 const DimensionsMap& dim_map = dst_to_src_dim_map;
1296 int num_squeezable_dims = NumSqueezableInnerDims(dim_map);
1308 int num_size_one_inner_dims = 0;
1309 for (
int i = 0; i < num_squeezable_dims; ++i) {
1310 const int dst_dim = IsColMajor ? i : NumDims - i - 1;
1311 if (dst.dims[dst_dim] != 1)
break;
1312 num_size_one_inner_dims++;
1316 if (num_size_one_inner_dims == NumDims) {
1317 *(dst.data + dst.offset) = *(src.data + src.offset);
1323 const int dst_inner_dim = IsColMajor ? num_size_one_inner_dims : NumDims - num_size_one_inner_dims - 1;
1326 const int src_dim_for_dst_inner_dim = NumDims == 0 ? 1 : dim_map[dst_inner_dim];
1329 IndexType dst_inner_dim_size = NumDims == 0 ? 1 : dst.dims[dst_inner_dim];
1335 const IndexType output_stride = NumDims == 0 ? 1 : dst.strides[dst_inner_dim];
1336 const IndexType input_stride = NumDims == 0 ? 1 : src.strides[src_dim_for_dst_inner_dim];
1337 for (
int i = num_size_one_inner_dims + 1; i < num_squeezable_dims; ++i) {
1338 const int dst_dim = IsColMajor ? i : NumDims - i - 1;
1339 const IndexType dst_stride = dst.strides[dst_dim];
1340 const IndexType src_stride = src.strides[dim_map[dst_dim]];
1341 if (dst_stride == dst_inner_dim_size * output_stride && src_stride == dst_inner_dim_size * input_stride) {
1342 dst_inner_dim_size *= dst.dims[dst_dim];
1343 ++num_size_one_inner_dims;
1350 IndexType input_offset = src.offset;
1351 IndexType output_offset = dst.offset;
1353 constexpr int at_least_1_dim = NumDims <= 1 ? 1 : NumDims - 1;
1354 array<BlockIteratorState, at_least_1_dim> it;
1358 for (
int i = num_size_one_inner_dims; i < NumDims - 1; ++i) {
1359 const int dst_dim = IsColMajor ? i + 1 : NumDims - i - 2;
1360 if (dst.dims[dst_dim] == 1)
continue;
1362 it[idx].size = dst.dims[dst_dim];
1363 it[idx].input_stride = src.strides[dim_map[dst_dim]];
1364 it[idx].output_stride = dst.strides[dst_dim];
1366 it[idx].input_span = it[idx].input_stride * (it[idx].size - 1);
1367 it[idx].output_span = it[idx].output_stride * (it[idx].size - 1);
1373 const IndexType block_total_size = NumDims == 0 ? 1 : dst.dims.TotalSize();
1375#define COPY_INNER_DIM(KIND) \
1376 IndexType num_copied = 0; \
1377 for (num_copied = 0; num_copied < block_total_size; num_copied += dst_inner_dim_size) { \
1378 LinCopy::template Run<KIND>(typename LinCopy::Dst(output_offset, output_stride, dst.data), \
1379 typename LinCopy::Src(input_offset, input_stride, src.data), dst_inner_dim_size); \
1381 for (int j = 0; j < idx; ++j) { \
1382 if (++it[j].count < it[j].size) { \
1383 input_offset += it[j].input_stride; \
1384 output_offset += it[j].output_stride; \
1388 input_offset -= it[j].input_span; \
1389 output_offset -= it[j].output_span; \
1394 if (input_stride == 1 && output_stride == 1) {
1395 COPY_INNER_DIM(LinCopy::Kind::Linear);
1396 }
else if (input_stride == 1 && output_stride == -1) {
1397 COPY_INNER_DIM(LinCopy::Kind::ReverseStore);
1398 }
else if (input_stride == -1 && output_stride == 1) {
1399 COPY_INNER_DIM(LinCopy::Kind::ReverseLoad);
1400 }
else if (input_stride == -1 && output_stride == -1) {
1401 COPY_INNER_DIM(LinCopy::Kind::ReverseBoth);
1402 }
else if (input_stride == 1 && output_stride != 1) {
1403 COPY_INNER_DIM(LinCopy::Kind::Scatter);
1404 }
else if (input_stride == 0 && output_stride == 1) {
1405 COPY_INNER_DIM(LinCopy::Kind::FillLinear);
1406 }
else if (input_stride == 0 && output_stride != 1) {
1407 COPY_INNER_DIM(LinCopy::Kind::FillScatter);
1408 }
else if (output_stride == 1) {
1409 COPY_INNER_DIM(LinCopy::Kind::Gather);
1411 COPY_INNER_DIM(LinCopy::Kind::Random);
1414#undef COPY_INNER_DIM
1419 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE IndexType Copy(
const Dst& dst,
const Src& src) {
1420 DimensionsMap dst_to_src_map;
1421 for (
int i = 0; i < NumDims; ++i) dst_to_src_map[i] = i;
1422 return Copy(dst, src, dst_to_src_map);
1426 struct BlockIteratorState {
1427 BlockIteratorState() =
default;
1430 IndexType count = 0;
1431 IndexType input_stride = 0;
1432 IndexType output_stride = 0;
1433 IndexType input_span = 0;
1434 IndexType output_span = 0;
1440 static int NumSqueezableInnerDims(
const DimensionsMap& dim_map) {
1441 int num_squeezable_dims = 0;
1442 for (
int i = 0; i < NumDims; ++i) {
1443 const int dim = IsColMajor ? i : NumDims - i - 1;
1444 if (dim_map[dim] != dim)
break;
1445 num_squeezable_dims++;
1447 return num_squeezable_dims;
1454template <
typename XprType>
1455struct TensorBlockRead {
1456 static constexpr bool Supported =
false;
1457 using Expression = XprType;
1458 using Index =
typename XprType::Index;
1459 explicit TensorBlockRead(
const XprType& expression) : m_expression(expression) {}
1460 Index innerSize()
const {
return NumTraits<Index>::highest(); }
1461 const Expression& expr(Index, Index)
const {
return m_expression; }
1464 const XprType& m_expression;
1467template <
typename Scalar,
int NumDims,
int Layout,
typename Index>
1468struct TensorBlockRead<TensorBlockView<Scalar, NumDims, Layout, Index>> {
1469 static constexpr bool Supported =
true;
1470 using XprType = TensorBlockView<Scalar, NumDims, Layout, Index>;
1471 using Expression = TensorBlockView<Scalar, 1, Layout, Index>;
1472 explicit TensorBlockRead(
const XprType& expression) : m_evaluator(expression, m_device), m_inner_size(1) {
1473 for (
int i = 0; i < NumDims; ++i) {
1474 const int dim = Layout ==
ColMajor ? i : NumDims - 1 - i;
1475 if (expression.dimensions()[dim] > 1 && expression.strides()[dim] != m_inner_size)
break;
1476 m_inner_size *= expression.dimensions()[dim];
1479 Index innerSize()
const {
return m_inner_size; }
1480 Expression expr(Index offset, Index size)
const {
1481 return Expression(m_evaluator.coeffAddress(offset), DSizes<Index, 1>(size));
1485 DefaultDevice m_device;
1486 TensorEvaluator<const XprType, DefaultDevice> m_evaluator;
1491template <
typename Functor>
1492class TensorBlockReadFunctor {
1494 explicit TensorBlockReadFunctor(
const Functor& functor) : m_functor(&functor) {}
1495 template <
typename... Args>
1496 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
auto operator()(Args&&... args)
const
1497 ->
decltype(std::declval<const Functor&>()(std::forward<Args>(args)...)) {
1498 return (*m_functor)(std::forward<Args>(args)...);
1500 template <
typename... Args,
typename F = Functor>
1501 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
auto packetOp(Args&&... args)
const
1502 ->
decltype(std::declval<const F&>().packetOp(std::forward<Args>(args)...)) {
1503 return m_functor->packetOp(std::forward<Args>(args)...);
1507 const Functor* m_functor;
1510template <
typename Functor>
1511struct functor_traits<TensorBlockReadFunctor<Functor>> : functor_traits<Functor> {};
1513template <
typename UnaryOp,
typename Arg>
1514struct TensorBlockRead<TensorCwiseUnaryOp<UnaryOp, Arg>> {
1515 using ArgRead = TensorBlockRead<remove_all_t<Arg>>;
1516 static constexpr bool Supported = ArgRead::Supported;
1517 using XprType = TensorCwiseUnaryOp<UnaryOp, Arg>;
1518 using Expression = TensorCwiseUnaryOp<TensorBlockReadFunctor<UnaryOp>,
const typename ArgRead::Expression>;
1519 using Index =
typename XprType::Index;
1520 explicit TensorBlockRead(
const XprType& expression)
1521 : m_arg(expression.nestedExpression()), m_functor(expression.functor()) {}
1522 Index innerSize()
const {
return m_arg.innerSize(); }
1523 Expression expr(Index offset, Index size)
const {
1524 return Expression(m_arg.expr(offset, size), TensorBlockReadFunctor<UnaryOp>(m_functor));
1532template <
typename BinaryOp,
typename Left,
typename Right>
1533struct TensorBlockRead<TensorCwiseBinaryOp<BinaryOp, Left, Right>> {
1534 using LeftRead = TensorBlockRead<remove_all_t<Left>>;
1535 using RightRead = TensorBlockRead<remove_all_t<Right>>;
1536 static constexpr bool Supported = LeftRead::Supported && RightRead::Supported;
1537 using XprType = TensorCwiseBinaryOp<BinaryOp, Left, Right>;
1538 using Expression = TensorCwiseBinaryOp<TensorBlockReadFunctor<BinaryOp>,
const typename LeftRead::Expression,
1539 const typename RightRead::Expression>;
1540 using Index =
typename XprType::Index;
1541 explicit TensorBlockRead(
const XprType& expression)
1542 : m_left(expression.lhsExpression()), m_right(expression.rhsExpression()), m_functor(expression.functor()) {}
1543 Index innerSize()
const {
return numext::mini(m_left.innerSize(), m_right.innerSize()); }
1544 Expression expr(Index offset, Index size)
const {
1545 return Expression(m_left.expr(offset, size), m_right.expr(offset, size),
1546 TensorBlockReadFunctor<BinaryOp>(m_functor));
1555template <
typename NullaryOp,
typename Arg>
1556struct TensorBlockRead<TensorCwiseNullaryOp<NullaryOp, Arg>> {
1557 using XprType = TensorCwiseNullaryOp<NullaryOp, Arg>;
1558 using Evaluator = TensorEvaluator<const XprType, DefaultDevice>;
1559 static constexpr bool Supported = Evaluator::IndexIndependentFunctor;
1560 using Index =
typename XprType::Index;
1561 using RunView = TensorBlockView<typename XprType::Scalar, 1, traits<XprType>::Layout, Index>;
1562 using Expression = TensorCwiseNullaryOp<NullaryOp, const RunView>;
1563 explicit TensorBlockRead(
const XprType& expression) : m_functor(expression.functor()) {}
1564 Index innerSize()
const {
return NumTraits<Index>::highest(); }
1565 Expression expr(Index, Index size)
const {
return Expression(RunView(
nullptr, DSizes<Index, 1>(size)), m_functor); }
1568 NullaryOp m_functor;
1571template <
typename TernaryOp,
typename Arg1,
typename Arg2,
typename Arg3>
1572struct TensorBlockRead<TensorCwiseTernaryOp<TernaryOp, Arg1, Arg2, Arg3>> {
1573 using Arg1Read = TensorBlockRead<remove_all_t<Arg1>>;
1574 using Arg2Read = TensorBlockRead<remove_all_t<Arg2>>;
1575 using Arg3Read = TensorBlockRead<remove_all_t<Arg3>>;
1576 static constexpr bool Supported = Arg1Read::Supported && Arg2Read::Supported && Arg3Read::Supported;
1577 using XprType = TensorCwiseTernaryOp<TernaryOp, Arg1, Arg2, Arg3>;
1578 using Expression = TensorCwiseTernaryOp<TensorBlockReadFunctor<TernaryOp>,
const typename Arg1Read::Expression,
1579 const typename Arg2Read::Expression,
const typename Arg3Read::Expression>;
1580 using Index =
typename XprType::Index;
1581 explicit TensorBlockRead(
const XprType& expression)
1582 : m_arg1(expression.arg1Expression()),
1583 m_arg2(expression.arg2Expression()),
1584 m_arg3(expression.arg3Expression()),
1585 m_functor(expression.functor()) {}
1586 Index innerSize()
const {
1587 return numext::mini(m_arg1.innerSize(), numext::mini(m_arg2.innerSize(), m_arg3.innerSize()));
1589 Expression expr(Index offset, Index size)
const {
1590 return Expression(m_arg1.expr(offset, size), m_arg2.expr(offset, size), m_arg3.expr(offset, size),
1591 TensorBlockReadFunctor<TernaryOp>(m_functor));
1598 TernaryOp m_functor;
1601template <
typename Cond,
typename Then,
typename Else>
1602struct TensorBlockRead<TensorSelectOp<Cond, Then, Else>> {
1603 using CondRead = TensorBlockRead<remove_all_t<Cond>>;
1604 using ThenRead = TensorBlockRead<remove_all_t<Then>>;
1605 using ElseRead = TensorBlockRead<remove_all_t<Else>>;
1606 static constexpr bool Supported = CondRead::Supported && ThenRead::Supported && ElseRead::Supported;
1607 using XprType = TensorSelectOp<Cond, Then, Else>;
1608 using Expression = TensorSelectOp<
const typename CondRead::Expression,
const typename ThenRead::Expression,
1609 const typename ElseRead::Expression>;
1610 using Index =
typename XprType::Index;
1611 explicit TensorBlockRead(
const XprType& expression)
1612 : m_cond(expression.ifExpression()), m_then(expression.thenExpression()), m_else(expression.elseExpression()) {}
1613 Index innerSize()
const {
1614 return numext::mini(m_cond.innerSize(), numext::mini(m_then.innerSize(), m_else.innerSize()));
1616 Expression expr(Index offset, Index size)
const {
1617 return Expression(m_cond.expr(offset, size), m_then.expr(offset, size), m_else.expr(offset, size));
1626template <
typename Scalar,
typename Arg>
1627struct TensorBlockRead<TensorConversionOp<Scalar, Arg>> {
1628 using ArgRead = TensorBlockRead<remove_all_t<Arg>>;
1629 static constexpr bool Supported = ArgRead::Supported;
1630 using XprType = TensorConversionOp<Scalar, Arg>;
1631 using Expression = TensorConversionOp<Scalar, const typename ArgRead::Expression>;
1632 using Index =
typename XprType::Index;
1633 explicit TensorBlockRead(
const XprType& expression) : m_arg(expression.expression()) {}
1634 Index innerSize()
const {
return m_arg.innerSize(); }
1635 Expression expr(Index offset, Index size)
const {
return Expression(m_arg.expr(offset, size)); }
1659template <
typename Scalar,
int NumDims,
typename TensorBlockExpr,
typename IndexType = Eigen::Index>
1660class TensorBlockAssignment {
1662 typedef TensorEvaluator<const TensorBlockExpr, DefaultDevice> TensorBlockEvaluator;
1664 typedef DSizes<IndexType, NumDims> Dimensions;
1665 using BlockRead = TensorBlockRead<TensorBlockExpr>;
1667 enum { Vectorizable = packet_traits<Scalar>::Vectorizable, PacketSize = packet_traits<Scalar>::size };
1669 template <
bool Vectorizable,
typename Evaluator>
1670 struct InnerDimAssign {
1671 EIGEN_ALWAYS_INLINE
static void Run(Scalar* target, IndexType count,
const Evaluator& eval, IndexType eval_offset) {
1672 for (IndexType i = 0; i < count; ++i) {
1673 target[i] = eval.coeff(eval_offset + i);
1678 template <
typename Evaluator>
1679 struct InnerDimAssign<true, Evaluator> {
1680 EIGEN_ALWAYS_INLINE
static void Run(Scalar* target, IndexType count,
const Evaluator& eval, IndexType eval_offset) {
1681 typedef typename packet_traits<Scalar>::type Packet;
1683 const IndexType unrolled_size = (4 * PacketSize) * (count / (4 * PacketSize));
1684 const IndexType vectorized_size = PacketSize * (count / PacketSize);
1687 for (; i < unrolled_size; i += 4 * PacketSize) {
1688 for (
int j = 0; j < 4; ++j) {
1689 const IndexType idx = eval_offset + i + j * PacketSize;
1690 Packet p = eval.template packet<Unaligned>(idx);
1691 pstoreu<Scalar>(target + i + j * PacketSize, p);
1695 for (; i < vectorized_size; i += PacketSize) {
1696 Packet p = eval.template packet<Unaligned>(eval_offset + i);
1697 pstoreu<Scalar>(target + i, p);
1700 for (; i < count; ++i) {
1701 target[i] = eval.coeff(eval_offset + i);
1706 template <
typename Evaluator>
1707 static EIGEN_STRONG_INLINE
void AssignInner(Scalar* target, IndexType count,
const Evaluator& eval,
const BlockRead&,
1708 IndexType offset, std::false_type) {
1709 InnerDimAssign<Vectorizable && Evaluator::PacketAccess, Evaluator>::Run(target, count, eval, offset);
1712 static EIGEN_STRONG_INLINE
void AssignInner(Scalar* target, IndexType count,
const TensorBlockEvaluator&,
1713 const BlockRead& reader, IndexType offset, std::true_type) {
1714 const auto expression = reader.expr(offset, count);
1715 using RunEvaluator = TensorEvaluator<const typename BlockRead::Expression, DefaultDevice>;
1716 const DefaultDevice device;
1717 const RunEvaluator eval(expression, device);
1718 InnerDimAssign<Vectorizable && RunEvaluator::PacketAccess, RunEvaluator>::Run(target, count, eval, 0);
1723 Target(
const Dimensions& target_dims,
const Dimensions& target_strides, Scalar* target_data,
1724 IndexType target_offset = 0)
1725 : dims(target_dims), strides(target_strides), data(target_data), offset(target_offset) {}
1733 static Target target(
const Dimensions& target_dims,
const Dimensions& target_strides, Scalar* target_data,
1734 IndexType target_offset = 0) {
1735 return Target(target_dims, target_strides, target_data, target_offset);
1738 template <
typename TargetDimsIndexType,
typename TargetStr
idesIndexType>
1739 static Target target(
const DSizes<TargetDimsIndexType, NumDims>& target_dims,
1740 const DSizes<TargetStridesIndexType, NumDims>& target_strides, Scalar* target_data,
1741 IndexType target_offset = 0) {
1743 return Target(Dimensions(target_dims), Dimensions(target_strides), target_data, target_offset);
1746 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void Run(
const Target& target,
const TensorBlockExpr& expr) {
1748 DefaultDevice default_device;
1749 TensorBlockEvaluator eval(expr, default_device);
1750 const BlockRead reader(expr);
1753 eigen_assert(dimensions_match(target.dims, eval.dimensions()));
1755 Run(target, eval, reader, bool_constant<BlockRead::Supported>());
1759 static void Run(
const Target& target,
const TensorBlockEvaluator& eval,
const BlockRead& reader, std::false_type) {
1760 RunImpl<false>(target, eval, reader);
1763 static void Run(
const Target& target,
const TensorBlockEvaluator& eval,
const BlockRead& reader, std::true_type) {
1764 const IndexType size = NumDims == 0 ? 1 : target.dims.TotalSize();
1765 if (reader.innerSize() >= size) {
1767 const auto expression = reader.expr(0, size);
1768 using RunEvaluator = TensorEvaluator<const typename BlockRead::Expression, DefaultDevice>;
1769 const DefaultDevice device;
1770 const RunEvaluator dense_eval(expression, device);
1771 RunImpl<false>(target, dense_eval, reader);
1773 RunImpl<true>(target, eval, reader);
1777 template <
bool UseInnerRuns,
typename Evaluator>
1778 static EIGEN_STRONG_INLINE
void RunImpl(
const Target& target,
const Evaluator& eval,
const BlockRead& reader) {
1779 static constexpr int Layout = Evaluator::Layout;
1780 static constexpr bool is_col_major = Layout ==
ColMajor;
1783 const IndexType output_size = NumDims == 0 ? 1 : target.dims.TotalSize();
1784 constexpr int inner_dim_idx = NumDims == 0 ? 0 : (is_col_major ? 0 : NumDims - 1);
1785 IndexType output_inner_dim_size = NumDims == 0 ? 1 : target.dims[inner_dim_idx];
1788 EIGEN_IF_CONSTEXPR (NumDims > 0) {
1789 eigen_assert(target.strides[inner_dim_idx] == 1);
1793 IndexType num_squeezed_dims = 0;
1794 for (Index i = 1; i < NumDims; ++i) {
1795 const Index dim = is_col_major ? i : NumDims - i - 1;
1796 const IndexType target_stride = target.strides[dim];
1798 if (output_inner_dim_size == target_stride &&
1799 (!UseInnerRuns || output_inner_dim_size * target.dims[dim] <= reader.innerSize())) {
1800 output_inner_dim_size *= target.dims[dim];
1801 num_squeezed_dims++;
1809 array<BlockIteratorState, NumDims> it;
1812 for (Index i = num_squeezed_dims; i < NumDims - 1; ++i) {
1813 const Index dim = is_col_major ? i + 1 : NumDims - i - 2;
1816 it[idx].size = target.dims[dim];
1817 it[idx].output_stride = target.strides[dim];
1818 it[idx].output_span = it[idx].output_stride * (it[idx].size - 1);
1824 IndexType input_offset = 0;
1825 IndexType output_offset = target.offset;
1828 for (IndexType i = 0; i < output_size; i += output_inner_dim_size) {
1830 AssignInner(target.data + output_offset, output_inner_dim_size, eval, reader, input_offset,
1831 bool_constant<UseInnerRuns>());
1834 input_offset += output_inner_dim_size;
1837 for (
int j = 0; j < idx; ++j) {
1838 if (++it[j].count < it[j].size) {
1839 output_offset += it[j].output_stride;
1843 output_offset -= it[j].output_span;
1849 struct BlockIteratorState {
1850 BlockIteratorState() =
default;
1852 IndexType count = 0;
1854 IndexType output_stride = 0;
1855 IndexType output_span = 0;
Namespace containing all symbols from the Eigen library.
The tensor evaluator class.
Definition TensorEvaluator.h:47