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:
ANIKET SHIVAM
2023-09-26 17:24:26 -04:00
committed by GitHub
parent e0aaa3c3b3
commit 90d3b0fb18
428 changed files with 22252 additions and 21761 deletions
@@ -46,7 +46,8 @@
#include "testbed.h"
#if defined(CUTLASS_ARCH_MMA_SM80_SUPPORTED)
#if defined(CUTLASS_ARCH_MMA_B1_AND_SM80_ENABLED)
////////////////////////////////////////////////////////////////////////////////
TEST(SM80_Device_Gemm_b1t_b1n_s32n_tensor_op_s32, 128x256x1024_64x64x1024) {
@@ -370,8 +371,12 @@ TEST(SM80_Device_Gemm_b1t_b1n_s32n_tensor_op_s32, 64x64x512_32x32x512) {
EXPECT_TRUE(test::gemm::device::TestAllGemmBasic<Gemm>());
}
#endif // defined(CUTLASS_ARCH_MMA_B1_AND_SM80_ENABLED)
////////////////////////////////////////////////////////////////////////////////
#if defined(CUTLASS_ARCH_MMA_B1_XOR_SM80_ENABLED)
TEST(SM80_Device_Gemm_XOR_b1t_b1n_s32n_tensor_op_s32, 128x256x1024_64x64x1024) {
using ElementOutput = int32_t;
using ElementAccumulator = int32_t;
@@ -694,6 +699,6 @@ TEST(SM80_Device_Gemm_XOR_b1t_b1n_s32n_tensor_op_s32, 64x64x512_32x32x512) {
EXPECT_TRUE(test::gemm::device::TestAllGemmBasic<Gemm>());
}
////////////////////////////////////////////////////////////////////////////////
#endif // defined(CUTLASS_ARCH_MMA_B1_XOR_SM80_ENABLED)
#endif // #if defined(CUTLASS_ARCH_MMA_SM80_SUPPORTED)
////////////////////////////////////////////////////////////////////////////////
@@ -47,10 +47,9 @@
#include "testbed.h"
#if defined(CUTLASS_ARCH_MMA_SM80_SUPPORTED)
////////////////////////////////////////////////////////////////////////////////
////////////////////////////////////////////////////////////////////////////////
#if defined(CUTLASS_ARCH_MMA_B1_XOR_SM80_ENABLED)
CUTLASS_TEST_L1(SM80_Device_Gemm_XOR_b1t_b1n_s32t_tensor_op_s32, 128x256x1024_64x64x1024, {
using ElementOutput = int32_t;
@@ -376,4 +375,4 @@ CUTLASS_TEST_L1(SM80_Device_Gemm_XOR_b1t_b1n_s32t_tensor_op_s32, 64x64x512_32x32
////////////////////////////////////////////////////////////////////////////////
#endif // #if defined(CUTLASS_ARCH_MMA_SM80_SUPPORTED)
#endif // #if defined(CUTLASS_ARCH_MMA_B1_XOR_SM80_ENABLED)
@@ -49,7 +49,6 @@
#include "testbed_interleaved.h"
#if defined(CUTLASS_ARCH_MMA_SM75_SUPPORTED)
////////////////////////////////////////////////////////////////////////////////
TEST(SM75_Device_Gemm_s4n_s4t_s4n_tensor_op_s32, 64x128x128_32x64x128) {
@@ -195,5 +194,4 @@ TEST(SM75_Device_Gemm_s4n_s4t_s4n_tensor_op_s32, 128x256x128_64x64x128) {
}
////////////////////////////////////////////////////////////////////////////////
#endif
@@ -49,7 +49,6 @@
#include "testbed.h"
#if defined(CUTLASS_ARCH_MMA_SM75_SUPPORTED)
/////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SM75_Device_Gemm_s4t_s4n_s32n_tensor_op_s32, 128x256x128_64x64x128) {
@@ -245,5 +244,4 @@ TEST(SM75_Device_Gemm_s4t_s4n_s32n_tensor_op_s32, 64x64x128_32x32x128) {
}
/////////////////////////////////////////////////////////////////////////////////////////////////
#endif
@@ -49,7 +49,6 @@
#include "cutlass/util/reference/host/gemm.h"
#include "testbed.h"
/////////////////////////////////////////////////////////////////////////////////////////////////
///////// WMMA Instruction Shape = 8x8x32, DataType/Instruction = s4 * s4 + s32 => s32 //////////
/////////////////////////////////////////////////////////////////////////////////////////////////
@@ -244,5 +243,4 @@ TEST(SM75_Device_Gemm_s4t_s4n_s32n_wmma_tensor_op_s32, 64x64x128_32x32x128_8x8x3
EXPECT_TRUE(test::gemm::device::TestAllGemm<Gemm>());
}
#endif //CUTLASS_SUBBYTE_INTEGER_MATRIX_MULTIPLY_ENABLED
@@ -49,7 +49,6 @@
#include "testbed.h"
#if defined(CUTLASS_ARCH_MMA_SM75_SUPPORTED)
/////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SM75_Device_Gemm_s4t_s4n_s32t_tensor_op_s32, 128x256x128_64x64x128) {
@@ -245,5 +244,4 @@ TEST(SM75_Device_Gemm_s4t_s4n_s32t_tensor_op_s32, 64x64x128_32x32x128) {
}
/////////////////////////////////////////////////////////////////////////////////////////////////
#endif
@@ -49,7 +49,6 @@
#include "cutlass/util/reference/host/gemm.h"
#include "testbed.h"
/////////////////////////////////////////////////////////////////////////////////////////////////
///////// WMMA Instruction Shape = 8x8x32, DataType/Instruction = s4 * s4 + s32 => s32 //////////
/////////////////////////////////////////////////////////////////////////////////////////////////
@@ -49,7 +49,6 @@
#include "testbed.h"
#if defined(CUTLASS_ARCH_MMA_SM75_SUPPORTED)
/////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SM75_Device_Gemm_s4t_s4n_s4n_tensor_op_s32, 128x256x128_64x64x128) {
@@ -49,7 +49,6 @@
#include "testbed.h"
#if defined(CUTLASS_ARCH_MMA_SM75_SUPPORTED)
/////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SM75_Device_Gemm_s4t_s4n_s4t_tensor_op_s32, 128x256x128_64x64x128) {
@@ -243,6 +242,7 @@ TEST(SM75_Device_Gemm_s4t_s4n_s4t_tensor_op_s32, 256x64x128_64x64x128) {
EXPECT_TRUE(test::gemm::device::TestAllGemmBasic<Gemm>());
}
TEST(SM75_Device_Gemm_s4t_s4n_s4t_tensor_op_s32, 64x128x128_32x64x128) {
using ElementOutput = cutlass::int4b_t;
@@ -340,5 +340,4 @@ TEST(SM75_Device_Gemm_s4t_s4n_s4t_tensor_op_s32, 64x64x128_32x32x128) {
}
/////////////////////////////////////////////////////////////////////////////////////////////////
#endif
@@ -49,7 +49,6 @@
#include "testbed_interleaved.h"
#if defined(CUTLASS_ARCH_MMA_SM75_SUPPORTED)
////////////////////////////////////////////////////////////////////////////////
TEST(SM75_Device_Gemm_s8n_s8t_s8n_tensor_op_s32, 32x64x64_16x32x64) {
@@ -289,5 +288,4 @@ TEST(SM75_Device_Gemm_s8n_s8t_s8n_tensor_op_s32, 128x256x64_64x64x64) {
}
////////////////////////////////////////////////////////////////////////////////
#endif
@@ -49,7 +49,6 @@
#include "testbed.h"
#if defined(CUTLASS_ARCH_MMA_SM75_SUPPORTED)
/////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SM75_Device_Gemm_s8t_s8n_s32n_tensor_op_s32, 128x256x64_64x64x64) {
@@ -245,5 +244,4 @@ TEST(SM75_Device_Gemm_s8t_s8n_s32n_tensor_op_s32, 64x64x64_32x32x64) {
}
/////////////////////////////////////////////////////////////////////////////////////////////////
#endif
@@ -49,7 +49,6 @@
#include "testbed.h"
#if defined(CUTLASS_ARCH_MMA_SM75_SUPPORTED)
/////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SM75_Device_Gemm_s8t_s8n_s32t_tensor_op_s32, 128x256x64_64x64x64) {
@@ -245,5 +244,4 @@ TEST(SM75_Device_Gemm_s8t_s8n_s32t_tensor_op_s32, 64x64x64_32x32x64) {
}
/////////////////////////////////////////////////////////////////////////////////////////////////
#endif
@@ -49,7 +49,6 @@
#include "testbed.h"
#if defined(CUTLASS_ARCH_MMA_SM75_SUPPORTED)
/////////////////////////////////////////////////////////////////////////////////////////////////
CUTLASS_TEST_L0(SM75_Device_Gemm_s8t_s8n_s8n_tensor_op_s32, 128x256x64_64x64x64, {
@@ -212,5 +211,4 @@ CUTLASS_TEST_L0(SM75_Device_Gemm_s8t_s8n_s8n_tensor_op_s32, 64x64x64_32x32x64, {
} )
/////////////////////////////////////////////////////////////////////////////////////////////////
#endif
@@ -49,7 +49,6 @@
#include "testbed.h"
#if defined(CUTLASS_ARCH_MMA_SM75_SUPPORTED)
/////////////////////////////////////////////////////////////////////////////////////////////////
CUTLASS_TEST_L0(SM75_Device_Gemm_s8t_s8n_s8t_tensor_op_s32, 128x256x64_64x64x64, {
@@ -186,5 +185,4 @@ CUTLASS_TEST_L0(SM75_Device_Gemm_s8t_s8n_s8t_tensor_op_s32, 64x64x64_32x32x64, {
} )
/////////////////////////////////////////////////////////////////////////////////////////////////
#endif
+12 -6
View File
@@ -544,7 +544,6 @@ struct TestbedImpl {
// Initialize the GEMM operator
//
typename Gemm::Arguments arguments;
cutlass::KernelHardwareInfo hw_info;
hw_info.device_id = 0;
if (not profiling) {
@@ -557,12 +556,12 @@ struct TestbedImpl {
}
typename Gemm::GemmKernel::TileScheduler::Arguments scheduler_args;
if constexpr (std::is_same_v<typename Gemm::GemmKernel::TileScheduleTag, cutlass::gemm::StreamKScheduler>) {
if constexpr (std::is_same_v<typename Gemm::GemmKernel::TileSchedulerTag, cutlass::gemm::StreamKScheduler>) {
scheduler_args = { static_cast<int>(splits) };
}
// DefaultEpilogue
arguments = typename Gemm::Arguments{
auto arguments = typename Gemm::Arguments {
cutlass::gemm::GemmUniversalMode::kGemm,
problem_size,
{
@@ -741,7 +740,7 @@ struct Testbed3xFusionOperation {
using ElementAux = non_void_t<typename FusionOp::ElementAux>;
using ElementAmax = non_void_t<typename FusionOp::ElementAmax>;
using LayoutTagAux = non_void_t<typename FusionOp::GmemLayoutTagAux, LayoutTagD>;
using ActivationFunctor = non_void_t<typename FusionOp::ActivationFn<ElementCompute>,
using ActivationFunctor = non_void_t<typename FusionOp::ActivationFn,
cutlass::epilogue::thread::Identity<ElementCompute>>;
static constexpr bool IsBiasEnabled = FusionOp::IsPerRowBiasSupported;
@@ -1152,7 +1151,7 @@ struct Testbed3xFusionOperation {
initialize(problem_size, alpha_, beta_);
typename Gemm::GemmKernel::TileScheduler::Arguments scheduler_args;
if constexpr (std::is_same_v<typename Gemm::GemmKernel::TileScheduleTag, cutlass::gemm::StreamKScheduler>) {
if constexpr (std::is_same_v<typename Gemm::GemmKernel::TileSchedulerTag, cutlass::gemm::StreamKScheduler>) {
scheduler_args = { static_cast<int>(splits) };
}
@@ -1208,6 +1207,13 @@ struct Testbed3xFusionOperation {
fusion_args.bias_ptr = bias.device_data();
}
// example of how to set kernel activation arguments
if constexpr (cute::is_same_v<ActivationFunctor, cutlass::epilogue::thread::ScaledGELU_taylor<ElementCompute>>) {
// see ActivationFunctor::Arguments in activation.h for definition
// if Arguments doesn't exist then fusion_args.activation is empty
fusion_args.activation.scale = ElementCompute(1);
}
if constexpr (IsAbsMaxEnabled) {
fusion_args.amax_D_ptr = abs_max_D.device_data();
}
@@ -1297,7 +1303,7 @@ bool TestAll(double alpha = 1.0, double beta = 0.0, Testbed testbed = {}) {
std::vector<int> problem_size_k = {max_alignment, TileShapeK * (Stages + 1) - max_alignment};
std::vector<int> problem_splits = {1};
if constexpr (std::is_same_v<typename Gemm::GemmKernel::TileScheduleTag, cutlass::gemm::StreamKScheduler>) {
if constexpr (std::is_same_v<typename Gemm::GemmKernel::TileSchedulerTag, cutlass::gemm::StreamKScheduler>) {
problem_splits.push_back(2);
problem_splits.push_back(3);
@@ -1316,7 +1316,7 @@ public:
}
typename Gemm::GemmKernel::TileScheduler::Arguments scheduler_args;
if constexpr (std::is_same_v<typename Gemm::GemmKernel::TileScheduleTag, cutlass::gemm::StreamKScheduler>) {
if constexpr (std::is_same_v<typename Gemm::GemmKernel::TileSchedulerTag, cutlass::gemm::StreamKScheduler>) {
scheduler_args = { splits };
}
@@ -454,57 +454,4 @@ using Sm90LinCombScalarReduce =
>;
} // namespace fusion
namespace collective {
template<
typename TileShape_MNK,
typename EpilogueTileType,
typename ElementC,
typename ElementD,
typename Schedule
>
struct EpilogueDescriptor{
using TileShape = TileShape_MNK;
using EpilogueTile =
decltype(detail::sm90_compute_tile_shape_or_override<ElementD, EpilogueTileType, Schedule>());
using DispatchPolicy =
decltype(detail::sm90_get_tma_dispatch_policy<TileShape_MNK,EpilogueTile,ElementC,ElementD, Schedule>());
constexpr static int StagesC = DispatchPolicy::StagesC;
constexpr static int StagesD = DispatchPolicy::StagesD;
};
template<
typename EpilogueDescriptor,
typename GmemLayoutTagAux,
typename ElementAux
>
struct AuxLoadDescriptor{
constexpr static int Stages = EpilogueDescriptor::StagesC;
using Element = ElementAux;
using Stride = gemm::TagToStrideC_t<GmemLayoutTagAux>;
using SmemLayoutAtom =
decltype(detail::sm90_get_epilogue_smem_swizzle_layout_atom<Stride, ElementAux, typename EpilogueDescriptor::EpilogueTile>());
using CopyOpS2R =
decltype(detail::sm90_get_smem_load_op_for_source<Stride, ElementAux>());
};
template<
typename EpilogueDescriptor,
typename GmemLayoutTagAux,
typename ElementAux
>
struct AuxStoreDescriptor{
constexpr static int Stages = EpilogueDescriptor::StagesD;
using Element = ElementAux;
using Stride = gemm::TagToStrideC_t<GmemLayoutTagAux>;
using SmemLayoutAtom =
decltype(detail::sm90_get_epilogue_smem_swizzle_layout_atom<Stride, ElementAux, typename EpilogueDescriptor::EpilogueTile>());
using CopyOpR2S =
decltype(detail::sm90_get_smem_store_op_for_accumulator<Stride, ElementAux>());
};
} // namespace collective
} // namespace cutlass::epilogue
@@ -70,10 +70,10 @@ TEST(SM90_Device_Gemm_f16t_f16n_f32t_tensor_op_gmma_f32_cooperative_epilogue, 25
using EpilogueSchedule = cutlass::epilogue::TmaWarpSpecializedCooperative;
using EpilogueTileType = cutlass::epilogue::collective::EpilogueTileAuto;
using EpilogueDescriptor = cutlass::epilogue::collective::EpilogueDescriptor<
using EpilogueDescriptor = cutlass::epilogue::collective::detail::EpilogueDescriptor<
TileShape_MNK, EpilogueTileType, cutlass::half_t, cutlass::half_t, EpilogueSchedule
>;
using AuxLoadDescriptor = cutlass::epilogue::collective::AuxLoadDescriptor<
using AuxLoadDescriptor = cutlass::epilogue::collective::detail::AuxLoadDescriptor<
EpilogueDescriptor, cutlass::layout::RowMajor, cutlass::half_t
>;
@@ -128,10 +128,10 @@ TEST(SM90_Device_Gemm_f16t_f16n_f32t_tensor_op_gmma_f32_cooperative_epilogue, 25
using EpilogueSchedule = cutlass::epilogue::TmaWarpSpecializedCooperative;
using EpilogueTileType = cutlass::epilogue::collective::EpilogueTileAuto;
using EpilogueDescriptor = cutlass::epilogue::collective::EpilogueDescriptor<
using EpilogueDescriptor = cutlass::epilogue::collective::detail::EpilogueDescriptor<
TileShape_MNK, EpilogueTileType, cutlass::half_t, cutlass::half_t, EpilogueSchedule
>;
using AuxLoadDescriptor = cutlass::epilogue::collective::AuxLoadDescriptor<
using AuxLoadDescriptor = cutlass::epilogue::collective::detail::AuxLoadDescriptor<
EpilogueDescriptor, cutlass::layout::ColumnMajor, cutlass::half_t
>;
@@ -185,10 +185,10 @@ TEST(SM90_Device_Gemm_f16t_f16n_f32t_tensor_op_gmma_f32_cooperative_epilogue, 12
using EpilogueSchedule = cutlass::epilogue::TmaWarpSpecializedCooperative;
using EpilogueTileType = cutlass::epilogue::collective::EpilogueTileAuto;
using EpilogueDescriptor = cutlass::epilogue::collective::EpilogueDescriptor<
using EpilogueDescriptor = cutlass::epilogue::collective::detail::EpilogueDescriptor<
TileShape_MNK, EpilogueTileType, cutlass::half_t, cutlass::half_t, EpilogueSchedule
>;
using AuxLoadDescriptor = cutlass::epilogue::collective::AuxLoadDescriptor<
using AuxLoadDescriptor = cutlass::epilogue::collective::detail::AuxLoadDescriptor<
EpilogueDescriptor, cutlass::layout::ColumnMajor, float
>;
@@ -72,10 +72,10 @@ TEST(SM90_Device_Gemm_f16t_f16n_f32t_tensor_op_gmma_f32_cooperative_epilogue, 25
using EpilogueSchedule = cutlass::epilogue::TmaWarpSpecializedCooperative;
using EpilogueTileType = cutlass::epilogue::collective::EpilogueTileAuto;
using EpilogueDescriptor = cutlass::epilogue::collective::EpilogueDescriptor<
using EpilogueDescriptor = cutlass::epilogue::collective::detail::EpilogueDescriptor<
TileShape_MNK, EpilogueTileType, cutlass::half_t, cutlass::half_t, EpilogueSchedule>;
using AuxLoadDescriptor = cutlass::epilogue::collective::AuxLoadDescriptor<
using AuxLoadDescriptor = cutlass::epilogue::collective::detail::AuxLoadDescriptor<
EpilogueDescriptor, cutlass::layout::RowMajor, cutlass::half_t>;
using FusionCallbacks = cutlass::epilogue::fusion::Sm90LinCombEVTDAG<
@@ -125,10 +125,10 @@ TEST(SM90_Device_Gemm_f16t_f16n_f32t_tensor_op_gmma_f32_cooperative_epilogue, 12
using EpilogueSchedule = cutlass::epilogue::TmaWarpSpecializedCooperative;
using EpilogueTileType = cutlass::epilogue::collective::EpilogueTileAuto;
using EpilogueDescriptor = cutlass::epilogue::collective::EpilogueDescriptor<
using EpilogueDescriptor = cutlass::epilogue::collective::detail::EpilogueDescriptor<
TileShape_MNK, EpilogueTileType, cutlass::half_t, cutlass::half_t, EpilogueSchedule>;
using AuxStoreDescriptor = cutlass::epilogue::collective::AuxStoreDescriptor<
using AuxStoreDescriptor = cutlass::epilogue::collective::detail::AuxStoreDescriptor<
EpilogueDescriptor, cutlass::layout::RowMajor, cutlass::half_t>;
using FusionCallbacks = cutlass::epilogue::fusion::Sm90LinCombDAGEVT<
@@ -71,7 +71,7 @@ TEST(SM90_Device_Gemm_f16t_f16n_f32t_tensor_op_gmma_f32_cooperative_epilogue, 25
using EpilogueSchedule = cutlass::epilogue::TmaWarpSpecializedCooperative;
using EpilogueTileType = cutlass::epilogue::collective::EpilogueTileAuto;
using EpilogueDescriptor = cutlass::epilogue::collective::EpilogueDescriptor<
using EpilogueDescriptor = cutlass::epilogue::collective::detail::EpilogueDescriptor<
TileShape_MNK, EpilogueTileType, cutlass::half_t, cutlass::half_t, EpilogueSchedule>;
using FusionCallbacks = cutlass::epilogue::fusion::Sm90LinCombPerColumnBias<
@@ -121,7 +121,7 @@ TEST(SM90_Device_Gemm_f16t_f16n_f32t_tensor_op_gmma_f32_cooperative_epilogue, 25
using EpilogueSchedule = cutlass::epilogue::TmaWarpSpecializedCooperative;
using EpilogueTileType = cutlass::epilogue::collective::EpilogueTileAuto;
using EpilogueDescriptor = cutlass::epilogue::collective::EpilogueDescriptor<
using EpilogueDescriptor = cutlass::epilogue::collective::detail::EpilogueDescriptor<
TileShape_MNK, EpilogueTileType, cutlass::half_t, cutlass::half_t, EpilogueSchedule>;
using FusionCallbacks = cutlass::epilogue::fusion::Sm90LinCombPerColumnBias<
@@ -71,10 +71,10 @@ TEST(SM90_Device_Gemm_f16t_f16n_f32t_tensor_op_gmma_f32_persistent_epilogue, 128
using EpilogueSchedule = cutlass::epilogue::TmaWarpSpecialized;
using EpilogueTileType = cutlass::epilogue::collective::EpilogueTileAuto;
using EpilogueDescriptor = cutlass::epilogue::collective::EpilogueDescriptor<
using EpilogueDescriptor = cutlass::epilogue::collective::detail::EpilogueDescriptor<
TileShape_MNK, EpilogueTileType, cutlass::half_t, cutlass::half_t, EpilogueSchedule
>;
using AuxLoadDescriptor = cutlass::epilogue::collective::AuxLoadDescriptor<
using AuxLoadDescriptor = cutlass::epilogue::collective::detail::AuxLoadDescriptor<
EpilogueDescriptor, cutlass::layout::RowMajor, cutlass::half_t
>;
@@ -127,10 +127,10 @@ TEST(SM90_Device_Gemm_f16t_f16n_f32t_tensor_op_gmma_f32_persistent_epilogue, 128
using EpilogueSchedule = cutlass::epilogue::TmaWarpSpecialized;
using EpilogueTileType = cutlass::epilogue::collective::EpilogueTileAuto;
using EpilogueDescriptor = cutlass::epilogue::collective::EpilogueDescriptor<
using EpilogueDescriptor = cutlass::epilogue::collective::detail::EpilogueDescriptor<
TileShape_MNK, EpilogueTileType, cutlass::half_t, cutlass::half_t, EpilogueSchedule
>;
using AuxLoadDescriptor = cutlass::epilogue::collective::AuxLoadDescriptor<
using AuxLoadDescriptor = cutlass::epilogue::collective::detail::AuxLoadDescriptor<
EpilogueDescriptor, cutlass::layout::ColumnMajor, cutlass::half_t
>;
@@ -182,10 +182,10 @@ TEST(SM90_Device_Gemm_f16t_f16n_f32t_tensor_op_gmma_f32_persistent_epilogue, 128
using EpilogueSchedule = cutlass::epilogue::TmaWarpSpecialized;
using EpilogueTileType = cutlass::epilogue::collective::EpilogueTileAuto;
using EpilogueDescriptor = cutlass::epilogue::collective::EpilogueDescriptor<
using EpilogueDescriptor = cutlass::epilogue::collective::detail::EpilogueDescriptor<
TileShape_MNK, EpilogueTileType, cutlass::half_t, cutlass::half_t, EpilogueSchedule
>;
using AuxLoadDescriptor = cutlass::epilogue::collective::AuxLoadDescriptor<
using AuxLoadDescriptor = cutlass::epilogue::collective::detail::AuxLoadDescriptor<
EpilogueDescriptor, cutlass::layout::ColumnMajor, float
>;
@@ -72,10 +72,10 @@ TEST(SM90_Device_Gemm_f16t_f16n_f32t_tensor_op_gmma_f32_persistent_epilogue, 128
using EpilogueSchedule = cutlass::epilogue::TmaWarpSpecialized;
using EpilogueTileType = cutlass::epilogue::collective::EpilogueTileAuto;
using EpilogueDescriptor = cutlass::epilogue::collective::EpilogueDescriptor<
using EpilogueDescriptor = cutlass::epilogue::collective::detail::EpilogueDescriptor<
TileShape_MNK, EpilogueTileType, cutlass::half_t, cutlass::half_t, EpilogueSchedule>;
using AuxLoadDescriptor = cutlass::epilogue::collective::AuxLoadDescriptor<
using AuxLoadDescriptor = cutlass::epilogue::collective::detail::AuxLoadDescriptor<
EpilogueDescriptor, cutlass::layout::RowMajor, cutlass::half_t>;
using FusionCallbacks = cutlass::epilogue::fusion::Sm90LinCombEVTDAG<
@@ -125,10 +125,10 @@ TEST(SM90_Device_Gemm_f16t_f16n_f32t_tensor_op_gmma_f32_persistent_epilogue, 128
using EpilogueSchedule = cutlass::epilogue::TmaWarpSpecialized;
using EpilogueTileType = cutlass::epilogue::collective::EpilogueTileAuto;
using EpilogueDescriptor = cutlass::epilogue::collective::EpilogueDescriptor<
using EpilogueDescriptor = cutlass::epilogue::collective::detail::EpilogueDescriptor<
TileShape_MNK, EpilogueTileType, cutlass::half_t, cutlass::half_t, EpilogueSchedule>;
using AuxStoreDescriptor = cutlass::epilogue::collective::AuxStoreDescriptor<
using AuxStoreDescriptor = cutlass::epilogue::collective::detail::AuxStoreDescriptor<
EpilogueDescriptor, cutlass::layout::RowMajor, cutlass::half_t>;
using FusionCallbacks = cutlass::epilogue::fusion::Sm90LinCombDAGEVT<
@@ -71,7 +71,7 @@ TEST(SM90_Device_Gemm_f16t_f16n_f32t_tensor_op_gmma_f32_persistent_epilogue, 128
using EpilogueSchedule = cutlass::epilogue::TmaWarpSpecialized;
using EpilogueTileType = cutlass::epilogue::collective::EpilogueTileAuto;
using EpilogueDescriptor = cutlass::epilogue::collective::EpilogueDescriptor<
using EpilogueDescriptor = cutlass::epilogue::collective::detail::EpilogueDescriptor<
TileShape_MNK, EpilogueTileType, cutlass::half_t, cutlass::half_t, EpilogueSchedule>;
using FusionCallbacks = cutlass::epilogue::fusion::Sm90LinCombPerColumnBias<
@@ -121,7 +121,7 @@ TEST(SM90_Device_Gemm_f16t_f16n_f32t_tensor_op_gmma_f32_persistent_epilogue, 128
using EpilogueSchedule = cutlass::epilogue::TmaWarpSpecialized;
using EpilogueTileType = cutlass::epilogue::collective::EpilogueTileAuto;
using EpilogueDescriptor = cutlass::epilogue::collective::EpilogueDescriptor<
using EpilogueDescriptor = cutlass::epilogue::collective::detail::EpilogueDescriptor<
TileShape_MNK, EpilogueTileType, cutlass::half_t, cutlass::half_t, EpilogueSchedule>;
using FusionCallbacks = cutlass::epilogue::fusion::Sm90LinCombPerColumnBias<
@@ -139,9 +139,9 @@ TEST(SM90_Device_Gemm_e4m3t_e4m3n_bf16n_tensor_op_gmma_f32_epilogue, 64x128x128_
using EpilogueSchedule = cutlass::epilogue::TmaWarpSpecialized;
using EpilogueTileType = cutlass::epilogue::collective::EpilogueTileAuto;
using EpilogueDescriptor = cutlass::epilogue::collective::EpilogueDescriptor<
using EpilogueDescriptor = cutlass::epilogue::collective::detail::EpilogueDescriptor<
TileShape_MNK, EpilogueTileType, cutlass::bfloat16_t, cutlass::bfloat16_t, EpilogueSchedule>;
using AuxStoreDescriptor = cutlass::epilogue::collective::AuxStoreDescriptor<
using AuxStoreDescriptor = cutlass::epilogue::collective::detail::AuxStoreDescriptor<
EpilogueDescriptor, cutlass::layout::RowMajor, cutlass::bfloat16_t>;
using FusionCallbacks = cutlass::epilogue::fusion::Sm90ScaledLinCombPerRowBiasEltActAmaxAux<
@@ -139,9 +139,9 @@ TEST(SM90_Device_Gemm_e4m3t_e4m3n_f32t_tensor_op_gmma_f32_cooperative_epilogue,
using EpilogueSchedule = cutlass::epilogue::TmaWarpSpecializedCooperative;
using EpilogueTileType = cutlass::epilogue::collective::EpilogueTileAuto;
using EpilogueDescriptor = cutlass::epilogue::collective::EpilogueDescriptor<
using EpilogueDescriptor = cutlass::epilogue::collective::detail::EpilogueDescriptor<
TileShape_MNK, EpilogueTileType, float, float, EpilogueSchedule>;
using AuxStoreDescriptor = cutlass::epilogue::collective::AuxStoreDescriptor<
using AuxStoreDescriptor = cutlass::epilogue::collective::detail::AuxStoreDescriptor<
EpilogueDescriptor, cutlass::layout::RowMajor, float>;
using FusionCallbacks = cutlass::epilogue::fusion::Sm90ScaledLinCombPerRowBiasEltActAmaxAux<
@@ -139,9 +139,9 @@ TEST(SM90_Device_Gemm_f8t_f8n_f8t_tensor_op_gmma_f32_persistent_epilogue, 64x128
using EpilogueSchedule = cutlass::epilogue::TmaWarpSpecialized;
using EpilogueTileType = cutlass::epilogue::collective::EpilogueTileAuto;
using EpilogueDescriptor = cutlass::epilogue::collective::EpilogueDescriptor<
using EpilogueDescriptor = cutlass::epilogue::collective::detail::EpilogueDescriptor<
TileShape_MNK, EpilogueTileType, cutlass::float_e4m3_t, cutlass::float_e4m3_t, EpilogueSchedule>;
using AuxStoreDescriptor = cutlass::epilogue::collective::AuxStoreDescriptor<
using AuxStoreDescriptor = cutlass::epilogue::collective::detail::AuxStoreDescriptor<
EpilogueDescriptor, cutlass::layout::RowMajor, cutlass::float_e4m3_t>;
using FusionCallbacks = cutlass::epilogue::fusion::Sm90ScaledLinCombPerRowBiasEltActAmaxAux<
@@ -569,7 +569,6 @@ TEST(SM75_gemm_threadblock_crosswise,
}
////////////////////////////////////////////////////////////////////////////////
TEST(SM75_gemm_threadblock_interleaved, tensor_op_32x32x64_16x16x64_8x8x16) {
using ElementA = uint8_t;
using LayoutA = cutlass::layout::ColumnMajorInterleaved<32>;
@@ -1793,6 +1792,7 @@ TEST(SM75_gemm_threadblock_interleaved,
}
////////////////////////////////////////////////////////////////////////////////
TEST(SM75_gemm_threadblock_crosswise, tensor_op_64x64x512_64x64x512_8x8x128) {
using ElementA = cutlass::uint1b_t;
using LayoutA = cutlass::layout::RowMajor;
@@ -193,7 +193,6 @@ TEST(SM75_gemm_threadblock_wmma_tensor_op_col_row_row_s8, 64x64x64_64x64x64_16x1
///////////////////////////////////////////////////////////////////////
#if defined(CUTLASS_SUBBYTE_INTEGER_MATRIX_MULTIPLY_ENABLED)
TEST(SM75_gemm_threadblock_wmma_tensor_op_row_col_row_s4, 64x64x128_64x64x128_8x8x32) {
using ElementA = cutlass::int4b_t;
using LayoutA = cutlass::layout::RowMajor;
@@ -262,6 +261,7 @@ TEST(SM75_gemm_threadblock_wmma_tensor_op_row_col_col_s4, 64x64x64_64x64x64_8x8x
problem_size.k(), alpha, beta)
.run(grid, block);
}
TEST(SM75_gemm_threadblock_wmma_tensor_op_row_col_row_b1, 64x64x512_64x64x512_8x8x128) {
using ElementA = cutlass::uint1b_t;
using LayoutA = cutlass::layout::RowMajor;
@@ -193,7 +193,6 @@ TEST(SM75_gemm_threadblock_singlestage_wmma_tensor_op_col_row_row_s8, 64x64x64_6
///////////////////////////////////////////////////////////////////////
#if defined(CUTLASS_SUBBYTE_INTEGER_MATRIX_MULTIPLY_ENABLED)
TEST(SM75_gemm_threadblock_singlestage_wmma_tensor_op_row_col_row_s4, 64x64x128_64x64x128_8x8x32) {
using ElementA = cutlass::int4b_t;
using LayoutA = cutlass::layout::RowMajor;
@@ -262,6 +261,7 @@ TEST(SM75_gemm_threadblock_singlestage_wmma_tensor_op_row_col_col_s4, 64x64x64_6
problem_size.k(), alpha, beta)
.run(grid, block);
}
TEST(SM75_gemm_threadblock_singlestage_wmma_tensor_op_row_col_row_b1, 64x64x512_64x64x512_8x8x128) {
using ElementA = cutlass::uint1b_t;
using LayoutA = cutlass::layout::RowMajor;
+1 -1
View File
@@ -326,7 +326,6 @@ TEST(SM75_warp_gemm_tensor_op_crosswise_f16, 128x128x64_16x16x64_16x8x8) {
}
////////////////////////////////////////////////////////////////////////////////
TEST(SM75_warp_gemm_tensor_op_crosswise_i8, 128x128x64_64x64x64_8x8x16) {
using Shape = cutlass::gemm::GemmShape<64, 64, 64>;
using InstructionShape = cutlass::gemm::GemmShape<8, 8, 16>;
@@ -746,6 +745,7 @@ TEST(SM75_warp_gemm_tensor_op_interleaved_i4, 128x128x128_16x16x128_8x8x32) {
}
////////////////////////////////////////////////////////////////////////////////
TEST(SM75_warp_gemm_tensor_op_crosswise_b1, 128x128x512_64x64x512_8x8x128) {
using Shape = cutlass::gemm::GemmShape<64, 64, 512>;
using InstructionShape = cutlass::gemm::GemmShape<8, 8, 128>;