* 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:
Yujia Zhai
2025-03-21 01:52:23 -04:00
committed by GitHub
co-authored by yuzhai Haicheng Wu Haicheng Wu
parent 8c4d1dc47d
commit 62750a2b75
334 changed files with 91517 additions and 2656 deletions
+5
View File
@@ -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>
+51 -5
View File
@@ -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{});
+146 -14
View File
@@ -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>>{});
+71 -21
View File
@@ -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("-------------------------------");
+2
View File
@@ -29,6 +29,8 @@
*
**************************************************************************************************/
//#define CUTLASS_DEBUG_TRACE_LEVEL 1
#include "cutlass_unit_test.h"
#include <cutlass/trace.h>
+1 -1
View File
@@ -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)
{
+12
View File
@@ -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
)
+514 -13
View File
@@ -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()
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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)
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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)
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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()
@@ -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)
@@ -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)
@@ -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)
@@ -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)
@@ -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)
@@ -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)
@@ -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)
@@ -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()
@@ -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))
@@ -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))
@@ -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))
@@ -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))
@@ -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))
@@ -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))
@@ -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()
@@ -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))
@@ -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))
@@ -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))
@@ -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))
@@ -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))
@@ -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))
@@ -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()
@@ -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))
@@ -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))
@@ -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))
@@ -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))
@@ -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))
@@ -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))
@@ -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))
@@ -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