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:
Yujia Zhai
2024-10-09 15:33:27 -04:00
committed by GitHub
co-authored by yuzhai Haicheng Wu
parent 0837a2a00a
commit cc3c29a81a
354 changed files with 105937 additions and 8197 deletions
@@ -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
+7 -71
View File
@@ -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
+137
View File
@@ -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[] = {&params_};
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();
+5 -1
View File
@@ -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