releaase 2.11 (#703)
This commit is contained in:
@@ -145,6 +145,7 @@ private:
|
||||
uint32_t predicates_[kAccessesPerVector];
|
||||
int filter_rs_;
|
||||
int filter_c_;
|
||||
int channels_per_group_;
|
||||
|
||||
//
|
||||
// Assertions
|
||||
@@ -175,6 +176,7 @@ public:
|
||||
|
||||
filter_c_ = threadblock_offset.row() + thread_coord.contiguous();
|
||||
Index column = threadblock_offset.column() + thread_coord.strided();
|
||||
channels_per_group_ = problem_size_.C / problem_size_.groups;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int s = 0; s < ThreadMap::Iterations::kStrided; ++s) {
|
||||
@@ -188,7 +190,7 @@ public:
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int v_idx = 0; v_idx < kAccessesPerVector; ++v_idx) {
|
||||
clear_mask(v_idx, filter_c_ + v_idx * AccessType::kElements >= problem_size_.C);
|
||||
clear_mask(v_idx, filter_c_ + v_idx * AccessType::kElements >= channels_per_group_);
|
||||
}
|
||||
|
||||
pointer_ += (
|
||||
@@ -229,7 +231,7 @@ public:
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int v_idx = 0; v_idx < kAccessesPerVector; ++v_idx) {
|
||||
clear_mask(v_idx, filter_c_ + v_idx * AccessType::kElements >= problem_size_.C);
|
||||
clear_mask(v_idx, filter_c_ + v_idx * AccessType::kElements >= channels_per_group_);
|
||||
}
|
||||
|
||||
pointer_ += next;
|
||||
|
||||
230
include/cutlass/conv/threadblock/depthwise_direct_conv_params.h
Normal file
230
include/cutlass/conv/threadblock/depthwise_direct_conv_params.h
Normal file
@@ -0,0 +1,230 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017 - 2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
/*!
|
||||
\file
|
||||
\brief Extracts the host-params objects into non-template code.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#define TRACE_CONV_PARAMS_INITIALIZERS_ENABLED 0
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/fast_math.h"
|
||||
#include "cutlass/layout/tensor.h"
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/layout/pitch_linear.h"
|
||||
#include "cutlass/conv/convolution.h"
|
||||
#include "cutlass/conv/conv2d_problem_size.h"
|
||||
|
||||
#if TRACE_CONV_PARAMS_INITIALIZERS_ENABLED
|
||||
#include <fstream>
|
||||
#endif
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace conv {
|
||||
namespace threadblock {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Parameters structure used for DepthwiseFpropActivationDirect2dConvTileAccessIteratorOptimized
|
||||
template<typename Layout_ = layout::TensorNHWC >
|
||||
struct Depthwise2dFpropDirectConvParams;
|
||||
|
||||
/// Parameters structure used for DepthwiseFpropActivationDirect2dConvTileAccessIteratorFixedStrideDilation
|
||||
template<typename Layout_ = layout::TensorNHWC >
|
||||
struct Depthwise2dFpropDirectConvActivationIteratorFixedStrideDilationParams;
|
||||
|
||||
/// Parameters structure used for DepthwiseFpropFilterDirectConvTileAccessIteratorOptimized
|
||||
template<typename Layout_ = layout::TensorNHWC >
|
||||
struct Depthwise2dFpropDirectConvFilterIteratorParams;
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Parameters structure used for DepthwiseFpropActivationDirect2dConvTileAccessIteratorOptimized
|
||||
template<>
|
||||
struct Depthwise2dFpropDirectConvParams<layout::TensorNHWC> {
|
||||
|
||||
using Layout = layout::TensorNHWC;
|
||||
|
||||
Layout layout;
|
||||
|
||||
int32_t activation_tile_h;
|
||||
int32_t activation_tile_w;
|
||||
int32_t activation_tile_hw;
|
||||
FastDivmod activation_tile_w_divmod;
|
||||
|
||||
int filter[2];
|
||||
int stride[2];
|
||||
int dilation[2];
|
||||
int inc_next[2];
|
||||
FastDivmod pq_divmod;
|
||||
FastDivmod q_divmod;
|
||||
|
||||
int activation_load_count;
|
||||
int activation_storage_elements;
|
||||
int activation_size;
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Depthwise2dFpropDirectConvParams() { }
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Depthwise2dFpropDirectConvParams(
|
||||
Conv2dProblemSize const &problem_size,
|
||||
Layout const &layout, ///< layout object
|
||||
MatrixCoord threadblock_shape, ///< CTA threadblock Shape
|
||||
Layout::TensorCoord threadblock_output_shape, ///< Output tile Shape per threadblock
|
||||
const int element_size_bits, ///< bits of activation element
|
||||
const int thread_count, ///< threads per threadblock
|
||||
const int thread_count_contiguous, ///< number of threads for continuous dimension
|
||||
const int element_per_load) ///< element per each load
|
||||
: layout(layout) {
|
||||
|
||||
filter[0] = problem_size.S;
|
||||
filter[1] = problem_size.R;
|
||||
|
||||
stride[0] = problem_size.stride_w;
|
||||
stride[1] = problem_size.stride_h;
|
||||
|
||||
dilation[0] = problem_size.dilation_w;
|
||||
dilation[1] = problem_size.dilation_h;
|
||||
|
||||
// Compute activation_tile size per threadblock because stride and dilation are runtime params.
|
||||
activation_tile_h = (threadblock_output_shape.h() - 1) * problem_size.stride_h +
|
||||
(problem_size.R - 1) * problem_size.dilation_h + 1;
|
||||
activation_tile_w = (threadblock_output_shape.w() - 1) * problem_size.stride_w +
|
||||
(problem_size.S - 1) * problem_size.dilation_w + 1;
|
||||
activation_tile_hw = activation_tile_h * activation_tile_w;
|
||||
|
||||
activation_tile_w_divmod = FastDivmod(activation_tile_w);
|
||||
|
||||
/// Below two values could not be templatized because the stride and dilation are runtime params
|
||||
activation_load_count = (thread_count_contiguous * activation_tile_hw + (thread_count - 1)) / thread_count;
|
||||
activation_storage_elements = activation_load_count * element_per_load * thread_count;
|
||||
activation_size = activation_storage_elements * element_size_bits / 8;
|
||||
|
||||
// Fastdivmod for output P, Q
|
||||
int tiles_p =
|
||||
(problem_size.P + (threadblock_output_shape.h() - 1)) / (threadblock_output_shape.h());
|
||||
int tiles_q = (problem_size.Q + (threadblock_output_shape.w() - 1)) /
|
||||
(threadblock_output_shape.w());
|
||||
|
||||
pq_divmod = FastDivmod(tiles_p * tiles_q);
|
||||
q_divmod = FastDivmod(tiles_q);
|
||||
|
||||
// next S
|
||||
inc_next[0] = problem_size.dilation_w;
|
||||
// next R
|
||||
inc_next[1] = (activation_tile_w * problem_size.dilation_h - (problem_size.S - 1) * problem_size.dilation_w);
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// Parameters structure used for DepthwiseFpropActivationDirect2dConvTileAccessIteratorFixedStrideDilation
|
||||
template <>
|
||||
struct Depthwise2dFpropDirectConvActivationIteratorFixedStrideDilationParams<layout::TensorNHWC> {
|
||||
using Layout = layout::TensorNHWC;
|
||||
|
||||
Layout layout;
|
||||
|
||||
FastDivmod pq_divmod;
|
||||
FastDivmod q_divmod;
|
||||
|
||||
int activation_size;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Depthwise2dFpropDirectConvActivationIteratorFixedStrideDilationParams() {}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Depthwise2dFpropDirectConvActivationIteratorFixedStrideDilationParams(
|
||||
Conv2dProblemSize const &problem_size,
|
||||
Layout const &layout, ///< Layout object
|
||||
MatrixCoord threadblock_shape, ///< Threadblock Shape
|
||||
Layout::TensorCoord threadblock_output_shape, ///< Output tile Shape per threadblock
|
||||
const int activation_size_ ///< Activation size loaded by iterator
|
||||
)
|
||||
: layout(layout),
|
||||
activation_size(activation_size_) {
|
||||
// Fastdivmod for output P, Q
|
||||
int tiles_p =
|
||||
(problem_size.P + (threadblock_output_shape.h() - 1)) / (threadblock_output_shape.h());
|
||||
int tiles_q =
|
||||
(problem_size.Q + (threadblock_output_shape.w() - 1)) / (threadblock_output_shape.w());
|
||||
|
||||
pq_divmod = FastDivmod(tiles_p * tiles_q);
|
||||
q_divmod = FastDivmod(tiles_q);
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Parameters structure used for DepthwiseFpropFilterDirectConvTileAccessIteratorOptimized
|
||||
template <>
|
||||
struct Depthwise2dFpropDirectConvFilterIteratorParams<layout::TensorNHWC> {
|
||||
using Layout = layout::TensorNHWC;
|
||||
|
||||
Layout layout;
|
||||
|
||||
int filter_size;
|
||||
|
||||
bool is_convolution;
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Depthwise2dFpropDirectConvFilterIteratorParams() {}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Depthwise2dFpropDirectConvFilterIteratorParams(
|
||||
Conv2dProblemSize const &problem_size,
|
||||
Layout const &layout, ///< Layout object
|
||||
MatrixCoord threadblock_shape, ///< Threadblock Shape
|
||||
const int filter_size_) ///< Filter size loaded by iterator
|
||||
: layout(layout),
|
||||
filter_size(filter_size_),
|
||||
is_convolution(problem_size.mode == Mode::kConvolution){}
|
||||
};
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace conv
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,314 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017 - 2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
/*! \file
|
||||
\brief Templates implementing loading of convolution tiles mapped to GEMM A (activation tile)
|
||||
matrix from memory.
|
||||
|
||||
This iterator assumes TensorNHWC layout of tensors in Global Memory.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/conv/conv2d_problem_size.h"
|
||||
#include "cutlass/conv/convolution.h"
|
||||
#include "cutlass/conv/threadblock/depthwise_direct_conv_params.h"
|
||||
#include "cutlass/coord.h"
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/layout/pitch_linear.h"
|
||||
#include "cutlass/layout/tensor.h"
|
||||
#include "cutlass/matrix_shape.h"
|
||||
#include "cutlass/predicate_vector.h"
|
||||
#include "cutlass/tensor_ref.h"
|
||||
#include "cutlass/tensor_view.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace conv {
|
||||
namespace threadblock {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename Shape_,
|
||||
typename OutputTileShape_,
|
||||
typename StrideShape_,
|
||||
typename DilationShape_,
|
||||
typename ActivationShape_,
|
||||
typename Element_,
|
||||
typename Layout_,
|
||||
typename ThreadMap_,
|
||||
typename AccessType_ = cutlass::AlignedArray<Element_, ThreadMap_::kElementsPerAccess> >
|
||||
class DepthwiseFpropActivationDirect2dConvTileAccessIteratorFixedStrideDilation {
|
||||
public:
|
||||
//
|
||||
// Types
|
||||
//
|
||||
|
||||
using Shape = Shape_;
|
||||
using OutputTileShape = OutputTileShape_;
|
||||
using Element = Element_;
|
||||
using Layout = Layout_;
|
||||
using TensorCoord = typename Layout::TensorCoord;
|
||||
using ThreadMap = ThreadMap_;
|
||||
using AccessType = AccessType_;
|
||||
using TensorRef = cutlass::TensorRef<Element, Layout>;
|
||||
using Index = typename Layout::Index;
|
||||
using LongIndex = typename Layout::LongIndex;
|
||||
static IteratorAlgorithm const kIteratorAlgorithm = conv::IteratorAlgorithm::kOptimized;
|
||||
static StrideSupport const kStrideSupport = conv::StrideSupport::kStrided;
|
||||
static int const kConvDim = 2;
|
||||
using ConvProblemSize = typename conv::Conv2dProblemSize;
|
||||
|
||||
// Compilation value of stride , dialtion and activation shape
|
||||
using StrideShape = StrideShape_;
|
||||
using DilationShape = DilationShape_;
|
||||
using ActivationShape = ActivationShape_;
|
||||
|
||||
|
||||
static int const kAccessesPerVector = ThreadMap::kElementsPerAccess / AccessType::kElements;
|
||||
static int const kActivationSize = ThreadMap::Iterations::kCount * ThreadMap::kElementsPerAccess * ThreadMap::kThreads *
|
||||
sizeof_bits<Element>::value / 8;
|
||||
|
||||
|
||||
static_assert(!(ThreadMap::kElementsPerAccess % AccessType::kElements),
|
||||
"Vectors implied by the thread map must be divisible by the access type.");
|
||||
|
||||
//
|
||||
// Simplifying assertions
|
||||
//
|
||||
static_assert(ThreadMap::Iterations::kContiguous == 1, "Require Iterations::kContiguous == 1");
|
||||
|
||||
static_assert(OutputTileShape::kN == 1, "Require OutputTileShape::kN == 1");
|
||||
static_assert(OutputTileShape::kC == Shape::kColumn, "Require OutputTile shape == channels per threadblock");
|
||||
|
||||
//
|
||||
// Parameters structure
|
||||
//
|
||||
|
||||
using Params = Depthwise2dFpropDirectConvActivationIteratorFixedStrideDilationParams<Layout>;
|
||||
|
||||
private:
|
||||
Conv2dProblemSize const &problem_size_;
|
||||
Params const ¶ms_;
|
||||
char const *pointer_;
|
||||
|
||||
// Base channels for current threadblock
|
||||
int base_c_;
|
||||
// Base activation index for current threadblock
|
||||
int offset_intial_npq_;
|
||||
// Base activation coord for current threadblock
|
||||
TensorCoord activatioin_base_;
|
||||
// Intial thread positioin
|
||||
int offset_initial_hwc_;
|
||||
// Overall load instruction per thread.
|
||||
int iterator_load_;
|
||||
// thread loading position.
|
||||
int iterator_hwc_;
|
||||
// activation N is inside the Tensor or not
|
||||
bool valid_n_;
|
||||
|
||||
public:
|
||||
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
DepthwiseFpropActivationDirect2dConvTileAccessIteratorFixedStrideDilation(
|
||||
Params const ¶ms,
|
||||
Conv2dProblemSize const &problem_size,
|
||||
Element const *ptr,
|
||||
int thread_idx,
|
||||
MatrixCoord const &threadblock_offset =
|
||||
MatrixCoord()
|
||||
)
|
||||
: params_(params),
|
||||
problem_size_(problem_size),
|
||||
pointer_(reinterpret_cast<char const *>(ptr)),
|
||||
offset_intial_npq_(threadblock_offset.row()),
|
||||
offset_initial_hwc_(thread_idx),
|
||||
iterator_load_(0) {
|
||||
|
||||
base_c_ = threadblock_offset.column();
|
||||
|
||||
set_iteration_index(0);
|
||||
|
||||
set_activation_coord(offset_intial_npq_);
|
||||
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
void set_activation_coord(int offset_npq) {
|
||||
int offset_inital_n, offset_inital_p, offset_inital_q;
|
||||
int residual;
|
||||
|
||||
params_.pq_divmod(offset_inital_n, residual, offset_npq);
|
||||
params_.q_divmod(offset_inital_p, offset_inital_q, residual);
|
||||
|
||||
int base_n = offset_inital_n;
|
||||
|
||||
int base_h =
|
||||
offset_inital_p * OutputTileShape::kH * StrideShape::kRow - problem_size_.pad_h;
|
||||
|
||||
int base_w =
|
||||
offset_inital_q * OutputTileShape::kW * StrideShape::kColumn - problem_size_.pad_w;
|
||||
|
||||
activatioin_base_ = TensorCoord(base_n, base_h, base_w, base_c_);
|
||||
|
||||
valid_n_ = activatioin_base_.n() < problem_size_.N;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static Params getParams(Conv2dProblemSize const &problem_size, Layout const &layout) {
|
||||
return Params(
|
||||
problem_size,
|
||||
layout,
|
||||
{Shape::kRow, Shape::kColumn},
|
||||
{OutputTileShape::kN, OutputTileShape::kH, OutputTileShape::kW, OutputTileShape::kC},
|
||||
kActivationSize);
|
||||
}
|
||||
|
||||
/// Overrides the internal iteration index
|
||||
CUTLASS_HOST_DEVICE
|
||||
void set_iteration_index(Index index) {
|
||||
iterator_hwc_ = offset_initial_hwc_ + index * ThreadMap::kThreads;
|
||||
iterator_load_ = index;
|
||||
}
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
pointer_ += pointer_offset * sizeof_bits<Element>::value / 8;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
void advance() {
|
||||
// Go to next threadblock
|
||||
offset_intial_npq_ += problem_size_.split_k_slices;
|
||||
|
||||
set_iteration_index(0);
|
||||
|
||||
set_activation_coord(offset_intial_npq_);
|
||||
}
|
||||
|
||||
/// Returns the coordinate in the activations tensor X that is currently pointed to
|
||||
/// by the iterator.
|
||||
CUTLASS_HOST_DEVICE
|
||||
TensorCoord at() const {
|
||||
int c = iterator_hwc_ % ThreadMap::Detail::ShapeVec::kContiguous ;
|
||||
int next = iterator_hwc_ / ThreadMap::Detail::ShapeVec::kContiguous ;
|
||||
int h = next / ActivationShape::kW;
|
||||
int w = next % ActivationShape::kW;
|
||||
|
||||
c = c * AccessType::kElements;
|
||||
|
||||
return activatioin_base_ + TensorCoord(0, h, w, c);
|
||||
}
|
||||
|
||||
/// Returns true if the current coordinate is within the activations tensor X
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool valid() const {
|
||||
TensorCoord coord = at();
|
||||
bool valid_c = coord.c() < problem_size_.C;
|
||||
bool valid_h = coord.h() >= 0 && coord.h() < problem_size_.H;
|
||||
bool valid_w = coord.w() >= 0 && coord.w() < problem_size_.W;
|
||||
return valid_n_ ? valid_c & valid_h & valid_w : 0;
|
||||
}
|
||||
|
||||
/// Returns a pointer to the vector starting at the current coordinate
|
||||
CUTLASS_HOST_DEVICE
|
||||
AccessType const *get() const {
|
||||
TensorCoord coord = at();
|
||||
LongIndex offset = params_.layout(coord);
|
||||
|
||||
AccessType const *ptr =
|
||||
reinterpret_cast<AccessType const *>(pointer_ + offset * sizeof_bits<Element>::value / 8);
|
||||
|
||||
return ptr;
|
||||
}
|
||||
|
||||
/// Increments to the next memory access
|
||||
CUTLASS_HOST_DEVICE
|
||||
DepthwiseFpropActivationDirect2dConvTileAccessIteratorFixedStrideDilation &operator++() {
|
||||
|
||||
++iterator_load_;
|
||||
iterator_hwc_ += ThreadMap::kThreads;
|
||||
|
||||
if (iterator_load_ < ThreadMap::Iterations::kCount) {
|
||||
return *this;
|
||||
}
|
||||
|
||||
iterator_load_ = 0;
|
||||
iterator_hwc_ = offset_initial_hwc_;
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Determines the activation size loaded by iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
int get_load_size() {
|
||||
return kActivationSize;
|
||||
}
|
||||
|
||||
/// Determines the iterations needed
|
||||
CUTLASS_HOST_DEVICE
|
||||
int get_iteration_num() {
|
||||
return ThreadMap::Iterations::kCount;
|
||||
}
|
||||
|
||||
/// Determines whether the Depthwise fprop can execute the given problem.
|
||||
CUTLASS_HOST_DEVICE
|
||||
static Status can_implement(Conv2dProblemSize const &problem_size) {
|
||||
|
||||
// check stride and dilation constraint
|
||||
if (problem_size.stride_h != StrideShape::kRow || problem_size.stride_w != StrideShape::kColumn) {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
if (problem_size.dilation_h != DilationShape::kRow || problem_size.dilation_w != DilationShape::kColumn) {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
// check alignment constraint on iterator's contiguous dimension
|
||||
if (problem_size.C % AccessType::kElements) {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace conv
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,291 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017 - 2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
/*! \file
|
||||
\brief Templates implementing loading of convolution tiles mapped to GEMM A (activation tile)
|
||||
matrix from memory.
|
||||
|
||||
This iterator assumes TensorNHWC layout of tensors in Global Memory.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/conv/conv2d_problem_size.h"
|
||||
#include "cutlass/conv/convolution.h"
|
||||
#include "cutlass/conv/threadblock/depthwise_direct_conv_params.h"
|
||||
#include "cutlass/coord.h"
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/layout/pitch_linear.h"
|
||||
#include "cutlass/layout/tensor.h"
|
||||
#include "cutlass/matrix_shape.h"
|
||||
#include "cutlass/predicate_vector.h"
|
||||
#include "cutlass/tensor_ref.h"
|
||||
#include "cutlass/tensor_view.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace conv {
|
||||
namespace threadblock {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename Shape_,
|
||||
typename OutputTileShape_,
|
||||
typename Element_,
|
||||
typename Layout_,
|
||||
typename ThreadMap_,
|
||||
typename AccessType_ = cutlass::AlignedArray<Element_, ThreadMap_::kElementsPerAccess> >
|
||||
class DepthwiseFpropActivationDirect2dConvTileAccessIteratorOptimized {
|
||||
public:
|
||||
//
|
||||
// Types
|
||||
//
|
||||
|
||||
using Shape = Shape_;
|
||||
using OutputTileShape = OutputTileShape_;
|
||||
using Element = Element_;
|
||||
using Layout = Layout_;
|
||||
using TensorCoord = typename Layout::TensorCoord;
|
||||
using ThreadMap = ThreadMap_;
|
||||
using AccessType = AccessType_;
|
||||
using TensorRef = cutlass::TensorRef<Element, Layout>;
|
||||
using Index = typename Layout::Index;
|
||||
using LongIndex = typename Layout::LongIndex;
|
||||
static IteratorAlgorithm const kIteratorAlgorithm = conv::IteratorAlgorithm::kOptimized;
|
||||
static StrideSupport const kStrideSupport = conv::StrideSupport::kStrided;
|
||||
static int const kConvDim = 2;
|
||||
using ConvProblemSize = typename conv::Conv2dProblemSize;
|
||||
|
||||
static int const kAccessesPerVector = ThreadMap::kElementsPerAccess / AccessType::kElements;
|
||||
|
||||
static_assert(!(ThreadMap::kElementsPerAccess % AccessType::kElements),
|
||||
"Vectors implied by the thread map must be divisible by the access type.");
|
||||
|
||||
//
|
||||
// Simplifying assertions
|
||||
//
|
||||
static_assert(ThreadMap::Iterations::kContiguous == 1, "Require Iterations::kContiguous == 1");
|
||||
|
||||
static_assert(OutputTileShape::kN == 1, "Require OutputTileShape::kN == 1");
|
||||
static_assert(OutputTileShape::kC == Shape::kColumn, "Require OutputTile shape == channels per threadblock");
|
||||
|
||||
//
|
||||
// Parameters structure
|
||||
//
|
||||
|
||||
using Params = Depthwise2dFpropDirectConvParams<Layout>;
|
||||
|
||||
private:
|
||||
Conv2dProblemSize const &problem_size_;
|
||||
Params const ¶ms_;
|
||||
char const *pointer_;
|
||||
|
||||
// Base channels for current threadblock
|
||||
int base_c_;
|
||||
// Base activation index for current threadblock
|
||||
int offset_intial_npq_;
|
||||
// Base activation coord for current threadblock
|
||||
TensorCoord activatioin_base_;
|
||||
// Intial thread positioin
|
||||
int offset_initial_hwc_;
|
||||
// Overall load instruction per thread.
|
||||
int iterator_load_;
|
||||
// thread loading position.
|
||||
int iterator_hwc_;
|
||||
// Number of loads for activations tensor X.
|
||||
const int number_of_loads_;
|
||||
|
||||
public:
|
||||
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
DepthwiseFpropActivationDirect2dConvTileAccessIteratorOptimized(
|
||||
Params const ¶ms,
|
||||
Conv2dProblemSize const &problem_size,
|
||||
Element const *ptr,
|
||||
int thread_idx,
|
||||
MatrixCoord const &threadblock_offset =
|
||||
MatrixCoord()
|
||||
)
|
||||
: params_(params),
|
||||
problem_size_(problem_size),
|
||||
pointer_(reinterpret_cast<char const *>(ptr)),
|
||||
offset_intial_npq_(threadblock_offset.row()),
|
||||
offset_initial_hwc_(thread_idx),
|
||||
iterator_load_(0),
|
||||
number_of_loads_(params.activation_load_count) {
|
||||
|
||||
base_c_ = threadblock_offset.column();
|
||||
|
||||
set_activation_coord(offset_intial_npq_);
|
||||
|
||||
set_iteration_index(0);
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
void set_activation_coord(int offset_npq) {
|
||||
int offset_inital_n, offset_inital_p, offset_inital_q;
|
||||
int residual;
|
||||
|
||||
params_.pq_divmod(offset_inital_n, residual, offset_npq);
|
||||
params_.q_divmod(offset_inital_p, offset_inital_q, residual);
|
||||
|
||||
int base_n = offset_inital_n;
|
||||
|
||||
int base_h =
|
||||
offset_inital_p * OutputTileShape::kH * problem_size_.stride_h - problem_size_.pad_h;
|
||||
|
||||
int base_w =
|
||||
offset_inital_q * OutputTileShape::kW * problem_size_.stride_w - problem_size_.pad_w;
|
||||
|
||||
activatioin_base_ = TensorCoord(base_n, base_h, base_w, base_c_);
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static Params getParams(Conv2dProblemSize const &problem_size, Layout const &layout) {
|
||||
return Params(
|
||||
problem_size,
|
||||
layout,
|
||||
{Shape::kRow, Shape::kColumn},
|
||||
{OutputTileShape::kN, OutputTileShape::kH, OutputTileShape::kW, OutputTileShape::kC},
|
||||
sizeof_bits<Element>::value,
|
||||
ThreadMap::kThreads,
|
||||
ThreadMap::Detail::ShapeVec::kContiguous,
|
||||
ThreadMap::kElementsPerAccess);
|
||||
}
|
||||
|
||||
/// Overrides the internal iteration index
|
||||
CUTLASS_HOST_DEVICE
|
||||
void set_iteration_index(Index index) {
|
||||
iterator_hwc_ = offset_initial_hwc_ + index * ThreadMap::kThreads;
|
||||
iterator_load_ = index;
|
||||
}
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
pointer_ += pointer_offset * sizeof_bits<Element>::value / 8;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
void advance() {
|
||||
// Go to next threadblock
|
||||
offset_intial_npq_ += problem_size_.split_k_slices;
|
||||
|
||||
set_activation_coord(offset_intial_npq_);
|
||||
}
|
||||
|
||||
/// Returns the coordinate in the activations tensor X that is currently pointed to
|
||||
/// by the iterator.
|
||||
CUTLASS_HOST_DEVICE
|
||||
TensorCoord at() const {
|
||||
|
||||
int c = iterator_hwc_ % ThreadMap::Detail::ShapeVec::kContiguous ;
|
||||
int next = iterator_hwc_ / ThreadMap::Detail::ShapeVec::kContiguous ;
|
||||
int h, w;
|
||||
params_.activation_tile_w_divmod(h, w, next) ;
|
||||
|
||||
c = c * AccessType::kElements;
|
||||
|
||||
return activatioin_base_ + TensorCoord(0, h, w, c);
|
||||
}
|
||||
|
||||
/// Returns true if the current coordinate is within the activations tensor X
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool valid() const {
|
||||
TensorCoord coord = at();
|
||||
|
||||
return coord.n() < problem_size_.N && coord.h() >= 0 && coord.h() < problem_size_.H &&
|
||||
coord.w() >= 0 && coord.w() < problem_size_.W && coord.c() < problem_size_.C;
|
||||
}
|
||||
|
||||
/// Returns a pointer to the vector starting at the current coordinate
|
||||
CUTLASS_HOST_DEVICE
|
||||
AccessType const *get() const {
|
||||
TensorCoord coord = at();
|
||||
LongIndex offset = params_.layout(coord);
|
||||
|
||||
AccessType const *ptr =
|
||||
reinterpret_cast<AccessType const *>(pointer_ + offset * sizeof_bits<Element>::value / 8);
|
||||
|
||||
return ptr;
|
||||
}
|
||||
|
||||
/// Increments to the next memory access
|
||||
CUTLASS_HOST_DEVICE
|
||||
DepthwiseFpropActivationDirect2dConvTileAccessIteratorOptimized &operator++() {
|
||||
|
||||
++iterator_load_;
|
||||
iterator_hwc_ += ThreadMap::kThreads;
|
||||
|
||||
if (iterator_load_ < number_of_loads_) {
|
||||
return *this;
|
||||
}
|
||||
|
||||
iterator_load_ = 0;
|
||||
iterator_hwc_ = offset_initial_hwc_;
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Determines the activation size loaded by iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
int get_load_size() {
|
||||
return params_.activation_size;
|
||||
}
|
||||
|
||||
/// Determines the iterations needed
|
||||
CUTLASS_HOST_DEVICE
|
||||
int get_iteration_num() {
|
||||
return number_of_loads_;
|
||||
}
|
||||
|
||||
/// Determines whether the Depthwise fprop can execute the given problem.
|
||||
CUTLASS_HOST_DEVICE
|
||||
static Status can_implement(Conv2dProblemSize const &problem_size) {
|
||||
// check alignment constraint on iterator's contiguous dimension
|
||||
if (problem_size.C % AccessType::kElements) {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace conv
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,551 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017 - 2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
/*! \file
|
||||
\brief Template for a multistage threadblock-scoped Implicit GEMM Convolution kernel.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/aligned_buffer.h"
|
||||
#include "cutlass/arch/memory.h"
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/matrix_shape.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/arch/cache_operation.h"
|
||||
#include "cutlass/conv/threadblock/depthwise_mma_base.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace conv {
|
||||
namespace threadblock {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Structure to compute the matrix product targeting CUDA cores and SIMT math
|
||||
/// instructions.
|
||||
template <
|
||||
/// Size of the Gemm problem - concept: gemm::GemmShape<>
|
||||
typename Shape_,
|
||||
/// Iterates over tiles of A operand in global memory
|
||||
// (concept: ReadableTileIterator | ForwardTileIterator |
|
||||
// MaskedTileIterator)
|
||||
typename IteratorA_,
|
||||
/// Iterates over tiles of A operand in shared memory
|
||||
/// (concept: WriteableTileIterator | RandomAccessTileIterator)
|
||||
typename SmemIteratorA_,
|
||||
/// Cache operation for operand A
|
||||
cutlass::arch::CacheOperation::Kind CacheOpA,
|
||||
/// Iterates over tiles of B operand in global memory
|
||||
// (concept: ReadableTileIterator | ForwardTileIterator |
|
||||
// MaskedTileIterator)
|
||||
typename IteratorB_,
|
||||
/// Iterates over tiles of B operand in shared memory
|
||||
/// (concept: WriteableTileIterator | RandomAccessTileIterator)
|
||||
typename SmemIteratorB_,
|
||||
/// Cache operation for operand B
|
||||
cutlass::arch::CacheOperation::Kind CacheOpB,
|
||||
/// Policy describing tuning details (concept: MmaPolicy)
|
||||
typename Policy_,
|
||||
/// Number of stages,
|
||||
int Stages,
|
||||
/// Epilogue stores the data into global memory
|
||||
typename Epilogue_,
|
||||
/// iterator implementation variants
|
||||
conv::IteratorAlgorithm IteratorAlgorithm_ = conv::IteratorAlgorithm::kOptimized,
|
||||
/// Used for partial specialization
|
||||
typename Enable = bool>
|
||||
class DepthwiseFpropDirectConvMultipleStage :
|
||||
public DepthwiseDirectConvMmaBase<Shape_, Policy_, Stages> {
|
||||
public:
|
||||
///< Base class
|
||||
using Base = DepthwiseDirectConvMmaBase<Shape_, Policy_, Stages>;
|
||||
///< Size of the Gemm problem - concept: gemm::GemmShape<>
|
||||
using Shape = Shape_;
|
||||
///< Iterates over tiles of A operand in global memory
|
||||
using IteratorA = IteratorA_;
|
||||
///< Iterates over tiles of B operand in global memory
|
||||
using IteratorB = IteratorB_;
|
||||
///< Policy describing tuning details
|
||||
using Policy = Policy_;
|
||||
|
||||
using Epilogue = Epilogue_;
|
||||
|
||||
using SmemIteratorA = SmemIteratorA_;
|
||||
using SmemIteratorB = SmemIteratorB_;
|
||||
|
||||
static cutlass::arch::CacheOperation::Kind const kCacheOpA = CacheOpA;
|
||||
static cutlass::arch::CacheOperation::Kind const kCacheOpB = CacheOpB;
|
||||
|
||||
static conv::IteratorAlgorithm const kItertorAlgorithm = IteratorAlgorithm_;
|
||||
|
||||
//
|
||||
// Dependent types
|
||||
//
|
||||
|
||||
/// Fragment of accumulator tile
|
||||
|
||||
using ElementC = typename Policy::Operator::ElementC;
|
||||
using FragmentC = typename Policy::Operator::FragmentC;
|
||||
|
||||
/// Warp-level Mma
|
||||
using Operator = typename Policy::Operator;
|
||||
|
||||
/// Internal structure exposed for introspection.
|
||||
struct Detail {
|
||||
|
||||
/// Number of cp.async instructions to load one stage of operand A
|
||||
static int const AsyncCopyIterationsPerStageA =
|
||||
IteratorA::ThreadMap::Iterations::kCount;
|
||||
|
||||
/// Number of cp.async instructions to load one stage of operand B
|
||||
static int const AsyncCopyIterationsPerStageB =
|
||||
IteratorB::ThreadMap::Iterations::kCount;
|
||||
|
||||
/// Number of stages
|
||||
static int const kStages = Stages;
|
||||
|
||||
/// Number of cp.async instructions to load on group of operand B
|
||||
static int const kAccessesPerGroupB =
|
||||
(AsyncCopyIterationsPerStageB + Base::kWarpGemmIterations - 1) / Base::kWarpGemmIterations;
|
||||
};
|
||||
|
||||
private:
|
||||
|
||||
using WarpLoadedFragmentA = typename Operator::FragmentA;
|
||||
using WarpLoadedFragmentB = typename Operator::FragmentB;
|
||||
using WarpTransformedFragmentA = typename Operator::TransformedFragmentA;
|
||||
using WarpTransformedFragmentB = typename Operator::TransformedFragmentB;
|
||||
|
||||
private:
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Iterator to write threadblock-scoped tile of A operand to shared memory
|
||||
SmemIteratorA smem_iterator_A_;
|
||||
|
||||
/// Iterator to write threadblock-scoped tile of B operand to shared memory
|
||||
SmemIteratorB smem_iterator_B_;
|
||||
|
||||
public:
|
||||
|
||||
/// Construct from tensor references
|
||||
CUTLASS_DEVICE
|
||||
DepthwiseFpropDirectConvMultipleStage(
|
||||
///< Shared storage needed for internal use by threadblock-scoped GEMM
|
||||
typename Base::SharedStorage &shared_storage,
|
||||
///< ID within the threadblock
|
||||
int thread_idx,
|
||||
///< ID of warp
|
||||
int warp_idx,
|
||||
///< ID of each thread within a warp
|
||||
int lane_idx
|
||||
):
|
||||
Base(shared_storage, thread_idx, warp_idx, lane_idx),
|
||||
smem_iterator_A_(shared_storage.operand_A_ref(), thread_idx),
|
||||
smem_iterator_B_(shared_storage.operand_B_ref(), thread_idx)
|
||||
{
|
||||
// Compute warp location within threadblock tile by mapping the warp_id to
|
||||
// three coordinates:
|
||||
// _m: the warp's position within the threadblock along the M dimension
|
||||
// _n: the warp's position within the threadblock along the N dimension
|
||||
// _k: the warp's position within the threadblock along the K dimension
|
||||
|
||||
int warp_idx_mn = warp_idx % (Base::WarpCount::kM * Base::WarpCount::kN);
|
||||
int warp_idx_k = warp_idx / (Base::WarpCount::kM * Base::WarpCount::kN);
|
||||
|
||||
int warp_idx_m = warp_idx_mn % Base::WarpCount::kM;
|
||||
int warp_idx_n = warp_idx_mn / Base::WarpCount::kM;
|
||||
|
||||
// Add per-warp offsets in units of warp-level tiles
|
||||
this->warp_tile_iterator_A_.add_tile_offset(
|
||||
{warp_idx_m, Base::kWarpGemmIterations * warp_idx_k});
|
||||
this->warp_tile_iterator_B_.add_tile_offset(
|
||||
{Base::kWarpGemmIterations * warp_idx_k, warp_idx_n});
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void copy_tiles_and_advance(IteratorA &iterator_A,
|
||||
IteratorB &iterator_B,
|
||||
int group_start_A = 0,
|
||||
int group_start_B = 0) {
|
||||
if (kItertorAlgorithm == conv::IteratorAlgorithm::kFixedStrideDilation) {
|
||||
// Number of iterators is a static value.
|
||||
iterator_A.set_iteration_index(group_start_A * IteratorA::kAccessesPerVector);
|
||||
this->smem_iterator_A_.set_iteration_index(group_start_A);
|
||||
|
||||
// Async Copy for operand A
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int j = 0; j < Detail::AsyncCopyIterationsPerStageA; ++j) {
|
||||
typename IteratorA::AccessType *dst_ptr =
|
||||
reinterpret_cast<typename IteratorA::AccessType *>(this->smem_iterator_A_.get());
|
||||
|
||||
int const kSrcBytes = sizeof_bits<typename IteratorA::Element>::value *
|
||||
IteratorA::ThreadMap::kElementsPerAccess /
|
||||
IteratorA::kAccessesPerVector / 8;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int v = 0; v < IteratorA::kAccessesPerVector; ++v) {
|
||||
cutlass::arch::cp_async_zfill<kSrcBytes, kCacheOpA>(
|
||||
dst_ptr + v, iterator_A.get(), iterator_A.valid());
|
||||
|
||||
++iterator_A;
|
||||
}
|
||||
++this->smem_iterator_A_;
|
||||
}
|
||||
} else {
|
||||
// Number of iterators is a runtime value.
|
||||
iterator_A.set_iteration_index(group_start_A * IteratorA::kAccessesPerVector);
|
||||
this->smem_iterator_A_.set_iteration_index(group_start_A);
|
||||
|
||||
// Async Copy for operand A
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int j = 0; j < iterator_A.get_iteration_num(); ++j) {
|
||||
typename IteratorA::AccessType *dst_ptr =
|
||||
reinterpret_cast<typename IteratorA::AccessType *>(this->smem_iterator_A_.get());
|
||||
|
||||
int const kSrcBytes = sizeof_bits<typename IteratorA::Element>::value *
|
||||
IteratorA::ThreadMap::kElementsPerAccess /
|
||||
IteratorA::kAccessesPerVector / 8;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int v = 0; v < IteratorA::kAccessesPerVector; ++v) {
|
||||
cutlass::arch::cp_async_zfill<kSrcBytes, kCacheOpA>(
|
||||
dst_ptr + v, iterator_A.get(), iterator_A.valid());
|
||||
|
||||
++iterator_A;
|
||||
}
|
||||
++this->smem_iterator_A_;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Perform a threadblock-scoped matrix multiply-accumulate
|
||||
CUTLASS_DEVICE
|
||||
void operator()(
|
||||
///< problem size of GEMM
|
||||
int gemm_k_iterations,
|
||||
///< destination accumulator tile
|
||||
FragmentC &accum,
|
||||
///< iterator over A operand in global memory
|
||||
IteratorA &iterator_A,
|
||||
///< Params of global memory iterator
|
||||
typename IteratorA::Params const &iterator_a_params,
|
||||
///< iterator over B operand in global memory
|
||||
IteratorB &iterator_B,
|
||||
///< Params of global memory iterator
|
||||
typename IteratorB::Params const &iterator_b_params,
|
||||
///< initial value of accumulator
|
||||
FragmentC const &src_accum,
|
||||
/// Epilogue
|
||||
Epilogue &epilogue,
|
||||
///< Output operator
|
||||
typename Epilogue::OutputOp const &output_op,
|
||||
///< Tile iterator for destination
|
||||
typename Epilogue::OutputTileIterator &destination_iterator,
|
||||
///< Threadblock tile coordinate in GEMM (in units of threadblock tiles)
|
||||
typename Epilogue::OutputTileIterator &source_iterator,
|
||||
|
||||
int split_k_slices = 1
|
||||
) {
|
||||
|
||||
//
|
||||
// Prologue
|
||||
//
|
||||
|
||||
// Issue several complete stages
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int stage = 0; stage < Base::kStages - 1; ++stage, --gemm_k_iterations) {
|
||||
|
||||
if (stage == 0) {
|
||||
iterator_B.set_iteration_index(0);
|
||||
this->smem_iterator_B_.set_iteration_index(0);
|
||||
|
||||
// Async Copy for operand B
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int j = 0; j < Detail::AsyncCopyIterationsPerStageB; ++j) {
|
||||
typename IteratorB::AccessType *dst_ptr =
|
||||
reinterpret_cast<typename IteratorB::AccessType *>(this->smem_iterator_B_.get());
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int v = 0; v < IteratorB::kAccessesPerVector; ++v) {
|
||||
int const kSrcBytes = sizeof_bits<typename IteratorB::Element>::value *
|
||||
IteratorB::ThreadMap::kElementsPerAccess /
|
||||
IteratorB::kAccessesPerVector / 8;
|
||||
|
||||
cutlass::arch::cp_async_zfill<kSrcBytes, kCacheOpB>(
|
||||
dst_ptr + v, iterator_B.get(), iterator_B.valid());
|
||||
|
||||
++iterator_B;
|
||||
}
|
||||
|
||||
++this->smem_iterator_B_;
|
||||
}
|
||||
}
|
||||
|
||||
if(kItertorAlgorithm == conv::IteratorAlgorithm::kFixedStrideDilation){
|
||||
// Number of iterators is compilation static.
|
||||
iterator_A.set_iteration_index(0);
|
||||
this->smem_iterator_A_.set_iteration_index(0);
|
||||
|
||||
// Async Copy for operand A
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int j = 0; j < Detail::AsyncCopyIterationsPerStageA; ++j) {
|
||||
typename IteratorA::AccessType *dst_ptr =
|
||||
reinterpret_cast<typename IteratorA::AccessType *>(this->smem_iterator_A_.get());
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int v = 0; v < IteratorA::kAccessesPerVector; ++v) {
|
||||
int const kSrcBytes = sizeof_bits<typename IteratorA::Element>::value *
|
||||
IteratorA::ThreadMap::kElementsPerAccess /
|
||||
IteratorA::kAccessesPerVector / 8;
|
||||
|
||||
cutlass::arch::cp_async_zfill<kSrcBytes, kCacheOpA>(
|
||||
dst_ptr + v, iterator_A.get(), iterator_A.valid());
|
||||
|
||||
++iterator_A;
|
||||
}
|
||||
|
||||
++this->smem_iterator_A_;
|
||||
}
|
||||
|
||||
} else {
|
||||
// Number of iterators is a runtime value.
|
||||
iterator_A.set_iteration_index(0);
|
||||
this->smem_iterator_A_.set_iteration_num(iterator_A.get_iteration_num());
|
||||
this->smem_iterator_A_.set_iteration_index(0);
|
||||
|
||||
|
||||
// Async Copy for operand A
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int j = 0; j < iterator_A.get_iteration_num(); ++j) {
|
||||
typename IteratorA::AccessType *dst_ptr =
|
||||
reinterpret_cast<typename IteratorA::AccessType *>(this->smem_iterator_A_.get());
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int v = 0; v < IteratorA::kAccessesPerVector; ++v) {
|
||||
int const kSrcBytes = sizeof_bits<typename IteratorA::Element>::value *
|
||||
IteratorA::ThreadMap::kElementsPerAccess /
|
||||
IteratorA::kAccessesPerVector / 8;
|
||||
|
||||
cutlass::arch::cp_async_zfill<kSrcBytes, kCacheOpA>(
|
||||
dst_ptr + v, iterator_A.get(), iterator_A.valid());
|
||||
|
||||
++iterator_A;
|
||||
}
|
||||
|
||||
++this->smem_iterator_A_;
|
||||
}
|
||||
}
|
||||
|
||||
// Move to the next stage
|
||||
iterator_A.advance();
|
||||
|
||||
this->smem_iterator_A_.add_tile_offset({1, 0});
|
||||
|
||||
// Inserts a fence to group cp.async instructions into stages.
|
||||
cutlass::arch::cp_async_fence();
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////
|
||||
// Waits until kStages-2 stages have committed.
|
||||
cutlass::arch::cp_async_wait<Base::kStages - 2>();
|
||||
__syncthreads();
|
||||
|
||||
// Pair of fragments used to overlap shared memory loads and math
|
||||
// instructions
|
||||
WarpLoadedFragmentA warp_loaded_frag_A[2];
|
||||
WarpLoadedFragmentB warp_loaded_frag_B[2];
|
||||
WarpTransformedFragmentA warp_transformed_frag_A[2];
|
||||
WarpTransformedFragmentB warp_transformed_frag_B[2];
|
||||
|
||||
Operator warp_mma;
|
||||
|
||||
this->warp_tile_iterator_A_.set_kgroup_index(0);
|
||||
this->warp_tile_iterator_B_.set_kgroup_index(0);
|
||||
|
||||
this->warp_tile_iterator_A_.setup_initial_status(iterator_a_params);
|
||||
|
||||
|
||||
this->warp_tile_iterator_A_.load(warp_loaded_frag_A[0]);
|
||||
this->warp_tile_iterator_B_.load(warp_loaded_frag_B[0]);
|
||||
|
||||
++this->warp_tile_iterator_A_;
|
||||
++this->warp_tile_iterator_B_;
|
||||
|
||||
int smem_write_stage_idx = Base::kStages - 1;
|
||||
int smem_read_stage_idx = 0;
|
||||
|
||||
warp_mma.transform(warp_transformed_frag_A[0], warp_transformed_frag_B[0],
|
||||
warp_loaded_frag_A[0], warp_loaded_frag_B[0]);
|
||||
|
||||
//
|
||||
// Mainloop
|
||||
//
|
||||
|
||||
unsigned int iterations = 0;
|
||||
constexpr int inner_loop_iterations = round_up(Base::kWarpGemmIterations, 2);
|
||||
|
||||
CUTLASS_GEMM_LOOP
|
||||
for (; gemm_k_iterations > (-Base::kStages + 1);) { // Each iteration is a cta tile.
|
||||
|
||||
accum.clear();
|
||||
|
||||
//
|
||||
// Loop over GEMM K dimension
|
||||
//
|
||||
|
||||
// Computes a warp-level GEMM on data held in shared memory
|
||||
// Each "warp_mma_k" refers to a warp-level matrix multiply-accumulate
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int warp_mma_k = 0; warp_mma_k < inner_loop_iterations; ++warp_mma_k) {
|
||||
if (Base::kWarpGemmIterations % 2 == 0 || warp_mma_k + 1 != Base::kWarpGemmIterations) {
|
||||
// Load warp-level tiles from shared memory, wrapping to k offset if
|
||||
// this is the last group as the case may be.
|
||||
|
||||
this->warp_tile_iterator_A_.set_kgroup_index((warp_mma_k + 1) % Shape::kK);
|
||||
this->warp_tile_iterator_B_.set_kgroup_index((warp_mma_k + 1) % Shape::kK);
|
||||
|
||||
this->warp_tile_iterator_A_.load(warp_loaded_frag_A[(warp_mma_k + 1) % 2]);
|
||||
this->warp_tile_iterator_B_.load(warp_loaded_frag_B[(warp_mma_k + 1) % 2]);
|
||||
|
||||
++this->warp_tile_iterator_A_;
|
||||
++this->warp_tile_iterator_B_;
|
||||
}
|
||||
|
||||
if (warp_mma_k > 0)
|
||||
warp_mma.transform(warp_transformed_frag_A[warp_mma_k % 2],
|
||||
warp_transformed_frag_B[warp_mma_k % 2],
|
||||
warp_loaded_frag_A[warp_mma_k % 2],
|
||||
warp_loaded_frag_B[warp_mma_k % 2]);
|
||||
|
||||
// Issue global->shared copies for the next stage
|
||||
int group_start_iteration_A, group_start_iteration_B;
|
||||
|
||||
if (warp_mma_k == 0) {
|
||||
group_start_iteration_A = 0;
|
||||
group_start_iteration_B = 0;
|
||||
copy_tiles_and_advance(
|
||||
iterator_A, iterator_B, group_start_iteration_A, group_start_iteration_B);
|
||||
}
|
||||
|
||||
if (warp_mma_k < Base::kWarpGemmIterations) {
|
||||
warp_mma(
|
||||
accum,
|
||||
warp_transformed_frag_A[warp_mma_k % 2],
|
||||
warp_transformed_frag_B[warp_mma_k % 2],
|
||||
accum
|
||||
);
|
||||
}
|
||||
|
||||
if (warp_mma_k + 1 == inner_loop_iterations)
|
||||
warp_mma.transform(warp_transformed_frag_A[(warp_mma_k + 1) % 2],
|
||||
warp_transformed_frag_B[(warp_mma_k + 1) % 2],
|
||||
warp_loaded_frag_A[(warp_mma_k + 1) % 2],
|
||||
warp_loaded_frag_B[(warp_mma_k + 1) % 2]);
|
||||
|
||||
if (warp_mma_k + 2 == inner_loop_iterations) {
|
||||
// Inserts a fence to group cp.async instructions into stages.
|
||||
cutlass::arch::cp_async_fence();
|
||||
|
||||
// Waits until kStages-2 stages of cp.async have committed
|
||||
arch::cp_async_wait<Base::kStages - 2>();
|
||||
__syncthreads();
|
||||
|
||||
// Move to the next cta
|
||||
iterator_A.advance();
|
||||
|
||||
this->smem_iterator_A_.add_tile_offset({1, 0});
|
||||
|
||||
// Add negative offsets to return iterators to the 'start' of the
|
||||
// circular buffer in shared memory
|
||||
if (smem_write_stage_idx == (Base::kStages - 1)) {
|
||||
this->smem_iterator_A_.add_tile_offset({-Base::kStages, 0});
|
||||
|
||||
smem_write_stage_idx = 0;
|
||||
} else {
|
||||
++smem_write_stage_idx;
|
||||
}
|
||||
|
||||
if (smem_read_stage_idx == (Base::kStages - 1)) {
|
||||
this->warp_tile_iterator_A_.advance(- (Base::kStages-1) * iterator_A.get_load_size());
|
||||
smem_read_stage_idx = 0;
|
||||
} else {
|
||||
this->warp_tile_iterator_A_.advance(iterator_A.get_load_size());
|
||||
++smem_read_stage_idx;
|
||||
}
|
||||
|
||||
if (kItertorAlgorithm == conv::IteratorAlgorithm::kFixedStrideDilation) {
|
||||
this->warp_tile_iterator_A_.setup_initial_status(iterator_a_params);
|
||||
}
|
||||
|
||||
// goback to start position. B has no multiple stage
|
||||
this->warp_tile_iterator_B_.add_tile_offset({-Policy::kPartitionsK * Shape::kK, 0});
|
||||
|
||||
--gemm_k_iterations;
|
||||
}
|
||||
}
|
||||
|
||||
//
|
||||
// Epilogue
|
||||
//
|
||||
int32_t smem_base_offset = iterator_B.get_load_size() + (iterations % Base::kStages) * iterator_A.get_load_size();
|
||||
|
||||
destination_iterator.set_tile_index(iterations * split_k_slices);
|
||||
|
||||
source_iterator.set_tile_index(iterations * split_k_slices);
|
||||
|
||||
epilogue(output_op, destination_iterator, accum, source_iterator, smem_base_offset);
|
||||
|
||||
++iterations;
|
||||
}
|
||||
|
||||
// Insert fence and wait for all outstanding cp.async operations to commit.
|
||||
cutlass::arch::cp_async_fence();
|
||||
cutlass::arch::cp_async_wait<0>();
|
||||
__syncthreads();
|
||||
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,261 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017 - 2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
/*! \file
|
||||
\brief Templates implementing loading of convolution tiles mapped to GEMM B (filter tile)
|
||||
matrix from memory.
|
||||
|
||||
This iterator assumes TensorNHWC layout of tensors in Global Memory.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/coord.h"
|
||||
#include "cutlass/predicate_vector.h"
|
||||
#include "cutlass/tensor_ref.h"
|
||||
#include "cutlass/tensor_view.h"
|
||||
#include "cutlass/layout/pitch_linear.h"
|
||||
#include "cutlass/layout/tensor.h"
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/conv/convolution.h"
|
||||
#include "cutlass/conv/conv2d_problem_size.h"
|
||||
#include "cutlass/conv/threadblock/conv2d_params.h"
|
||||
#include "cutlass/conv/threadblock/conv2d_fprop_filter_tile_access_iterator_analytic.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace conv {
|
||||
namespace threadblock {
|
||||
|
||||
template <typename Shape_,
|
||||
typename Element_,
|
||||
typename Layout_,
|
||||
typename ThreadMap_,
|
||||
typename AccessType_ = cutlass::AlignedArray<Element_, ThreadMap_::kElementsPerAccess> >
|
||||
class DepthwiseFpropFilterDirectConvTileAccessIteratorOptimized {
|
||||
public:
|
||||
//
|
||||
// Types
|
||||
//
|
||||
|
||||
using Shape = Shape_;
|
||||
using Element = Element_;
|
||||
using Layout = Layout_;
|
||||
using ThreadMap = ThreadMap_;
|
||||
using AccessType = AccessType_;
|
||||
using TensorRef = cutlass::TensorRef<Element, Layout>;
|
||||
using TensorCoord = typename Layout::TensorCoord;
|
||||
using Index = typename Layout::Index;
|
||||
using LongIndex = typename Layout::LongIndex;
|
||||
static IteratorAlgorithm const kIteratorAlgorithm = conv::IteratorAlgorithm::kOptimized;
|
||||
static StrideSupport const kStrideSupport = conv::StrideSupport::kStrided;
|
||||
static int const kConvDim = 2;
|
||||
using ConvProblemSize = typename conv::Conv2dProblemSize;
|
||||
|
||||
static int const kAccessesPerVector = ThreadMap::kElementsPerAccess / AccessType::kElements;
|
||||
|
||||
static int const kFilterSize = ThreadMap::Iterations::kCount * ThreadMap::kElementsPerAccess * ThreadMap::kThreads *
|
||||
sizeof_bits<Element>::value / 8;
|
||||
|
||||
static_assert(!(ThreadMap::kElementsPerAccess % AccessType::kElements),
|
||||
"Vectors implied by the thread map must be divisible by the access type.");
|
||||
|
||||
//
|
||||
// Simplifying assertions
|
||||
//
|
||||
static_assert(ThreadMap::Iterations::kContiguous == 1,
|
||||
"Require Iterations::kContiguous == 1");
|
||||
|
||||
//
|
||||
// Parameters structure
|
||||
//
|
||||
using Params = Depthwise2dFpropDirectConvFilterIteratorParams<Layout>;
|
||||
|
||||
protected:
|
||||
|
||||
Conv2dProblemSize const &problem_size_;
|
||||
Params const ¶ms_;
|
||||
LongIndex iteration_contiguous_;
|
||||
LongIndex iteration_strided_;
|
||||
LongIndex iteration_vector_;
|
||||
char const *pointer_;
|
||||
|
||||
int filter_k_;
|
||||
int offset_trs_[ThreadMap::Iterations::kStrided];
|
||||
|
||||
public:
|
||||
|
||||
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
DepthwiseFpropFilterDirectConvTileAccessIteratorOptimized(
|
||||
Params const ¶ms,
|
||||
Conv2dProblemSize const &problem_size,
|
||||
Element const *ptr,
|
||||
int thread_idx,
|
||||
MatrixCoord const &threadblock_offset = MatrixCoord()
|
||||
):
|
||||
params_(params),
|
||||
problem_size_(problem_size),
|
||||
pointer_(reinterpret_cast<char const *>(ptr)),
|
||||
filter_k_(0) {
|
||||
|
||||
layout::PitchLinearCoord thread_coord = ThreadMap::initial_offset(thread_idx);
|
||||
|
||||
filter_k_ = threadblock_offset.column() + thread_coord.contiguous();
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int s = 0; s < ThreadMap::Iterations::kStrided; ++s) {
|
||||
offset_trs_[s] = threadblock_offset.row() + thread_coord.strided() + s * ThreadMap::Delta::kStrided;
|
||||
}
|
||||
|
||||
set_iteration_index(0);
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static Params getParams(Conv2dProblemSize const &problem_size, Layout const &layout) {
|
||||
return Params(problem_size, layout, {Shape::kRow, Shape::kColumn}, kFilterSize);
|
||||
}
|
||||
|
||||
/// Overrides the internal iteration index
|
||||
CUTLASS_HOST_DEVICE
|
||||
void set_iteration_index(Index index) {
|
||||
iteration_vector_ = index % kAccessesPerVector;
|
||||
int residual_access = index / kAccessesPerVector;
|
||||
iteration_contiguous_ = residual_access % ThreadMap::Iterations::kContiguous;
|
||||
iteration_strided_ = residual_access / ThreadMap::Iterations::kContiguous;
|
||||
}
|
||||
|
||||
/// Adds a pointer offset in units of Element
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_pointer_offset(LongIndex pointer_offset) {
|
||||
pointer_ += pointer_offset * 8 / sizeof_bits<Element>::value;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
void advance() {
|
||||
// Do nothing because the filter is persistent in the SMEM
|
||||
}
|
||||
|
||||
/// Returns the coordinate in the filter tensor W that is currently pointed to
|
||||
/// by the iterator.
|
||||
CUTLASS_HOST_DEVICE
|
||||
TensorCoord at() const {
|
||||
|
||||
int k = filter_k_ + iteration_vector_ * AccessType::kElements;
|
||||
int trs = offset_trs_[iteration_strided_];
|
||||
|
||||
return TensorCoord(k, trs, 0 , 0); // As a 2D-matrix
|
||||
}
|
||||
|
||||
/// Returns true if the current coordinate is within the activations tensor W
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool valid() const {
|
||||
|
||||
TensorCoord coord = at();
|
||||
|
||||
return coord.n() < problem_size_.K &&
|
||||
coord.h() < Shape::kColumn;
|
||||
}
|
||||
|
||||
/// Returns a pointer to the vector starting at the current coordinate
|
||||
CUTLASS_HOST_DEVICE
|
||||
AccessType const *get() const {
|
||||
TensorCoord coord = at();
|
||||
int64_t offset = coord.n();
|
||||
if (params_.is_convolution) {
|
||||
offset += (Shape::kColumn - coord.h() - 1)* problem_size_.K;
|
||||
} else {
|
||||
offset += coord.h() * problem_size_.K;
|
||||
}
|
||||
|
||||
return reinterpret_cast<AccessType const *>(pointer_ +
|
||||
offset * sizeof_bits<Element>::value / 8);
|
||||
}
|
||||
|
||||
/// Increments to the next memory access
|
||||
CUTLASS_HOST_DEVICE
|
||||
DepthwiseFpropFilterDirectConvTileAccessIteratorOptimized &operator++() {
|
||||
++iteration_vector_;
|
||||
if (iteration_vector_ < kAccessesPerVector) {
|
||||
return *this;
|
||||
}
|
||||
iteration_vector_ = 0;
|
||||
|
||||
++iteration_contiguous_;
|
||||
if (iteration_contiguous_ < ThreadMap::Iterations::kContiguous) {
|
||||
return *this;
|
||||
}
|
||||
iteration_contiguous_ = 0;
|
||||
|
||||
++iteration_strided_;
|
||||
if (iteration_strided_ < ThreadMap::Iterations::kStrided) {
|
||||
return *this;
|
||||
}
|
||||
iteration_strided_ = 0;
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Determines the filter size loaded by iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
int get_load_size() {
|
||||
return kFilterSize;
|
||||
}
|
||||
|
||||
/// Determines whether the Implicit GEMM can execute the given problem.
|
||||
CUTLASS_HOST_DEVICE
|
||||
static Status can_implement(Conv2dProblemSize const &problem_size) {
|
||||
|
||||
// check alignment constraint on iterator's contiguous dimension
|
||||
if (problem_size.K % AccessType::kElements) {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
// check whether runtime filter size is same as templated filter size.
|
||||
if ((problem_size.R * problem_size.S) != Shape::kColumn) {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace conv
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
229
include/cutlass/conv/threadblock/depthwise_mma_base.h
Normal file
229
include/cutlass/conv/threadblock/depthwise_mma_base.h
Normal file
@@ -0,0 +1,229 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017 - 2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
/*! \file
|
||||
\brief Template for a directconv threadblock-scoped Depthwise kernel.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/aligned_buffer.h"
|
||||
#include "cutlass/arch/memory.h"
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/matrix_shape.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace conv {
|
||||
namespace threadblock {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Policy object describing MmaTensorOp
|
||||
template <
|
||||
/// Warp-level GEMM operator (concept: gemm::warp::Mma)
|
||||
typename Operator_,
|
||||
/// Padding used for A operand in shared memory (concept: MatrixShape)
|
||||
typename SmemPaddingA_,
|
||||
/// Padding used for B operand in shared memory (concept: MatrixShape)
|
||||
typename SmemPaddingB_,
|
||||
///
|
||||
typename ThreadMapA_,
|
||||
///
|
||||
typename ThreadMapB_,
|
||||
/// Number of partitions of K dimension of GEMM
|
||||
int PartitionsK = 1>
|
||||
struct DepthwiseDirectConvMmaPolicy {
|
||||
/// Warp-level GEMM operator (concept: gemm::warp::MmaTensorOp or gemm::warp::MmaSimt)
|
||||
using Operator = Operator_;
|
||||
|
||||
/// Padding used for A operand in shared memory
|
||||
using SmemPaddingA = SmemPaddingA_;
|
||||
|
||||
/// Padding used for B operand in shared memory
|
||||
using SmemPaddingB = SmemPaddingB_;
|
||||
|
||||
using ThreadMapA = ThreadMapA_;
|
||||
using ThreadMapB = ThreadMapB_;
|
||||
|
||||
/// Number of partitions of K dimension
|
||||
static int const kPartitionsK = PartitionsK;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Structure to compute the matrix product targeting CUDA cores and SIMT math
|
||||
/// instructions.
|
||||
template <
|
||||
/// Size of the Gemm problem - concept: gemm::GemmShape<>
|
||||
typename Shape_,
|
||||
/// Policy describing tuning details (concept: MmaPolicy)
|
||||
typename Policy_,
|
||||
/// Number of stages,
|
||||
int Stages,
|
||||
/// Used for partial specialization
|
||||
typename Enable = bool>
|
||||
class DepthwiseDirectConvMmaBase {
|
||||
public:
|
||||
///< Size of the Gemm problem - concept: gemm::GemmShape<>
|
||||
using Shape = Shape_;
|
||||
|
||||
///< Policy describing tuning details
|
||||
using Policy = Policy_;
|
||||
|
||||
//
|
||||
// Dependent types
|
||||
//
|
||||
|
||||
/// Warp-level Mma
|
||||
using Operator = typename Policy::Operator;
|
||||
|
||||
/// Shape describing the overall GEMM computed from shared memory
|
||||
/// by each warp.
|
||||
using WarpGemm = typename Policy::Operator::Shape;
|
||||
|
||||
/// Shape describing the number of warps filling the CTA
|
||||
using WarpCount = cutlass::gemm::
|
||||
GemmShape<Shape::kM / WarpGemm::kM, Shape::kN / WarpGemm::kN, Shape::kK / WarpGemm::kK>;
|
||||
|
||||
/// Number of warp-level GEMM oeprations
|
||||
/// kWarpGemmIterations could be even and odd.
|
||||
static int const kWarpGemmIterations = (WarpGemm::kK / Operator::Policy::MmaShape::kK);
|
||||
|
||||
/// Number of stages
|
||||
static int const kStages = Stages;
|
||||
|
||||
/// Tensor reference to the A operand
|
||||
using TensorRefA = TensorRef<typename Operator::ElementA, typename Operator::LayoutA>;
|
||||
|
||||
/// Tensor reference to the B operand
|
||||
using TensorRefB = TensorRef<typename Operator::ElementB, typename Operator::LayoutB>;
|
||||
|
||||
static_assert(kWarpGemmIterations > 1,
|
||||
"The pipelined structure requires at least two warp-level "
|
||||
"GEMM operations.");
|
||||
|
||||
//
|
||||
// Nested structs
|
||||
//
|
||||
|
||||
/// Shared storage object needed by threadblock-scoped GEMM
|
||||
class SharedStorage {
|
||||
public:
|
||||
//
|
||||
// Type definitions
|
||||
//
|
||||
|
||||
/// Shape of the A matrix operand in shared memory
|
||||
using ShapeA = MatrixShape<1, // Not determined at compile-time :(
|
||||
Shape::kN + Policy::SmemPaddingA::kRow>;
|
||||
|
||||
/// Shape of the B matrix operand in shared memory
|
||||
using ShapeB = MatrixShape<Policy::ThreadMapB::StorageShape::kStrided +
|
||||
Policy::SmemPaddingB::kRow, // filter_rs_size
|
||||
Policy::ThreadMapB::StorageShape::kContiguous +
|
||||
Policy::SmemPaddingB::kColumn>; // Tile N = 64?
|
||||
|
||||
public:
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
// Let persistent B matrix in front of dynamic matrix A
|
||||
/// Buffer for B operand
|
||||
AlignedBuffer<typename Operator::ElementB, ShapeB::kCount> operand_B;
|
||||
|
||||
/// Buffer for A operand
|
||||
/// Not be determined at compile-time -- Just to get a Smem start address.
|
||||
AlignedBuffer<typename Operator::ElementA, 1> operand_A;
|
||||
public:
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
/// Returns a layout object for the A matrix
|
||||
CUTLASS_DEVICE
|
||||
static typename Operator::LayoutA LayoutA() {
|
||||
return Operator::LayoutA::packed({ShapeA::kRow, ShapeA::kColumn});
|
||||
}
|
||||
|
||||
/// Returns a layout object for the B matrix
|
||||
CUTLASS_HOST_DEVICE
|
||||
static typename Operator::LayoutB LayoutB() {
|
||||
return Operator::LayoutB::packed({ShapeB::kRow, ShapeB::kColumn});
|
||||
}
|
||||
|
||||
/// Returns a TensorRef to the A operand
|
||||
CUTLASS_HOST_DEVICE
|
||||
TensorRefA operand_A_ref() { return TensorRefA{operand_A.data(), LayoutA()}; }
|
||||
|
||||
/// Returns a TensorRef to the B operand
|
||||
CUTLASS_HOST_DEVICE
|
||||
TensorRefB operand_B_ref() { return TensorRefB{operand_B.data(), LayoutB()}; }
|
||||
};
|
||||
|
||||
protected:
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Iterator to load a warp-scoped tile of A operand from shared memory
|
||||
typename Operator::IteratorA warp_tile_iterator_A_;
|
||||
|
||||
/// Iterator to load a warp-scoped tile of B operand from shared memory
|
||||
typename Operator::IteratorB warp_tile_iterator_B_;
|
||||
|
||||
public:
|
||||
/// Construct from tensor references
|
||||
CUTLASS_DEVICE
|
||||
DepthwiseDirectConvMmaBase(
|
||||
///< Shared storage needed for internal use by threadblock-scoped GEMM
|
||||
SharedStorage &shared_storage,
|
||||
///< ID within the threadblock
|
||||
int thread_idx,
|
||||
///< ID of warp
|
||||
int warp_idx,
|
||||
///< ID of each thread within a warp
|
||||
int lane_idx)
|
||||
: warp_tile_iterator_A_(shared_storage.operand_A_ref(), lane_idx),
|
||||
warp_tile_iterator_B_(shared_storage.operand_B_ref(), lane_idx) {}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace conv
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -44,11 +44,17 @@
|
||||
#include "cutlass/matrix_shape.h"
|
||||
|
||||
#include "cutlass/gemm/warp/mma.h"
|
||||
|
||||
#include "cutlass/conv/convolution.h"
|
||||
#include "cutlass/conv/warp/mma_depthwise_simt.h"
|
||||
|
||||
#include "cutlass/gemm/threadblock/mma_pipelined.h"
|
||||
#include "cutlass/gemm/threadblock/mma_singlestage.h"
|
||||
|
||||
#include "cutlass/gemm/threadblock/mma_base.h"
|
||||
#include "cutlass/conv/warp/mma_depthwise_simt.h"
|
||||
#include "cutlass/conv/threadblock/depthwise_mma_base.h"
|
||||
|
||||
#include "cutlass/transform/threadblock/regular_tile_access_iterator_pitch_linear_direct_conv.h"
|
||||
|
||||
#include "cutlass/arch/cache_operation.h"
|
||||
|
||||
@@ -58,6 +64,95 @@ namespace cutlass {
|
||||
namespace conv {
|
||||
namespace threadblock {
|
||||
|
||||
namespace detail {
|
||||
//
|
||||
// Convert a WarpShapeM which is the whole tile of elements into the number of elements (2D) held by
|
||||
// each partitions within warp.
|
||||
// The goal is for each thread's tile of elements to be as square as
|
||||
// possible for performance (4x4 will be faster than 2x8).
|
||||
template<int WarpShapeM, // The number of elements (1D) contained in the entire warp
|
||||
int WarpNumThreadsM> // The number of partitions within the warp
|
||||
struct SimtWarpShape {
|
||||
// kP * kQ * WarpNumThreadsM = WarpShapeM
|
||||
// If needed, enable more specializations.
|
||||
};
|
||||
template <>
|
||||
struct SimtWarpShape<4, 4> {
|
||||
static constexpr int kP = 1;
|
||||
static constexpr int kQ = 1;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct SimtWarpShape<4, 2> {
|
||||
static constexpr int kP = 2;
|
||||
static constexpr int kQ = 1;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct SimtWarpShape<4, 1> {
|
||||
static constexpr int kP = 2;
|
||||
static constexpr int kQ = 2;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct SimtWarpShape<8, 1> {
|
||||
static constexpr int kP = 2;
|
||||
static constexpr int kQ = 4;
|
||||
};
|
||||
template <>
|
||||
struct SimtWarpShape<8, 2> {
|
||||
static constexpr int kP = 2;
|
||||
static constexpr int kQ = 2;
|
||||
};
|
||||
template <>
|
||||
struct SimtWarpShape<8, 4> {
|
||||
static constexpr int kP = 1;
|
||||
static constexpr int kQ = 2;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct SimtWarpShape<16, 1> {
|
||||
static constexpr int kP = 4;
|
||||
static constexpr int kQ = 4;
|
||||
};
|
||||
template <>
|
||||
struct SimtWarpShape<16, 2> {
|
||||
static constexpr int kP = 2;
|
||||
static constexpr int kQ = 4;
|
||||
};
|
||||
template <>
|
||||
struct SimtWarpShape<16, 4> {
|
||||
static constexpr int kP = 2;
|
||||
static constexpr int kQ = 2;
|
||||
};
|
||||
|
||||
template <int WarpNumThreadsM>
|
||||
struct SimtWarpShape<25, WarpNumThreadsM> {
|
||||
static_assert(WarpNumThreadsM == 1, "WarpShapeM could not be evenly splited by threads");
|
||||
static constexpr int kP = 5;
|
||||
static constexpr int kQ = 5;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct SimtWarpShape<32, 1> {
|
||||
static constexpr int kP = 4;
|
||||
static constexpr int kQ = 8;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct SimtWarpShape<32, 2> {
|
||||
static constexpr int kP = 4;
|
||||
static constexpr int kQ = 4;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct SimtWarpShape<32, 4> {
|
||||
static constexpr int kP = 2;
|
||||
static constexpr int kQ = 4;
|
||||
};
|
||||
|
||||
} // namespace detail
|
||||
|
||||
template <
|
||||
/// Shape of threadblock-scoped matrix multiply operator
|
||||
typename Shape,
|
||||
@@ -114,6 +209,74 @@ struct DepthwiseMmaCoreWithLaneAccessSize;
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
/// Shape of threadblock-scoped matrix multiply operator
|
||||
typename Shape,
|
||||
/// Shape of threadblock-scoped output tile
|
||||
typename ThreadBlockOutputShape,
|
||||
/// Shape of filter shape per threadblock
|
||||
typename FilterShape,
|
||||
/// Shape of warp-level matrix multiply operator
|
||||
typename WarpShape,
|
||||
/// Shape of one matrix production operation (concept: GemmShape)
|
||||
typename InstructionShape,
|
||||
/// Element data type of A operand
|
||||
typename ElementA,
|
||||
/// Layout of operand A
|
||||
typename LayoutA,
|
||||
/// Element data type of B operand
|
||||
typename ElementB,
|
||||
/// Layout of operand B
|
||||
typename LayoutB,
|
||||
/// Data type of accumulator
|
||||
typename ElementC,
|
||||
/// Layout of accumulator
|
||||
typename LayoutC,
|
||||
/// Indicates type of math operator (arch::OpClassSimt or arch::OpClassTensorOp)
|
||||
typename OperatorClass,
|
||||
/// Size of a warp-scoped per thread access
|
||||
int kLaneAccessSizeA_ = 0,
|
||||
/// Size of a warp-scoped per thread access
|
||||
int kLaneAccessSizeB_ = 0,
|
||||
/// Number of stages
|
||||
int Stages = 2,
|
||||
/// Operation performed by MMA
|
||||
typename Operator = typename platform::conditional<
|
||||
(platform::is_same<OperatorClass,
|
||||
cutlass::arch::OpClassTensorOp>::value) &&
|
||||
(platform::is_same<ElementA, int8_t>::value ||
|
||||
platform::is_same<ElementA, int4b_t>::value ||
|
||||
platform::is_same<ElementA, uint8_t>::value ||
|
||||
platform::is_same<ElementA, uint4b_t>::value),
|
||||
cutlass::arch::OpMultiplyAddSaturate,
|
||||
cutlass::arch::OpMultiplyAdd>::type,
|
||||
/// Iterator algo type
|
||||
conv::IteratorAlgorithm IteratorAlgorithm = IteratorAlgorithm::kAnalytic,
|
||||
/// Stride ( MatrixShape<Height, Width> )
|
||||
typename StrideShape = cutlass::MatrixShape<-1, -1>,
|
||||
/// Dilation ( MatrixShape<Height, Width> )
|
||||
typename DilationShape = cutlass::MatrixShape<-1, -1>,
|
||||
/// Activation Shape loaded by threadblock
|
||||
typename ActivationShape = cutlass::conv::TensorNHWCShape<-1,-1,-1,-1>,
|
||||
/// Store the accumulators in row major or column major. Row major is used
|
||||
/// when output layout is interleaved.
|
||||
bool AccumulatorsInRowMajor = false,
|
||||
/// Cache operation of operand A
|
||||
cutlass::arch::CacheOperation::Kind CacheOpA =
|
||||
cutlass::arch::CacheOperation::Global,
|
||||
/// Cache operation of operand B
|
||||
cutlass::arch::CacheOperation::Kind CacheOpB =
|
||||
cutlass::arch::CacheOperation::Global,
|
||||
/// per-element transformation for elements of A
|
||||
ComplexTransform TransformA = ComplexTransform::kNone,
|
||||
/// per-element transformation for elements of B
|
||||
ComplexTransform TransformB = ComplexTransform::kNone,
|
||||
bool IsComplex = false // (is_complex<ElementA>::value || is_complex<ElementB>::value)
|
||||
>
|
||||
struct DepthwiseDirectConvMmaCoreWithLaneAccessSize;
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
/// Shape of threadblock-scoped matrix multiply operator
|
||||
typename Shape,
|
||||
@@ -332,6 +495,458 @@ struct DepthwiseMmaCoreWithLaneAccessSize<Shape_,
|
||||
>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization:
|
||||
///
|
||||
/// A: row-major
|
||||
/// B: row-major
|
||||
/// Operator: simt class
|
||||
///
|
||||
/// This uses the default warp-level operator given tile sizes
|
||||
template <
|
||||
/// Shape of threadblock-scoped matrix multiply operator (concept:
|
||||
/// GemmShape)
|
||||
typename Shape_,
|
||||
/// Shape of threadblock-scoped output tile (concept: TensorNHWCShape)
|
||||
typename ThreadBlockOutputShape_,
|
||||
/// Shape of filter shape per threadblock
|
||||
typename FilterShape_,
|
||||
/// Shape of warp-level matrix multiply operator (concept: GemmShape)
|
||||
typename WarpShape_,
|
||||
/// Data type of A operand
|
||||
typename ElementA_,
|
||||
/// Data type of B operand
|
||||
typename ElementB_,
|
||||
/// Data type of accumulator
|
||||
typename ElementC_,
|
||||
/// Layout of accumulator
|
||||
typename LayoutC_,
|
||||
/// Size of a warp-scoped per thread access
|
||||
int kLaneAccessSizeA_,
|
||||
/// Number of stages
|
||||
int Stages_,
|
||||
/// Operation performed by GEMM
|
||||
typename Operator_>
|
||||
struct DepthwiseDirectConvMmaCoreWithLaneAccessSize<Shape_,
|
||||
ThreadBlockOutputShape_,
|
||||
FilterShape_,
|
||||
WarpShape_,
|
||||
cutlass::gemm::GemmShape<1, 1, 1>,
|
||||
ElementA_,
|
||||
layout::RowMajor,
|
||||
ElementB_,
|
||||
layout::ColumnMajor,
|
||||
ElementC_,
|
||||
LayoutC_,
|
||||
arch::OpClassSimt,
|
||||
kLaneAccessSizeA_,
|
||||
128,
|
||||
Stages_,
|
||||
Operator_> {
|
||||
using Shape = Shape_;
|
||||
using FilterShape = FilterShape_;
|
||||
using WarpShape = WarpShape_;
|
||||
using InstructionShape = cutlass::gemm::GemmShape<1, 1, 1>;
|
||||
using ElementA = ElementA_;
|
||||
using LayoutA = layout::RowMajor;
|
||||
using ElementB = ElementB_;
|
||||
using LayoutB = layout::ColumnMajor;
|
||||
using ElementC = ElementC_;
|
||||
using LayoutC = LayoutC_;
|
||||
using OperatorClass = arch::OpClassSimt;
|
||||
|
||||
static int const kLaneAccessSizeB = 128;
|
||||
|
||||
// Divisility requirements
|
||||
static_assert( kLaneAccessSizeB > 0,
|
||||
"Size of a warp-scoped per thread access should be larger then ZERO" );
|
||||
|
||||
/// Default Operator
|
||||
using Operator = Operator_;
|
||||
|
||||
/// Number of warps present
|
||||
using WarpCount = cutlass::gemm::GemmShape<
|
||||
Shape::kM / WarpShape::kM,
|
||||
Shape::kN / WarpShape::kN,
|
||||
1
|
||||
>;
|
||||
|
||||
// Divisility requirements
|
||||
static_assert(
|
||||
!(Shape::kM % WarpShape::kM) &&
|
||||
!(Shape::kN % WarpShape::kN),
|
||||
"Threadblock-scoped GEMM should be divisible by warp-scoped GEMM size."
|
||||
);
|
||||
|
||||
/// Number of threads per warp
|
||||
static int const kWarpSize = cutlass::gemm::warp::WarpSize<arch::OpClassSimt>::value;
|
||||
|
||||
/// Number of threads total
|
||||
static int const kThreads = WarpCount::kCount * kWarpSize;
|
||||
|
||||
// For Gmem load
|
||||
static int const kElementsPerAccessA = 128 / sizeof_bits<ElementA>::value;
|
||||
static int const kElementsPerAccessB = 128 / sizeof_bits<ElementB>::value;
|
||||
|
||||
//
|
||||
// Shared memory layouts
|
||||
//
|
||||
|
||||
using SmemLayoutA = layout::RowMajor;
|
||||
using SmemLayoutB = layout::RowMajor;
|
||||
|
||||
|
||||
//
|
||||
// Iterators to write to shared memory
|
||||
//
|
||||
|
||||
/// ThreadMap of iterator A
|
||||
using IteratorThreadMapA = transform::PitchLinearStripminedThreadMap<
|
||||
layout::PitchLinearShape<Shape::kN, 1>, // Set kStrided = 1 because activation shape is runtime value.
|
||||
kThreads,
|
||||
kElementsPerAccessA
|
||||
>;
|
||||
|
||||
/// ThreadMap of iterator A
|
||||
using SmemThreadMapA = IteratorThreadMapA;
|
||||
|
||||
/// Shared memory iterator to A operand
|
||||
using SmemIteratorA = transform::threadblock::RegularTileAccessIteratorDirectConv<
|
||||
MatrixShape<1, Shape::kN>, // set kRow is 1 because it is a runtime value
|
||||
ElementA,
|
||||
SmemLayoutA,
|
||||
0,
|
||||
SmemThreadMapA, // was IteratorThreadMapA
|
||||
true // Dynamic iterations.
|
||||
>;
|
||||
|
||||
/// ThreadMap of iterator B
|
||||
using IteratorThreadMapB = transform::PitchLinearStripminedThreadMap<
|
||||
layout::PitchLinearShape<Shape::kN, FilterShape::kCount>,
|
||||
kThreads,
|
||||
kElementsPerAccessB
|
||||
>;
|
||||
|
||||
/// Transpose the ThreadMap of iterator B
|
||||
using SmemThreadMapB = IteratorThreadMapB;
|
||||
|
||||
/// Shared memory iterator to B operand
|
||||
using SmemIteratorB = transform::threadblock::RegularTileAccessIteratorDirectConv<
|
||||
MatrixShape<FilterShape::kCount, Shape::kN>,
|
||||
ElementB,
|
||||
SmemLayoutB,
|
||||
0,
|
||||
SmemThreadMapB, // was IteratorThreadMapB
|
||||
false // static iterations.
|
||||
>;
|
||||
|
||||
//
|
||||
// Warp-level matrix multiply operator
|
||||
//
|
||||
// Groups per threads
|
||||
// Fp32: 2 groups
|
||||
// Fp16: 2 groups
|
||||
static const int GroupsPerThread = sizeof(ElementB) > 1 ? 2 : 4;
|
||||
// Define the warp-level op
|
||||
static const int WarpNumThreadsN = cutlass::const_min(WarpShape::kN / GroupsPerThread, kWarpSize);
|
||||
static const int WarpNumThreadsM = kWarpSize / WarpNumThreadsN;
|
||||
|
||||
static_assert(!(WarpShape::kM % WarpNumThreadsM) && !(WarpShape::kN % WarpNumThreadsN),
|
||||
"WarpShape must be divisible by ThreadTile shape.");
|
||||
|
||||
// Get output P, Q per thread
|
||||
static const int TileP = cutlass::conv::threadblock::detail::SimtWarpShape<WarpShape::kM, WarpNumThreadsM>::kP;
|
||||
static const int TileQ = cutlass::conv::threadblock::detail::SimtWarpShape<WarpShape::kM, WarpNumThreadsM>::kQ;
|
||||
|
||||
static const int LaneLayout = 1;
|
||||
static const int numElementsB = kLaneAccessSizeB / sizeof_bits<ElementB>::value;
|
||||
static const int LaneN = cutlass::const_min(numElementsB, WarpShape::kN / WarpNumThreadsN);
|
||||
|
||||
// Define the output tile computed by each thread
|
||||
using ThreadOutputShape = cutlass::conv::TensorNHWCShape<1, TileP, TileQ, LaneN>;
|
||||
|
||||
// Fetch the channel with same access size
|
||||
static const int LaneM = LaneN;
|
||||
|
||||
// No paddings
|
||||
static int const kPaddingM = 0;
|
||||
static int const kPaddingN = 0;
|
||||
|
||||
static_assert(!(kPaddingM % LaneM) && !(kPaddingN % LaneN),
|
||||
"Padding must be divisible by Lane");
|
||||
|
||||
// these should have max of thread tile also
|
||||
using LaneMmaShape = cutlass::gemm::GemmShape<
|
||||
LaneM,
|
||||
LaneN,
|
||||
1>;
|
||||
|
||||
using Policy = cutlass::gemm::warp::MmaSimtPolicy<
|
||||
cutlass::MatrixShape<WarpNumThreadsM, WarpNumThreadsN>, // WarpShape
|
||||
cutlass::layout::RowMajorInterleaved<LaneLayout>, // LaneLayout
|
||||
LaneMmaShape
|
||||
>;
|
||||
|
||||
using MmaWarpSimt = cutlass::conv::warp::MmaDepthwiseDirectConvSimt<
|
||||
WarpShape, /// Size of the Gemm problem - concept: gemm::GemmShape<>
|
||||
FilterShape, /// Shape of filter shape per threadblock - concept: gemm::GemmShape<Depth, Height, Width>
|
||||
ThreadOutputShape, /// Size of the output tile computed by thread - concept: conv::TensorNHWCShape<>
|
||||
ThreadBlockOutputShape_, /// Size of the output tile computed by threadblock - concept: conv::TensorNHWCShape<>
|
||||
ElementA, /// Data type of A elements
|
||||
SmemLayoutA, /// Layout of A matrix (concept: MatrixLayout)
|
||||
ElementB, /// Data type of B elements
|
||||
SmemLayoutB, /// Layout of B matrix (concept: MatrixLayout)
|
||||
ElementC, /// Element type of C matrix
|
||||
LayoutC, /// Layout of C matrix (concept: MatrixLayout)
|
||||
Policy /// Policy describing warp-level MmaSimtOp (concept: MmaSimtOp policy)
|
||||
>;
|
||||
|
||||
/// Policy used to define MmaPipelined
|
||||
using MmaPolicy = cutlass::conv::threadblock::DepthwiseDirectConvMmaPolicy<
|
||||
MmaWarpSimt,
|
||||
MatrixShape<kPaddingM, 0>, // skew for A matrix to avoid SMEM bank conflicts
|
||||
MatrixShape<0, kPaddingN>, // skew for B matrix to avoid SMEM bank conflicts
|
||||
IteratorThreadMapA,
|
||||
IteratorThreadMapB,
|
||||
WarpCount::kK
|
||||
>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization:
|
||||
///
|
||||
/// A: row-major
|
||||
/// B: row-major
|
||||
/// Operator: simt class
|
||||
///
|
||||
/// This uses the default warp-level operator given tile sizes
|
||||
template <
|
||||
/// Shape of threadblock-scoped matrix multiply operator (concept:
|
||||
/// GemmShape)
|
||||
typename Shape_,
|
||||
/// Shape of threadblock-scoped output tile (concept: TensorNHWCShape)
|
||||
typename ThreadBlockOutputShape_,
|
||||
/// Shape of filter shape per threadblock
|
||||
typename FilterShape_,
|
||||
/// Shape of warp-level matrix multiply operator (concept: GemmShape)
|
||||
typename WarpShape_,
|
||||
/// Data type of A operand
|
||||
typename ElementA_,
|
||||
/// Data type of B operand
|
||||
typename ElementB_,
|
||||
/// Data type of accumulator
|
||||
typename ElementC_,
|
||||
/// Layout of accumulator
|
||||
typename LayoutC_,
|
||||
/// Size of a warp-scoped per thread access
|
||||
int kLaneAccessSizeA_,
|
||||
/// Number of stages
|
||||
int Stages_,
|
||||
/// Operation performed by GEMM
|
||||
typename Operator_,
|
||||
/// Stride ( MatrixShape<Height, Width> )
|
||||
typename StrideShape_,
|
||||
/// Dilation ( MatrixShape<Height, Width> )
|
||||
typename DilationShape_,
|
||||
/// Activation Shape loaded by threadblock
|
||||
typename ActivationShape_>
|
||||
struct DepthwiseDirectConvMmaCoreWithLaneAccessSize<Shape_,
|
||||
ThreadBlockOutputShape_,
|
||||
FilterShape_,
|
||||
WarpShape_,
|
||||
cutlass::gemm::GemmShape<1, 1, 1>,
|
||||
ElementA_,
|
||||
layout::RowMajor,
|
||||
ElementB_,
|
||||
layout::ColumnMajor,
|
||||
ElementC_,
|
||||
LayoutC_,
|
||||
arch::OpClassSimt,
|
||||
kLaneAccessSizeA_,
|
||||
128,
|
||||
Stages_,
|
||||
Operator_,
|
||||
IteratorAlgorithm::kFixedStrideDilation,
|
||||
StrideShape_,
|
||||
DilationShape_,
|
||||
ActivationShape_> {
|
||||
using Shape = Shape_;
|
||||
using FilterShape = FilterShape_;
|
||||
using WarpShape = WarpShape_;
|
||||
using InstructionShape = cutlass::gemm::GemmShape<1, 1, 1>;
|
||||
using ElementA = ElementA_;
|
||||
using LayoutA = layout::RowMajor;
|
||||
using ElementB = ElementB_;
|
||||
using LayoutB = layout::ColumnMajor;
|
||||
using ElementC = ElementC_;
|
||||
using LayoutC = LayoutC_;
|
||||
using OperatorClass = arch::OpClassSimt;
|
||||
using StrideShape = StrideShape_;
|
||||
using DilationShape = DilationShape_;
|
||||
using ThreadBlockOutputShape = ThreadBlockOutputShape_;
|
||||
using ActivationShape = ActivationShape_;
|
||||
|
||||
static int const kLaneAccessSizeB = 128;
|
||||
|
||||
// Divisility requirements
|
||||
static_assert( kLaneAccessSizeB > 0,
|
||||
"Size of a warp-scoped per thread access should be larger then ZERO" );
|
||||
|
||||
/// Default Operator
|
||||
using Operator = Operator_;
|
||||
|
||||
/// Number of warps present
|
||||
using WarpCount = cutlass::gemm::GemmShape<
|
||||
Shape::kM / WarpShape::kM,
|
||||
Shape::kN / WarpShape::kN,
|
||||
1
|
||||
>;
|
||||
|
||||
// Divisility requirements
|
||||
static_assert(
|
||||
!(Shape::kM % WarpShape::kM) &&
|
||||
!(Shape::kN % WarpShape::kN),
|
||||
"Threadblock-scoped GEMM should be divisible by warp-scoped GEMM size."
|
||||
);
|
||||
|
||||
/// Number of threads per warp
|
||||
static int const kWarpSize = cutlass::gemm::warp::WarpSize<arch::OpClassSimt>::value;
|
||||
|
||||
/// Number of threads total
|
||||
static int const kThreads = WarpCount::kCount * kWarpSize;
|
||||
|
||||
// For Gmem load
|
||||
static int const kElementsPerAccessA = 128 / sizeof_bits<ElementA>::value;
|
||||
static int const kElementsPerAccessB = 128 / sizeof_bits<ElementB>::value;
|
||||
|
||||
//
|
||||
// Shared memory layouts
|
||||
//
|
||||
|
||||
using SmemLayoutA = layout::RowMajor;
|
||||
using SmemLayoutB = layout::RowMajor;
|
||||
|
||||
|
||||
//
|
||||
// Iterators to write to shared memory
|
||||
//
|
||||
|
||||
/// ThreadMap of iterator A
|
||||
using IteratorThreadMapA = transform::PitchLinearStripminedThreadMap<
|
||||
layout::PitchLinearShape<ActivationShape::kC, ActivationShape::kNHW>,
|
||||
kThreads,
|
||||
kElementsPerAccessA
|
||||
>;
|
||||
|
||||
/// ThreadMap of iterator A
|
||||
using SmemThreadMapA = IteratorThreadMapA;
|
||||
|
||||
/// Shared memory iterator to A operand
|
||||
using SmemIteratorA = transform::threadblock::RegularTileAccessIteratorDirectConv<
|
||||
MatrixShape<ActivationShape::kNHW, ActivationShape::kC>,
|
||||
ElementA,
|
||||
SmemLayoutA,
|
||||
0,
|
||||
SmemThreadMapA, // was IteratorThreadMapA
|
||||
false // static iterations.
|
||||
>;
|
||||
|
||||
/// ThreadMap of iterator B
|
||||
using IteratorThreadMapB = transform::PitchLinearStripminedThreadMap<
|
||||
layout::PitchLinearShape<Shape::kN, FilterShape::kCount>,
|
||||
kThreads,
|
||||
kElementsPerAccessB
|
||||
>;
|
||||
|
||||
/// Transpose the ThreadMap of iterator B
|
||||
using SmemThreadMapB = IteratorThreadMapB;
|
||||
|
||||
/// Shared memory iterator to B operand
|
||||
using SmemIteratorB = transform::threadblock::RegularTileAccessIteratorDirectConv<
|
||||
MatrixShape<FilterShape::kCount, Shape::kN>,
|
||||
ElementB,
|
||||
SmemLayoutB,
|
||||
0,
|
||||
SmemThreadMapB, // was IteratorThreadMapB
|
||||
false // static iterations.
|
||||
>;
|
||||
|
||||
//
|
||||
// Warp-level matrix multiply operator
|
||||
//
|
||||
// Groups per threads
|
||||
// Fp32: 2 groups
|
||||
// Fp16: 2 groups
|
||||
static const int GroupsPerThread = sizeof(ElementB) > 1 ? 2 : 4;
|
||||
// Define the warp-level op
|
||||
static const int WarpNumThreadsN = cutlass::const_min(WarpShape::kN / GroupsPerThread, kWarpSize);
|
||||
static const int WarpNumThreadsM = kWarpSize / WarpNumThreadsN;
|
||||
|
||||
static const int TileP = cutlass::conv::threadblock::detail::SimtWarpShape<WarpShape::kM, WarpNumThreadsM>::kP;
|
||||
static const int TileQ = cutlass::conv::threadblock::detail::SimtWarpShape<WarpShape::kM, WarpNumThreadsM>::kQ;
|
||||
|
||||
static_assert(!(WarpShape::kM % WarpNumThreadsM) && !(WarpShape::kN % WarpNumThreadsN),
|
||||
"WarpShape must be divisible by ThreadTile shape.");
|
||||
|
||||
static const int LaneLayout = 1;
|
||||
static const int numElementsB = kLaneAccessSizeB / sizeof_bits<ElementB>::value;
|
||||
static const int LaneN = cutlass::const_min(numElementsB, WarpShape::kN / WarpNumThreadsN);
|
||||
|
||||
// Define the output tile computed by each thread
|
||||
using ThreadOutputShape = cutlass::conv::TensorNHWCShape<1, TileP, TileQ, LaneN>;
|
||||
|
||||
// Fetch the channel with same access size
|
||||
static const int LaneM = LaneN;
|
||||
|
||||
// No paddings
|
||||
static int const kPaddingM = 0;
|
||||
static int const kPaddingN = 0;
|
||||
|
||||
static_assert(!(kPaddingM % LaneM) && !(kPaddingN % LaneN),
|
||||
"Padding must be divisible by Lane");
|
||||
|
||||
// these should have max of thread tile also
|
||||
using LaneMmaShape = cutlass::gemm::GemmShape<
|
||||
LaneM,
|
||||
LaneN,
|
||||
1>;
|
||||
|
||||
using Policy = cutlass::gemm::warp::MmaSimtPolicy<
|
||||
cutlass::MatrixShape<WarpNumThreadsM, WarpNumThreadsN>, // WarpShape
|
||||
cutlass::layout::RowMajorInterleaved<LaneLayout>, // LaneLayout
|
||||
LaneMmaShape
|
||||
>;
|
||||
|
||||
using MmaWarpSimt = cutlass::conv::warp::MmaDepthwiseDirectConvSimt<
|
||||
WarpShape, /// Size of the Gemm problem - concept: gemm::GemmShape<>
|
||||
FilterShape, /// Shape of filter shape per threadblock - concept: gemm::GemmShape<Depth, Height, Width>
|
||||
ThreadOutputShape, /// Size of the output tile computed by thread - concept: conv::TensorNHWCShape<>
|
||||
ThreadBlockOutputShape, /// Size of the output tile computed by threadblock - concept: conv::TensorNHWCShape<>
|
||||
ElementA, /// Data type of A elements
|
||||
SmemLayoutA, /// Layout of A matrix (concept: MatrixLayout)
|
||||
ElementB, /// Data type of B elements
|
||||
SmemLayoutB, /// Layout of B matrix (concept: MatrixLayout)
|
||||
ElementC, /// Element type of C matrix
|
||||
LayoutC, /// Layout of C matrix (concept: MatrixLayout)
|
||||
Policy, /// Policy describing warp-level MmaSimtOp (concept: MmaSimtOp policy)
|
||||
IteratorAlgorithm::kFixedStrideDilation, /// Iterator algo type
|
||||
StrideShape, /// Stride ( MatrixShape<Height, Width> )
|
||||
DilationShape, /// Dilation ( MatrixShape<Height, Width> )
|
||||
ActivationShape /// Activation Shape loaded by threadblock
|
||||
>;
|
||||
|
||||
/// Policy used to define MmaPipelined
|
||||
using MmaPolicy = cutlass::conv::threadblock::DepthwiseDirectConvMmaPolicy<
|
||||
MmaWarpSimt,
|
||||
MatrixShape<kPaddingM, 0>, // skew for A matrix to avoid SMEM bank conflicts
|
||||
MatrixShape<0, kPaddingN>, // skew for B matrix to avoid SMEM bank conflicts
|
||||
IteratorThreadMapA,
|
||||
IteratorThreadMapB,
|
||||
WarpCount::kK
|
||||
>;
|
||||
};
|
||||
} // namespace threadblock
|
||||
} // namespace conv
|
||||
} // namespace cutlass
|
||||
|
||||
@@ -165,7 +165,29 @@ struct StridedDgradIdentityThreadblockSwizzle :
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Threadblock swizzling function for GEMMs
|
||||
template <int N = 1, int Output_N = 1, int Output_P = 1, int Output_Q = 1>
|
||||
struct DepthwiseDirect2dConvIdentityThreadblockSwizzle
|
||||
: public gemm::threadblock::GemmIdentityThreadblockSwizzle<N> {
|
||||
CUTLASS_HOST_DEVICE
|
||||
DepthwiseDirect2dConvIdentityThreadblockSwizzle() {}
|
||||
|
||||
/// Returns the shape of the problem in units of logical tiles
|
||||
CUTLASS_HOST_DEVICE
|
||||
gemm::GemmCoord get_tiled_shape(cutlass::conv::Operator conv_operator,
|
||||
cutlass::conv::Conv2dProblemSize const &problem_size,
|
||||
gemm::GemmCoord tile_size,
|
||||
int split_k_slices) const {
|
||||
|
||||
gemm::GemmCoord implicit_gemm_problem_size =
|
||||
cutlass::conv::implicit_gemm_problem_size(conv_operator, problem_size);
|
||||
|
||||
return gemm::GemmCoord(1,
|
||||
(implicit_gemm_problem_size.n() + tile_size.n() - 1) / tile_size.n(),
|
||||
split_k_slices);
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace gemm
|
||||
} // namespace conv
|
||||
} // namespace cutlass
|
||||
|
||||
Reference in New Issue
Block a user