CUTLASS 3.2.1 (#1113)
* Updates for 3.2.1 release. * Minor fix in gemm op profiler for raster order. * Add scheduler mapping for raster order in the kernels.
This commit is contained in:
@@ -31,7 +31,9 @@
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/fast_math.h"
|
||||
#include "cutlass/gemm_coord.hpp"
|
||||
#include "cutlass/kernel_hardware_info.hpp"
|
||||
#include "cutlass/gemm/kernel/tile_scheduler_params.h"
|
||||
#include "cute/layout.hpp"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/arch/cluster_sm90.hpp"
|
||||
@@ -57,30 +59,15 @@ public:
|
||||
bool is_valid_tile = false;
|
||||
};
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
enum class RasterOrder {
|
||||
AlongM,
|
||||
AlongN
|
||||
};
|
||||
using Params = PersistentTileSchedulerSm90Params;
|
||||
using RasterOrder = typename Params::RasterOrder;
|
||||
using RasterOrderOptions = typename Params::RasterOrderOptions;
|
||||
|
||||
struct Arguments {
|
||||
int max_swizzle_size = 1;
|
||||
RasterOrderOptions raster_order = RasterOrderOptions::Heuristic;
|
||||
};
|
||||
|
||||
struct Params {
|
||||
|
||||
FastDivmodU64 divmod_cluster_shape_major_{};
|
||||
FastDivmodU64 divmod_cluster_shape_minor_{};
|
||||
FastDivmodU64 divmod_batch_{};
|
||||
FastDivmodU64 divmod_cluster_blk_major_{};
|
||||
|
||||
uint64_t blocks_per_problem_ = 0;
|
||||
int32_t log_swizzle_size_ = 0;
|
||||
RasterOrder raster_order_ = RasterOrder::AlongN;
|
||||
};
|
||||
// Sink scheduler params as a member
|
||||
Params scheduler_params;
|
||||
|
||||
@@ -102,40 +89,18 @@ public:
|
||||
static_assert(cute::is_static<TileShape>::value);
|
||||
static_assert(cute::is_static<ClusterShape>::value);
|
||||
|
||||
// Round up to nearest multiple of cluster dim along each mode
|
||||
auto [problem_blocks_m, problem_blocks_n, problem_blocks_l] = get_tiled_cta_shape_mnl(
|
||||
problem_shape_mnkl, tile_shape, cluster_shape);
|
||||
dim3 problem_blocks = get_tiled_cta_shape_mnl(problem_shape_mnkl, tile_shape, cluster_shape);
|
||||
|
||||
// Round up to nearest multiple of swizzle_size along each mode
|
||||
auto log_swizzle_size = get_log_swizzle_size(problem_blocks_m, problem_blocks_n, arguments.max_swizzle_size);
|
||||
problem_blocks_m = round_up(problem_blocks_m, (1 << log_swizzle_size) * cute::size<0>(cluster_shape));
|
||||
problem_blocks_n = round_up(problem_blocks_n, (1 << log_swizzle_size) * cute::size<1>(cluster_shape));
|
||||
|
||||
Params params;
|
||||
params.initialize(
|
||||
problem_blocks,
|
||||
to_gemm_coord(cluster_shape),
|
||||
hw_info,
|
||||
arguments.max_swizzle_size,
|
||||
arguments.raster_order
|
||||
);
|
||||
|
||||
RasterOrder raster_order;
|
||||
raster_order = get_rasterization_order(problem_shape_mnkl, tile_shape);
|
||||
if (raster_order == RasterOrder::AlongN) {
|
||||
return {
|
||||
FastDivmodU64(cute::size<1>(cluster_shape)),
|
||||
FastDivmodU64(cute::size<0>(cluster_shape)),
|
||||
FastDivmodU64(problem_blocks_m * problem_blocks_n),
|
||||
FastDivmodU64(problem_blocks_n / cute::size<1>(cluster_shape)),
|
||||
problem_blocks_m * problem_blocks_n * problem_blocks_l,
|
||||
log_swizzle_size,
|
||||
raster_order
|
||||
};
|
||||
}
|
||||
else {
|
||||
return {
|
||||
FastDivmodU64(cute::size<0>(cluster_shape)),
|
||||
FastDivmodU64(cute::size<1>(cluster_shape)),
|
||||
FastDivmodU64(problem_blocks_m * problem_blocks_n),
|
||||
FastDivmodU64(problem_blocks_m / cute::size<0>(cluster_shape)),
|
||||
problem_blocks_m * problem_blocks_n * problem_blocks_l,
|
||||
log_swizzle_size,
|
||||
raster_order
|
||||
};
|
||||
}
|
||||
return params;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
@@ -146,10 +111,10 @@ public:
|
||||
// like blockIdx and gridDim, with __CUDA_ARCH__.
|
||||
#if defined(__CUDA_ARCH__)
|
||||
if (params_.raster_order_ == RasterOrder::AlongN) {
|
||||
current_work_linear_idx_ = static_cast<uint64_t>(int(blockIdx.x) + (int(blockIdx.y) * int(gridDim.x)));
|
||||
current_work_linear_idx_ = uint64_t(blockIdx.x) + uint64_t(blockIdx.y) * uint64_t(gridDim.x);
|
||||
}
|
||||
else {
|
||||
current_work_linear_idx_ = static_cast<uint64_t>((int(blockIdx.x) * int(gridDim.y)) + int(blockIdx.y));
|
||||
current_work_linear_idx_ = uint64_t(blockIdx.x) * uint64_t(gridDim.y) + uint64_t(blockIdx.y);
|
||||
}
|
||||
#else
|
||||
CUTLASS_ASSERT(false && "This line should never be reached");
|
||||
@@ -187,7 +152,7 @@ public:
|
||||
// MSVC requires protecting use of CUDA-specific nonstandard syntax,
|
||||
// like blockIdx and gridDim, with __CUDA_ARCH__.
|
||||
#if defined(__CUDA_ARCH__)
|
||||
current_work_linear_idx_ += static_cast<uint64_t>(int(gridDim.x) * int(gridDim.y) * int(gridDim.z)) * advance_count;
|
||||
current_work_linear_idx_ += uint64_t(gridDim.x) * uint64_t(gridDim.y) * uint64_t(gridDim.z) * uint64_t(advance_count);
|
||||
#else
|
||||
CUTLASS_ASSERT(false && "This line should never be reached");
|
||||
#endif
|
||||
@@ -246,35 +211,14 @@ public:
|
||||
CUTLASS_HOST_DEVICE static
|
||||
dim3
|
||||
get_tiled_cta_shape_mnl(ProblemShapeMNKL problem_shape_mnkl, BlockShape cta_shape, ClusterShape cluster_shape) {
|
||||
// Across M and N is our Cluster tile, so we must round up the blocks to the nearest whole number of Cluster tiles
|
||||
auto cta_m = cute::size(cute::ceil_div(cute::shape<0>(problem_shape_mnkl), cute::shape<0>(cta_shape)));
|
||||
auto cta_n = cute::size(cute::ceil_div(cute::shape<1>(problem_shape_mnkl), cute::shape<1>(cta_shape)));
|
||||
|
||||
// Round up to nearest multiple of cluster dim along each mode
|
||||
int problem_blocks_m = round_up(cta_m, cute::size<0>(cluster_shape));
|
||||
int problem_blocks_n = round_up(cta_n, cute::size<1>(cluster_shape));
|
||||
|
||||
// Cluster tile does not span the batch mode, so no extra rounding up required for it
|
||||
int problem_blocks_l = int(cute::size<3>(problem_shape_mnkl));
|
||||
return {uint32_t(problem_blocks_m), uint32_t(problem_blocks_n), uint32_t(problem_blocks_l)};
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static int32_t
|
||||
get_log_swizzle_size(int problem_ctas_m, int problem_ctas_n, int max_swizzle_size) {
|
||||
int min_cta_dim = min(problem_ctas_m, problem_ctas_n);
|
||||
if (max_swizzle_size >= 8 && min_cta_dim >= 6) {
|
||||
return 3;
|
||||
}
|
||||
else if (max_swizzle_size >= 4 && min_cta_dim >= 3) {
|
||||
return 2;
|
||||
}
|
||||
else if (max_swizzle_size >= 2 && min_cta_dim >= 2) {
|
||||
return 1;
|
||||
}
|
||||
else {
|
||||
return 0;
|
||||
}
|
||||
return Params::get_tiled_cta_shape_mnl(
|
||||
to_gemm_coord(problem_shape_mnkl),
|
||||
to_gemm_coord(cluster_shape),
|
||||
cta_m, cta_n
|
||||
);
|
||||
}
|
||||
|
||||
// Given the inputs, computes the physical grid we should launch.
|
||||
@@ -289,111 +233,17 @@ public:
|
||||
Arguments arguments,
|
||||
bool truncate_by_problem_size=true) {
|
||||
|
||||
int const sm_count = hw_info.sm_count;
|
||||
CUTLASS_TRACE_HOST("get_grid_shape(): Persistent schedule grid plan using SM count = " << sm_count);
|
||||
auto problem_shape_mnkl = cute::append<4>(problem_shape_mnk, cute::Int<1>{});
|
||||
dim3 problem_blocks = get_tiled_cta_shape_mnl(problem_shape_mnkl, cta_shape, cluster_shape);
|
||||
|
||||
// Compute the total number of output tiles our problem has
|
||||
auto problem_shape_MNKL = cute::append<4>(problem_shape_mnk, cute::Int<1>{});
|
||||
auto [problem_blocks_m, problem_blocks_n, problem_blocks_l] =
|
||||
get_tiled_cta_shape_mnl(problem_shape_MNKL, cta_shape, cluster_shape);
|
||||
|
||||
// Round up to nearest multiple of swizzle_size along each mode
|
||||
auto swizzle_size = 1 << get_log_swizzle_size(problem_blocks_m, problem_blocks_n, arguments.max_swizzle_size);
|
||||
problem_blocks_m = round_up(problem_blocks_m, swizzle_size * cute::size<0>(cluster_shape));
|
||||
problem_blocks_n = round_up(problem_blocks_n, swizzle_size * cute::size<1>(cluster_shape));
|
||||
|
||||
int problem_blocks_total = problem_blocks_m * problem_blocks_n * problem_blocks_l;
|
||||
|
||||
RasterOrder raster_order;
|
||||
raster_order = get_rasterization_order(problem_shape_mnk, cta_shape);
|
||||
dim3 launch_grid;
|
||||
|
||||
if (raster_order == RasterOrder::AlongN) {
|
||||
launch_grid = dim3(cute::size<0>(cluster_shape), 1, 1);
|
||||
}
|
||||
else {
|
||||
launch_grid = dim3(1, cute::size<1>(cluster_shape), 1);
|
||||
}
|
||||
|
||||
auto possibly_truncate = [&](int x, int y) {
|
||||
if (truncate_by_problem_size) {
|
||||
return std::min(x, y);
|
||||
}
|
||||
else {
|
||||
return x;
|
||||
}
|
||||
};
|
||||
|
||||
// The else path is generic, however, we can avoid some divs if we know cluster size is 1
|
||||
if constexpr (size(cluster_shape) == 1) {
|
||||
if (raster_order == RasterOrder::AlongN) {
|
||||
launch_grid.y = possibly_truncate(sm_count, problem_blocks_total);
|
||||
}
|
||||
else {
|
||||
launch_grid.x = possibly_truncate(sm_count, problem_blocks_total);
|
||||
}
|
||||
}
|
||||
else {
|
||||
/*
|
||||
* Optimal grid size calculation is based on
|
||||
* GH100: 8 GPCs, 72 TPCs (9 TPCs/GPC), 2 SMs/TPC, 144 SMs per full GPU
|
||||
* Hence, maximum SMs per GPC = 18
|
||||
*/
|
||||
constexpr int max_sm_per_gpc = 18;
|
||||
// Provided SM count could possibly be less than the assumed maximum SMs per GPC
|
||||
int const min_num_gpc = sm_count < max_sm_per_gpc ? 1 : sm_count / max_sm_per_gpc;
|
||||
int const max_cta_occupancy_per_gpc = max_sm_per_gpc - (max_sm_per_gpc % size(cluster_shape));
|
||||
int cta_per_device = min_num_gpc * max_cta_occupancy_per_gpc;
|
||||
|
||||
// The calculation below allows for larger grid size launch for different GPUs.
|
||||
int const num_gpc_residual = sm_count < max_sm_per_gpc ? 0 : sm_count % max_sm_per_gpc;
|
||||
int const max_cta_occupancy_per_residual_gpc = num_gpc_residual - (num_gpc_residual % size(cluster_shape));
|
||||
cta_per_device += max_cta_occupancy_per_residual_gpc;
|
||||
|
||||
cta_per_device = sm_count < cta_per_device ? sm_count : cta_per_device;
|
||||
|
||||
if (raster_order == RasterOrder::AlongN) {
|
||||
launch_grid.y = possibly_truncate(
|
||||
cta_per_device / cute::size<0>(cluster_shape),
|
||||
problem_blocks_total / cute::size<0>(cluster_shape));
|
||||
}
|
||||
else {
|
||||
launch_grid.x = possibly_truncate(
|
||||
cta_per_device / cute::size<1>(cluster_shape),
|
||||
problem_blocks_total / cute::size<1>(cluster_shape));
|
||||
}
|
||||
}
|
||||
return launch_grid;
|
||||
}
|
||||
|
||||
template <class ProblemShapeMNKL, class BlockShape>
|
||||
CUTLASS_HOST_DEVICE static RasterOrder get_rasterization_order(ProblemShapeMNKL problem_shape_mnkl, BlockShape cta_shape) {
|
||||
auto tiles_m = cute::size(cute::ceil_div(cute::shape<0>(problem_shape_mnkl), cute::shape<0>(cta_shape)));
|
||||
auto tiles_n = cute::size(cute::ceil_div(cute::shape<1>(problem_shape_mnkl), cute::shape<1>(cta_shape)));
|
||||
|
||||
if (tiles_n > tiles_m) {
|
||||
return RasterOrder::AlongM;
|
||||
}
|
||||
|
||||
return RasterOrder::AlongN;
|
||||
}
|
||||
|
||||
// Splits an input tensor with MxK according to the splitting configuration specified by work_tile_info.
|
||||
// Since the basic tile scheduler does not split output tiles, this method is a no-op.
|
||||
template<class Engine, class Layout>
|
||||
CUTLASS_DEVICE
|
||||
static auto
|
||||
split_MK(cute::Tensor<Engine, Layout> const& tensor, WorkTileInfo const&) {
|
||||
return tensor;
|
||||
}
|
||||
|
||||
// Splits an input tensor with NxK tiles according to the splitting configuration specified by work_tile_info.
|
||||
// Since the basic tile scheduler does not split output tiles, this method is a no-op.
|
||||
template<class Engine, class Layout>
|
||||
CUTLASS_DEVICE
|
||||
static auto
|
||||
split_NK(cute::Tensor<Engine, Layout> const& tensor, WorkTileInfo const&) {
|
||||
return tensor;
|
||||
return Params::get_grid_shape(
|
||||
problem_blocks,
|
||||
to_gemm_coord(cluster_shape),
|
||||
hw_info,
|
||||
arguments.max_swizzle_size,
|
||||
arguments.raster_order,
|
||||
/* truncate_by_problem_size = */true
|
||||
);
|
||||
}
|
||||
|
||||
// Returns whether the block assigned this work should compute the epilogue for the corresponding
|
||||
@@ -441,6 +291,13 @@ public:
|
||||
// space of the output tile assigned to the work unit.
|
||||
return cute::size(cute::ceil_div(cute::get<2>(problem_shape), cute::get<2>(tile_shape)));
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static uint32_t
|
||||
get_work_k_tile_start(WorkTileInfo const&) {
|
||||
// All work units returned by this scheduler start from K tile 0
|
||||
return 0u;
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace cutlass::gemm::kernel::detail
|
||||
|
||||
Reference in New Issue
Block a user