11#ifndef EIGEN_TENSOR_TENSOR_PATCH_H
12#define EIGEN_TENSOR_TENSOR_PATCH_H
15#include "./InternalHeaderCheck.h"
20template <
typename PatchDim,
typename XprType>
21struct traits<TensorPatchOp<PatchDim, XprType> > : traits<XprType> {
22 typedef typename XprType::Scalar Scalar;
23 typedef traits<XprType> XprTraits;
24 typedef typename XprTraits::StorageKind StorageKind;
25 typedef typename XprTraits::Index Index;
26 static constexpr int NumDimensions = XprTraits::NumDimensions + 1;
27 static constexpr int Layout = XprTraits::Layout;
28 typedef typename XprTraits::PointerType PointerType;
31template <
typename PatchDim,
typename XprType>
32struct eval<TensorPatchOp<PatchDim, XprType>, Eigen::Dense> {
33 typedef const TensorPatchOp<PatchDim, XprType>& type;
43template <
typename PatchDim,
typename XprType>
44class TensorPatchOp :
public TensorBase<TensorPatchOp<PatchDim, XprType>, ReadOnlyAccessors> {
46 typedef typename Eigen::internal::traits<TensorPatchOp>::Scalar Scalar;
48 typedef typename XprType::CoeffReturnType CoeffReturnType;
49 typedef typename Eigen::internal::ref_selector<TensorPatchOp>::type Nested;
50 typedef typename Eigen::internal::traits<TensorPatchOp>::StorageKind StorageKind;
51 typedef typename Eigen::internal::traits<TensorPatchOp>::Index Index;
53 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorPatchOp(
const XprType& expr,
const PatchDim& patch_dims)
54 : m_xpr(expr), m_patch_dims(patch_dims) {}
56 EIGEN_DEVICE_FUNC
const PatchDim& patch_dims()
const {
return m_patch_dims; }
58 EIGEN_DEVICE_FUNC
const internal::remove_all_t<typename XprType::Nested>& expression()
const {
return m_xpr; }
61 typename XprType::Nested m_xpr;
62 const PatchDim m_patch_dims;
66template <
typename PatchDim,
typename ArgType,
typename Device>
69 typedef typename XprType::Index Index;
70 static constexpr int NumDims = internal::array_size<typename TensorEvaluator<ArgType, Device>::Dimensions>::value + 1;
72 typedef typename XprType::Scalar Scalar;
74 typedef typename PacketType<CoeffReturnType, Device>::type PacketReturnType;
75 static constexpr int PacketSize = PacketType<CoeffReturnType, Device>::size;
76 typedef StorageMemory<CoeffReturnType, Device> Storage;
77 typedef typename Storage::Type EvaluatorPointerType;
79 static constexpr int Layout = TensorEvaluator<ArgType, Device>::Layout;
82 PacketAccess = TensorEvaluator<ArgType, Device>::PacketAccess,
87 BlockAccess = NumDims > 1,
90 PreferBlockAccess =
true,
96 typedef internal::TensorBlockDescriptor<NumDims, Index> TensorBlockDesc;
97 typedef internal::TensorBlockScratchAllocator<Device> TensorBlockScratch;
98 typedef typename internal::TensorMaterializedBlock<CoeffReturnType, NumDims, Layout, Index> TensorBlock;
101 EIGEN_STRONG_INLINE TensorEvaluator(
const XprType& op,
const Device& device)
102 : m_impl(op.expression(), device), m_device(device) {
103 Index num_patches = 1;
104 const typename TensorEvaluator<ArgType, Device>::Dimensions& input_dims = m_impl.dimensions();
105 const PatchDim& patch_dims = op.patch_dims();
106 EIGEN_IF_CONSTEXPR (
static_cast<int>(Layout) ==
static_cast<int>(
ColMajor)) {
107 for (
int i = 0; i < NumDims - 1; ++i) {
108 m_dimensions[i] = patch_dims[i];
109 num_patches *= (input_dims[i] - patch_dims[i] + 1);
111 m_dimensions[NumDims - 1] = num_patches;
113 m_inputStrides[0] = 1;
114 m_patchStrides[0] = 1;
115 for (
int i = 1; i < NumDims - 1; ++i) {
116 m_inputStrides[i] = m_inputStrides[i - 1] * input_dims[i - 1];
117 m_patchStrides[i] = m_patchStrides[i - 1] * (input_dims[i - 1] - patch_dims[i - 1] + 1);
119 m_outputStrides[0] = 1;
120 for (
int i = 1; i < NumDims; ++i) {
121 m_outputStrides[i] = m_outputStrides[i - 1] * m_dimensions[i - 1];
124 for (
int i = 0; i < NumDims - 1; ++i) {
125 m_dimensions[i + 1] = patch_dims[i];
126 num_patches *= (input_dims[i] - patch_dims[i] + 1);
128 m_dimensions[0] = num_patches;
130 m_inputStrides[NumDims - 2] = 1;
131 m_patchStrides[NumDims - 2] = 1;
132 for (
int i = NumDims - 3; i >= 0; --i) {
133 m_inputStrides[i] = m_inputStrides[i + 1] * input_dims[i + 1];
134 m_patchStrides[i] = m_patchStrides[i + 1] * (input_dims[i + 1] - patch_dims[i + 1] + 1);
136 m_outputStrides[NumDims - 1] = 1;
137 for (
int i = NumDims - 2; i >= 0; --i) {
138 m_outputStrides[i] = m_outputStrides[i + 1] * m_dimensions[i + 1];
143 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
const Dimensions& dimensions()
const {
return m_dimensions; }
145 EIGEN_STRONG_INLINE
bool evalSubExprsIfNeeded(EvaluatorPointerType ) {
146 m_impl.evalSubExprsIfNeeded(
nullptr);
150 EIGEN_STRONG_INLINE
void cleanup() { m_impl.cleanup(); }
152 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffReturnType coeff(Index index)
const {
153 Index output_stride_index = (
static_cast<int>(Layout) ==
static_cast<int>(
ColMajor)) ? NumDims - 1 : 0;
155 Index patchIndex = index / m_outputStrides[output_stride_index];
157 Index patchOffset = index - patchIndex * m_outputStrides[output_stride_index];
158 Index inputIndex = 0;
159 EIGEN_IF_CONSTEXPR (
static_cast<int>(Layout) ==
static_cast<int>(
ColMajor)) {
161 for (
int i = NumDims - 2; i > 0; --i) {
162 const Index patchIdx = patchIndex / m_patchStrides[i];
163 patchIndex -= patchIdx * m_patchStrides[i];
164 const Index offsetIdx = patchOffset / m_outputStrides[i];
165 patchOffset -= offsetIdx * m_outputStrides[i];
166 inputIndex += (patchIdx + offsetIdx) * m_inputStrides[i];
170 for (
int i = 0; i < NumDims - 2; ++i) {
171 const Index patchIdx = patchIndex / m_patchStrides[i];
172 patchIndex -= patchIdx * m_patchStrides[i];
173 const Index offsetIdx = patchOffset / m_outputStrides[i + 1];
174 patchOffset -= offsetIdx * m_outputStrides[i + 1];
175 inputIndex += (patchIdx + offsetIdx) * m_inputStrides[i];
178 inputIndex += (patchIndex + patchOffset);
179 return m_impl.coeff(inputIndex);
182 template <
int LoadMode>
183 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketReturnType packet(Index index)
const {
184 eigen_assert(index + PacketSize - 1 < dimensions().TotalSize());
186 Index output_stride_index = (
static_cast<int>(Layout) ==
static_cast<int>(
ColMajor)) ? NumDims - 1 : 0;
187 Index indices[2] = {index, index + PacketSize - 1};
188 Index patchIndices[2] = {indices[0] / m_outputStrides[output_stride_index],
189 indices[1] / m_outputStrides[output_stride_index]};
190 Index patchOffsets[2] = {indices[0] - patchIndices[0] * m_outputStrides[output_stride_index],
191 indices[1] - patchIndices[1] * m_outputStrides[output_stride_index]};
193 Index inputIndices[2] = {0, 0};
194 EIGEN_IF_CONSTEXPR (
static_cast<int>(Layout) ==
static_cast<int>(
ColMajor)) {
196 for (
int i = NumDims - 2; i > 0; --i) {
197 const Index patchIdx[2] = {patchIndices[0] / m_patchStrides[i], patchIndices[1] / m_patchStrides[i]};
198 patchIndices[0] -= patchIdx[0] * m_patchStrides[i];
199 patchIndices[1] -= patchIdx[1] * m_patchStrides[i];
201 const Index offsetIdx[2] = {patchOffsets[0] / m_outputStrides[i], patchOffsets[1] / m_outputStrides[i]};
202 patchOffsets[0] -= offsetIdx[0] * m_outputStrides[i];
203 patchOffsets[1] -= offsetIdx[1] * m_outputStrides[i];
205 inputIndices[0] += (patchIdx[0] + offsetIdx[0]) * m_inputStrides[i];
206 inputIndices[1] += (patchIdx[1] + offsetIdx[1]) * m_inputStrides[i];
210 for (
int i = 0; i < NumDims - 2; ++i) {
211 const Index patchIdx[2] = {patchIndices[0] / m_patchStrides[i], patchIndices[1] / m_patchStrides[i]};
212 patchIndices[0] -= patchIdx[0] * m_patchStrides[i];
213 patchIndices[1] -= patchIdx[1] * m_patchStrides[i];
215 const Index offsetIdx[2] = {patchOffsets[0] / m_outputStrides[i + 1], patchOffsets[1] / m_outputStrides[i + 1]};
216 patchOffsets[0] -= offsetIdx[0] * m_outputStrides[i + 1];
217 patchOffsets[1] -= offsetIdx[1] * m_outputStrides[i + 1];
219 inputIndices[0] += (patchIdx[0] + offsetIdx[0]) * m_inputStrides[i];
220 inputIndices[1] += (patchIdx[1] + offsetIdx[1]) * m_inputStrides[i];
223 inputIndices[0] += (patchIndices[0] + patchOffsets[0]);
224 inputIndices[1] += (patchIndices[1] + patchOffsets[1]);
226 if (inputIndices[1] - inputIndices[0] == PacketSize - 1) {
227 PacketReturnType rslt = m_impl.template packet<Unaligned>(inputIndices[0]);
230 EIGEN_ALIGN_TO_BOUNDARY(internal::unpacket_traits<PacketReturnType>::alignment)
231 CoeffReturnType values[PacketSize];
232 values[0] = m_impl.coeff(inputIndices[0]);
233 values[PacketSize - 1] = m_impl.coeff(inputIndices[1]);
235 for (
int i = 1; i < PacketSize - 1; ++i) {
236 values[i] = coeff(index + i);
238 PacketReturnType rslt = internal::pload<PacketReturnType>(values);
243 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE internal::TensorBlockResourceRequirements getResourceRequirements()
const {
244 const size_t target_size = m_device.firstLevelCacheSize();
249 const TensorOpCost cost_per_coeff =
250 m_impl.costPerCoeff(
false) + TensorOpCost(0,
sizeof(CoeffReturnType), 0);
251 return internal::TensorBlockResourceRequirements::withShapeAndSize<Scalar>(
252 internal::TensorBlockShapeType::kSkewedInnerDims, target_size, cost_per_coeff);
255 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorBlock block(TensorBlockDesc& desc, TensorBlockScratch& scratch,
256 bool =
false)
const {
257 constexpr bool is_col_major =
static_cast<int>(Layout) ==
static_cast<int>(
ColMajor);
259 if (desc.size() == 0) {
260 return TensorBlock(internal::TensorBlockKind::kView,
nullptr, desc.dimensions());
263 typename TensorBlock::Storage block_storage = TensorBlock::prepareStorage(desc, scratch);
264 CoeffReturnType* block_buffer = block_storage.data();
267 array<Index, NumDims> coords;
268 Index remaining = desc.offset();
269 EIGEN_IF_CONSTEXPR (is_col_major) {
270 for (
int i = NumDims - 1; i > 0; --i) {
271 coords[i] = remaining / m_outputStrides[i];
272 remaining -= coords[i] * m_outputStrides[i];
274 coords[0] = remaining;
276 for (
int i = 0; i < NumDims - 1; ++i) {
277 coords[i] = remaining / m_outputStrides[i];
278 remaining -= coords[i] * m_outputStrides[i];
280 coords[NumDims - 1] = remaining;
283 const int patch_dim = is_col_major ? NumDims - 1 : 0;
284 const int inner_dim = is_col_major ? 0 : NumDims - 1;
286 const Index num_patches_in_block = desc.dimension(patch_dim);
289 const Index inner_size = desc.dimension(inner_dim);
294 array<Index, NumDims> mid_sizes;
295 array<Index, NumDims> mid_src_stride;
296 array<Index, NumDims> mid_count;
298 Index src_corner = 0;
299 for (
int k = 0; k < NumDims - 1; ++k) {
300 const int d = is_col_major ? k : NumDims - 1 - k;
301 const int in_d = is_col_major ? d : d - 1;
302 src_corner += coords[d] * m_inputStrides[in_d];
304 mid_sizes[num_mid] = desc.dimension(d);
305 mid_src_stride[num_mid] = m_inputStrides[in_d];
314 for (Index p = 0; p < num_patches_in_block; ++p) {
316 Index patch_index = coords[patch_dim] + p;
318 EIGEN_IF_CONSTEXPR (is_col_major) {
319 for (
int i = NumDims - 2; i > 0; --i) {
320 const Index idx = patch_index / m_patchStrides[i];
321 patch_index -= idx * m_patchStrides[i];
322 src_patch += idx * m_inputStrides[i];
325 for (
int i = 0; i < NumDims - 2; ++i) {
326 const Index idx = patch_index / m_patchStrides[i];
327 patch_index -= idx * m_patchStrides[i];
328 src_patch += idx * m_inputStrides[i];
331 src_patch += patch_index;
333 Index src = src_patch + src_corner;
334 for (
int k = 0; k < num_mid; ++k) mid_count[k] = 0;
336 for (Index j = 0; j < inner_size; ++j) {
337 block_buffer[dst + j] = m_impl.coeff(src + j);
341 for (; k < num_mid; ++k) {
342 if (++mid_count[k] < mid_sizes[k]) {
343 src += mid_src_stride[k];
347 src -= mid_src_stride[k] * (mid_sizes[k] - 1);
349 if (k == num_mid)
break;
352 eigen_assert(dst == desc.size());
354 return block_storage.AsTensorMaterializedBlock();
357 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorOpCost costPerCoeff(
bool vectorized)
const {
358 const double compute_cost = NumDims * (TensorOpCost::DivCost<Index>() + TensorOpCost::MulCost<Index>() +
359 2 * TensorOpCost::AddCost<Index>());
360 return m_impl.costPerCoeff(vectorized) + TensorOpCost(0, 0, compute_cost, vectorized, PacketSize);
363 EIGEN_DEVICE_FUNC EvaluatorPointerType data()
const {
return nullptr; }
366 Dimensions m_dimensions;
367 array<Index, NumDims> m_outputStrides;
368 array<Index, NumDims - 1> m_inputStrides;
369 array<Index, NumDims - 1> m_patchStrides;
371 TensorEvaluator<ArgType, Device> m_impl;
372 const Device EIGEN_DEVICE_REF m_device;
The tensor base class.
Definition TensorForwardDeclarations.h:69
Tensor patch class.
Definition TensorPatch.h:44
Namespace containing all symbols from the Eigen library.
The tensor evaluator class.
Definition TensorEvaluator.h:47