11#ifndef EIGEN_TENSOR_TENSOR_IMAGE_PATCH_H
12#define EIGEN_TENSOR_TENSOR_IMAGE_PATCH_H
15#include "./InternalHeaderCheck.h"
21template <DenseIndex Rows, DenseIndex Cols,
typename XprType>
22struct traits<TensorImagePatchOp<Rows, Cols, XprType>> : traits<XprType> {
23 typedef std::remove_const_t<typename XprType::Scalar> Scalar;
24 typedef traits<XprType> XprTraits;
25 typedef typename XprTraits::StorageKind StorageKind;
26 typedef typename XprTraits::Index Index;
27 static constexpr int NumDimensions = XprTraits::NumDimensions + 1;
28 static constexpr int Layout = XprTraits::Layout;
29 typedef typename XprTraits::PointerType PointerType;
32template <DenseIndex Rows, DenseIndex Cols,
typename XprType>
33struct eval<TensorImagePatchOp<Rows, Cols, XprType>, Eigen::Dense> {
34 typedef const TensorImagePatchOp<Rows, Cols, XprType>& type;
53template <DenseIndex Rows, DenseIndex Cols,
typename XprType>
54class TensorImagePatchOp :
public TensorBase<TensorImagePatchOp<Rows, Cols, XprType>, ReadOnlyAccessors> {
56 typedef typename Eigen::internal::traits<TensorImagePatchOp>::Scalar Scalar;
58 typedef typename XprType::CoeffReturnType CoeffReturnType;
59 typedef typename Eigen::internal::ref_selector<TensorImagePatchOp>::type Nested;
60 typedef typename Eigen::internal::traits<TensorImagePatchOp>::StorageKind StorageKind;
61 typedef typename Eigen::internal::traits<TensorImagePatchOp>::Index Index;
63 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorImagePatchOp(
const XprType& expr, DenseIndex patch_rows,
64 DenseIndex patch_cols, DenseIndex row_strides,
65 DenseIndex col_strides, DenseIndex in_row_strides,
66 DenseIndex in_col_strides, DenseIndex row_inflate_strides,
67 DenseIndex col_inflate_strides, PaddingType padding_type,
70 m_patch_rows(patch_rows),
71 m_patch_cols(patch_cols),
72 m_row_strides(row_strides),
73 m_col_strides(col_strides),
74 m_in_row_strides(in_row_strides),
75 m_in_col_strides(in_col_strides),
76 m_row_inflate_strides(row_inflate_strides),
77 m_col_inflate_strides(col_inflate_strides),
78 m_padding_explicit(
false),
83 m_padding_type(padding_type),
84 m_padding_value(padding_value) {}
86 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorImagePatchOp(
const XprType& expr, DenseIndex patch_rows,
87 DenseIndex patch_cols, DenseIndex row_strides,
88 DenseIndex col_strides, DenseIndex in_row_strides,
89 DenseIndex in_col_strides, DenseIndex row_inflate_strides,
90 DenseIndex col_inflate_strides, DenseIndex padding_top,
91 DenseIndex padding_bottom, DenseIndex padding_left,
92 DenseIndex padding_right, Scalar padding_value)
94 m_patch_rows(patch_rows),
95 m_patch_cols(patch_cols),
96 m_row_strides(row_strides),
97 m_col_strides(col_strides),
98 m_in_row_strides(in_row_strides),
99 m_in_col_strides(in_col_strides),
100 m_row_inflate_strides(row_inflate_strides),
101 m_col_inflate_strides(col_inflate_strides),
102 m_padding_explicit(
true),
103 m_padding_top(padding_top),
104 m_padding_bottom(padding_bottom),
105 m_padding_left(padding_left),
106 m_padding_right(padding_right),
107 m_padding_type(PADDING_VALID),
108 m_padding_value(padding_value) {}
110 EIGEN_DEVICE_FUNC DenseIndex patch_rows()
const {
return m_patch_rows; }
111 EIGEN_DEVICE_FUNC DenseIndex patch_cols()
const {
return m_patch_cols; }
112 EIGEN_DEVICE_FUNC DenseIndex row_strides()
const {
return m_row_strides; }
113 EIGEN_DEVICE_FUNC DenseIndex col_strides()
const {
return m_col_strides; }
114 EIGEN_DEVICE_FUNC DenseIndex in_row_strides()
const {
return m_in_row_strides; }
115 EIGEN_DEVICE_FUNC DenseIndex in_col_strides()
const {
return m_in_col_strides; }
116 EIGEN_DEVICE_FUNC DenseIndex row_inflate_strides()
const {
return m_row_inflate_strides; }
117 EIGEN_DEVICE_FUNC DenseIndex col_inflate_strides()
const {
return m_col_inflate_strides; }
118 EIGEN_DEVICE_FUNC
bool padding_explicit()
const {
return m_padding_explicit; }
119 EIGEN_DEVICE_FUNC DenseIndex padding_top()
const {
return m_padding_top; }
120 EIGEN_DEVICE_FUNC DenseIndex padding_bottom()
const {
return m_padding_bottom; }
121 EIGEN_DEVICE_FUNC DenseIndex padding_left()
const {
return m_padding_left; }
122 EIGEN_DEVICE_FUNC DenseIndex padding_right()
const {
return m_padding_right; }
123 EIGEN_DEVICE_FUNC PaddingType padding_type()
const {
return m_padding_type; }
124 EIGEN_DEVICE_FUNC Scalar padding_value()
const {
return m_padding_value; }
126 EIGEN_DEVICE_FUNC
const internal::remove_all_t<typename XprType::Nested>& expression()
const {
return m_xpr; }
129 typename XprType::Nested m_xpr;
130 const DenseIndex m_patch_rows;
131 const DenseIndex m_patch_cols;
132 const DenseIndex m_row_strides;
133 const DenseIndex m_col_strides;
134 const DenseIndex m_in_row_strides;
135 const DenseIndex m_in_col_strides;
136 const DenseIndex m_row_inflate_strides;
137 const DenseIndex m_col_inflate_strides;
138 const bool m_padding_explicit;
139 const DenseIndex m_padding_top;
140 const DenseIndex m_padding_bottom;
141 const DenseIndex m_padding_left;
142 const DenseIndex m_padding_right;
143 const PaddingType m_padding_type;
144 const Scalar m_padding_value;
148template <DenseIndex Rows, DenseIndex Cols,
typename ArgType,
typename Device>
151 typedef typename XprType::Index Index;
152 static constexpr int NumInputDims =
153 internal::array_size<typename TensorEvaluator<ArgType, Device>::Dimensions>::value;
154 static constexpr int NumDims = NumInputDims + 1;
156 typedef std::remove_const_t<typename XprType::Scalar> Scalar;
160 typedef typename PacketType<CoeffReturnType, Device>::type PacketReturnType;
161 static constexpr int PacketSize = PacketType<CoeffReturnType, Device>::size;
162 typedef StorageMemory<CoeffReturnType, Device> Storage;
163 typedef typename Storage::Type EvaluatorPointerType;
165 static constexpr int Layout = TensorEvaluator<ArgType, Device>::Layout;
168 PacketAccess = TensorEvaluator<ArgType, Device>::PacketAccess,
174 PreferBlockAccess =
true,
180 typedef internal::TensorBlockDescriptor<NumDims, Index> TensorBlockDesc;
181 typedef internal::TensorBlockScratchAllocator<Device> TensorBlockScratch;
182 typedef typename internal::TensorMaterializedBlock<Scalar, NumDims, Layout, Index> TensorBlock;
185 EIGEN_STRONG_INLINE TensorEvaluator(
const XprType& op,
const Device& device)
186 : m_device(device), m_impl(op.expression(), device) {
187 EIGEN_STATIC_ASSERT((NumDims >= 4), YOU_MADE_A_PROGRAMMING_MISTAKE);
189 m_paddingValue = op.padding_value();
191 const typename TensorEvaluator<ArgType, Device>::Dimensions& input_dims = m_impl.dimensions();
194 EIGEN_IF_CONSTEXPR (
static_cast<int>(Layout) ==
static_cast<int>(
ColMajor)) {
195 m_inputDepth = input_dims[0];
196 m_inputRows = input_dims[1];
197 m_inputCols = input_dims[2];
199 m_inputDepth = input_dims[NumInputDims - 1];
200 m_inputRows = input_dims[NumInputDims - 2];
201 m_inputCols = input_dims[NumInputDims - 3];
204 m_row_strides = op.row_strides();
205 m_col_strides = op.col_strides();
208 m_in_row_strides = op.in_row_strides();
209 m_in_col_strides = op.in_col_strides();
210 m_row_inflate_strides = op.row_inflate_strides();
211 m_col_inflate_strides = op.col_inflate_strides();
225 m_input_rows_eff = (m_inputRows - 1) * m_row_inflate_strides + 1;
226 m_input_cols_eff = (m_inputCols - 1) * m_col_inflate_strides + 1;
227 m_patch_rows_eff = op.patch_rows() + (op.patch_rows() - 1) * (m_in_row_strides - 1);
228 m_patch_cols_eff = op.patch_cols() + (op.patch_cols() - 1) * (m_in_col_strides - 1);
230 if (op.padding_explicit()) {
231 m_outputRows = numext::ceil((m_input_rows_eff + op.padding_top() + op.padding_bottom() - m_patch_rows_eff + 1.f) /
232 static_cast<float>(m_row_strides));
233 m_outputCols = numext::ceil((m_input_cols_eff + op.padding_left() + op.padding_right() - m_patch_cols_eff + 1.f) /
234 static_cast<float>(m_col_strides));
235 m_rowPaddingTop = op.padding_top();
236 m_colPaddingLeft = op.padding_left();
239 switch (op.padding_type()) {
241 m_outputRows = numext::ceil((m_input_rows_eff - m_patch_rows_eff + 1.f) /
static_cast<float>(m_row_strides));
242 m_outputCols = numext::ceil((m_input_cols_eff - m_patch_cols_eff + 1.f) /
static_cast<float>(m_col_strides));
245 numext::maxi<Index>(0, ((m_outputRows - 1) * m_row_strides + m_patch_rows_eff - m_input_rows_eff) / 2);
247 numext::maxi<Index>(0, ((m_outputCols - 1) * m_col_strides + m_patch_cols_eff - m_input_cols_eff) / 2);
250 m_outputRows = numext::ceil(m_input_rows_eff /
static_cast<float>(m_row_strides));
251 m_outputCols = numext::ceil(m_input_cols_eff /
static_cast<float>(m_col_strides));
253 m_rowPaddingTop = ((m_outputRows - 1) * m_row_strides + m_patch_rows_eff - m_input_rows_eff) / 2;
254 m_colPaddingLeft = ((m_outputCols - 1) * m_col_strides + m_patch_cols_eff - m_input_cols_eff) / 2;
257 m_rowPaddingTop = numext::maxi<Index>(0, m_rowPaddingTop);
258 m_colPaddingLeft = numext::maxi<Index>(0, m_colPaddingLeft);
261 eigen_assert(
false &&
"unexpected padding");
266 eigen_assert(m_outputRows > 0);
267 eigen_assert(m_outputCols > 0);
270 EIGEN_IF_CONSTEXPR (
static_cast<int>(Layout) ==
static_cast<int>(
ColMajor)) {
277 m_dimensions[0] = input_dims[0];
278 m_dimensions[1] = op.patch_rows();
279 m_dimensions[2] = op.patch_cols();
280 m_dimensions[3] = m_outputRows * m_outputCols;
281 for (
int i = 4; i < NumDims; ++i) {
282 m_dimensions[i] = input_dims[i - 1];
291 m_dimensions[NumDims - 1] = input_dims[NumInputDims - 1];
292 m_dimensions[NumDims - 2] = op.patch_rows();
293 m_dimensions[NumDims - 3] = op.patch_cols();
294 m_dimensions[NumDims - 4] = m_outputRows * m_outputCols;
295 for (
int i = NumDims - 5; i >= 0; --i) {
296 m_dimensions[i] = input_dims[i];
301 EIGEN_IF_CONSTEXPR (
static_cast<int>(Layout) ==
static_cast<int>(
ColMajor)) {
302 m_colStride = m_dimensions[1];
303 m_patchStride = m_colStride * m_dimensions[2] * m_dimensions[0];
304 m_otherStride = m_patchStride * m_dimensions[3];
306 m_colStride = m_dimensions[NumDims - 2];
307 m_patchStride = m_colStride * m_dimensions[NumDims - 3] * m_dimensions[NumDims - 1];
308 m_otherStride = m_patchStride * m_dimensions[NumDims - 4];
312 m_rowInputStride = m_inputDepth;
313 m_colInputStride = m_inputDepth * m_inputRows;
314 m_patchInputStride = m_inputDepth * m_inputRows * m_inputCols;
317 m_fastOtherStride = internal::TensorIntDivisor<Index>(m_otherStride);
318 m_fastPatchStride = internal::TensorIntDivisor<Index>(m_patchStride);
319 m_fastColStride = internal::TensorIntDivisor<Index>(m_colStride);
320 m_fastInflateRowStride = internal::TensorIntDivisor<Index>(m_row_inflate_strides);
321 m_fastInflateColStride = internal::TensorIntDivisor<Index>(m_col_inflate_strides);
322 m_fastInputColsEff = internal::TensorIntDivisor<Index>(m_input_cols_eff);
325 m_fastOutputRows = internal::TensorIntDivisor<Index>(m_outputRows);
326 EIGEN_IF_CONSTEXPR (
static_cast<int>(Layout) ==
static_cast<int>(
ColMajor)) {
327 m_fastOutputDepth = internal::TensorIntDivisor<Index>(m_dimensions[0]);
329 m_fastOutputDepth = internal::TensorIntDivisor<Index>(m_dimensions[NumDims - 1]);
333 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
const Dimensions& dimensions()
const {
return m_dimensions; }
335 EIGEN_STRONG_INLINE
bool evalSubExprsIfNeeded(EvaluatorPointerType ) {
336 m_impl.evalSubExprsIfNeeded(
nullptr);
340#ifdef EIGEN_USE_THREADS
341 template <
typename EvalSubExprsCallback>
342 EIGEN_STRONG_INLINE
void evalSubExprsIfNeededAsync(EvaluatorPointerType, EvalSubExprsCallback done) {
343 m_impl.evalSubExprsIfNeededAsync(
nullptr, [done](
bool) { done(
true); });
347 EIGEN_STRONG_INLINE
void cleanup() { m_impl.cleanup(); }
349 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffReturnType coeff(Index index)
const {
351 Index otherIndex, patch2DIndex;
352 EIGEN_IF_CONSTEXPR (NumDims == 4) {
354 patch2DIndex = index / m_fastPatchStride;
356 otherIndex = index / m_fastOtherStride;
357 patch2DIndex = (index - otherIndex * m_otherStride) / m_fastPatchStride;
362 constexpr int depth_index =
static_cast<int>(Layout) ==
static_cast<int>(
ColMajor) ? 0 : NumDims - 1;
363 const Index patchRemainder = index - otherIndex * m_otherStride - patch2DIndex * m_patchStride;
364 const Index patchOffset = patchRemainder / m_fastOutputDepth;
365 const Index depth = patchRemainder - patchOffset * m_dimensions[depth_index];
368 const Index colIndex = patch2DIndex / m_fastOutputRows;
369 const Index colOffset = patchOffset / m_fastColStride;
370 const Index inputCol = colIndex * m_col_strides + colOffset * m_in_col_strides - m_colPaddingLeft;
371 const Index origInputCol =
372 (m_col_inflate_strides == 1) ? inputCol : ((inputCol >= 0) ? (inputCol / m_fastInflateColStride) : 0);
373 if (inputCol < 0 || inputCol >= m_input_cols_eff ||
374 ((m_col_inflate_strides != 1) && (inputCol != origInputCol * m_col_inflate_strides))) {
375 return Scalar(m_paddingValue);
379 const Index rowIndex = patch2DIndex - colIndex * m_outputRows;
380 const Index rowOffset = patchOffset - colOffset * m_colStride;
381 const Index inputRow = rowIndex * m_row_strides + rowOffset * m_in_row_strides - m_rowPaddingTop;
382 const Index origInputRow =
383 (m_row_inflate_strides == 1) ? inputRow : ((inputRow >= 0) ? (inputRow / m_fastInflateRowStride) : 0);
384 if (inputRow < 0 || inputRow >= m_input_rows_eff ||
385 ((m_row_inflate_strides != 1) && (inputRow != origInputRow * m_row_inflate_strides))) {
386 return Scalar(m_paddingValue);
389 const Index inputIndex =
390 depth + origInputRow * m_rowInputStride + origInputCol * m_colInputStride + otherIndex * m_patchInputStride;
391 return m_impl.coeff(inputIndex);
394 template <
int LoadMode>
395 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketReturnType packet(Index index)
const {
396 eigen_assert(index + PacketSize - 1 < dimensions().TotalSize());
398 constexpr int depth_index =
static_cast<int>(Layout) ==
static_cast<int>(
ColMajor) ? 0 : NumDims - 1;
399 const Index lastIdx = index + PacketSize - 1;
404 Index otherIndex, patch2DIndex, patchRemainder0, patchRemainder1;
405 EIGEN_IF_CONSTEXPR (NumDims == 4) {
407 patch2DIndex = index / m_fastPatchStride;
408 const Index patchBase = patch2DIndex * m_patchStride;
409 if (lastIdx >= patchBase + m_patchStride) {
410 return packetWithPossibleZero(index);
412 patchRemainder0 = index - patchBase;
413 patchRemainder1 = lastIdx - patchBase;
415 otherIndex = index / m_fastOtherStride;
416 const Index otherBase = otherIndex * m_otherStride;
417 if (lastIdx >= otherBase + m_otherStride) {
418 return packetWithPossibleZero(index);
420 const Index patchBase0 = index - otherBase;
421 patch2DIndex = patchBase0 / m_fastPatchStride;
422 const Index patchStart = patch2DIndex * m_patchStride;
423 if (lastIdx - otherBase >= patchStart + m_patchStride) {
424 return packetWithPossibleZero(index);
426 patchRemainder0 = patchBase0 - patchStart;
427 patchRemainder1 = lastIdx - otherBase - patchStart;
432 const Index patchOffset0 = patchRemainder0 / m_fastOutputDepth;
433 const Index colIndex = patch2DIndex / m_fastOutputRows;
438 const Index outputDepth = m_dimensions[depth_index];
439 if (patchRemainder1 < (patchOffset0 + 1) * outputDepth) {
440 const Index colOffset = patchOffset0 / m_fastColStride;
441 const Index rowIndex = patch2DIndex - colIndex * m_outputRows;
442 const Index rowOffset = patchOffset0 - colOffset * m_colStride;
444 const Index inputCol = colIndex * m_col_strides + colOffset * m_in_col_strides - m_colPaddingLeft;
445 const Index inputRow = rowIndex * m_row_strides + rowOffset * m_in_row_strides - m_rowPaddingTop;
448 if (inputCol < 0 || inputCol >= m_input_cols_eff) {
449 return internal::pset1<PacketReturnType>(Scalar(m_paddingValue));
451 if (m_col_inflate_strides != 1) {
452 const Index origCol = inputCol / m_fastInflateColStride;
453 if (inputCol != origCol * m_col_inflate_strides) {
454 return internal::pset1<PacketReturnType>(Scalar(m_paddingValue));
459 if (inputRow < 0 || inputRow >= m_input_rows_eff) {
460 return internal::pset1<PacketReturnType>(Scalar(m_paddingValue));
462 if (m_row_inflate_strides != 1) {
463 const Index origRow = inputRow / m_fastInflateRowStride;
464 if (inputRow != origRow * m_row_inflate_strides) {
465 return internal::pset1<PacketReturnType>(Scalar(m_paddingValue));
470 const Index origInputCol = (m_col_inflate_strides == 1) ? inputCol : inputCol / m_fastInflateColStride;
471 const Index origInputRow = (m_row_inflate_strides == 1) ? inputRow : inputRow / m_fastInflateRowStride;
473 const Index depth = patchRemainder0 - patchOffset0 * outputDepth;
474 const Index inputIndex =
475 depth + origInputRow * m_rowInputStride + origInputCol * m_colInputStride + otherIndex * m_patchInputStride;
476 return m_impl.template packet<Unaligned>(inputIndex);
480 if (m_in_row_strides != 1 || m_in_col_strides != 1 || m_row_inflate_strides != 1 || m_col_inflate_strides != 1) {
481 return packetWithPossibleZero(index);
486 const Index patchOffset1 = patchRemainder1 / m_fastOutputDepth;
487 const Index colOffset0 = patchOffset0 / m_fastColStride;
491 const Index colBound = (colOffset0 + 1) * m_colStride;
492 const bool sameCol = (patchOffset1 < colBound);
495 const Index inputCol0 = colIndex * m_col_strides + colOffset0 - m_colPaddingLeft;
497 if (inputCol0 < 0 || inputCol0 >= m_inputCols) {
498 return internal::pset1<PacketReturnType>(Scalar(m_paddingValue));
501 const Index rowIndex = patch2DIndex - colIndex * m_outputRows;
502 const Index rowOffset0 = patchOffset0 - colOffset0 * m_colStride;
503 const Index rowOffset1 = patchOffset1 - colOffset0 * m_colStride;
504 eigen_assert(rowOffset0 <= rowOffset1);
506 const Index inputRow0 = rowIndex * m_row_strides + rowOffset0 - m_rowPaddingTop;
507 const Index inputRow1 = rowIndex * m_row_strides + rowOffset1 - m_rowPaddingTop;
509 if (inputRow1 < 0 || inputRow0 >= m_inputRows) {
510 return internal::pset1<PacketReturnType>(Scalar(m_paddingValue));
513 if (inputRow0 >= 0 && inputRow1 < m_inputRows) {
515 const Index depth = patchRemainder0 - patchOffset0 * outputDepth;
516 const Index inputIndex =
517 depth + inputRow0 * m_rowInputStride + inputCol0 * m_colInputStride + otherIndex * m_patchInputStride;
518 return m_impl.template packet<Unaligned>(inputIndex);
523 const Index colOffset1 = patchOffset1 / m_fastColStride;
524 const Index inputCol1 = colIndex * m_col_strides + colOffset1 - m_colPaddingLeft;
525 if (inputCol1 < 0 || inputCol0 >= m_inputCols) {
526 return internal::pset1<PacketReturnType>(Scalar(m_paddingValue));
530 return packetWithPossibleZero(index);
533 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE internal::TensorBlockResourceRequirements getResourceRequirements()
const {
534 const size_t target_size = m_device.firstLevelCacheSize();
540 const TensorOpCost cost_per_coeff = m_impl.costPerCoeff(
false) + TensorOpCost(0,
sizeof(Scalar), 0);
541 return internal::TensorBlockResourceRequirements::withShapeAndSize<Scalar>(
542 internal::TensorBlockShapeType::kSkewedInnerDims, target_size, cost_per_coeff);
549 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorBlock block(TensorBlockDesc& desc, TensorBlockScratch& scratch,
550 bool =
false)
const {
551 constexpr bool is_col_major =
static_cast<int>(Layout) ==
static_cast<int>(
ColMajor);
553 if (desc.size() == 0) {
554 return TensorBlock(internal::TensorBlockKind::kView,
nullptr, desc.dimensions());
557 typename TensorBlock::Storage block_storage = TensorBlock::prepareStorage(desc, scratch);
558 Scalar* block_buffer = block_storage.data();
561 const DSizes<Index, NumDims> output_strides = internal::strides<Layout>(m_dimensions);
562 array<Index, NumDims> coords;
563 Index remaining = desc.offset();
564 EIGEN_IF_CONSTEXPR (is_col_major) {
565 for (
int i = NumDims - 1; i > 0; --i) {
566 coords[i] = remaining / output_strides[i];
567 remaining -= coords[i] * output_strides[i];
569 coords[0] = remaining;
571 for (
int i = 0; i < NumDims - 1; ++i) {
572 coords[i] = remaining / output_strides[i];
573 remaining -= coords[i] * output_strides[i];
575 coords[NumDims - 1] = remaining;
579 const int dd = is_col_major ? 0 : NumDims - 1;
580 const int rd = is_col_major ? 1 : NumDims - 2;
581 const int cd = is_col_major ? 2 : NumDims - 3;
582 const int pd = is_col_major ? 3 : NumDims - 4;
584 const Index depth_start = coords[dd];
585 const Index depth_size = desc.dimension(dd);
586 const Index row_start = coords[rd];
587 const Index row_size = desc.dimension(rd);
588 const Index col_start = coords[cd];
589 const Index col_size = desc.dimension(cd);
590 const Index patch_start = coords[pd];
591 const Index patch_size = desc.dimension(pd);
595 array<Index, NumDims> other_sizes;
596 array<Index, NumDims> other_src_stride;
597 array<Index, NumDims> other_count;
601 Index in_stride = m_patchInputStride;
602 for (
int k = 4; k < NumDims; ++k) {
603 const int d = is_col_major ? k : NumDims - 1 - k;
604 other_sizes[num_other] = desc.dimension(d);
605 other_src_stride[num_other] = in_stride;
606 other_count[num_other] = 0;
607 src_other += coords[d] * in_stride;
608 in_stride *= m_dimensions[d];
613 typedef internal::StridedLinearBufferCopy<Scalar, Index> LinCopy;
620 for (Index p = 0; p < patch_size; ++p) {
621 const Index patch2DIndex = patch_start + p;
622 const Index colIndex = patch2DIndex / m_fastOutputRows;
623 const Index rowIndex = patch2DIndex - colIndex * m_outputRows;
625 for (Index c = 0; c < col_size; ++c) {
626 const Index colOffset = col_start + c;
627 const Index inputCol = colIndex * m_col_strides + colOffset * m_in_col_strides - m_colPaddingLeft;
628 Index origInputCol = inputCol;
629 bool col_valid = inputCol >= 0 && inputCol < m_input_cols_eff;
630 if (col_valid && m_col_inflate_strides != 1) {
631 origInputCol = inputCol / m_fastInflateColStride;
632 col_valid = (inputCol == origInputCol * m_col_inflate_strides);
635 for (Index r = 0; r < row_size; ++r) {
636 const Index rowOffset = row_start + r;
637 bool valid = col_valid;
638 Index origInputRow = 0;
640 const Index inputRow = rowIndex * m_row_strides + rowOffset * m_in_row_strides - m_rowPaddingTop;
641 valid = inputRow >= 0 && inputRow < m_input_rows_eff;
643 origInputRow = inputRow;
644 if (m_row_inflate_strides != 1) {
645 origInputRow = inputRow / m_fastInflateRowStride;
646 valid = (inputRow == origInputRow * m_row_inflate_strides);
652 depth_start + origInputRow * m_rowInputStride + origInputCol * m_colInputStride + src_other;
653 for (Index d = 0; d < depth_size; ++d) {
654 block_buffer[dst + d] = m_impl.coeff(src + d);
657 LinCopy::template Run<LinCopy::Kind::FillLinear>(
typename LinCopy::Dst(dst, 1, block_buffer),
658 typename LinCopy::Src(0, 0, &m_paddingValue),
667 for (; k < num_other; ++k) {
668 if (++other_count[k] < other_sizes[k]) {
669 src_other += other_src_stride[k];
673 src_other -= other_src_stride[k] * (other_sizes[k] - 1);
675 if (k == num_other)
break;
677 eigen_assert(dst == desc.size());
679 return block_storage.AsTensorMaterializedBlock();
682 EIGEN_DEVICE_FUNC EvaluatorPointerType data()
const {
return nullptr; }
684 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
const TensorEvaluator<ArgType, Device>& impl()
const {
return m_impl; }
685 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index rowPaddingTop()
const {
return m_rowPaddingTop; }
686 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index colPaddingLeft()
const {
return m_colPaddingLeft; }
687 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index outputRows()
const {
return m_outputRows; }
688 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index outputCols()
const {
return m_outputCols; }
689 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index userRowStride()
const {
return m_row_strides; }
690 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index userColStride()
const {
return m_col_strides; }
691 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index userInRowStride()
const {
return m_in_row_strides; }
692 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index userInColStride()
const {
return m_in_col_strides; }
693 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index rowInflateStride()
const {
return m_row_inflate_strides; }
694 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index colInflateStride()
const {
return m_col_inflate_strides; }
696 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorOpCost costPerCoeff(
bool vectorized)
const {
700 const double compute_cost =
701 5 * TensorOpCost::DivCost<Index>() + 12 * TensorOpCost::MulCost<Index>() + 8 * TensorOpCost::AddCost<Index>();
702 return m_impl.costPerCoeff(vectorized) + TensorOpCost(0, 0, compute_cost, vectorized, PacketSize);
706 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketReturnType packetWithPossibleZero(Index index)
const {
707 EIGEN_ALIGN_TO_BOUNDARY(internal::unpacket_traits<PacketReturnType>::alignment)
708 std::remove_const_t<CoeffReturnType> values[PacketSize];
710 for (
int i = 0; i < PacketSize; ++i) {
711 values[i] = coeff(index + i);
713 PacketReturnType rslt = internal::pload<PacketReturnType>(values);
717 Dimensions m_dimensions;
725 Index m_in_row_strides;
726 Index m_in_col_strides;
727 Index m_row_inflate_strides;
728 Index m_col_inflate_strides;
730 Index m_input_rows_eff;
731 Index m_input_cols_eff;
732 Index m_patch_rows_eff;
733 Index m_patch_cols_eff;
735 internal::TensorIntDivisor<Index> m_fastOtherStride;
736 internal::TensorIntDivisor<Index> m_fastPatchStride;
737 internal::TensorIntDivisor<Index> m_fastColStride;
738 internal::TensorIntDivisor<Index> m_fastInflateRowStride;
739 internal::TensorIntDivisor<Index> m_fastInflateColStride;
740 internal::TensorIntDivisor<Index> m_fastInputColsEff;
742 Index m_rowInputStride;
743 Index m_colInputStride;
744 Index m_patchInputStride;
753 Index m_rowPaddingTop;
754 Index m_colPaddingLeft;
756 internal::TensorIntDivisor<Index> m_fastOutputRows;
757 internal::TensorIntDivisor<Index> m_fastOutputDepth;
759 Scalar m_paddingValue;
761 const Device EIGEN_DEVICE_REF m_device;
762 TensorEvaluator<ArgType, Device> m_impl;
The tensor base class.
Definition TensorForwardDeclarations.h:69
Patch extraction specialized for image processing. This assumes that the input has at least 3 dimensi...
Definition TensorImagePatch.h:54
Namespace containing all symbols from the Eigen library.
The tensor evaluator class.
Definition TensorEvaluator.h:47