CUTLASS 2.2 (#96)
Adds support for NVIDIA Ampere Architecture features. CUDA 11 Toolkit recommended.
This commit is contained in:
@@ -1,5 +1,5 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2019, NVIDIA CORPORATION. All rights reserved.
|
||||
* Copyright (c) 2017-2020, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without modification, are permitted
|
||||
* provided that the following conditions are met:
|
||||
@@ -99,61 +99,13 @@ int RematerializeBlockDimZ() {
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Threadblock swizzling function for GEMMs
|
||||
template <int N = 1>
|
||||
struct GemmIdentityThreadblockSwizzle {
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
GemmIdentityThreadblockSwizzle() { }
|
||||
|
||||
int const kTile = 1;
|
||||
|
||||
/// Returns the shape of the problem in units of logical tiles
|
||||
CUTLASS_HOST_DEVICE
|
||||
GemmCoord get_tiled_shape(
|
||||
GemmCoord problem_size,
|
||||
GemmCoord tile_size,
|
||||
int split_k_slices) const {
|
||||
|
||||
return GemmCoord(
|
||||
(problem_size.m() + tile_size.m() - 1) / tile_size.m(),
|
||||
(problem_size.n() + tile_size.n() - 1) / tile_size.n(),
|
||||
split_k_slices);
|
||||
}
|
||||
|
||||
/// Computes CUDA grid dimensions given a size in units of logical tiles
|
||||
CUTLASS_HOST_DEVICE
|
||||
dim3 get_grid_shape(GemmCoord tiled_shape) const {
|
||||
return dim3(tiled_shape.m() * kTile, (tiled_shape.n() + kTile - 1) / kTile, tiled_shape.k());
|
||||
}
|
||||
|
||||
/// Obtains the threadblock offset (in units of threadblock-scoped tiles)
|
||||
CUTLASS_DEVICE
|
||||
GemmCoord get_tile_offset() const {
|
||||
|
||||
int block_idx_x = RematerializeBlockIdxX();
|
||||
int block_idx_y = RematerializeBlockIdxY();
|
||||
|
||||
return GemmCoord{
|
||||
(block_idx_x / kTile),
|
||||
(block_idx_y * kTile) + (block_idx_x % kTile),
|
||||
RematerializeBlockIdxZ()
|
||||
};
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// A special version of GemmIdentityThreadblockSwizzle. See the choice of kTile below.
|
||||
template <typename LayoutA_, typename LayoutB_>
|
||||
struct GemmCohortThreadblockSwizzle
|
||||
{
|
||||
const int kTile =
|
||||
(platform::is_same<LayoutA_, cutlass::layout::RowMajor>::value ||
|
||||
platform::is_same<LayoutB_, cutlass::layout::ColumnMajor>::value)
|
||||
? 4
|
||||
: 1;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
GemmCohortThreadblockSwizzle() { }
|
||||
int const kTile = N;
|
||||
|
||||
/// Returns the shape of the problem in units of logical tiles
|
||||
CUTLASS_HOST_DEVICE
|
||||
@@ -271,8 +223,11 @@ struct GemmBatchedIdentityThreadblockSwizzle {
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Threadblock swizzling function for split-K GEMMs
|
||||
template <int N = 1>
|
||||
struct GemmSplitKIdentityThreadblockSwizzle {
|
||||
|
||||
int const kTile = N;
|
||||
|
||||
/// Returns the shape of the problem in units of logical tiles
|
||||
CUTLASS_HOST_DEVICE
|
||||
GemmCoord get_tiled_shape(
|
||||
@@ -289,16 +244,20 @@ struct GemmSplitKIdentityThreadblockSwizzle {
|
||||
/// Computes CUDA grid dimensions given a size in units of logical tiles
|
||||
CUTLASS_HOST_DEVICE
|
||||
dim3 get_grid_shape(GemmCoord tiled_shape) const {
|
||||
return dim3(tiled_shape.m(), tiled_shape.n(), tiled_shape.k());
|
||||
return dim3(tiled_shape.m() * kTile, (tiled_shape.n() + kTile - 1) / kTile, tiled_shape.k());
|
||||
}
|
||||
|
||||
|
||||
/// Obtains the threadblock offset (in units of threadblock-scoped tiles)
|
||||
CUTLASS_DEVICE
|
||||
GemmCoord get_tile_offset() const {
|
||||
|
||||
int block_idx_x = RematerializeBlockIdxX();
|
||||
int block_idx_y = RematerializeBlockIdxY();
|
||||
|
||||
return GemmCoord{
|
||||
RematerializeBlockIdxX(),
|
||||
RematerializeBlockIdxY(),
|
||||
(block_idx_x / kTile),
|
||||
(block_idx_y * kTile) + (block_idx_x % kTile),
|
||||
RematerializeBlockIdxZ()
|
||||
};
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user