CUTLASS 3.6.0 (#1850)
* v3.6 * update changelog * update readme * fix typo * fixing typos * hopper gemm with weight prefetch --------- Co-authored-by: yuzhai <yuzhai@nvidia.com> Co-authored-by: Haicheng Wu <haichengw@nvidia.com>
This commit is contained in:
co-authored by
yuzhai
Haicheng Wu
parent
0837a2a00a
commit
cc3c29a81a
@@ -41,8 +41,8 @@
|
||||
#include "cute/algorithm/functional.hpp"
|
||||
#include "cute/algorithm/gemm.hpp"
|
||||
|
||||
#include "cutlass/conv/detail.hpp"
|
||||
#include "cutlass/conv/convolution.h"
|
||||
#include "cutlass/conv/convnd_problem_shape.hpp"
|
||||
#include "cutlass/conv/dispatch_policy.hpp"
|
||||
#include "cutlass/pipeline/pipeline.hpp"
|
||||
#include "cutlass/util/packed_stride.hpp"
|
||||
@@ -103,6 +103,8 @@ struct CollectiveConv<
|
||||
|
||||
using PipelineParams = typename MainloopPipeline::Params;
|
||||
using PipelineState = typename cutlass::PipelineState<DispatchPolicy::Stages>;
|
||||
|
||||
using ProblemShape = ConvProblemShape<ConvOp, NumSpatialDimensions>;
|
||||
|
||||
// TODO: move pipeline mode tiling into the collective setup phase instead
|
||||
static_assert(rank(SmemLayoutA{}) == 3, "SmemLayout must be rank 3 (M/N, K, PIPE)");
|
||||
@@ -143,7 +145,7 @@ struct CollectiveConv<
|
||||
|
||||
struct SharedStorage
|
||||
{
|
||||
struct TensorStorage : cute::aligned_struct<128> {
|
||||
struct TensorStorage : cute::aligned_struct<128, _0> {
|
||||
cute::array_aligned<typename TiledMma::ValTypeA, cute::cosize_v<SmemLayoutA>> smem_A;
|
||||
cute::array_aligned<typename TiledMma::ValTypeB, cute::cosize_v<SmemLayoutB>> smem_B;
|
||||
} tensors;
|
||||
@@ -162,8 +164,6 @@ struct CollectiveConv<
|
||||
|
||||
// Host side kernel arguments
|
||||
struct Arguments {
|
||||
using ProblemShape = ConvProblemShape<ConvOp, NumSpatialDimensions>;
|
||||
ProblemShape problem_shape{};
|
||||
ElementA const* ptr_A{nullptr};
|
||||
ElementB const* ptr_B{nullptr};
|
||||
};
|
||||
@@ -175,7 +175,7 @@ private:
|
||||
// Get tma_load_a instantce.
|
||||
template <class TensorA>
|
||||
static constexpr auto
|
||||
get_tma_load_a_instance(TensorA const& tensor_a, typename Arguments::ProblemShape const& problem_shape) {
|
||||
get_tma_load_a_instance(TensorA const& tensor_a, ProblemShape const& problem_shape) {
|
||||
if constexpr (is_im2col_A) {
|
||||
// compute the upper and lower corners based on the conv padding
|
||||
auto lower_corner_whd = detail::compute_lower_corner_whd(problem_shape);
|
||||
@@ -218,7 +218,7 @@ private:
|
||||
// Get tma_load_b instantce.
|
||||
template <class TensorB>
|
||||
static constexpr auto
|
||||
get_tma_load_b_instance(TensorB const& tensor_b, typename Arguments::ProblemShape const& problem_shape) {
|
||||
get_tma_load_b_instance(TensorB const& tensor_b, ProblemShape const& problem_shape) {
|
||||
// TMA im2col mode for tensor B in wgrad kernel.
|
||||
if constexpr (is_im2col_B) {
|
||||
// compute the upper and lower corners based on the conv padding
|
||||
@@ -250,24 +250,25 @@ private:
|
||||
}
|
||||
}
|
||||
|
||||
public:
|
||||
|
||||
// Performs im2col transformations on the input of type ConvProblemShape
|
||||
static constexpr auto
|
||||
get_problem_shape_MNKL(typename Arguments::ProblemShape const& problem_shape) {
|
||||
get_problem_shape_MNKL(ProblemShape const& problem_shape) {
|
||||
|
||||
if constexpr (is_im2col_A || is_im2col_B) {
|
||||
// transformation + im2col linearization
|
||||
return problem_shape.get_linearized_problem_shape_MNKL();
|
||||
return cutlass::conv::detail::get_linearized_problem_shape_MNKL(problem_shape);
|
||||
}
|
||||
else {
|
||||
// transformation
|
||||
return problem_shape.get_transformed_problem_shape_MNKL();
|
||||
return cutlass::conv::detail::get_transformed_problem_shape_MNKL(problem_shape);
|
||||
}
|
||||
}
|
||||
|
||||
public:
|
||||
|
||||
// Device side kernel params
|
||||
struct Params {
|
||||
using _Submode = decltype(take<0,NumTensorDimensions-1>(typename Arguments::ProblemShape::TensorExtent{}));
|
||||
using ProblemShape = decltype(get_problem_shape_MNKL(typename Arguments::ProblemShape{}));
|
||||
using _Submode = decltype(take<0,NumTensorDimensions-1>(typename ProblemShape::TensorExtent{}));
|
||||
|
||||
// Assumption: StrideA is congruent with Problem_MK
|
||||
// Select TMA load type according to convolution operator.
|
||||
@@ -294,7 +295,6 @@ public:
|
||||
// Members
|
||||
TMA_A tma_load_a;
|
||||
TMA_B tma_load_b;
|
||||
ProblemShape problem_shape;
|
||||
uint32_t tma_transaction_bytes = TmaTransactionBytes;
|
||||
};
|
||||
|
||||
@@ -304,19 +304,19 @@ public:
|
||||
|
||||
// Lowers the host side user facing arguments to the kernel facing lauch params
|
||||
static constexpr Params
|
||||
to_underlying_arguments(Arguments const& args, void* workspace) {
|
||||
to_underlying_arguments(ProblemShape const& problem_shape, Arguments const& args, void* workspace) {
|
||||
(void) workspace;
|
||||
// from the flat problem shape arrays of ConvProblemShape<ConvOp, N>, create a rank-3 MNK problem shape tuple
|
||||
// tma desc creation depends on the original untransformed domain.
|
||||
|
||||
// A extents.
|
||||
auto shape_A_orig = args.problem_shape.get_shape_A();
|
||||
auto shape_A_orig = problem_shape.get_shape_A();
|
||||
// B extents.
|
||||
auto shape_B_orig = args.problem_shape.get_shape_B();
|
||||
auto shape_B_orig = problem_shape.get_shape_B();
|
||||
|
||||
// Fill inferred cute strides from flat stride arrays
|
||||
auto dA = make_cute_packed_stride(StrideA{}, args.problem_shape.stride_A, ConvOp);
|
||||
auto dB = make_cute_packed_stride(StrideB{}, args.problem_shape.stride_B, ConvOp);
|
||||
auto dA = make_cute_packed_stride(StrideA{}, problem_shape.stride_A, ConvOp);
|
||||
auto dB = make_cute_packed_stride(StrideB{}, problem_shape.stride_B, ConvOp);
|
||||
|
||||
auto ptr_A = reinterpret_cast<InternalElementA const*>(args.ptr_A);
|
||||
auto ptr_B = reinterpret_cast<InternalElementB const*>(args.ptr_B);
|
||||
@@ -324,20 +324,17 @@ public:
|
||||
Tensor tensor_a = make_tensor(make_gmem_ptr(ptr_A), make_layout(shape_A_orig, dA));
|
||||
Tensor tensor_b = make_tensor(make_gmem_ptr(ptr_B), make_layout(shape_B_orig, dB));
|
||||
|
||||
auto tma_load_a = get_tma_load_a_instance(tensor_a, args.problem_shape);
|
||||
auto tma_load_b = get_tma_load_b_instance(tensor_b, args.problem_shape);
|
||||
|
||||
auto problem_shape_mnkl = get_problem_shape_MNKL(args.problem_shape);
|
||||
auto tma_load_a = get_tma_load_a_instance(tensor_a, problem_shape);
|
||||
auto tma_load_b = get_tma_load_b_instance(tensor_b, problem_shape);
|
||||
|
||||
return {
|
||||
tma_load_a,
|
||||
tma_load_b,
|
||||
problem_shape_mnkl,
|
||||
TmaTransactionBytes
|
||||
};
|
||||
}
|
||||
|
||||
template<class ProblemShape>
|
||||
|
||||
template <class ProblemShape>
|
||||
static bool
|
||||
can_implement(
|
||||
ProblemShape const& problem_shape,
|
||||
@@ -345,14 +342,14 @@ public:
|
||||
// Activation and Filter channel mode extents much match
|
||||
bool implementable = true;
|
||||
// channel mode is major
|
||||
implementable &= args.problem_shape.stride_A[NumTensorDimensions-1] == 1;
|
||||
implementable &= args.problem_shape.stride_B[NumTensorDimensions-1] == 1;
|
||||
implementable &= problem_shape.stride_A[NumTensorDimensions-1] == 1;
|
||||
implementable &= problem_shape.stride_B[NumTensorDimensions-1] == 1;
|
||||
|
||||
constexpr int tma_alignment_bits = 128;
|
||||
// A extents.
|
||||
auto shape_A_orig = args.problem_shape.get_shape_A();
|
||||
auto shape_A_orig = problem_shape.get_shape_A();
|
||||
// B extents.
|
||||
auto shape_B_orig = args.problem_shape.get_shape_B();
|
||||
auto shape_B_orig = problem_shape.get_shape_B();
|
||||
constexpr int min_tma_aligned_elements_A = tma_alignment_bits / cutlass::sizeof_bits<ElementA>::value;
|
||||
implementable = implementable && cutlass::detail::check_alignment<min_tma_aligned_elements_A>(shape_A_orig, StrideA{});
|
||||
constexpr int min_tma_aligned_elements_B = tma_alignment_bits / cutlass::sizeof_bits<ElementB>::value;
|
||||
@@ -375,61 +372,6 @@ public:
|
||||
return false;
|
||||
}
|
||||
|
||||
if (is_im2col_A || is_im2col_B) {
|
||||
// Check valid corner values for TMA_LOAD_IM2COL, signed int ranging from [-corner_limit, corner_limit - 1]
|
||||
constexpr int32_t corner_limit = 1 << (16 / NumSpatialDimensions - 1);
|
||||
auto lower_corner_whd = detail::compute_lower_corner_whd(problem_shape);
|
||||
for (int i = 0; i < problem_shape.RankS; ++i) {
|
||||
implementable = implementable && lower_corner_whd[i] >= -corner_limit && lower_corner_whd[i] <= (corner_limit - 1);
|
||||
}
|
||||
auto upper_corner_whd = detail::compute_upper_corner_whd(problem_shape);
|
||||
for (int i = 0; i < problem_shape.RankS; ++i) {
|
||||
implementable = implementable && upper_corner_whd[i] >= -corner_limit && upper_corner_whd[i] <= (corner_limit - 1);
|
||||
}
|
||||
|
||||
if (!implementable) {
|
||||
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Padding values don't meet requirements for TMA LOAD IM2COL.\n");
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
// Wgrad kernels don't support non-packed output strides, non-packed tensor A stride (linearized)
|
||||
if constexpr (ConvOp == conv::Operator::kWgrad) {
|
||||
|
||||
const auto & input_shape = problem_shape.shape_A;
|
||||
const auto & input_stride = problem_shape.stride_A;
|
||||
|
||||
implementable &= input_stride[ProblemShape::RankT - 1] == 1;
|
||||
int input_shape_size = 1;
|
||||
for (int i = ProblemShape::RankT - 2; i >= 0; --i) {
|
||||
input_shape_size *= input_shape[i + 1];
|
||||
implementable &= input_stride[i] == input_shape_size;
|
||||
}
|
||||
|
||||
const auto & output_shape = problem_shape.shape_C;
|
||||
const auto & output_stride = problem_shape.stride_C;
|
||||
|
||||
implementable &= output_stride[ProblemShape::RankT - 1] == 1;
|
||||
int output_shape_size = 1;
|
||||
for (int i = ProblemShape::RankT - 2; i >= 0; --i) {
|
||||
output_shape_size *= output_shape[i + 1];
|
||||
implementable &= output_stride[i] == output_shape_size;
|
||||
}
|
||||
|
||||
if (!implementable) {
|
||||
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Wgrad kernels don't support non-packed output strides.\n");
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
// Conv kernels only support cross correlation mode currently.
|
||||
implementable &= problem_shape.mode == cutlass::conv::Mode::kCrossCorrelation;
|
||||
|
||||
if (!implementable) {
|
||||
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Conv kernels only support cross correlation mode currently.\n");
|
||||
return false;
|
||||
}
|
||||
|
||||
if (problem_shape.groups > 1) {
|
||||
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: This kernel does not support conv groups > 1.\n");
|
||||
return false;
|
||||
@@ -445,24 +387,53 @@ public:
|
||||
cute::prefetch_tma_descriptor(mainloop_params.tma_load_b.get_tma_descriptor());
|
||||
}
|
||||
|
||||
/// Set up the data needed by this collective for load and mma.
|
||||
/// Returns a tuple of tensors. The collective and the kernel layer have the contract
|
||||
/// Returned tuple must contain at least two elements, with the first two elements being:
|
||||
/// gA_mk - The tma tensor, A after a local tile so it has shape (BLK_M,BLK_K,m,k)
|
||||
/// gB_nk - The tma tensor, B after a local tile so it has shape (BLK_N,BLK_K,n,k)
|
||||
/// The rest of the tensors can be specified as needed by this collective.
|
||||
/// The dimensions of gA_mk and gA_nk do not contain L to maintain consistency with
|
||||
/// StrideA and StrideB set up for TMA
|
||||
template <class ProblemShapeMNKL>
|
||||
CUTLASS_DEVICE auto
|
||||
load_init(ProblemShapeMNKL const& problem_shape_MNKL, Params const& mainloop_params){
|
||||
//load_init(ProblemShapeMNKL const& problem_shape_MNKL, Params const& mainloop_params) const {
|
||||
using X = Underscore;
|
||||
// Separate out problem shape for convenience
|
||||
auto [M, N, K, L] = problem_shape_MNKL;
|
||||
|
||||
// TMA requires special handling of strides to deal with coord codomain mapping
|
||||
// Represent the full tensors -- get these from TMA
|
||||
Tensor mA_mk = mainloop_params.tma_load_a.get_tma_tensor(make_shape(M,K)); // (m,k)
|
||||
Tensor mB_nk = mainloop_params.tma_load_b.get_tma_tensor(make_shape(N,K)); // (n,k)
|
||||
|
||||
// Make tiled views, defer the slice
|
||||
Tensor gA_mk = local_tile(mA_mk, TileShape{}, make_coord(_,_,_), Step<_1, X,_1>{}); // (BLK_M,BLK_K,m,k)
|
||||
Tensor gB_nk = local_tile(mB_nk, TileShape{}, make_coord(_,_,_), Step< X,_1,_1>{}); // (BLK_N,BLK_K,n,k)
|
||||
|
||||
return cute::make_tuple(gA_mk, gB_nk);
|
||||
}
|
||||
|
||||
/// Perform a collective-scoped matrix multiply-accumulate
|
||||
/// Producer Perspective
|
||||
template <
|
||||
class TensorA, class TMA_LOAD_A,
|
||||
class TensorB, class TMA_LOAD_B,
|
||||
class KTileIterator
|
||||
class TensorA, class TensorB,
|
||||
class KTileIterator, class BlockCoord
|
||||
>
|
||||
CUTLASS_DEVICE void
|
||||
load(MainloopPipeline pipeline,
|
||||
PipelineState smem_pipe_producer_state,
|
||||
TensorA const& gA, TMA_LOAD_A& tma_load_a,
|
||||
TensorB const& gB, TMA_LOAD_B& tma_load_b,
|
||||
KTileIterator k_tile_iter, int k_tile_count,
|
||||
int thread_idx,
|
||||
uint32_t block_rank_in_cluster,
|
||||
TensorStorage& shared_tensors) {
|
||||
int lane_predicate = cute::elect_one_sync();
|
||||
load(
|
||||
Params const& mainloop_params,
|
||||
MainloopPipeline pipeline,
|
||||
PipelineState smem_pipe_producer_state,
|
||||
cute::tuple<TensorA, TensorB> const& load_inputs,
|
||||
BlockCoord const& blk_coord,
|
||||
KTileIterator k_tile_iter, int k_tile_count,
|
||||
int thread_idx,
|
||||
uint32_t block_rank_in_cluster,
|
||||
TensorStorage& shared_tensors) {
|
||||
|
||||
int lane_predicate = cute::elect_one_sync();
|
||||
if (lane_predicate) {
|
||||
Tensor sA = make_tensor(make_smem_ptr(shared_tensors.smem_A.data()), SmemLayoutA{}); // (BLK_M,BLK_K,PIPE)
|
||||
Tensor sB = make_tensor(make_smem_ptr(shared_tensors.smem_B.data()), SmemLayoutB{}); // (BLK_N,BLK_K,PIPE)
|
||||
@@ -470,11 +441,19 @@ public:
|
||||
//
|
||||
// Prepare the TMA loads for A and B
|
||||
//
|
||||
|
||||
constexpr uint32_t cluster_shape_x = get<0>(ClusterShape());
|
||||
|
||||
uint2 cluster_local_block_id = {block_rank_in_cluster % cluster_shape_x, block_rank_in_cluster / cluster_shape_x};
|
||||
auto block_tma_a = tma_load_a.get_slice(cluster_local_block_id.y);
|
||||
auto block_tma_b = tma_load_b.get_slice(cluster_local_block_id.x);
|
||||
auto block_tma_a = mainloop_params.tma_load_a.get_slice(cluster_local_block_id.y);
|
||||
auto block_tma_b = mainloop_params.tma_load_b.get_slice(cluster_local_block_id.x);
|
||||
|
||||
auto [gA_mk, gB_nk] = load_inputs;
|
||||
|
||||
// Partition the inputs based on the current block coordinates.
|
||||
auto [m_coord, n_coord, k_coord, l_coord] = blk_coord;
|
||||
|
||||
Tensor gA = gA_mk(_,_,m_coord,_); // (BLK_M,BLK_K,k)
|
||||
Tensor gB = gB_nk(_,_,n_coord,_); // (BLK_N,BLK_K,k)
|
||||
|
||||
// Applies the mapping from block_tma_a
|
||||
Tensor tAgA = block_tma_a.partition_S(gA); // (TMA,TMA_M,TMA_K,k)
|
||||
@@ -518,8 +497,9 @@ public:
|
||||
BarrierType* tma_barrier = pipeline.producer_get_barrier(smem_pipe_producer_state);
|
||||
|
||||
int write_stage = smem_pipe_producer_state.index();
|
||||
copy(tma_load_a.with(*tma_barrier, mcast_mask_a), tAgA(_,_,_,*k_tile_iter), tAsA(_,_,_,write_stage));
|
||||
copy(tma_load_b.with(*tma_barrier, mcast_mask_b), tBgB(_,_,_,*k_tile_iter), tBsB(_,_,_,write_stage));
|
||||
|
||||
copy(mainloop_params.tma_load_a.with(*tma_barrier, mcast_mask_a), tAgA(_,_,_,*k_tile_iter), tAsA(_,_,_,write_stage));
|
||||
copy(mainloop_params.tma_load_b.with(*tma_barrier, mcast_mask_b), tBgB(_,_,_,*k_tile_iter), tBsB(_,_,_,write_stage));
|
||||
++k_tile_iter;
|
||||
|
||||
// Advance smem_pipe_producer_state
|
||||
|
||||
@@ -43,6 +43,7 @@
|
||||
#include <initializer_list>
|
||||
#endif
|
||||
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::conv {
|
||||
@@ -54,15 +55,17 @@ namespace cutlass::conv {
|
||||
// Supports asymmetric padding, traversal strides, dilations, and all conv algorithm types.
|
||||
template <
|
||||
conv::Operator ConvOp_,
|
||||
int NumSpatialDimensions
|
||||
int NumSpatialDimensions_
|
||||
>
|
||||
struct ConvProblemShape {
|
||||
//
|
||||
// Alias types for members
|
||||
//
|
||||
static constexpr int RankS = NumSpatialDimensions;
|
||||
static constexpr int RankT = NumSpatialDimensions + 2;
|
||||
|
||||
static constexpr int RankS = NumSpatialDimensions_;
|
||||
static constexpr int RankT = NumSpatialDimensions_ + 2;
|
||||
static constexpr conv::Operator ConvOp = ConvOp_;
|
||||
static constexpr int NumSpatialDimensions = NumSpatialDimensions_;
|
||||
using SpatialExtent = cute::array<int, RankS>;
|
||||
using TensorExtent = cute::array<int, RankT>;
|
||||
using TensorStride = cute::array<int64_t, RankT>;
|
||||
@@ -352,71 +355,6 @@ struct ConvProblemShape {
|
||||
}
|
||||
}
|
||||
|
||||
// Get problem shape MNKL according to following table:
|
||||
// | | Fprop | Dgrad | Wgrad |
|
||||
// | ---- | --------- | -------- | -------- |
|
||||
// | Shape_M | (Q,P,Z,N) | (W/V,H/U,D/O,N) | (K) |
|
||||
// | Shape_N | (K) | (C) | (C,S,R,T) |
|
||||
// | Shape_K | (C,S,R,T) | (K,S,R,T) | (Q,P,Z,N) |
|
||||
// | Shape_L | _1 | (V,U,O) | _1 |
|
||||
CUTLASS_HOST_DEVICE
|
||||
constexpr auto
|
||||
get_transformed_problem_shape_MNKL() const {
|
||||
using cute::insert;
|
||||
using cute::make_shape;
|
||||
using cute::reverse;
|
||||
using cute::take;
|
||||
|
||||
if constexpr (ConvOp == conv::Operator::kWgrad) {
|
||||
auto M_xformed = shape_C[0];
|
||||
auto N_xformed = reverse(take<1, RankT>(shape_C));
|
||||
auto K_xformed = reverse(take<0, RankT - 1>(shape_A));
|
||||
auto L_xformed = cute::Int<1>{};
|
||||
|
||||
return make_shape(M_xformed, N_xformed, K_xformed, L_xformed);
|
||||
}
|
||||
else if constexpr (ConvOp == conv::Operator::kFprop){
|
||||
auto M_xformed = reverse(take<0, RankT - 1>(shape_C));
|
||||
auto N_xformed = shape_C[RankT - 1];
|
||||
auto K_xformed = reverse(take<1, RankT>(shape_B));
|
||||
auto L_xformed = cute::Int<1>{};
|
||||
|
||||
return make_shape(M_xformed, N_xformed, K_xformed, L_xformed);
|
||||
}
|
||||
else if constexpr (ConvOp == conv::Operator::kDgrad) {
|
||||
auto L_xformed = reverse(traversal_stride); // (V,U,O)
|
||||
auto M_xformed = ceil_div(reverse(take<0,RankT - 1>(shape_C)), L_xformed);
|
||||
auto N_xformed = shape_C[RankT - 1];
|
||||
// shape_B: [K,T,R,S,C], K_xformed: [K,S,R,T]
|
||||
auto K_xformed = insert<0>(
|
||||
(reverse(take<1,RankT - 1>(shape_B))),
|
||||
shape_B[0]);
|
||||
|
||||
return make_shape(M_xformed, N_xformed, K_xformed, L_xformed);
|
||||
}
|
||||
}
|
||||
|
||||
// Assuming im2col linearization
|
||||
// Get problem shape MNKL according to following table:
|
||||
// | | Fprop | Dgrad | Wgrad |
|
||||
// | ---- | --------- | -------- | -------- |
|
||||
// | Shape_M | (Q*P*Z*N) | ([W/V]*[H/U]*[D/O]*N) | (K) |
|
||||
// | Shape_N | (K) | (C) | (C,S,R,T) |
|
||||
// | Shape_K | (C,S,R,T) | (K,S,R,T) | (Q*P*Z*N) |
|
||||
// | Shape_L | _1 | (V*U*O) | _1 |
|
||||
CUTLASS_HOST_DEVICE
|
||||
constexpr auto
|
||||
get_linearized_problem_shape_MNKL() const {
|
||||
auto [M, N, K, L] = get_transformed_problem_shape_MNKL();
|
||||
|
||||
if constexpr (ConvOp == conv::Operator::kFprop || ConvOp == conv::Operator::kDgrad) {
|
||||
return cute::make_shape(cute::product(M), N, K, cute::product(L));
|
||||
}
|
||||
else if constexpr (ConvOp == conv::Operator::kWgrad) {
|
||||
return cute::make_shape(M, N, cute::product(K), L);
|
||||
}
|
||||
}
|
||||
|
||||
// Get A extents.
|
||||
// fprop: A extents array contains [N,D,H,W,C]. Turn that into ((W,H,D,N), (C))
|
||||
// dgrad: A extents array contains [N,Z,P,Q,K]. Turn that into ((Q,P,Z,N), (K))
|
||||
@@ -578,9 +516,7 @@ private:
|
||||
// calculate n,z,p,q,k.
|
||||
// a helper lambda to compute a single spatial extent of the nzpqk tensor
|
||||
auto nzpqk_extent = [](int act_ext, int filter_ext, int pad_total, int dilation, int tstride) {
|
||||
auto tmp = act_ext + pad_total - ((filter_ext -1) * dilation + 1);
|
||||
CUTLASS_ASSERT(tmp % tstride == 0);
|
||||
return 1 + tmp / tstride;
|
||||
return 1 + (act_ext + pad_total - ((filter_ext -1) * dilation + 1)) / tstride;
|
||||
};
|
||||
|
||||
shape_xformed_act[0] = shape_act[0]; // Activation N extent
|
||||
|
||||
@@ -0,0 +1,137 @@
|
||||
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2023 - 2024 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/conv/convnd_problem_shape.hpp"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::conv::detail {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// Helper function to get the problem shape
|
||||
template <typename T, class ProblemShape>
|
||||
auto get_problem_shape_MNKL_helper(ProblemShape const& problem_shape, cute::true_type) {
|
||||
return T::get_problem_shape_MNKL(problem_shape);
|
||||
}
|
||||
|
||||
template <typename T, class ProblemShape>
|
||||
ProblemShape get_problem_shape_MNKL_helper(ProblemShape const& problem_shape, cute::false_type) {
|
||||
return problem_shape;
|
||||
}
|
||||
|
||||
// Get problem shape MNKL according to following table:
|
||||
// | | Fprop | Dgrad | Wgrad |
|
||||
// | ---- | --------- | -------- | -------- |
|
||||
// | Shape_M | (Q,P,Z,N) | (W/V,H/U,D/O,N) | (K) |
|
||||
// | Shape_N | (K) | (C) | (C,S,R,T) |
|
||||
// | Shape_K | (C,S,R,T) | (K,S,R,T) | (Q,P,Z,N) |
|
||||
// | Shape_L | _1 | (V,U,O) | _1 |
|
||||
|
||||
template <class ProblemShape>
|
||||
CUTLASS_HOST_DEVICE
|
||||
constexpr auto
|
||||
get_transformed_problem_shape_MNKL(ProblemShape const& problem_shape) {
|
||||
return problem_shape;
|
||||
}
|
||||
|
||||
|
||||
template <conv::Operator ConvOp, int SpatialDim>
|
||||
CUTLASS_HOST_DEVICE
|
||||
constexpr auto
|
||||
get_transformed_problem_shape_MNKL(ConvProblemShape<ConvOp, SpatialDim> const& problem_shape) {
|
||||
using cute::insert;
|
||||
using cute::make_shape;
|
||||
using cute::reverse;
|
||||
using cute::take;
|
||||
|
||||
constexpr int RankT = SpatialDim + 2;
|
||||
|
||||
if constexpr (ConvOp == conv::Operator::kWgrad) {
|
||||
auto M_xformed = problem_shape.shape_C[0];
|
||||
auto N_xformed = reverse(take<1, RankT>(problem_shape.shape_C));
|
||||
auto K_xformed = reverse(take<0, RankT - 1>(problem_shape.shape_A));
|
||||
auto L_xformed = cute::Int<1>{};
|
||||
|
||||
return make_shape(M_xformed, N_xformed, K_xformed, L_xformed);
|
||||
}
|
||||
else if constexpr (ConvOp == conv::Operator::kFprop){
|
||||
auto M_xformed = reverse(take<0, RankT - 1>(problem_shape.shape_C));
|
||||
auto N_xformed = problem_shape.shape_C[RankT - 1];
|
||||
auto K_xformed = reverse(take<1, RankT>(problem_shape.shape_B));
|
||||
auto L_xformed = cute::Int<1>{};
|
||||
|
||||
return make_shape(M_xformed, N_xformed, K_xformed, L_xformed);
|
||||
}
|
||||
else if constexpr (ConvOp == conv::Operator::kDgrad) {
|
||||
auto L_xformed = reverse(problem_shape.traversal_stride); // (V,U,O)
|
||||
auto M_xformed = ceil_div(reverse(take<0,RankT - 1>(problem_shape.shape_C)), L_xformed);
|
||||
auto N_xformed = problem_shape.shape_C[RankT - 1];
|
||||
// shape_B: [K,T,R,S,C], K_xformed: [K,S,R,T]
|
||||
auto K_xformed = insert<0>(
|
||||
(reverse(take<1,RankT - 1>(problem_shape.shape_B))),
|
||||
problem_shape.shape_B[0]);
|
||||
|
||||
return make_shape(M_xformed, N_xformed, K_xformed, L_xformed);
|
||||
}
|
||||
}
|
||||
|
||||
// Assuming im2col linearization
|
||||
// Get problem shape MNKL according to following table:
|
||||
// | | Fprop | Dgrad | Wgrad |
|
||||
// | ---- | --------- | -------- | -------- |
|
||||
// | Shape_M | (Q*P*Z*N) | ([W/V]*[H/U]*[D/O]*N) | (K) |
|
||||
// | Shape_N | (K) | (C) | (C,S,R,T) |
|
||||
// | Shape_K | (C,S,R,T) | (K,S,R,T) | (Q*P*Z*N) |
|
||||
// | Shape_L | _1 | (V*U*O) | _1 |
|
||||
template <conv::Operator ConvOp, int SpatialDim>
|
||||
CUTLASS_HOST_DEVICE
|
||||
constexpr auto
|
||||
get_linearized_problem_shape_MNKL(ConvProblemShape<ConvOp, SpatialDim> const& problem_shape) {
|
||||
|
||||
auto [M, N, K, L] = get_transformed_problem_shape_MNKL(problem_shape);
|
||||
|
||||
if constexpr (ConvOp == conv::Operator::kFprop || ConvOp == conv::Operator::kDgrad) {
|
||||
return cute::make_shape(cute::product(M), N, K, cute::product(L));
|
||||
}
|
||||
else if constexpr (ConvOp == conv::Operator::kWgrad) {
|
||||
return cute::make_shape(M, N, cute::product(K), L);
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::conv::detail
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -61,7 +61,7 @@ template <class ConvKernel_>
|
||||
class ConvUniversalAdapter
|
||||
{
|
||||
public:
|
||||
using ConvKernel = ConvKernel_;
|
||||
using ConvKernel = GetUnderlyingKernel_t<ConvKernel_>;
|
||||
using TileShape = typename ConvKernel::TileShape;
|
||||
using ElementA = typename ConvKernel::ElementA;
|
||||
using ElementB = typename ConvKernel::ElementB;
|
||||
@@ -76,7 +76,7 @@ public:
|
||||
|
||||
// Tease out meta-information about the conv algorithm
|
||||
static constexpr conv::Operator kConvolutionalOperator = DispatchPolicy::ConvOp;
|
||||
static constexpr int NumSpatialDimensions = ConvKernel::NumSpatialDimensions;
|
||||
static constexpr int NumSpatialDimensions = CollectiveMainloop::NumSpatialDimensions;
|
||||
|
||||
// If our TiledMMA's instruction thread layout size is larger than 1, we know its a tensorop!
|
||||
using OperatorClass = cute::conditional_t<
|
||||
@@ -121,13 +121,13 @@ public:
|
||||
static int constexpr kStages = CollectiveMainloop::DispatchPolicy::Stages;
|
||||
|
||||
// Inspect TiledCopy for A and B to compute the alignment size
|
||||
static int constexpr kAlignmentA = detail::get_alignment_count_from_gmem_tiled_copy<
|
||||
static int constexpr kAlignmentA = cutlass::detail::get_alignment_count_from_gmem_tiled_copy<
|
||||
typename CollectiveMainloop::GmemTiledCopyA, ElementA>();
|
||||
static int constexpr kAlignmentB = detail::get_alignment_count_from_gmem_tiled_copy<
|
||||
static int constexpr kAlignmentB = cutlass::detail::get_alignment_count_from_gmem_tiled_copy<
|
||||
typename CollectiveMainloop::GmemTiledCopyB, ElementB>();
|
||||
static int constexpr kAlignmentC = detail::get_alignment_count_from_gmem_tiled_copy<
|
||||
static int constexpr kAlignmentC = cutlass::detail::get_alignment_count_from_gmem_tiled_copy<
|
||||
typename CollectiveEpilogue::GmemTiledCopyC, ElementC>();
|
||||
static int constexpr kAlignmentD = detail::get_alignment_count_from_gmem_tiled_copy<
|
||||
static int constexpr kAlignmentD = cutlass::detail::get_alignment_count_from_gmem_tiled_copy<
|
||||
typename CollectiveEpilogue::GmemTiledCopyD, ElementD>();
|
||||
|
||||
using EpilogueOutputOp = typename CollectiveEpilogue::ThreadEpilogueOp;
|
||||
@@ -297,8 +297,9 @@ public:
|
||||
Status launch_result;
|
||||
// Use extended launch API only for mainloops that use it
|
||||
if constexpr (ConvKernel::ArchTag::kMinComputeCapability >= 90) {
|
||||
constexpr bool is_static_1x1x1 = cute::is_static_v<typename ConvKernel::DispatchPolicy::ClusterShape> and
|
||||
cute::size(typename ConvKernel::DispatchPolicy::ClusterShape{}) == 1;
|
||||
[[maybe_unused]] constexpr bool is_static_1x1x1 =
|
||||
cute::is_static_v<typename ConvKernel::DispatchPolicy::ClusterShape> and
|
||||
cute::size(typename ConvKernel::DispatchPolicy::ClusterShape{}) == 1;
|
||||
dim3 cluster(cute::size<0>(typename ConvKernel::DispatchPolicy::ClusterShape{}),
|
||||
cute::size<1>(typename ConvKernel::DispatchPolicy::ClusterShape{}),
|
||||
cute::size<2>(typename ConvKernel::DispatchPolicy::ClusterShape{}));
|
||||
|
||||
@@ -211,6 +211,7 @@ public:
|
||||
dim3 grid = ReorderKernel::get_grid_shape(params_);
|
||||
dim3 block = ReorderKernel::get_block_shape();
|
||||
|
||||
cutlass::arch::synclog_setup();
|
||||
cutlass::Kernel<ReorderKernel><<<grid, block, 0, stream>>>(params_);
|
||||
}
|
||||
|
||||
@@ -229,6 +230,7 @@ public:
|
||||
if (status != cudaSuccess)
|
||||
return Status::kErrorInternal;
|
||||
|
||||
cutlass::arch::synclog_setup();
|
||||
cutlass::Kernel<UnderlyingKernel><<<grid, block, smem_size, stream>>>(params_);
|
||||
|
||||
cudaError_t result = cudaGetLastError();
|
||||
|
||||
@@ -53,7 +53,7 @@ template<typename ImplicitGemmKernel_>
|
||||
class ImplicitGemmConvolution {
|
||||
public:
|
||||
|
||||
using UnderlyingKernel = ImplicitGemmKernel_;
|
||||
using UnderlyingKernel = GetUnderlyingKernel_t<ImplicitGemmKernel_>;
|
||||
|
||||
using ElementA = typename UnderlyingKernel::ElementA;
|
||||
using LayoutA = typename UnderlyingKernel::LayoutA;
|
||||
@@ -103,7 +103,6 @@ public:
|
||||
|
||||
/// 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) {
|
||||
@@ -164,9 +163,8 @@ public:
|
||||
// check for unsupported problem sizes for strided dgrad / deconv implementation
|
||||
if ((kConvolutionalOperator == conv::Operator::kDgrad || kConvolutionalOperator == conv::Operator::kDeconv) &&
|
||||
kStrideSupport == conv::StrideSupport::kStrided) {
|
||||
|
||||
// split-k (serial or parallel) is not supported for strided dgrad / deconv
|
||||
if(args.problem_size.split_k_slices > 1) {
|
||||
if(args.problem_size.split_k_slices > 1 && (args.problem_size.stride().at(args.problem_size.stride().max_dim_index()) > 1)) {
|
||||
return Status::kErrorNotSupported;
|
||||
}
|
||||
|
||||
@@ -291,7 +289,7 @@ public:
|
||||
}
|
||||
|
||||
/// Runs the kernel using initialized state.
|
||||
Status run(cudaStream_t stream = nullptr, CudaHostAdapter *cuda_adapter = nullptr) {
|
||||
Status run(cudaStream_t stream = nullptr, CudaHostAdapter *cuda_adapter = nullptr, int32_t kernel_index = 0) {
|
||||
|
||||
|
||||
ThreadblockSwizzle threadblock_swizzle;
|
||||
@@ -311,7 +309,7 @@ public:
|
||||
|
||||
void* kernel_params[] = {¶ms_};
|
||||
launch_result = cuda_adapter->launch(
|
||||
grid, dim3(1,1,1), block, smem_size, stream, kernel_params, 0
|
||||
grid, dim3(1,1,1), block, smem_size, stream, kernel_params, kernel_index
|
||||
);
|
||||
}
|
||||
else {
|
||||
@@ -319,6 +317,7 @@ public:
|
||||
}
|
||||
}
|
||||
else {
|
||||
cutlass::arch::synclog_setup();
|
||||
cutlass::Kernel<UnderlyingKernel><<<grid, block, smem_size, stream>>>(params_);
|
||||
}
|
||||
|
||||
@@ -333,20 +332,20 @@ public:
|
||||
}
|
||||
|
||||
/// Runs the kernel using initialized state.
|
||||
Status operator()(cudaStream_t stream = nullptr, CudaHostAdapter *cuda_adapter = nullptr) {
|
||||
return run(stream, cuda_adapter);
|
||||
Status operator()(cudaStream_t stream = nullptr, CudaHostAdapter *cuda_adapter = nullptr, int32_t kernel_index = 0) {
|
||||
return run(stream, cuda_adapter, kernel_index);
|
||||
}
|
||||
|
||||
/// Runs the kernel using initialized state.
|
||||
Status operator()(
|
||||
Arguments const &args,
|
||||
void *workspace = nullptr,
|
||||
cudaStream_t stream = nullptr, CudaHostAdapter *cuda_adapter = nullptr) {
|
||||
cudaStream_t stream = nullptr, CudaHostAdapter *cuda_adapter = nullptr, int32_t kernel_index = 0) {
|
||||
|
||||
Status status = initialize(args, workspace, stream, cuda_adapter);
|
||||
|
||||
if (status == Status::kSuccess) {
|
||||
status = run(stream, cuda_adapter);
|
||||
status = run(stream, cuda_adapter, kernel_index);
|
||||
}
|
||||
|
||||
return status;
|
||||
|
||||
@@ -231,6 +231,7 @@ public:
|
||||
|
||||
int smem_size = int(sizeof(typename ImplicitGemmFusionKernel::SharedStorage));
|
||||
|
||||
cutlass::arch::synclog_setup();
|
||||
cutlass::Kernel<ImplicitGemmFusionKernel><<<grid, block, smem_size, stream>>>(params_);
|
||||
|
||||
cudaError_t result = cudaGetLastError();
|
||||
|
||||
@@ -37,6 +37,8 @@
|
||||
#include "cute/layout.hpp"
|
||||
#include "cute/numeric/integral_constant.hpp"
|
||||
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
@@ -48,7 +50,7 @@ namespace cutlass::conv {
|
||||
//
|
||||
// Policies for categorical dispatch of mainloop against kernel grid schedules
|
||||
//
|
||||
struct KernelImplicitTmaWarpSpecializedSm90 { };
|
||||
struct KernelImplicitTmaWarpSpecializedSm90 : cutlass::gemm::KernelTmaWarpSpecialized { };
|
||||
struct KernelImplicitTmaWarpSpecializedSm90Cooperative { };
|
||||
struct KernelImplicitTmaWarpSpecializedSm90Pingpong { };
|
||||
|
||||
@@ -84,3 +86,5 @@ struct MainloopSm90TmaGmmaWarpSpecializedImplicitGemm {
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::conv
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -30,6 +30,7 @@
|
||||
**************************************************************************************************/
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/conv/convnd_problem_shape.hpp"
|
||||
#include "cutlass/detail/dependent_false.hpp"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -43,6 +44,7 @@ namespace cutlass::conv::kernel {
|
||||
* a composition of a collective mainloop and a collective epilogue.
|
||||
**/
|
||||
template <
|
||||
class ProblemShape_,
|
||||
class CollectiveMainloop_,
|
||||
class CollectiveEpilogue_,
|
||||
class TileSchedulerTag_ = void,
|
||||
|
||||
@@ -37,9 +37,12 @@
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/arch/cluster_sm90.hpp"
|
||||
|
||||
#include "cutlass/conv/detail.hpp"
|
||||
#include "cutlass/conv/convolution.h"
|
||||
#include "cutlass/conv/dispatch_policy.hpp"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
#include "cutlass/pipeline/sm90_pipeline.hpp"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/kernel/tile_scheduler.hpp"
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
@@ -49,365 +52,25 @@ namespace cutlass::conv::kernel {
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
class ProblemShape_,
|
||||
class CollectiveMainloop_,
|
||||
class CollectiveEpilogue_,
|
||||
class TileSchedulerTag_
|
||||
class TileScheduler_
|
||||
>
|
||||
class ConvUniversal<
|
||||
ProblemShape_,
|
||||
CollectiveMainloop_,
|
||||
CollectiveEpilogue_,
|
||||
TileSchedulerTag_,
|
||||
cute::enable_if_t<cute::is_base_of_v<cutlass::conv::KernelImplicitTmaWarpSpecializedSm90,
|
||||
typename CollectiveMainloop_::DispatchPolicy::Schedule>>>
|
||||
{
|
||||
public:
|
||||
//
|
||||
// Type Aliases
|
||||
//
|
||||
|
||||
// Mainloop derived types
|
||||
using CollectiveMainloop = CollectiveMainloop_;
|
||||
using TileShape = typename CollectiveMainloop::TileShape;
|
||||
using TiledMma = typename CollectiveMainloop::TiledMma;
|
||||
using ArchTag = typename CollectiveMainloop::ArchTag;
|
||||
using ElementA = typename CollectiveMainloop::ElementA;
|
||||
using StrideA = typename CollectiveMainloop::StrideA;
|
||||
using ElementB = typename CollectiveMainloop::ElementB;
|
||||
using StrideB = typename CollectiveMainloop::StrideB;
|
||||
using DispatchPolicy = typename CollectiveMainloop::DispatchPolicy;
|
||||
using ElementAccumulator = typename CollectiveMainloop::ElementAccumulator;
|
||||
using ClusterShape = typename DispatchPolicy::ClusterShape;
|
||||
using MainloopArguments = typename CollectiveMainloop::Arguments;
|
||||
using MainloopParams = typename CollectiveMainloop::Params;
|
||||
static constexpr int NumSpatialDimensions = CollectiveMainloop::NumSpatialDimensions;
|
||||
static_assert(ArchTag::kMinComputeCapability >= 90);
|
||||
// Epilogue derived types
|
||||
using CollectiveEpilogue = CollectiveEpilogue_;
|
||||
using ElementC = typename CollectiveEpilogue::ElementC;
|
||||
using StrideC = typename CollectiveEpilogue::StrideC;
|
||||
using ElementD = typename CollectiveEpilogue::ElementD;
|
||||
using StrideD = typename CollectiveEpilogue::StrideD;
|
||||
using EpilogueArguments = typename CollectiveEpilogue::Arguments;
|
||||
using EpilogueParams = typename CollectiveEpilogue::Params;
|
||||
|
||||
using TileSchedulerTag = TileSchedulerTag_;
|
||||
static_assert(cute::is_void_v<TileSchedulerTag>,
|
||||
"TMA warp-specialized kernel does not support specializing the tile scheduler.");
|
||||
using TileScheduler = typename cutlass::gemm::kernel::detail::TileSchedulerSelector<
|
||||
TileSchedulerTag, ArchTag, TileShape, ClusterShape>::Scheduler;
|
||||
using TileSchedulerArguments = typename TileScheduler::Arguments;
|
||||
|
||||
// Kernel level shared memory storage
|
||||
struct SharedStorage {
|
||||
union TensorStorage {
|
||||
using MainloopTensorStorage = typename CollectiveMainloop::TensorStorage;
|
||||
using EpilogueTensorStorage = typename CollectiveEpilogue::TensorStorage;
|
||||
|
||||
MainloopTensorStorage mainloop;
|
||||
EpilogueTensorStorage epilogue;
|
||||
} tensors;
|
||||
|
||||
struct PipelineStorage : cute::aligned_struct<16> {
|
||||
using MainloopPipelineStorage = typename CollectiveMainloop::PipelineStorage;
|
||||
using EpiLoadPipelineStorage = typename CollectiveEpilogue::PipelineStorage;
|
||||
|
||||
alignas(16) MainloopPipelineStorage mainloop;
|
||||
alignas(16) EpiLoadPipelineStorage epi_load;
|
||||
} pipelines;
|
||||
};
|
||||
|
||||
static constexpr int SharedStorageSize = sizeof(SharedStorage);
|
||||
static constexpr uint32_t NumLoadWarpGroups = 1;
|
||||
static constexpr uint32_t NumMmaWarpGroups = 1;
|
||||
static constexpr uint32_t MaxThreadsPerBlock = CUTE_STATIC_V(size(TiledMma{})) + (NumLoadWarpGroups * NumThreadsPerWarpGroup);
|
||||
static constexpr uint32_t MinBlocksPerMultiprocessor = 1;
|
||||
|
||||
// Host facing host arguments
|
||||
struct Arguments {
|
||||
MainloopArguments mainloop{};
|
||||
EpilogueArguments epilogue{};
|
||||
KernelHardwareInfo hw_info{};
|
||||
TileSchedulerArguments scheduler{};
|
||||
};
|
||||
|
||||
// Kernel device entry point API
|
||||
struct Params {
|
||||
MainloopParams mainloop;
|
||||
EpilogueParams epilogue;
|
||||
};
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
// Map user facing arguments to device facing params
|
||||
static Params
|
||||
to_underlying_arguments(Arguments const& args, void* workspace) {
|
||||
(void) workspace;
|
||||
auto mainloop_params = CollectiveMainloop::to_underlying_arguments(args.mainloop, workspace);
|
||||
auto problem_shape_MNKL = args.mainloop.problem_shape.get_transformed_problem_shape_MNKL();
|
||||
|
||||
return {
|
||||
mainloop_params,
|
||||
CollectiveEpilogue::to_underlying_arguments(problem_shape_MNKL, args.epilogue, workspace)
|
||||
};
|
||||
}
|
||||
|
||||
// Given arguemnts, returns true if the kernel can successfully compute upon them. False otherwise.
|
||||
static bool
|
||||
can_implement(Arguments const& args) {
|
||||
bool implementable = true;
|
||||
implementable &= CollectiveMainloop::can_implement(args.mainloop.problem_shape, args.mainloop);
|
||||
implementable &= CollectiveEpilogue::can_implement(args.mainloop.problem_shape.get_transformed_problem_shape_MNKL(), args.epilogue);
|
||||
return implementable;
|
||||
}
|
||||
|
||||
static size_t
|
||||
get_workspace_size(Arguments const& args) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
static cutlass::Status
|
||||
initialize_workspace(Arguments const& args, void* workspace = nullptr, cudaStream_t stream = nullptr,
|
||||
CudaHostAdapter* cuda_adapter = nullptr) {
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
// Computes the kernel launch grid shape based on runtime parameters
|
||||
static dim3
|
||||
get_grid_shape(Params const& params) {
|
||||
return cutlass::gemm::kernel::detail::PersistentTileSchedulerSm90::get_tiled_cta_shape_mnl(
|
||||
params.mainloop.problem_shape, TileShape{}, ClusterShape{});
|
||||
}
|
||||
|
||||
static dim3
|
||||
get_block_shape() {
|
||||
return dim3(MaxThreadsPerBlock, 1, 1);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void
|
||||
operator()(Params const& params, char* smem_buf) {
|
||||
using namespace cute;
|
||||
using X = Underscore;
|
||||
|
||||
// Any Tensor Op MMA Atom in the WGMMA ISA is arch conditional to sm90a.
|
||||
#if ! defined(__CUDA_ARCH_FEAT_SM90_ALL)
|
||||
if constexpr(size<0>(typename TiledMma::AtomShape_MNK{}) == 64) {
|
||||
printf("ERROR : Arch conditional MMA instruction used without targeting sm90a compute capability. Aborting.\n");
|
||||
return;
|
||||
}
|
||||
#endif
|
||||
|
||||
enum class WarpGroupRole {
|
||||
Producer = 0,
|
||||
Consumer = 1,
|
||||
};
|
||||
|
||||
enum class ProducerWarpRole {
|
||||
MainloopEpilogue = 0,
|
||||
Warp1 = 1,
|
||||
Warp2 = 2,
|
||||
Warp3 = 3
|
||||
};
|
||||
|
||||
// Kernel level shared memory storage
|
||||
SharedStorage& shared_storage = *reinterpret_cast<SharedStorage*>(smem_buf);
|
||||
|
||||
int thread_idx = int(threadIdx.x);
|
||||
int lane_idx = canonical_lane_idx();
|
||||
int warp_idx = canonical_warp_idx_sync();
|
||||
int warp_idx_in_warp_group = warp_idx % NumWarpsPerWarpGroup;
|
||||
int warp_group_thread_idx = thread_idx % NumThreadsPerWarpGroup;
|
||||
auto warp_group_role = WarpGroupRole(canonical_warp_group_idx());
|
||||
auto producer_warp_role = ProducerWarpRole(warp_idx_in_warp_group);
|
||||
int lane_predicate = cute::elect_one_sync();
|
||||
uint32_t block_rank_in_cluster = cute::block_rank_in_cluster();
|
||||
|
||||
// Issue Tma Descriptor Prefetch from a single thread
|
||||
if ((warp_idx == 0) && lane_predicate) {
|
||||
CollectiveMainloop::prefetch_tma_descriptors(params.mainloop);
|
||||
CollectiveEpilogue::prefetch_tma_descriptors(params.epilogue);
|
||||
}
|
||||
|
||||
// Mainloop Load pipeline
|
||||
using MainloopPipeline = typename CollectiveMainloop::MainloopPipeline;
|
||||
typename MainloopPipeline::Params mainloop_pipeline_params;
|
||||
if (warp_group_role == WarpGroupRole::Producer && producer_warp_role == ProducerWarpRole::MainloopEpilogue) {
|
||||
mainloop_pipeline_params.role = MainloopPipeline::ThreadCategory::Producer;
|
||||
}
|
||||
if (warp_group_role == WarpGroupRole::Consumer) {
|
||||
mainloop_pipeline_params.role = MainloopPipeline::ThreadCategory::Consumer;
|
||||
}
|
||||
mainloop_pipeline_params.is_leader = warp_group_thread_idx == 0;
|
||||
mainloop_pipeline_params.num_consumers = NumThreadsPerWarpGroup;
|
||||
mainloop_pipeline_params.transaction_bytes = params.mainloop.tma_transaction_bytes;
|
||||
MainloopPipeline mainloop_pipeline(shared_storage.pipelines.mainloop, mainloop_pipeline_params, ClusterShape{});
|
||||
|
||||
// Epilogue Load pipeline
|
||||
using EpiLoadPipeline = typename CollectiveEpilogue::LoadPipeline;
|
||||
typename EpiLoadPipeline::Params epi_load_pipeline_params;
|
||||
if (warp_group_role == WarpGroupRole::Producer && producer_warp_role == ProducerWarpRole::MainloopEpilogue) {
|
||||
epi_load_pipeline_params.role = EpiLoadPipeline::ThreadCategory::Producer;
|
||||
}
|
||||
if (warp_group_role == WarpGroupRole::Consumer) {
|
||||
epi_load_pipeline_params.role = EpiLoadPipeline::ThreadCategory::Consumer;
|
||||
}
|
||||
epi_load_pipeline_params.dst_blockid = cute::block_rank_in_cluster();
|
||||
epi_load_pipeline_params.producer_arv_count = NumThreadsPerWarp;
|
||||
epi_load_pipeline_params.consumer_arv_count = NumThreadsPerWarpGroup;
|
||||
if constexpr (CollectiveEpilogue::RequiresTransactionBytes) {
|
||||
epi_load_pipeline_params.transaction_bytes = params.epilogue.tma_transaction_bytes;
|
||||
}
|
||||
EpiLoadPipeline epi_load_pipeline(shared_storage.pipelines.epi_load, epi_load_pipeline_params);
|
||||
|
||||
// Epilogue Store pipeline
|
||||
using EpiStorePipeline = typename CollectiveEpilogue::StorePipeline;
|
||||
typename EpiStorePipeline::Params epi_store_pipeline_params;
|
||||
epi_store_pipeline_params.always_wait = true;
|
||||
EpiStorePipeline epi_store_pipeline(epi_store_pipeline_params);
|
||||
|
||||
// Initialize starting pipeline states for the collectives
|
||||
// Epilogue store pipe is producer-only (consumer is TMA unit, waits via scoreboarding)
|
||||
typename CollectiveMainloop::PipelineState mainloop_pipe_consumer_state;
|
||||
typename CollectiveEpilogue::LoadPipelineState epi_load_pipe_consumer_state;
|
||||
|
||||
// For the DMA Load (producer) we start with an opposite phase
|
||||
// i.e., we skip all waits since we know that the buffer is indeed empty
|
||||
PipelineState mainloop_pipe_producer_state = cutlass::make_producer_start_state<MainloopPipeline>();
|
||||
PipelineState epi_load_pipe_producer_state = cutlass::make_producer_start_state<EpiLoadPipeline>();
|
||||
PipelineState epi_store_pipe_producer_state = cutlass::make_producer_start_state<EpiStorePipeline>();
|
||||
|
||||
auto cluster_wait_fn = [&] () {
|
||||
// We need this to guarantee that the Pipeline init is visible
|
||||
// To all producers and consumer thread blocks in the Cluster
|
||||
if constexpr (size(ClusterShape{}) > 1) {
|
||||
cute::cluster_arrive_relaxed();
|
||||
return [] () { cute::cluster_wait(); };
|
||||
}
|
||||
else {
|
||||
__syncthreads();
|
||||
return [] () {}; // do nothing
|
||||
}
|
||||
} ();
|
||||
|
||||
// Separate out problem shape for convenience
|
||||
auto problem_shape_MNKL = append<4>(params.mainloop.problem_shape, _1{});
|
||||
auto [M, N, K, L] = problem_shape_MNKL;
|
||||
|
||||
// TMA requires special handling of strides to deal with coord codomain mapping
|
||||
// Represent the full tensors -- get these from TMA
|
||||
Tensor mA_mk = params.mainloop.tma_load_a.get_tma_tensor(make_shape(M, K));
|
||||
Tensor mB_nk = params.mainloop.tma_load_b.get_tma_tensor(make_shape(N, K));
|
||||
|
||||
// Get the appropriate blocks for this thread block -- potential for thread block locality
|
||||
auto cta_tile_shape = TileShape{}; // (BLK_M,BLK_N,BLK_K)
|
||||
TiledMma tiled_mma;
|
||||
|
||||
// Make tiled views, defer the slice
|
||||
Tensor gA_mk = local_tile(mA_mk, cta_tile_shape, make_coord(_,_,_), Step<_1, X,_1>{}); // (BLK_M,BLK_K,m,k)
|
||||
Tensor gB_nk = local_tile(mB_nk, cta_tile_shape, make_coord(_,_,_), Step< X,_1,_1>{}); // (BLK_N,BLK_K,n,k)
|
||||
|
||||
// Compute m_coord, n_coord, and l_coord with their post-tiled shapes
|
||||
auto m_coord = idx2crd(int(blockIdx.x), shape<2>(gA_mk));
|
||||
auto n_coord = idx2crd(int(blockIdx.y), shape<2>(gB_nk), compact_col_major(shape<2>(gB_nk)));
|
||||
|
||||
// The output shape M is linearized so the output coord M here should also be linearized.
|
||||
auto output_tile_coord = make_coord(int(blockIdx.x), n_coord, _, Int<0>{});
|
||||
|
||||
// Slice with m_coord and n_coord
|
||||
Tensor gA = gA_mk(_,_,m_coord,_); // (BLK_M,BLK_K,k)
|
||||
Tensor gB = gB_nk(_,_,n_coord,_); // (BLK_N,BLK_K,k)
|
||||
|
||||
// Get pipeline iterators and increments from tensor shapes
|
||||
auto k_tile_iter = cute::make_coord_iterator(shape<2>(gA));
|
||||
auto k_tile_count = size<2>(gA);
|
||||
|
||||
// In a warp specialized kernel, collectives expose data movement and compute operations separately
|
||||
CollectiveMainloop collective_mainloop;
|
||||
CollectiveEpilogue collective_epilogue{params.epilogue, shared_storage.tensors.epilogue};
|
||||
|
||||
// Wait for all thread blocks in Cluster
|
||||
cluster_wait_fn();
|
||||
|
||||
if (warp_group_role == WarpGroupRole::Producer) {
|
||||
if (producer_warp_role == ProducerWarpRole::MainloopEpilogue) {
|
||||
collective_mainloop.load(
|
||||
mainloop_pipeline,
|
||||
mainloop_pipe_producer_state,
|
||||
gA, params.mainloop.tma_load_a,
|
||||
gB, params.mainloop.tma_load_b,
|
||||
k_tile_iter, k_tile_count,
|
||||
lane_idx,
|
||||
block_rank_in_cluster,
|
||||
shared_storage.tensors.mainloop
|
||||
);
|
||||
// Update starting mainloop pipeline state for the pipeline drain
|
||||
mainloop_pipe_producer_state.advance(k_tile_count);
|
||||
// Make sure mainloop consumer has been waited upon before issuing epilogue load
|
||||
collective_mainloop.load_tail(mainloop_pipeline, mainloop_pipe_producer_state);
|
||||
|
||||
if (collective_epilogue.is_producer_load_needed()) {
|
||||
epi_load_pipe_producer_state = collective_epilogue.load(
|
||||
epi_load_pipeline,
|
||||
epi_load_pipe_producer_state,
|
||||
problem_shape_MNKL,
|
||||
cta_tile_shape,
|
||||
output_tile_coord,
|
||||
tiled_mma,
|
||||
lane_idx,
|
||||
shared_storage.tensors.epilogue
|
||||
);
|
||||
collective_epilogue.load_tail(epi_load_pipeline, epi_load_pipe_producer_state);
|
||||
}
|
||||
}
|
||||
}
|
||||
else if (warp_group_role == WarpGroupRole::Consumer) {
|
||||
Tensor accumulators = partition_fragment_C(tiled_mma, take<0,2>(cta_tile_shape)); // (MMA,MMA_M,MMA_N)
|
||||
|
||||
collective_mainloop.mma(
|
||||
mainloop_pipeline,
|
||||
mainloop_pipe_consumer_state,
|
||||
accumulators,
|
||||
k_tile_count,
|
||||
thread_idx,
|
||||
shared_storage.tensors.mainloop,
|
||||
params.mainloop
|
||||
);
|
||||
|
||||
// Make sure the math instructions are done and free buffers before entering the epilogue
|
||||
collective_mainloop.mma_tail(
|
||||
mainloop_pipeline,
|
||||
mainloop_pipe_consumer_state,
|
||||
k_tile_count
|
||||
);
|
||||
|
||||
// Epilogue and write to gD
|
||||
auto [epi_load_pipe_consumer_state_next, epi_store_pipe_producer_state_next] =
|
||||
collective_epilogue.store(
|
||||
epi_load_pipeline,
|
||||
epi_load_pipe_consumer_state,
|
||||
epi_store_pipeline,
|
||||
epi_store_pipe_producer_state,
|
||||
problem_shape_MNKL,
|
||||
cta_tile_shape,
|
||||
output_tile_coord,
|
||||
accumulators,
|
||||
tiled_mma,
|
||||
warp_group_thread_idx,
|
||||
shared_storage.tensors.epilogue
|
||||
);
|
||||
|
||||
collective_epilogue.store_tail(
|
||||
epi_load_pipeline,
|
||||
epi_load_pipe_consumer_state_next,
|
||||
epi_store_pipeline,
|
||||
epi_store_pipe_producer_state_next
|
||||
);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
TileScheduler_,
|
||||
cute::enable_if_t<cute::is_base_of_v<KernelImplicitTmaWarpSpecializedSm90, typename CollectiveMainloop_::DispatchPolicy::Schedule>>
|
||||
> : public cutlass::gemm::kernel::GemmUniversal<
|
||||
ProblemShape_,
|
||||
CollectiveMainloop_,
|
||||
CollectiveEpilogue_,
|
||||
TileScheduler_
|
||||
>
|
||||
{};
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::conv::kernel
|
||||
|
||||
|
||||
Reference in New Issue
Block a user