v3.9 (#2185)
* v3.8 update x * fix blackwell gg * doc change * doc change * doc change --------- Co-authored-by: yuzhai <yuzhai@nvidia.com> Co-authored-by: Haicheng Wu <haichengw@nvidia.com> Co-authored-by: Haicheng Wu <57973641+hwu36@users.noreply.github.com>
This commit is contained in:
co-authored by
yuzhai
Haicheng Wu
Haicheng Wu
parent
8c4d1dc47d
commit
62750a2b75
@@ -29,6 +29,8 @@
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
//#define CUTLASS_DEBUG_TRACE_LEVEL 1
|
||||
|
||||
#include "cutlass_unit_test.h"
|
||||
|
||||
#include <cutlass/trace.h>
|
||||
@@ -68,6 +70,9 @@ test_complement(Layout const& layout, CoTarget const& cotarget)
|
||||
if constexpr (is_static<decltype(stride(completed))>::value) { // If we can apply complement again
|
||||
EXPECT_EQ(size(complement(completed)), 1); // There's no more codomain left over
|
||||
}
|
||||
if constexpr (is_static<decltype(result)>::value && is_static<decltype(layout)>::value) {
|
||||
EXPECT_TRUE(bool(complement(complement(result,cosize(layout)),cotarget) == result));
|
||||
}
|
||||
}
|
||||
|
||||
template <class Layout>
|
||||
|
||||
@@ -29,6 +29,10 @@
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
//#define CUTLASS_DEBUG_TRACE_LEVEL 1
|
||||
|
||||
#include "cutlass_unit_test.h"
|
||||
|
||||
#include <cutlass/trace.h>
|
||||
|
||||
#include <cute/layout.hpp>
|
||||
@@ -38,8 +42,6 @@
|
||||
#include <cute/tensor.hpp>
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass_unit_test.h"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
|
||||
@@ -273,14 +275,14 @@ TEST(CuTe_core, Composition)
|
||||
|
||||
{
|
||||
auto a = make_layout(Shape<_8,_8>{});
|
||||
auto b = make_layout(Shape<Shape<_2, _2,_2>, Shape<_2,_2, _2>>{},
|
||||
auto b = make_layout(Shape <Shape <_2, _2,_2>, Shape <_2,_2, _2>>{},
|
||||
Stride<Stride<_1,_16,_4>, Stride<_8,_2,_32>>{});
|
||||
test_composition(a, b);
|
||||
}
|
||||
|
||||
{
|
||||
auto a = make_layout(Shape<_8,_8>{}, Stride<_8,_1>{});
|
||||
auto b = make_layout(Shape<Shape<_2, _2,_2>, Shape<_2,_2, _2>>{},
|
||||
auto b = make_layout(Shape <Shape <_2, _2,_2>, Shape <_2,_2, _2>>{},
|
||||
Stride<Stride<_1,_16,_4>, Stride<_8,_2,_32>>{});
|
||||
|
||||
test_composition(a, b);
|
||||
@@ -423,7 +425,7 @@ TEST(CuTe_core, Composition)
|
||||
test_composition(a, b);
|
||||
}
|
||||
|
||||
// Capping a Layout with 1:0 forces divisibility and extends in stride-0
|
||||
// Capping a Layout with 1:0 extends in stride-0
|
||||
{
|
||||
auto a = make_layout(Shape<_4,_3,_1>{}, Stride<_3,_1,_0>{});
|
||||
auto b = make_layout(Shape<_24>{});
|
||||
@@ -431,6 +433,50 @@ TEST(CuTe_core, Composition)
|
||||
test_composition(a, b);
|
||||
}
|
||||
|
||||
{
|
||||
auto a = make_layout(Shape<_4,_3,_1>{}, Stride<_3,_1,_0>{});
|
||||
auto b = make_layout(Shape<_4>{});
|
||||
|
||||
test_composition(a, b);
|
||||
}
|
||||
|
||||
// Pre-coalesced LHS
|
||||
{
|
||||
auto a = make_layout(Shape<_4,_6,_8>{}, Stride<_1,_4,_7>{});
|
||||
auto b = make_layout(_6{}, _1{});
|
||||
|
||||
test_composition(a, b);
|
||||
}
|
||||
|
||||
// Mid-layout truncation
|
||||
{
|
||||
auto a = make_layout(Shape<_4,_6,_8,_10>{}, Stride<_2,_3,_5,_7>{});
|
||||
auto b = make_layout(_6{}, _12{});
|
||||
|
||||
test_composition(a, b);
|
||||
}
|
||||
|
||||
{
|
||||
auto a = make_layout(Shape<_8,_8>{}, Stride<_8,_1>{});
|
||||
auto b = make_layout(_2{}, _3{});
|
||||
|
||||
test_composition(a, b);
|
||||
}
|
||||
|
||||
{
|
||||
auto a = make_layout(Shape<_8,_8>{}, Stride<_8,_1>{});
|
||||
auto b = make_layout(_3{}, _3{});
|
||||
|
||||
test_composition(a, b);
|
||||
}
|
||||
|
||||
// Should fail to a static divisibility condition
|
||||
// {
|
||||
// auto a = make_layout(Shape<_8,_8>{}, Stride<_8,_1>{});
|
||||
// auto b = make_layout(_4{}, _3{});
|
||||
// test_composition(a, b);
|
||||
// }
|
||||
|
||||
{
|
||||
auto a = make_layout(3, _1{});
|
||||
auto b = make_layout(_4{}, _1{});
|
||||
|
||||
@@ -29,6 +29,8 @@
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
//#define CUTLASS_DEBUG_TRACE_LEVEL 1
|
||||
|
||||
#include "cutlass_unit_test.h"
|
||||
|
||||
#include <cutlass/trace.h>
|
||||
@@ -50,18 +52,31 @@ test_left_inverse(Layout const& layout)
|
||||
CUTLASS_TRACE_HOST(layout << " ^ -1\n" << " => \n" << inv_layout);
|
||||
|
||||
for (int i = 0; i < size(layout); ++i) {
|
||||
//printf("%3d: %3d %3d\n", i, int(layout(i)), int(inv_layout(layout(i))));
|
||||
EXPECT_EQ(inv_layout(layout(i)), i);
|
||||
//printf("%3d: %3d %3d %3d\n", i, int(layout(i)), int(inv_layout(layout(i))), int(layout(inv_layout(layout(i)))));
|
||||
EXPECT_EQ(layout(inv_layout(layout(i))), layout(i));
|
||||
}
|
||||
|
||||
CUTLASS_TRACE_HOST("Composition: " << coalesce(composition(inv_layout, layout)));
|
||||
CUTLASS_TRACE_HOST("Composition: " << coalesce(composition(layout, composition(inv_layout, layout))));
|
||||
}
|
||||
|
||||
TEST(CuTe_core, Inverse_left)
|
||||
{
|
||||
CUTLASS_TRACE_HOST("-------------------------------");
|
||||
CUTLASS_TRACE_HOST("LEFT INVERSE" );
|
||||
CUTLASS_TRACE_HOST("-------------------------------");
|
||||
|
||||
CUTLASS_TRACE_HOST("-------------------------------");
|
||||
CUTLASS_TRACE_HOST("Simple tests" );
|
||||
CUTLASS_TRACE_HOST("-------------------------------");
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape <_1>,
|
||||
Stride<_0>>{};
|
||||
auto layout = Layout<_1, _0>{};
|
||||
|
||||
test_left_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<_1, _1>{};
|
||||
|
||||
test_left_inverse(layout);
|
||||
}
|
||||
@@ -74,8 +89,15 @@ TEST(CuTe_core, Inverse_left)
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape <_1>,
|
||||
Stride<_1>>{};
|
||||
auto layout = Layout<Shape <Shape <_3,_7>>,
|
||||
Stride<Stride<_0,_0>>>{};
|
||||
|
||||
test_left_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape <_4>,
|
||||
Stride<_0>>{};
|
||||
|
||||
test_left_inverse(layout);
|
||||
}
|
||||
@@ -94,6 +116,13 @@ TEST(CuTe_core, Inverse_left)
|
||||
test_left_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape <_2,_4>,
|
||||
Stride<_0,_2>>{};
|
||||
|
||||
test_left_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape <_8, _4>>{};
|
||||
|
||||
@@ -120,6 +149,13 @@ TEST(CuTe_core, Inverse_left)
|
||||
test_left_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape <_2,_4,_4,_6>,
|
||||
Stride<_4,_1,_0,_8>>{};
|
||||
|
||||
test_left_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape <_4, _2>,
|
||||
Stride<_1,_16>>{};
|
||||
@@ -127,9 +163,105 @@ TEST(CuTe_core, Inverse_left)
|
||||
test_left_inverse(layout);
|
||||
}
|
||||
|
||||
//
|
||||
// Swizzle left_inverse
|
||||
//
|
||||
{
|
||||
auto layout = Layout<Shape <_4, _2>,
|
||||
Stride<_1, _5>>{};
|
||||
|
||||
test_left_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape<_128,_128>,Stride<_65536,_1>>{};
|
||||
|
||||
test_left_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape<_128,_160>,Stride<_65536,_1>>{};
|
||||
|
||||
test_left_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape<_128,_3,_160>,Stride<_65536,_512,_1>>{};
|
||||
|
||||
test_left_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape<_128, _64>, Stride<Int<131072>, Int<2>>>{};
|
||||
|
||||
test_left_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape<_32,_4,_4,_4>, Stride<_262144,_4,Int<8388608>,_1>>{};
|
||||
|
||||
test_left_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape<_2,_2,_2>, Stride<_4,_0,_1>>{};
|
||||
|
||||
test_left_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape <Shape <Shape <Shape <Shape < _32, _4>, _1>, Shape < _32, _2>>, _4>, _1, Shape <_2, _2>, _2>,
|
||||
Stride<Stride<Stride<Stride<Stride<C<262144>, _4>, _0>, Stride<C<0>, C<1>>>, C<8388608>>, _0, Stride<_2, _16>, _32>>{};
|
||||
|
||||
test_left_inverse(layout);
|
||||
}
|
||||
|
||||
// CUTLASS_TRACE_HOST("-------------------------------");
|
||||
// CUTLASS_TRACE_HOST("Dynamic shapes/strides" );
|
||||
// CUTLASS_TRACE_HOST("-------------------------------");
|
||||
|
||||
// {
|
||||
// auto layout = make_layout(Shape<_4, _2>{}, make_stride(Int<1>{}, 4));
|
||||
|
||||
// test_left_inverse(layout);
|
||||
// }
|
||||
|
||||
// {
|
||||
// auto layout = make_layout(make_shape(_4{}, 2), make_stride(Int<1>{}, 4));
|
||||
|
||||
// test_left_inverse(layout);
|
||||
// }
|
||||
|
||||
// {
|
||||
// auto layout = make_layout(make_shape(4, 2), make_stride(Int<1>{}, 4));
|
||||
|
||||
// test_left_inverse(layout);
|
||||
// }
|
||||
|
||||
// {
|
||||
// auto layout = make_layout(Shape<_2, _4>{}, make_stride(4, Int<1>{}));
|
||||
|
||||
// test_left_inverse(layout);
|
||||
// }
|
||||
|
||||
// {
|
||||
// auto layout = make_layout(make_shape(2, Int<4>{}), make_stride(4, Int<1>{}));
|
||||
|
||||
// test_left_inverse(layout);
|
||||
// }
|
||||
|
||||
// {
|
||||
// auto layout = make_layout(make_shape(2, 4), make_stride(4, Int<1>{}));
|
||||
|
||||
// test_left_inverse(layout);
|
||||
// }
|
||||
|
||||
// {
|
||||
// auto layout = make_layout(make_shape(2, 4), make_stride(4, 1));
|
||||
|
||||
// test_left_inverse(layout);
|
||||
// }
|
||||
|
||||
CUTLASS_TRACE_HOST("-------------------------------");
|
||||
CUTLASS_TRACE_HOST("Swizzle layouts" );
|
||||
CUTLASS_TRACE_HOST("-------------------------------");
|
||||
|
||||
{
|
||||
auto layout = ComposedLayout<Swizzle<1,0,2>, _0, Layout<Shape <_4, _4>,
|
||||
@@ -152,10 +284,10 @@ TEST(CuTe_core, Inverse_left)
|
||||
test_left_inverse(layout);
|
||||
}
|
||||
|
||||
//
|
||||
// Negative strides (beta support)
|
||||
// Post-conditions/layout indexing aren't generalized enough to support these yet
|
||||
// However, the composition post-condition is general enough.
|
||||
CUTLASS_TRACE_HOST("-------------------------------");
|
||||
CUTLASS_TRACE_HOST("BETA: Negative strides" );
|
||||
CUTLASS_TRACE_HOST("-------------------------------");
|
||||
|
||||
{
|
||||
auto layout = make_layout(Shape<_4>{}, Stride<Int<-1>>{});
|
||||
|
||||
|
||||
@@ -29,6 +29,8 @@
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
//#define CUTLASS_DEBUG_TRACE_LEVEL 1
|
||||
|
||||
#include "cutlass_unit_test.h"
|
||||
|
||||
#include <cutlass/trace.h>
|
||||
@@ -41,16 +43,6 @@
|
||||
|
||||
using namespace cute;
|
||||
|
||||
template <class Layout, class InvLayout>
|
||||
void
|
||||
test_postconditions(Layout const& layout, InvLayout const& inv_layout)
|
||||
{
|
||||
for (int i = 0; i < size(inv_layout); ++i) {
|
||||
//printf("%3d: %3d %3d\n", i, int(inv_layout(i)), int(layout(inv_layout(i))));
|
||||
EXPECT_EQ(layout(inv_layout(i)), i);
|
||||
}
|
||||
}
|
||||
|
||||
template <class Layout>
|
||||
void
|
||||
test_right_inverse(Layout const& layout)
|
||||
@@ -58,9 +50,13 @@ test_right_inverse(Layout const& layout)
|
||||
auto inv_layout = right_inverse(layout);
|
||||
|
||||
CUTLASS_TRACE_HOST(layout << " ^ -1\n" << " => \n" << inv_layout);
|
||||
CUTLASS_TRACE_HOST("Composition: " << coalesce(composition(layout, inv_layout)) << std::endl);
|
||||
|
||||
test_postconditions(layout, inv_layout);
|
||||
for (int i = 0; i < size(inv_layout); ++i) {
|
||||
//printf("%3d: %3d %3d\n", i, int(inv_layout(i)), int(layout(inv_layout(i))));
|
||||
EXPECT_EQ(layout(inv_layout(i)), i);
|
||||
}
|
||||
|
||||
CUTLASS_TRACE_HOST("Composition: " << coalesce(composition(layout, inv_layout)) << std::endl);
|
||||
}
|
||||
|
||||
TEST(CuTe_core, Inverse_right)
|
||||
@@ -85,13 +81,6 @@ TEST(CuTe_core, Inverse_right)
|
||||
test_right_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape <_4>,
|
||||
Stride<_0>>{};
|
||||
|
||||
test_right_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape <Shape <_1,_1>>,
|
||||
Stride<Stride<_0,_0>>>{};
|
||||
@@ -107,8 +96,8 @@ TEST(CuTe_core, Inverse_right)
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape <_1>,
|
||||
Stride<_1>>{};
|
||||
auto layout = Layout<Shape <_4>,
|
||||
Stride<_0>>{};
|
||||
|
||||
test_right_inverse(layout);
|
||||
}
|
||||
@@ -181,6 +170,49 @@ TEST(CuTe_core, Inverse_right)
|
||||
test_right_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape<_128,_128>,Stride<_65536,_1>>{};
|
||||
|
||||
test_right_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape<_128,_160>,Stride<_65536,_1>>{};
|
||||
|
||||
test_right_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape<_128,_3,_160>,Stride<_65536,_512,_1>>{};
|
||||
|
||||
test_right_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape<_128, _64>, Stride<Int<131072>, Int<2>>>{};
|
||||
|
||||
test_right_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape<_32,_4,_4,_4>, Stride<_262144,_4,Int<8388608>,_1>>{};
|
||||
|
||||
test_right_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape<_2,_2,_2>, Stride<_4,_0,_1>>{};
|
||||
|
||||
test_right_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape <Shape <Shape <Shape <Shape < _32, _4>, _1>, Shape < _32, _2>>, _4>, _1, Shape <_2, _2>, _2>,
|
||||
Stride<Stride<Stride<Stride<Stride<C<262144>, _4>, _0>, Stride<C<0>, C<1>>>, C<8388608>>, _0, Stride<_2, _16>, _32>>{};
|
||||
|
||||
test_right_inverse(layout);
|
||||
}
|
||||
|
||||
CUTLASS_TRACE_HOST("-------------------------------");
|
||||
CUTLASS_TRACE_HOST("Dynamic shapes/strides" );
|
||||
CUTLASS_TRACE_HOST("-------------------------------");
|
||||
@@ -203,6 +235,24 @@ TEST(CuTe_core, Inverse_right)
|
||||
test_right_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = make_layout(Shape<_2, _4>{}, make_stride(4, Int<1>{}));
|
||||
|
||||
test_right_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = make_layout(make_shape(2, Int<4>{}), make_stride(4, Int<1>{}));
|
||||
|
||||
test_right_inverse(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = make_layout(make_shape(2, 4), make_stride(4, Int<1>{}));
|
||||
|
||||
test_right_inverse(layout);
|
||||
}
|
||||
|
||||
CUTLASS_TRACE_HOST("-------------------------------");
|
||||
CUTLASS_TRACE_HOST("Swizzle layouts" );
|
||||
CUTLASS_TRACE_HOST("-------------------------------");
|
||||
|
||||
@@ -29,6 +29,8 @@
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
//#define CUTLASS_DEBUG_TRACE_LEVEL 1
|
||||
|
||||
#include "cutlass_unit_test.h"
|
||||
|
||||
#include <cutlass/trace.h>
|
||||
|
||||
@@ -131,7 +131,7 @@ tma_test_device_cute(T const* g_in, T* g_out,
|
||||
for (int stage = 0; stage < size<1>(tAgA); ++stage)
|
||||
{
|
||||
// Set the bytes transferred in this TMA transaction (may involve multiple issues)
|
||||
constexpr int kTmaTransactionBytes = sizeof(ArrayEngine<T, CUTE_STATIC_V(size(filter_zeros(sA)))>);
|
||||
constexpr int kTmaTransactionBytes = sizeof(make_tensor_like(tensor<0>(tAsA)));
|
||||
|
||||
if (threadIdx.x == 0)
|
||||
{
|
||||
|
||||
@@ -53,6 +53,13 @@ endfunction()
|
||||
|
||||
add_subdirectory(sm100_blockscaled_tensorop_gemm)
|
||||
add_subdirectory(sm100_tensorop_gemm)
|
||||
add_subdirectory(sm100_blockscaled_sparse_tensorop_gemm)
|
||||
add_subdirectory(sm100_sparse_tensorop_gemm)
|
||||
add_subdirectory(sm120_blockscaled_sparse_tensorop_gemm)
|
||||
add_subdirectory(sm120_sparse_tensorop_gemm)
|
||||
add_subdirectory(sm120_tensorop_gemm)
|
||||
add_subdirectory(sm120_blockscaled_tensorop_gemm)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable(
|
||||
cutlass_test_unit_gemm_device_simt
|
||||
|
||||
@@ -330,6 +337,11 @@ cutlass_test_unit_gemm_device_add_executable(
|
||||
cutlass_test_unit_gemm_device_add_executable(
|
||||
cutlass_test_unit_gemm_device_tensorop_sm90_group_gemm
|
||||
sm90_gemm_f16_f16_f16_tensor_op_f32_group_gemm.cu
|
||||
)
|
||||
|
||||
# Group Gemm pingpong test
|
||||
cutlass_test_unit_gemm_device_add_executable(
|
||||
cutlass_test_unit_gemm_device_tensorop_sm90_group_gemm_pingpong
|
||||
sm90_gemm_f16_f16_f16_tensor_op_f32_group_gemm_pingpong.cu
|
||||
)
|
||||
|
||||
|
||||
@@ -892,6 +892,83 @@ struct HostCollectiveMainloop<ScheduleType_, Gemm, ElementA_, ElementB_,
|
||||
using HostCollectiveMainloopSparse<Gemm, ElementA_, ElementB_>::HostCollectiveMainloopSparse;
|
||||
};
|
||||
|
||||
//
|
||||
// Sparse MMA input Operands : A_compressed, B, metadata
|
||||
//
|
||||
// Structured Sparse Gemm Input Operands
|
||||
|
||||
template<
|
||||
class Gemm,
|
||||
int SchedulerPipelineStageCount_,
|
||||
int AccumulatorPipelineStageCount_,
|
||||
typename ElementA_,
|
||||
typename ElementB_
|
||||
>
|
||||
struct HostCollectiveMainloop<cutlass::gemm::KernelSparseTmaWarpSpecializedSm100<SchedulerPipelineStageCount_,
|
||||
AccumulatorPipelineStageCount_>,
|
||||
Gemm, ElementA_, ElementB_>
|
||||
: HostCollectiveMainloopSparse<Gemm, ElementA_, ElementB_>
|
||||
{
|
||||
using HostCollectiveMainloopSparse<Gemm, ElementA_, ElementB_>::HostCollectiveMainloopSparse;
|
||||
};
|
||||
|
||||
//
|
||||
// Sparse Gemm Input Operands : A , B, E
|
||||
//
|
||||
template<
|
||||
class Gemm,
|
||||
int SchedulerPipelineStageCount_,
|
||||
class ElementA_,
|
||||
class ElementB_
|
||||
>
|
||||
struct HostCollectiveMainloop<cutlass::gemm::KernelTmaWarpSpecializedCooperativeSparseSm120<SchedulerPipelineStageCount_, false /*isAsymmetric*/>,
|
||||
Gemm, ElementA_, ElementB_> : public
|
||||
HostCollectiveMainloop<cutlass::gemm::KernelSparseTmaWarpSpecializedSm100<0/*SchedulerPipelineStageCount_*/,
|
||||
0/*AccumulatorPipelineStageCount_*/>,
|
||||
Gemm, ElementA_, ElementB_> {
|
||||
using Base = HostCollectiveMainloop<cutlass::gemm::KernelSparseTmaWarpSpecializedSm100<0,0>,
|
||||
Gemm, ElementA_, ElementB_ >;
|
||||
HostCollectiveMainloop(
|
||||
CheckEquality check_relative_equality_ = CheckEquality::EXACT,
|
||||
cutlass::Distribution::Kind init_A_ = cutlass::Distribution::Uniform,
|
||||
cutlass::Distribution::Kind init_B_ = cutlass::Distribution::Uniform,
|
||||
uint64_t seed_ = Base::kDefaultSeed,
|
||||
typename Base::LayoutTagA::Stride stride_factor_A_ = typename Base::LayoutTagA::Stride(),
|
||||
typename Base::LayoutTagB::Stride stride_factor_B_ = typename Base::LayoutTagB::Stride(),
|
||||
typename Base::LayoutTagE::Stride stride_factor_E_ = typename Base::LayoutTagE::Stride()
|
||||
) : Base::HostCollectiveMainloop(check_relative_equality_, init_A_, init_B_, seed_, stride_factor_A_,
|
||||
stride_factor_B_,
|
||||
stride_factor_E_) {}
|
||||
};
|
||||
|
||||
//
|
||||
// Sparse Gemm Input Operands : A , B, E
|
||||
//
|
||||
template<
|
||||
class Gemm,
|
||||
int SchedulerPipelineStageCount_,
|
||||
class ElementA_,
|
||||
class ElementB_
|
||||
>
|
||||
struct HostCollectiveMainloop<cutlass::gemm::KernelTmaWarpSpecializedCooperativeSparseSm120<SchedulerPipelineStageCount_, true /*isAsymmetric*/>,
|
||||
Gemm, ElementA_, ElementB_> : public
|
||||
HostCollectiveMainloop<cutlass::gemm::KernelSparseTmaWarpSpecializedSm100<0/*SchedulerPipelineStageCount_*/,
|
||||
0/*AccumulatorPipelineStageCount_*/>,
|
||||
Gemm, ElementA_, ElementB_> {
|
||||
using Base = HostCollectiveMainloop<cutlass::gemm::KernelSparseTmaWarpSpecializedSm100<0,0>,
|
||||
Gemm, ElementA_, ElementB_ >;
|
||||
HostCollectiveMainloop(
|
||||
CheckEquality check_relative_equality_ = CheckEquality::EXACT,
|
||||
cutlass::Distribution::Kind init_A_ = cutlass::Distribution::Uniform,
|
||||
cutlass::Distribution::Kind init_B_ = cutlass::Distribution::Uniform,
|
||||
uint64_t seed_ = Base::kDefaultSeed,
|
||||
typename Base::LayoutTagA::Stride stride_factor_A_ = typename Base::LayoutTagA::Stride(),
|
||||
typename Base::LayoutTagB::Stride stride_factor_B_ = typename Base::LayoutTagB::Stride(),
|
||||
typename Base::LayoutTagE::Stride stride_factor_E_ = typename Base::LayoutTagE::Stride()
|
||||
) : Base::HostCollectiveMainloop(check_relative_equality_, init_A_, init_B_, seed_, stride_factor_A_,
|
||||
stride_factor_B_,
|
||||
stride_factor_E_) {}
|
||||
};
|
||||
|
||||
//
|
||||
// Block Scaled Gemm Input Operands : A , B, scalefactorA, scalefactorB
|
||||
@@ -923,10 +1000,10 @@ struct HostCollectiveMainloop<cutlass::gemm::KernelTmaWarpSpecializedBlockScaled
|
||||
static constexpr int SFVecSize = Gemm::GemmKernel::CollectiveMainloop::SFVecSize;
|
||||
|
||||
using ElementSF = typename Gemm::GemmKernel::CollectiveMainloop::ElementSF;
|
||||
using Sm100BlkScaledConfig = typename Gemm::GemmKernel::CollectiveMainloop::Sm100BlkScaledConfig;
|
||||
using Blk_MN = typename Sm100BlkScaledConfig::Blk_MN;
|
||||
using Blk_SF = typename Sm100BlkScaledConfig::Blk_SF;
|
||||
using SfAtom = typename Sm100BlkScaledConfig::SfAtom;
|
||||
using Sm1xxBlkScaledConfig = typename Gemm::GemmKernel::CollectiveMainloop::Sm1xxBlkScaledConfig;
|
||||
using Blk_MN = typename Sm1xxBlkScaledConfig::Blk_MN;
|
||||
using Blk_SF = typename Sm1xxBlkScaledConfig::Blk_SF;
|
||||
using SfAtom = typename Sm1xxBlkScaledConfig::SfAtom;
|
||||
using LayoutSFA = typename Gemm::GemmKernel::CollectiveMainloop::LayoutSFA;
|
||||
using LayoutSFB = typename Gemm::GemmKernel::CollectiveMainloop::LayoutSFB;
|
||||
|
||||
@@ -1015,8 +1092,8 @@ struct HostCollectiveMainloop<cutlass::gemm::KernelTmaWarpSpecializedBlockScaled
|
||||
auto k_blks = cutlass::ceil_div(K, size<1>(shape(SfAtom{})));
|
||||
auto m_blks = cutlass::ceil_div(M, Blk_MN{});
|
||||
auto n_blks = cutlass::ceil_div(N, Blk_MN{});
|
||||
layout_sfa = Sm100BlkScaledConfig::tile_atom_to_shape_SFA(problem_shape_MNKL);
|
||||
layout_sfb = Sm100BlkScaledConfig::tile_atom_to_shape_SFB(problem_shape_MNKL);
|
||||
layout_sfa = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(problem_shape_MNKL);
|
||||
layout_sfb = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(problem_shape_MNKL);
|
||||
|
||||
// 2.x host tensor does not natively contain a batch stride or coord, so we spoof if by folding it into the outer mode
|
||||
auto sfa_coord = cutlass::make_Coord(m_blks * Blk_MN{} * L, k_blks * Blk_SF{});
|
||||
@@ -1098,6 +1175,413 @@ struct HostCollectiveMainloop<cutlass::gemm::KernelTmaWarpSpecializedBlockScaled
|
||||
};
|
||||
|
||||
|
||||
//
|
||||
// Block Scaled Gemm Input Operands : A , B, scalefactorA, scalefactorB
|
||||
//
|
||||
template<
|
||||
class Gemm,
|
||||
int SchedulerPipelineStageCount_,
|
||||
class ElementA_,
|
||||
class ElementB_
|
||||
>
|
||||
struct HostCollectiveMainloop<cutlass::gemm::KernelTmaWarpSpecializedPingpongBlockScaledSm120<SchedulerPipelineStageCount_>,
|
||||
Gemm, ElementA_, ElementB_> : public
|
||||
HostCollectiveMainloop<cutlass::gemm::KernelTmaWarpSpecializedBlockScaledSm100<0,0>,
|
||||
Gemm, ElementA_, ElementB_> {
|
||||
using Base = HostCollectiveMainloop<cutlass::gemm::KernelTmaWarpSpecializedBlockScaledSm100<0,0>,
|
||||
Gemm, ElementA_, ElementB_>;
|
||||
HostCollectiveMainloop(
|
||||
CheckEquality check_relative_equality_ = CheckEquality::EXACT,
|
||||
cutlass::Distribution::Kind init_A_ = cutlass::Distribution::Uniform,
|
||||
cutlass::Distribution::Kind init_B_ = cutlass::Distribution::Uniform,
|
||||
uint64_t seed_ = Base::kDefaultSeed,
|
||||
typename Base::LayoutTagA::Stride stride_factor_A_ = typename Base::LayoutTagA::Stride(),
|
||||
typename Base::LayoutTagB::Stride stride_factor_B_ = typename Base::LayoutTagB::Stride()
|
||||
) : Base::HostCollectiveMainloop(check_relative_equality_, init_A_, init_B_, seed_, stride_factor_A_, stride_factor_B_) {}
|
||||
};
|
||||
|
||||
//
|
||||
// Block Scaled Gemm Input Operands : A , B, scalefactorA, scalefactorB
|
||||
//
|
||||
template<
|
||||
class Gemm,
|
||||
int SchedulerPipelineStageCount_,
|
||||
class ElementA_,
|
||||
class ElementB_
|
||||
>
|
||||
struct HostCollectiveMainloop<cutlass::gemm::KernelTmaWarpSpecializedCooperativeBlockScaledSm120<SchedulerPipelineStageCount_>,
|
||||
Gemm, ElementA_, ElementB_> : public
|
||||
HostCollectiveMainloop<cutlass::gemm::KernelTmaWarpSpecializedBlockScaledSm100<0,0>,
|
||||
Gemm, ElementA_, ElementB_> {
|
||||
using Base = HostCollectiveMainloop<cutlass::gemm::KernelTmaWarpSpecializedBlockScaledSm100<0,0>,
|
||||
Gemm, ElementA_, ElementB_>;
|
||||
HostCollectiveMainloop(
|
||||
CheckEquality check_relative_equality_ = CheckEquality::EXACT,
|
||||
cutlass::Distribution::Kind init_A_ = cutlass::Distribution::Uniform,
|
||||
cutlass::Distribution::Kind init_B_ = cutlass::Distribution::Uniform,
|
||||
uint64_t seed_ = Base::kDefaultSeed,
|
||||
typename Base::LayoutTagA::Stride stride_factor_A_ = typename Base::LayoutTagA::Stride(),
|
||||
typename Base::LayoutTagB::Stride stride_factor_B_ = typename Base::LayoutTagB::Stride()
|
||||
) : Base::HostCollectiveMainloop(check_relative_equality_, init_A_, init_B_, seed_, stride_factor_A_, stride_factor_B_) {}
|
||||
};
|
||||
|
||||
//
|
||||
// Block Scaled Structured Sparse Gemm Input Operands : A_compressed, B, metadata, scalefactorA, scalefactorB
|
||||
//
|
||||
template<
|
||||
class Gemm,
|
||||
int SchedulerPipelineStageCount_,
|
||||
int AccumulatorPipelineStageCount_,
|
||||
typename ElementA_,
|
||||
typename ElementB_
|
||||
>
|
||||
struct HostCollectiveMainloop<cutlass::gemm::KernelSparseTmaWarpSpecializedBlockScaledSm100<SchedulerPipelineStageCount_,
|
||||
AccumulatorPipelineStageCount_>,
|
||||
Gemm, ElementA_, ElementB_> {
|
||||
// Kernel data types
|
||||
using ElementA = ElementA_;
|
||||
// CuTe layout A for the kernel's sparse tensorA.
|
||||
using LayoutA = typename Gemm::GemmKernel::CollectiveMainloop::LayoutA;
|
||||
using ElementB = ElementB_;
|
||||
using StrideB = typename Gemm::GemmKernel::StrideB;
|
||||
using ScheduleType = typename Gemm::GemmKernel::CollectiveMainloop::DispatchPolicy::Schedule;
|
||||
|
||||
using ElementE = typename Gemm::GemmKernel::CollectiveMainloop::ElementE;
|
||||
// CuTe layout E for the kernel's metadata tensor.
|
||||
using LayoutE = typename Gemm::GemmKernel::CollectiveMainloop::LayoutE;
|
||||
using ElementAccumulator = typename Gemm::GemmKernel::ElementAccumulator;
|
||||
using ElementScalingFactor = ElementAccumulator;
|
||||
using ProblemShapeType = typename Gemm::GemmKernel::ProblemShape;
|
||||
using EpilogueOutputOp = typename Gemm::EpilogueOutputOp;
|
||||
using SparseConfig = typename Gemm::GemmKernel::CollectiveMainloop::SparseConfig;
|
||||
|
||||
// The following typenames are for the reference host tensors. They are non-sparse tensors.
|
||||
using LayoutTagA = decltype(SparseConfig::deduce_layoutA_tag(LayoutA{}));
|
||||
using StrideA = cutlass::gemm::TagToStrideA_t<LayoutTagA>;
|
||||
// We don't care about the actual strideE for the host tensor, but just need one to allocate memory.
|
||||
using StrideE = StrideA;
|
||||
|
||||
static constexpr int SFVecSize = Gemm::GemmKernel::CollectiveMainloop::SFVecSize;
|
||||
// Deduce Cutlass Layouts (RowMajor & ColumnMajor)
|
||||
using LayoutTagB = cutlass::detail::StrideToLayoutTagB_t<StrideB>;
|
||||
|
||||
using LayoutTagE = cutlass::detail::StrideToLayoutTagA_t<StrideE>;
|
||||
|
||||
using ElementSF = typename Gemm::GemmKernel::CollectiveMainloop::ElementSF;
|
||||
using Sm1xxBlkScaledConfig = typename Gemm::GemmKernel::CollectiveMainloop::Sm1xxBlkScaledConfig;
|
||||
using Blk_MN = typename Sm1xxBlkScaledConfig::Blk_MN;
|
||||
using Blk_SF = typename Sm1xxBlkScaledConfig::Blk_SF;
|
||||
using SfAtom = typename Sm1xxBlkScaledConfig::SfAtom;
|
||||
using LayoutSFA = typename Gemm::GemmKernel::CollectiveMainloop::LayoutSFA;
|
||||
using LayoutSFB = typename Gemm::GemmKernel::CollectiveMainloop::LayoutSFB;
|
||||
|
||||
using CompressorUtility = cutlass::transform::kernel::StructuredSparseCompressorUtility<
|
||||
cute::Shape<int, int, int, int>,
|
||||
ElementA,
|
||||
LayoutTagA,
|
||||
SparseConfig>;
|
||||
using CompressorKernel = cutlass::transform::kernel::StructuredSparseCompressor<
|
||||
cute::Shape<int, int, int, int>,
|
||||
ElementA,
|
||||
LayoutTagA,
|
||||
SparseConfig,
|
||||
cutlass::arch::Sm100>;
|
||||
|
||||
using Compressor = cutlass::transform::device::TransformUniversalAdapter<CompressorKernel>;
|
||||
|
||||
using Arguments = typename Gemm::GemmKernel::MainloopArguments;
|
||||
// Whether to use relative equality checks
|
||||
CheckEquality check_relative_equality = CheckEquality::EXACT;
|
||||
|
||||
StrideA stride_a;
|
||||
StrideA stride_a_compressed;
|
||||
StrideB stride_b;
|
||||
StrideE stride_e;
|
||||
|
||||
LayoutA layout_a;
|
||||
LayoutE layout_e;
|
||||
LayoutSFA layout_sfa;
|
||||
LayoutSFB layout_sfb;
|
||||
|
||||
typename LayoutTagA::Stride stride_factor_A;
|
||||
typename LayoutTagB::Stride stride_factor_B;
|
||||
typename LayoutTagE::Stride stride_factor_E;
|
||||
|
||||
cutlass::Distribution::Kind init_A;
|
||||
cutlass::Distribution::Kind init_B;
|
||||
|
||||
cutlass::HostTensor<ElementA, LayoutTagA> tensor_A;
|
||||
cutlass::HostTensor<ElementA, LayoutTagA> tensor_A_Comp;
|
||||
cutlass::HostTensor<ElementB, LayoutTagB> tensor_B;
|
||||
cutlass::HostTensor<ElementE, LayoutTagE> tensor_E;
|
||||
cutlass::HostTensor<ElementSF, LayoutTagA> tensor_SFA;
|
||||
cutlass::HostTensor<ElementSF, LayoutTagB> tensor_SFB;
|
||||
|
||||
uint64_t seed;
|
||||
static constexpr uint64_t kDefaultSeed = 4096;
|
||||
|
||||
// Note: this limitation comes from testbed / not the library
|
||||
static_assert(is_row_or_col_major<StrideA>(),
|
||||
"ERROR : A Layout is neither Row / Column Major)");
|
||||
static_assert(is_row_or_col_major<StrideB>(),
|
||||
"ERROR : B Layout is neither Row / Column Major)");
|
||||
|
||||
HostCollectiveMainloop(
|
||||
CheckEquality check_relative_equality_ = CheckEquality::EXACT,
|
||||
cutlass::Distribution::Kind init_A_ = cutlass::Distribution::Uniform,
|
||||
cutlass::Distribution::Kind init_B_ = cutlass::Distribution::Uniform,
|
||||
uint64_t seed_ = kDefaultSeed,
|
||||
typename LayoutTagA::Stride stride_factor_A_ = typename LayoutTagA::Stride(),
|
||||
typename LayoutTagB::Stride stride_factor_B_ = typename LayoutTagB::Stride(),
|
||||
typename LayoutTagE::Stride stride_factor_E_ = typename LayoutTagE::Stride()
|
||||
):
|
||||
check_relative_equality(check_relative_equality_),
|
||||
stride_factor_A(stride_factor_A_),
|
||||
stride_factor_B(stride_factor_B_),
|
||||
stride_factor_E(stride_factor_E_),
|
||||
init_A(init_A_), init_B(init_B_), seed(seed_) { }
|
||||
|
||||
template<class ProblemShapeType>
|
||||
bool initialize(ProblemShapeType problem_size) {
|
||||
#if (CUTLASS_DEBUG_TRACE_LEVEL > 1)
|
||||
CUTLASS_TRACE_HOST("HostCollectiveMainloop (KernelSparseTmaWarpSpecializedBlockScaledSm100)::initialize");
|
||||
#endif
|
||||
//
|
||||
// Allocate the GEMM workspace
|
||||
//
|
||||
auto problem_shape_MNKL = cute::append<4>(problem_size, 1);
|
||||
auto M = cute::size<0>(problem_shape_MNKL);
|
||||
auto N = cute::size<1>(problem_shape_MNKL);
|
||||
auto K = cute::size<2>(problem_shape_MNKL);
|
||||
auto L = cute::size<3>(problem_shape_MNKL);
|
||||
|
||||
stride_a = cutlass::make_cute_packed_stride(StrideA{}, cute::make_shape(M, K, L));
|
||||
stride_b = cutlass::make_cute_packed_stride(StrideB{}, cute::make_shape(N, K, L));
|
||||
|
||||
CompressorUtility compressor_utility(problem_shape_MNKL, stride_a);
|
||||
|
||||
// TensorE
|
||||
// In unit of ElementE (uint8_t), after alignment requirement
|
||||
// M-dim: TensorEAtom_M alignment
|
||||
// K-dim: TensorEAtom_K alignment
|
||||
int KAlignedE = compressor_utility.get_metadata_k_physical();
|
||||
int MAlignedE = compressor_utility.get_metadata_m_physical();
|
||||
|
||||
// TensorA Compressed
|
||||
// In unit of ElementARaw, after alignment requirement
|
||||
// M-dim: TMA alignment
|
||||
// K-dim: TMA alignment
|
||||
int KAlignedAC = compressor_utility.get_tensorA_k_physical();
|
||||
int MAlignedAC = compressor_utility.get_tensorA_m_physical();
|
||||
|
||||
stride_a_compressed = cutlass::make_cute_packed_stride(StrideA{}, cute::make_shape(M, KAlignedAC, L));
|
||||
stride_e = cutlass::make_cute_packed_stride(StrideE{}, cute::make_shape(MAlignedE, KAlignedE, L));
|
||||
|
||||
auto a_coord = cutlass::make_Coord(M * L, K);
|
||||
auto b_coord = cutlass::make_Coord(K, N * L);
|
||||
auto e_coord = cutlass::make_Coord(MAlignedE * L, KAlignedE);
|
||||
auto a_comp_coord = cutlass::make_Coord(MAlignedAC * L, KAlignedAC);
|
||||
|
||||
tensor_A.resize(a_coord, cutlass::layout::Affine2Layout_Factory<LayoutTagA>::layout_factory(a_coord, stride_factor_A));
|
||||
tensor_A_Comp.resize(a_comp_coord, cutlass::layout::Affine2Layout_Factory<LayoutTagA>::layout_factory(a_comp_coord, stride_factor_A));
|
||||
tensor_B.resize(b_coord, cutlass::layout::Affine2Layout_Factory<LayoutTagB>::layout_factory(b_coord, stride_factor_B));
|
||||
tensor_E.resize(e_coord, cutlass::layout::Affine2Layout_Factory<LayoutTagE>::layout_factory(e_coord, stride_factor_E));
|
||||
|
||||
EXPECT_TRUE(initialize_tensor(tensor_A.host_view(), init_A, seed + 2022));
|
||||
EXPECT_TRUE(initialize_tensor(tensor_B.host_view(), init_B, seed + 2021));
|
||||
|
||||
// It is possible to randomly initialize to all zeros, so override this with non-zeros
|
||||
// in the upper left corner of each operand.
|
||||
tensor_A.host_view().at({0, 0}) = ElementA(1);
|
||||
tensor_B.host_view().at({0, 0}) = ElementB(1);
|
||||
|
||||
compressor_utility.structure_sparse_zero_mask_fill(tensor_A.host_data(), static_cast<int>(seed + 2023));
|
||||
|
||||
tensor_A.sync_device();
|
||||
tensor_B.sync_device();
|
||||
tensor_E.sync_device();
|
||||
tensor_A_Comp.sync_device();
|
||||
|
||||
cutlass::Status status {cutlass::Status::kSuccess };
|
||||
|
||||
cutlass::KernelHardwareInfo hw_info;
|
||||
hw_info.device_id = 0;
|
||||
hw_info.sm_count = cutlass::KernelHardwareInfo::query_device_multiprocessor_count(hw_info.device_id);
|
||||
typename Compressor::Arguments arguments{
|
||||
{M, N, K, L},
|
||||
{tensor_A.device_data(),
|
||||
stride_a,
|
||||
tensor_A_Comp.device_data(),
|
||||
tensor_E.device_data()},
|
||||
{hw_info}
|
||||
};
|
||||
|
||||
Compressor compressor_op;
|
||||
size_t workspace_size = Compressor::get_workspace_size(arguments);
|
||||
cutlass::device_memory::allocation<uint8_t> workspace(workspace_size);
|
||||
|
||||
status = compressor_op.can_implement(arguments);
|
||||
if (status != cutlass::Status::kSuccess) {
|
||||
return false;
|
||||
}
|
||||
|
||||
status = compressor_op.initialize(arguments, workspace.get());
|
||||
if (status != cutlass::Status::kSuccess) {
|
||||
return false;
|
||||
}
|
||||
|
||||
status = compressor_op.run();
|
||||
|
||||
auto result = cudaDeviceSynchronize();
|
||||
if (result != cudaSuccess) {
|
||||
EXPECT_EQ(result, cudaSuccess) << "Error at Kernel Sync.";
|
||||
return false;
|
||||
}
|
||||
|
||||
layout_a = SparseConfig::fill_layoutA(problem_shape_MNKL);
|
||||
layout_e = SparseConfig::fill_layoutE(problem_shape_MNKL);
|
||||
|
||||
tensor_E.sync_host();
|
||||
tensor_A_Comp.sync_host();
|
||||
|
||||
using namespace cute;
|
||||
auto k_blks = cutlass::ceil_div(K, size<1>(shape(SfAtom{})));
|
||||
auto m_blks = cutlass::ceil_div(M, Blk_MN{});
|
||||
auto n_blks = cutlass::ceil_div(N, Blk_MN{});
|
||||
layout_sfa = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(problem_shape_MNKL);
|
||||
layout_sfb = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(problem_shape_MNKL);
|
||||
|
||||
// 2.x host tensor does not natively contain a batch stride or coord, so we spoof if by folding it into the outer mode
|
||||
auto sfa_coord = cutlass::make_Coord(m_blks * Blk_MN{} * L, k_blks * Blk_SF{});
|
||||
auto sfb_coord = cutlass::make_Coord(n_blks * Blk_MN{} * L, k_blks * Blk_SF{});
|
||||
|
||||
tensor_SFA.resize(sfa_coord, cutlass::layout::Affine2Layout_Factory<LayoutTagA>::layout_factory(sfa_coord, stride_factor_A));
|
||||
tensor_SFB.resize(sfb_coord, cutlass::layout::Affine2Layout_Factory<LayoutTagB>::layout_factory(sfb_coord, stride_factor_B));
|
||||
|
||||
EXPECT_TRUE(initialize_tensor(tensor_SFA.host_view(), init_A, seed + 2024));
|
||||
EXPECT_TRUE(initialize_tensor(tensor_SFB.host_view(), init_B, seed + 2025));
|
||||
|
||||
// It is possible to randomly initialize to all zeros, so override this with non-zeros
|
||||
// in the upper left corner of each operand.
|
||||
tensor_SFA.host_view().at({0, 0}) = ElementSF(1);
|
||||
tensor_SFB.host_view().at({0, 0}) = ElementSF(1);
|
||||
|
||||
tensor_SFA.sync_device();
|
||||
tensor_SFB.sync_device();
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
Arguments to_args() {
|
||||
using ArrayElementA = typename Gemm::GemmKernel::CollectiveMainloop::ArrayElementA;
|
||||
using ArrayElementB = typename Gemm::GemmKernel::CollectiveMainloop::ArrayElementB;
|
||||
return {
|
||||
reinterpret_cast<ArrayElementA *>(tensor_A_Comp.device_data()), layout_a,
|
||||
reinterpret_cast<ArrayElementB *>(tensor_B.device_data()), stride_b,
|
||||
tensor_E.device_data(), layout_e,
|
||||
tensor_SFA.device_data(), layout_sfa,
|
||||
tensor_SFB.device_data(), layout_sfb
|
||||
};
|
||||
}
|
||||
|
||||
auto to_host_args(ProblemShapeType problem_size) {
|
||||
using namespace cute;
|
||||
//
|
||||
// Allocate the GEMM workspace
|
||||
//
|
||||
auto problem_shape_MNKL = cute::append<4>(problem_size, 1);
|
||||
auto M = cute::size<0>(problem_shape_MNKL);
|
||||
auto N = cute::size<1>(problem_shape_MNKL);
|
||||
auto K = cute::size<2>(problem_shape_MNKL);
|
||||
auto L = cute::size<3>(problem_shape_MNKL);
|
||||
auto A = make_tensor(make_iterator(tensor_A.host_data()),
|
||||
make_layout(make_shape(M, K, L), stride_a));
|
||||
auto SfA = make_tensor(tensor_SFA.host_data(), layout_sfa);
|
||||
|
||||
auto B = make_tensor(make_iterator(tensor_B.host_data()),
|
||||
make_layout(make_shape(N, K, L), stride_b));
|
||||
auto SfB = make_tensor(tensor_SFB.host_data(), layout_sfb);
|
||||
|
||||
// return {A, SfA, B, SfB};
|
||||
cutlass::reference::host::GettMainloopParams<ElementAccumulator,
|
||||
decltype(A),
|
||||
decltype(B),
|
||||
decltype(SfA),
|
||||
decltype(SfB)
|
||||
>
|
||||
mainloop_params{A, SfA, B, SfB};
|
||||
return mainloop_params;
|
||||
}
|
||||
|
||||
void print_tensors(std::ofstream& file) {
|
||||
file << "A =\n" << tensor_A.host_view()
|
||||
<< "\nB =\n" << tensor_B.host_view()
|
||||
<< "\nSFA =\n" << tensor_SFA.host_view()
|
||||
<< "\nSFB =\n" << tensor_SFB.host_view();
|
||||
}
|
||||
|
||||
bool compare_reference(
|
||||
cute::Shape<int,int,int,int> problem_shape_MNKL) {
|
||||
auto [M, N, K, L] = problem_shape_MNKL;
|
||||
|
||||
EXPECT_GT(cutlass::reference::host::TensorNorm(tensor_A.host_view()), 0);
|
||||
EXPECT_GT(cutlass::reference::host::TensorNorm(tensor_B.host_view()), 0);
|
||||
EXPECT_GT(cutlass::reference::host::TensorNorm(tensor_SFA.host_view()), 0);
|
||||
EXPECT_GT(cutlass::reference::host::TensorNorm(tensor_SFB.host_view()), 0);
|
||||
return true;
|
||||
}
|
||||
};
|
||||
|
||||
template<
|
||||
class Gemm,
|
||||
int SchedulerPipelineStageCount_,
|
||||
class ElementA_,
|
||||
class ElementB_
|
||||
>
|
||||
struct HostCollectiveMainloop<cutlass::gemm::KernelTmaWarpSpecializedCooperativeSparseBlockScaledSm120<SchedulerPipelineStageCount_, true>,
|
||||
Gemm, ElementA_, ElementB_> : public
|
||||
HostCollectiveMainloop<cutlass::gemm::KernelSparseTmaWarpSpecializedBlockScaledSm100<0,0>,
|
||||
Gemm, ElementA_, ElementB_> {
|
||||
using Base = HostCollectiveMainloop<cutlass::gemm::KernelSparseTmaWarpSpecializedBlockScaledSm100<0,0>,
|
||||
Gemm, ElementA_, ElementB_>;
|
||||
HostCollectiveMainloop(
|
||||
CheckEquality check_relative_equality_ = CheckEquality::EXACT,
|
||||
cutlass::Distribution::Kind init_A_ = cutlass::Distribution::Uniform,
|
||||
cutlass::Distribution::Kind init_B_ = cutlass::Distribution::Uniform,
|
||||
uint64_t seed_ = Base::kDefaultSeed,
|
||||
typename Base::LayoutTagA::Stride stride_factor_A_ = typename Base::LayoutTagA::Stride(),
|
||||
typename Base::LayoutTagB::Stride stride_factor_B_ = typename Base::LayoutTagB::Stride(),
|
||||
typename Base::LayoutTagE::Stride stride_factor_E_ = typename Base::LayoutTagE::Stride()
|
||||
) : Base::HostCollectiveMainloop(check_relative_equality_, init_A_, init_B_, seed_, stride_factor_A_,
|
||||
stride_factor_B_,
|
||||
stride_factor_E_) {}
|
||||
};
|
||||
|
||||
template<
|
||||
class Gemm,
|
||||
int SchedulerPipelineStageCount_,
|
||||
class ElementA_,
|
||||
class ElementB_
|
||||
>
|
||||
struct HostCollectiveMainloop<cutlass::gemm::KernelTmaWarpSpecializedCooperativeSparseBlockScaledSm120<SchedulerPipelineStageCount_, false>,
|
||||
Gemm, ElementA_, ElementB_> : public
|
||||
HostCollectiveMainloop<cutlass::gemm::KernelSparseTmaWarpSpecializedBlockScaledSm100<0,0>,
|
||||
Gemm, ElementA_, ElementB_> {
|
||||
using Base = HostCollectiveMainloop<cutlass::gemm::KernelSparseTmaWarpSpecializedBlockScaledSm100<0,0>,
|
||||
Gemm, ElementA_, ElementB_>;
|
||||
HostCollectiveMainloop(
|
||||
CheckEquality check_relative_equality_ = CheckEquality::EXACT,
|
||||
cutlass::Distribution::Kind init_A_ = cutlass::Distribution::Uniform,
|
||||
cutlass::Distribution::Kind init_B_ = cutlass::Distribution::Uniform,
|
||||
uint64_t seed_ = Base::kDefaultSeed,
|
||||
typename Base::LayoutTagA::Stride stride_factor_A_ = typename Base::LayoutTagA::Stride(),
|
||||
typename Base::LayoutTagB::Stride stride_factor_B_ = typename Base::LayoutTagB::Stride(),
|
||||
typename Base::LayoutTagE::Stride stride_factor_E_ = typename Base::LayoutTagE::Stride()
|
||||
) : Base::HostCollectiveMainloop(check_relative_equality_, init_A_, init_B_, seed_, stride_factor_A_,
|
||||
stride_factor_B_,
|
||||
stride_factor_E_) {}
|
||||
};
|
||||
|
||||
template<class Gemm>
|
||||
struct HostCollectiveDefaultEpilogue {
|
||||
// fusion types are potentially void if the fusion is not supported
|
||||
@@ -1391,13 +1875,13 @@ struct HostCollectiveEpilogue {
|
||||
static constexpr bool IsBlockScaleSupported = FusionOp::IsBlockScaleSupported;
|
||||
static constexpr SfStrategy SfGenStrategy = (!IsBlockScaleSupported) ? SfStrategy::None : SfStrategy::SfDGen;
|
||||
static constexpr int32_t SFD_VectorSize = IsBlockScaleSupported ? FusionOp::SFVecSize : 1;
|
||||
static constexpr bool IsKMajorSFD = cute::is_same_v<typename FusionOp::GmemLayoutTagScalefactor, cutlass::layout::RowMajor>;
|
||||
using ElementSFD = non_void_t<typename FusionOp::ElementBlockScaleFactor, ElementD>;
|
||||
using Sm100BlockScaledOutputConfig = cutlass::detail::Sm100BlockScaledOutputConfig<
|
||||
SFD_VectorSize
|
||||
>;
|
||||
using Blk_MN = typename Sm100BlockScaledOutputConfig::Blk_MN;
|
||||
using Blk_SF = typename Sm100BlockScaledOutputConfig::Blk_SF;
|
||||
using OutputSFAtom = typename Sm100BlockScaledOutputConfig::SfAtom;
|
||||
using Sm1xxBlockScaledOutputConfig= cutlass::detail::Sm1xxBlockScaledOutputConfig<SFD_VectorSize,
|
||||
IsKMajorSFD ? cute::UMMA::Major::K : cute::UMMA::Major::MN>;
|
||||
using Blk_MN = typename Sm1xxBlockScaledOutputConfig::Blk_MN;
|
||||
using Blk_SF = typename Sm1xxBlockScaledOutputConfig::Blk_SF;
|
||||
using OutputSFAtom = typename Sm1xxBlockScaledOutputConfig::SfAtom;
|
||||
cutlass::HostTensor<ElementSFD, LayoutTagD> tensor_SFD;
|
||||
cutlass::HostTensor<ElementSFD, LayoutTagD> reference_SFD;
|
||||
|
||||
@@ -1693,7 +2177,12 @@ struct HostCollectiveEpilogue {
|
||||
auto m_blks = cutlass::ceil_div(M, cute::size<0>(cute::shape(OutputSFAtom{})));
|
||||
auto n_blks = cutlass::ceil_div(N, cute::size<1>(cute::shape(OutputSFAtom{})));
|
||||
auto sfd_coord = [&] () {
|
||||
if constexpr (IsKMajorSFD) {
|
||||
return cutlass::make_Coord(m_blks * Blk_MN{} * L, n_blks * Blk_SF{});
|
||||
}
|
||||
else {
|
||||
return cutlass::make_Coord(m_blks * Blk_SF{} * L, n_blks * Blk_MN{});
|
||||
}
|
||||
}();
|
||||
tensor_SFD.resize(sfd_coord, cutlass::layout::Affine2Layout_Factory<LayoutTagD>::layout_factory(sfd_coord, stride_factor_D));
|
||||
reference_SFD.resize(sfd_coord, cutlass::layout::Affine2Layout_Factory<LayoutTagD>::layout_factory(sfd_coord, stride_factor_D), false);
|
||||
@@ -2062,7 +2551,7 @@ struct HostCollectiveEpilogue {
|
||||
auto SfD = [&](){
|
||||
if constexpr (IsBlockScaleSupported) {
|
||||
auto tensor = make_tensor(detail::make_iterator(reference_SFD.host_data()),
|
||||
Sm100BlockScaledOutputConfig::tile_atom_to_shape_SFD(problem_shape_MNKL));
|
||||
Sm1xxBlockScaledOutputConfig::tile_atom_to_shape_SFD(problem_shape_MNKL));
|
||||
return tensor;
|
||||
}
|
||||
else {
|
||||
@@ -3154,6 +3643,18 @@ bool TestSmall(double alpha = 1.0, double beta = cute::is_same_v<typename Gemm::
|
||||
max_alignment_n = std::max(Gemm::kAlignmentA, Gemm::kAlignmentB);
|
||||
max_alignment_m = std::max(Gemm::kAlignmentA, Gemm::kAlignmentB);
|
||||
}
|
||||
// Alignment for SFD
|
||||
if constexpr (detail::IsSfdEpi<typename Gemm::GemmKernel::CollectiveEpilogue>::value) {
|
||||
using GmemLayoutTagScalefactor = typename Gemm::GemmKernel::CollectiveEpilogue::FusionCallbacks::Operation::GmemLayoutTagScalefactor;
|
||||
constexpr int SFDVecSize = Gemm::GemmKernel::CollectiveEpilogue::FusionCallbacks::Operation::SFVecSize;
|
||||
if constexpr (cute::is_same_v<GmemLayoutTagScalefactor, cutlass::layout::RowMajor>) {
|
||||
max_alignment_n = std::lcm(max_alignment_n, SFDVecSize);
|
||||
}
|
||||
else {
|
||||
max_alignment_m = std::lcm(max_alignment_m, SFDVecSize);
|
||||
}
|
||||
}
|
||||
|
||||
float waves[] = {0.5, 1.25, 2.5};
|
||||
int cluster_m = 1;
|
||||
int cluster_n = 1;
|
||||
|
||||
@@ -552,10 +552,10 @@ struct HostCollectiveMainloop<cutlass::gemm::KernelPtrArrayTmaWarpSpecializedBlo
|
||||
static constexpr int SFVecSize = Gemm::GemmKernel::CollectiveMainloop::SFVecSize;
|
||||
|
||||
using ElementSF = typename Gemm::GemmKernel::ElementSF;
|
||||
using Sm100BlkScaledConfig = typename Gemm::GemmKernel::CollectiveMainloop::Sm100BlkScaledConfig;
|
||||
using Blk_MN = typename Sm100BlkScaledConfig::Blk_MN;
|
||||
using Blk_SF = typename Sm100BlkScaledConfig::Blk_SF;
|
||||
using SfAtom = typename Sm100BlkScaledConfig::SfAtom;
|
||||
using Sm1xxBlkScaledConfig = typename Gemm::GemmKernel::CollectiveMainloop::Sm1xxBlkScaledConfig;
|
||||
using Blk_MN = typename Sm1xxBlkScaledConfig::Blk_MN;
|
||||
using Blk_SF = typename Sm1xxBlkScaledConfig::Blk_SF;
|
||||
using SfAtom = typename Sm1xxBlkScaledConfig::SfAtom;
|
||||
using LayoutSFA = typename Gemm::GemmKernel::CollectiveMainloop::LayoutSFA;
|
||||
using InternalLayoutSFA = typename Gemm::GemmKernel::CollectiveMainloop::InternalLayoutSFA;
|
||||
using LayoutSFB = typename Gemm::GemmKernel::CollectiveMainloop::LayoutSFB;
|
||||
@@ -662,8 +662,8 @@ struct HostCollectiveMainloop<cutlass::gemm::KernelPtrArrayTmaWarpSpecializedBlo
|
||||
auto k_blks = cutlass::ceil_div(K, size<1>(shape(SfAtom{})));
|
||||
auto m_blks = cutlass::ceil_div(M, Blk_MN{});
|
||||
auto n_blks = cutlass::ceil_div(N, Blk_MN{});
|
||||
layout_sfa_host.push_back(Sm100BlkScaledConfig::tile_atom_to_shape_SFA(cute::make_shape(M, N, K, 1)));
|
||||
layout_sfb_host.push_back(Sm100BlkScaledConfig::tile_atom_to_shape_SFB(cute::make_shape(M, N, K, 1)));
|
||||
layout_sfa_host.push_back(Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(cute::make_shape(M, N, K, 1)));
|
||||
layout_sfb_host.push_back(Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(cute::make_shape(M, N, K, 1)));
|
||||
|
||||
// 2.x host tensor does not natively contain a batch stride or coord, so we spoof if by folding it into the outer mode
|
||||
auto sfa_coord = cutlass::make_Coord(m_blks * Blk_MN{}, k_blks * Blk_SF{});
|
||||
@@ -1110,12 +1110,12 @@ struct HostCollectiveEpilogue {
|
||||
static constexpr SfStrategy SfGenStrategy = (!IsBlockScaleSupported) ? SfStrategy::None : SfStrategy::SfDGen;
|
||||
static constexpr int32_t SFD_VectorSize = IsBlockScaleSupported ? FusionOp::SFVecSize : 1;
|
||||
using ElementSFD = non_void_t<cute::remove_pointer_t<typename FusionOp::ElementBlockScaleFactor>, ElementD>;
|
||||
using Sm100BlockScaledOutputConfig = cutlass::detail::Sm100BlockScaledOutputConfig<
|
||||
using Sm1xxBlockScaledOutputConfig= cutlass::detail::Sm1xxBlockScaledOutputConfig<
|
||||
SFD_VectorSize
|
||||
>;
|
||||
using Blk_MN = typename Sm100BlockScaledOutputConfig::Blk_MN;
|
||||
using Blk_SF = typename Sm100BlockScaledOutputConfig::Blk_SF;
|
||||
using OutputSFAtom = typename Sm100BlockScaledOutputConfig::SfAtom;
|
||||
using Blk_MN = typename Sm1xxBlockScaledOutputConfig::Blk_MN;
|
||||
using Blk_SF = typename Sm1xxBlockScaledOutputConfig::Blk_SF;
|
||||
using OutputSFAtom = typename Sm1xxBlockScaledOutputConfig::SfAtom;
|
||||
std::vector<cutlass::HostTensor<ElementSFD, LayoutTagD>> tensors_SFD;
|
||||
std::vector<cutlass::HostTensor<ElementSFD, LayoutTagD>> references_SFD;
|
||||
cutlass::DeviceAllocation<ElementSFD *> device_tensors_SFD;
|
||||
@@ -1711,7 +1711,7 @@ struct HostCollectiveEpilogue {
|
||||
auto SfD = [&](){
|
||||
if constexpr (IsBlockScaleSupported) {
|
||||
auto tensor = make_tensor(detail::make_iterator(references_SFD[batch].host_data()),
|
||||
Sm100BlockScaledOutputConfig::tile_atom_to_shape_SFD(problem_shape_MNKL));
|
||||
Sm1xxBlockScaledOutputConfig::tile_atom_to_shape_SFD(problem_shape_MNKL));
|
||||
return tensor;
|
||||
}
|
||||
else {
|
||||
|
||||
@@ -0,0 +1,174 @@
|
||||
# Copyright (c) 2025 - 2025 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.
|
||||
|
||||
if (CUTLASS_NVCC_ARCHS MATCHES 100a)
|
||||
|
||||
add_custom_target(
|
||||
cutlass_test_unit_gemm_device_sm100_bssp
|
||||
DEPENDS
|
||||
cutlass_test_unit_gemm_device_sm100_bssp_nvf4_nvf4_f32_f32_f32_o
|
||||
cutlass_test_unit_gemm_device_sm100_bssp_nvf4_nvf4_f32_f16_f16_o
|
||||
cutlass_test_unit_gemm_device_sm100_bssp_nvf4_nvf4_f32_f16_nvf4_o
|
||||
cutlass_test_unit_gemm_device_sm100_bssp_mxf8_mxf8_f32_f32_f32_q
|
||||
cutlass_test_unit_gemm_device_sm100_bssp_mxf8_mxf8_f32_f16_f16_q
|
||||
cutlass_test_unit_gemm_device_sm100_bssp_mxf8_mxf8_f32_f16_mxf8_q
|
||||
cutlass_test_unit_gemm_device_sm100_bssp_mxf8_mxf4_f32_q
|
||||
cutlass_test_unit_gemm_device_sm100_bssp_mxf4_mxf4_f32_q
|
||||
cutlass_test_unit_gemm_device_sm100_bssp_mxf4_mxf4_f32_o
|
||||
cutlass_test_unit_gemm_device_sm100_bssp_streamk
|
||||
)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable_split_file(
|
||||
cutlass_test_unit_gemm_device_sm100_bssp_nvf4_nvf4_f32_f32_f32_o
|
||||
|
||||
BATCH_SOURCES ON
|
||||
BATCH_SIZE 1
|
||||
|
||||
sm100_bssp_gemm_nvf4_nvf4_f32_f32_f32_o_tnt.cu
|
||||
sm100_bssp_gemm_nvf4_nvf4_f32_void_f32_o_tnt.cu
|
||||
|
||||
sm100_bssp_gemm_nvf4_nvf4_f32_f32_f32_o_tnn.cu
|
||||
sm100_bssp_gemm_nvf4_nvf4_f32_void_f32_o_tnn.cu
|
||||
)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable_split_file(
|
||||
cutlass_test_unit_gemm_device_sm100_bssp_nvf4_nvf4_f32_f16_f16_o
|
||||
|
||||
BATCH_SOURCES ON
|
||||
BATCH_SIZE 1
|
||||
|
||||
sm100_bssp_gemm_nvf4_nvf4_f32_f16_f16_o_tnt.cu
|
||||
sm100_bssp_gemm_nvf4_nvf4_f32_void_f16_o_tnt.cu
|
||||
|
||||
sm100_bssp_gemm_nvf4_nvf4_f32_f16_f16_o_tnn.cu
|
||||
sm100_bssp_gemm_nvf4_nvf4_f32_void_f16_o_tnn.cu
|
||||
)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable_split_file(
|
||||
cutlass_test_unit_gemm_device_sm100_bssp_nvf4_nvf4_f32_f16_nvf4_o
|
||||
|
||||
BATCH_SOURCES ON
|
||||
BATCH_SIZE 1
|
||||
|
||||
sm100_bssp_gemm_nvf4_nvf4_f32_f16_nvf4_o_tnt_sfd.cu
|
||||
sm100_bssp_gemm_nvf4_nvf4_f32_void_nvf4_o_tnt_sfd.cu
|
||||
|
||||
sm100_bssp_gemm_nvf4_nvf4_f32_f16_nvf4_o_tnn_sfd.cu
|
||||
sm100_bssp_gemm_nvf4_nvf4_f32_void_nvf4_o_tnn_sfd.cu
|
||||
)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable_split_file(
|
||||
cutlass_test_unit_gemm_device_sm100_bssp_mxf8_mxf8_f32_f32_f32_q
|
||||
|
||||
BATCH_SOURCES ON
|
||||
BATCH_SIZE 1
|
||||
|
||||
sm100_bssp_gemm_mxf8_mxf8_f32_f32_f32_q_tnt.cu
|
||||
sm100_bssp_gemm_mxf8_mxf8_f32_void_f32_q_tnt.cu
|
||||
|
||||
sm100_bssp_gemm_mxf8_mxf8_f32_f32_f32_q_tnn.cu
|
||||
sm100_bssp_gemm_mxf8_mxf8_f32_void_f32_q_tnn.cu
|
||||
)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable_split_file(
|
||||
cutlass_test_unit_gemm_device_sm100_bssp_mxf8_mxf8_f32_f16_f16_q
|
||||
|
||||
BATCH_SOURCES ON
|
||||
BATCH_SIZE 1
|
||||
|
||||
sm100_bssp_gemm_mxf8_mxf8_f32_f16_f16_q_tnt.cu
|
||||
sm100_bssp_gemm_mxf8_mxf8_f32_void_f16_q_tnt.cu
|
||||
|
||||
sm100_bssp_gemm_mxf8_mxf8_f32_f16_f16_q_tnn.cu
|
||||
sm100_bssp_gemm_mxf8_mxf8_f32_void_f16_q_tnn.cu
|
||||
)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable_split_file(
|
||||
cutlass_test_unit_gemm_device_sm100_bssp_mxf8_mxf8_f32_f16_mxf8_q
|
||||
|
||||
BATCH_SOURCES ON
|
||||
BATCH_SIZE 1
|
||||
|
||||
sm100_bssp_gemm_mxf8_mxf8_f32_f16_mxf8_q_tnt_sfd.cu
|
||||
sm100_bssp_gemm_mxf8_mxf8_f32_void_mxf8_q_tnt_sfd.cu
|
||||
|
||||
sm100_bssp_gemm_mxf8_mxf8_f32_f16_mxf8_q_tnn_sfd.cu
|
||||
sm100_bssp_gemm_mxf8_mxf8_f32_void_mxf8_q_tnn_sfd.cu
|
||||
|
||||
sm100_bssp_gemm_mxf8_mxf8_f32_f16_mxf8_q_ttt_sfd.cu
|
||||
sm100_bssp_gemm_mxf8_mxf8_f32_f16_mxf8_q_ttn_sfd.cu
|
||||
sm100_bssp_gemm_mxf8_mxf8_f32_f16_mxf8_q_nnt_sfd.cu
|
||||
sm100_bssp_gemm_mxf8_mxf8_f32_f16_mxf8_q_nnn_sfd.cu
|
||||
)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable_split_file(
|
||||
cutlass_test_unit_gemm_device_sm100_bssp_mxf4_mxf4_f32_o
|
||||
|
||||
BATCH_SOURCES ON
|
||||
BATCH_SIZE 1
|
||||
|
||||
sm100_bssp_gemm_mxf4_mxf4_f32_f16_f16_o_tnn.cu
|
||||
sm100_bssp_gemm_mxf4_mxf4_f32_f16_f16_o_tnt.cu
|
||||
|
||||
sm100_bssp_gemm_mxf4_mxf4_f32_f32_f32_o_tnt.cu
|
||||
sm100_bssp_gemm_mxf4_mxf4_f32_f32_f32_o_tnn.cu
|
||||
)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable_split_file(
|
||||
cutlass_test_unit_gemm_device_sm100_bssp_mxf8_mxf4_f32_q
|
||||
|
||||
BATCH_SOURCES ON
|
||||
BATCH_SIZE 1
|
||||
|
||||
sm100_bssp_gemm_mxf8_mxf4_f32_f16_mxf8_q_tnt.cu
|
||||
sm100_bssp_gemm_mxf8_mxf4_f32_f16_f16_q_tnt.cu
|
||||
sm100_bssp_gemm_mxf8_mxf4_f32_f32_f32_q_tnt.cu
|
||||
)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable_split_file(
|
||||
cutlass_test_unit_gemm_device_sm100_bssp_mxf4_mxf4_f32_q
|
||||
|
||||
BATCH_SOURCES ON
|
||||
BATCH_SIZE 1
|
||||
|
||||
sm100_bssp_gemm_mxf4_mxf4_f32_f16_mxf8_q_tnt.cu
|
||||
sm100_bssp_gemm_mxf4_mxf4_f32_f16_f16_q_tnt.cu
|
||||
sm100_bssp_gemm_mxf4_mxf4_f32_f32_f32_q_tnt.cu
|
||||
)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable_split_file(
|
||||
cutlass_test_unit_gemm_device_sm100_bssp_streamk
|
||||
|
||||
BATCH_SOURCES ON
|
||||
BATCH_SIZE 1
|
||||
|
||||
sm100_bssp_gemm_nvf4_nvf4_f32_f16_nvf4_o_tnt_streamk.cu
|
||||
sm100_bssp_gemm_mxf8_mxf8_f32_f16_mxf8_q_tnt_streamk.cu
|
||||
)
|
||||
|
||||
endif()
|
||||
+1451
File diff suppressed because it is too large
Load Diff
+1451
File diff suppressed because it is too large
Load Diff
+1102
File diff suppressed because it is too large
Load Diff
+1102
File diff suppressed because it is too large
Load Diff
+1451
File diff suppressed because it is too large
Load Diff
+1453
File diff suppressed because it is too large
Load Diff
+1102
File diff suppressed because it is too large
Load Diff
+1102
File diff suppressed because it is too large
Load Diff
+1102
File diff suppressed because it is too large
Load Diff
+1102
File diff suppressed because it is too large
Load Diff
+580
@@ -0,0 +1,580 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnn_vs64in
|
||||
// 2. 128x192_tnn_vs64in
|
||||
// 3. 128x256_tnn_vs64in
|
||||
// 4. 256x128_tnn_vs64in
|
||||
// 5. 256x192_tnn_vs64in
|
||||
// 6. 256x256_tnn_vs64in
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_128x128x256_0_vs64_tnn_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_128x192x256_0_vs64_tnn_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_128x256x256_0_vs64_tnn_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_256x128x256_0_vs64_tnn_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_256x192x256_0_vs64_tnn_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_256x256x256_0_vs64_tnn_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_128x128x256_0_vs64_tnn_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_128x128x256_0_vs64_tnn_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_128x192x256_0_vs64_tnn_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_128x192x256_0_vs64_tnn_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_128x256x256_0_vs64_tnn_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_128x256x256_0_vs64_tnn_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_256x128x256_0_vs64_tnn_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_256x128x256_0_vs64_tnn_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_256x192x256_0_vs64_tnn_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_256x192x256_0_vs64_tnn_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_256x256x256_0_vs64_tnn_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_256x256x256_0_vs64_tnn_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+580
@@ -0,0 +1,580 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnt_vs64in
|
||||
// 2. 128x192_tnt_vs64in
|
||||
// 3. 128x256_tnt_vs64in
|
||||
// 4. 256x128_tnt_vs64in
|
||||
// 5. 256x192_tnt_vs64in
|
||||
// 6. 256x256_tnt_vs64in
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_128x128x256_0_vs64_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_128x192x256_0_vs64_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_128x256x256_0_vs64_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_256x128x256_0_vs64_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_256x192x256_0_vs64_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_256x256x256_0_vs64_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_128x128x256_0_vs64_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_128x128x256_0_vs64_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_128x192x256_0_vs64_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_128x192x256_0_vs64_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_128x256x256_0_vs64_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_128x256x256_0_vs64_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_256x128x256_0_vs64_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_256x128x256_0_vs64_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_256x192x256_0_vs64_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_256x192x256_0_vs64_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_256x256x256_0_vs64_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_f16_256x256x256_0_vs64_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+614
@@ -0,0 +1,614 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test Config
|
||||
// 1. layout {nnn_InputVs64_OutputVs64}
|
||||
// 2. tilesize {128x128, 128x192, 128x256, 256x128, 256x192, 256x256}
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_nnn_vs64in_vs64out
|
||||
// 2. 128x192_nnn_vs64in_vs64out
|
||||
// 3. 128x256_nnn_vs64in_vs64out
|
||||
// 4. 256x128_nnn_vs64in_vs64out
|
||||
// 5. 256x192_nnn_vs64in_vs64out
|
||||
// 6. 256x256_nnn_vs64in_vs64out
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x128x256_0_vs64_nnn_align32_q_1sm_epiVs64n {
|
||||
|
||||
using LayoutA = cutlass::layout::ColumnMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x192x256_0_vs64_nnn_align32_q_1sm_epiVs64n {
|
||||
|
||||
using LayoutA = cutlass::layout::ColumnMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x256x256_0_vs64_nnn_align32_q_1sm_epiVs64n {
|
||||
|
||||
using LayoutA = cutlass::layout::ColumnMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x128x256_0_vs64_nnn_align32_q_2sm_epiVs64n {
|
||||
|
||||
using LayoutA = cutlass::layout::ColumnMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x192x256_0_vs64_nnn_align32_q_2sm_epiVs64n {
|
||||
|
||||
using LayoutA = cutlass::layout::ColumnMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x256x256_0_vs64_nnn_align32_q_2sm_epiVs64n {
|
||||
|
||||
using LayoutA = cutlass::layout::ColumnMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x128x256_0_vs64_nnn_align32_q_1sm_epiVs64n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x128x256_0_vs64_nnn_align32_q_1sm_epiVs64n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x192x256_0_vs64_nnn_align32_q_1sm_epiVs64n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x192x256_0_vs64_nnn_align32_q_1sm_epiVs64n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x256x256_0_vs64_nnn_align32_q_1sm_epiVs64n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x256x256_0_vs64_nnn_align32_q_1sm_epiVs64n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x128x256_0_vs64_nnn_align32_q_2sm_epiVs64n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x128x256_0_vs64_nnn_align32_q_2sm_epiVs64n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x192x256_0_vs64_nnn_align32_q_2sm_epiVs64n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x192x256_0_vs64_nnn_align32_q_2sm_epiVs64n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x256x256_0_vs64_nnn_align32_q_2sm_epiVs64n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x256x256_0_vs64_nnn_align32_q_2sm_epiVs64n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+614
@@ -0,0 +1,614 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test Config
|
||||
// 1. layout {nnt_InputVs64_OutputVs64}
|
||||
// 2. tilesize {128x128, 128x192, 128x256, 256x128, 256x192, 256x256}
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_nnt_vs64in_vs64out
|
||||
// 2. 128x192_nnt_vs64in_vs64out
|
||||
// 3. 128x256_nnt_vs64in_vs64out
|
||||
// 4. 256x128_nnt_vs64in_vs64out
|
||||
// 5. 256x192_nnt_vs64in_vs64out
|
||||
// 6. 256x256_nnt_vs64in_vs64out
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x128x256_0_vs64_nnt_align32_q_1sm_epiVs64t {
|
||||
|
||||
using LayoutA = cutlass::layout::ColumnMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x192x256_0_vs64_nnt_align32_q_1sm_epiVs64t {
|
||||
|
||||
using LayoutA = cutlass::layout::ColumnMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x256x256_0_vs64_nnt_align32_q_1sm_epiVs64t {
|
||||
|
||||
using LayoutA = cutlass::layout::ColumnMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x128x256_0_vs64_nnt_align32_q_2sm_epiVs64t {
|
||||
|
||||
using LayoutA = cutlass::layout::ColumnMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x192x256_0_vs64_nnt_align32_q_2sm_epiVs64t {
|
||||
|
||||
using LayoutA = cutlass::layout::ColumnMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x256x256_0_vs64_nnt_align32_q_2sm_epiVs64t {
|
||||
|
||||
using LayoutA = cutlass::layout::ColumnMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x128x256_0_vs64_nnt_align32_q_1sm_epiVs64t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x128x256_0_vs64_nnt_align32_q_1sm_epiVs64t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x192x256_0_vs64_nnt_align32_q_1sm_epiVs64t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x192x256_0_vs64_nnt_align32_q_1sm_epiVs64t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x256x256_0_vs64_nnt_align32_q_1sm_epiVs64t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x256x256_0_vs64_nnt_align32_q_1sm_epiVs64t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x128x256_0_vs64_nnt_align32_q_2sm_epiVs64t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x128x256_0_vs64_nnt_align32_q_2sm_epiVs64t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x192x256_0_vs64_nnt_align32_q_2sm_epiVs64t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x192x256_0_vs64_nnt_align32_q_2sm_epiVs64t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x256x256_0_vs64_nnt_align32_q_2sm_epiVs64t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x256x256_0_vs64_nnt_align32_q_2sm_epiVs64t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+614
@@ -0,0 +1,614 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test Config
|
||||
// 1. layout {tnn_InputVs64_OutputVs64}
|
||||
// 2. tilesize {128x128, 128x192, 128x256, 256x128, 256x192, 256x256}
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnn_vs64in_vs64out
|
||||
// 2. 128x192_tnn_vs64in_vs64out
|
||||
// 3. 128x256_tnn_vs64in_vs64out
|
||||
// 4. 256x128_tnn_vs64in_vs64out
|
||||
// 5. 256x192_tnn_vs64in_vs64out
|
||||
// 6. 256x256_tnn_vs64in_vs64out
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x128x256_0_vs64_tnn_align32_q_1sm_epiVs64n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 16;
|
||||
constexpr int kAlignmentC = 16;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x192x256_0_vs64_tnn_align32_q_1sm_epiVs64n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 16;
|
||||
constexpr int kAlignmentC = 16;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x256x256_0_vs64_tnn_align32_q_1sm_epiVs64n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 16;
|
||||
constexpr int kAlignmentC = 16;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x128x256_0_vs64_tnn_align32_q_2sm_epiVs64n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 16;
|
||||
constexpr int kAlignmentC = 16;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x192x256_0_vs64_tnn_align32_q_2sm_epiVs64n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 16;
|
||||
constexpr int kAlignmentC = 16;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x256x256_0_vs64_tnn_align32_q_2sm_epiVs64n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 16;
|
||||
constexpr int kAlignmentC = 16;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x128x256_0_vs64_tnn_align32_q_1sm_epiVs64n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x128x256_0_vs64_tnn_align32_q_1sm_epiVs64n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x192x256_0_vs64_tnn_align32_q_1sm_epiVs64n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x192x256_0_vs64_tnn_align32_q_1sm_epiVs64n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x256x256_0_vs64_tnn_align32_q_1sm_epiVs64n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x256x256_0_vs64_tnn_align32_q_1sm_epiVs64n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x128x256_0_vs64_tnn_align32_q_2sm_epiVs64n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x128x256_0_vs64_tnn_align32_q_2sm_epiVs64n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x192x256_0_vs64_tnn_align32_q_2sm_epiVs64n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x192x256_0_vs64_tnn_align32_q_2sm_epiVs64n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x256x256_0_vs64_tnn_align32_q_2sm_epiVs64n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x256x256_0_vs64_tnn_align32_q_2sm_epiVs64n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+614
@@ -0,0 +1,614 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test Config
|
||||
// 1. layout {tnt_InputVs64_OutputVs64}
|
||||
// 2. tilesize {128x128, 128x192, 128x256, 256x128, 256x192, 256x256}
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnt_vs64in_vs64out [NOT SUPPORTED]
|
||||
// 2. 128x192_tnt_vs64in_vs64out
|
||||
// 3. 128x256_tnt_vs64in_vs64out
|
||||
// 4. 256x128_tnt_vs64in_vs64out [NOT SUPPORTED]
|
||||
// 5. 256x192_tnt_vs64in_vs64out
|
||||
// 6. 256x256_tnt_vs64in_vs64out
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x128x256_0_vs64_tnt_align32_q_1sm_epiVs64t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x192x256_0_vs64_tnt_align32_q_1sm_epiVs64t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x256x256_0_vs64_tnt_align32_q_1sm_epiVs64t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x128x256_0_vs64_tnt_align32_q_2sm_epiVs64t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x192x256_0_vs64_tnt_align32_q_2sm_epiVs64t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x256x256_0_vs64_tnt_align32_q_2sm_epiVs64t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x128x256_0_vs64_tnt_align32_q_1sm_epiVs64t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x128x256_0_vs64_tnt_align32_q_1sm_epiVs64t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x192x256_0_vs64_tnt_align32_q_1sm_epiVs64t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x192x256_0_vs64_tnt_align32_q_1sm_epiVs64t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x256x256_0_vs64_tnt_align32_q_1sm_epiVs64t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x256x256_0_vs64_tnt_align32_q_1sm_epiVs64t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x128x256_0_vs64_tnt_align32_q_2sm_epiVs64t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x128x256_0_vs64_tnt_align32_q_2sm_epiVs64t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x192x256_0_vs64_tnt_align32_q_2sm_epiVs64t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x192x256_0_vs64_tnt_align32_q_2sm_epiVs64t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x256x256_0_vs64_tnt_align32_q_2sm_epiVs64t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x256x256_0_vs64_tnt_align32_q_2sm_epiVs64t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+538
@@ -0,0 +1,538 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnt_vs64in
|
||||
// 2. 128x192_tnt_vs64in
|
||||
// 3. 128x256_tnt_vs64in
|
||||
// 4. 256x128_tnt_vs64in
|
||||
// 5. 256x192_tnt_vs64in
|
||||
// 6. 256x256_tnt_vs64in
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_e4m3_128x128x256_0_vs64_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::StreamKScheduler>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_e4m3_128x192x256_0_vs64_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::StreamKScheduler>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_e4m3_128x256x256_0_vs64_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::StreamKScheduler>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_e4m3_256x128x256_0_vs64_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::StreamKScheduler>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_e4m3_256x192x256_0_vs64_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::StreamKScheduler>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_e4m3_256x256x256_0_vs64_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::StreamKScheduler>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_e4m3_128x128x256_0_vs64_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_e4m3_128x128x256_0_vs64_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{1536}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_e4m3_128x192x256_0_vs64_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_e4m3_128x192x256_0_vs64_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{1536}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_e4m3_128x256x256_0_vs64_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_e4m3_128x256x256_0_vs64_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{1536}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_e4m3_256x128x256_0_vs64_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_e4m3_256x128x256_0_vs64_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{1536}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_e4m3_256x192x256_0_vs64_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_e4m3_256x192x256_0_vs64_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{1536}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_e4m3_256x256x256_0_vs64_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_e4m3_256x256x256_0_vs64_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{1536}));
|
||||
}
|
||||
|
||||
#endif // #if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
+614
@@ -0,0 +1,614 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test Config
|
||||
// 1. layout {ttn_InputVs64_OutputVs64}
|
||||
// 2. tilesize {128x128, 128x192, 128x256, 256x128, 256x192, 256x256}
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_ttn_vs64in_vs64out
|
||||
// 2. 128x192_ttn_vs64in_vs64out
|
||||
// 3. 128x256_ttn_vs64in_vs64out
|
||||
// 4. 256x128_ttn_vs64in_vs64out
|
||||
// 5. 256x192_ttn_vs64in_vs64out
|
||||
// 6. 256x256_ttn_vs64in_vs64out
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x128x256_0_vs64_ttn_align32_q_1sm_epiVs64n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::RowMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x192x256_0_vs64_ttn_align32_q_1sm_epiVs64n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::RowMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x256x256_0_vs64_ttn_align32_q_1sm_epiVs64n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::RowMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x128x256_0_vs64_ttn_align32_q_2sm_epiVs64n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::RowMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x192x256_0_vs64_ttn_align32_q_2sm_epiVs64n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::RowMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x256x256_0_vs64_ttn_align32_q_2sm_epiVs64n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::RowMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x128x256_0_vs64_ttn_align32_q_1sm_epiVs64n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x128x256_0_vs64_ttn_align32_q_1sm_epiVs64n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x192x256_0_vs64_ttn_align32_q_1sm_epiVs64n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x192x256_0_vs64_ttn_align32_q_1sm_epiVs64n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x256x256_0_vs64_ttn_align32_q_1sm_epiVs64n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x256x256_0_vs64_ttn_align32_q_1sm_epiVs64n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x128x256_0_vs64_ttn_align32_q_2sm_epiVs64n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x128x256_0_vs64_ttn_align32_q_2sm_epiVs64n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x192x256_0_vs64_ttn_align32_q_2sm_epiVs64n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x192x256_0_vs64_ttn_align32_q_2sm_epiVs64n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x256x256_0_vs64_ttn_align32_q_2sm_epiVs64n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x256x256_0_vs64_ttn_align32_q_2sm_epiVs64n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+614
@@ -0,0 +1,614 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test Config
|
||||
// 1. layout {ttt_InputVs64_OutputVs64}
|
||||
// 2. tilesize {128x128, 128x192, 128x256, 256x128, 256x192, 256x256}
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_ttt_vs64in_vs64out [NOT SUPPORTED]
|
||||
// 2. 128x192_ttt_vs64in_vs64out
|
||||
// 3. 128x256_ttt_vs64in_vs64out
|
||||
// 4. 256x128_ttt_vs64in_vs64out [NOT SUPPORTED]
|
||||
// 5. 256x192_ttt_vs64in_vs64out
|
||||
// 6. 256x256_ttt_vs64in_vs64out
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x128x256_0_vs64_ttt_align32_q_1sm_epiVs64t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::RowMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x192x256_0_vs64_ttt_align32_q_1sm_epiVs64t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::RowMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x256x256_0_vs64_ttt_align32_q_1sm_epiVs64t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::RowMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x128x256_0_vs64_ttt_align32_q_2sm_epiVs64t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::RowMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x192x256_0_vs64_ttt_align32_q_2sm_epiVs64t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::RowMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x256x256_0_vs64_ttt_align32_q_2sm_epiVs64t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::RowMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x128x256_0_vs64_ttt_align32_q_1sm_epiVs64t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x128x256_0_vs64_ttt_align32_q_1sm_epiVs64t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x192x256_0_vs64_ttt_align32_q_1sm_epiVs64t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x192x256_0_vs64_ttt_align32_q_1sm_epiVs64t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x256x256_0_vs64_ttt_align32_q_1sm_epiVs64t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_128x256x256_0_vs64_ttt_align32_q_1sm_epiVs64t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x128x256_0_vs64_ttt_align32_q_2sm_epiVs64t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x128x256_0_vs64_ttt_align32_q_2sm_epiVs64t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x192x256_0_vs64_ttt_align32_q_2sm_epiVs64t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x192x256_0_vs64_ttt_align32_q_2sm_epiVs64t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x256x256_0_vs64_ttt_align32_q_2sm_epiVs64t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f16_ue8m0xe4m3_256x256x256_0_vs64_ttt_align32_q_2sm_epiVs64t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+580
@@ -0,0 +1,580 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnn_vs64in
|
||||
// 2. 128x192_tnn_vs64in
|
||||
// 3. 128x256_tnn_vs64in
|
||||
// 4. 256x128_tnn_vs64in
|
||||
// 5. 256x192_tnn_vs64in
|
||||
// 6. 256x256_tnn_vs64in
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_128x128x256_0_vs64_tnn_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_128x192x256_0_vs64_tnn_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_128x256x256_0_vs64_tnn_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_256x128x256_0_vs64_tnn_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_256x192x256_0_vs64_tnn_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_256x256x256_0_vs64_tnn_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_128x128x256_0_vs64_tnn_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_128x128x256_0_vs64_tnn_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_128x192x256_0_vs64_tnn_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_128x192x256_0_vs64_tnn_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_128x256x256_0_vs64_tnn_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_128x256x256_0_vs64_tnn_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_256x128x256_0_vs64_tnn_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_256x128x256_0_vs64_tnn_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_256x192x256_0_vs64_tnn_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_256x192x256_0_vs64_tnn_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_256x256x256_0_vs64_tnn_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_256x256x256_0_vs64_tnn_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+580
@@ -0,0 +1,580 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnt_vs64in
|
||||
// 2. 128x192_tnt_vs64in
|
||||
// 3. 128x256_tnt_vs64in
|
||||
// 4. 256x128_tnt_vs64in
|
||||
// 5. 256x192_tnt_vs64in
|
||||
// 6. 256x256_tnt_vs64in
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_128x128x256_0_vs64_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_128x192x256_0_vs64_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_128x256x256_0_vs64_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_256x128x256_0_vs64_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_256x192x256_0_vs64_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_256x256x256_0_vs64_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_128x128x256_0_vs64_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_128x128x256_0_vs64_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_128x192x256_0_vs64_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_128x192x256_0_vs64_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_128x256x256_0_vs64_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_128x256x256_0_vs64_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_256x128x256_0_vs64_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_256x128x256_0_vs64_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_256x192x256_0_vs64_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_256x192x256_0_vs64_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_256x256x256_0_vs64_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_f32_f32_256x256x256_0_vs64_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+580
@@ -0,0 +1,580 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnn_vs64in
|
||||
// 2. 128x192_tnn_vs64in
|
||||
// 3. 128x256_tnn_vs64in
|
||||
// 4. 256x128_tnn_vs64in
|
||||
// 5. 256x192_tnn_vs64in
|
||||
// 6. 256x256_tnn_vs64in
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_128x128x256_0_vs64_tnn_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_128x192x256_0_vs64_tnn_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_128x256x256_0_vs64_tnn_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_256x128x256_0_vs64_tnn_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_256x192x256_0_vs64_tnn_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_256x256x256_0_vs64_tnn_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_128x128x256_0_vs64_tnn_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_128x128x256_0_vs64_tnn_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_128x192x256_0_vs64_tnn_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_128x192x256_0_vs64_tnn_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_128x256x256_0_vs64_tnn_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_128x256x256_0_vs64_tnn_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_256x128x256_0_vs64_tnn_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_256x128x256_0_vs64_tnn_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_256x192x256_0_vs64_tnn_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_256x192x256_0_vs64_tnn_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_256x256x256_0_vs64_tnn_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_256x256x256_0_vs64_tnn_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+580
@@ -0,0 +1,580 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnt_vs64in
|
||||
// 2. 128x192_tnt_vs64in
|
||||
// 3. 128x256_tnt_vs64in
|
||||
// 4. 256x128_tnt_vs64in
|
||||
// 5. 256x192_tnt_vs64in
|
||||
// 6. 256x256_tnt_vs64in
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_128x128x256_0_vs64_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_128x192x256_0_vs64_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_128x256x256_0_vs64_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_256x128x256_0_vs64_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_256x192x256_0_vs64_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_256x256x256_0_vs64_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_128x128x256_0_vs64_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_128x128x256_0_vs64_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_128x192x256_0_vs64_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_128x192x256_0_vs64_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_128x256x256_0_vs64_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_128x256x256_0_vs64_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_256x128x256_0_vs64_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_256x128x256_0_vs64_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_256x192x256_0_vs64_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_256x192x256_0_vs64_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_256x256x256_0_vs64_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f16_256x256x256_0_vs64_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+580
@@ -0,0 +1,580 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnn_vs64in
|
||||
// 2. 128x192_tnn_vs64in
|
||||
// 3. 128x256_tnn_vs64in
|
||||
// 4. 256x128_tnn_vs64in
|
||||
// 5. 256x192_tnn_vs64in
|
||||
// 6. 256x256_tnn_vs64in
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_128x128x256_0_vs64_tnn_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_128x192x256_0_vs64_tnn_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_128x256x256_0_vs64_tnn_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_256x128x256_0_vs64_tnn_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_256x192x256_0_vs64_tnn_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_256x256x256_0_vs64_tnn_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_128x128x256_0_vs64_tnn_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_128x128x256_0_vs64_tnn_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_128x192x256_0_vs64_tnn_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_128x192x256_0_vs64_tnn_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_128x256x256_0_vs64_tnn_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_128x256x256_0_vs64_tnn_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_256x128x256_0_vs64_tnn_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_256x128x256_0_vs64_tnn_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_256x192x256_0_vs64_tnn_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_256x192x256_0_vs64_tnn_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_256x256x256_0_vs64_tnn_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_256x256x256_0_vs64_tnn_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+580
@@ -0,0 +1,580 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnt_vs64in
|
||||
// 2. 128x192_tnt_vs64in
|
||||
// 3. 128x256_tnt_vs64in
|
||||
// 4. 256x128_tnt_vs64in
|
||||
// 5. 256x192_tnt_vs64in
|
||||
// 6. 256x256_tnt_vs64in
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_128x128x256_0_vs64_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_128x192x256_0_vs64_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_128x256x256_0_vs64_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_256x128x256_0_vs64_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_256x192x256_0_vs64_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_256x256x256_0_vs64_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_128x128x256_0_vs64_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_128x128x256_0_vs64_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_128x192x256_0_vs64_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_128x192x256_0_vs64_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_128x256x256_0_vs64_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_128x256x256_0_vs64_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_256x128x256_0_vs64_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_256x128x256_0_vs64_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_256x192x256_0_vs64_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_256x192x256_0_vs64_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_256x256x256_0_vs64_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_f32_256x256x256_0_vs64_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+614
@@ -0,0 +1,614 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test Config
|
||||
// 1. layout {tnn_InputVs64_OutputVs64}
|
||||
// 2. tilesize {128x128, 128x192, 128x256, 256x128, 256x192, 256x256}
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnn_vs64in_vs64out
|
||||
// 2. 128x192_tnn_vs64in_vs64out
|
||||
// 3. 128x256_tnn_vs64in_vs64out
|
||||
// 4. 256x128_tnn_vs64in_vs64out
|
||||
// 5. 256x192_tnn_vs64in_vs64out
|
||||
// 6. 256x256_tnn_vs64in_vs64out
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_128x128x256_0_vs64_tnn_align32_q_1sm_epiVs64n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 16;
|
||||
constexpr int kAlignmentC = 16;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_128x192x256_0_vs64_tnn_align32_q_1sm_epiVs64n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 16;
|
||||
constexpr int kAlignmentC = 16;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_128x256x256_0_vs64_tnn_align32_q_1sm_epiVs64n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 16;
|
||||
constexpr int kAlignmentC = 16;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_256x128x256_0_vs64_tnn_align32_q_2sm_epiVs64n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 16;
|
||||
constexpr int kAlignmentC = 16;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_256x192x256_0_vs64_tnn_align32_q_2sm_epiVs64n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 16;
|
||||
constexpr int kAlignmentC = 16;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_256x256x256_0_vs64_tnn_align32_q_2sm_epiVs64n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 16;
|
||||
constexpr int kAlignmentC = 16;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_128x128x256_0_vs64_tnn_align32_q_1sm_epiVs64n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_128x128x256_0_vs64_tnn_align32_q_1sm_epiVs64n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_128x192x256_0_vs64_tnn_align32_q_1sm_epiVs64n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_128x192x256_0_vs64_tnn_align32_q_1sm_epiVs64n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_128x256x256_0_vs64_tnn_align32_q_1sm_epiVs64n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_128x256x256_0_vs64_tnn_align32_q_1sm_epiVs64n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_256x128x256_0_vs64_tnn_align32_q_2sm_epiVs64n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_256x128x256_0_vs64_tnn_align32_q_2sm_epiVs64n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_256x192x256_0_vs64_tnn_align32_q_2sm_epiVs64n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_256x192x256_0_vs64_tnn_align32_q_2sm_epiVs64n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_256x256x256_0_vs64_tnn_align32_q_2sm_epiVs64n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_256x256x256_0_vs64_tnn_align32_q_2sm_epiVs64n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+614
@@ -0,0 +1,614 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test Config
|
||||
// 1. layout {tnt_InputVs64_OutputVs64}
|
||||
// 2. tilesize {128x128, 128x192, 128x256, 256x128, 256x192, 256x256}
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnt_vs64in_vs64out
|
||||
// 2. 128x192_tnt_vs64in_vs64out
|
||||
// 3. 128x256_tnt_vs64in_vs64out
|
||||
// 4. 256x128_tnt_vs64in_vs64out
|
||||
// 5. 256x192_tnt_vs64in_vs64out
|
||||
// 6. 256x256_tnt_vs64in_vs64out
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_128x128x256_0_vs64_tnt_align32_q_1sm_epiVs64t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_128x192x256_0_vs64_tnt_align32_q_1sm_epiVs64t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_128x256x256_0_vs64_tnt_align32_q_1sm_epiVs64t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_256x128x256_0_vs64_tnt_align32_q_2sm_epiVs64t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_256x192x256_0_vs64_tnt_align32_q_2sm_epiVs64t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_256x256x256_0_vs64_tnt_align32_q_2sm_epiVs64t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementPairB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 64;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmMxf8f6f4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmMxf8f6f4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_128x128x256_0_vs64_tnt_align32_q_1sm_epiVs64t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_128x128x256_0_vs64_tnt_align32_q_1sm_epiVs64t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_128x192x256_0_vs64_tnt_align32_q_1sm_epiVs64t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_128x192x256_0_vs64_tnt_align32_q_1sm_epiVs64t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_128x256x256_0_vs64_tnt_align32_q_1sm_epiVs64t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_128x256x256_0_vs64_tnt_align32_q_1sm_epiVs64t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_256x128x256_0_vs64_tnt_align32_q_2sm_epiVs64t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_256x128x256_0_vs64_tnt_align32_q_2sm_epiVs64t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_256x192x256_0_vs64_tnt_align32_q_2sm_epiVs64t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x192x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_256x192x256_0_vs64_tnt_align32_q_2sm_epiVs64t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_256x256x256_0_vs64_tnt_align32_q_2sm_epiVs64t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x64bsspgemm_ue8m0xe4m3_ue8m0xe4m3_f32_void_ue8m0xe4m3_256x256x256_0_vs64_tnt_align32_q_2sm_epiVs64t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+754
@@ -0,0 +1,754 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnn_vs32in
|
||||
// 2. 128x192_tnn_vs32in
|
||||
// 3. 128x256_tnn_vs32in
|
||||
// 4. 256x128_tnn_vs32in
|
||||
// 5. 256x192_tnn_vs32in
|
||||
// 6. 256x256_tnn_vs32in
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_128x128x256_0_vs32_tnn_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_128x192x256_0_vs32_tnn_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_128x256x256_0_vs32_tnn_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_128x256x512_0_vs32_tnn_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _512>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_256x128x256_0_vs32_tnn_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_256x192x256_0_vs32_tnn_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_256x256x256_0_vs32_tnn_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.2
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_256x256x512_0_vs32_tnn_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _512>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_128x128x256_0_vs32_tnn_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_128x128x256_0_vs32_tnn_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_128x192x256_0_vs32_tnn_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_128x192x256_0_vs32_tnn_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_128x256x256_0_vs32_tnn_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_128x256x256_0_vs32_tnn_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 3.2
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_128x256x512_0_vs32_tnn_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_128x256x512_0_vs32_tnn_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{512, 3840}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_256x128x256_0_vs32_tnn_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_256x128x256_0_vs32_tnn_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_256x192x256_0_vs32_tnn_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_256x192x256_0_vs32_tnn_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_256x256x256_0_vs32_tnn_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_256x256x256_0_vs32_tnn_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 6.2
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_256x256x512_0_vs32_tnn_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_256x256x512_0_vs32_tnn_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{512, 3840}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+754
@@ -0,0 +1,754 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnt_vs32in
|
||||
// 2. 128x192_tnt_vs32in
|
||||
// 3. 128x256_tnt_vs32in
|
||||
// 4. 256x128_tnt_vs32in
|
||||
// 5. 256x192_tnt_vs32in
|
||||
// 6. 256x256_tnt_vs32in
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_128x128x256_0_vs32_tnt_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_128x192x256_0_vs32_tnt_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_128x256x256_0_vs32_tnt_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.2
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_128x256x512_0_vs32_tnt_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _512>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_256x128x256_0_vs32_tnt_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_256x192x256_0_vs32_tnt_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_256x256x256_0_vs32_tnt_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.2
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_256x256x512_0_vs32_tnt_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _512>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_128x128x256_0_vs32_tnt_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_128x128x256_0_vs32_tnt_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_128x192x256_0_vs32_tnt_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_128x192x256_0_vs32_tnt_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_128x256x256_0_vs32_tnt_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_128x256x256_0_vs32_tnt_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 3.2
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_128x256x512_0_vs32_tnt_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_128x256x512_0_vs32_tnt_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{512, 3840}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_256x128x256_0_vs32_tnt_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_256x128x256_0_vs32_tnt_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_256x192x256_0_vs32_tnt_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_256x192x256_0_vs32_tnt_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_256x256x256_0_vs32_tnt_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_256x256x256_0_vs32_tnt_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 6.2
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_256x256x512_0_vs32_tnt_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_f16_256x256x512_0_vs32_tnt_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{512, 3840}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+797
@@ -0,0 +1,797 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test Config
|
||||
// 1. layout {tnn_InputVs32_OutputVs32}
|
||||
// 2. tilesize {128x128, 128x192, 128x256, 256x128, 256x192, 256x256}
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnn_vs32in_vs32out
|
||||
// 2. 128x192_tnn_vs32in_vs32out
|
||||
// 3. 128x256_tnn_vs32in_vs32out
|
||||
// 4. 256x128_tnn_vs32in_vs32out
|
||||
// 5. 256x192_tnn_vs32in_vs32out
|
||||
// 6. 256x256_tnn_vs32in_vs32out
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_128x128x256_0_vs32_tnn_align64_o_1sm_epiVs32n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_128x192x256_0_vs32_tnn_align64_o_1sm_epiVs32n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_128x256x256_0_vs32_tnn_align64_o_1sm_epiVs32n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.2
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_128x256x512_0_vs32_tnn_align64_o_1sm_epiVs32n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _512>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_256x128x256_0_vs32_tnn_align64_o_2sm_epiVs32n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_256x192x256_0_vs32_tnn_align64_o_2sm_epiVs32n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_256x256x256_0_vs32_tnn_align64_o_2sm_epiVs32n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_256x256x512_0_vs32_tnn_align64_o_2sm_epiVs32n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _512>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_128x128x256_0_vs32_tnn_align64_o_1sm_epiVs32n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_128x128x256_0_vs32_tnn_align64_o_1sm_epiVs32n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_128x192x256_0_vs32_tnn_align64_o_1sm_epiVs32n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_128x192x256_0_vs32_tnn_align64_o_1sm_epiVs32n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_128x256x256_0_vs32_tnn_align64_o_1sm_epiVs32n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_128x256x256_0_vs32_tnn_align64_o_1sm_epiVs32n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 3.2
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_128x256x512_0_vs32_tnn_align64_o_1sm_epiVs32n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_128x256x512_0_vs32_tnn_align64_o_1sm_epiVs32n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{512, 3840}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_256x128x256_0_vs32_tnn_align64_o_2sm_epiVs32n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_256x128x256_0_vs32_tnn_align64_o_2sm_epiVs32n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_256x192x256_0_vs32_tnn_align64_o_2sm_epiVs32n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_256x192x256_0_vs32_tnn_align64_o_2sm_epiVs32n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_256x256x256_0_vs32_tnn_align64_o_2sm_epiVs32n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_256x256x256_0_vs32_tnn_align64_o_2sm_epiVs32n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 6.2
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_256x256x512_0_vs32_tnn_align64_o_2sm_epiVs32n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_256x256x512_0_vs32_tnn_align64_o_2sm_epiVs32n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{512, 3840}));
|
||||
}
|
||||
#endif
|
||||
+798
@@ -0,0 +1,798 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test Config
|
||||
// 1. layout {tnt_InputVs32_OutputVs32}
|
||||
// 2. tilesize {128x128, 128x192, 128x256, 256x128, 256x192, 256x256}
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnt_vs32in_vs32out
|
||||
// 2. 128x192_tnt_vs32in_vs32out
|
||||
// 3. 128x256_tnt_vs32in_vs32out
|
||||
// 4. 256x128_tnt_vs32in_vs32out
|
||||
// 5. 256x192_tnt_vs32in_vs32out
|
||||
// 6. 256x256_tnt_vs32in_vs32out
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_128x128x256_0_vs32_tnt_align64_o_1sm_epiVs32t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_128x192x256_0_vs32_tnt_align64_o_1sm_epiVs32t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_128x256x256_0_vs32_tnt_align64_o_1sm_epiVs32t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.2
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_128x256x512_0_vs32_tnt_align64_o_1sm_epiVs32t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _512>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_256x128x256_0_vs32_tnt_align64_o_2sm_epiVs32t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_256x192x256_0_vs32_tnt_align64_o_2sm_epiVs32t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_256x256x256_0_vs32_tnt_align64_o_2sm_epiVs32t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_256x256x512_0_vs32_tnt_align64_o_2sm_epiVs32t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _512>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_128x128x256_0_vs32_tnt_align64_o_1sm_epiVs32t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_128x128x256_0_vs32_tnt_align64_o_1sm_epiVs32t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_128x192x256_0_vs32_tnt_align64_o_1sm_epiVs32t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_128x192x256_0_vs32_tnt_align64_o_1sm_epiVs32t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_128x256x256_0_vs32_tnt_align64_o_1sm_epiVs32t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_128x256x256_0_vs32_tnt_align64_o_1sm_epiVs32t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 3.2
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_128x256x512_0_vs32_tnt_align64_o_1sm_epiVs32t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_128x256x512_0_vs32_tnt_align64_o_1sm_epiVs32t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{512, 3840}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_256x128x256_0_vs32_tnt_align64_o_2sm_epiVs32t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_256x128x256_0_vs32_tnt_align64_o_2sm_epiVs32t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_256x192x256_0_vs32_tnt_align64_o_2sm_epiVs32t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_256x192x256_0_vs32_tnt_align64_o_2sm_epiVs32t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_256x256x256_0_vs32_tnt_align64_o_2sm_epiVs32t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_256x256x256_0_vs32_tnt_align64_o_2sm_epiVs32t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_256x256x512_0_vs32_tnt_align64_o_2sm_epiVs32t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_ue4m3xe2m1_256x256x512_0_vs32_tnt_align64_o_2sm_epiVs32t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{512, 3840}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+580
@@ -0,0 +1,580 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnt_vs32in
|
||||
// 2. 128x192_tnt_vs32in
|
||||
// 3. 128x256_tnt_vs32in
|
||||
// 4. 256x128_tnt_vs32in
|
||||
// 5. 256x192_tnt_vs32in
|
||||
// 6. 256x256_tnt_vs32in
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_e2m1_128x128x256_0_vs32_tnt_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::StreamKScheduler>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_e2m1_128x192x256_0_vs32_tnt_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::StreamKScheduler>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_e2m1_128x256x256_0_vs32_tnt_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::StreamKScheduler>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_e2m1_256x128x256_0_vs32_tnt_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::StreamKScheduler>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_e2m1_256x192x256_0_vs32_tnt_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::StreamKScheduler>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_e2m1_256x256x256_0_vs32_tnt_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::StreamKScheduler>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_e2m1_128x128x256_0_vs32_tnt_align64_o_1sm, streamk) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_e2m1_128x128x256_0_vs32_tnt_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{1536}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_e2m1_128x192x256_0_vs32_tnt_align64_o_1sm, streamk) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_e2m1_128x192x256_0_vs32_tnt_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{1536}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_e2m1_128x256x256_0_vs32_tnt_align64_o_1sm, streamk) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_e2m1_128x256x256_0_vs32_tnt_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{1536}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_e2m1_256x128x256_0_vs32_tnt_align64_o_2sm, streamk) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_e2m1_256x128x256_0_vs32_tnt_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{1536}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_e2m1_256x192x256_0_vs32_tnt_align64_o_2sm, streamk) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_e2m1_256x192x256_0_vs32_tnt_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{1536}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_e2m1_256x256x256_0_vs32_tnt_align64_o_2sm, streamk) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f16_e2m1_256x256x256_0_vs32_tnt_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{1536}));
|
||||
}
|
||||
|
||||
#endif // #if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
+754
@@ -0,0 +1,754 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnn_vs32in
|
||||
// 2. 128x192_tnn_vs32in
|
||||
// 3. 128x256_tnn_vs32in
|
||||
// 4. 256x128_tnn_vs32in
|
||||
// 5. 256x192_tnn_vs32in
|
||||
// 6. 256x256_tnn_vs32in
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_128x128x256_0_vs32_tnn_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_128x192x256_0_vs32_tnn_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_128x256x256_0_vs32_tnn_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.2
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_128x256x512_0_vs32_tnn_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _512>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_256x128x256_0_vs32_tnn_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_256x192x256_0_vs32_tnn_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_256x256x256_0_vs32_tnn_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_256x256x512_0_vs32_tnn_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _512>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_128x128x256_0_vs32_tnn_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_128x128x256_0_vs32_tnn_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_128x192x256_0_vs32_tnn_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_128x192x256_0_vs32_tnn_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_128x256x256_0_vs32_tnn_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_128x256x256_0_vs32_tnn_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 3.2
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_128x256x512_0_vs32_tnn_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_128x256x512_0_vs32_tnn_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{512, 3840}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_256x128x256_0_vs32_tnn_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_256x128x256_0_vs32_tnn_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_256x192x256_0_vs32_tnn_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_256x192x256_0_vs32_tnn_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_256x256x256_0_vs32_tnn_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_256x256x256_0_vs32_tnn_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 6.2
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_256x256x512_0_vs32_tnn_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_256x256x512_0_vs32_tnn_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{512, 3840}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+756
@@ -0,0 +1,756 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnt_vs32in
|
||||
// 2. 128x192_tnt_vs32in
|
||||
// 3. 128x256_tnt_vs32in
|
||||
// 4. 256x128_tnt_vs32in
|
||||
// 5. 256x192_tnt_vs32in
|
||||
// 6. 256x256_tnt_vs32in
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_128x128x256_0_vs32_tnt_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_128x192x256_0_vs32_tnt_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_128x256x256_0_vs32_tnt_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.2
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_128x256x512_0_vs32_tnt_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _512>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_256x128x256_0_vs32_tnt_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_256x192x256_0_vs32_tnt_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_256x256x256_0_vs32_tnt_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.2
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_256x256x512_0_vs32_tnt_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _512>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_128x128x256_0_vs32_tnt_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_128x128x256_0_vs32_tnt_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_128x192x256_0_vs32_tnt_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_128x192x256_0_vs32_tnt_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_128x256x256_0_vs32_tnt_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_128x256x256_0_vs32_tnt_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 3.2
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_128x256x512_0_vs32_tnt_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_128x256x512_0_vs32_tnt_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{384, 3840}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_256x128x256_0_vs32_tnt_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_256x128x256_0_vs32_tnt_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_256x192x256_0_vs32_tnt_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_256x192x256_0_vs32_tnt_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_256x256x256_0_vs32_tnt_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_256x256x256_0_vs32_tnt_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 6.2
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_256x256x512_0_vs32_tnt_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_f32_f32_256x256x512_0_vs32_tnt_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{512, 3840}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+754
@@ -0,0 +1,754 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnn_vs32in
|
||||
// 2. 128x192_tnn_vs32in
|
||||
// 3. 128x256_tnn_vs32in
|
||||
// 4. 256x128_tnn_vs32in
|
||||
// 5. 256x192_tnn_vs32in
|
||||
// 6. 256x256_tnn_vs32in
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_128x128x256_0_vs32_tnn_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_128x192x256_0_vs32_tnn_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cute::Shape<_128, _8>;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_128x256x256_0_vs32_tnn_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_128x256x512_0_vs32_tnn_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _512>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_256x128x256_0_vs32_tnn_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cute::Shape<_128, _8>;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_256x192x256_0_vs32_tnn_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cute::Shape<_128, _8>;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_256x256x256_0_vs32_tnn_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.2
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_256x256x512_0_vs32_tnn_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _512>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_128x128x256_0_vs32_tnn_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_128x128x256_0_vs32_tnn_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_128x192x256_0_vs32_tnn_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_128x192x256_0_vs32_tnn_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_128x256x256_0_vs32_tnn_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_128x256x256_0_vs32_tnn_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 3.2
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_128x256x512_0_vs32_tnn_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_128x256x512_0_vs32_tnn_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{512, 3840}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_256x128x256_0_vs32_tnn_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_256x128x256_0_vs32_tnn_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_256x192x256_0_vs32_tnn_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_256x192x256_0_vs32_tnn_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_256x256x256_0_vs32_tnn_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_256x256x256_0_vs32_tnn_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 6.2
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_256x256x512_0_vs32_tnn_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_256x256x512_0_vs32_tnn_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{512, 3840}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+754
@@ -0,0 +1,754 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnt_vs32in
|
||||
// 2. 128x192_tnt_vs32in
|
||||
// 3. 128x256_tnt_vs32in
|
||||
// 4. 256x128_tnt_vs32in
|
||||
// 5. 256x192_tnt_vs32in
|
||||
// 6. 256x256_tnt_vs32in
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_128x128x256_0_vs32_tnt_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_128x192x256_0_vs32_tnt_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cute::Shape<_128, _8>;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_128x256x256_0_vs32_tnt_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.2
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_128x256x512_0_vs32_tnt_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _512>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_256x128x256_0_vs32_tnt_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cute::Shape<_128, _8>;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_256x192x256_0_vs32_tnt_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cute::Shape<_128, _8>;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_256x256x256_0_vs32_tnt_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.2
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_256x256x512_0_vs32_tnt_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _512>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_128x128x256_0_vs32_tnt_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_128x128x256_0_vs32_tnt_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_128x192x256_0_vs32_tnt_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_128x192x256_0_vs32_tnt_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_128x256x256_0_vs32_tnt_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_128x256x256_0_vs32_tnt_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 3.2
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_128x256x512_0_vs32_tnt_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_128x256x512_0_vs32_tnt_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{512, 3840}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_256x128x256_0_vs32_tnt_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_256x128x256_0_vs32_tnt_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_256x192x256_0_vs32_tnt_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_256x192x256_0_vs32_tnt_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_256x256x256_0_vs32_tnt_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_256x256x256_0_vs32_tnt_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 6.2
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_256x256x512_0_vs32_tnt_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f16_256x256x512_0_vs32_tnt_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{512, 3840}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+754
@@ -0,0 +1,754 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnn_vs32in
|
||||
// 2. 128x192_tnn_vs32in
|
||||
// 3. 128x256_tnn_vs32in
|
||||
// 4. 256x128_tnn_vs32in
|
||||
// 5. 256x192_tnn_vs32in
|
||||
// 6. 256x256_tnn_vs32in
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_128x128x256_0_vs32_tnn_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_128x192x256_0_vs32_tnn_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_128x256x256_0_vs32_tnn_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.2
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_128x256x512_0_vs32_tnn_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _512>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_256x128x256_0_vs32_tnn_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_256x192x256_0_vs32_tnn_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_256x256x256_0_vs32_tnn_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_256x256x512_0_vs32_tnn_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _512>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_128x128x256_0_vs32_tnn_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_128x128x256_0_vs32_tnn_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_128x192x256_0_vs32_tnn_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_128x192x256_0_vs32_tnn_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_128x256x256_0_vs32_tnn_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_128x256x256_0_vs32_tnn_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 3.2
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_128x256x512_0_vs32_tnn_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_128x256x512_0_vs32_tnn_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{512, 3840}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_256x128x256_0_vs32_tnn_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_256x128x256_0_vs32_tnn_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_256x192x256_0_vs32_tnn_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_256x192x256_0_vs32_tnn_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_256x256x256_0_vs32_tnn_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_256x256x256_0_vs32_tnn_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 6.2
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_256x256x512_0_vs32_tnn_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_256x256x512_0_vs32_tnn_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{512, 3840}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+756
@@ -0,0 +1,756 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnt_vs32in
|
||||
// 2. 128x192_tnt_vs32in
|
||||
// 3. 128x256_tnt_vs32in
|
||||
// 4. 256x128_tnt_vs32in
|
||||
// 5. 256x192_tnt_vs32in
|
||||
// 6. 256x256_tnt_vs32in
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_128x128x256_0_vs32_tnt_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_128x192x256_0_vs32_tnt_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_128x256x256_0_vs32_tnt_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.2
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_128x256x512_0_vs32_tnt_align64_o_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _512>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_256x128x256_0_vs32_tnt_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_256x192x256_0_vs32_tnt_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_256x256x256_0_vs32_tnt_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.2
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_256x256x512_0_vs32_tnt_align64_o_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _512>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_128x128x256_0_vs32_tnt_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_128x128x256_0_vs32_tnt_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_128x192x256_0_vs32_tnt_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_128x192x256_0_vs32_tnt_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_128x256x256_0_vs32_tnt_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_128x256x256_0_vs32_tnt_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 3.2
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_128x256x512_0_vs32_tnt_align64_o_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_128x256x512_0_vs32_tnt_align64_o_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{384, 3840}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_256x128x256_0_vs32_tnt_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_256x128x256_0_vs32_tnt_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_256x192x256_0_vs32_tnt_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_256x192x256_0_vs32_tnt_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_256x256x256_0_vs32_tnt_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_256x256x256_0_vs32_tnt_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 6.2
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_256x256x512_0_vs32_tnt_align64_o_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_f32_256x256x512_0_vs32_tnt_align64_o_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{512, 3840}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+798
@@ -0,0 +1,798 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test Config
|
||||
// 1. layout {tnn_InputVs32_OutputVs32}
|
||||
// 2. tilesize {128x128, 128x192, 128x256, 256x128, 256x192, 256x256}
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnn_vs32in_vs32out
|
||||
// 2. 128x192_tnn_vs32in_vs32out
|
||||
// 3. 128x256_tnn_vs32in_vs32out
|
||||
// 4. 256x128_tnn_vs32in_vs32out
|
||||
// 5. 256x192_tnn_vs32in_vs32out
|
||||
// 6. 256x256_tnn_vs32in_vs32out
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_128x128x256_0_vs32_tnn_align64_o_1sm_epiVs32n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_128x192x256_0_vs32_tnn_align64_o_1sm_epiVs32n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_128x256x256_0_vs32_tnn_align64_o_1sm_epiVs32n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.2
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_128x256x512_0_vs32_tnn_align64_o_1sm_epiVs32n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _512>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_256x128x256_0_vs32_tnn_align64_o_2sm_epiVs32n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_256x192x256_0_vs32_tnn_align64_o_2sm_epiVs32n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_256x256x256_0_vs32_tnn_align64_o_2sm_epiVs32n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.2
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_256x256x512_0_vs32_tnn_align64_o_2sm_epiVs32n {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _512>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_128x128x256_0_vs32_tnn_align64_o_1sm_epiVs32n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_128x128x256_0_vs32_tnn_align64_o_1sm_epiVs32n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_128x192x256_0_vs32_tnn_align64_o_1sm_epiVs32n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_128x192x256_0_vs32_tnn_align64_o_1sm_epiVs32n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_128x256x256_0_vs32_tnn_align64_o_1sm_epiVs32n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_128x256x256_0_vs32_tnn_align64_o_1sm_epiVs32n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 3.2
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_128x256x512_0_vs32_tnn_align64_o_1sm_epiVs32n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_128x256x512_0_vs32_tnn_align64_o_1sm_epiVs32n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{512, 3840}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_256x128x256_0_vs32_tnn_align64_o_2sm_epiVs32n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_256x128x256_0_vs32_tnn_align64_o_2sm_epiVs32n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_256x192x256_0_vs32_tnn_align64_o_2sm_epiVs32n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_256x192x256_0_vs32_tnn_align64_o_2sm_epiVs32n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_256x256x256_0_vs32_tnn_align64_o_2sm_epiVs32n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_256x256x256_0_vs32_tnn_align64_o_2sm_epiVs32n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 6.2
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_256x256x512_0_vs32_tnn_align64_o_2sm_epiVs32n, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_256x256x512_0_vs32_tnn_align64_o_2sm_epiVs32n;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{512, 3840}));
|
||||
}
|
||||
|
||||
#endif
|
||||
+798
@@ -0,0 +1,798 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test Config
|
||||
// 1. layout {tnt_InputVs32_OutputVs32}
|
||||
// 2. tilesize {128x128, 128x192, 128x256, 256x128, 256x192, 256x256}
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnt_vs32in_vs32out
|
||||
// 2. 128x192_tnt_vs32in_vs32out
|
||||
// 3. 128x256_tnt_vs32in_vs32out
|
||||
// 4. 256x128_tnt_vs32in_vs32out
|
||||
// 5. 256x192_tnt_vs32in_vs32out
|
||||
// 6. 256x256_tnt_vs32in_vs32out
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_128x128x256_0_vs32_tnt_align64_o_1sm_epiVs32t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentC = 32;
|
||||
constexpr int kAlignmentD = 32;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_128x192x256_0_vs32_tnt_align64_o_1sm_epiVs32t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentC = 32;
|
||||
constexpr int kAlignmentD = 32;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_128x256x256_0_vs32_tnt_align64_o_1sm_epiVs32t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentC = 32;
|
||||
constexpr int kAlignmentD = 32;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.2
|
||||
namespace cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_128x256x512_0_vs32_tnt_align64_o_1sm_epiVs32t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentC = 32;
|
||||
constexpr int kAlignmentD = 32;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _512>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_256x128x256_0_vs32_tnt_align64_o_2sm_epiVs32t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentC = 32;
|
||||
constexpr int kAlignmentD = 32;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 5.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_256x192x256_0_vs32_tnt_align64_o_2sm_epiVs32t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentC = 32;
|
||||
constexpr int kAlignmentD = 32;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _192, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_256x256x256_0_vs32_tnt_align64_o_2sm_epiVs32t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentC = 32;
|
||||
constexpr int kAlignmentD = 32;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 6.2
|
||||
namespace cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_256x256x512_0_vs32_tnt_align64_o_2sm_epiVs32t {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 64;
|
||||
constexpr int kAlignmentB = 32;
|
||||
constexpr int kAlignmentC = 32;
|
||||
constexpr int kAlignmentD = 32;
|
||||
constexpr int SFDVectorSize = 32;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _512>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassTag = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2SmNvf4;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmNvf4Sm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType,
|
||||
cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFDVectorSize,
|
||||
ElementD,
|
||||
ElementEpilogueCompute,
|
||||
ElementSF,
|
||||
LayoutD,
|
||||
ElementBias,
|
||||
ElementC,
|
||||
ElementEpilogueCompute>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassTag,
|
||||
ElementPairA, LayoutA, kAlignmentA,
|
||||
ElementPairB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_128x128x256_0_vs32_tnt_align64_o_1sm_epiVs32t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_128x128x256_0_vs32_tnt_align64_o_1sm_epiVs32t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_128x192x256_0_vs32_tnt_align64_o_1sm_epiVs32t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_128x192x256_0_vs32_tnt_align64_o_1sm_epiVs32t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_128x256x256_0_vs32_tnt_align64_o_1sm_epiVs32t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_128x256x256_0_vs32_tnt_align64_o_1sm_epiVs32t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 3.2
|
||||
TEST(cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_128x256x512_0_vs32_tnt_align64_o_1sm_epiVs32t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s128x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_128x256x512_0_vs32_tnt_align64_o_1sm_epiVs32t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{512, 3840}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_256x128x256_0_vs32_tnt_align64_o_2sm_epiVs32t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_256x128x256_0_vs32_tnt_align64_o_2sm_epiVs32t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 5.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_256x192x256_0_vs32_tnt_align64_o_2sm_epiVs32t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x128x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_256x192x256_0_vs32_tnt_align64_o_2sm_epiVs32t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 6.
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_256x256x256_0_vs32_tnt_align64_o_2sm_epiVs32t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_256x256x256_0_vs32_tnt_align64_o_2sm_epiVs32t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 3840}));
|
||||
}
|
||||
|
||||
// 6.2
|
||||
TEST(cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_256x256x512_0_vs32_tnt_align64_o_2sm_epiVs32t, sfd_fusion) {
|
||||
namespace gemm = cutlass3x_sm100_bssptensorop_s256x256x128bsspgemm_ue4m3xe2m1_ue4m3xe2m1_f32_void_ue4m3xe2m1_256x256x512_0_vs32_tnt_align64_o_2sm_epiVs32t;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmallFusion<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{512, 3840}));
|
||||
}
|
||||
|
||||
#endif
|
||||
@@ -26,10 +26,6 @@
|
||||
# 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.
|
||||
|
||||
#
|
||||
|
||||
#
|
||||
|
||||
if (CUTLASS_NVCC_ARCHS MATCHES 100a)
|
||||
add_custom_target(
|
||||
cutlass_test_unit_gemm_device_sm100_blockscaled
|
||||
|
||||
@@ -112,6 +112,60 @@ TEST(SM100_Device_Gemm_e4m3t_e4m3n_f32t_tensorop_2sm_f32_auto, 512x512x128_4x4x1
|
||||
EXPECT_TRUE(pass);
|
||||
}
|
||||
|
||||
/// A Row B Col
|
||||
TEST(SM100_Device_Gemm_e4m3t_e4m3n_f32t_tensorop_2sm_f32_auto, 512x256x128_4x4x1) {
|
||||
using ElementA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
using ElementCompute = float;
|
||||
using ElementAccumulator = float;
|
||||
using ElementAccumulator = float;
|
||||
using GmemLayoutA = cutlass::layout::RowMajor;
|
||||
using GmemLayoutB = cutlass::layout::ColumnMajor;
|
||||
using GmemLayoutC = cutlass::layout::RowMajor;
|
||||
using MmaTileShape_MNK = Shape<_256,_64,_128>;
|
||||
using ClusterShape_MNK = Shape<_4,_4,_1>;
|
||||
|
||||
//
|
||||
// Construct CollectiveEpilogue
|
||||
//
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassTensorOp,
|
||||
MmaTileShape_MNK, ClusterShape_MNK,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, GmemLayoutC, 4,
|
||||
ElementD, GmemLayoutC, 4,
|
||||
cutlass::epilogue::collective::EpilogueScheduleAuto
|
||||
>::CollectiveOp;
|
||||
|
||||
//
|
||||
// Construct CollectiveMainloop
|
||||
//
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementA, GmemLayoutA, 16,
|
||||
ElementB, GmemLayoutB, 16,
|
||||
ElementAccumulator,
|
||||
MmaTileShape_MNK, ClusterShape_MNK,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::collective::KernelScheduleAuto
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
auto pass = test::gemm::device::TestSmallFusion<Gemm>(1.0, 0);
|
||||
EXPECT_TRUE(pass);
|
||||
}
|
||||
|
||||
|
||||
/// A Col B Row
|
||||
TEST(SM100_Device_Gemm_e4m3n_e4m3t_f32t_tensorop_2sm_f32_auto, 512x512x128_4x4x1) {
|
||||
using ElementA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
|
||||
@@ -227,6 +227,62 @@ TEST(SM100Only_Device_Gemm_e4m3t_e4m3t_f32t_tensorop_2sm_f32_group, 512x512x128_
|
||||
EXPECT_TRUE(pass);
|
||||
}
|
||||
|
||||
/// A Row B Row
|
||||
TEST(SM100Only_Device_Gemm_e4m3t_e4m3t_f32t_tensorop_2sm_f32_group, 512x256x128_4x4x1) {
|
||||
using ElementA = cutlass::float_e4m3_t;
|
||||
using ElementB = cutlass::float_e4m3_t;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
using ElementCompute = float;
|
||||
using ElementAccumulator = float;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
using MmaTypePairA = decltype(cute::make_tuple(ElementA{}, ElementSF{}));
|
||||
using MmaTypePairB = decltype(cute::make_tuple(ElementB{}, ElementSF{}));
|
||||
using ElementAccumulator = float;
|
||||
using GmemLayoutA = cutlass::layout::RowMajor;
|
||||
using GmemLayoutB = cutlass::layout::RowMajor;
|
||||
using GmemLayoutC = cutlass::layout::RowMajor;
|
||||
using MmaTileShape_MNK = Shape<_256,_64,_128>;
|
||||
using ClusterShape_MNK = Shape<_4,_4,_1>;
|
||||
|
||||
//
|
||||
// Construct CollectiveEpilogue
|
||||
//
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassTensorOp,
|
||||
MmaTileShape_MNK, ClusterShape_MNK,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, GmemLayoutC *, 16,
|
||||
ElementD, GmemLayoutC *, 16,
|
||||
cutlass::epilogue::PtrArrayTmaWarpSpecialized2Sm
|
||||
>::CollectiveOp;
|
||||
|
||||
//
|
||||
// Construct CollectiveMainloop
|
||||
//
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
MmaTypePairA, GmemLayoutA *, 16,
|
||||
MmaTypePairB, GmemLayoutB *, 16,
|
||||
ElementAccumulator,
|
||||
MmaTileShape_MNK, ClusterShape_MNK,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelPtrArrayTmaWarpSpecialized2SmMxf8f6f4Sm100
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cutlass::gemm::GroupProblemShape<Shape<int,int,int>>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
auto pass = test::gemm::device::TestSmall<Gemm>(1.0, 0);
|
||||
EXPECT_TRUE(pass);
|
||||
}
|
||||
|
||||
/// A Col B Col
|
||||
TEST(SM100Only_Device_Gemm_e4m3n_e4m3n_f32t_tensorop_2sm_f32_group, 512x512x128_4x4x1) {
|
||||
using ElementA = cutlass::float_e4m3_t;
|
||||
|
||||
@@ -0,0 +1,80 @@
|
||||
# Copyright (c) 2025 - 2025 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.
|
||||
|
||||
if (CUTLASS_NVCC_ARCHS MATCHES 100a)
|
||||
|
||||
add_custom_target(
|
||||
cutlass_test_unit_gemm_device_sm100_sp
|
||||
DEPENDS
|
||||
cutlass_test_unit_gemm_device_sm100_sp_general
|
||||
cutlass_test_unit_gemm_device_sm100_sp_qmma_variance
|
||||
cutlass_test_unit_gemm_device_sm100_sp_streamk
|
||||
)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable_split_file(
|
||||
cutlass_test_unit_gemm_device_sm100_sp_general
|
||||
|
||||
# No batching of source to control compiler memory usage
|
||||
BATCH_SOURCES ON
|
||||
BATCH_SIZE 1
|
||||
|
||||
sm100_sp_gemm_s8_s8_s32_s8_s8_imma.cu
|
||||
sm100_sp_gemm_f8_f8_f32_f16_f8_qmma.cu
|
||||
sm100_sp_gemm_f8_f8_f32_f16_f16_qmma.cu
|
||||
sm100_sp_gemm_f8_f8_f32_f32_f32_qmma.cu
|
||||
sm100_sp_gemm_f32_f32_f32_f32_f32_tfmma.cu
|
||||
sm100_sp_gemm_f16_f16_f32_f16_f16_hmma.cu
|
||||
)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable_split_file(
|
||||
cutlass_test_unit_gemm_device_sm100_sp_qmma_variance
|
||||
|
||||
# No batching of source to control compiler memory usage
|
||||
BATCH_SOURCES ON
|
||||
BATCH_SIZE 1
|
||||
|
||||
sm100_sp_gemm_f4_f4_f32_f16_f8_qmma.cu
|
||||
sm100_sp_gemm_f4_f4_f32_f16_f16_qmma.cu
|
||||
sm100_sp_gemm_f4_f4_f32_f32_f32_qmma.cu
|
||||
|
||||
sm100_sp_gemm_f6_f6_f32_f16_f8_qmma.cu
|
||||
sm100_sp_gemm_f6_f6_f32_f16_f16_qmma.cu
|
||||
sm100_sp_gemm_f6_f6_f32_f32_f32_qmma.cu
|
||||
)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable_split_file(
|
||||
cutlass_test_unit_gemm_device_sm100_sp_streamk
|
||||
|
||||
# No batching of source to control compiler memory usage
|
||||
BATCH_SOURCES ON
|
||||
BATCH_SIZE 1
|
||||
|
||||
sm100_sp_gemm_f16_f16_f32_f32_f32_streamk.cu
|
||||
)
|
||||
|
||||
endif()
|
||||
+1258
File diff suppressed because it is too large
Load Diff
+565
@@ -0,0 +1,565 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
///////////////////////////////////////////////////// 128x128x64 //////////////////////////////////////////////////////
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
TEST(SM100_Device_Sparse_Gemm_f16t_f16n_f32t_tensorop_1cta_f32_streamk, 128x128x64_1x1x1) {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using MmaTileShape = Shape<_128,_128,_64>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
constexpr int ALIGNMENT_C = 4;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassSparseTensorOp,
|
||||
MmaTileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
float, float,
|
||||
float, LayoutC, ALIGNMENT_C,
|
||||
float, LayoutC, ALIGNMENT_C,
|
||||
cutlass::epilogue::TmaWarpSpecialized1Sm
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassSparseTensorOp,
|
||||
cutlass::half_t, LayoutATag, 16,
|
||||
cutlass::half_t, LayoutBTag, 8,
|
||||
float,
|
||||
MmaTileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>,
|
||||
cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::StreamKScheduler
|
||||
>;
|
||||
|
||||
using namespace test::gemm::device;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
bool result = TestSmall<Gemm>(1.0, 0.0, CheckEquality::EXACT, ScalarLoc::ON_DEVICE, VectorScale::ENABLED, {64, 1536});
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
TEST(SM100_Device_Sparse_Gemm_f16n_f16t_f32n_tensorop_1cta_f32_streamk, 256x256x64_2x2x1) {
|
||||
using LayoutATag = cutlass::layout::ColumnMajor;
|
||||
using LayoutBTag = cutlass::layout::RowMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using MmaTileShape = Shape<_128,_128,_64>;
|
||||
using ClusterShape = Shape<_2,_2,_1>;
|
||||
|
||||
constexpr int ALIGNMENT_C = 4;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassSparseTensorOp,
|
||||
MmaTileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
float, float,
|
||||
float, LayoutC, ALIGNMENT_C,
|
||||
float, LayoutC, ALIGNMENT_C,
|
||||
cutlass::epilogue::TmaWarpSpecialized1Sm
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassSparseTensorOp,
|
||||
cutlass::half_t, LayoutATag, 16,
|
||||
cutlass::half_t, LayoutBTag, 8,
|
||||
float,
|
||||
MmaTileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>,
|
||||
cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::StreamKScheduler
|
||||
>;
|
||||
|
||||
using namespace test::gemm::device;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
bool result = TestSmall<Gemm>(1.0, 0.0, CheckEquality::EXACT, ScalarLoc::ON_DEVICE, VectorScale::ENABLED, {64, 1536});
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
|
||||
TEST(SM100_Device_Sparse_Gemm_f16t_f16n_f32t_tensorop_2cta_f32_streamk, 256x256x64_2x2x1) {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using MmaTileShape = Shape<_256,_128,_64>;
|
||||
using ClusterShape = Shape<_2,_2,_1>;
|
||||
|
||||
constexpr int ALIGNMENT_C = 4;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassSparseTensorOp,
|
||||
MmaTileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
float, float,
|
||||
float, LayoutC, ALIGNMENT_C,
|
||||
float, LayoutC, ALIGNMENT_C,
|
||||
cutlass::epilogue::TmaWarpSpecialized2Sm
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassSparseTensorOp,
|
||||
cutlass::half_t, LayoutATag, 16,
|
||||
cutlass::half_t, LayoutBTag, 8,
|
||||
float,
|
||||
MmaTileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>,
|
||||
cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::StreamKScheduler
|
||||
>;
|
||||
|
||||
using namespace test::gemm::device;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
bool result = TestSmall<Gemm>(1.0, 0.0, CheckEquality::EXACT, ScalarLoc::ON_DEVICE, VectorScale::ENABLED, {64, 1536});
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
TEST(SM100_Device_Sparse_Gemm_f16n_f16t_f32n_tensorop_2cta_f32_streamk, 512x512x64_4x4x1) {
|
||||
using LayoutATag = cutlass::layout::ColumnMajor;
|
||||
using LayoutBTag = cutlass::layout::RowMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using MmaTileShape = Shape<_256,_128,_64>;
|
||||
using ClusterShape = Shape<_4,_4,_1>;
|
||||
|
||||
constexpr int ALIGNMENT_C = 4;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassSparseTensorOp,
|
||||
MmaTileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
float, float,
|
||||
float, LayoutC, ALIGNMENT_C,
|
||||
float, LayoutC, ALIGNMENT_C,
|
||||
cutlass::epilogue::TmaWarpSpecialized2Sm
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassSparseTensorOp,
|
||||
cutlass::half_t, LayoutATag, 16,
|
||||
cutlass::half_t, LayoutBTag, 8,
|
||||
float,
|
||||
MmaTileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>,
|
||||
cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::StreamKScheduler
|
||||
>;
|
||||
|
||||
using namespace test::gemm::device;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
bool result = TestSmall<Gemm>(1.0, 0.0, CheckEquality::EXACT, ScalarLoc::ON_DEVICE, VectorScale::ENABLED, {64, 1536});
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
TEST(SM100_Device_Sparse_Gemm_f16t_f16n_f32n_tensorop_2cta_f32_streamk, 256x512x128_2x4x1) {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using MmaTileShape = Shape<_256,_128,_128>;
|
||||
using ClusterShape = Shape<_2,_4,_1>;
|
||||
|
||||
constexpr int ALIGNMENT_C = 4;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassSparseTensorOp,
|
||||
MmaTileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
float, float,
|
||||
float, LayoutC, ALIGNMENT_C,
|
||||
float, LayoutC, ALIGNMENT_C,
|
||||
cutlass::epilogue::TmaWarpSpecialized2Sm
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassSparseTensorOp,
|
||||
cutlass::half_t, LayoutATag, 16,
|
||||
cutlass::half_t, LayoutBTag, 8,
|
||||
float,
|
||||
MmaTileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>,
|
||||
cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::StreamKScheduler
|
||||
>;
|
||||
|
||||
using namespace test::gemm::device;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
bool result = TestSmall<Gemm>(1.0, 0.0, CheckEquality::EXACT, ScalarLoc::ON_DEVICE, VectorScale::ENABLED, {64, 1536});
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
TEST(SM100_Device_Sparse_Gemm_f16t_f16n_f32t_tensorop_2cta_f32_streamk, 256x256x64_2x1x1) {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using MmaTileShape = Shape<_256,_256,_64>;
|
||||
using ClusterShape = Shape<_2,_1,_1>;
|
||||
|
||||
constexpr int ALIGNMENT_C = 4;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassSparseTensorOp,
|
||||
MmaTileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
float, float,
|
||||
float, LayoutC, ALIGNMENT_C,
|
||||
float, LayoutC, ALIGNMENT_C,
|
||||
cutlass::epilogue::TmaWarpSpecialized2Sm
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassSparseTensorOp,
|
||||
cutlass::half_t, LayoutATag, 16,
|
||||
cutlass::half_t, LayoutBTag, 8,
|
||||
float,
|
||||
MmaTileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>,
|
||||
cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::StreamKScheduler
|
||||
>;
|
||||
|
||||
using namespace test::gemm::device;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
bool result = TestSmall<Gemm>(1.0, 0.0, CheckEquality::EXACT, ScalarLoc::ON_DEVICE, VectorScale::ENABLED, {64, 1536});
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
// Enable this after linearized scheduler is functional again.
|
||||
#if 0
|
||||
TEST(SM100_Device_Sparse_Gemm_f16t_f16n_f32t_tensorop_1cta_f32_linearized, 128x128x64_1x1x1) {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using MmaTileShape = Shape<_128,_128,_64>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
constexpr int ALIGNMENT_C = 4;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassSparseTensorOp,
|
||||
MmaTileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
float, float,
|
||||
float, LayoutC, ALIGNMENT_C,
|
||||
float, LayoutC, ALIGNMENT_C,
|
||||
cutlass::epilogue::TmaWarpSpecialized1Sm
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassSparseTensorOp,
|
||||
cutlass::half_t, LayoutATag, 16,
|
||||
cutlass::half_t, LayoutBTag, 8,
|
||||
float,
|
||||
MmaTileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>,
|
||||
cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::LinearizedScheduler
|
||||
>;
|
||||
|
||||
using namespace test::gemm::device;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
bool result = TestSmall<Gemm>(1.0, 0.0, CheckEquality::EXACT, ScalarLoc::ON_DEVICE, VectorScale::ENABLED, {64, 1024, 2048});
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
TEST(SM100_Device_Sparse_Gemm_f16n_f16t_f32n_tensorop_1cta_f32_linearized, 256x256x64_2x2x1) {
|
||||
using LayoutATag = cutlass::layout::ColumnMajor;
|
||||
using LayoutBTag = cutlass::layout::RowMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using MmaTileShape = Shape<_128,_128,_64>;
|
||||
using ClusterShape = Shape<_2,_2,_1>;
|
||||
|
||||
constexpr int ALIGNMENT_C = 4;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassSparseTensorOp,
|
||||
MmaTileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
float, float,
|
||||
float, LayoutC, ALIGNMENT_C,
|
||||
float, LayoutC, ALIGNMENT_C,
|
||||
cutlass::epilogue::TmaWarpSpecialized1Sm
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassSparseTensorOp,
|
||||
cutlass::half_t, LayoutATag, 16,
|
||||
cutlass::half_t, LayoutBTag, 8,
|
||||
float,
|
||||
MmaTileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>,
|
||||
cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::LinearizedScheduler
|
||||
>;
|
||||
|
||||
using namespace test::gemm::device;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
bool result = TestSmall<Gemm>(1.0, 0.0, CheckEquality::EXACT, ScalarLoc::ON_DEVICE, VectorScale::ENABLED, {64, 1024, 2048});
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
|
||||
TEST(SM100_Device_Sparse_Gemm_f16t_f16n_f32t_tensorop_2cta_f32_linearized, 256x256x64_2x2x1) {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using MmaTileShape = Shape<_256,_128,_64>;
|
||||
using ClusterShape = Shape<_2,_2,_1>;
|
||||
|
||||
constexpr int ALIGNMENT_C = 4;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassSparseTensorOp,
|
||||
MmaTileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
float, float,
|
||||
float, LayoutC, ALIGNMENT_C,
|
||||
float, LayoutC, ALIGNMENT_C,
|
||||
cutlass::epilogue::TmaWarpSpecialized2Sm
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassSparseTensorOp,
|
||||
cutlass::half_t, LayoutATag, 16,
|
||||
cutlass::half_t, LayoutBTag, 8,
|
||||
float,
|
||||
MmaTileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>,
|
||||
cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::LinearizedScheduler
|
||||
>;
|
||||
|
||||
using namespace test::gemm::device;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
bool result = TestSmall<Gemm>(1.0, 0.0, CheckEquality::EXACT, ScalarLoc::ON_DEVICE, VectorScale::ENABLED, {64, 1024, 2048});
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
TEST(SM100_Device_Sparse_Gemm_f16n_f16t_f32n_tensorop_2cta_f32_linearized, 512x512x64_4x4x1) {
|
||||
using LayoutATag = cutlass::layout::ColumnMajor;
|
||||
using LayoutBTag = cutlass::layout::RowMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using MmaTileShape = Shape<_256,_128,_64>;
|
||||
using ClusterShape = Shape<_4,_4,_1>;
|
||||
|
||||
constexpr int ALIGNMENT_C = 4;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassSparseTensorOp,
|
||||
MmaTileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
float, float,
|
||||
float, LayoutC, ALIGNMENT_C,
|
||||
float, LayoutC, ALIGNMENT_C,
|
||||
cutlass::epilogue::TmaWarpSpecialized2Sm
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassSparseTensorOp,
|
||||
cutlass::half_t, LayoutATag, 16,
|
||||
cutlass::half_t, LayoutBTag, 8,
|
||||
float,
|
||||
MmaTileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>,
|
||||
cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::LinearizedScheduler
|
||||
>;
|
||||
|
||||
using namespace test::gemm::device;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
bool result = TestSmall<Gemm>(1.0, 0.0, CheckEquality::EXACT, ScalarLoc::ON_DEVICE, VectorScale::ENABLED, {64, 1024, 2048});
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
TEST(SM100_Device_Sparse_Gemm_f16t_f16n_f32n_tensorop_2cta_f32_linearized, 256x512x128_2x4x1) {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using MmaTileShape = Shape<_256,_128,_128>;
|
||||
using ClusterShape = Shape<_2,_4,_1>;
|
||||
|
||||
constexpr int ALIGNMENT_C = 4;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassSparseTensorOp,
|
||||
MmaTileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
float, float,
|
||||
float, LayoutC, ALIGNMENT_C,
|
||||
float, LayoutC, ALIGNMENT_C,
|
||||
cutlass::epilogue::TmaWarpSpecialized2Sm
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassSparseTensorOp,
|
||||
cutlass::half_t, LayoutATag, 16,
|
||||
cutlass::half_t, LayoutBTag, 8,
|
||||
float,
|
||||
MmaTileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>,
|
||||
cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::LinearizedScheduler
|
||||
>;
|
||||
|
||||
using namespace test::gemm::device;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
bool result = TestSmall<Gemm>(1.0, 0.0, CheckEquality::EXACT, ScalarLoc::ON_DEVICE, VectorScale::ENABLED, {64, 1024, 2048});
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
TEST(SM100_Device_Sparse_Gemm_f16t_f16n_f32t_tensorop_2cta_f32_linearized, 256x256x64_2x1x1) {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using MmaTileShape = Shape<_256,_256,_64>;
|
||||
using ClusterShape = Shape<_2,_1,_1>;
|
||||
|
||||
constexpr int ALIGNMENT_C = 4;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassSparseTensorOp,
|
||||
MmaTileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
float, float,
|
||||
float, LayoutC, ALIGNMENT_C,
|
||||
float, LayoutC, ALIGNMENT_C,
|
||||
cutlass::epilogue::TmaWarpSpecialized2Sm
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm100, cutlass::arch::OpClassSparseTensorOp,
|
||||
cutlass::half_t, LayoutATag, 16,
|
||||
cutlass::half_t, LayoutBTag, 8,
|
||||
float,
|
||||
MmaTileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>,
|
||||
cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::LinearizedScheduler
|
||||
>;
|
||||
|
||||
using namespace test::gemm::device;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
bool result = TestSmall<Gemm>(1.0, 0.0, CheckEquality::EXACT, ScalarLoc::ON_DEVICE, VectorScale::ENABLED, {64, 1024, 2048});
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
#endif
|
||||
|
||||
#endif // #if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
+1259
File diff suppressed because it is too large
Load Diff
+705
@@ -0,0 +1,705 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnt
|
||||
// 2. 128x256_tnt
|
||||
// 3. 256x128_tnt
|
||||
// 4. 256x256_tnt
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_sptensorop_s128x128x64spgemm_e2m1_e2m1_f32_f16_f16_128x128x256_0_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_sptensorop_s128x256x64spgemm_e2m1_e2m1_f32_f16_f16_128x256x256_0_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_sptensorop_s256x128x64spgemm_e2m1_e2m1_f32_f16_f16_256x128x256_0_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_sptensorop_s256x256x64spgemm_e2m1_e2m1_f32_f16_f16_256x256x256_0_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_sptensorop_s128x128x64spgemm_e2m1_e2m1_f32_f16_f16_128x128x256_0_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s128x128x64spgemm_e2m1_e2m1_f32_f16_f16_128x128x256_0_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_sptensorop_s128x256x64spgemm_e2m1_e2m1_f32_f16_f16_128x256x256_0_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s128x256x64spgemm_e2m1_e2m1_f32_f16_f16_128x256x256_0_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_sptensorop_s256x128x64spgemm_e2m1_e2m1_f32_f16_f16_256x128x256_0_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s256x128x64spgemm_e2m1_e2m1_f32_f16_f16_256x128x256_0_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_sptensorop_s256x256x64spgemm_e2m1_e2m1_f32_f16_f16_256x256x256_0_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s256x256x64spgemm_e2m1_e2m1_f32_f16_f16_256x256x256_0_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_sptensorop_s128x128x64spgemm_e2m1_e2m1_f32_void_f16_128x128x256_0_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_sptensorop_s128x256x64spgemm_e2m1_e2m1_f32_void_f16_128x256x256_0_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_sptensorop_s256x128x64spgemm_e2m1_e2m1_f32_void_f16_256x128x256_0_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_sptensorop_s256x256x64spgemm_e2m1_e2m1_f32_void_f16_256x256x256_0_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_sptensorop_s128x128x64spgemm_e2m1_e2m1_f32_void_f16_128x128x256_0_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s128x128x64spgemm_e2m1_e2m1_f32_void_f16_128x128x256_0_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_sptensorop_s128x256x64spgemm_e2m1_e2m1_f32_void_f16_128x256x256_0_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s128x256x64spgemm_e2m1_e2m1_f32_void_f16_128x256x256_0_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_sptensorop_s256x128x64spgemm_e2m1_e2m1_f32_void_f16_256x128x256_0_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s256x128x64spgemm_e2m1_e2m1_f32_void_f16_256x128x256_0_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_sptensorop_s256x256x64spgemm_e2m1_e2m1_f32_void_f16_256x256x256_0_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s256x256x64spgemm_e2m1_e2m1_f32_void_f16_256x256x256_0_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
#endif // #if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
+705
@@ -0,0 +1,705 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnt
|
||||
// 2. 128x256_tnt
|
||||
// 3. 256x128_tnt
|
||||
// 4. 256x256_tnt
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_sptensorop_s128x128x64spgemm_e2m1_e2m1_f32_f16_e4m3_128x128x256_0_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_sptensorop_s128x256x64spgemm_e2m1_e2m1_f32_f16_e4m3_128x256x256_0_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_sptensorop_s256x128x64spgemm_e2m1_e2m1_f32_f16_e4m3_256x128x256_0_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_sptensorop_s256x256x64spgemm_e2m1_e2m1_f32_f16_e4m3_256x256x256_0_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_sptensorop_s128x128x64spgemm_e2m1_e2m1_f32_f16_e4m3_128x128x256_0_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s128x128x64spgemm_e2m1_e2m1_f32_f16_e4m3_128x128x256_0_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_sptensorop_s128x256x64spgemm_e2m1_e2m1_f32_f16_e4m3_128x256x256_0_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s128x256x64spgemm_e2m1_e2m1_f32_f16_e4m3_128x256x256_0_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_sptensorop_s256x128x64spgemm_e2m1_e2m1_f32_f16_e4m3_256x128x256_0_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s256x128x64spgemm_e2m1_e2m1_f32_f16_e4m3_256x128x256_0_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_sptensorop_s256x256x64spgemm_e2m1_e2m1_f32_f16_e4m3_256x256x256_0_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s256x256x64spgemm_e2m1_e2m1_f32_f16_e4m3_256x256x256_0_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_sptensorop_s128x128x64spgemm_e2m1_e2m1_f32_void_e4m3_128x128x256_0_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_sptensorop_s128x256x64spgemm_e2m1_e2m1_f32_void_e4m3_128x256x256_0_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_sptensorop_s256x128x64spgemm_e2m1_e2m1_f32_void_e4m3_256x128x256_0_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_sptensorop_s256x256x64spgemm_e2m1_e2m1_f32_void_e4m3_256x256x256_0_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_sptensorop_s128x128x64spgemm_e2m1_e2m1_f32_void_e4m3_128x128x256_0_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s128x128x64spgemm_e2m1_e2m1_f32_void_e4m3_128x128x256_0_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_sptensorop_s128x256x64spgemm_e2m1_e2m1_f32_void_e4m3_128x256x256_0_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s128x256x64spgemm_e2m1_e2m1_f32_void_e4m3_128x256x256_0_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_sptensorop_s256x128x64spgemm_e2m1_e2m1_f32_void_e4m3_256x128x256_0_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s256x128x64spgemm_e2m1_e2m1_f32_void_e4m3_256x128x256_0_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_sptensorop_s256x256x64spgemm_e2m1_e2m1_f32_void_e4m3_256x256x256_0_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s256x256x64spgemm_e2m1_e2m1_f32_void_e4m3_256x256x256_0_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
#endif // #if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
+705
@@ -0,0 +1,705 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnt
|
||||
// 2. 128x256_tnt
|
||||
// 3. 256x128_tnt
|
||||
// 4. 256x256_tnt
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_sptensorop_s128x128x64spgemm_e2m1_e2m1_f32_f32_f32_128x128x256_0_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_sptensorop_s128x256x64spgemm_e2m1_e2m1_f32_f32_f32_128x256x256_0_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_sptensorop_s256x128x64spgemm_e2m1_e2m1_f32_f32_f32_256x128x256_0_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_sptensorop_s256x256x64spgemm_e2m1_e2m1_f32_f32_f32_256x256x256_0_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_sptensorop_s128x128x64spgemm_e2m1_e2m1_f32_f32_f32_128x128x256_0_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s128x128x64spgemm_e2m1_e2m1_f32_f32_f32_128x128x256_0_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_sptensorop_s128x256x64spgemm_e2m1_e2m1_f32_f32_f32_128x256x256_0_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s128x256x64spgemm_e2m1_e2m1_f32_f32_f32_128x256x256_0_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_sptensorop_s256x128x64spgemm_e2m1_e2m1_f32_f32_f32_256x128x256_0_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s256x128x64spgemm_e2m1_e2m1_f32_f32_f32_256x128x256_0_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_sptensorop_s256x256x64spgemm_e2m1_e2m1_f32_f32_f32_256x256x256_0_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s256x256x64spgemm_e2m1_e2m1_f32_f32_f32_256x256x256_0_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_sptensorop_s128x128x64spgemm_e2m1_e2m1_f32_void_f32_128x128x256_0_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_sptensorop_s128x256x64spgemm_e2m1_e2m1_f32_void_f32_128x256x256_0_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_sptensorop_s256x128x64spgemm_e2m1_e2m1_f32_void_f32_256x128x256_0_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_sptensorop_s256x256x64spgemm_e2m1_e2m1_f32_void_f32_256x256x256_0_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_sptensorop_s128x128x64spgemm_e2m1_e2m1_f32_void_f32_128x128x256_0_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s128x128x64spgemm_e2m1_e2m1_f32_void_f32_128x128x256_0_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_sptensorop_s128x256x64spgemm_e2m1_e2m1_f32_void_f32_128x256x256_0_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s128x256x64spgemm_e2m1_e2m1_f32_void_f32_128x256x256_0_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_sptensorop_s256x128x64spgemm_e2m1_e2m1_f32_void_f32_256x128x256_0_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s256x128x64spgemm_e2m1_e2m1_f32_void_f32_256x128x256_0_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_sptensorop_s256x256x64spgemm_e2m1_e2m1_f32_void_f32_256x256x256_0_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s256x256x64spgemm_e2m1_e2m1_f32_void_f32_256x256x256_0_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
#endif // #if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
+705
@@ -0,0 +1,705 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnt
|
||||
// 2. 128x256_tnt
|
||||
// 3. 256x128_tnt
|
||||
// 4. 256x256_tnt
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_sptensorop_s128x128x64spgemm_e3m2_e3m2_f32_f16_f16_128x128x256_0_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_sptensorop_s128x256x64spgemm_e3m2_e3m2_f32_f16_f16_128x256x256_0_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_sptensorop_s256x128x64spgemm_e3m2_e3m2_f32_f16_f16_256x128x256_0_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_sptensorop_s256x256x64spgemm_e3m2_e3m2_f32_f16_f16_256x256x256_0_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_sptensorop_s128x128x64spgemm_e3m2_e3m2_f32_f16_f16_128x128x256_0_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s128x128x64spgemm_e3m2_e3m2_f32_f16_f16_128x128x256_0_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_sptensorop_s128x256x64spgemm_e3m2_e3m2_f32_f16_f16_128x256x256_0_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s128x256x64spgemm_e3m2_e3m2_f32_f16_f16_128x256x256_0_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_sptensorop_s256x128x64spgemm_e3m2_e3m2_f32_f16_f16_256x128x256_0_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s256x128x64spgemm_e3m2_e3m2_f32_f16_f16_256x128x256_0_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_sptensorop_s256x256x64spgemm_e3m2_e3m2_f32_f16_f16_256x256x256_0_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s256x256x64spgemm_e3m2_e3m2_f32_f16_f16_256x256x256_0_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_sptensorop_s128x128x64spgemm_e3m2_e3m2_f32_void_f16_128x128x256_0_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_sptensorop_s128x256x64spgemm_e3m2_e3m2_f32_void_f16_128x256x256_0_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_sptensorop_s256x128x64spgemm_e3m2_e3m2_f32_void_f16_256x128x256_0_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_sptensorop_s256x256x64spgemm_e3m2_e3m2_f32_void_f16_256x256x256_0_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_sptensorop_s128x128x64spgemm_e3m2_e3m2_f32_void_f16_128x128x256_0_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s128x128x64spgemm_e3m2_e3m2_f32_void_f16_128x128x256_0_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_sptensorop_s128x256x64spgemm_e3m2_e3m2_f32_void_f16_128x256x256_0_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s128x256x64spgemm_e3m2_e3m2_f32_void_f16_128x256x256_0_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_sptensorop_s256x128x64spgemm_e3m2_e3m2_f32_void_f16_256x128x256_0_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s256x128x64spgemm_e3m2_e3m2_f32_void_f16_256x128x256_0_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_sptensorop_s256x256x64spgemm_e3m2_e3m2_f32_void_f16_256x256x256_0_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s256x256x64spgemm_e3m2_e3m2_f32_void_f16_256x256x256_0_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
#endif // #if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
+705
@@ -0,0 +1,705 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnt
|
||||
// 2. 128x256_tnt
|
||||
// 3. 256x128_tnt
|
||||
// 4. 256x256_tnt
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_sptensorop_s128x128x64spgemm_e3m2_e3m2_f32_f16_e4m3_128x128x256_0_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_sptensorop_s128x256x64spgemm_e3m2_e3m2_f32_f16_e4m3_128x256x256_0_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_sptensorop_s256x128x64spgemm_e3m2_e3m2_f32_f16_e4m3_256x128x256_0_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_sptensorop_s256x256x64spgemm_e3m2_e3m2_f32_f16_e4m3_256x256x256_0_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_sptensorop_s128x128x64spgemm_e3m2_e3m2_f32_f16_e4m3_128x128x256_0_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s128x128x64spgemm_e3m2_e3m2_f32_f16_e4m3_128x128x256_0_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_sptensorop_s128x256x64spgemm_e3m2_e3m2_f32_f16_e4m3_128x256x256_0_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s128x256x64spgemm_e3m2_e3m2_f32_f16_e4m3_128x256x256_0_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_sptensorop_s256x128x64spgemm_e3m2_e3m2_f32_f16_e4m3_256x128x256_0_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s256x128x64spgemm_e3m2_e3m2_f32_f16_e4m3_256x128x256_0_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_sptensorop_s256x256x64spgemm_e3m2_e3m2_f32_f16_e4m3_256x256x256_0_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s256x256x64spgemm_e3m2_e3m2_f32_f16_e4m3_256x256x256_0_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_sptensorop_s128x128x64spgemm_e3m2_e3m2_f32_void_e4m3_128x128x256_0_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_sptensorop_s128x256x64spgemm_e3m2_e3m2_f32_void_e4m3_128x256x256_0_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_sptensorop_s256x128x64spgemm_e3m2_e3m2_f32_void_e4m3_256x128x256_0_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_sptensorop_s256x256x64spgemm_e3m2_e3m2_f32_void_e4m3_256x256x256_0_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_sptensorop_s128x128x64spgemm_e3m2_e3m2_f32_void_e4m3_128x128x256_0_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s128x128x64spgemm_e3m2_e3m2_f32_void_e4m3_128x128x256_0_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_sptensorop_s128x256x64spgemm_e3m2_e3m2_f32_void_e4m3_128x256x256_0_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s128x256x64spgemm_e3m2_e3m2_f32_void_e4m3_128x256x256_0_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_sptensorop_s256x128x64spgemm_e3m2_e3m2_f32_void_e4m3_256x128x256_0_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s256x128x64spgemm_e3m2_e3m2_f32_void_e4m3_256x128x256_0_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_sptensorop_s256x256x64spgemm_e3m2_e3m2_f32_void_e4m3_256x256x256_0_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s256x256x64spgemm_e3m2_e3m2_f32_void_e4m3_256x256x256_0_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
#endif // #if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
+705
@@ -0,0 +1,705 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
using namespace cute;
|
||||
|
||||
// * Test list
|
||||
// 1. 128x128_tnt
|
||||
// 2. 128x256_tnt
|
||||
// 3. 256x128_tnt
|
||||
// 4. 256x256_tnt
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_sptensorop_s128x128x64spgemm_e3m2_e3m2_f32_f32_f32_128x128x256_0_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_sptensorop_s128x256x64spgemm_e3m2_e3m2_f32_f32_f32_128x256x256_0_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_sptensorop_s256x128x64spgemm_e3m2_e3m2_f32_f32_f32_256x128x256_0_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_sptensorop_s256x256x64spgemm_e3m2_e3m2_f32_f32_f32_256x256x256_0_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_sptensorop_s128x128x64spgemm_e3m2_e3m2_f32_f32_f32_128x128x256_0_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s128x128x64spgemm_e3m2_e3m2_f32_f32_f32_128x128x256_0_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_sptensorop_s128x256x64spgemm_e3m2_e3m2_f32_f32_f32_128x256x256_0_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s128x256x64spgemm_e3m2_e3m2_f32_f32_f32_128x256x256_0_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_sptensorop_s256x128x64spgemm_e3m2_e3m2_f32_f32_f32_256x128x256_0_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s256x128x64spgemm_e3m2_e3m2_f32_f32_f32_256x128x256_0_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_sptensorop_s256x256x64spgemm_e3m2_e3m2_f32_f32_f32_256x256x256_0_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s256x256x64spgemm_e3m2_e3m2_f32_f32_f32_256x256x256_0_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 1,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
|
||||
// 1.
|
||||
namespace cutlass3x_sm100_sptensorop_s128x128x64spgemm_e3m2_e3m2_f32_void_f32_128x128x256_0_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 2.
|
||||
namespace cutlass3x_sm100_sptensorop_s128x256x64spgemm_e3m2_e3m2_f32_void_f32_128x256x256_0_tnt_align32_q_1sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_1, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_128, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized1Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized1SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 3.
|
||||
namespace cutlass3x_sm100_sptensorop_s256x128x64spgemm_e3m2_e3m2_f32_void_f32_256x128x256_0_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _128, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 4.
|
||||
namespace cutlass3x_sm100_sptensorop_s256x256x64spgemm_e3m2_e3m2_f32_void_f32_256x256x256_0_tnt_align32_q_2sm {
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
|
||||
constexpr int kAlignmentA = 256;
|
||||
constexpr int kAlignmentB = 128;
|
||||
constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
constexpr int kAlignmentC = cute::is_same_v<ElementC, void> ? kAlignmentD : 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
|
||||
using MmaTileShape = Shape<_256, _256, _256>;
|
||||
using ArchTag = cutlass::arch::Sm100;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassSparseTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecialized2SmSm100;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = float;
|
||||
using TileScheduler = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutC, kAlignmentC,
|
||||
ElementD, LayoutD, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveoutEpi<CollectiveEpilogue>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
ElementB, LayoutB, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MmaTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
cute::Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
void>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
}
|
||||
|
||||
// 1.
|
||||
TEST(cutlass3x_sm100_sptensorop_s128x128x64spgemm_e3m2_e3m2_f32_void_f32_128x128x256_0_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s128x128x64spgemm_e3m2_e3m2_f32_void_f32_128x128x256_0_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 2.
|
||||
TEST(cutlass3x_sm100_sptensorop_s128x256x64spgemm_e3m2_e3m2_f32_void_f32_128x256x256_0_tnt_align32_q_1sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s128x256x64spgemm_e3m2_e3m2_f32_void_f32_128x256x256_0_tnt_align32_q_1sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 3.
|
||||
TEST(cutlass3x_sm100_sptensorop_s256x128x64spgemm_e3m2_e3m2_f32_void_f32_256x128x256_0_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s256x128x64spgemm_e3m2_e3m2_f32_void_f32_256x128x256_0_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
// 4.
|
||||
TEST(cutlass3x_sm100_sptensorop_s256x256x64spgemm_e3m2_e3m2_f32_void_f32_256x256x256_0_tnt_align32_q_2sm, functional) {
|
||||
namespace gemm = cutlass3x_sm100_sptensorop_s256x256x64spgemm_e3m2_e3m2_f32_void_f32_256x256x256_0_tnt_align32_q_2sm;
|
||||
EXPECT_TRUE(test::gemm::device::TestSmall<gemm::Gemm>(
|
||||
1, 0,
|
||||
test::gemm::device::CheckEquality::RELATIVE,
|
||||
test::gemm::device::ScalarLoc::ON_DEVICE,
|
||||
test::gemm::device::VectorScale::ENABLED,
|
||||
{256, 2560}));
|
||||
}
|
||||
|
||||
#endif // #if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
|
||||
+1260
File diff suppressed because it is too large
Load Diff
+1260
File diff suppressed because it is too large
Load Diff
+1254
File diff suppressed because it is too large
Load Diff
+1259
File diff suppressed because it is too large
Load Diff
@@ -26,9 +26,7 @@
|
||||
# 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.
|
||||
|
||||
#
|
||||
|
||||
#
|
||||
if (CUTLASS_NVCC_ARCHS MATCHES 100a)
|
||||
add_custom_target(
|
||||
cutlass_test_unit_gemm_device_sm100_tensorop
|
||||
|
||||
@@ -29,8 +29,6 @@
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
|
||||
@@ -26,9 +26,7 @@
|
||||
# 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.
|
||||
|
||||
#
|
||||
|
||||
#
|
||||
if (CUTLASS_NVCC_ARCHS MATCHES 100a)
|
||||
add_custom_target(
|
||||
cutlass_test_unit_gemm_device_sm100_tensorop_narrow_precision
|
||||
|
||||
@@ -0,0 +1,67 @@
|
||||
# Copyright (c) 2024 - 2025 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.
|
||||
|
||||
if (CUTLASS_NVCC_ARCHS MATCHES 120a)
|
||||
|
||||
add_custom_target(
|
||||
cutlass_test_unit_gemm_device_sm120_bssp
|
||||
DEPENDS
|
||||
cutlass_test_unit_gemm_device_sm120_bssp_general
|
||||
cutlass_test_unit_gemm_device_sm120_bssp_stream_k
|
||||
cutlass_test_unit_gemm_device_sm120_bssp_epilogue_fusion
|
||||
)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable(
|
||||
cutlass_test_unit_gemm_device_sm120_bssp_general
|
||||
|
||||
sm120_bssp_gemm_f4_f4_f32_tensor_op.cu
|
||||
sm120_bssp_gemm_f6_f4_f32_tensor_op.cu
|
||||
sm120_bssp_gemm_f8_f6_f32_tensor_op.cu
|
||||
sm120_bssp_gemm_f4t_f4n_f4t_tensor_op.cu
|
||||
sm120_bssp_gemm_f8t_f8n_f8t_tensor_op.cu
|
||||
)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable(
|
||||
cutlass_test_unit_gemm_device_sm120_bssp_epilogue_fusion
|
||||
|
||||
# No batching of source to control compiler memory usage
|
||||
BATCH_SOURCES ON
|
||||
BATCH_SIZE 1
|
||||
sm120_bssp_gemm_f4_f4_f32_tensor_op_epilogue_fusion.cu
|
||||
)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable(
|
||||
cutlass_test_unit_gemm_device_sm120_bssp_stream_k
|
||||
|
||||
# No batching of source to control compiler memory usage
|
||||
BATCH_SOURCES ON
|
||||
BATCH_SIZE 1
|
||||
sm120_bssp_gemm_f4_f4_f32_tensor_op_f32_stream_k.cu
|
||||
)
|
||||
|
||||
endif()
|
||||
+251
@@ -0,0 +1,251 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace kernel_1 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 64 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 64 bytes with sparse ratio 4:2.
|
||||
static constexpr int AlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_64,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledSparseTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelSparseTmaWarpSpecializedMxf8f6f4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
} // kernel_1
|
||||
|
||||
namespace kernel_2 {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes with sparse ratio 4:2.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledSparseTensorOp,
|
||||
ElementPairA, LayoutATag, AlignmentA,
|
||||
ElementPairB, LayoutBTag, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelSparseTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_2
|
||||
|
||||
namespace kernel_3 {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes with sparse ratio 4:2.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledSparseTensorOp,
|
||||
ElementPairA, LayoutATag, AlignmentA,
|
||||
ElementPairB, LayoutBTag, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelSparseTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_3
|
||||
|
||||
TEST(SM120_Device_Sparse_BlockScaled_Gemm_fe2m1t_fe2m1n_f32n_tensor_op_f32, 128x64x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_1::Gemm, false /*force_legacy_epilogue*/, false /*apply_alignment_offset*/>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
TEST(SM120_Device_Sparse_BlockScaled_VS32_Gemm_e2m1t_e2m1n_f32n_tensorop_op_f32, 128x128x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_2::Gemm, false /*force_legacy_epilogue*/, false /*apply_alignment_offset*/, true /*batched_gemm*/>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
TEST(SM120_Device_Sparse_BlockScaled_VS64_Gemm_e2m1t_e2m1n_f32n_tensorop_op_f32, 128x128x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_3::Gemm, false /*force_legacy_epilogue*/, false /*apply_alignment_offset*/, true /*batched_gemm*/>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
+592
@@ -0,0 +1,592 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// D = gelu(alpha * accum + beta * C + per-row bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: fp32
|
||||
namespace kernel_1 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutCTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutDTag = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::GELU, ElementD, ElementCompute, ElementBias, ElementAccumulator, ElementCompute>;
|
||||
|
||||
static constexpr int AlignmentA = 64 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 64 bytes with sparse ratio 4:2.
|
||||
static constexpr int AlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_64,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutCTag, AlignmentC,
|
||||
ElementD, LayoutDTag, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledSparseTensorOp,
|
||||
ElementPairA, LayoutATag, AlignmentA,
|
||||
ElementPairB, LayoutBTag, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelSparseTmaWarpSpecializedMxf8f6f4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
} // kernel_1
|
||||
|
||||
// D = gelu(alpha * accum + beta * C + per-col bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: fp32
|
||||
namespace kernel_2 {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutCTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutDTag = cutlass::layout::ColumnMajor;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes with sparse ratio 4:2.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerColBiasEltAct<
|
||||
cutlass::epilogue::thread::GELU, ElementD, ElementCompute, ElementBias, ElementAccumulator, ElementCompute>;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutCTag, AlignmentC,
|
||||
ElementD, LayoutDTag, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledSparseTensorOp,
|
||||
ElementPairA, LayoutATag, AlignmentA,
|
||||
ElementPairB, LayoutBTag, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelSparseTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_2
|
||||
|
||||
// D = clamp(alpha * accum + beta * C + per-col bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: fp32
|
||||
namespace kernel_3 {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutCTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutDTag = cutlass::layout::ColumnMajor;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes with sparse ratio 4:2.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerColBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp, ElementD, ElementCompute, ElementBias, ElementAccumulator, ElementCompute>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutCTag, AlignmentC,
|
||||
ElementD, LayoutDTag, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledSparseTensorOp,
|
||||
ElementPairA, LayoutATag, AlignmentA,
|
||||
ElementPairB, LayoutBTag, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelSparseTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_3
|
||||
|
||||
namespace kernel_4 {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutCTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutDTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutSFDTag = cutlass::layout::ColumnMajor;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
|
||||
static constexpr int AlignmentA = 64 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 64 bytes with sparse ratio 4:2.
|
||||
static constexpr int AlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
constexpr int SFVectorSize = 64;
|
||||
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerColBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFVectorSize,
|
||||
ElementD,
|
||||
ElementCompute,
|
||||
ElementSF,
|
||||
LayoutSFDTag,
|
||||
ElementBias,
|
||||
ElementC
|
||||
>;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutCTag, AlignmentC,
|
||||
ElementD, LayoutDTag, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledSparseTensorOp,
|
||||
ElementPairA, LayoutATag, AlignmentA,
|
||||
ElementPairB, LayoutBTag, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelSparseTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_4
|
||||
|
||||
// D = clamp(alpha * accum + beta * C + per-row bias)
|
||||
// C: fp16
|
||||
// Acc: fp32
|
||||
// Bias: fp16
|
||||
// Scale (alpha, beta): fp32
|
||||
// Scale factor: fp8
|
||||
// D: fp4 with SF VEC32
|
||||
namespace kernel_5 {
|
||||
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutCTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutDTag = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
|
||||
static constexpr int kAlignmentA = 64 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 64 bytes with sparse ratio 4:2.
|
||||
static constexpr int kAlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int kAlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using ClusterShape = Shape<cute::Int<1>, cute::Int<1>, cute::Int<1>>;
|
||||
using MainloopTileShape = Shape<cute::Int<128>, cute::Int<128>, cute::Int<256>>;
|
||||
using EpilogueTileShape = Shape<cute::Int<128>, cute::Int<128>, cute::Int<256>>;
|
||||
using ArchTag = cutlass::arch::Sm120;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecializedNvf4Sm120;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileSchedulerTag = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
EpilogueTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutCTag, kAlignmentC,
|
||||
ElementD, LayoutDTag, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
, cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp, 32, ElementD, float, cutlass::float_ue4m3_t, LayoutDTag, cutlass::half_t, ElementC, float>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutATag, kAlignmentA,
|
||||
ElementB, LayoutBTag, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MainloopTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
template <class T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
} // kernel_5
|
||||
|
||||
|
||||
// D = alpha * accum + beta * C + per-row bias
|
||||
// C: fp16
|
||||
// Acc: fp32
|
||||
// Bias: fp16
|
||||
// Scale (alpha, beta): fp32
|
||||
// Scale factor: fp8
|
||||
// D: fp4 with SF VEC32
|
||||
namespace kernel_6 {
|
||||
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutCTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutDTag = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
|
||||
static constexpr int kAlignmentA = 64 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 64 bytes with sparse ratio 4:2.
|
||||
static constexpr int kAlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int kAlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int kAlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using ClusterShape = Shape<cute::Int<1>, cute::Int<1>, cute::Int<1>>;
|
||||
using MainloopTileShape = Shape<cute::Int<128>, cute::Int<128>, cute::Int<256>>;
|
||||
using EpilogueTileShape = Shape<cute::Int<128>, cute::Int<128>, cute::Int<256>>;
|
||||
using ArchTag = cutlass::arch::Sm120;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecializedNvf4Sm120;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::half_t;
|
||||
using TileSchedulerTag = cutlass::gemm::PersistentScheduler;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
EpilogueTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutCTag, kAlignmentC,
|
||||
ElementD, LayoutDTag, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
, cutlass::epilogue::fusion::LinCombPerRowBiasBlockScaleFactor<
|
||||
32, ElementD, float, cutlass::float_ue4m3_t, LayoutDTag, cutlass::half_t, ElementC, float>
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutATag, kAlignmentA,
|
||||
ElementB, LayoutBTag, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MainloopTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
template <class T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
} // kernel_6
|
||||
|
||||
|
||||
// D = gelu(alpha * accum + beta * C + per-row bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: fp32
|
||||
TEST(SM120_Device_Sparse_BlockScaled_Gemm_fe2m1t_fe2m1n_f32n_tensor_op_f32, 128x64x256_per_row_bias_gelu) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_1::Gemm, false /*force_legacy_epilogue*/, false /*apply_alignment_offset*/>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
// D = gelu(alpha * accum + beta * C + per-col bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: fp32
|
||||
TEST(SM120_Device_Sparse_BlockScaled_VS32_Gemm_e2m1t_e2m1n_f32n_tensorop_op_f32, 128x128x256_alpha_beta_per_col_bias_gelu) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_2::Gemm, false /*force_legacy_epilogue*/, false /*apply_alignment_offset*/>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
// D = clamp(alpha * accum + beta * C + per-col bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: fp32
|
||||
TEST(SM120_Device_Sparse_BlockScaled_VS64_Gemm_e2m1t_e2m1n_f32n_tensorop_op_f32, 128x128x256_alpha_beta_per_col_bias_clamp) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_3::Gemm, false /*force_legacy_epilogue*/, false /*apply_alignment_offset*/>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
// D = clamp(alpha * accum + beta * C + per-col bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// Scale factor: fp8
|
||||
// D: fp4
|
||||
TEST(SM120_Device_Sparse_Gemm_fe2m1t_fe2m1n_f32n_tensor_op_f4, 128x128x256_column_major_bias_clamp) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_4::Gemm, false /*force_legacy_epilogue*/, false /*apply_alignment_offset*/>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
|
||||
// D = clamp(alpha * accum + beta * C + per-row bias)
|
||||
// C: fp16
|
||||
// Bias: fp16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// Scale factor: fp8
|
||||
// D: fp4 SF VEC32
|
||||
TEST(SM120_Device_Sparse_Gemm_fe2m1t_fe2m1n_f32n_tensor_op_bsf4_bs32_clamp, 128x128x256) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_5::Gemm,
|
||||
false /*force_legacy_epilogue*/,
|
||||
false /*apply_alignment_offset*/>(1.0, 0);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
// D = alpha * accum + beta * C + per-row bias
|
||||
// C: fp16
|
||||
// Bias: fp16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// Scale factor: fp8
|
||||
// D: fp4 SF VEC32
|
||||
TEST(SM120_Device_Sparse_Gemm_fe2m1t_fe2m1n_f32n_tensor_op_bsf4_bs32, 128x128x256) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_6::Gemm,
|
||||
false /*force_legacy_epilogue*/,
|
||||
false /*apply_alignment_offset*/>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
+124
@@ -0,0 +1,124 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace kernel_1 {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutCTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutDTag = cutlass::layout::ColumnMajor;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using ElementA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
|
||||
static constexpr int AlignmentA = 256;
|
||||
static constexpr int AlignmentB = 128;
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutCTag, AlignmentC,
|
||||
ElementD, LayoutDTag, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledSparseTensorOp,
|
||||
ElementPairA, LayoutATag, AlignmentA,
|
||||
ElementPairB, LayoutBTag, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelSparseTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::StreamKScheduler
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
} // kernel_1
|
||||
|
||||
|
||||
TEST(SM120_Device_Sparse_Gemm_e2m1t_e2m1n_f32n_vs32_tensorop_op_f32_stream_k, 128x128x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_1::Gemm, true>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
+138
@@ -0,0 +1,138 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace kernel_1 {
|
||||
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutCTag = cutlass::layout::RowMajor;
|
||||
using LayoutDTag = cutlass::layout::RowMajor;
|
||||
using LayoutSFDTag = cutlass::layout::RowMajor;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::bfloat16_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
using ElementPairA = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 64 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 64 bytes with sparse ratio 4:2.
|
||||
static constexpr int AlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
static constexpr int SFVectorSize = 64;
|
||||
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerColBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFVectorSize,
|
||||
ElementD,
|
||||
ElementCompute,
|
||||
ElementSF,
|
||||
LayoutSFDTag,
|
||||
ElementBias,
|
||||
ElementC>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
float, float,
|
||||
ElementC, LayoutCTag, AlignmentC,
|
||||
ElementD, LayoutDTag, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledSparseTensorOp,
|
||||
ElementPairA, LayoutATag, AlignmentA,
|
||||
ElementPairB, LayoutBTag, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelSparseTmaWarpSpecializedMxf8f6f4Acc2x4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_1
|
||||
|
||||
TEST(SM120_Device_Sparse_Gemm_fe2m1t_fe2m1n_fe2m1t_tensor_op_f32, 128x128x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_1::Gemm, false /*force_legacy_epilogue*/, false /*apply_alignment_offset*/>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
+121
@@ -0,0 +1,121 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace kernel_1 {
|
||||
using ElementA = cutlass::float_e2m3_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
using ElementPairA = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 96 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 96 bytes with sparse ratio 4:2.
|
||||
static constexpr int AlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_64,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledSparseTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelSparseTmaWarpSpecializedMxf8f6f4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_1
|
||||
|
||||
TEST(SM120_Device_Sparse_BlockedScalar_Gemm_fe2m3t_fe2m1n_f32n_tensor_op_f32, 128x64x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_1::Gemm, false /*force_legacy_epilogue*/, false /*apply_alignment_offset*/>(2.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
+122
@@ -0,0 +1,122 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace kernel_1 {
|
||||
using ElementA = cutlass::float_e5m2_t;
|
||||
using ElementB = cutlass::float_e2m3_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 256;
|
||||
static constexpr int AlignmentB = 128;
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_64,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledSparseTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelSparseTmaWarpSpecializedMxf8f6f4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
} // kernel_1
|
||||
|
||||
|
||||
TEST(SM120_Device_Sparse_BlockScaled_Gemm_fe5m2t_fe2m3n_f32n_tensor_op_f32, 128x64x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_1::Gemm, false /*force_legacy_epilogue*/, false /*apply_alignment_offset*/>(1.0, 0.0);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
+157
@@ -0,0 +1,157 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace kernel_1 {
|
||||
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutCTag = cutlass::layout::RowMajor;
|
||||
using LayoutDTag = cutlass::layout::RowMajor;
|
||||
using LayoutSFDTag = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementA = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementB = cutlass::mx_float8_t<cutlass::float_e4m3_t>;
|
||||
using ElementC = cutlass::bfloat16_t;
|
||||
using ElementD = cutlass::float_e4m3_t;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
|
||||
constexpr int kAlignmentA = 32;
|
||||
constexpr int kAlignmentB = 16;
|
||||
constexpr int kAlignmentC = 1;
|
||||
constexpr int kAlignmentD = 4;
|
||||
|
||||
using ProblemShape = Shape<int,int,int,int>;
|
||||
using ClusterShape = Shape<cute::Int<1>, cute::Int<1>, cute::Int<1>>;
|
||||
using MainloopTileShape = Shape<cute::Int<128>, cute::Int<128>, cute::Int<256>>;
|
||||
using EpilogueTileShape = Shape<cute::Int<128>, cute::Int<128>, cute::Int<256>>;
|
||||
using ArchTag = cutlass::arch::Sm120;
|
||||
using OpClassEpilogue = cutlass::arch::OpClassTensorOp;
|
||||
using OpClassMainLoop = cutlass::arch::OpClassBlockScaledSparseTensorOp;
|
||||
using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
|
||||
using EpilogueScheduleType = cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120;
|
||||
using KernelScheduleType = cutlass::gemm::KernelSparseTmaWarpSpecializedMxf8f6f4Acc2x4Sm120;
|
||||
using ElementAccumulator = float;
|
||||
using ElementEpilogueCompute = float;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
using TileScheduler = void;
|
||||
static constexpr int SFVectorSize = 64;
|
||||
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerColBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::Clamp,
|
||||
SFVectorSize,
|
||||
ElementD,
|
||||
ElementCompute,
|
||||
ElementSF,
|
||||
LayoutSFDTag,
|
||||
ElementBias,
|
||||
ElementC>;
|
||||
|
||||
using CollectiveEpilogue =
|
||||
typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassEpilogue,
|
||||
EpilogueTileShape,
|
||||
ClusterShape,
|
||||
EpilogueTile,
|
||||
ElementAccumulator,
|
||||
ElementEpilogueCompute,
|
||||
ElementC, LayoutCTag, kAlignmentC,
|
||||
ElementD, LayoutDTag, kAlignmentD,
|
||||
EpilogueScheduleType
|
||||
, FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using StageCount = cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>;
|
||||
|
||||
using CollectiveMainloop =
|
||||
typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
ArchTag,
|
||||
OpClassMainLoop,
|
||||
ElementA, LayoutATag, kAlignmentA,
|
||||
ElementB, LayoutBTag, kAlignmentB,
|
||||
ElementAccumulator,
|
||||
MainloopTileShape,
|
||||
ClusterShape,
|
||||
StageCount,
|
||||
KernelScheduleType
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
TileScheduler
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_1
|
||||
|
||||
TEST(SM120_Device_Sparse_Gemm_fe4m3t_fe4m3n_fe4m3t_tensor_op_f32, 128x128x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_1::Gemm, false /*force_legacy_epilogue*/, false /*apply_alignment_offset*/>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
@@ -0,0 +1,68 @@
|
||||
# Copyright (c) 2024 - 2025 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.
|
||||
|
||||
#
|
||||
#
|
||||
|
||||
if (CUTLASS_NVCC_ARCHS MATCHES 120a)
|
||||
|
||||
add_custom_target(
|
||||
cutlass_test_unit_gemm_device_sm120_bs
|
||||
DEPENDS
|
||||
cutlass_test_unit_bs_gemm_device_tensorop_epilogue_fusion_sm120
|
||||
cutlass_test_unit_bs_gemm_device_tensorop_sm120
|
||||
cutlass_test_unit_bs_gemm_device_tensorop_sm120_stream_k
|
||||
)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable(
|
||||
cutlass_test_unit_bs_gemm_device_tensorop_epilogue_fusion_sm120
|
||||
|
||||
BATCH_SOURCES ON
|
||||
BATCH_SIZE 1
|
||||
|
||||
sm120_bs_gemm_f4_f4_f32_f32_epilogue_fusion.cu
|
||||
sm120_bs_gemm_f4_f4_f32_f4_epilogue_fusion.cu
|
||||
sm120_bs_gemm_f4_f4_f32_bf16_epilogue_fusion.cu
|
||||
)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable(
|
||||
cutlass_test_unit_bs_gemm_device_tensorop_sm120
|
||||
|
||||
sm120_bs_gemm_f4_f4_f32_bf16.cu
|
||||
sm120_bs_gemm_f4_f4_f32_f16.cu
|
||||
sm120_bs_gemm_f4_f4_f32_f32.cu
|
||||
sm120_bs_gemm_f4_f4_f32_f32_narrow_output.cu
|
||||
sm120_bs_gemm_f4_f4_f32_epilogue.cu
|
||||
)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable(
|
||||
cutlass_test_unit_bs_gemm_device_tensorop_sm120_stream_k
|
||||
sm120_bs_gemm_f4_f4_f32_f32_stream_k.cu
|
||||
)
|
||||
|
||||
endif()
|
||||
+188
@@ -0,0 +1,188 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace kernel_1 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::bfloat16_t;
|
||||
using ElementD = cutlass::bfloat16_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using TileSchedulerTag = cutlass::gemm::PersistentScheduler; // both void (default) and PersistentScheduler map to dynamic scheduler with CLC query
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
TileSchedulerTag>;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_1
|
||||
|
||||
namespace kernel_3 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::bfloat16_t;
|
||||
using ElementD = cutlass::bfloat16_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 64 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_128>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedMxf8f6f4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using TileSchedulerTag = cutlass::gemm::PersistentScheduler; // both void (default) and PersistentScheduler map to dynamic scheduler with CLC query
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
TileSchedulerTag>;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_3
|
||||
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_f32n_vs16_tensor_op_bf16, 128x128x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_1::Gemm, true>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_f32n_vs32_tensor_op_bf16, 128x128x128) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_3::Gemm, true>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
+385
@@ -0,0 +1,385 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// D = relu(alpha * accum + beta * C + per-col bias)
|
||||
// C: bf16
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: bf16
|
||||
namespace kernel_1 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::bfloat16_t;
|
||||
using ElementD = cutlass::bfloat16_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerColBiasEltAct<
|
||||
cutlass::epilogue::thread::ReLU, ElementD, ElementCompute, ElementBias, ElementAccumulator, ElementCompute>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue>;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
|
||||
} // kernel_1
|
||||
|
||||
// D = gelu(alpha * accum + beta * C + per-col bias)
|
||||
// C: bf16
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: bf16
|
||||
namespace kernel_2 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::bfloat16_t;
|
||||
using ElementD = cutlass::bfloat16_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerColBiasEltAct<
|
||||
cutlass::epilogue::thread::GELU, ElementD, ElementCompute, ElementBias, ElementAccumulator, ElementCompute>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue>;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
|
||||
} // kernel_2
|
||||
|
||||
// D = relu(alpha * accum + beta * C + per-row bias)
|
||||
// C: bf16
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: bf16
|
||||
namespace kernel_3 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::bfloat16_t;
|
||||
using ElementD = cutlass::bfloat16_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::ReLu, ElementD, ElementCompute, ElementBias, ElementAccumulator, ElementCompute>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue>;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
|
||||
} // kernel_3
|
||||
|
||||
// Aux = alpha * accum + beta * C + per-row bias
|
||||
// D = gelu(Aux)
|
||||
// C: bf16
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: bf16
|
||||
namespace kernel_4 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::bfloat16_t;
|
||||
using ElementD = cutlass::bfloat16_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
using ElementAux = cutlass::bfloat16_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerRowBiasEltActAux<
|
||||
LayoutC, cutlass::epilogue::thread::GELU, ElementD, ElementCompute, ElementAux, ElementBias, ElementAccumulator, ElementCompute>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue>;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
|
||||
} // kernel_4
|
||||
|
||||
// D = relu(alpha * accum + beta * C + per-col bias)
|
||||
// C: bf16
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: bf16
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_bf16n_vs16_tensor_op_f32_f32_epilogue, 128x128x256_alpha_beta_per_col_bias_relu) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_1::Gemm, false, false>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
// D = gelu(alpha * accum + beta * C + per-col bias)
|
||||
// C: bf16
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: bf16
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_bf16n_vs16_tensor_op_f32_f32_epilogue, 128x128x256_alpha_beta_per_col_bias_gelu) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_2::Gemm, false, false>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
// D = relu(alpha * accum + beta * C + per-row bias)
|
||||
// C: bf16
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: bf16
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_bf16n_vs16_tensor_op_f32_f32_epilogue, 128x128x256_alpha_beta_per_row_bias_relu) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_3::Gemm, false, false>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
// Aux = alpha * accum + beta * C + per-row bias
|
||||
// D = gelu(Aux)
|
||||
// C: bf16
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: bf16
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_bf16n_vs16_tensor_op_f32_f32_epilogue, 128x128x256_alpha_beta_per_row_bias_gelu) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_4::Gemm, false, false>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
+590
@@ -0,0 +1,590 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace kernel_1 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
using LayoutSFD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
constexpr int SFVectorSize = 16;
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombBlockScaleFactor<
|
||||
SFVectorSize,
|
||||
ElementD,
|
||||
ElementCompute,
|
||||
ElementSF,
|
||||
LayoutSFD
|
||||
>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::PersistentScheduler>; // both void (default) and PersistentScheduler map to dynamic scheduler with CLC query
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
|
||||
} // kernel_1
|
||||
|
||||
namespace kernel_2 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
using LayoutSFD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
constexpr int SFVectorSize = 32;
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombBlockScaleFactor<
|
||||
SFVectorSize,
|
||||
ElementD,
|
||||
ElementCompute,
|
||||
ElementSF,
|
||||
LayoutSFD
|
||||
>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedMxf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::PersistentScheduler>; // both void (default) and PersistentScheduler map to dynamic scheduler with CLC query
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_2
|
||||
|
||||
namespace kernel_3 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
using LayoutSFD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
constexpr int SFVectorSize = 16;
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombBlockScaleFactor<
|
||||
SFVectorSize,
|
||||
ElementD,
|
||||
ElementCompute,
|
||||
ElementSF,
|
||||
LayoutSFD
|
||||
>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::PersistentScheduler>; // both void (default) and PersistentScheduler map to dynamic scheduler with CLC query
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
|
||||
} // kernel_3
|
||||
|
||||
namespace kernel_4 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
using LayoutSFD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
constexpr int SFVectorSize = 32;
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombBlockScaleFactor<
|
||||
SFVectorSize,
|
||||
ElementD,
|
||||
ElementCompute,
|
||||
ElementSF,
|
||||
LayoutSFD
|
||||
>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedMxf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::PersistentScheduler>; // both void (default) and PersistentScheduler map to dynamic scheduler with CLC query
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_4
|
||||
|
||||
namespace kernel_5 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::bfloat16_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
using LayoutSFD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
constexpr int SFVectorSize = 16;
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombBlockScaleFactor<
|
||||
SFVectorSize,
|
||||
ElementD,
|
||||
ElementCompute,
|
||||
ElementSF,
|
||||
LayoutSFD
|
||||
>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::PersistentScheduler>; // both void (default) and PersistentScheduler map to dynamic scheduler with CLC query
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
|
||||
} // kernel_5
|
||||
|
||||
namespace kernel_6 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::bfloat16_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
using LayoutSFD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
constexpr int SFVectorSize = 32;
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombBlockScaleFactor<
|
||||
SFVectorSize,
|
||||
ElementD,
|
||||
ElementCompute,
|
||||
ElementSF,
|
||||
LayoutSFD
|
||||
>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedMxf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::PersistentScheduler>; // both void (default) and PersistentScheduler map to dynamic scheduler with CLC query
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_6
|
||||
|
||||
namespace kernel_7 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
using LayoutSFD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
constexpr int SFVectorSize = 16;
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombBlockScaleFactor<
|
||||
SFVectorSize,
|
||||
ElementD,
|
||||
ElementCompute,
|
||||
ElementSF,
|
||||
LayoutSFD,
|
||||
ElementC
|
||||
>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::PersistentScheduler>; // both void (default) and PersistentScheduler map to dynamic scheduler with CLC query
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
|
||||
} // kernel_7
|
||||
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_fe2m1t_vs16_tensor_op_f32_f32_epilogue_vs16, 128x128x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_1::Gemm>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_fe2m1t_vs32_tensor_op_f32_f32_epilogue_vs32, 128x128x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_2::Gemm>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
// ==== mixed datatypes for C (fp16/bf16) / D (fp32) matrices ==== //
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_fe2m1t_vs16_tensor_op_f16_f32_epilogue_vs16, 128x128x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_3::Gemm>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_fe2m1t_vs32_tensor_op_f16_f32_epilogue_vs32, 128x128x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_4::Gemm>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_fe2m1t_vs16_tensor_op_bf16_f32_epilogue_vs16, 128x128x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_5::Gemm>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_fe2m1t_vs32_tensor_op_bf16_f32_epilogue_vs32, 128x128x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_6::Gemm>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_fe2m1t_vs32_tensor_op_void_f32_epilogue_vs32, 128x128x256) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_7::Gemm>(1.0, 0.0);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
@@ -0,0 +1,253 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace kernel_1 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using TileSchedulerTag = cutlass::gemm::PersistentScheduler; // both void (default) and PersistentScheduler map to dynamic scheduler with CLC query
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
TileSchedulerTag>;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_1
|
||||
|
||||
namespace kernel_2 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedMxf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_2
|
||||
|
||||
|
||||
namespace kernel_3 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 64 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_128>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedMxf8f6f4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using TileSchedulerTag = cutlass::gemm::PersistentScheduler; // both void (default) and PersistentScheduler map to dynamic scheduler with CLC query
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
TileSchedulerTag>;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_3
|
||||
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_f32n_vs16_tensor_op_f16, 128x128x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_1::Gemm>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_f32n_vs32_tensor_op_f16_static_sched, 128x128x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_2::Gemm>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_f32n_vs32_tensor_op_f16, 128x128x128) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_3::Gemm>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
@@ -0,0 +1,122 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace kernel_1 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedMxf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_1
|
||||
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_f32n_vs32_tensor_op_f32_static_sched, 128x128x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_1::Gemm, true>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
+544
@@ -0,0 +1,544 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// D = relu(alpha * accum + beta * C + per-row bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: fp32
|
||||
namespace kernel_1 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::ReLu, ElementD, ElementCompute, ElementBias, ElementAccumulator, ElementCompute>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue>;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
|
||||
} // kernel_1
|
||||
|
||||
// D = gelu(alpha * accum + beta * C + per-row bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: fp32
|
||||
namespace kernel_2 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::GELU, ElementD, ElementCompute, ElementBias, ElementAccumulator, ElementCompute>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue>;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
|
||||
} // kernel_2
|
||||
|
||||
// D = gelu(alpha * accum + beta * C + per-col bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: fp32
|
||||
namespace kernel_3 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerColBiasEltAct<
|
||||
cutlass::epilogue::thread::GELU, ElementD, ElementCompute, ElementBias, ElementAccumulator, ElementCompute>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue>;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
|
||||
} // kernel_3
|
||||
|
||||
// D = relu(alpha * accum + beta * C + per-col bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: fp32
|
||||
namespace kernel_4 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerColBiasEltAct<
|
||||
cutlass::epilogue::thread::ReLU, ElementD, ElementCompute, ElementBias, ElementAccumulator, ElementCompute>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue>;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
|
||||
} // kernel_4
|
||||
|
||||
// D = clamp(alpha * accum + beta * C + per-row bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: fp32
|
||||
namespace kernel_5 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp, ElementD, ElementCompute, ElementBias, ElementAccumulator, ElementCompute>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue>;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
|
||||
} // kernel_5
|
||||
|
||||
// D = clamp(alpha * accum + beta * C + per-col bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: fp32
|
||||
namespace kernel_6 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerColBiasEltAct<
|
||||
cutlass::epilogue::thread::Clamp, ElementD, ElementCompute, ElementBias, ElementAccumulator, ElementCompute>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue>;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
|
||||
} // kernel_6
|
||||
|
||||
// D = relu(alpha * accum + beta * C + per-row bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: fp32
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_f32n_vs16_tensor_op_f32_f32_epilogue, 128x128x256_per_row_bias_relu) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_1::Gemm, false, false>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
// D = gelu(alpha * accum + beta * C + per-row bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: fp32
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_f32n_vs16_tensor_op_f32_f32_epilogue, 128x128x256_per_row_bias_gelu) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_2::Gemm, false, false>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
// D = gelu(alpha * accum + beta * C + per-col bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: fp32
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_f32n_vs16_tensor_op_f32_f32_epilogue, 128x128x256_alpha_beta_per_col_bias_gelu) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_3::Gemm, false, false>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
// D = relu(alpha * accum + beta * C + per-col bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: fp32
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_f32n_vs16_tensor_op_f32_f32_epilogue, 128x128x256_alpha_beta_per_col_bias_relu) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_4::Gemm, false, false>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
// D = clamp(alpha * accum + beta * C + per-row bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: fp32
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_f32n_vs16_tensor_op_f32_f32_epilogue, 128x128x256_per_row_bias_clamp) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_5::Gemm, false, false>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
// D = clamp(alpha * accum + beta * C + per-col bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: fp32
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_f32n_vs16_tensor_op_f32_f32_epilogue, 128x128x256_alpha_beta_per_col_bias_clamp) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_6::Gemm, false, false>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
+218
@@ -0,0 +1,218 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace kernel_1 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::bfloat16_t;
|
||||
using ElementD = cutlass::float_e2m3_t;
|
||||
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
constexpr int SFVectorSize = 16;
|
||||
using LayoutSFD = cutlass::layout::RowMajor;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
using GmemLayoutSFC = cutlass::layout::RowMajor;
|
||||
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerColBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::ReLU,
|
||||
SFVectorSize,
|
||||
ElementD,
|
||||
ElementCompute,
|
||||
ElementSF, LayoutSFD,
|
||||
ElementBias,
|
||||
ElementC>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::PersistentScheduler>; // both void (default) and PersistentScheduler map to dynamic scheduler with CLC query
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_1
|
||||
|
||||
namespace kernel_2 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::bfloat16_t;
|
||||
using ElementD = cutlass::float_e2m3_t;
|
||||
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::mx_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
constexpr int SFVectorSize = 32;
|
||||
using LayoutSFD = cutlass::layout::RowMajor;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
using GmemLayoutSFC = cutlass::layout::RowMajor;
|
||||
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerColBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::ReLU,
|
||||
SFVectorSize,
|
||||
ElementD,
|
||||
ElementCompute,
|
||||
ElementSF, LayoutSFD,
|
||||
ElementBias,
|
||||
ElementC>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative
|
||||
,FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedMxf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::PersistentScheduler>; // both void (default) and PersistentScheduler map to dynamic scheduler with CLC query
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_2
|
||||
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_f32n_vs16_tensor_op_f32_fe2m3n, 128x128x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_1::Gemm, false, false>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_f32n_vs32_tensor_op_f32_fe2m3n, 128x128x256) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_2::Gemm, false, false>(1.0, 0);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
+123
@@ -0,0 +1,123 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace kernel_1 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::StreamKScheduler
|
||||
>;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_1
|
||||
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_f32n_vs16_tensor_op_f32_stream_k, 128x128x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_1::Gemm, true>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
+590
@@ -0,0 +1,590 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// D = alpha * accum + beta * C + per-row bias
|
||||
// C: bf16
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: bf16
|
||||
namespace kernel_1 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::bfloat16_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
using LayoutSFD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
constexpr int SFVectorSize = 16;
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerRowBiasBlockScaleFactor<
|
||||
SFVectorSize,
|
||||
ElementD,
|
||||
ElementCompute,
|
||||
ElementSF, LayoutSFD,
|
||||
ElementBias,
|
||||
ElementC>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue>;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
|
||||
} // kernel_1
|
||||
|
||||
// D = relu(alpha * accum + beta * C + per-row bias)
|
||||
// C: bf16
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: bf16
|
||||
namespace kernel_2 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::bfloat16_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
using LayoutSFD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
constexpr int SFVectorSize = 16;
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::ReLU,
|
||||
SFVectorSize,
|
||||
ElementD,
|
||||
ElementCompute,
|
||||
ElementSF, LayoutSFD,
|
||||
ElementBias,
|
||||
ElementC>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue>;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
|
||||
} // kernel_2
|
||||
|
||||
// D = gelu(alpha * accum + beta * C + per-row bias)
|
||||
// C: bf16
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: bf16
|
||||
namespace kernel_3 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::bfloat16_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
using LayoutSFD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
constexpr int SFVectorSize = 16;
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerRowBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::GELU,
|
||||
SFVectorSize,
|
||||
ElementD,
|
||||
ElementCompute,
|
||||
ElementSF, LayoutSFD,
|
||||
ElementBias,
|
||||
ElementC>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue>;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_3
|
||||
|
||||
// D = alpha * accum + beta * C + per-col bias
|
||||
// C: bf16
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: bf16
|
||||
namespace kernel_4 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::bfloat16_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
using LayoutSFD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
constexpr int SFVectorSize = 16;
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerColBiasBlockScaleFactor<
|
||||
SFVectorSize,
|
||||
ElementD,
|
||||
ElementCompute,
|
||||
ElementSF, LayoutSFD,
|
||||
ElementBias,
|
||||
ElementC>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue>;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_4
|
||||
|
||||
// D = relu(alpha * accum + beta * C + per-col bias)
|
||||
// C: bf16
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: bf16
|
||||
namespace kernel_5 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::bfloat16_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
using LayoutSFD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
constexpr int SFVectorSize = 16;
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerColBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::ReLU,
|
||||
SFVectorSize,
|
||||
ElementD,
|
||||
ElementCompute,
|
||||
ElementSF, LayoutSFD,
|
||||
ElementBias,
|
||||
ElementC>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue>;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_5
|
||||
|
||||
// D = gelu(alpha * accum + beta * C + per-col bias)
|
||||
// C: bf16
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: bf16
|
||||
namespace kernel_6 {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::bfloat16_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::RowMajor;
|
||||
using LayoutD = cutlass::layout::RowMajor;
|
||||
using LayoutSFD = cutlass::layout::RowMajor;
|
||||
|
||||
using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 16 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
constexpr int SFVectorSize = 16;
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerColBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::GELU,
|
||||
SFVectorSize,
|
||||
ElementD,
|
||||
ElementCompute,
|
||||
ElementSF, LayoutSFD,
|
||||
ElementBias,
|
||||
ElementC>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA, LayoutA, AlignmentA,
|
||||
ElementPairB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue>;
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_6
|
||||
|
||||
////////////////////////////////////////////////////////////////////
|
||||
// EVT for epilogue with scale factor
|
||||
////////////////////////////////////////////////////////////////////
|
||||
|
||||
// D = alpha * accum + beta * C + per-row bias
|
||||
// C: bf16
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: bf16
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_fe2m1t_vs16_tensor_op_f32_f32_epilogue, 128x128x256_alpha_beta_per_row_bias) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_1::Gemm, false, false>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
// D = relu(alpha * accum + beta * C + per-row bias)
|
||||
// C: bf16
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: bf16
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_fe2m1t_vs16_tensor_op_f32_f32_epilogue, 128x128x256_alpha_beta_per_row_bias_relu) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_2::Gemm, false, false>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
// D = gelu(alpha * accum + beta * C + per-row bias)
|
||||
// C: bf16
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: bf16
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_fe2m1t_vs16_tensor_op_f32_f32_epilogue, 128x128x256_alpha_beta_per_row_bias_gelu) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_3::Gemm, false, false>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
// D = alpha * accum + beta * C + per-col bias
|
||||
// C: bf16
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: bf16
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_fe2m1t_vs16_tensor_op_f32_f32_epilogue, 128x128x256_alpha_beta_per_col_bias) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_4::Gemm, false, false>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
// D = relu(alpha * accum + beta * C + per-col bias)
|
||||
// C: bf16
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: bf16
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_fe2m1t_vs16_tensor_op_f32_f32_epilogue, 128x128x256_alpha_beta_per_col_bias_relu) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_5::Gemm, false, false>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
// D = gelu(alpha * accum + beta * C + per-col bias)
|
||||
// C: bf16
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: bf16
|
||||
TEST(SM120_Device_Blockscaled_Gemm_fe2m1t_fe2m1n_fe2m1t_vs16_tensor_op_f32_f32_epilogue, 128x128x256_alpha_beta_per_col_bias_gelu) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_6::Gemm, false, false>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
@@ -0,0 +1,68 @@
|
||||
# Copyright (c) 2025 - 2025 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.
|
||||
|
||||
if (CUTLASS_NVCC_ARCHS MATCHES 120a)
|
||||
|
||||
add_custom_target(
|
||||
cutlass_test_unit_gemm_device_sm120_sptensorop
|
||||
DEPENDS
|
||||
cutlass_test_unit_sparse_gemm_device_tensorop_sm120
|
||||
cutlass_test_unit_sparse_gemm_device_tensorop_sm120_stream_k
|
||||
cutlass_test_unit_sparse_gemm_device_tensorop_sm120_epilogue_fusion
|
||||
)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable(
|
||||
cutlass_test_unit_sparse_gemm_device_tensorop_sm120
|
||||
|
||||
sm120_sparse_gemm_f4_f4_f32_tensor_op.cu
|
||||
sm120_sparse_gemm_f6_f4_f32_tensor_op.cu
|
||||
sm120_sparse_gemm_f8_f6_f32_tensor_op.cu
|
||||
sm120_sparse_gemm_f4_f4_f16_tensor_op.cu
|
||||
sm120_sparse_gemm_f6_f4_f16_tensor_op.cu
|
||||
sm120_sparse_gemm_f8_f6_f16_tensor_op.cu
|
||||
)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable(
|
||||
cutlass_test_unit_sparse_gemm_device_tensorop_sm120_epilogue_fusion
|
||||
|
||||
# No batching of source to control compiler memory usage
|
||||
BATCH_SOURCES ON
|
||||
BATCH_SIZE 1
|
||||
sm120_sparse_gemm_f4_f4_f32_tensor_op_epilogue_fusion.cu
|
||||
)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable(
|
||||
cutlass_test_unit_sparse_gemm_device_tensorop_sm120_stream_k
|
||||
|
||||
# No batching of source to control compiler memory usage
|
||||
BATCH_SOURCES ON
|
||||
BATCH_SIZE 1
|
||||
sm120_sparse_gemm_f4_f4_f32_tensor_op_f32_stream_k.cu
|
||||
)
|
||||
|
||||
endif()
|
||||
+119
@@ -0,0 +1,119 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace kernel_1 {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using TileShape = Shape<_128,_64,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = half_t;
|
||||
using ElementCompute = float;
|
||||
|
||||
static constexpr int AlignmentA = 64 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 64 bytes with sparse ratio 4:2.
|
||||
static constexpr int AlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassSparseTensorOp,
|
||||
ElementA, LayoutATag, AlignmentA,
|
||||
ElementB, LayoutBTag, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::collective::KernelScheduleAuto
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_1
|
||||
|
||||
TEST(SM120_Device_Sparse_Gemm_fe2m1t_fe2m1n_f16n_tensor_op_f32, 128x64x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_1::Gemm, false /*force_legacy_epilogue*/, false /*apply_alignment_offset*/>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
+118
@@ -0,0 +1,118 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace kernel_1 {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using TileShape = Shape<_128,_64,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = float;
|
||||
|
||||
static constexpr int AlignmentA = 64 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 64 bytes with sparse ratio 4:2.
|
||||
static constexpr int AlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
float, float,
|
||||
float, LayoutC, AlignmentC,
|
||||
float, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassSparseTensorOp,
|
||||
ElementA, LayoutATag, AlignmentA,
|
||||
ElementB, LayoutBTag, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::collective::KernelScheduleAuto
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_1
|
||||
|
||||
TEST(SM120_Device_Sparse_Gemm_fe2m1t_fe2m1n_f32n_tensor_op_f32, 128x64x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_1::Gemm, false /*force_legacy_epilogue*/, false /*apply_alignment_offset*/>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
+593
@@ -0,0 +1,593 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// D = relu(alpha * accum + beta * C + per-row bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: fp32
|
||||
namespace kernel_1 {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutCTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutDTag = cutlass::layout::ColumnMajor;
|
||||
|
||||
using TileShape = Shape<_128,_64,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
|
||||
static constexpr int AlignmentA = 64 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 64 bytes with sparse ratio 4:2.
|
||||
static constexpr int AlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerRowBiasEltAct<
|
||||
cutlass::epilogue::thread::ReLu, ElementD, ElementCompute, ElementBias, ElementAccumulator, ElementCompute>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
float, float,
|
||||
float, LayoutCTag, AlignmentC,
|
||||
float, LayoutDTag, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassSparseTensorOp,
|
||||
ElementA, LayoutATag, AlignmentA,
|
||||
ElementB, LayoutBTag, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::collective::KernelScheduleAuto
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_1
|
||||
|
||||
// D = alpha * accum + beta * C
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// Scale factor: fp8
|
||||
// D: fp4
|
||||
namespace kernel_2 {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutCTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutDTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutSFDTag = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
using ElementC = cutlass::bfloat16_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
static constexpr int AlignmentA = 64 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 64 bytes with sparse ratio 4:2.
|
||||
static constexpr int AlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
constexpr int SFVectorSize = 64;
|
||||
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombBlockScaleFactor<
|
||||
SFVectorSize,
|
||||
ElementD,
|
||||
ElementCompute,
|
||||
ElementSF,
|
||||
LayoutSFDTag,
|
||||
ElementC
|
||||
>;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutCTag, AlignmentC,
|
||||
ElementD, LayoutDTag, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassSparseTensorOp,
|
||||
ElementA, LayoutATag, AlignmentA,
|
||||
ElementB, LayoutBTag, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelScheduleSparseF8f6f4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
} // kernel_2
|
||||
|
||||
// D = alpha * accum + beta * C + per-column bias
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// Scale factor: fp8
|
||||
// D: fp4
|
||||
namespace kernel_3 {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutCTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutDTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutSFDTag = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
using ElementC = cutlass::bfloat16_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
static constexpr int AlignmentA = 64 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 64 bytes with sparse ratio 4:2.
|
||||
static constexpr int AlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
constexpr int SFVectorSize = 32;
|
||||
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerColBiasBlockScaleFactor<
|
||||
SFVectorSize,
|
||||
ElementD,
|
||||
ElementCompute,
|
||||
ElementSF,
|
||||
LayoutSFDTag,
|
||||
ElementBias,
|
||||
ElementC
|
||||
>;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutCTag, AlignmentC,
|
||||
ElementD, LayoutDTag, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassSparseTensorOp,
|
||||
ElementA, LayoutATag, AlignmentA,
|
||||
ElementB, LayoutBTag, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelScheduleSparseF8f6f4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
} // kernel_3
|
||||
|
||||
// D = relu(alpha * accum + beta * C + per-column bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// Scale factor: fp8
|
||||
// D: fp4
|
||||
namespace kernel_4 {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutCTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutDTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutSFDTag = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
using ElementC = cutlass::bfloat16_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
static constexpr int AlignmentA = 64 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 64 bytes with sparse ratio 4:2.
|
||||
static constexpr int AlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
constexpr int SFVectorSize = 64;
|
||||
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerColBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::ReLU,
|
||||
SFVectorSize,
|
||||
ElementD,
|
||||
ElementCompute,
|
||||
ElementSF,
|
||||
LayoutSFDTag,
|
||||
ElementBias,
|
||||
ElementC
|
||||
>;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutCTag, AlignmentC,
|
||||
ElementD, LayoutDTag, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassSparseTensorOp,
|
||||
ElementA, LayoutATag, AlignmentA,
|
||||
ElementB, LayoutBTag, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelScheduleSparseF8f6f4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
} // kernel_4
|
||||
|
||||
// D = gelu(alpha * accum + beta * C + per-column bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// Scale factor: fp8
|
||||
// D: fp4
|
||||
namespace kernel_5 {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutCTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutDTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutSFDTag = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
using ElementC = cutlass::bfloat16_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
static constexpr int AlignmentA = 64 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 64 bytes with sparse ratio 4:2.
|
||||
static constexpr int AlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
constexpr int SFVectorSize = 32;
|
||||
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerColBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::GELU,
|
||||
SFVectorSize,
|
||||
ElementD,
|
||||
ElementCompute,
|
||||
ElementSF,
|
||||
LayoutSFDTag,
|
||||
ElementBias,
|
||||
ElementC
|
||||
>;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutCTag, AlignmentC,
|
||||
ElementD, LayoutDTag, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassSparseTensorOp,
|
||||
ElementA, LayoutATag, AlignmentA,
|
||||
ElementB, LayoutBTag, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelScheduleSparseF8f6f4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <class T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
} // kernel_5
|
||||
|
||||
// D = gelu(alpha * accum + beta * C + per-column bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// Scale factor: fp8
|
||||
// D: fp4 with SF VEC16
|
||||
namespace kernel_6 {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutCTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutDTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutSFDTag = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using ElementBias = cutlass::bfloat16_t;
|
||||
using ElementC = cutlass::bfloat16_t;
|
||||
using ElementD = cutlass::float_e2m1_t;
|
||||
using ElementSF = cutlass::float_ue8m0_t;
|
||||
|
||||
static constexpr int AlignmentA = 64 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 64 bytes with sparse ratio 4:2.
|
||||
static constexpr int AlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
constexpr int SFVectorSize = 16;
|
||||
|
||||
using FusionOperation = cutlass::epilogue::fusion::LinCombPerColBiasEltActBlockScaleFactor<
|
||||
cutlass::epilogue::thread::GELU,
|
||||
SFVectorSize,
|
||||
ElementD,
|
||||
ElementCompute,
|
||||
ElementSF,
|
||||
LayoutSFDTag,
|
||||
ElementBias,
|
||||
ElementC
|
||||
>;
|
||||
|
||||
using TileShape = Shape<_128,_128,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutCTag, AlignmentC,
|
||||
ElementD, LayoutDTag, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120,
|
||||
FusionOperation
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassSparseTensorOp,
|
||||
ElementA, LayoutATag, AlignmentA,
|
||||
ElementB, LayoutBTag, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::KernelScheduleSparseF8f6f4Sm120
|
||||
>::CollectiveOp;
|
||||
|
||||
template <class T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
} // kernel_6
|
||||
|
||||
// D = relu(alpha * accum + beta * C + per-row bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// D: fp32
|
||||
TEST(SM120_Device_Sparse_Gemm_fe2m1t_fe2m1n_f32n_tensor_op_f32, 128x64x256_per_row_bias_relu) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_1::Gemm, false /*force_legacy_epilogue*/, false /*apply_alignment_offset*/>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
// D = alpha * accum + beta * C
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// Scale factor: fp8
|
||||
// D: fp4
|
||||
TEST(SM120_Device_Sparse_Gemm_fe2m1t_fe2m1n_f32n_tensor_op_f4, 128x128x256_column_major) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_2::Gemm, false /*force_legacy_epilogue*/, false /*apply_alignment_offset*/>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
// D = alpha * accum + beta * C + per-column bias
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// Scale factor: fp8
|
||||
// D: fp4
|
||||
TEST(SM120_Device_Sparse_Gemm_fe2m1t_fe2m1n_f32n_tensor_op_f4, 128x128x256_column_major_bias) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_3::Gemm, false /*force_legacy_epilogue*/, false /*apply_alignment_offset*/>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
// D = relu(alpha * accum + beta * C + per-column bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// Scale factor: fp8
|
||||
// D: fp4
|
||||
TEST(SM120_Device_Sparse_Gemm_fe2m1t_fe2m1n_f32n_tensor_op_f4, 128x128x256_column_major_bias_relu) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_4::Gemm, false /*force_legacy_epilogue*/, false /*apply_alignment_offset*/>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
// D = gelu(alpha * accum + beta * C + per-column bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// Scale factor: fp8
|
||||
// D: fp4
|
||||
TEST(SM120_Device_Sparse_Gemm_fe2m1t_fe2m1n_f32n_tensor_op_f4, 128x128x256_column_major_bias_gelu) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_5::Gemm, false /*force_legacy_epilogue*/, false /*apply_alignment_offset*/>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
// D = gelu(alpha * accum + beta * C + per-column bias)
|
||||
// C: fp32
|
||||
// Bias: bf16
|
||||
// Acc: fp32
|
||||
// Scale (alpha, beta): fp32
|
||||
// Scale factor: fp8
|
||||
// D: fp4 SF VEC16
|
||||
TEST(SM120_Device_Sparse_Gemm_fe2m1t_fe2m1n_f32n_tensor_op_f4, 128x128x256_column_major_sf_vec16_bias_gelu) {
|
||||
bool result = test::gemm::device::TestSmallFusion<kernel_6::Gemm, false /*force_legacy_epilogue*/, false /*apply_alignment_offset*/>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
+118
@@ -0,0 +1,118 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace kernel_1 {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutCTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutDTag = cutlass::layout::ColumnMajor;
|
||||
|
||||
using TileShape = Shape<_128,_64,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = float;
|
||||
|
||||
static constexpr int AlignmentA = 64 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 64 bytes with sparse ratio 4:2.
|
||||
static constexpr int AlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
float, float,
|
||||
float, LayoutCTag, AlignmentC,
|
||||
float, LayoutDTag, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassSparseTensorOp,
|
||||
ElementA, LayoutATag, AlignmentA,
|
||||
ElementB, LayoutBTag, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::collective::KernelScheduleAuto
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue,
|
||||
cutlass::gemm::StreamKScheduler
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
} // kernel_1
|
||||
|
||||
TEST(SM120_Device_Sparse_Gemm_e2m1t_e2m1n_f32n_vs64_tensor_op_f32_stream_k, 128x128x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_1::Gemm, true>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
+120
@@ -0,0 +1,120 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace kernel_1 {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using TileShape = Shape<_128,_64,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using ElementA = cutlass::float_e3m2_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = half_t;
|
||||
using ElementCompute = float;
|
||||
|
||||
static constexpr int AlignmentA = 96 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 96 bytes with sparse ratio 4:2.
|
||||
static constexpr int AlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassSparseTensorOp,
|
||||
ElementA, LayoutATag, AlignmentA,
|
||||
ElementB, LayoutBTag, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::collective::KernelScheduleAuto
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_1
|
||||
|
||||
TEST(SM120_Device_Sparse_Gemm_fe2m3t_fe2m1n_f16n_tensor_op_f32, 128x64x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_1::Gemm, false /*force_legacy_epilogue*/, false /*apply_alignment_offset*/>(2.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
+118
@@ -0,0 +1,118 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace kernel_1 {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using TileShape = Shape<_128,_64,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using ElementA = cutlass::float_e2m3_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = float;
|
||||
|
||||
static constexpr int AlignmentA = 96 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 96 bytes with sparse ratio 4:2.
|
||||
static constexpr int AlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
float, float,
|
||||
float, LayoutC, AlignmentC,
|
||||
float, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassSparseTensorOp,
|
||||
ElementA, LayoutATag, AlignmentA,
|
||||
ElementB, LayoutBTag, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::collective::KernelScheduleAuto
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_1
|
||||
|
||||
TEST(SM120_Device_Sparse_Gemm_fe2m3t_fe2m1n_f32n_tensor_op_f32, 128x64x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_1::Gemm, false /*force_legacy_epilogue*/, false /*apply_alignment_offset*/>(2.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
+118
@@ -0,0 +1,118 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace kernel_1 {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using TileShape = Shape<_128,_64,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using ElementA = cutlass::float_e5m2_t;
|
||||
using ElementB = cutlass::float_e3m2_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = half_t;
|
||||
using ElementCompute = float;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes with sparse ratio 4:2.
|
||||
static constexpr int AlignmentB = 96 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 96 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassSparseTensorOp,
|
||||
ElementA, LayoutATag, AlignmentA,
|
||||
ElementB, LayoutBTag, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::collective::KernelScheduleAuto
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
} // kernel_1
|
||||
|
||||
TEST(SM120_Device_Sparse_Gemm_fe5m2t_fe3m2n_f16n_tensor_op_f32, 128x64x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_1::Gemm, false /*force_legacy_epilogue*/, false /*apply_alignment_offset*/>(2.0, 0.0);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
+178
@@ -0,0 +1,178 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace kernel_1 {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using TileShape = Shape<_128,_64,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using ElementA = cutlass::float_e5m2_t;
|
||||
using ElementB = cutlass::float_e2m3_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = float;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes with sparse ratio 4:2.
|
||||
static constexpr int AlignmentB = 96 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 96 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
float, float,
|
||||
float, LayoutC, AlignmentC,
|
||||
float, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassSparseTensorOp,
|
||||
ElementA, LayoutATag, AlignmentA,
|
||||
ElementB, LayoutBTag, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::collective::KernelScheduleAuto
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_1
|
||||
|
||||
namespace kernel_2 {
|
||||
using LayoutATag = cutlass::layout::RowMajor;
|
||||
using LayoutBTag = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
using TileShape = Shape<_128,_64,_256>; // M, N, K
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using ElementA = cutlass::float_e5m2_t;
|
||||
using ElementB = cutlass::float_e2m3_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = float;
|
||||
|
||||
static constexpr int AlignmentA = 16 * 8 * 2 / cutlass::sizeof_bits<ElementA>::value; // Align to 16 bytes with sparse ratio 4:2.
|
||||
static constexpr int AlignmentB = 96 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 96 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
float, float,
|
||||
float, LayoutC, AlignmentC,
|
||||
float, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::SparseTmaWarpSpecializedCooperativeSm120
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassSparseTensorOp,
|
||||
ElementA, LayoutATag, AlignmentA,
|
||||
ElementB, LayoutBTag, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::collective::KernelScheduleAuto
|
||||
>::CollectiveOp;
|
||||
|
||||
template <typename T>
|
||||
struct dummy {
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
};
|
||||
using GemmKernel = typename dummy<void>::GemmKernel;
|
||||
using Gemm = typename dummy<void>::Gemm;
|
||||
|
||||
} // kernel_2
|
||||
|
||||
TEST(SM120_Device_Sparse_Gemm_fe5m2t_fe2m3n_f32n_tensor_op_f32, 128x64x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_1::Gemm, false /*force_legacy_epilogue*/, false /*apply_alignment_offset*/>(1.0, 0.0);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
TEST(SM120_Device_Sparse_BlockScaled_Gemm_fe5m2t_fe2m3n_f32n_tensor_op_f32, 128x64x256) {
|
||||
bool result = test::gemm::device::TestSmall<kernel_2::Gemm, false /*force_legacy_epilogue*/, false /*apply_alignment_offset*/>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
@@ -0,0 +1,62 @@
|
||||
# Copyright (c) 2025 - 2025 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.
|
||||
|
||||
if (CUTLASS_NVCC_ARCHS MATCHES 120a)
|
||||
|
||||
add_custom_target(
|
||||
cutlass_test_unit_gemm_device_sm120_tensorop
|
||||
DEPENDS
|
||||
cutlass_test_unit_gemm_device_tensorop_f32_sm120
|
||||
cutlass_test_unit_gemm_device_tensorop_f16_sm120
|
||||
)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable(
|
||||
cutlass_test_unit_gemm_device_tensorop_f32_sm120
|
||||
|
||||
sm120_gemm_f4_f6_f32_tensor_op_narrow_output.cu
|
||||
sm120_gemm_f4_f6_f32_tensor_op.cu
|
||||
sm120_gemm_f4_f8_f32_tensor_op.cu
|
||||
sm120_gemm_f6_f8_f32_tensor_op.cu
|
||||
sm120_gemm_f4_f4_f32_tensor_op.cu
|
||||
sm120_gemm_f6_f6_f32_tensor_op.cu
|
||||
sm120_gemm_f8_f8_f32_tensor_op.cu
|
||||
)
|
||||
|
||||
cutlass_test_unit_gemm_device_add_executable(
|
||||
cutlass_test_unit_gemm_device_tensorop_f16_sm120
|
||||
|
||||
sm120_gemm_f4_f6_f16_tensor_op_narrow_output.cu
|
||||
sm120_gemm_f4_f6_f16_tensor_op.cu
|
||||
sm120_gemm_f4_f8_f16_tensor_op.cu
|
||||
sm120_gemm_f6_f8_f16_tensor_op.cu
|
||||
sm120_gemm_f4_f4_f16_tensor_op.cu
|
||||
sm120_gemm_f6_f6_f16_tensor_op.cu
|
||||
sm120_gemm_f8_f8_f16_tensor_op.cu
|
||||
)
|
||||
|
||||
endif()
|
||||
@@ -0,0 +1,265 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
TEST(SM120_Device_Gemm_fe2m1t_fe2m1n_f16n_void_f32_tensor_op, 128x64x128_1x1x1) {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = void;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = half_t;
|
||||
using ElementCompute = float;
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
static constexpr int AlignmentA = 64 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_64,_128>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
ElementA, LayoutA, AlignmentA,
|
||||
ElementB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::collective::KernelScheduleAuto
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
bool result = test::gemm::device::TestSmall<Gemm, true>(1.0, 0.0);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
TEST(SM120_Device_Gemm_fe2m1t_fe2m1n_f16n_void_f16_tensor_op, 128x64x128_1x1x1) {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = void;
|
||||
using ElementD = cutlass::half_t;
|
||||
using ElementAccumulator = cutlass::half_t;
|
||||
using ElementCompute = float;
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
static constexpr int AlignmentA = 64 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_64,_128>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
ElementA, LayoutA, AlignmentA,
|
||||
ElementB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::collective::KernelScheduleAuto
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
bool result = test::gemm::device::TestSmall<Gemm, true>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
TEST(SM120_Device_Gemm_fe2m1t_fe2m1n_f16n_tensor_op_f32, 128x64x128_1x1x1) {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = half_t;
|
||||
using ElementCompute = float;
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
static constexpr int AlignmentA = 64 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_64,_128>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
ElementA, LayoutA, AlignmentA,
|
||||
ElementB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::collective::KernelScheduleAuto
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
bool result = test::gemm::device::TestSmall<Gemm, true>(1.0, 0.0);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
TEST(SM120_Device_Gemm_fe2m1t_fe2m1n_f16n_tensor_op_f16, 128x64x128_1x1x1) {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
using ElementAccumulator = cutlass::half_t;
|
||||
using ElementCompute = float;
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
static constexpr int AlignmentA = 64 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_64,_128>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
ElementA, LayoutA, AlignmentA,
|
||||
ElementB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::collective::KernelScheduleAuto
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
bool result = test::gemm::device::TestSmall<Gemm, true>(1.0, 0.5);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
@@ -0,0 +1,112 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
|
||||
#include <iostream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.hpp"
|
||||
#include "cutlass/epilogue/collective/collective_builder.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cutlass/epilogue/collective/default_epilogue.hpp"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "../../../common/cutlass_unit_test.h"
|
||||
|
||||
#include "../gemm_testbed_3x.hpp"
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
|
||||
using namespace cute;
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
TEST(SM120_Device_Gemm_fe2m1t_fe2m1n_f32n_tensor_op_f32, 128x64x128_1x1x1) {
|
||||
using ElementA = cutlass::float_e2m1_t;
|
||||
using ElementB = cutlass::float_e2m1_t;
|
||||
using ElementC = float;
|
||||
using ElementD = float;
|
||||
using ElementAccumulator = float;
|
||||
using ElementCompute = float;
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
using LayoutC = cutlass::layout::ColumnMajor;
|
||||
using LayoutD = cutlass::layout::ColumnMajor;
|
||||
|
||||
static constexpr int AlignmentA = 64 * 8 / cutlass::sizeof_bits<ElementA>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentB = 64 * 8 / cutlass::sizeof_bits<ElementB>::value; // Align to 64 bytes.
|
||||
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
|
||||
static constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
|
||||
|
||||
using TileShape = Shape<_128,_64,_128>;
|
||||
using ClusterShape = Shape<_1,_1,_1>;
|
||||
|
||||
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::epilogue::collective::EpilogueTileAuto,
|
||||
ElementAccumulator, ElementCompute,
|
||||
ElementC, LayoutC, AlignmentC,
|
||||
ElementD, LayoutD, AlignmentD,
|
||||
cutlass::epilogue::TmaWarpSpecializedCooperative
|
||||
>::CollectiveOp;
|
||||
|
||||
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
|
||||
cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
|
||||
ElementA, LayoutA, AlignmentA,
|
||||
ElementB, LayoutB, AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape, ClusterShape,
|
||||
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
|
||||
cutlass::gemm::collective::KernelScheduleAuto
|
||||
>::CollectiveOp;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
Shape<int,int,int,int>,
|
||||
CollectiveMainloop,
|
||||
CollectiveEpilogue
|
||||
>;
|
||||
|
||||
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
|
||||
bool result = test::gemm::device::TestSmall<Gemm, true>(1.0, 0.0);
|
||||
EXPECT_TRUE(result);
|
||||
}
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
#endif // (defined(CUTLASS_ARCH_MMA_SM120_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM121_SUPPORTED))
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user