CUTLASS 3.3.0 (#1167)
* Release 3.3.0 Adds support for mixed precision GEMMs On Hopper and Ampere Adds support for < 16B aligned GEMMs on Hopper Enhancements to EVT Enhancements to Python interface Enhancements to Sub-byte type handling in CuTe Several other bug-fixes and performance improvements. * minor doc update
This commit is contained in:
@@ -43,23 +43,42 @@ using namespace cute;
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
//
|
||||
// Policies for categorical dispatch of mainloop against kernel grid schedules
|
||||
// Kernel schedule policies (the base class tags, one for each kernel layer file)
|
||||
//
|
||||
struct KernelMultistage { };
|
||||
struct KernelCpAsyncWarpSpecialized { };
|
||||
struct KernelCpAsyncWarpSpecializedPingpong { };
|
||||
struct KernelCpAsyncWarpSpecializedCooperative { };
|
||||
struct KernelTma { };
|
||||
struct KernelTmaWarpSpecialized { };
|
||||
struct KernelTmaWarpSpecializedPingpong { };
|
||||
struct KernelTmaWarpSpecializedCooperative { };
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
//
|
||||
// Builder dispatch policies (not a part of the main CUTLASS layers, simply used to opt into
|
||||
// specific collective builder dispatches)
|
||||
//
|
||||
|
||||
// FP8 related policies (including Fast Accumulation)
|
||||
struct KernelTmaWarpSpecializedFP8FastAccum : KernelTmaWarpSpecialized { };
|
||||
struct KernelTmaWarpSpecializedPingpongFP8FastAccum : KernelTmaWarpSpecializedPingpong { };
|
||||
struct KernelTmaWarpSpecializedCooperativeFP8FastAccum: KernelTmaWarpSpecializedCooperative { };
|
||||
|
||||
// Policies to opt into mixed type GEMMs
|
||||
struct KernelTmaWarpSpecializedMixedInput : KernelTmaWarpSpecialized { };
|
||||
struct KernelTmaWarpSpecializedPingpongMixedInput : KernelTmaWarpSpecializedPingpong { };
|
||||
struct KernelTmaWarpSpecializedCooperativeMixedInput: KernelTmaWarpSpecializedCooperative { };
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// Policies for dispatch of epilogue
|
||||
struct EpilogueDefault { };
|
||||
struct EpilogueTransposed { };
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
//
|
||||
// Collective Mainloop Policies
|
||||
//
|
||||
@@ -98,28 +117,30 @@ struct MainloopSm80CpAsync {
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
};
|
||||
|
||||
// n-buffer in smem (cp.async), pipelined with Hopper GMMA, WITHOUT predicated gmem loads
|
||||
// n-buffer in smem (cp.async), pipelined with Hopper GMMA, with predicated gmem loads, warp specialized dynamic schedule
|
||||
template<
|
||||
int Stages_,
|
||||
class ClusterShape_ = Shape<_1,_1,_1>
|
||||
class ClusterShape_ = Shape<_1,_1,_1>,
|
||||
class KernelSchedule = KernelCpAsyncWarpSpecialized
|
||||
>
|
||||
struct MainloopSm90CpAsyncGmmaUnpredicated {
|
||||
struct MainloopSm90CpAsyncGmmaWarpSpecialized {
|
||||
constexpr static int Stages = Stages_;
|
||||
using ClusterShape = ClusterShape_;
|
||||
using ArchTag = arch::Sm90;
|
||||
using Schedule = KernelMultistage;
|
||||
using Schedule = KernelSchedule;
|
||||
};
|
||||
|
||||
// n-buffer in smem (cp.async), pipelined with Hopper GMMA, with predicated gmem loads
|
||||
// n-buffer in smem (cp.async), pipelined with Hopper GMMA, with predicated gmem loads, warp specialized dynamic schedule
|
||||
template<
|
||||
int Stages_,
|
||||
class ClusterShape_ = Shape<_1,_1,_1>
|
||||
class ClusterShape_ = Shape<_1,_1,_1>,
|
||||
class KernelSchedule = KernelCpAsyncWarpSpecialized
|
||||
>
|
||||
struct MainloopSm90CpAsyncGmma {
|
||||
struct MainloopSm90CpAsyncGmmaRmemAWarpSpecialized {
|
||||
constexpr static int Stages = Stages_;
|
||||
using ClusterShape = ClusterShape_;
|
||||
using ArchTag = arch::Sm90;
|
||||
using Schedule = KernelMultistage;
|
||||
using Schedule = KernelSchedule;
|
||||
};
|
||||
|
||||
// n-buffer in smem (Hopper TMA), pipelined with Hopper GMMA and TMA, static schedule between TMA and GMMA
|
||||
@@ -154,13 +175,11 @@ struct MainloopSm90TmaGmmaWarpSpecialized {
|
||||
template<
|
||||
int Stages_,
|
||||
class ClusterShape_ = Shape<_1,_1,_1>,
|
||||
class KernelSchedule = KernelTmaWarpSpecialized,
|
||||
int PipelineAsyncMmaStages_ = 0
|
||||
class KernelSchedule = KernelTmaWarpSpecialized
|
||||
>
|
||||
struct MainloopSm90TmaGmmaRmemAWarpSpecialized {
|
||||
constexpr static int Stages = Stages_;
|
||||
using ClusterShape = ClusterShape_;
|
||||
constexpr static int PipelineAsyncMmaStages = PipelineAsyncMmaStages_;
|
||||
using ArchTag = arch::Sm90;
|
||||
using Schedule = KernelSchedule;
|
||||
static_assert(
|
||||
@@ -170,6 +189,26 @@ struct MainloopSm90TmaGmmaRmemAWarpSpecialized {
|
||||
"KernelSchedule must be one of the warp specialized policies");
|
||||
};
|
||||
|
||||
template<
|
||||
int Stages_,
|
||||
class ClusterShape_ = Shape<_1,_1,_1>,
|
||||
class KernelSchedule = KernelTmaWarpSpecialized
|
||||
>
|
||||
struct MainloopSm90TmaGmmaRmemAWarpSpecializedMixedInput {
|
||||
constexpr static int Stages = Stages_;
|
||||
using ClusterShape = ClusterShape_;
|
||||
using ArchTag = arch::Sm90;
|
||||
using Schedule = KernelSchedule;
|
||||
static_assert(
|
||||
cute::is_same_v<Schedule, KernelTmaWarpSpecialized> ||
|
||||
cute::is_same_v<Schedule, KernelTmaWarpSpecializedMixedInput> ||
|
||||
cute::is_same_v<Schedule, KernelTmaWarpSpecializedPingpong> ||
|
||||
cute::is_same_v<Schedule, KernelTmaWarpSpecializedPingpongMixedInput> ||
|
||||
cute::is_same_v<Schedule, KernelTmaWarpSpecializedCooperative> ||
|
||||
cute::is_same_v<Schedule, KernelTmaWarpSpecializedCooperativeMixedInput>,
|
||||
"KernelSchedule must be one of the warp specialized policies");
|
||||
};
|
||||
|
||||
// n-buffer in smem (Hopper TMA), pipelined with Hopper GMMA and TMA, Warp specialized dynamic schedule
|
||||
// For FP8 kernels
|
||||
template<
|
||||
|
||||
Reference in New Issue
Block a user