releaase 2.11 (#703)
This commit is contained in:
@@ -247,7 +247,7 @@ public:
|
||||
CUTLASS_HOST_DEVICE
|
||||
cutlass::Tensor4DCoord filter_extent() const {
|
||||
|
||||
return cutlass::Tensor4DCoord ({K, R, S, C});
|
||||
return cutlass::Tensor4DCoord ({K, R, S, C / groups});
|
||||
}
|
||||
|
||||
/// Returns output extent as Tensor4DCoord
|
||||
@@ -336,7 +336,7 @@ cutlass::gemm::GemmCoord implicit_gemm_problem_size(
|
||||
return gemm::GemmCoord(
|
||||
problem_size.N * problem_size.P * problem_size.Q,
|
||||
problem_size.K,
|
||||
problem_size.R * problem_size.S * problem_size.C
|
||||
problem_size.R * problem_size.S * problem_size.C / problem_size.groups
|
||||
);
|
||||
case Operator::kDgrad:
|
||||
return gemm::GemmCoord(
|
||||
@@ -451,6 +451,18 @@ int implicit_gemm_k_iterations(
|
||||
default:
|
||||
break;
|
||||
}
|
||||
} else if (algorithm == IteratorAlgorithm::kOptimized) {
|
||||
// Current optimized iterator only support GroupMode::kSingleGroup
|
||||
if (group_mode == GroupMode::kSingleGroup) {
|
||||
switch (conv_operator) {
|
||||
case Operator::kFprop:
|
||||
iterations = problem_size.R * problem_size.S * ((channels_per_group + threadblock_K - 1) / threadblock_K);
|
||||
break;
|
||||
|
||||
default:
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
@@ -459,6 +471,25 @@ int implicit_gemm_k_iterations(
|
||||
}
|
||||
|
||||
|
||||
template <int N = 1, int Output_P = 1, int Output_Q = 1>
|
||||
CUTLASS_HOST_DEVICE
|
||||
int depthwise_gemm_k_iterations(
|
||||
Operator conv_operator,
|
||||
int threadblock_K,
|
||||
Conv2dProblemSize const &problem_size,
|
||||
IteratorAlgorithm algorithm = IteratorAlgorithm::kAnalytic,
|
||||
GroupMode group_mode = GroupMode::kNone,
|
||||
int threadblock_N = 0) {
|
||||
|
||||
int n = problem_size.N;
|
||||
int p = (problem_size.P + Output_P - 1) / Output_P;
|
||||
int q = (problem_size.Q + Output_Q - 1) / Output_Q;
|
||||
|
||||
int iterations = (n * p * q + problem_size.split_k_slices - 1) / problem_size.split_k_slices;
|
||||
return iterations;
|
||||
}
|
||||
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
int implicit_gemm_k_iterations_per_channel(
|
||||
Operator conv_operator,
|
||||
|
||||
@@ -100,14 +100,16 @@ enum class IteratorAlgorithm {
|
||||
kAnalytic, ///< functionally correct in all cases but lower performance
|
||||
kOptimized, ///< optimized for R <= 32, S <= 32 and unity-stride dgrad
|
||||
kFixedChannels, ///< Analytic algorithm optimized for fixed channel count (C == AccessSize)
|
||||
kFewChannels ///< Analytic algorithm optimized for few channels (C divisible by AccessSize)
|
||||
kFewChannels, ///< Analytic algorithm optimized for few channels (C divisible by AccessSize)
|
||||
kFixedStrideDilation ///< Optimized for fixed stride and dilation
|
||||
};
|
||||
|
||||
/// Distinguishes among partial specializations that accelerate certain problems where convolution
|
||||
/// stride is unit.
|
||||
enum class StrideSupport {
|
||||
kStrided, ///< arbitrary convolution stride
|
||||
kUnity ///< unit convolution stride
|
||||
kUnity, ///< unit convolution stride
|
||||
kFixed ///< fixed convolution stride
|
||||
};
|
||||
|
||||
/// Identifies split-K mode
|
||||
@@ -125,6 +127,38 @@ enum class GroupMode {
|
||||
kDepthwise ///< One CTA calculates cta_n groups (problem_size.C == problem_size.K == problem_size.groups)
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Shape of a tensor
|
||||
template <
|
||||
int N = 1,
|
||||
int H = 1,
|
||||
int W = 1,
|
||||
int C = 1
|
||||
>
|
||||
struct TensorNHWCShape {
|
||||
static int const kN = N;
|
||||
static int const kH = H;
|
||||
static int const kW = W;
|
||||
static int const kC = C;
|
||||
|
||||
static int const kHW = H * W;
|
||||
static int const kNHW = N * kHW;
|
||||
static int const kNHWC = N * H * W * C;
|
||||
|
||||
static int const kCount = kNHWC;
|
||||
|
||||
//
|
||||
// Static member functions
|
||||
//
|
||||
|
||||
/// Returns a Coord object
|
||||
CUTLASS_HOST_DEVICE
|
||||
static Coord<4> toCoord() {
|
||||
return make_Coord(kN, kH, kW, kC);
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace conv
|
||||
|
||||
@@ -0,0 +1,269 @@
|
||||
/***************************************************************************************************
|
||||
* 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 device-level Depthwise Convolution
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <limits>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/device_kernel.h"
|
||||
#include "cutlass/conv/convolution.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace conv {
|
||||
namespace device {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template<typename DirectConvolutionKernel_>
|
||||
class DirectConvolution {
|
||||
public:
|
||||
|
||||
using UnderlyingKernel = DirectConvolutionKernel_;
|
||||
|
||||
using ElementA = typename UnderlyingKernel::ElementA;
|
||||
using LayoutA = typename UnderlyingKernel::LayoutA;
|
||||
using ElementB = typename UnderlyingKernel::ElementB;
|
||||
using LayoutB = typename UnderlyingKernel::LayoutB;
|
||||
using ElementC = typename UnderlyingKernel::ElementC;
|
||||
using LayoutC = typename UnderlyingKernel::LayoutC;
|
||||
using ElementAccumulator = typename UnderlyingKernel::ElementAccumulator;
|
||||
using ElementCompute = typename UnderlyingKernel::ElementCompute;
|
||||
using OperatorClass = typename UnderlyingKernel::OperatorClass;
|
||||
using ArchTag = typename UnderlyingKernel::ArchTag;
|
||||
using ThreadblockShape = typename UnderlyingKernel::ThreadblockShape;
|
||||
using WarpShape = typename UnderlyingKernel::WarpShape;
|
||||
using InstructionShape = typename UnderlyingKernel::InstructionShape;
|
||||
using ThreadblockSwizzle = typename UnderlyingKernel::ThreadblockSwizzle;
|
||||
using EpilogueOutputOp = typename UnderlyingKernel::EpilogueOutputOp;
|
||||
static int const kStages = UnderlyingKernel::kStages;
|
||||
static int const kConvDim = UnderlyingKernel::kConvDim;
|
||||
using WarpMmaOperator = typename UnderlyingKernel::WarpMmaOperator;
|
||||
using ArchMmaOperator = typename UnderlyingKernel::ArchMmaOperator;
|
||||
using MathOperator = typename UnderlyingKernel::MathOperator;
|
||||
|
||||
static cutlass::conv::Operator const kConvolutionalOperator = UnderlyingKernel::kConvolutionalOperator;
|
||||
static cutlass::conv::IteratorAlgorithm const kIteratorAlgorithm = UnderlyingKernel::kIteratorAlgorithm;
|
||||
static cutlass::conv::StrideSupport const kStrideSupport = UnderlyingKernel::kStrideSupport;
|
||||
static cutlass::conv::GroupMode const kGroupMode = UnderlyingKernel::kGroupMode;
|
||||
|
||||
static int const kWarpCount =
|
||||
(ThreadblockShape::kM / WarpShape::kM) *
|
||||
(ThreadblockShape::kN / WarpShape::kN) *
|
||||
(ThreadblockShape::kK / WarpShape::kK);
|
||||
|
||||
/// Argument structure
|
||||
using Arguments = typename UnderlyingKernel::Arguments;
|
||||
|
||||
using ReorderKernel = typename UnderlyingKernel::ReorderKernel;
|
||||
|
||||
private:
|
||||
|
||||
/// Kernel parameters object
|
||||
typename UnderlyingKernel::Params params_;
|
||||
|
||||
public:
|
||||
|
||||
/// Constructs Implicit GEMM
|
||||
DirectConvolution() { }
|
||||
|
||||
/// Determines whether the Implicit GEMM can execute the given problem.
|
||||
static Status can_implement(Arguments const &args) {
|
||||
|
||||
// dispatch to iterators
|
||||
Status status = UnderlyingKernel::Mma::IteratorA::can_implement(args.problem_size);
|
||||
if (Status::kSuccess != status) {
|
||||
return status;
|
||||
}
|
||||
|
||||
status = UnderlyingKernel::Mma::IteratorB::can_implement(args.problem_size);
|
||||
if (Status::kSuccess != status) {
|
||||
return status;
|
||||
}
|
||||
|
||||
if (kGroupMode != conv::GroupMode::kDepthwise) {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
// C and K should be multiple of groups
|
||||
if (args.problem_size.K != args.problem_size.groups &&
|
||||
args.problem_size.C != args.problem_size.groups) {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
|
||||
static int const kAlignmentC = UnderlyingKernel::Epilogue::OutputTileIterator::kElementsPerAccess;
|
||||
if (kConvolutionalOperator == conv::Operator::kFprop) {
|
||||
if (args.problem_size.K % kAlignmentC)
|
||||
return Status::kErrorMisalignedOperand;
|
||||
} else if (kConvolutionalOperator == conv::Operator::kDgrad) {
|
||||
if (args.problem_size.C % kAlignmentC)
|
||||
return Status::kErrorMisalignedOperand;
|
||||
} else if (kConvolutionalOperator == conv::Operator::kWgrad) {
|
||||
if (args.problem_size.C % kAlignmentC)
|
||||
return Status::kErrorMisalignedOperand;
|
||||
}
|
||||
|
||||
// Determine grid shape
|
||||
ThreadblockSwizzle threadblock_swizzle;
|
||||
|
||||
dim3 grid = threadblock_swizzle.get_grid_shape(
|
||||
threadblock_swizzle.get_tiled_shape(
|
||||
kConvolutionalOperator,
|
||||
args.problem_size,
|
||||
{ThreadblockShape::kM, ThreadblockShape::kN, ThreadblockShape::kK},
|
||||
args.problem_size.split_k_slices));
|
||||
|
||||
if (!(grid.y <= std::numeric_limits<uint16_t>::max() &&
|
||||
grid.z <= std::numeric_limits<uint16_t>::max())) {
|
||||
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
/// Gets the workspace size
|
||||
static size_t get_workspace_size(Arguments const &args) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
/// Initializes GEMM state from arguments.
|
||||
Status initialize(
|
||||
Arguments const &args,
|
||||
void *workspace = nullptr,
|
||||
cudaStream_t stream = nullptr) {
|
||||
|
||||
// initialize the params structure from the arguments
|
||||
params_ = typename UnderlyingKernel::Params(
|
||||
args,
|
||||
static_cast<int *>(workspace)
|
||||
);
|
||||
|
||||
int smem_size = int(sizeof(typename UnderlyingKernel::SharedStorage));
|
||||
|
||||
if (smem_size >= (48 << 10)) {
|
||||
cudaError_t result = cudaFuncSetAttribute(cutlass::Kernel<UnderlyingKernel>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
smem_size);
|
||||
|
||||
if (result != cudaSuccess) {
|
||||
return Status::kErrorInternal;
|
||||
}
|
||||
}
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
/// Initializes GEMM state from arguments.
|
||||
Status update(Arguments const &args, void *workspace = nullptr) {
|
||||
|
||||
// update the params structure from the arguments
|
||||
params_.ptr_A = args.ref_A.data();
|
||||
params_.ptr_B = args.ref_B.data();
|
||||
params_.ptr_C = args.ref_C.data();
|
||||
params_.ptr_D = args.ref_D.data();
|
||||
params_.output_op = args.output_op;
|
||||
params_.ptr_reordered_B = args.ref_reordered_B.data();;
|
||||
params_.semaphore = static_cast<int *>(workspace);
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
/// Runs the kernel using initialized state.
|
||||
Status run(cudaStream_t stream = nullptr) {
|
||||
|
||||
// Launch reorder kernel
|
||||
if (params_.ptr_reordered_B != nullptr) {
|
||||
dim3 grid = ReorderKernel::get_grid_shape(params_);
|
||||
dim3 block = ReorderKernel::get_block_shape();
|
||||
|
||||
cutlass::Kernel<ReorderKernel><<<grid, block, 0, stream>>>(params_);
|
||||
}
|
||||
|
||||
// Launch main kernel
|
||||
ThreadblockSwizzle threadblock_swizzle;
|
||||
|
||||
dim3 grid = threadblock_swizzle.get_grid_shape(params_.grid_tiled_shape);
|
||||
dim3 block(32 * kWarpCount, 1, 1);
|
||||
|
||||
// Dynamic SMEM size based on input params.
|
||||
int smem_size = int(params_.get_smem_size());
|
||||
|
||||
// Make sure we can use that much shared memory.
|
||||
cudaError_t status =
|
||||
cudaFuncSetAttribute(cutlass::Kernel<UnderlyingKernel>, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
|
||||
if (status != cudaSuccess)
|
||||
return Status::kErrorInternal;
|
||||
|
||||
|
||||
cutlass::Kernel<UnderlyingKernel><<<grid, block, smem_size, stream>>>(params_);
|
||||
|
||||
cudaError_t result = cudaGetLastError();
|
||||
|
||||
return result == cudaSuccess ? Status::kSuccess : Status::kErrorInternal;
|
||||
}
|
||||
|
||||
/// Runs the kernel using initialized state.
|
||||
Status operator()(cudaStream_t stream = nullptr) {
|
||||
return run(stream);
|
||||
}
|
||||
|
||||
/// Runs the kernel using initialized state.
|
||||
Status operator()(
|
||||
Arguments const &args,
|
||||
void *workspace = nullptr,
|
||||
cudaStream_t stream = nullptr) {
|
||||
|
||||
Status status = initialize(args, workspace, stream);
|
||||
|
||||
if (status == Status::kSuccess) {
|
||||
status = run(stream);
|
||||
}
|
||||
|
||||
return status;
|
||||
}
|
||||
|
||||
int get_smem_size() { return int(params_.get_smem_size()); }
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -52,33 +52,33 @@ template<typename ImplicitGemmKernel_>
|
||||
class ImplicitGemmConvolution {
|
||||
public:
|
||||
|
||||
using ImplicitGemmKernel = ImplicitGemmKernel_;
|
||||
using UnderlyingKernel = ImplicitGemmKernel_;
|
||||
|
||||
using ElementA = typename ImplicitGemmKernel::ElementA;
|
||||
using LayoutA = typename ImplicitGemmKernel::LayoutA;
|
||||
using ElementB = typename ImplicitGemmKernel::ElementB;
|
||||
using LayoutB = typename ImplicitGemmKernel::LayoutB;
|
||||
using ElementC = typename ImplicitGemmKernel::ElementC;
|
||||
using LayoutC = typename ImplicitGemmKernel::LayoutC;
|
||||
using ElementAccumulator = typename ImplicitGemmKernel::ElementAccumulator;
|
||||
using ElementCompute = typename ImplicitGemmKernel::ElementCompute;
|
||||
using OperatorClass = typename ImplicitGemmKernel::OperatorClass;
|
||||
using ArchTag = typename ImplicitGemmKernel::ArchTag;
|
||||
using ThreadblockShape = typename ImplicitGemmKernel::ThreadblockShape;
|
||||
using WarpShape = typename ImplicitGemmKernel::WarpShape;
|
||||
using InstructionShape = typename ImplicitGemmKernel::InstructionShape;
|
||||
using ThreadblockSwizzle = typename ImplicitGemmKernel::ThreadblockSwizzle;
|
||||
using EpilogueOutputOp = typename ImplicitGemmKernel::EpilogueOutputOp;
|
||||
static int const kStages = ImplicitGemmKernel::kStages;
|
||||
static int const kConvDim = ImplicitGemmKernel::kConvDim;
|
||||
using WarpMmaOperator = typename ImplicitGemmKernel::WarpMmaOperator;
|
||||
using ArchMmaOperator = typename ImplicitGemmKernel::ArchMmaOperator;
|
||||
using MathOperator = typename ImplicitGemmKernel::MathOperator;
|
||||
using ElementA = typename UnderlyingKernel::ElementA;
|
||||
using LayoutA = typename UnderlyingKernel::LayoutA;
|
||||
using ElementB = typename UnderlyingKernel::ElementB;
|
||||
using LayoutB = typename UnderlyingKernel::LayoutB;
|
||||
using ElementC = typename UnderlyingKernel::ElementC;
|
||||
using LayoutC = typename UnderlyingKernel::LayoutC;
|
||||
using ElementAccumulator = typename UnderlyingKernel::ElementAccumulator;
|
||||
using ElementCompute = typename UnderlyingKernel::ElementCompute;
|
||||
using OperatorClass = typename UnderlyingKernel::OperatorClass;
|
||||
using ArchTag = typename UnderlyingKernel::ArchTag;
|
||||
using ThreadblockShape = typename UnderlyingKernel::ThreadblockShape;
|
||||
using WarpShape = typename UnderlyingKernel::WarpShape;
|
||||
using InstructionShape = typename UnderlyingKernel::InstructionShape;
|
||||
using ThreadblockSwizzle = typename UnderlyingKernel::ThreadblockSwizzle;
|
||||
using EpilogueOutputOp = typename UnderlyingKernel::EpilogueOutputOp;
|
||||
static int const kStages = UnderlyingKernel::kStages;
|
||||
static int const kConvDim = UnderlyingKernel::kConvDim;
|
||||
using WarpMmaOperator = typename UnderlyingKernel::WarpMmaOperator;
|
||||
using ArchMmaOperator = typename UnderlyingKernel::ArchMmaOperator;
|
||||
using MathOperator = typename UnderlyingKernel::MathOperator;
|
||||
|
||||
static cutlass::conv::Operator const kConvolutionalOperator = ImplicitGemmKernel::kConvolutionalOperator;
|
||||
static cutlass::conv::IteratorAlgorithm const kIteratorAlgorithm = ImplicitGemmKernel::kIteratorAlgorithm;
|
||||
static cutlass::conv::StrideSupport const kStrideSupport = ImplicitGemmKernel::kStrideSupport;
|
||||
static cutlass::conv::GroupMode const kGroupMode = ImplicitGemmKernel::kGroupMode;
|
||||
static cutlass::conv::Operator const kConvolutionalOperator = UnderlyingKernel::kConvolutionalOperator;
|
||||
static cutlass::conv::IteratorAlgorithm const kIteratorAlgorithm = UnderlyingKernel::kIteratorAlgorithm;
|
||||
static cutlass::conv::StrideSupport const kStrideSupport = UnderlyingKernel::kStrideSupport;
|
||||
static cutlass::conv::GroupMode const kGroupMode = UnderlyingKernel::kGroupMode;
|
||||
|
||||
static int const kWarpCount =
|
||||
(ThreadblockShape::kM / WarpShape::kM) *
|
||||
@@ -86,12 +86,12 @@ public:
|
||||
(ThreadblockShape::kK / WarpShape::kK);
|
||||
|
||||
/// Argument structure
|
||||
using Arguments = typename ImplicitGemmKernel::Arguments;
|
||||
using Arguments = typename UnderlyingKernel::Arguments;
|
||||
|
||||
private:
|
||||
|
||||
/// Kernel parameters object
|
||||
typename ImplicitGemmKernel::Params params_;
|
||||
typename UnderlyingKernel::Params params_;
|
||||
|
||||
public:
|
||||
|
||||
@@ -102,12 +102,12 @@ public:
|
||||
static Status can_implement(Arguments const &args) {
|
||||
|
||||
// dispatch to iterators
|
||||
Status status = ImplicitGemmKernel::Mma::IteratorA::can_implement(args.problem_size);
|
||||
Status status = UnderlyingKernel::Mma::IteratorA::can_implement(args.problem_size);
|
||||
if (Status::kSuccess != status) {
|
||||
return status;
|
||||
}
|
||||
|
||||
status = ImplicitGemmKernel::Mma::IteratorB::can_implement(args.problem_size);
|
||||
status = UnderlyingKernel::Mma::IteratorB::can_implement(args.problem_size);
|
||||
if (Status::kSuccess != status) {
|
||||
return status;
|
||||
}
|
||||
@@ -138,9 +138,15 @@ public:
|
||||
if (kGroupMode == conv::GroupMode::kMultipleGroup && ThreadblockShape::kN % k_per_group) {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
|
||||
// current optimized iterator algo only supports SingleGroup mode
|
||||
if (kIteratorAlgorithm == IteratorAlgorithm::kOptimized &&
|
||||
kGroupMode != conv::GroupMode::kSingleGroup) {
|
||||
return Status::kErrorInvalidProblem;
|
||||
}
|
||||
}
|
||||
|
||||
static int const kAlignmentC = ImplicitGemmKernel::Epilogue::OutputTileIterator::kElementsPerAccess;
|
||||
static int const kAlignmentC = UnderlyingKernel::Epilogue::OutputTileIterator::kElementsPerAccess;
|
||||
if (kConvolutionalOperator == conv::Operator::kFprop) {
|
||||
if (args.problem_size.K % kAlignmentC)
|
||||
return Status::kErrorMisalignedOperand;
|
||||
@@ -249,15 +255,15 @@ public:
|
||||
}
|
||||
|
||||
// initialize the params structure from the arguments
|
||||
params_ = typename ImplicitGemmKernel::Params(
|
||||
params_ = typename UnderlyingKernel::Params(
|
||||
args,
|
||||
static_cast<int *>(workspace)
|
||||
);
|
||||
|
||||
int smem_size = int(sizeof(typename ImplicitGemmKernel::SharedStorage));
|
||||
int smem_size = int(sizeof(typename UnderlyingKernel::SharedStorage));
|
||||
|
||||
if (smem_size >= (48 << 10)) {
|
||||
cudaError_t result = cudaFuncSetAttribute(cutlass::Kernel<ImplicitGemmKernel>,
|
||||
cudaError_t result = cudaFuncSetAttribute(cutlass::Kernel<UnderlyingKernel>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
smem_size);
|
||||
|
||||
@@ -292,9 +298,9 @@ public:
|
||||
dim3 grid = threadblock_swizzle.get_grid_shape(params_.grid_tiled_shape);
|
||||
dim3 block(32 * kWarpCount, 1, 1);
|
||||
|
||||
int smem_size = int(sizeof(typename ImplicitGemmKernel::SharedStorage));
|
||||
int smem_size = int(sizeof(typename UnderlyingKernel::SharedStorage));
|
||||
|
||||
cutlass::Kernel<ImplicitGemmKernel><<<grid, block, smem_size, stream>>>(params_);
|
||||
cutlass::Kernel<UnderlyingKernel><<<grid, block, smem_size, stream>>>(params_);
|
||||
|
||||
cudaError_t result = cudaGetLastError();
|
||||
|
||||
|
||||
@@ -89,7 +89,7 @@ template <
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Defines a kernel for Conv2dGroupFprop specialization for Analytic IteratorAlgorithm and multistage
|
||||
/// pipeline.
|
||||
/// pipeline that supports all GroupMode.
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
@@ -135,6 +135,13 @@ struct DefaultConv2dGroupFprop <
|
||||
AlignmentB
|
||||
> {
|
||||
|
||||
static_assert(std::is_same<LayoutA, cutlass::layout::TensorNHWC>::value,
|
||||
"Current group conv only support NHWC layout");
|
||||
static_assert(std::is_same<LayoutB, cutlass::layout::TensorNHWC>::value,
|
||||
"Current group conv only support NHWC layout");
|
||||
static_assert(std::is_same<LayoutC, cutlass::layout::TensorNHWC>::value,
|
||||
"Current group conv only support NHWC layout");
|
||||
|
||||
// Define the core components from GEMM
|
||||
using MmaCore = typename cutlass::gemm::threadblock::DefaultMmaCore<
|
||||
ThreadblockShape, WarpShape, InstructionShape, ElementA, layout::RowMajor,
|
||||
@@ -215,6 +222,267 @@ struct DefaultConv2dGroupFprop <
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Defines a kernel for Conv2dGroupFprop specialization for Optimized IteratorAlgorithm and multistage
|
||||
/// pipeline that supports GroupMode::kSingleGroup.
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
typename ElementB,
|
||||
typename LayoutB,
|
||||
typename ElementC,
|
||||
typename LayoutC,
|
||||
typename ElementAccumulator,
|
||||
typename ArchTag,
|
||||
typename ThreadblockShape,
|
||||
typename WarpShape,
|
||||
typename InstructionShape,
|
||||
typename EpilogueOutputOp,
|
||||
typename ThreadblockSwizzle,
|
||||
int Stages,
|
||||
typename MathOperatorTag,
|
||||
conv::StrideSupport StrideSupport,
|
||||
int AlignmentA,
|
||||
int AlignmentB
|
||||
>
|
||||
struct DefaultConv2dGroupFprop <
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
ElementAccumulator,
|
||||
arch::OpClassTensorOp,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
Stages,
|
||||
MathOperatorTag,
|
||||
GroupMode::kSingleGroup,
|
||||
IteratorAlgorithm::kOptimized,
|
||||
StrideSupport,
|
||||
AlignmentA,
|
||||
AlignmentB
|
||||
> {
|
||||
|
||||
static_assert(std::is_same<LayoutA, cutlass::layout::TensorNHWC>::value,
|
||||
"Current group conv only support NHWC layout");
|
||||
static_assert(std::is_same<LayoutB, cutlass::layout::TensorNHWC>::value,
|
||||
"Current group conv only support NHWC layout");
|
||||
static_assert(std::is_same<LayoutC, cutlass::layout::TensorNHWC>::value,
|
||||
"Current group conv only support NHWC layout");
|
||||
|
||||
// Define the core components from GEMM
|
||||
using MmaCore = typename cutlass::gemm::threadblock::DefaultMmaCore<
|
||||
ThreadblockShape, WarpShape, InstructionShape, ElementA, layout::RowMajor,
|
||||
ElementB, layout::ColumnMajor, ElementAccumulator, layout::RowMajor, arch::OpClassTensorOp,
|
||||
Stages, MathOperatorTag>;
|
||||
|
||||
// Define iterators over tiles from the A operand
|
||||
using ThreadMapA = typename MmaCore::IteratorThreadMapA;
|
||||
using AccessTypeA = cutlass::AlignedArray<ElementA, AlignmentA>;
|
||||
using IteratorA =
|
||||
cutlass::conv::threadblock::Conv2dFpropActivationTileAccessIteratorOptimized<
|
||||
cutlass::MatrixShape<ThreadblockShape::kM, ThreadblockShape::kK>,
|
||||
ElementA, LayoutA,
|
||||
ThreadMapA,
|
||||
AccessTypeA
|
||||
>;
|
||||
|
||||
using SmemIteratorA = typename MmaCore::SmemIteratorA;
|
||||
|
||||
// Define iterators over tiles from the B operand
|
||||
using ThreadMapB = typename MmaCore::IteratorThreadMapB;
|
||||
using AccessTypeB = cutlass::AlignedArray<ElementB, AlignmentB>;
|
||||
using IteratorB =
|
||||
cutlass::conv::threadblock::Conv2dFpropFilterTileAccessIteratorOptimized<
|
||||
cutlass::MatrixShape<ThreadblockShape::kK, ThreadblockShape::kN>,
|
||||
ElementB, LayoutB,
|
||||
ThreadMapB,
|
||||
AccessTypeB
|
||||
>;
|
||||
|
||||
using SmemIteratorB = typename MmaCore::SmemIteratorB;
|
||||
|
||||
// Warp-level GEMM components
|
||||
using WarpMmaTensorOp = typename MmaCore::MmaTensorOp;
|
||||
using MmaPolicy = typename MmaCore::MmaPolicy;
|
||||
|
||||
static cutlass::arch::CacheOperation::Kind const CacheOpB =
|
||||
((sizeof_bits<ElementB>::value * AlignmentB) == 128)
|
||||
? cutlass::arch::CacheOperation::Global
|
||||
: cutlass::arch::CacheOperation::Always;
|
||||
|
||||
// Define the Mma
|
||||
using Mma = threadblock::ImplicitGemmMultistage<
|
||||
ThreadblockShape,
|
||||
IteratorA,
|
||||
SmemIteratorA,
|
||||
arch::CacheOperation::Always,
|
||||
IteratorB,
|
||||
SmemIteratorB,
|
||||
CacheOpB,
|
||||
MmaPolicy,
|
||||
Stages
|
||||
>;
|
||||
|
||||
static const int kPartitionsK = ThreadblockShape::kK / WarpShape::kK;
|
||||
|
||||
// Define the epilogue
|
||||
using Epilogue = typename epilogue::threadblock::DefaultEpilogueTensorOp<
|
||||
ThreadblockShape,
|
||||
WarpMmaTensorOp,
|
||||
kPartitionsK,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
using Kernel = cutlass::conv::kernel::ImplicitGemmConvolution<
|
||||
Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
conv::Operator::kFprop,
|
||||
Conv2dProblemSize,
|
||||
GroupMode::kSingleGroup
|
||||
>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Defines a kernel for Conv2dGroupFprop specialization for Optimized IteratorAlgorithm and
|
||||
/// 2 stage pipeline that supports GroupMode::kSingleGroup.
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
typename ElementB,
|
||||
typename LayoutB,
|
||||
typename ElementC,
|
||||
typename LayoutC,
|
||||
typename ElementAccumulator,
|
||||
typename ArchTag,
|
||||
typename ThreadblockShape,
|
||||
typename WarpShape,
|
||||
typename InstructionShape,
|
||||
typename EpilogueOutputOp,
|
||||
typename ThreadblockSwizzle,
|
||||
typename MathOperatorTag,
|
||||
conv::StrideSupport StrideSupport,
|
||||
int AlignmentA,
|
||||
int AlignmentB
|
||||
>
|
||||
struct DefaultConv2dGroupFprop <
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
ElementAccumulator,
|
||||
arch::OpClassTensorOp,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
2,
|
||||
MathOperatorTag,
|
||||
GroupMode::kSingleGroup,
|
||||
IteratorAlgorithm::kOptimized,
|
||||
StrideSupport,
|
||||
AlignmentA,
|
||||
AlignmentB
|
||||
> {
|
||||
|
||||
static_assert(std::is_same<LayoutA, cutlass::layout::TensorNHWC>::value,
|
||||
"Current group conv only support NHWC layout");
|
||||
static_assert(std::is_same<LayoutB, cutlass::layout::TensorNHWC>::value,
|
||||
"Current group conv only support NHWC layout");
|
||||
static_assert(std::is_same<LayoutC, cutlass::layout::TensorNHWC>::value,
|
||||
"Current group conv only support NHWC layout");
|
||||
|
||||
// Define the core components from GEMM
|
||||
using MmaCore = typename cutlass::gemm::threadblock::DefaultMmaCore<
|
||||
ThreadblockShape, WarpShape, InstructionShape, ElementA, layout::RowMajor,
|
||||
ElementB, layout::ColumnMajor, ElementAccumulator, layout::RowMajor, arch::OpClassTensorOp,
|
||||
2, MathOperatorTag>;
|
||||
|
||||
// Define iterators over tiles from the A operand
|
||||
using ThreadMapA = typename MmaCore::IteratorThreadMapA;
|
||||
using AccessTypeA = cutlass::AlignedArray<ElementA, AlignmentA>;
|
||||
using IteratorA =
|
||||
cutlass::conv::threadblock::TileIterator<
|
||||
cutlass::conv::threadblock::Conv2dFpropActivationTileAccessIteratorOptimized<
|
||||
cutlass::MatrixShape<ThreadblockShape::kM, ThreadblockShape::kK>,
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ThreadMapA,
|
||||
AccessTypeA
|
||||
>
|
||||
>;
|
||||
|
||||
using SmemIteratorA = typename MmaCore::SmemIteratorA;
|
||||
|
||||
// Define iterators over tiles from the B operand
|
||||
using ThreadMapB = typename MmaCore::IteratorThreadMapB;
|
||||
using AccessTypeB = cutlass::AlignedArray<ElementB, AlignmentB>;
|
||||
using IteratorB =
|
||||
cutlass::conv::threadblock::TileIterator<
|
||||
cutlass::conv::threadblock::Conv2dFpropFilterTileAccessIteratorOptimized<
|
||||
cutlass::MatrixShape<ThreadblockShape::kK, ThreadblockShape::kN>,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ThreadMapB,
|
||||
AccessTypeB
|
||||
>
|
||||
>;
|
||||
|
||||
using SmemIteratorB = typename MmaCore::SmemIteratorB;
|
||||
|
||||
// Warp-level GEMM components
|
||||
using WarpMmaTensorOp = typename MmaCore::MmaTensorOp;
|
||||
using MmaPolicy = typename MmaCore::MmaPolicy;
|
||||
|
||||
// Define the Mma
|
||||
using Mma = threadblock::ImplicitGemmPipelined<
|
||||
ThreadblockShape,
|
||||
IteratorA,
|
||||
SmemIteratorA,
|
||||
IteratorB,
|
||||
SmemIteratorB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
MmaPolicy
|
||||
>;
|
||||
|
||||
static const int kPartitionsK = ThreadblockShape::kK / WarpShape::kK;
|
||||
|
||||
// Define the epilogue
|
||||
using Epilogue = typename detail::DefaultConvEpilogue<
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpMmaTensorOp,
|
||||
kPartitionsK,
|
||||
EpilogueOutputOp
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
using Kernel = cutlass::conv::kernel::ImplicitGemmConvolution<
|
||||
Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
conv::Operator::kFprop,
|
||||
Conv2dProblemSize,
|
||||
GroupMode::kSingleGroup
|
||||
>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace kernel
|
||||
} // namespace conv
|
||||
} // namespace cutlass
|
||||
|
||||
@@ -39,14 +39,21 @@
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/conv/kernel/default_conv2d.h"
|
||||
#include "cutlass/conv/kernel/direct_convolution.h"
|
||||
|
||||
#include "cutlass/conv/threadblock/depthwise_mma_core_with_lane_access_size.h"
|
||||
|
||||
#include "cutlass/conv/threadblock/conv2d_fprop_activation_tile_access_iterator_analytic.h"
|
||||
|
||||
#include "cutlass/conv/threadblock/conv2d_fprop_filter_tile_access_iterator_analytic.h"
|
||||
#include "cutlass/conv/threadblock/depthwise_fprop_pipelined.h"
|
||||
|
||||
// Direct Conv Related Header files
|
||||
#include "cutlass/conv/threadblock/depthwise_fprop_activation_tile_access_iterator_direct_conv_optimized.h"
|
||||
#include "cutlass/conv/threadblock/depthwise_fprop_activation_tile_access_iterator_direct_conv_fixed_stride_dilation.h"
|
||||
|
||||
#include "cutlass/conv/threadblock/depthwise_fprop_filter_tile_access_iterator_direct_conv_optimized.h"
|
||||
#include "cutlass/conv/threadblock/depthwise_fprop_direct_conv_multistage.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
@@ -54,7 +61,7 @@ namespace conv {
|
||||
namespace kernel {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// Defines a kernel for Conv2dFprop
|
||||
/// Defines a kernel for DepthwiseFprop
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
@@ -80,12 +87,43 @@ template <
|
||||
int AlignmentB = cutlass::sizeof_bits<ElementB>::value / cutlass::sizeof_bits<ElementB>::value
|
||||
> struct DefaultDepthwiseFprop;
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// Defines a kernel for DepthwiseFprop with direct convolution algorithm
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
typename ElementB,
|
||||
typename LayoutB,
|
||||
typename ElementC,
|
||||
typename LayoutC,
|
||||
typename ElementAccumulator,
|
||||
typename OperatorClass,
|
||||
typename ArchTag,
|
||||
typename ThreadblockShape,
|
||||
typename ThreadBlockOutputShape,
|
||||
typename FilterShape,
|
||||
typename WarpShape,
|
||||
typename InstructionShape,
|
||||
typename EpilogueOutputOp,
|
||||
typename ThreadblockSwizzle,
|
||||
int Stages,
|
||||
typename MathOperatorTag,
|
||||
conv::IteratorAlgorithm IteratorAlgorithm = IteratorAlgorithm::kAnalytic,
|
||||
conv::StrideSupport StrideSupport = StrideSupport::kStrided,
|
||||
// MatrixShape<Height, Width>
|
||||
typename StrideShape = cutlass::MatrixShape<-1, -1>,
|
||||
// MatrixShape< Height, Width>
|
||||
typename DilationShape = cutlass::MatrixShape<-1, -1>,
|
||||
/// Access granularity of A matrix in units of elements
|
||||
int AlignmentA = 128 / cutlass::sizeof_bits<ElementA>::value,
|
||||
/// Access granularity of B matrix in units of elements
|
||||
int AlignmentB = 128 / cutlass::sizeof_bits<ElementB>::value
|
||||
> struct DefaultDepthwiseDirect2dConvFprop;
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
// OpClassSimt convolutions
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Defines a kernel for Depthwise specialization for Analytic IteratorAlgorithm,
|
||||
/// 2 stage pipeline, and FFMA-based mainloop for SM50
|
||||
/// Defines a kernel for Depthwise specialization for Analytic IteratorAlgorithm
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
@@ -210,6 +248,338 @@ struct DefaultDepthwiseFprop <
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// Defines a kernel for Depthwise specialization for direct 2d conv implementation,
|
||||
/// multiple stage pipeline, and SIMT-based mainloop
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
typename ElementB,
|
||||
typename LayoutB,
|
||||
typename ElementC,
|
||||
typename LayoutC,
|
||||
typename ElementAccumulator,
|
||||
typename ArchTag,
|
||||
typename ThreadblockShape,
|
||||
typename ThreadBlockOutputShape,
|
||||
typename FilterShape,
|
||||
typename WarpShape,
|
||||
typename InstructionShape,
|
||||
typename EpilogueOutputOp,
|
||||
typename ThreadblockSwizzle,
|
||||
int Stages,
|
||||
typename MathOperatorTag,
|
||||
conv::StrideSupport StrideSupport,
|
||||
typename StrideShape,
|
||||
typename DilationShape,
|
||||
int AlignmentA,
|
||||
int AlignmentB
|
||||
>
|
||||
struct DefaultDepthwiseDirect2dConvFprop <
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
ElementAccumulator,
|
||||
arch::OpClassSimt,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
ThreadBlockOutputShape,
|
||||
FilterShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
Stages,
|
||||
MathOperatorTag,
|
||||
IteratorAlgorithm::kOptimized,
|
||||
StrideSupport,
|
||||
StrideShape,
|
||||
DilationShape,
|
||||
AlignmentA,
|
||||
AlignmentB
|
||||
> {
|
||||
// One warp handles the entrie groups per cta.
|
||||
static_assert(ThreadblockShape::kN == WarpShape::kN,
|
||||
"ThreadblockShape::kN should be same as WarpShape::kN ");
|
||||
static_assert(ThreadblockShape::kK == FilterShape::kCount && WarpShape::kK == FilterShape::kCount,
|
||||
"ThreadblockShape::kK and WarpShape::kK should be same as filter size");
|
||||
static_assert(ThreadblockShape::kM % WarpShape::kM == 0,
|
||||
"ThreadblockShape::kM must be divisible by WarpShape shape::kM");
|
||||
static_assert(ThreadBlockOutputShape::kN, "ThreadBlockOutputShape::kN should be 1");
|
||||
|
||||
// Define the core components from GEMM
|
||||
using MmaCore = typename cutlass::conv::threadblock::DepthwiseDirectConvMmaCoreWithLaneAccessSize<
|
||||
ThreadblockShape,
|
||||
ThreadBlockOutputShape,
|
||||
FilterShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
ElementA,
|
||||
layout::RowMajor,
|
||||
ElementB,
|
||||
layout::ColumnMajor,
|
||||
ElementAccumulator,
|
||||
layout::RowMajor,
|
||||
arch::OpClassSimt,
|
||||
128,
|
||||
128,
|
||||
Stages,
|
||||
MathOperatorTag>;
|
||||
|
||||
// Define iterators over tiles from the A operand
|
||||
using ThreadMapA = typename MmaCore::IteratorThreadMapA;
|
||||
using IteratorA =
|
||||
cutlass::conv::threadblock::DepthwiseFpropActivationDirect2dConvTileAccessIteratorOptimized<
|
||||
cutlass::MatrixShape<ThreadblockShape::kM,ThreadblockShape::kN>, // < outputShape:KMNK, groups per cta>
|
||||
ThreadBlockOutputShape,
|
||||
ElementA, LayoutA,
|
||||
ThreadMapA
|
||||
>;
|
||||
|
||||
using SmemIteratorA = typename MmaCore::SmemIteratorA;
|
||||
|
||||
// Define iterators over tiles from the B operand
|
||||
using ThreadMapB = typename MmaCore::IteratorThreadMapB;
|
||||
using AccessTypeB = cutlass::AlignedArray<ElementB, AlignmentB>;
|
||||
using IteratorB =
|
||||
cutlass::conv::threadblock::DepthwiseFpropFilterDirectConvTileAccessIteratorOptimized<
|
||||
cutlass::MatrixShape<ThreadblockShape::kN, FilterShape::kCount>,
|
||||
ElementB, LayoutB,
|
||||
ThreadMapB
|
||||
>;
|
||||
|
||||
using SmemIteratorB = typename MmaCore::SmemIteratorB;
|
||||
|
||||
// Warp-level GEMM components
|
||||
using WarpMmaSimtOp = typename MmaCore::MmaWarpSimt;
|
||||
using MmaPolicy = typename MmaCore::MmaPolicy;
|
||||
using ThreadOutputShape = typename MmaCore::ThreadOutputShape;
|
||||
static cutlass::arch::CacheOperation::Kind const CacheOpA =
|
||||
((sizeof_bits<ElementA>::value * AlignmentA) == 128)
|
||||
? cutlass::arch::CacheOperation::Global
|
||||
: cutlass::arch::CacheOperation::Always;
|
||||
|
||||
static cutlass::arch::CacheOperation::Kind const CacheOpB =
|
||||
((sizeof_bits<ElementB>::value * AlignmentB) == 128)
|
||||
? cutlass::arch::CacheOperation::Global
|
||||
: cutlass::arch::CacheOperation::Always;
|
||||
|
||||
// Define the epilogue
|
||||
using Epilogue = typename epilogue::threadblock::DefaultDirectConvEpilogueSimt<
|
||||
ThreadblockShape, // < outputShape:KMNK, groups per cta>
|
||||
WarpMmaSimtOp,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount,
|
||||
ThreadOutputShape,
|
||||
ThreadBlockOutputShape
|
||||
>::Epilogue;
|
||||
|
||||
// Define the Mma
|
||||
using Mma = threadblock::DepthwiseFpropDirectConvMultipleStage<
|
||||
ThreadblockShape,
|
||||
IteratorA,
|
||||
SmemIteratorA,
|
||||
CacheOpA,
|
||||
IteratorB,
|
||||
SmemIteratorB,
|
||||
CacheOpB,
|
||||
MmaPolicy,
|
||||
Stages,
|
||||
Epilogue
|
||||
>;
|
||||
|
||||
// Define the kernel
|
||||
using Kernel = cutlass::conv::kernel::DirectConvolution<
|
||||
Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
conv::Operator::kFprop,
|
||||
Conv2dProblemSize,
|
||||
cutlass::conv::GroupMode::kDepthwise,
|
||||
ThreadBlockOutputShape
|
||||
>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// Defines a kernel for Depthwise specialization for direct 2d conv implementation,
|
||||
/// multiple stage pipeline, and SIMT-based mainloop
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
typename ElementB,
|
||||
typename LayoutB,
|
||||
typename ElementC,
|
||||
typename LayoutC,
|
||||
typename ElementAccumulator,
|
||||
typename ArchTag,
|
||||
typename ThreadblockShape,
|
||||
typename ThreadBlockOutputShape,
|
||||
typename FilterShape,
|
||||
typename WarpShape,
|
||||
typename InstructionShape,
|
||||
typename EpilogueOutputOp,
|
||||
typename ThreadblockSwizzle,
|
||||
int Stages,
|
||||
typename MathOperatorTag,
|
||||
conv::StrideSupport StrideSupport,
|
||||
typename StrideShape,
|
||||
typename DilationShape,
|
||||
int AlignmentA,
|
||||
int AlignmentB
|
||||
>
|
||||
struct DefaultDepthwiseDirect2dConvFprop <
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
ElementAccumulator,
|
||||
arch::OpClassSimt,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
ThreadBlockOutputShape,
|
||||
FilterShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
Stages,
|
||||
MathOperatorTag,
|
||||
IteratorAlgorithm::kFixedStrideDilation,
|
||||
StrideSupport,
|
||||
StrideShape,
|
||||
DilationShape,
|
||||
AlignmentA,
|
||||
AlignmentB,
|
||||
> {
|
||||
|
||||
|
||||
|
||||
// One warp handles the entrie groups per cta.
|
||||
static_assert(ThreadblockShape::kN == WarpShape::kN,
|
||||
"ThreadblockShape::kN should be same as WarpShape::kN ");
|
||||
static_assert(ThreadblockShape::kK == FilterShape::kCount && WarpShape::kK == FilterShape::kCount,
|
||||
"ThreadblockShape::kK and WarpShape::kK should be same as filter size");
|
||||
static_assert(ThreadblockShape::kM % WarpShape::kM == 0,
|
||||
"ThreadblockShape::kM must be divisible by WarpShape shape::kM");
|
||||
static_assert(ThreadBlockOutputShape::kN, "ThreadBlockOutputShape::kN should be 1");
|
||||
|
||||
static_assert(StrideShape::kRow >= 0 && StrideShape::kColumn >= 0, "Stride should be fixed");
|
||||
static_assert(DilationShape::kRow >= 0 && DilationShape::kColumn >= 0, "Stride should be fixed");
|
||||
|
||||
// Activations loaded by threadblock
|
||||
static int const ActivationShapeH = (ThreadBlockOutputShape::kH - 1) * StrideShape::kRow +
|
||||
(FilterShape::kRow - 1) * DilationShape::kRow + 1;
|
||||
|
||||
static int const ActivationShapeW = (ThreadBlockOutputShape::kW - 1) * StrideShape::kColumn +
|
||||
(FilterShape::kColumn - 1) * DilationShape::kColumn + 1;
|
||||
|
||||
using ActivationShape =
|
||||
cutlass::conv::TensorNHWCShape<1, ActivationShapeH, ActivationShapeW, ThreadblockShape::kN >;
|
||||
|
||||
// Define the core components from GEMM
|
||||
using MmaCore = typename cutlass::conv::threadblock::DepthwiseDirectConvMmaCoreWithLaneAccessSize<
|
||||
ThreadblockShape,
|
||||
ThreadBlockOutputShape,
|
||||
FilterShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
ElementA,
|
||||
layout::RowMajor,
|
||||
ElementB,
|
||||
layout::ColumnMajor,
|
||||
ElementAccumulator,
|
||||
layout::RowMajor,
|
||||
arch::OpClassSimt,
|
||||
128,
|
||||
128,
|
||||
Stages,
|
||||
MathOperatorTag,
|
||||
IteratorAlgorithm::kFixedStrideDilation,
|
||||
StrideShape,
|
||||
DilationShape,
|
||||
ActivationShape>;
|
||||
|
||||
// Define iterators over tiles from the A operand
|
||||
using ThreadMapA = typename MmaCore::IteratorThreadMapA;
|
||||
using IteratorA =
|
||||
cutlass::conv::threadblock::DepthwiseFpropActivationDirect2dConvTileAccessIteratorFixedStrideDilation<
|
||||
cutlass::MatrixShape<ThreadblockShape::kM,ThreadblockShape::kN>, // < outputShape:KMNK, groups per cta>
|
||||
ThreadBlockOutputShape,
|
||||
StrideShape,
|
||||
DilationShape,
|
||||
ActivationShape,
|
||||
ElementA, LayoutA,
|
||||
ThreadMapA
|
||||
>;
|
||||
|
||||
using SmemIteratorA = typename MmaCore::SmemIteratorA;
|
||||
|
||||
// Define iterators over tiles from the B operand
|
||||
using ThreadMapB = typename MmaCore::IteratorThreadMapB;
|
||||
using AccessTypeB = cutlass::AlignedArray<ElementB, AlignmentB>;
|
||||
using IteratorB =
|
||||
cutlass::conv::threadblock::DepthwiseFpropFilterDirectConvTileAccessIteratorOptimized<
|
||||
cutlass::MatrixShape<ThreadblockShape::kN, FilterShape::kCount>,
|
||||
ElementB, LayoutB,
|
||||
ThreadMapB
|
||||
>;
|
||||
|
||||
using SmemIteratorB = typename MmaCore::SmemIteratorB;
|
||||
|
||||
// Warp-level GEMM components
|
||||
using WarpMmaSimtOp = typename MmaCore::MmaWarpSimt;
|
||||
using MmaPolicy = typename MmaCore::MmaPolicy;
|
||||
using ThreadOutputShape = typename MmaCore::ThreadOutputShape;
|
||||
static cutlass::arch::CacheOperation::Kind const CacheOpA =
|
||||
((sizeof_bits<ElementA>::value * AlignmentA) == 128)
|
||||
? cutlass::arch::CacheOperation::Global
|
||||
: cutlass::arch::CacheOperation::Always;
|
||||
|
||||
static cutlass::arch::CacheOperation::Kind const CacheOpB =
|
||||
((sizeof_bits<ElementB>::value * AlignmentB) == 128)
|
||||
? cutlass::arch::CacheOperation::Global
|
||||
: cutlass::arch::CacheOperation::Always;
|
||||
|
||||
// Define the epilogue
|
||||
using Epilogue = typename epilogue::threadblock::DefaultDirectConvEpilogueSimt<
|
||||
ThreadblockShape, // < outputShape:KMNK, groups per cta>
|
||||
WarpMmaSimtOp,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount,
|
||||
ThreadOutputShape,
|
||||
ThreadBlockOutputShape
|
||||
>::Epilogue;
|
||||
|
||||
// Define the Mma
|
||||
using Mma = threadblock::DepthwiseFpropDirectConvMultipleStage<
|
||||
ThreadblockShape,
|
||||
IteratorA,
|
||||
SmemIteratorA,
|
||||
CacheOpA,
|
||||
IteratorB,
|
||||
SmemIteratorB,
|
||||
CacheOpB,
|
||||
MmaPolicy,
|
||||
Stages,
|
||||
Epilogue,
|
||||
IteratorAlgorithm::kFixedStrideDilation
|
||||
>;
|
||||
|
||||
// Define the kernel
|
||||
using Kernel = cutlass::conv::kernel::DirectConvolution<
|
||||
Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
conv::Operator::kFprop,
|
||||
Conv2dProblemSize,
|
||||
cutlass::conv::GroupMode::kDepthwise,
|
||||
ThreadBlockOutputShape
|
||||
>;
|
||||
};
|
||||
|
||||
} // namespace kernel
|
||||
} // namespace conv
|
||||
|
||||
@@ -0,0 +1,505 @@
|
||||
/***************************************************************************************************
|
||||
* 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 multi-staged Depthwise Convolution kernel.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
|
||||
#include "cutlass/aligned_buffer.h"
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/matrix_shape.h"
|
||||
#include "cutlass/semaphore.h"
|
||||
#include "cutlass/tensor_ref.h"
|
||||
#include "cutlass/layout/tensor.h"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/conv/convolution.h"
|
||||
#include "cutlass/conv/conv2d_problem_size.h"
|
||||
#include "cutlass/conv/conv3d_problem_size.h"
|
||||
#include "cutlass/epilogue/threadblock/output_iterator_parameter.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace conv {
|
||||
namespace kernel {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Parameters structure
|
||||
template <typename Mma_, ///! Threadblock-scoped matrix multiply-accumulate
|
||||
typename Epilogue_, ///! Epilogue
|
||||
typename ThreadblockSwizzle_, ///! Threadblock swizzling function
|
||||
conv::Operator ConvOperator, ///! Convolutional operator (Fprop, Dgrad, Wgrad)
|
||||
typename Arguments_, ///! Kernel Arguments
|
||||
typename ConvOutputIteratorParameter_, ///! Output Iterator Params
|
||||
typename ConvProblemSize_ = Conv2dProblemSize, ///! Convolutional operator on 2D or 3D problem
|
||||
conv::GroupMode GroupMode_ = conv::GroupMode::kNone, ///! Group mode
|
||||
typename ThreadBlockOutputShape_ = cutlass::conv::TensorNHWCShape<1, 1, 1, 1> > ///! OutputShape per ThreadBlock
|
||||
struct DirectConvolutionParams {
|
||||
using Mma = Mma_;
|
||||
using Epilogue = Epilogue_;
|
||||
using EpilogueOutputOp = typename Epilogue::OutputOp;
|
||||
using ThreadblockSwizzle = ThreadblockSwizzle_;
|
||||
using ThreadBlockOutputShape = ThreadBlockOutputShape_;
|
||||
static Operator const kConvolutionalOperator = ConvOperator;
|
||||
using ConvProblemSize = ConvProblemSize_;
|
||||
using Arguments = Arguments_;
|
||||
using ConvOutputIteratorParameter = ConvOutputIteratorParameter_;
|
||||
|
||||
using ThreadblockShape = typename Mma::Shape;
|
||||
static IteratorAlgorithm const kIteratorAlgorithm = Mma::IteratorA::kIteratorAlgorithm;
|
||||
static conv::GroupMode const kGroupMode = GroupMode_;
|
||||
static int const kStages = Mma::kStages;
|
||||
|
||||
ConvProblemSize problem_size;
|
||||
cutlass::gemm::GemmCoord grid_tiled_shape;
|
||||
gemm::GemmCoord implicit_gemm_problem_size;
|
||||
int swizzle_log_tile;
|
||||
int smem_size_;
|
||||
|
||||
int gemm_k_iterations;
|
||||
int gemm_k_iterations_per_channel;
|
||||
typename Mma::IteratorA::Params iterator_A;
|
||||
typename Mma::IteratorA::Element const *ptr_A;
|
||||
typename Mma::IteratorB::Params iterator_B;
|
||||
typename Mma::IteratorB::Element const *ptr_B;
|
||||
typename Mma::IteratorB::Element *ptr_reordered_B;
|
||||
typename Epilogue::OutputTileIterator::Params iterator_C;
|
||||
typename Epilogue::OutputTileIterator::Element *ptr_C;
|
||||
typename Epilogue::OutputTileIterator::Params iterator_D;
|
||||
typename Epilogue::OutputTileIterator::Element *ptr_D;
|
||||
typename EpilogueOutputOp::Params output_op;
|
||||
int *semaphore;
|
||||
SplitKMode split_k_mode;
|
||||
int split_k_slices;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
DirectConvolutionParams() : swizzle_log_tile(0), gemm_k_iterations(0) {}
|
||||
|
||||
///
|
||||
CUTLASS_HOST_DEVICE
|
||||
DirectConvolutionParams(Arguments const &args, int *semaphore = nullptr)
|
||||
: problem_size(args.problem_size),
|
||||
implicit_gemm_problem_size(
|
||||
cutlass::conv::implicit_gemm_problem_size(kConvolutionalOperator, args.problem_size)),
|
||||
iterator_A(Mma::IteratorA::getParams(args.problem_size, args.ref_A.layout())),
|
||||
ptr_A(args.ref_A.data()),
|
||||
iterator_B(Mma::IteratorB::getParams(args.problem_size, args.ref_B.layout())),
|
||||
ptr_B(args.ref_B.data()),
|
||||
ptr_reordered_B(args.ref_reordered_B.data()),
|
||||
iterator_C(ConvOutputIteratorParameter::layout(args.ref_C), args.problem_size),
|
||||
ptr_C(args.ref_C.data()),
|
||||
iterator_D(ConvOutputIteratorParameter::layout(args.ref_D), args.problem_size),
|
||||
ptr_D(args.ref_D.data()),
|
||||
output_op(args.output_op),
|
||||
semaphore(semaphore),
|
||||
split_k_mode(args.split_k_mode),
|
||||
split_k_slices(args.problem_size.split_k_slices) {
|
||||
gemm_k_iterations =
|
||||
depthwise_gemm_k_iterations<ThreadBlockOutputShape::kN,
|
||||
ThreadBlockOutputShape::kH,
|
||||
ThreadBlockOutputShape::kW>(kConvolutionalOperator,
|
||||
ThreadblockShape::kK,
|
||||
args.problem_size,
|
||||
kIteratorAlgorithm,
|
||||
kGroupMode,
|
||||
ThreadblockShape::kN);
|
||||
|
||||
gemm_k_iterations_per_channel = implicit_gemm_k_iterations_per_channel(
|
||||
kConvolutionalOperator, ThreadblockShape::kK, args.problem_size, kIteratorAlgorithm);
|
||||
|
||||
ThreadblockSwizzle threadblock_swizzle;
|
||||
|
||||
grid_tiled_shape = threadblock_swizzle.get_tiled_shape(
|
||||
kConvolutionalOperator,
|
||||
problem_size,
|
||||
{ThreadblockShape::kM, ThreadblockShape::kN, ThreadblockShape::kK},
|
||||
args.problem_size.split_k_slices);
|
||||
|
||||
swizzle_log_tile = threadblock_swizzle.get_log_tile(grid_tiled_shape);
|
||||
|
||||
// Dynamic SMEM usage because stride and dilation are runtime params.
|
||||
smem_size_ = (iterator_A.activation_size * kStages + iterator_B.filter_size);
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
int get_smem_size() {
|
||||
// Dynamic Smem Size
|
||||
return smem_size_;
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
template <typename Params_, typename ElementB_>
|
||||
struct ReorderKernel {
|
||||
using Params = Params_;
|
||||
using ElementB = ElementB_;
|
||||
|
||||
union SharedStorage {};
|
||||
|
||||
static unsigned int const kReorderKernelThreadPerCTA = 128;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
ReorderKernel() {}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static dim3 get_grid_shape(Params const ¶ms) {
|
||||
return dim3{static_cast<unsigned int>(
|
||||
(params.problem_size.filter_size() + kReorderKernelThreadPerCTA - 1) /
|
||||
kReorderKernelThreadPerCTA),
|
||||
1,
|
||||
1};
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static dim3 get_block_shape() { return dim3{kReorderKernelThreadPerCTA, 1, 1}; }
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
void operator()(Params const ¶ms, SharedStorage &shared_storage) {
|
||||
int64_t m = static_cast<int64_t>(params.problem_size.groups);
|
||||
int64_t n = static_cast<int64_t>(params.problem_size.filter_size() / params.problem_size.K);
|
||||
const ElementB *src_with_type = static_cast<const ElementB *>(params.ptr_B);
|
||||
ElementB *dst_with_type = static_cast<ElementB *>(params.ptr_reordered_B);
|
||||
|
||||
int64_t linear_index = blockIdx.x * kReorderKernelThreadPerCTA + threadIdx.x;
|
||||
int64_t index_m = linear_index / n;
|
||||
int64_t index_n = linear_index % n;
|
||||
int64_t new_linear_index = index_m + index_n * m;
|
||||
|
||||
if (linear_index < m * n) {
|
||||
dst_with_type[new_linear_index] = src_with_type[linear_index];
|
||||
}
|
||||
return;
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename Mma_, ///! Threadblock-scoped matrix multiply-accumulate
|
||||
typename Epilogue_, ///! Epilogue
|
||||
typename ThreadblockSwizzle_, ///! Threadblock swizzling function
|
||||
conv::Operator ConvOperator, ///! Convolutional operator (Fprop, Dgrad, Wgrad)
|
||||
typename ConvProblemSize_ = Conv2dProblemSize, ///! Convolutional operator on 2D or 3D problem
|
||||
conv::GroupMode GroupMode_ = conv::GroupMode::kNone, ///! Group mode
|
||||
typename ThreadBlockOutputShape_ = cutlass::conv::TensorNHWCShape<1, 1, 1, 1>
|
||||
>
|
||||
struct DirectConvolution {
|
||||
|
||||
using Mma = Mma_;
|
||||
using Epilogue = Epilogue_;
|
||||
using EpilogueOutputOp = typename Epilogue::OutputOp;
|
||||
using ThreadblockSwizzle = ThreadblockSwizzle_;
|
||||
using ThreadBlockOutputShape = ThreadBlockOutputShape_;
|
||||
static Operator const kConvolutionalOperator = ConvOperator;
|
||||
|
||||
using ElementA = typename Mma::IteratorA::Element;
|
||||
using LayoutA = typename Mma::IteratorA::Layout;
|
||||
using ElementB = typename Mma::IteratorB::Element;
|
||||
using LayoutB = typename Mma::IteratorB::Layout;
|
||||
using ElementC = typename EpilogueOutputOp::ElementOutput;
|
||||
|
||||
/// Set output tensor C layout
|
||||
using LayoutC = LayoutA;
|
||||
|
||||
using ElementAccumulator = typename EpilogueOutputOp::ElementAccumulator;
|
||||
using ElementCompute = typename EpilogueOutputOp::ElementCompute;
|
||||
|
||||
using WarpMmaOperator = typename Mma::Policy::Operator;
|
||||
|
||||
using ArchMmaOperator = typename WarpMmaOperator::ArchMmaOperator;
|
||||
using MathOperator = typename ArchMmaOperator::Operator;
|
||||
|
||||
using OperatorClass = typename WarpMmaOperator::OperatorClass;
|
||||
using ArchTag = typename WarpMmaOperator::ArchTag;
|
||||
|
||||
using ThreadblockShape = typename Mma::Shape;
|
||||
using WarpShape = typename WarpMmaOperator::Shape;
|
||||
using InstructionShape = typename cutlass::gemm::GemmShape<1, 1, 1>;
|
||||
|
||||
static int const kStages = Mma::kStages;
|
||||
static IteratorAlgorithm const kIteratorAlgorithm = Mma::IteratorA::kIteratorAlgorithm;
|
||||
static StrideSupport const kStrideSupport = Mma::IteratorA::kStrideSupport;
|
||||
|
||||
/// Warp count (concept: GemmShape)
|
||||
using WarpCount = typename Mma::WarpCount;
|
||||
static int const kThreadCount = 32 * WarpCount::kCount;
|
||||
|
||||
using TensorRefA = typename Mma::IteratorA::TensorRef;
|
||||
using TensorRefB = typename Mma::IteratorB::TensorRef;
|
||||
using TensorRefC = cutlass::TensorRef<ElementC, LayoutC>;
|
||||
|
||||
/// Check iterator A and B convolution dimension are the same and
|
||||
// set device::ImplicitGemmConvolution::kConvDim
|
||||
static_assert(Mma::IteratorA::kConvDim == Mma::IteratorB::kConvDim,
|
||||
"Convolution on different different dimensions is not supported");
|
||||
static int const kConvDim = Mma::IteratorA::kConvDim;
|
||||
|
||||
/// Conv dimension and problem size structure (Conv2d or Conv3d)
|
||||
using ConvProblemSize = ConvProblemSize_;
|
||||
|
||||
static conv::GroupMode const kGroupMode = GroupMode_;
|
||||
|
||||
|
||||
//
|
||||
//
|
||||
//
|
||||
using ConvOutputIteratorParameter = epilogue::threadblock::ConvOutputIteratorParameter<
|
||||
LayoutC,
|
||||
typename Epilogue::OutputTileIterator::Layout,
|
||||
TensorRefC,
|
||||
ConvOperator,
|
||||
ConvProblemSize
|
||||
>;
|
||||
|
||||
|
||||
/// Argument structure
|
||||
struct Arguments {
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
ConvProblemSize problem_size;
|
||||
TensorRefA ref_A;
|
||||
TensorRefB ref_B;
|
||||
TensorRefB ref_reordered_B;
|
||||
TensorRefC ref_C;
|
||||
TensorRefC ref_D;
|
||||
typename EpilogueOutputOp::Params output_op;
|
||||
SplitKMode split_k_mode;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
/// Default ctor
|
||||
CUTLASS_HOST_DEVICE
|
||||
Arguments() { }
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Arguments(
|
||||
ConvProblemSize const & problem_size
|
||||
):
|
||||
problem_size(problem_size) { }
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Arguments(
|
||||
ConvProblemSize const & problem_size,
|
||||
TensorRefA const & ref_A,
|
||||
TensorRefB const & ref_B,
|
||||
TensorRefC const & ref_C,
|
||||
TensorRefC const & ref_D,
|
||||
typename EpilogueOutputOp::Params const & output_op,
|
||||
TensorRefB const & ref_reordered_B = nullptr,
|
||||
SplitKMode const & split_k_mode = SplitKMode::kSerial
|
||||
):
|
||||
problem_size(problem_size),
|
||||
ref_A(ref_A),
|
||||
ref_B(ref_B),
|
||||
ref_C(ref_C),
|
||||
ref_D(ref_D),
|
||||
output_op(output_op),
|
||||
ref_reordered_B(ref_reordered_B),
|
||||
split_k_mode(split_k_mode)
|
||||
{
|
||||
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
using Params =
|
||||
typename cutlass::conv::kernel::DirectConvolutionParams<Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
kConvolutionalOperator,
|
||||
Arguments,
|
||||
ConvOutputIteratorParameter,
|
||||
ConvProblemSize,
|
||||
kGroupMode,
|
||||
ThreadBlockOutputShape>;
|
||||
|
||||
using ReorderKernel = typename cutlass::conv::kernel::ReorderKernel<Params, ElementB>;
|
||||
|
||||
/// Shared memory storage structure
|
||||
union SharedStorage {
|
||||
typename Mma::SharedStorage main_loop;
|
||||
typename Epilogue::SharedStorage epilogue;
|
||||
};
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
DirectConvolution() { }
|
||||
|
||||
/// Executes one ImplicitGEMM
|
||||
CUTLASS_DEVICE
|
||||
void operator()(Params const ¶ms, SharedStorage &shared_storage) {
|
||||
|
||||
// Compute threadblock location
|
||||
ThreadblockSwizzle threadblock_swizzle;
|
||||
|
||||
cutlass::gemm::GemmCoord threadblock_tile_idx =
|
||||
threadblock_swizzle.get_tile_offset(params.swizzle_log_tile);
|
||||
|
||||
// Early exit if threadblock is out of range
|
||||
if (params.grid_tiled_shape.m() <= threadblock_tile_idx.m() ||
|
||||
params.grid_tiled_shape.n() <= threadblock_tile_idx.n()) {
|
||||
|
||||
return;
|
||||
}
|
||||
|
||||
// Compute position within threadblock
|
||||
int thread_idx = threadIdx.x;
|
||||
int iterator_column_offset = 0;
|
||||
int filter_row_offset = 0;
|
||||
if (kGroupMode != GroupMode::kNone) {
|
||||
if (kGroupMode == GroupMode::kDepthwise) {
|
||||
iterator_column_offset += threadblock_tile_idx.n() * Mma::Shape::kN;
|
||||
}
|
||||
}
|
||||
|
||||
// Construct iterators to A and B operands
|
||||
typename Mma::IteratorA iterator_A(
|
||||
params.iterator_A,
|
||||
params.problem_size,
|
||||
params.ptr_A,
|
||||
thread_idx,
|
||||
MatrixCoord(
|
||||
threadblock_tile_idx.m() + threadblock_tile_idx.k(),
|
||||
iterator_column_offset
|
||||
)
|
||||
);
|
||||
|
||||
typename Mma::IteratorB iterator_B(
|
||||
params.iterator_B,
|
||||
params.problem_size,
|
||||
params.ptr_reordered_B,
|
||||
thread_idx,
|
||||
MatrixCoord(
|
||||
filter_row_offset,
|
||||
iterator_column_offset
|
||||
)
|
||||
);
|
||||
|
||||
// Broadcast the warp_id computed by lane 0 to ensure dependent code
|
||||
// is compiled as warp-uniform.
|
||||
int warp_idx = __shfl_sync(0xffffffff, threadIdx.x / 32, 0);
|
||||
int lane_idx = threadIdx.x % 32;
|
||||
|
||||
//
|
||||
// Main loop
|
||||
//
|
||||
|
||||
// Construct thread-scoped matrix multiply
|
||||
Mma mma(shared_storage.main_loop, thread_idx, warp_idx, lane_idx);
|
||||
|
||||
typename Mma::FragmentC accumulators;
|
||||
|
||||
accumulators.clear();
|
||||
|
||||
//
|
||||
// Epilogue
|
||||
//
|
||||
|
||||
EpilogueOutputOp output_op(params.output_op);
|
||||
|
||||
// Compute logical position within grid
|
||||
threadblock_tile_idx =
|
||||
threadblock_swizzle.get_tile_offset(params.swizzle_log_tile);
|
||||
|
||||
|
||||
MatrixCoord threadblock_offset(
|
||||
threadblock_tile_idx.m() + threadblock_tile_idx.k(),
|
||||
threadblock_tile_idx.n() * Mma::Shape::kN
|
||||
);
|
||||
|
||||
// Tile iterator writing to destination tensor
|
||||
typename Epilogue::OutputTileIterator iterator_D(
|
||||
params.iterator_D,
|
||||
params.ptr_D,
|
||||
ConvOutputIteratorParameter::extent(params.problem_size),
|
||||
thread_idx,
|
||||
threadblock_offset
|
||||
);
|
||||
|
||||
// Tile iterator reading from source accumulator tensor
|
||||
typename Epilogue::OutputTileIterator iterator_C(
|
||||
params.iterator_C,
|
||||
params.ptr_C,
|
||||
ConvOutputIteratorParameter::extent(params.problem_size),
|
||||
thread_idx,
|
||||
threadblock_offset
|
||||
);
|
||||
|
||||
|
||||
// Construct the epilogue
|
||||
Epilogue epilogue(
|
||||
shared_storage.epilogue,
|
||||
thread_idx,
|
||||
warp_idx,
|
||||
lane_idx);
|
||||
|
||||
|
||||
// Compute threadblock-scoped matrix multiply-add
|
||||
// Epilogue is fused in the mainloop
|
||||
mma(params.gemm_k_iterations,
|
||||
accumulators,
|
||||
iterator_A,
|
||||
params.iterator_A,
|
||||
iterator_B,
|
||||
params.iterator_B,
|
||||
accumulators,
|
||||
epilogue,
|
||||
output_op,
|
||||
iterator_D,
|
||||
iterator_C,
|
||||
params.split_k_slices);
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace kernel
|
||||
} // namespace conv
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,325 @@
|
||||
/***************************************************************************************************
|
||||
* 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 exposing architecture support for depthwise convolution
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/tensor_ref.h"
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/arch/mma.h"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/gemm/thread/mma.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace conv {
|
||||
namespace thread {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// MMA operation
|
||||
template <
|
||||
/// Size of the matrix product (concept: GemmShape)
|
||||
typename Shape_,
|
||||
/// Number of threads participating
|
||||
int kThreads_,
|
||||
/// Data type of A elements
|
||||
typename ElementA,
|
||||
/// Data type of B elements
|
||||
typename ElementB,
|
||||
/// Element type of C matrix
|
||||
typename ElementC,
|
||||
/// Inner product operator
|
||||
typename Operator
|
||||
>
|
||||
struct ElementwiseInnerProduct;
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// General implementation
|
||||
template <
|
||||
/// Size of the matrix product (concept: GemmShape)
|
||||
typename Shape_,
|
||||
/// Data type of A elements
|
||||
typename ElementA_,
|
||||
/// Data type of B elements
|
||||
typename ElementB_,
|
||||
/// Element type of C matrix
|
||||
typename ElementC_>
|
||||
struct ElementwiseInnerProduct<Shape_, 1, ElementA_, ElementB_, ElementC_, arch::OpMultiplyAdd> {
|
||||
using Shape = Shape_;
|
||||
using Operator = arch::OpMultiplyAdd;
|
||||
using ElementC = ElementC_;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
void operator()(Array<ElementC_, Shape::kN> &d,
|
||||
Array<ElementA_, Shape::kN> const &a,
|
||||
Array<ElementB_, Shape::kN> const &b,
|
||||
Array<ElementC_, Shape::kN> const &c) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < Shape::kN; ++i) {
|
||||
d[i] = a[i] * b[i] + c[i];
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// Specialization of half_t
|
||||
template <>
|
||||
struct ElementwiseInnerProduct<
|
||||
gemm::GemmShape<2, 2, 1>,
|
||||
1,
|
||||
half_t,
|
||||
half_t,
|
||||
half_t,
|
||||
arch::OpMultiplyAdd> {
|
||||
|
||||
using Shape = gemm::GemmShape<2, 2, 1>;
|
||||
using Operator = arch::OpMultiplyAdd;
|
||||
using ElementC = half_t;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
void operator()(
|
||||
Array<half_t, 2> &d,
|
||||
Array<half_t, 2> const &a,
|
||||
Array<half_t, 2> const &b,
|
||||
Array<half_t, 2> const &c
|
||||
) {
|
||||
|
||||
#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 600))
|
||||
|
||||
__half2 const & A = reinterpret_cast<__half2 const &>(a);
|
||||
__half2 const & B = reinterpret_cast<__half2 const &>(b);
|
||||
__half2 const & C = reinterpret_cast<__half2 const &>(c);
|
||||
|
||||
__half2 tmp_D = __hfma2(A, B, C);
|
||||
|
||||
d = reinterpret_cast<Array<half_t, 2> const &>(tmp_D);
|
||||
|
||||
#else
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < 2; ++i) {
|
||||
d[i] = a[i] * b[i] + c[i];
|
||||
}
|
||||
#endif
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Structure to compute the matrix product
|
||||
template <
|
||||
/// Size of the Gemm problem - concept: gemm::GemmShape<>
|
||||
typename Shape,
|
||||
/// Data type of A elements
|
||||
typename ElementA,
|
||||
/// Data type of B elements
|
||||
typename ElementB,
|
||||
/// Element type of C matrix
|
||||
typename ElementC,
|
||||
/// Concept: arch::OpMultiplyAdd or arch::Mma<>
|
||||
typename Operator = arch::OpMultiplyAdd,
|
||||
/// Used for partial specialization
|
||||
typename Enable = bool
|
||||
>
|
||||
struct DepthwiseDirectConvElementwiseInnerProduct;
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Gemplate that handles all packed matrix layouts
|
||||
template <
|
||||
/// Size of the Gemm problem - concept: gemm::GemmShape<>
|
||||
typename Shape_,
|
||||
/// Data type of A elements
|
||||
typename ElementA_,
|
||||
/// Data type of B elements
|
||||
typename ElementB_,
|
||||
/// Element type of C matrix
|
||||
typename ElementC_,
|
||||
/// Operator used to compute GEMM
|
||||
typename Operator_
|
||||
>
|
||||
struct DepthwiseDirectConvElementwiseInnerProductGeneric {
|
||||
|
||||
/// Size of the Gemm problem - concept: gemm::GemmShape<>
|
||||
using Shape = Shape_;
|
||||
|
||||
/// Data type of operand A
|
||||
using ElementA = ElementA_;
|
||||
|
||||
/// Data type of operand B
|
||||
using ElementB = ElementB_;
|
||||
|
||||
/// Element type of operand C
|
||||
using ElementC = ElementC_;
|
||||
|
||||
/// Underlying mathematical operator
|
||||
using Operator = Operator_;
|
||||
|
||||
/// A operand storage
|
||||
using FragmentA = Array<ElementA, Shape::kMN>;
|
||||
|
||||
/// B operand storage
|
||||
using FragmentB = Array<ElementB, Shape::kN>;
|
||||
|
||||
/// C operand storage
|
||||
using FragmentC = Array<ElementC, Shape::kMN>;
|
||||
|
||||
/// Instruction
|
||||
using MmaOp = cutlass::conv::thread::ElementwiseInnerProduct<
|
||||
gemm::GemmShape<Shape::kN, Shape::kN, 1>,
|
||||
1,
|
||||
ElementA,
|
||||
ElementB,
|
||||
ElementC,
|
||||
Operator>;
|
||||
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
/// Computes a matrix product D = A * B + C
|
||||
CUTLASS_HOST_DEVICE
|
||||
void operator()(
|
||||
FragmentC & D,
|
||||
FragmentA const & A,
|
||||
FragmentB const & B,
|
||||
FragmentC const & C) {
|
||||
Array<ElementC, Shape::kN> *ptr_D = reinterpret_cast<Array<ElementC, Shape::kN> *>(&D);
|
||||
Array<ElementA, Shape::kN> const *ptr_A =
|
||||
reinterpret_cast<Array<ElementA, Shape::kN> const *>(&A);
|
||||
Array<ElementB, Shape::kN> const *ptr_B =
|
||||
reinterpret_cast<Array<ElementB, Shape::kN> const *>(&B);
|
||||
|
||||
MmaOp mma_op;
|
||||
|
||||
// Copy accumulators
|
||||
D = C;
|
||||
|
||||
// Compute matrix product
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int n = 0; n < Shape::kN / MmaOp::Shape::kN; ++n) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int m = 0; m < Shape::kM; ++m) {
|
||||
|
||||
Array<ElementC, MmaOp::Shape::kN> tmpD = ptr_D[m * Shape::kN / MmaOp::Shape::kN + n];
|
||||
Array<ElementA, MmaOp::Shape::kN> tmpA = ptr_A[m * Shape::kN / MmaOp::Shape::kN + n];
|
||||
Array<ElementB, MmaOp::Shape::kN> tmpB = ptr_B[n];
|
||||
|
||||
mma_op(tmpD, tmpA, tmpB, tmpD);
|
||||
|
||||
ptr_D[m * Shape::kN / MmaOp::Shape::kN + n] = tmpD;
|
||||
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Structure to compute the matrix product
|
||||
template <
|
||||
/// Size of the Gemm problem - concept: gemm::GemmShape<>
|
||||
typename Shape_,
|
||||
/// Data type of A elements
|
||||
typename ElementA_,
|
||||
/// Data type of B elements
|
||||
typename ElementB_,
|
||||
/// Element type of C matrix
|
||||
typename ElementC_
|
||||
>
|
||||
struct DepthwiseDirectConvElementwiseInnerProduct<
|
||||
Shape_,
|
||||
ElementA_,
|
||||
ElementB_,
|
||||
ElementC_,
|
||||
arch::OpMultiplyAdd
|
||||
> {
|
||||
/// Size of the Gemm problem - concept: gemm::GemmShape<>
|
||||
using Shape = Shape_;
|
||||
|
||||
/// Data type of operand A
|
||||
using ElementA = ElementA_;
|
||||
|
||||
/// Data type of operand B
|
||||
using ElementB = ElementB_;
|
||||
|
||||
/// Element type of operand C
|
||||
using ElementC = ElementC_;
|
||||
|
||||
/// Underlying mathematical operator
|
||||
using Operator = arch::OpMultiplyAdd;
|
||||
|
||||
/// A operand storage
|
||||
using FragmentA =
|
||||
Array<ElementA, Shape::kMN>; // output_tile_size per thread * groups_per_thread
|
||||
|
||||
/// B operand storage
|
||||
using FragmentB = Array<ElementB, Shape::kN>; // 1 * groups_per_thread
|
||||
|
||||
/// C operand storage
|
||||
using FragmentC =
|
||||
Array<ElementC, Shape::kMN>; // output_tile_size per thread * groups_per_thread
|
||||
|
||||
static bool const use_optimized = 0;
|
||||
|
||||
using ArchMmaOperator = DepthwiseDirectConvElementwiseInnerProductGeneric<Shape,
|
||||
ElementA,
|
||||
ElementB,
|
||||
ElementC,
|
||||
Operator>;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
/// Computes a matrix product D = A * B + C
|
||||
CUTLASS_HOST_DEVICE
|
||||
void operator()(
|
||||
FragmentC & D,
|
||||
FragmentA const & A,
|
||||
FragmentB const & B,
|
||||
FragmentC const & C) {
|
||||
|
||||
ArchMmaOperator mma;
|
||||
|
||||
mma(D, A, B, C);
|
||||
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace thread
|
||||
} // namespace conv
|
||||
} // namespace cutlass
|
||||
+4
-2
@@ -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;
|
||||
|
||||
@@ -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
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
+314
@@ -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
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
+291
@@ -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
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
+261
@@ -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
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -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
|
||||
|
||||
@@ -42,6 +42,9 @@
|
||||
#include "cutlass/gemm/warp/mma.h"
|
||||
|
||||
#include "cutlass/gemm/thread/mma.h"
|
||||
#include "cutlass/conv/convolution.h"
|
||||
#include "cutlass/conv/thread/depthwise_mma.h"
|
||||
|
||||
|
||||
#include "cutlass/gemm/warp/mma_simt_tile_iterator.h"
|
||||
#include "cutlass/gemm/warp/mma_simt_policy.h"
|
||||
@@ -91,7 +94,7 @@ class MmaDepthwiseSimt
|
||||
|
||||
public:
|
||||
/// Shape of warp-level matrix operation (concept: GemmShape)
|
||||
using Shape = Shape_; // < 64, 16 , 8>
|
||||
using Shape = Shape_;
|
||||
|
||||
/// Data type of multiplicand A
|
||||
using ElementA = ElementA_;
|
||||
@@ -156,8 +159,223 @@ public:
|
||||
MmaDepthwiseSimt():Base() {}
|
||||
};
|
||||
|
||||
/// Structure to compute the matrix product targeting CUDA cores and SIMT math instructions.
|
||||
template <
|
||||
/// Size of the Gemm problem - concept: gemm::GemmShape<>
|
||||
typename Shape_,
|
||||
/// Shape of filter shape per threadblock - concept: gemm::GemmShape<Depth, Height, Width>
|
||||
typename FilterShape_,
|
||||
/// Shape of the output tile computed by thread- concept: conv::TensorNHWCShape<>
|
||||
typename ThreadOutputShape_,
|
||||
/// Shape of the output tile computed by threadblock - concept: conv::TensorNHWCShape<>
|
||||
typename ThreadBlockOutputShape_,
|
||||
/// Data type of A elements
|
||||
typename ElementA_,
|
||||
/// Layout of A matrix (concept: MatrixLayout)
|
||||
typename LayoutA_,
|
||||
/// Data type of B elements
|
||||
typename ElementB_,
|
||||
/// Layout of B matrix (concept: MatrixLayout)
|
||||
typename LayoutB_,
|
||||
/// Element type of C matrix
|
||||
typename ElementC_,
|
||||
/// Layout of C matrix (concept: MatrixLayout)
|
||||
typename LayoutC_,
|
||||
/// Shape of the warp in units of thread (concept: MmaSimtPolicy)
|
||||
typename Policy_,
|
||||
/// 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>,
|
||||
/// Number of partitions along K dimension
|
||||
int PartitionsK = 1,
|
||||
/// Complex transformation on operand A
|
||||
ComplexTransform TransformA = ComplexTransform::kNone,
|
||||
/// Complex transformation on operand B
|
||||
ComplexTransform TransformB = ComplexTransform::kNone,
|
||||
/// Used for partial specialization
|
||||
typename Enable = bool>
|
||||
class MmaDepthwiseDirectConvSimt {
|
||||
public:
|
||||
/// Shape of warp-level matrix operation (concept: GemmShape)
|
||||
using Shape = Shape_;
|
||||
|
||||
/// Shape of filter shape per threadblock - concept: gemm::GemmShape<Depth, Height, Width>
|
||||
using FilterShape = FilterShape_;
|
||||
|
||||
/// Shape of the output tile computed by thread- concept: conv::TensorNHWCShape<>
|
||||
using ThreadOutputShape = ThreadOutputShape_;
|
||||
|
||||
/// Shape of the output tile computed by threadblock - concept: conv::TensorNHWCShape<>
|
||||
using ThreadBlockOutputShape = ThreadBlockOutputShape_;
|
||||
|
||||
/// Data type of multiplicand A
|
||||
using ElementA = ElementA_;
|
||||
|
||||
/// Layout of multiplicand A
|
||||
using LayoutA = LayoutA_;
|
||||
|
||||
/// Data type of multiplicand B
|
||||
using ElementB = ElementB_;
|
||||
|
||||
/// Layout of multiplicand B
|
||||
using LayoutB = LayoutB_;
|
||||
|
||||
/// Data type of accumulator matrix C
|
||||
using ElementC = ElementC_;
|
||||
|
||||
/// Layout of accumulator matrix C
|
||||
using LayoutC = LayoutC_;
|
||||
|
||||
/// Shape of the warp in units of thread (concept: MmaLanePolicySimt)
|
||||
using Policy = Policy_;
|
||||
|
||||
/// Iterator algo type
|
||||
static conv::IteratorAlgorithm const IteratorAlgorithm = IteratorAlgorithm_;
|
||||
|
||||
/// Stride ( MatrixShape<Height, Width> )
|
||||
using StrideShape = StrideShape_;
|
||||
|
||||
/// Dilation ( MatrixShape<Height, Width> )
|
||||
using DilationShape = DilationShape_;
|
||||
|
||||
/// Activation Shape loaded by threadblock
|
||||
using ActivationShape = ActivationShape_;
|
||||
|
||||
/// Indicates class of matrix operator
|
||||
using OperatorClass = arch::OpClassSimt;
|
||||
|
||||
/// Hard-coded for now
|
||||
using ArchTag = arch::Sm50;
|
||||
|
||||
/// Complex transform on A operand
|
||||
static ComplexTransform const kTransformA = TransformA;
|
||||
|
||||
/// Complex transform on B operand
|
||||
static ComplexTransform const kTransformB = TransformB;
|
||||
|
||||
static constexpr bool use_dp4a = (platform::is_same< layout::ColumnMajorInterleaved<4>, LayoutA>::value ||
|
||||
platform::is_same< layout::RowMajorInterleaved<4>, LayoutA >::value) &&
|
||||
platform::is_same< ElementA, int8_t >::value &&
|
||||
platform::is_same< ElementB, int8_t >::value;
|
||||
|
||||
using dp4a_type = typename platform::conditional< use_dp4a , int8_t, bool >::type;
|
||||
|
||||
/// Thread-level matrix multiply accumulate operator
|
||||
using ThreadMma = cutlass::conv::thread::DepthwiseDirectConvElementwiseInnerProduct<
|
||||
cutlass::gemm::GemmShape<
|
||||
Shape::kM / Policy::WarpShape::kRow, // number of output pixels proccessed per thread
|
||||
Shape::kN / Policy::WarpShape::kColumn, // number of channels proccessed per thread
|
||||
1>,
|
||||
ElementA,
|
||||
ElementB,
|
||||
ElementC,
|
||||
arch::OpMultiplyAdd,
|
||||
dp4a_type
|
||||
>;
|
||||
|
||||
/// Underlying matrix multiply operator (concept: arch::Mma)
|
||||
using ArchMmaOperator = typename ThreadMma::ArchMmaOperator;
|
||||
|
||||
/// Indicates math operator
|
||||
using MathOperator = typename ArchMmaOperator::Operator;
|
||||
|
||||
/// Shape of the underlying instruction
|
||||
using InstructionShape = cutlass::gemm::GemmShape<1,1,use_dp4a ? 4 : 1>;
|
||||
|
||||
public:
|
||||
|
||||
/// Iterates over the A operand in memory
|
||||
using IteratorA = cutlass::conv::warp::DepthwiseDirect2dConvSimtTileIterator<
|
||||
MatrixShape<Shape::kM, Shape::kN>, // <output tile=(P*Q), output channels> per warp
|
||||
FilterShape,
|
||||
ThreadOutputShape,
|
||||
ThreadBlockOutputShape,
|
||||
cutlass::gemm::Operand::kA,
|
||||
ElementA,
|
||||
Policy,
|
||||
IteratorAlgorithm,
|
||||
StrideShape,
|
||||
DilationShape,
|
||||
ActivationShape,
|
||||
PartitionsK,
|
||||
Shape::kK
|
||||
>;
|
||||
|
||||
/// Storage for A tile
|
||||
using FragmentA = typename IteratorA::Fragment;
|
||||
|
||||
/// Storage for transformed A tile
|
||||
using TransformedFragmentA = FragmentA;
|
||||
|
||||
/// Iterates over the B operand in memory
|
||||
using IteratorB = cutlass::gemm::warp::MmaSimtTileIterator<
|
||||
MatrixShape<1, Shape::kN>,
|
||||
cutlass::gemm::Operand::kB,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
Policy,
|
||||
PartitionsK,
|
||||
Shape::kK
|
||||
>;
|
||||
|
||||
/// Storage for B tile
|
||||
using FragmentB = typename IteratorB::Fragment;
|
||||
|
||||
/// Storage for transformed A tile
|
||||
using TransformedFragmentB = FragmentB;
|
||||
|
||||
/// Iterates over the C operand in memory
|
||||
using IteratorC = cutlass::gemm::warp::MmaSimtTileIterator<
|
||||
MatrixShape<Shape::kM, Shape::kN>,
|
||||
cutlass::gemm::Operand::kC,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
Policy
|
||||
>;
|
||||
|
||||
/// Storage for C tile
|
||||
using FragmentC = typename ThreadMma::FragmentC;
|
||||
|
||||
public:
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
/// Ctor
|
||||
CUTLASS_DEVICE
|
||||
MmaDepthwiseDirectConvSimt() {}
|
||||
|
||||
/// Performs a warp-level matrix multiply-accumulate operation
|
||||
CUTLASS_DEVICE
|
||||
void operator()(
|
||||
FragmentC &d,
|
||||
FragmentA a,
|
||||
FragmentB b,
|
||||
FragmentC const &c, int group_idx = 0) const {
|
||||
|
||||
ThreadMma mma;
|
||||
|
||||
mma(d, a, b, c);
|
||||
}
|
||||
|
||||
/// Transform the mma operands to the required types
|
||||
CUTLASS_DEVICE
|
||||
void transform(TransformedFragmentA &dst_A, TransformedFragmentB &dst_B,
|
||||
FragmentA const &A, FragmentB const &B) const {
|
||||
//TODO: Implement this
|
||||
dst_A = A;
|
||||
dst_B = B;
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace warp
|
||||
} // namespace gemm
|
||||
} // namespace conv
|
||||
} // namespace cutlass
|
||||
|
||||
@@ -40,6 +40,8 @@
|
||||
#include "cutlass/tensor_ref.h"
|
||||
#include "cutlass/matrix_shape.h"
|
||||
|
||||
#include "cutlass/conv/convolution.h"
|
||||
|
||||
#include "cutlass/arch/memory_sm75.h"
|
||||
|
||||
#include "cutlass/layout/matrix.h"
|
||||
@@ -250,6 +252,611 @@ private:
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
/// Size of the matrix to load (concept: MatrixShape)
|
||||
typename Shape_,
|
||||
/// Size of filter (concept: gemm::GemmShape<Depth, Height, Width>)
|
||||
typename FilterShape_,
|
||||
/// Size of the matrix to load (concept: MatrixShape)
|
||||
typename ThreadOutputShape_,
|
||||
/// Size of the matrix to load (concept: MatrixShape)
|
||||
typename ThreadBlockOutputShape_,
|
||||
/// Operand identity
|
||||
cutlass::gemm::Operand Operand,
|
||||
/// Data type of A elements
|
||||
typename Element_,
|
||||
/// Shape of the warp in units of thread (concept: MmaSimtPolicy)
|
||||
typename Policy_,
|
||||
/// 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>,
|
||||
/// Number of partitions along K dimension - used in sliced-K
|
||||
int PartitionsK = 1,
|
||||
/// Group Size along kPartition - used in sliced-K
|
||||
int PartitionGroupSize = 1>
|
||||
class DepthwiseDirect2dConvSimtTileIterator;
|
||||
|
||||
|
||||
/// Specialization for A operands of row-major layouts
|
||||
///
|
||||
/// Concept: MutableRandomAccessContiguousTileIteratorConcept
|
||||
///
|
||||
template <
|
||||
/// Size of the matrix to load (concept: MatrixShape)
|
||||
typename Shape_,
|
||||
/// Size of filter (concept: gemm::GemmShape<Depth, Height, Width>)
|
||||
typename FilterShape_,
|
||||
/// Size of the matrix to load (concept: TensorNHWC)
|
||||
typename ThreadOutputShape_,
|
||||
/// Size of the matrix to load (concept: TensorNHWC)
|
||||
typename ThreadBlockOutputShape_,
|
||||
/// Data type of A elements
|
||||
typename Element_,
|
||||
/// Shape of the warp in units of thread (concept: MmaSimtPolicy)
|
||||
typename Policy_,
|
||||
/// Iterator algo type
|
||||
conv::IteratorAlgorithm IteratorAlgorithm,
|
||||
/// Stride ( MatrixShape<Height, Width> )
|
||||
typename StrideShape,
|
||||
/// Dilation ( MatrixShape<Height, Width> )
|
||||
typename DilationShape,
|
||||
/// Activation Shape loaded by threadblock
|
||||
typename ActivationShape,
|
||||
/// Number of partitions along K dimension - used in sliced-K
|
||||
int PartitionsK,
|
||||
/// Group Size along kPartition - used in sliced-K
|
||||
int PartitionGroupSize>
|
||||
class DepthwiseDirect2dConvSimtTileIterator<Shape_,
|
||||
FilterShape_,
|
||||
ThreadOutputShape_,
|
||||
ThreadBlockOutputShape_,
|
||||
cutlass::gemm::Operand::kA,
|
||||
Element_,
|
||||
Policy_,
|
||||
IteratorAlgorithm,
|
||||
StrideShape,
|
||||
DilationShape,
|
||||
ActivationShape,
|
||||
PartitionsK,
|
||||
PartitionGroupSize> {
|
||||
public:
|
||||
/// Shape of tile to load (concept: MatrixShape)
|
||||
using Shape = Shape_;
|
||||
|
||||
/// Shape of filter (concept: gemm::GemmShape<Depth, Height, Width>)
|
||||
using FilterShape = FilterShape_;
|
||||
|
||||
/// Shape of tile to load (concept: TensorNHWC)
|
||||
using ThreadOutputShape = ThreadOutputShape_;
|
||||
|
||||
/// Shape of tile to load (concept: TensorNHWC)
|
||||
using ThreadBlockOutputShape = ThreadBlockOutputShape_;
|
||||
|
||||
/// Operand tag
|
||||
static cutlass::gemm::Operand const kOperand = cutlass::gemm::Operand::kA;
|
||||
|
||||
/// Element type
|
||||
using Element = Element_;
|
||||
|
||||
/// Layout of policy
|
||||
using Layout = layout::RowMajor;
|
||||
|
||||
/// Decomposition of elements among threads
|
||||
using Policy = Policy_;
|
||||
|
||||
/// TensorRef type for loading element from a tensor
|
||||
using TensorRef = TensorRef<Element, Layout>;
|
||||
|
||||
/// Index type
|
||||
using Index = typename TensorRef::Index;
|
||||
|
||||
/// Long Index type
|
||||
using LongIndex = typename TensorRef::LongIndex;
|
||||
|
||||
/// Coordinate for an element in the tensor
|
||||
using TensorCoord = typename TensorRef::TensorCoord;
|
||||
|
||||
//
|
||||
// Derived quantities
|
||||
//
|
||||
|
||||
static_assert(!(Shape::kRow % Policy::WarpShape::kRow),
|
||||
"The warp-level GEMM M size must be divisible by the number of threads arranged along the M dimension.");
|
||||
|
||||
static_assert(Shape::kRow > 0, "Shape::kRow must be greater than zero.");
|
||||
static_assert(Shape::kColumn > 0, "Shape::kColumn must be greater than zero.");
|
||||
static_assert(Policy::WarpShape::kRow > 0, "Policy::WarpShape::kRow must be greater than zero.");
|
||||
static_assert(Shape::kRow / Policy::WarpShape::kRow > 0, "Shape::kRow / Policy::WarpShape::kRow must be greater than zero.");
|
||||
|
||||
// Thread-level shape of a fragment
|
||||
using ThreadShape = MatrixShape<
|
||||
ThreadOutputShape::kNHW, // Output tile shape Computed by current threads
|
||||
ThreadOutputShape::kC
|
||||
>;
|
||||
|
||||
static_assert(!(ThreadShape::kColumn % Policy::LaneMmaShape::kN),
|
||||
"Thread-level GEMM must be divisible by Policy::LaneMmaShape.");
|
||||
|
||||
/// Number of individual loads
|
||||
using Iterations = MatrixShape<
|
||||
ThreadShape::kRow,
|
||||
ThreadShape::kColumn / Policy::LaneMmaShape::kN
|
||||
>;
|
||||
|
||||
using ThreadTileCount = MatrixShape<
|
||||
ThreadBlockOutputShape::kH / ThreadOutputShape::kH,
|
||||
ThreadBlockOutputShape::kW / ThreadOutputShape::kW
|
||||
>;
|
||||
|
||||
/// Fragment object holding a thread's part of a tile
|
||||
using Fragment = Array<Element, ThreadShape::kCount>;
|
||||
|
||||
protected:
|
||||
|
||||
/// Internal reference
|
||||
cutlass::TensorRef<Array<Element, Policy::LaneMmaShape::kN>, layout::RowMajor> ref_;
|
||||
|
||||
int activation_offset[ThreadOutputShape::kH][ThreadOutputShape::kW][Iterations::kColumn];
|
||||
int iterator_r_;
|
||||
int iterator_s_;
|
||||
int iterator_offset_;
|
||||
|
||||
int inc_next_s_ ;
|
||||
int inc_next_r_ ;
|
||||
|
||||
MatrixCoord lane_offset_;
|
||||
public:
|
||||
|
||||
/// Default ctor constructs null iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
DepthwiseDirect2dConvSimtTileIterator() { }
|
||||
|
||||
/// Constructor from TensorRef
|
||||
CUTLASS_HOST_DEVICE
|
||||
DepthwiseDirect2dConvSimtTileIterator(
|
||||
TensorRef ref,
|
||||
int lane_id
|
||||
) {
|
||||
|
||||
// compute offset based on thread ID and lane layout
|
||||
typename Policy::LaneLayout lane_layout = Policy::get_lane_layout();
|
||||
|
||||
// Set channel offset
|
||||
lane_offset_ = lane_layout.inverse(lane_id) * MatrixCoord(0, Policy::LaneMmaShape::kN);
|
||||
|
||||
ref.add_coord_offset(lane_offset_);
|
||||
|
||||
ref_.reset(reinterpret_cast<Array<Element, Policy::LaneMmaShape::kN> *>(ref.data()),
|
||||
ref.stride(0) / Policy::LaneMmaShape::kN);
|
||||
|
||||
iterator_r_ = 0;
|
||||
iterator_s_ = 0;
|
||||
iterator_offset_ = 0;
|
||||
}
|
||||
|
||||
/// Adds a pointer offset to internal pointer(s) to advance through memory
|
||||
CUTLASS_HOST_DEVICE
|
||||
DepthwiseDirect2dConvSimtTileIterator &add_pointer_offset(LongIndex offset) {
|
||||
ref_.add_pointer_offset(offset);
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory at the location pointed to by the iterator.
|
||||
template<typename Params>
|
||||
CUTLASS_HOST_DEVICE
|
||||
void setup_initial_status(Params const& params) {
|
||||
|
||||
inc_next_s_ = params.inc_next[0];
|
||||
inc_next_r_ = params.inc_next[1];
|
||||
|
||||
// Get base HW offset of current threads
|
||||
int threadgroup = threadIdx.x / (ThreadBlockOutputShape::kC / ThreadOutputShape::kC);
|
||||
int base_p_ =
|
||||
(threadgroup / (ThreadTileCount::kColumn)) * ThreadOutputShape::kH;
|
||||
int base_q_ =
|
||||
(threadgroup % (ThreadTileCount::kColumn)) * ThreadOutputShape::kW;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int p = 0; p < ThreadOutputShape::kH; ++p) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int q = 0; q < ThreadOutputShape::kW; ++q) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int col = 0; col < Iterations::kColumn; ++col) {
|
||||
int base_w = (base_q_ + q) * params.stride[0];
|
||||
int base_h = (base_p_ + p) * params.stride[1];
|
||||
|
||||
int offset = base_h * params.activation_tile_w + base_w;
|
||||
activation_offset[p][q][col] = offset;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
/// Advances an iterator along logical dimensions of matrix in units of whole tiles
|
||||
CUTLASS_HOST_DEVICE
|
||||
DepthwiseDirect2dConvSimtTileIterator &add_tile_offset(TensorCoord const &coord) {
|
||||
// Set warp row and col start
|
||||
lane_offset_ = MatrixCoord({lane_offset_.row() + coord.row() * Shape::kRow, lane_offset_.column()});
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances an iterator along logical dimensions of matrix in units of whole tiles
|
||||
CUTLASS_HOST_DEVICE
|
||||
void advance(int32_t pointer_offset) {
|
||||
ref_.reset(ref_.data() + pointer_offset / sizeof(Element) / Policy::LaneMmaShape::kN);
|
||||
iterator_s_ = 0;
|
||||
iterator_r_ = 0;
|
||||
iterator_offset_ = 0;
|
||||
}
|
||||
|
||||
/// Advances the iterator along the advance dimension
|
||||
CUTLASS_HOST_DEVICE
|
||||
DepthwiseDirect2dConvSimtTileIterator &operator++() {
|
||||
++iterator_s_;
|
||||
if (iterator_s_ < FilterShape::kColumn) {
|
||||
iterator_offset_ += inc_next_s_;
|
||||
|
||||
return *this;
|
||||
}
|
||||
|
||||
iterator_s_ = 0;
|
||||
|
||||
++iterator_r_;
|
||||
if (iterator_r_ < FilterShape::kRow) {
|
||||
iterator_offset_ += inc_next_r_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
iterator_r_ = 0;
|
||||
iterator_offset_ = 0;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the iterator along the advance dimension
|
||||
CUTLASS_HOST_DEVICE
|
||||
DepthwiseDirect2dConvSimtTileIterator & operator--() {
|
||||
// Do nothing
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory at the location pointed to by the iterator. (vector loads)
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load_with_pointer_offset(Fragment &frag, Index pointer_offset) const {
|
||||
|
||||
Array<Element, Policy::LaneMmaShape::kN> *dst_ptr =
|
||||
reinterpret_cast<Array<Element, Policy::LaneMmaShape::kN> *>(&frag);
|
||||
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int p = 0; p < ThreadOutputShape::kH; ++p) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int q = 0; q < ThreadOutputShape::kW; ++q) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int n = 0; n < Iterations::kColumn; ++n) {
|
||||
void const *ptr = ref_.data() +
|
||||
ref_.offset({activation_offset[p][q][n] + (iterator_offset_),
|
||||
n * Policy::WarpShape::kColumn}) +
|
||||
pointer_offset / Policy::LaneMmaShape::kN;
|
||||
arch::shared_load(dst_ptr[n + q + p * ThreadOutputShape::kW], ptr);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory at the location pointed to by the iterator.
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load(Fragment &frag) const {
|
||||
load_with_pointer_offset(frag, 0);
|
||||
}
|
||||
|
||||
/// Stores a fragment to memory at the location pointed to by the iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
void store_with_pointer_offset(Fragment const &frag, Index pointer_offset) const {
|
||||
// Do nothing at present.
|
||||
}
|
||||
|
||||
/// Stores a fragment to memory at the location pointed to by the iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
void store(Fragment const &frag, Index pointer_offset) const {
|
||||
store_with_pointer_offset(frag, 0);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void set_kgroup_index(int k_group) {
|
||||
// no operation here
|
||||
}
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// Specialization for A operands of row-major layouts
|
||||
///
|
||||
/// Concept: MutableRandomAccessContiguousTileIteratorConcept
|
||||
///
|
||||
template <
|
||||
/// Size of the matrix to load (concept: MatrixShape)
|
||||
typename Shape_,
|
||||
/// Size of filter (concept: gemm::GemmShape<Depth, Height, Width>)
|
||||
typename FilterShape_,
|
||||
/// Size of the matrix to load (concept: TensorNHWC)
|
||||
typename ThreadOutputShape_,
|
||||
/// Size of the matrix to load (concept: TensorNHWC)
|
||||
typename ThreadBlockOutputShape_,
|
||||
/// Data type of A elements
|
||||
typename Element_,
|
||||
/// Shape of the warp in units of thread (concept: MmaSimtPolicy)
|
||||
typename Policy_,
|
||||
/// Stride ( MatrixShape<Height, Width> )
|
||||
typename StrideShape_,
|
||||
/// Dilation ( MatrixShape<Height, Width> )
|
||||
typename DilationShape_,
|
||||
/// Activation Shape loaded by threadblock
|
||||
typename ActivationShape_,
|
||||
/// Number of partitions along K dimension - used in sliced-K
|
||||
int PartitionsK,
|
||||
/// Group Size along kPartition - used in sliced-K
|
||||
int PartitionGroupSize>
|
||||
class DepthwiseDirect2dConvSimtTileIterator<Shape_,
|
||||
FilterShape_,
|
||||
ThreadOutputShape_,
|
||||
ThreadBlockOutputShape_,
|
||||
cutlass::gemm::Operand::kA,
|
||||
Element_,
|
||||
Policy_,
|
||||
IteratorAlgorithm::kFixedStrideDilation,
|
||||
StrideShape_,
|
||||
DilationShape_,
|
||||
ActivationShape_,
|
||||
PartitionsK,
|
||||
PartitionGroupSize> {
|
||||
public:
|
||||
/// Shape of tile to load (concept: MatrixShape)
|
||||
using Shape = Shape_;
|
||||
|
||||
/// Shape of filter (concept: gemm::GemmShape<Depth, Height, Width>)
|
||||
using FilterShape = FilterShape_;
|
||||
|
||||
/// Shape of tile to load (concept: TensorNHWC)
|
||||
using ThreadOutputShape = ThreadOutputShape_;
|
||||
|
||||
/// Shape of tile to load (concept: TensorNHWC)
|
||||
using ThreadBlockOutputShape = ThreadBlockOutputShape_;
|
||||
|
||||
/// Stride ( MatrixShape<Height, Width> )
|
||||
using StrideShape = StrideShape_;
|
||||
|
||||
/// Dilation ( MatrixShape<Height, Width> )
|
||||
using DilationShape = DilationShape_;
|
||||
|
||||
/// Activation Shape loaded by threadblock
|
||||
using ActivationShape = ActivationShape_;
|
||||
|
||||
/// Operand tag
|
||||
static cutlass::gemm::Operand const kOperand = cutlass::gemm::Operand::kA;
|
||||
|
||||
/// Element type
|
||||
using Element = Element_;
|
||||
|
||||
/// Layout of policy
|
||||
using Layout = layout::RowMajor;
|
||||
|
||||
/// Decomposition of elements among threads
|
||||
using Policy = Policy_;
|
||||
|
||||
/// TensorRef type for loading element from a tensor
|
||||
using TensorRef = TensorRef<Element, Layout>;
|
||||
|
||||
/// Index type
|
||||
using Index = typename TensorRef::Index;
|
||||
|
||||
/// Long Index type
|
||||
using LongIndex = typename TensorRef::LongIndex;
|
||||
|
||||
/// Coordinate for an element in the tensor
|
||||
using TensorCoord = typename TensorRef::TensorCoord;
|
||||
|
||||
//
|
||||
// Derived quantities
|
||||
//
|
||||
|
||||
static_assert(!(Shape::kRow % Policy::WarpShape::kRow),
|
||||
"The warp-level GEMM M size must be divisible by the number of threads arranged "
|
||||
"along the M dimension.");
|
||||
|
||||
static_assert(Shape::kRow > 0, "Shape::kRow must be greater than zero.");
|
||||
static_assert(Shape::kColumn > 0, "Shape::kColumn must be greater than zero.");
|
||||
static_assert(Policy::WarpShape::kRow > 0, "Policy::WarpShape::kRow must be greater than zero.");
|
||||
static_assert(Shape::kRow / Policy::WarpShape::kRow > 0,
|
||||
"Shape::kRow / Policy::WarpShape::kRow must be greater than zero.");
|
||||
|
||||
// Activations loaded by threadblock
|
||||
static int const ThreadActivationShapeH = (ThreadOutputShape::kH - 1) * StrideShape::kRow +
|
||||
(FilterShape::kRow - 1) * DilationShape::kRow + 1;
|
||||
|
||||
static int const ThreadActivationShapeW = (ThreadOutputShape::kW - 1) * StrideShape::kColumn +
|
||||
(FilterShape::kColumn - 1) * DilationShape::kColumn + 1;
|
||||
|
||||
using ThreadActivationShape = cutlass::conv::
|
||||
TensorNHWCShape<1, ThreadActivationShapeH, ThreadActivationShapeW, ThreadOutputShape::kC>;
|
||||
|
||||
// Thread-level shape of a fragment
|
||||
using ThreadShape =
|
||||
MatrixShape<ThreadOutputShape::kNHW,
|
||||
ThreadOutputShape::kC>;
|
||||
|
||||
static_assert(!(ThreadShape::kColumn % Policy::LaneMmaShape::kN),
|
||||
"Thread-level GEMM must be divisible by Policy::LaneMmaShape.");
|
||||
|
||||
/// Number of individual loads
|
||||
using Iterations =
|
||||
MatrixShape<ThreadShape::kRow, ThreadShape::kColumn / Policy::LaneMmaShape::kN>;
|
||||
|
||||
using ThreadTileCount = MatrixShape<ThreadBlockOutputShape::kH / ThreadOutputShape::kH,
|
||||
ThreadBlockOutputShape::kW / ThreadOutputShape::kW>;
|
||||
|
||||
/// Fragment object holding a thread's part of a tile
|
||||
using Fragment = Array<Element, ThreadShape::kCount>;
|
||||
|
||||
protected:
|
||||
/// Internal reference
|
||||
cutlass::TensorRef<Array<Element, Policy::LaneMmaShape::kN>, layout::RowMajor> ref_;
|
||||
|
||||
Array<Element, Policy::LaneMmaShape::kN>
|
||||
activation[ThreadActivationShape::kH][ThreadActivationShape::kW][Iterations::kColumn];
|
||||
int iterator_r_;
|
||||
int iterator_s_;
|
||||
|
||||
|
||||
MatrixCoord lane_offset_;
|
||||
|
||||
public:
|
||||
/// Default ctor constructs null iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
DepthwiseDirect2dConvSimtTileIterator() {}
|
||||
|
||||
/// Constructor from TensorRef
|
||||
CUTLASS_HOST_DEVICE
|
||||
DepthwiseDirect2dConvSimtTileIterator(TensorRef ref, int lane_id) {
|
||||
// compute offset based on thread ID and lane layout
|
||||
typename Policy::LaneLayout lane_layout = Policy::get_lane_layout();
|
||||
|
||||
// Set channel offset
|
||||
lane_offset_ = lane_layout.inverse(lane_id) * MatrixCoord(0, Policy::LaneMmaShape::kN);
|
||||
|
||||
ref.add_coord_offset(lane_offset_);
|
||||
|
||||
ref_.reset(reinterpret_cast<Array<Element, Policy::LaneMmaShape::kN> *>(ref.data()),
|
||||
ref.stride(0) / Policy::LaneMmaShape::kN);
|
||||
|
||||
iterator_r_ = 0;
|
||||
iterator_s_ = 0;
|
||||
}
|
||||
|
||||
/// Adds a pointer offset to internal pointer(s) to advance through memory
|
||||
CUTLASS_HOST_DEVICE
|
||||
DepthwiseDirect2dConvSimtTileIterator &add_pointer_offset(LongIndex offset) {
|
||||
ref_.add_pointer_offset(offset);
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory at the location pointed to by the iterator.
|
||||
template <typename Params>
|
||||
CUTLASS_HOST_DEVICE void setup_initial_status(
|
||||
Params const ¶ms) {
|
||||
|
||||
// Get base HW offset of current threads
|
||||
int threadgroup = threadIdx.x / (ThreadBlockOutputShape::kC / ThreadOutputShape::kC);
|
||||
int base_h =
|
||||
(threadgroup / (ThreadTileCount::kColumn)) * ThreadOutputShape::kH * StrideShape::kRow;
|
||||
int base_w =
|
||||
(threadgroup % (ThreadTileCount::kColumn)) * ThreadOutputShape::kW * StrideShape::kColumn;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int h = 0; h < ThreadActivationShape::kH; ++h) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int w = 0; w < ThreadActivationShape::kW; ++w) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int col = 0; col < Iterations::kColumn; ++col) {
|
||||
int offset = (base_h + h) * ActivationShape::kW + (base_w + w);
|
||||
|
||||
void const *ptr = ref_.data() + ref_.offset({offset, col * Policy::WarpShape::kColumn});
|
||||
arch::shared_load(activation[h][w][col], ptr);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Advances an iterator along logical dimensions of matrix in units of whole tiles
|
||||
CUTLASS_HOST_DEVICE
|
||||
DepthwiseDirect2dConvSimtTileIterator &add_tile_offset(TensorCoord const &coord) {
|
||||
// Set warp row and col start
|
||||
lane_offset_ =
|
||||
MatrixCoord({lane_offset_.row() + coord.row() * Shape::kRow, lane_offset_.column()});
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances an iterator along logical dimensions of matrix in units of whole tiles
|
||||
CUTLASS_HOST_DEVICE
|
||||
void advance(int32_t pointer_offset) {
|
||||
ref_.reset(ref_.data() + pointer_offset / sizeof(Element) / Policy::LaneMmaShape::kN);
|
||||
iterator_s_ = 0;
|
||||
iterator_r_ = 0;
|
||||
}
|
||||
|
||||
/// Advances the iterator along the advance dimension
|
||||
CUTLASS_HOST_DEVICE
|
||||
DepthwiseDirect2dConvSimtTileIterator &operator++() {
|
||||
++iterator_s_;
|
||||
if (iterator_s_ < FilterShape::kColumn) {
|
||||
return *this;
|
||||
}
|
||||
|
||||
iterator_s_ = 0;
|
||||
|
||||
++iterator_r_;
|
||||
if (iterator_r_ < FilterShape::kRow) {
|
||||
return *this;
|
||||
}
|
||||
|
||||
iterator_r_ = 0;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Advances the iterator along the advance dimension
|
||||
CUTLASS_HOST_DEVICE
|
||||
DepthwiseDirect2dConvSimtTileIterator &operator--() {
|
||||
// Do nothing
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory at the location pointed to by the iterator. (vector loads)
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load_with_pointer_offset(Fragment &frag, Index pointer_offset) const {
|
||||
Array<Element, Policy::LaneMmaShape::kN> *dst_ptr =
|
||||
reinterpret_cast<Array<Element, Policy::LaneMmaShape::kN> *>(&frag);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int p = 0; p < ThreadOutputShape::kH; ++p) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int q = 0; q < ThreadOutputShape::kW; ++q) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int n = 0; n < Iterations::kColumn; ++n) {
|
||||
const int h = p * StrideShape::kRow + iterator_r_ * DilationShape::kRow;
|
||||
const int w = q * StrideShape::kColumn + iterator_s_ * DilationShape::kColumn;
|
||||
|
||||
dst_ptr[n + q + p * ThreadOutputShape::kW] = activation[h][w][n];
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Loads a fragment from memory at the location pointed to by the iterator.
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load(Fragment &frag) const { load_with_pointer_offset(frag, 0); }
|
||||
|
||||
/// Stores a fragment to memory at the location pointed to by the iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
void store_with_pointer_offset(Fragment const &frag, Index pointer_offset) const {
|
||||
// Do nothing at present.
|
||||
}
|
||||
|
||||
/// Stores a fragment to memory at the location pointed to by the iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
void store(Fragment const &frag, Index pointer_offset) const {
|
||||
store_with_pointer_offset(frag, 0);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void set_kgroup_index(int k_group) {
|
||||
// no operation here
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace warp
|
||||
} // namespace gemm
|
||||
} // namespace conv
|
||||
} // namespace cutlass
|
||||
|
||||
Reference in New Issue
Block a user