CUTLASS 3.5.1 (#1623)

* CUTLASS 3.5.1

* updates, optimizations, fixes
This commit is contained in:
Vijay Thakkar
2024-07-29 08:46:24 -04:00
committed by GitHub
parent 56b46e2d13
commit be60a0b272
312 changed files with 19793 additions and 6775 deletions
@@ -246,6 +246,9 @@ compute_lower_srt(ConvProblemShape<ConvOp, NumSpatialDimensions> const& problem_
return lower;
}
template <class CopyOp> struct is_im2col_load { static constexpr bool value = false; };
template <> struct is_im2col_load<SM90_TMA_LOAD_IM2COL > { static constexpr bool value = true; };
template <> struct is_im2col_load<SM90_TMA_LOAD_IM2COL_MULTICAST> { static constexpr bool value = true; };
/////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace cutlass::conv::collective::detail
@@ -131,6 +131,9 @@ struct CollectiveConv<
&& (cute::is_same_v<GmemTiledCopyB, SM90_TMA_LOAD> || cute::is_same_v<GmemTiledCopyB, SM90_TMA_LOAD_MULTICAST>)),
"GmemTiledCopyB - invalid SM90 TMA copy atom specified.");
static constexpr bool is_im2col_A = detail::is_im2col_load<GmemTiledCopyA>::value;
static constexpr bool is_im2col_B = detail::is_im2col_load<GmemTiledCopyB>::value;
// TMA converts f32 input to tf32 when copying from GMEM to SMEM
// For all other types, cast to size equivalent uint type to avoid any rounding by TMA.
static constexpr bool ConvertF32toTF32A = cute::is_same_v<float, ElementA>;
@@ -169,12 +172,11 @@ private:
// Note that for fprop and dgrad kernel, the tma load mode is im2col for tensor A and tiled for
// tensor B while for wgrad kernel, the tma load mode is tiled for tensor A and im2col for tensor
// B since operand A, B is swapped.
// 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) {
if constexpr (ConvOp == conv::Operator::kFprop || ConvOp == conv::Operator::kDgrad) {
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);
auto upper_corner_whd = detail::compute_upper_corner_whd(problem_shape);
@@ -203,7 +205,7 @@ private:
shape(stride_srt));
}
// TMA tiled mode for tensor A in wgrad kernel.
else if constexpr (ConvOp == conv::Operator::kWgrad) {
else {
return make_tma_copy(
GmemTiledCopyA{},
tensor_a,
@@ -217,16 +219,8 @@ private:
template <class TensorB>
static constexpr auto
get_tma_load_b_instance(TensorB const& tensor_b, typename Arguments::ProblemShape const& problem_shape) {
if constexpr (ConvOp == conv::Operator::kFprop || ConvOp == conv::Operator::kDgrad) {
return make_tma_copy(
GmemTiledCopyB{},
tensor_b,
SmemLayoutB{}(_,_,_0{}),
make_shape(shape<1>(TileShape{}), shape<2>(TileShape{})),
size<0>(ClusterShape{}));
}
// TMA im2col mode for tensor B in wgrad kernel.
else if constexpr (ConvOp == conv::Operator::kWgrad) {
if constexpr (is_im2col_B) {
// compute the upper and lower corners based on the conv padding
auto lower_corner_whd = detail::compute_lower_corner_whd(problem_shape);
auto upper_corner_whd = detail::compute_upper_corner_whd(problem_shape);
@@ -246,6 +240,26 @@ private:
shape(lower_srt),
cute::reverse(shape(problem_shape.dilation)));
}
else {
return make_tma_copy(
GmemTiledCopyB{},
tensor_b,
SmemLayoutB{}(_,_,_0{}),
make_shape(shape<1>(TileShape{}), shape<2>(TileShape{})),
size<0>(ClusterShape{}));
}
}
static constexpr auto
get_problem_shape_MNKL(typename Arguments::ProblemShape const& problem_shape) {
if constexpr (is_im2col_A || is_im2col_B) {
// transformation + im2col linearization
return problem_shape.get_linearized_problem_shape_MNKL();
}
else {
// transformation
return problem_shape.get_transformed_problem_shape_MNKL();
}
}
public:
@@ -253,9 +267,7 @@ public:
// Device side kernel params
struct Params {
using _Submode = decltype(take<0,NumTensorDimensions-1>(typename Arguments::ProblemShape::TensorExtent{}));
using ProblemShape = cute::conditional_t<DispatchPolicy::ConvOp == conv::Operator::kWgrad,
Shape<int, _Submode, _Submode>,
Shape<_Submode, int, _Submode>>;
using ProblemShape = decltype(get_problem_shape_MNKL(typename Arguments::ProblemShape{}));
// Assumption: StrideA is congruent with Problem_MK
// Select TMA load type according to convolution operator.
@@ -283,6 +295,7 @@ public:
TMA_A tma_load_a;
TMA_B tma_load_b;
ProblemShape problem_shape;
uint32_t tma_transaction_bytes = TmaTransactionBytes;
};
//
@@ -314,17 +327,18 @@ public:
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_mnk = args.problem_shape.get_transformed_problem_shape_MNK();
auto problem_shape_mnkl = get_problem_shape_MNKL(args.problem_shape);
return {
tma_load_a,
tma_load_b,
problem_shape_mnk
problem_shape_mnkl,
TmaTransactionBytes
};
}
template<class ProblemShape>
CUTLASS_HOST_DEVICE static bool
static bool
can_implement(
ProblemShape const& problem_shape,
Arguments const& args) {
@@ -389,13 +403,12 @@ public:
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 thad_idx,
int thread_idx,
uint32_t block_rank_in_cluster,
TensorStorage& shared_tensors) {
int warp_idx = canonical_warp_idx_sync();
int warp_idx_in_warp_group = warp_idx % 4;
int lane_predicate = cute::elect_one_sync();
if (warp_idx_in_warp_group == 0 and lane_predicate) {
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)
@@ -403,7 +416,8 @@ public:
// Prepare the TMA loads for A and B
//
dim3 cluster_local_block_id = cute::block_id_in_cluster();
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);
@@ -462,12 +476,10 @@ public:
/// Perform a Producer Epilogue to prevent early exit of blocks in a Cluster
CUTLASS_DEVICE void
load_tail(MainloopPipeline pipeline, PipelineState smem_pipe_producer_state) {
int warp_idx = canonical_warp_idx_sync();
int warp_idx_in_warp_group = warp_idx % 4;
int lane_predicate = cute::elect_one_sync();
// Issue the epilogue waits
if (warp_idx_in_warp_group == 0 and lane_predicate) {
if (lane_predicate) {
/* This helps avoid early exit of blocks in Cluster
* Waits for all stages to either be released (all
* Consumer UNLOCKs), or if the stage was never used
+65 -14
View File
@@ -352,15 +352,16 @@ struct ConvProblemShape {
}
}
// Get problem shape MNK according to following table:
// | | Fprop | Dgrad | Wgrad |
// | ---- | --------- | -------- | -------- |
// | Shape_M | (Q,P,Z,N) | (W,H,D,N) | (K) |
// | Shape_N | (K) | (C) | (C,S,R,T) |
// | Shape_K | (C,S,R,T) | (K,S,R,T) | (Q,P,Z,N) |
// 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_MNK() const {
get_transformed_problem_shape_MNKL() const {
using cute::insert;
using cute::make_shape;
using cute::reverse;
@@ -370,32 +371,56 @@ struct ConvProblemShape {
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);
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);
return make_shape(M_xformed, N_xformed, K_xformed, L_xformed);
}
else if constexpr (ConvOp == conv::Operator::kDgrad) {
auto M_xformed = reverse(take<0,RankT - 1>(shape_C));
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);
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))
// wgrad: A extents array contains [N,Z,P,Q,K]. Turn that into ((K), (Q,P,Z,N))
// dgrad: A extents array contains [N,Z,P,Q,K]. Turn that into ((Q,P,Z,N), (K))
// wgrad: A extents array contains [N,Z,P,Q,K]. Turn that into ((K), (Q,P,Z,N))
CUTLASS_HOST_DEVICE
constexpr auto
get_shape_A() const {
@@ -418,8 +443,8 @@ struct ConvProblemShape {
// Get B extents.
// fprop: B extents array contains [K,T,R,S,C]. Turn that into ((K), (C,S,R,T))
// wgrad: B extents array contains [N,D,H,W,C]. Turn that into ((C), (W,H,D,N))
// dgrad: B extents array contains [K,T,R,S,C]. Turn that into ((C), (K,S,R,T))
// wgrad: B extents array contains [N,D,H,W,C]. Turn that into ((C), (W,H,D,N))
CUTLASS_HOST_DEVICE
constexpr auto
get_shape_B() const {
@@ -447,6 +472,30 @@ struct ConvProblemShape {
}
}
// Get C extents.
// fprop: C extents array contains [N,Z,P,Q,K]. Turn that into ((Q,P,Z,N), (K))
// dgrad: C extents array contains [N,D,H,W,C]. Turn that into ((W,H,D,N), (C))
// wgrad: C extents array contains [K,T,R,S,C]. Turn that into ((K), (C,S,R,T))
CUTLASS_HOST_DEVICE
constexpr auto
get_shape_C() const {
using cute::make_shape;
using cute::reverse;
using cute::take;
if constexpr (ConvOp == conv::Operator::kFprop ||
ConvOp == conv::Operator::kDgrad) {
return make_shape(
reverse(take<0, RankT - 1>(shape_C)),
shape_C[RankT - 1]);
}
else if constexpr (ConvOp == conv::Operator::kWgrad) {
return make_shape(
shape_C[0],
reverse(take<1, RankT>(shape_C)));
}
}
// Static method that returns the canonical strides of tensors (layouts are right major and compact)
CUTLASS_HOST_DEVICE
static constexpr TensorStride
@@ -529,7 +578,9 @@ 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) {
return 1 + (act_ext + pad_total - ((filter_ext -1) * dilation + 1)) / tstride;
auto tmp = act_ext + pad_total - ((filter_ext -1) * dilation + 1);
CUTLASS_ASSERT(tmp % tstride == 0);
return 1 + tmp / tstride;
};
shape_xformed_act[0] = shape_act[0]; // Activation N extent
@@ -228,29 +228,18 @@ public:
/// Initializes conv state from arguments.
Status
initialize(
Arguments const& args,
void* workspace = nullptr,
cudaStream_t stream = nullptr,
CudaHostAdapter *cuda_adapter = nullptr) {
Arguments const& args,
void* workspace = nullptr,
cudaStream_t stream = nullptr,
CudaHostAdapter *cuda_adapter = nullptr) {
CUTLASS_TRACE_HOST("ConvUniversal::initialize() - workspace "
<< workspace << ", stream: " << (stream ? "non-null" : "null"));
size_t workspace_bytes = ConvKernel::get_workspace_size(args);
CUTLASS_TRACE_HOST(" workspace_bytes: " << workspace_bytes);
if (workspace_bytes) {
if (!workspace) {
CUTLASS_TRACE_HOST(" error: device workspace must not be null");
return Status::kErrorWorkspaceNull;
}
CUTLASS_TRACE_HOST(" clearing device workspace");
cudaError_t result = cudaMemsetAsync(workspace, 0, workspace_bytes, stream);
if (cudaSuccess != result) {
result = cudaGetLastError(); // to clear the error bit
CUTLASS_TRACE_HOST(" cudaMemsetAsync() returned error " << cudaGetErrorString(result));
return Status::kErrorInternal;
}
// Initialize the workspace
Status status = ConvKernel::initialize_workspace(args, workspace, stream, cuda_adapter);
if (status != Status::kSuccess) {
return status;
}
// Initialize the Params structure
@@ -297,7 +286,7 @@ public:
/// Primary run() entry point API that is static allowing users to create and manage their own params.
/// Supplied params struct must be construct by calling ConvKernel::to_underling_arguments()
static Status
run(Params& params, cudaStream_t stream = nullptr, CudaHostAdapter *cuda_adapter = nullptr) {
run(Params& params, cudaStream_t stream = nullptr, CudaHostAdapter *cuda_adapter = nullptr, int32_t kernel_index = 0) {
CUTLASS_TRACE_HOST("ConvUniversal::run()");
dim3 const block = ConvKernel::get_block_shape();
dim3 const grid = get_grid_shape(params);
@@ -319,9 +308,13 @@ public:
CUTLASS_ASSERT(cuda_adapter);
if (cuda_adapter) {
launch_result = cuda_adapter->launch(
grid, cluster, block, smem_size, stream, kernel_params, 0
);
launch_result = cuda_adapter->launch(grid,
cluster,
block,
smem_size,
stream,
kernel_params,
kernel_index);
}
else {
return Status::kErrorInternal;
@@ -379,11 +372,12 @@ public:
Arguments const& args,
void* workspace = nullptr,
cudaStream_t stream = nullptr,
CudaHostAdapter *cuda_adapter = nullptr
CudaHostAdapter *cuda_adapter = nullptr,
int32_t kernel_index = 0
) {
Status status = initialize(args, workspace, stream, cuda_adapter);
if (Status::kSuccess == status) {
status = run(params_, stream, cuda_adapter);
status = run(params_, stream, cuda_adapter, kernel_index);
}
return status;
}
@@ -197,7 +197,7 @@ public:
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_.ptr_reordered_B = args.ref_reordered_B.data();
params_.semaphore = static_cast<int *>(workspace);
return Status::kSuccess;
+3
View File
@@ -31,6 +31,7 @@
#pragma once
#include "cutlass/conv/convolution.h"
#include "cutlass/epilogue/thread/activation.h"
#include "cutlass/arch/arch.h"
#include "cute/layout.hpp"
@@ -38,6 +39,8 @@
//////////////////////////////////////////////////////////////////////////////
//////////////////////////////////////////////////////////////////////////////
namespace cutlass::conv {
//////////////////////////////////////////////////////////////////////////////
+9 -2
View File
@@ -114,7 +114,10 @@ template <
typename ElementTensor,
typename ElementVector,
typename OutputOp,
int ElementsPerAccess
int ElementsPerAccess,
typename PermuteDLayout = layout::NoPermute,
conv::StrideSupport StrideSupport = conv::StrideSupport::kUnity,
int Rank = 4
>
struct DefaultConvEpilogueWithBroadcastSimt {
using Epilogue = typename epilogue::threadblock::DefaultEpilogueWithBroadcastSimt<
@@ -124,7 +127,11 @@ struct DefaultConvEpilogueWithBroadcastSimt {
ElementTensor,
ElementVector,
OutputOp,
ElementsPerAccess
ElementsPerAccess,
false,
PermuteDLayout,
StrideSupport,
Rank
>::Epilogue;
};
@@ -197,7 +197,10 @@ struct DefaultConv3dFpropWithBroadcast <
typename EpilogueOutputOp::ElementT,
typename EpilogueOutputOp::ElementVector,
EpilogueOutputOp,
ImplicitGemmBase::Epilogue::kElementsPerAccess
ImplicitGemmBase::Epilogue::kElementsPerAccess,
layout::NoPermute,
StrideSupport,
5
>::Epilogue;
// Define the kernel
+20 -4
View File
@@ -181,7 +181,11 @@ struct DefaultDeconv2d <
ThreadblockShape,
WarpMmaSimtOp,
EpilogueOutputOp,
EpilogueOutputOp::kCount
EpilogueOutputOp::kCount,
false,
layout::NoPermute,
StrideSupport::kStrided,
4
>::Epilogue;
// Define the kernel
@@ -405,7 +409,11 @@ struct DefaultDeconv2d <
ThreadblockShape,
WarpMmaSimtOp,
EpilogueOutputOp,
EpilogueOutputOp::kCount
EpilogueOutputOp::kCount,
false,
layout::NoPermute,
StrideSupport::kStrided,
4
>::Epilogue;
// Define the kernel
@@ -627,7 +635,11 @@ struct DefaultDeconv2d <
ThreadblockShape,
WarpMmaSimtOp,
EpilogueOutputOp,
EpilogueOutputOp::kCount
EpilogueOutputOp::kCount,
false,
layout::NoPermute,
StrideSupport::kStrided,
4
>::Epilogue;
// Define the kernel
@@ -852,7 +864,11 @@ struct DefaultDeconv2d <
ThreadblockShape,
WarpMmaSimtOp,
EpilogueOutputOp,
EpilogueOutputOp::kCount
EpilogueOutputOp::kCount,
false,
layout::NoPermute,
StrideSupport::kStrided,
4
>::Epilogue;
// Define the kernel
+20 -4
View File
@@ -170,7 +170,11 @@ struct DefaultDeconv3d <
ThreadblockShape,
WarpMmaSimtOp,
EpilogueOutputOp,
EpilogueOutputOp::kCount
EpilogueOutputOp::kCount,
false,
layout::NoPermute,
StrideSupport::kStrided,
5
>::Epilogue;
// Define the kernel
@@ -282,7 +286,11 @@ struct DefaultDeconv3d <
ThreadblockShape,
WarpMmaSimtOp,
EpilogueOutputOp,
EpilogueOutputOp::kCount
EpilogueOutputOp::kCount,
false,
layout::NoPermute,
StrideSupport::kStrided,
5
>::Epilogue;
// Define the kernel
@@ -389,7 +397,11 @@ struct DefaultDeconv3d <
ThreadblockShape,
WarpMmaSimtOp,
EpilogueOutputOp,
EpilogueOutputOp::kCount
EpilogueOutputOp::kCount,
false,
layout::NoPermute,
StrideSupport::kStrided,
5
>::Epilogue;
// Define the kernel
@@ -501,7 +513,11 @@ struct DefaultDeconv3d <
ThreadblockShape,
WarpMmaSimtOp,
EpilogueOutputOp,
EpilogueOutputOp::kCount
EpilogueOutputOp::kCount,
false,
layout::NoPermute,
StrideSupport::kStrided,
5
>::Epilogue;
// Define the kernel
@@ -196,7 +196,10 @@ struct DefaultDeconv3dWithBroadcast <
typename EpilogueOutputOp::ElementT,
typename EpilogueOutputOp::ElementVector,
EpilogueOutputOp,
ImplicitGemmBase::Epilogue::kElementsPerAccess
ImplicitGemmBase::Epilogue::kElementsPerAccess,
layout::NoPermute,
StrideSupport::kStrided,
5
>::Epilogue;
// Define the kernel
@@ -273,7 +276,7 @@ struct DefaultDeconv3dWithBroadcast <
>::Kernel;
// Define epilogue
using Epilogue = typename cutlass::conv::kernel::detail::DefaultConvEpilogueWithBroadcastSimtStridedDgrad<
using Epilogue = typename cutlass::conv::kernel::detail::DefaultConvEpilogueWithBroadcastSimt<
ArchTag,
typename ImplicitGemmBase::Epilogue::Shape,
typename ImplicitGemmBase::Epilogue::WarpMmaOperator,
@@ -281,7 +284,10 @@ struct DefaultDeconv3dWithBroadcast <
typename EpilogueOutputOp::ElementT,
typename EpilogueOutputOp::ElementVector,
EpilogueOutputOp,
ImplicitGemmBase::Epilogue::kElementsPerAccess
ImplicitGemmBase::Epilogue::kElementsPerAccess,
layout::NoPermute,
StrideSupport::kStrided,
5
>::Epilogue;
// Define the kernel
@@ -233,9 +233,9 @@ struct ImplicitGemmConvolution {
ptr_A(args.ref_A.data()),
iterator_B(args.problem_size, args.ref_B.layout()),
ptr_B(args.ref_B.data()),
iterator_C(ConvOutputIteratorParameter::layout(args.ref_C), args.problem_size),
iterator_C(ConvOutputIteratorParameter::layout(args.ref_C), implicit_gemm_tensor_c_extent(kConvolutionalOperator, args.problem_size)),
ptr_C(args.ref_C.data()),
iterator_D(ConvOutputIteratorParameter::layout(args.ref_D), args.problem_size),
iterator_D(ConvOutputIteratorParameter::layout(args.ref_D), implicit_gemm_tensor_c_extent(kConvolutionalOperator, args.problem_size)),
ptr_D(args.ref_D.data()),
output_op(args.output_op),
semaphore(semaphore),
@@ -257,9 +257,9 @@ struct ImplicitGemmConvolutionWithFusedEpilogue {
ptr_A(args.ref_A.data()),
iterator_B(args.problem_size, args.ref_B.layout()),
ptr_B(args.ref_B.data()),
iterator_C(ConvOutputIteratorParameter::layout(args.ref_C)),
iterator_C(ConvOutputIteratorParameter::layout(args.ref_C), implicit_gemm_tensor_c_extent(kConvolutionalOperator, args.problem_size)),
ptr_C(args.ref_C.data()),
iterator_D(ConvOutputIteratorParameter::layout(args.ref_D)),
iterator_D(ConvOutputIteratorParameter::layout(args.ref_D), implicit_gemm_tensor_c_extent(kConvolutionalOperator, args.problem_size)),
ptr_D(args.ref_D.data()),
output_op(args.output_op),
semaphore(semaphore),
@@ -51,12 +51,12 @@ namespace cutlass::conv::kernel {
template <
class CollectiveMainloop_,
class CollectiveEpilogue_,
class TileSchedulerTag
class TileSchedulerTag_
>
class ConvUniversal<
CollectiveMainloop_,
CollectiveEpilogue_,
TileSchedulerTag,
TileSchedulerTag_,
cute::enable_if_t<cute::is_base_of_v<cutlass::conv::KernelImplicitTmaWarpSpecializedSm90,
typename CollectiveMainloop_::DispatchPolicy::Schedule>>>
{
@@ -90,6 +90,7 @@ public:
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<
@@ -144,7 +145,7 @@ public:
to_underlying_arguments(Arguments const& args, void* workspace) {
(void) workspace;
auto mainloop_params = CollectiveMainloop::to_underlying_arguments(args.mainloop, workspace);
auto problem_shape_MNKL = append<4>(mainloop_params.problem_shape, Int<1>{});
auto problem_shape_MNKL = args.mainloop.problem_shape.get_transformed_problem_shape_MNKL();
return {
mainloop_params,
@@ -157,7 +158,7 @@ public:
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_MNK(), args.epilogue);
implementable &= CollectiveEpilogue::can_implement(args.mainloop.problem_shape.get_transformed_problem_shape_MNKL(), args.epilogue);
return implementable;
}
@@ -166,19 +167,17 @@ public:
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) {
// The CONV mainloop params problem shape will be the cute::Shape<> rank-3 MNK tuple we want for grid planning
// Although conv problems do not have an L mode, we add it here to comply with the scheduler API
auto linear_problem_shape_MNKL = make_shape(
size<0>(params.mainloop.problem_shape), // M mode is linearized.
shape<1>(params.mainloop.problem_shape),
shape<2>(params.mainloop.problem_shape),
Int<1>{});
return cutlass::gemm::kernel::detail::PersistentTileSchedulerSm90::get_tiled_cta_shape_mnl(
linear_problem_shape_MNKL, TileShape{}, ClusterShape{});
params.mainloop.problem_shape, TileShape{}, ClusterShape{});
}
static dim3
@@ -205,14 +204,25 @@ public:
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) {
@@ -223,7 +233,7 @@ public:
// Mainloop Load pipeline
using MainloopPipeline = typename CollectiveMainloop::MainloopPipeline;
typename MainloopPipeline::Params mainloop_pipeline_params;
if (warp_group_role == WarpGroupRole::Producer) {
if (warp_group_role == WarpGroupRole::Producer && producer_warp_role == ProducerWarpRole::MainloopEpilogue) {
mainloop_pipeline_params.role = MainloopPipeline::ThreadCategory::Producer;
}
if (warp_group_role == WarpGroupRole::Consumer) {
@@ -231,22 +241,24 @@ public:
}
mainloop_pipeline_params.is_leader = warp_group_thread_idx == 0;
mainloop_pipeline_params.num_consumers = NumThreadsPerWarpGroup;
mainloop_pipeline_params.transaction_bytes = CollectiveMainloop::TmaTransactionBytes;
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) {
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 = 1; // 1 thread issues TMA load
epi_load_pipeline_params.producer_arv_count = NumThreadsPerWarp;
epi_load_pipeline_params.consumer_arv_count = NumThreadsPerWarpGroup;
epi_load_pipeline_params.transaction_bytes = CollectiveEpilogue::TmaTransactionBytes;
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
@@ -266,16 +278,26 @@ public:
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 M = get<0>(params.mainloop.problem_shape);
auto N = get<1>(params.mainloop.problem_shape);
auto K = get<2>(params.mainloop.problem_shape);
// output strides are coalesced so we linearize the output shape to match the shape/stride profiles
auto linear_problem_shape_MNKL = make_shape(size(M), N, K, Int<1>{});
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, size(K)));
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
@@ -288,7 +310,8 @@ public:
// 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));
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>{});
@@ -300,51 +323,43 @@ public:
auto k_tile_iter = cute::make_coord_iterator(shape<2>(gA));
auto k_tile_count = size<2>(gA);
auto c_tile_count = CollectiveEpilogue::get_load_pipe_increment(cta_tile_shape);
auto d_tile_count = CollectiveEpilogue::get_store_pipe_increment(cta_tile_shape);
// Make sure pipeline init is visible to all producers and consumer CTAs in cluster
if constexpr (size(ClusterShape{}) > 1) {
cute::cluster_arrive_relaxed();
cute::cluster_wait();
}
else {
__syncthreads();
}
// 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};
if (warp_group_role == WarpGroupRole::Producer) {
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,
thread_idx,
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);
// Wait for all thread blocks in Cluster
cluster_wait_fn();
if (collective_epilogue.is_producer_load_needed()) {
collective_epilogue.load(
epi_load_pipeline,
epi_load_pipe_producer_state,
linear_problem_shape_MNKL,
cta_tile_shape,
output_tile_coord,
tiled_mma,
warp_group_thread_idx,
shared_storage.tensors.epilogue
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 load pipeline state for the pipeline drain
epi_load_pipe_producer_state.advance(c_tile_count);
collective_epilogue.load_tail(epi_load_pipeline, epi_load_pipe_producer_state);
// 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) {
@@ -368,12 +383,13 @@ public:
);
// 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,
linear_problem_shape_MNKL,
problem_shape_MNKL,
cta_tile_shape,
output_tile_coord,
accumulators,
@@ -381,6 +397,13 @@ public:
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
);
}
}
};
@@ -251,7 +251,7 @@ public:
static Status can_implement(Conv3dProblemSize const &problem_size) {
// check alignment constraint on iterator's contiguous dimension
if (problem_size.C % (128/sizeof_bits<Element>::value)) {
if (problem_size.C % AccessType::kElements) {
return Status::kErrorInvalidProblem;
}
@@ -272,7 +272,7 @@ public:
static Status can_implement(Conv3dProblemSize const &problem_size) {
// check alignment constraint on iterator's contiguous dimension
if (problem_size.C % (128/sizeof_bits<Element>::value)) {
if (problem_size.C % AccessType::kElements) {
return Status::kErrorInvalidProblem;
}
@@ -325,7 +325,7 @@ public:
static Status can_implement(ConvProblemSize const &problem_size) {
// check alignment constraint on iterator's contiguous dimension
if (problem_size.K % (128/sizeof_bits<Element>::value)) {
if (problem_size.K % AccessType::kElements) {
return Status::kErrorInvalidProblem;
}
@@ -466,7 +466,7 @@ public:
}
// check alignment constraint on iterator's contiguous dimension
if (problem_size.K % (128/sizeof_bits<Element>::value)) {
if (problem_size.K % AccessType::kElements) {
return Status::kErrorNotSupported;
}
@@ -272,7 +272,7 @@ public:
static Status can_implement(ConvProblemSize const &problem_size) {
// check alignment constraint on iterator's contiguous dimension
if (problem_size.C % (128/sizeof_bits<Element>::value)) {
if (problem_size.C % AccessType::kElements) {
return Status::kErrorInvalidProblem;
}
@@ -455,7 +455,7 @@ public:
static Status can_implement(Conv3dProblemSize const &problem_size) {
// check alignment constraint on iterator's contiguous dimension
if (problem_size.C % (128/sizeof_bits<Element>::value)) {
if (problem_size.C % AccessType::kElements) {
return Status::kErrorInvalidProblem;
}
@@ -238,11 +238,10 @@ public:
/// Determines whether the Implicit GEMM can execute the given problem.
CUTLASS_HOST_DEVICE
static Status can_implement(ConvProblemSize const &problem_size) {
auto input_channels = (IsDeconv ? problem_size.K : problem_size.C);
auto output_channels = (IsDeconv ? problem_size.C : problem_size.K);
// check alignment constraint on iterator's contiguous dimension
if (input_channels % (128/sizeof_bits<Element>::value)) {
if (input_channels % AccessType::kElements) {
return Status::kErrorInvalidProblem;
}
return Status::kSuccess;
@@ -260,14 +260,12 @@ public:
/// Determines whether the Implicit GEMM can execute the given problem.
CUTLASS_HOST_DEVICE
static Status can_implement(Conv3dProblemSize const &problem_size) {
auto input_channels = (IsDeconv ? problem_size.K : problem_size.C);
// check alignment constraint on iterator's contiguous dimension
if (input_channels % (128/sizeof_bits<Element>::value)) {
if (input_channels % AccessType::kElements) {
return Status::kErrorInvalidProblem;
}
return Status::kSuccess;
}
};
@@ -270,7 +270,7 @@ public:
static Status can_implement(Conv3dProblemSize const &problem_size) {
// check alignment constraint on iterator's contiguous dimension
if (problem_size.C % (128/sizeof_bits<Element>::value)) {
if (problem_size.C % AccessType::kElements) {
return Status::kErrorInvalidProblem;
}
@@ -250,7 +250,7 @@ public:
fast_divmod(p, q, residual, problem_size_.Q, params_.q_mul, params_.q_shr);
int d = z * problem_size_.stride_d + precomputed_filter_t_[iteration_contiguous_];
int h = p * problem_size_.stride_h + precomputed_filter_r_[iteration_contiguous_];;
int h = p * problem_size_.stride_h + precomputed_filter_r_[iteration_contiguous_];
int w = q * problem_size_.stride_w + precomputed_filter_s_[iteration_contiguous_];
return TensorCoord(n, d, h, w, filter_c_[iteration_contiguous_]);
@@ -300,7 +300,7 @@ public:
static Status can_implement(Conv3dProblemSize const &problem_size) {
// check alignment constraint on iterator's contiguous dimension
if (problem_size.C % (128/sizeof_bits<Element>::value)) {
if (problem_size.C % AccessType::kElements) {
return Status::kErrorInvalidProblem;
}
@@ -248,7 +248,7 @@ public:
static Status can_implement(Conv3dProblemSize const &problem_size) {
// check alignment constraint on iterator's contiguous dimension
if (problem_size.K % (128/sizeof_bits<Element>::value)) {
if (problem_size.K % AccessType::kElements) {
return Status::kErrorInvalidProblem;
}
@@ -291,7 +291,7 @@ public:
static Status can_implement(Conv3dProblemSize const &problem_size) {
// check alignment constraint on iterator's contiguous dimension
if (problem_size.K % (128/sizeof_bits<Element>::value)) {
if (problem_size.K % AccessType::kElements) {
return Status::kErrorInvalidProblem;
}