CUTLASS 3.8 Release (#2059)
* CUTLASS 3.8 Release * update * Update README.md * Revert "Update README.md" This reverts commit b353e36fe83e0815f99b44e46c0c95494c44726b. * update * update --------- Co-authored-by: Haicheng Wu <57973641+hwu36@users.noreply.github.com> Co-authored-by: Haicheng Wu <haichengw@nvidia.com>
This commit is contained in:
co-authored by
Haicheng Wu
Haicheng Wu
parent
9eb01fa0b0
commit
389e493055
@@ -111,6 +111,18 @@ struct ElementScalarType<Gemm, Default, std::void_t<typename Gemm::EpilogueOutpu
|
||||
using Type = typename Gemm::EpilogueOutputOp::ElementScalar;
|
||||
};
|
||||
|
||||
|
||||
template <typename Gemm, typename = void>
|
||||
struct IsF8F6F4Kernel {
|
||||
static constexpr bool value = false;
|
||||
};
|
||||
|
||||
template <typename Gemm>
|
||||
struct IsF8F6F4Kernel<Gemm, std::void_t<decltype(Gemm::GemmKernel::CollectiveMainloop::IsF8F6F4)>> {
|
||||
static constexpr bool value = true;
|
||||
};
|
||||
|
||||
|
||||
// The maximum swizzle size to use
|
||||
//
|
||||
// This class, like Splits above makes it harder to confuse
|
||||
@@ -212,9 +224,26 @@ bool initialize_tensor(
|
||||
scope_max = 2;
|
||||
scope_min = 0;
|
||||
}
|
||||
|
||||
else if (bits_input <= 6) {
|
||||
scope_max = 2;
|
||||
scope_min = -2;
|
||||
}
|
||||
|
||||
else if (bits_input <= 8) {
|
||||
|
||||
if constexpr (
|
||||
cute::is_same_v<Element, cutlass::float_ue8m0_t>){
|
||||
scope_max = 4;
|
||||
scope_min = 1;
|
||||
}
|
||||
else {
|
||||
|
||||
scope_max = 1;
|
||||
scope_min = -1;
|
||||
|
||||
}
|
||||
|
||||
}
|
||||
else{
|
||||
scope_max = 4;
|
||||
@@ -487,6 +516,277 @@ struct HostCollectiveMainloop {
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
//
|
||||
// Block Scaled Gemm Input Operands : A , B, scalefactorA, scalefactorB
|
||||
//
|
||||
template<
|
||||
class Gemm,
|
||||
int SchedulerPipelineStageCount_,
|
||||
int AccumulatorPipelineStageCount_,
|
||||
class ElementA_,
|
||||
class ElementB_
|
||||
>
|
||||
struct HostCollectiveMainloop<cutlass::gemm::KernelPtrArrayTmaWarpSpecializedBlockScaledSm100<
|
||||
SchedulerPipelineStageCount_,
|
||||
AccumulatorPipelineStageCount_>,
|
||||
Gemm, ElementA_, ElementB_> {
|
||||
// Kernel data types
|
||||
using ElementA = ElementA_;
|
||||
using StrideA = typename Gemm::GemmKernel::StrideA;
|
||||
using InternalStrideA = typename Gemm::GemmKernel::InternalStrideA;
|
||||
using ElementB = ElementB_;
|
||||
using StrideB = typename Gemm::GemmKernel::StrideB;
|
||||
using InternalStrideB = typename Gemm::GemmKernel::InternalStrideB;
|
||||
using ScheduleType = typename Gemm::GemmKernel::CollectiveMainloop::DispatchPolicy::Schedule;
|
||||
using LayoutTagA = cutlass::detail::StrideToLayoutTagA_t<StrideA>;
|
||||
using LayoutTagB = cutlass::detail::StrideToLayoutTagB_t<StrideB>;
|
||||
|
||||
static constexpr bool IsGroupGemm = !cute::is_same_v<StrideA, InternalStrideA>;
|
||||
|
||||
using ElementAccumulator = typename Gemm::GemmKernel::ElementAccumulator;
|
||||
using ElementScalingFactor = ElementAccumulator;
|
||||
using ProblemShapeType = typename Gemm::GemmKernel::ProblemShape;
|
||||
using EpilogueOutputOp = typename Gemm::EpilogueOutputOp;
|
||||
|
||||
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 LayoutSFA = typename Gemm::GemmKernel::CollectiveMainloop::LayoutSFA;
|
||||
using InternalLayoutSFA = typename Gemm::GemmKernel::CollectiveMainloop::InternalLayoutSFA;
|
||||
using LayoutSFB = typename Gemm::GemmKernel::CollectiveMainloop::LayoutSFB;
|
||||
using InternalLayoutSFB = typename Gemm::GemmKernel::CollectiveMainloop::InternalLayoutSFB;
|
||||
|
||||
using Arguments = typename Gemm::GemmKernel::MainloopArguments;
|
||||
|
||||
// Whether to use relative equality checks
|
||||
CheckEquality check_relative_equality = CheckEquality::EXACT;
|
||||
|
||||
std::vector<InternalStrideA> stride_a_host;
|
||||
std::vector<InternalStrideB> stride_b_host;
|
||||
cutlass::DeviceAllocation<InternalStrideA> stride_a_device;
|
||||
cutlass::DeviceAllocation<InternalStrideB> stride_b_device;
|
||||
|
||||
std::vector<InternalLayoutSFA> layout_sfa_host;
|
||||
std::vector<InternalLayoutSFB> layout_sfb_host;
|
||||
cutlass::DeviceAllocation<InternalLayoutSFA> layout_sfa_device;
|
||||
cutlass::DeviceAllocation<InternalLayoutSFB> layout_sfb_device;
|
||||
|
||||
typename LayoutTagA::Stride stride_factor_A;
|
||||
typename LayoutTagB::Stride stride_factor_B;
|
||||
|
||||
cutlass::Distribution::Kind init_A;
|
||||
cutlass::Distribution::Kind init_B;
|
||||
|
||||
std::vector<cutlass::HostTensor<ElementA, LayoutTagA>> tensors_A;
|
||||
std::vector<cutlass::HostTensor<ElementB, LayoutTagB>> tensors_B;
|
||||
std::vector<cutlass::HostTensor<ElementSF, LayoutTagA>> tensors_SFA;
|
||||
std::vector<cutlass::HostTensor<ElementSF, LayoutTagB>> tensors_SFB;
|
||||
|
||||
cutlass::DeviceAllocation<const ElementA *> device_tensors_A;
|
||||
cutlass::DeviceAllocation<const ElementB *> device_tensors_B;
|
||||
cutlass::DeviceAllocation<const ElementSF *> device_tensors_SFA;
|
||||
cutlass::DeviceAllocation<const ElementSF *> device_tensors_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<InternalStrideA>(),
|
||||
"ERROR : A Layout is neither Row / Column Major)");
|
||||
static_assert(is_row_or_col_major<InternalStrideB>(),
|
||||
"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()
|
||||
):
|
||||
check_relative_equality(check_relative_equality_),
|
||||
stride_factor_A(stride_factor_A_),
|
||||
stride_factor_B(stride_factor_B_),
|
||||
init_A(init_A_), init_B(init_B_), seed(seed_) { }
|
||||
|
||||
template<class ProblemShapeType>
|
||||
bool initialize(ProblemShapeType problem_shapes) {
|
||||
//
|
||||
// Allocate the GEMM workspace
|
||||
//
|
||||
tensors_A.clear();
|
||||
tensors_B.clear();
|
||||
stride_a_host.clear();
|
||||
stride_b_host.clear();
|
||||
tensors_SFA.clear();
|
||||
tensors_SFB.clear();
|
||||
layout_sfa_host.clear();
|
||||
layout_sfb_host.clear();
|
||||
|
||||
auto [M, N, K, L] = cute::append<4>(problem_shapes.get_host_problem_shape(0), 1);
|
||||
L = std::max(problem_shapes.groups(), L);
|
||||
|
||||
for (int32_t i = 0; i < L; ++i) {
|
||||
auto [M, N, K, mock_L] = cute::append<4>(problem_shapes.get_host_problem_shape(i), 1);
|
||||
|
||||
stride_a_host.push_back(cutlass::make_cute_packed_stride(InternalStrideA{}, {M, K, 1}));
|
||||
stride_b_host.push_back(cutlass::make_cute_packed_stride(InternalStrideB{}, {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 a_coord = cutlass::make_Coord(M, K);
|
||||
// Cutlass has Row/Col major refers to MxK times KxN matrix product,
|
||||
// so the HostTensorB should be treated as KxN in "coord"'s view
|
||||
auto b_coord = cutlass::make_Coord(K, N);
|
||||
|
||||
tensors_A.push_back(cutlass::HostTensor<ElementA, LayoutTagA>(a_coord, cutlass::layout::Affine2Layout_Factory<LayoutTagA>::layout_factory(a_coord, stride_factor_A)));
|
||||
tensors_B.push_back(cutlass::HostTensor<ElementB, LayoutTagB>(b_coord, cutlass::layout::Affine2Layout_Factory<LayoutTagB>::layout_factory(b_coord, stride_factor_B)));
|
||||
|
||||
EXPECT_TRUE(initialize_tensor(tensors_A[i].host_view(), init_A, seed + 2022 + i));
|
||||
EXPECT_TRUE(initialize_tensor(tensors_B[i].host_view(), init_B, seed + 2021 + i));
|
||||
|
||||
// It is possible to randomly initialize to all zeros, so override this with non-zeros
|
||||
// in the upper left corner of each operand.
|
||||
tensors_A[i].host_view().at({0, 0}) = ElementA(1);
|
||||
tensors_B[i].host_view().at({0, 0}) = ElementB(1);
|
||||
|
||||
tensors_A[i].sync_device();
|
||||
tensors_B[i].sync_device();
|
||||
|
||||
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_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)));
|
||||
|
||||
// 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{});
|
||||
auto sfb_coord = cutlass::make_Coord(n_blks * Blk_MN{}, k_blks * Blk_SF{});
|
||||
|
||||
tensors_SFA.push_back(cutlass::HostTensor<ElementSF, LayoutTagA>(sfa_coord, cutlass::layout::Affine2Layout_Factory<LayoutTagA>::layout_factory(sfa_coord, stride_factor_A)));
|
||||
tensors_SFB.push_back(cutlass::HostTensor<ElementSF, LayoutTagB>(sfb_coord, cutlass::layout::Affine2Layout_Factory<LayoutTagB>::layout_factory(sfb_coord, stride_factor_B)));
|
||||
|
||||
EXPECT_TRUE(initialize_tensor(tensors_SFA[i].host_view(), init_A, seed + 2024 + i));
|
||||
EXPECT_TRUE(initialize_tensor(tensors_SFB[i].host_view(), init_B, seed + 2025 + i));
|
||||
|
||||
// It is possible to randomly initialize to all zeros, so override this with non-zeros
|
||||
// in the upper left corner of each operand.
|
||||
tensors_SFA[i].host_view().at({0, 0}) = ElementSF(1);
|
||||
tensors_SFB[i].host_view().at({0, 0}) = ElementSF(1);
|
||||
|
||||
tensors_SFA[i].sync_device();
|
||||
tensors_SFB[i].sync_device();
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
Arguments to_args(ProblemShapeType problem_shapes) {
|
||||
auto [M, N, K, L] = cute::append<4>(problem_shapes.get_host_problem_shape(0), 1);
|
||||
L = std::max(problem_shapes.groups(), L);
|
||||
|
||||
std::vector<ElementA *> ptr_A_host(L);
|
||||
std::vector<ElementB *> ptr_B_host(L);
|
||||
std::vector<ElementSF *> ptr_SFA_host(L);
|
||||
std::vector<ElementSF *> ptr_SFB_host(L);
|
||||
|
||||
for (int32_t i = 0; i < L; ++i) {
|
||||
ptr_A_host.at(i) = tensors_A[i].device_data();
|
||||
ptr_B_host.at(i) = tensors_B[i].device_data();
|
||||
ptr_SFA_host.at(i) = tensors_SFA[i].device_data();
|
||||
ptr_SFB_host.at(i) = tensors_SFB[i].device_data();
|
||||
}
|
||||
|
||||
device_tensors_A.reset(L);
|
||||
device_tensors_A.copy_from_host(ptr_A_host.data());
|
||||
|
||||
device_tensors_B.reset(L);
|
||||
device_tensors_B.copy_from_host(ptr_B_host.data());
|
||||
|
||||
device_tensors_SFA.reset(L);
|
||||
device_tensors_SFA.copy_from_host(ptr_SFA_host.data());
|
||||
|
||||
device_tensors_SFB.reset(L);
|
||||
device_tensors_SFB.copy_from_host(ptr_SFB_host.data());
|
||||
|
||||
stride_a_device.reset(problem_shapes.groups());
|
||||
stride_a_device.copy_from_host(stride_a_host.data());
|
||||
|
||||
stride_b_device.reset(problem_shapes.groups());
|
||||
stride_b_device.copy_from_host(stride_b_host.data());
|
||||
|
||||
layout_sfa_device.reset(problem_shapes.groups());
|
||||
layout_sfa_device.copy_from_host(layout_sfa_host.data());
|
||||
|
||||
layout_sfb_device.reset(problem_shapes.groups());
|
||||
layout_sfb_device.copy_from_host(layout_sfb_host.data());
|
||||
|
||||
if constexpr (IsGroupGemm) {
|
||||
return Arguments{
|
||||
device_tensors_A.get(), stride_a_device.get(),
|
||||
device_tensors_B.get(), stride_b_device.get(),
|
||||
device_tensors_SFA.get(), layout_sfa_device.get(),
|
||||
device_tensors_SFB.get(), layout_sfb_device.get()
|
||||
};
|
||||
}
|
||||
else {
|
||||
return Arguments{
|
||||
device_tensors_A.get(), stride_a_host[0],
|
||||
device_tensors_B.get(), stride_b_host[0],
|
||||
device_tensors_SFA.get(), layout_sfa_host[0],
|
||||
device_tensors_SFB.get(), layout_sfb_host[0]
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
auto to_host_args(ProblemShapeType problem_shapes, int batch) {
|
||||
using namespace cute;
|
||||
//
|
||||
// Allocate the GEMM workspace
|
||||
//
|
||||
auto [M, N, K, L] = cute::append<4>(problem_shapes.get_host_problem_shape(batch), 1);
|
||||
auto A = make_tensor(make_iterator(tensors_A[batch].host_data()),
|
||||
make_layout(make_shape(M, K, 1), stride_a_host[batch]));
|
||||
auto SfA = make_tensor(tensors_SFA[batch].host_data(), layout_sfa_host[batch]);
|
||||
|
||||
auto B = make_tensor(make_iterator(tensors_B[batch].host_data()),
|
||||
make_layout(make_shape(N, K, 1), stride_b_host[batch]));
|
||||
auto SfB = make_tensor(tensors_SFB[batch].host_data(), layout_sfb_host[batch]);
|
||||
|
||||
return cutlass::reference::host::GettMainloopParams<ElementAccumulator,
|
||||
decltype(A),
|
||||
decltype(B),
|
||||
decltype(SfA),
|
||||
decltype(SfB)
|
||||
>
|
||||
{A, SfA, B, SfB};
|
||||
}
|
||||
|
||||
void print_tensors(std::ofstream& file, int batch) {
|
||||
file << "A =\n" << tensors_A[batch].host_view()
|
||||
<< "\nB =\n" << tensors_B[batch].host_view()
|
||||
<< "\nSFA =\n" << tensors_SFA[batch].host_view()
|
||||
<< "\nSFB =\n" << tensors_SFB[batch].host_view();
|
||||
}
|
||||
|
||||
bool compare_reference(
|
||||
ProblemShapeType problem_shapes, int batch) {
|
||||
|
||||
EXPECT_GT(cutlass::reference::host::TensorNorm(tensors_A[batch].host_view()), 0);
|
||||
EXPECT_GT(cutlass::reference::host::TensorNorm(tensors_B[batch].host_view()), 0);
|
||||
EXPECT_GT(cutlass::reference::host::TensorNorm(tensors_SFA[batch].host_view()), 0);
|
||||
EXPECT_GT(cutlass::reference::host::TensorNorm(tensors_SFB[batch].host_view()), 0);
|
||||
return true;
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
template<class Gemm>
|
||||
struct HostCollectiveDefaultEpilogue {
|
||||
// fusion types are potentially void if the fusion is not supported
|
||||
@@ -803,6 +1103,24 @@ struct HostCollectiveEpilogue {
|
||||
using FusionOp = typename Gemm::EpilogueOutputOp;
|
||||
static_assert(cute::is_base_of_v<cutlass::epilogue::fusion::FusionOperation, FusionOp>);
|
||||
|
||||
|
||||
// Scale factor Generation related
|
||||
using SfStrategy = cutlass::reference::host::SfStrategy;
|
||||
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;
|
||||
using ElementSFD = non_void_t<cute::remove_pointer_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;
|
||||
std::vector<cutlass::HostTensor<ElementSFD, LayoutTagD>> tensors_SFD;
|
||||
std::vector<cutlass::HostTensor<ElementSFD, LayoutTagD>> references_SFD;
|
||||
cutlass::DeviceAllocation<ElementSFD *> device_tensors_SFD;
|
||||
|
||||
|
||||
using ElementCompute = typename FusionOp::ElementCompute;
|
||||
using ElementScalar = typename FusionOp::ElementScalar;
|
||||
using ElementBias = non_void_t<typename FusionOp::ElementBias>;
|
||||
@@ -904,6 +1222,11 @@ struct HostCollectiveEpilogue {
|
||||
references_D.clear();
|
||||
stride_c_host.clear();
|
||||
stride_d_host.clear();
|
||||
|
||||
tensors_SFD.clear();
|
||||
references_SFD.clear();
|
||||
|
||||
|
||||
auto [M, N, K, L] = cute::append<4>(problem_shapes.get_host_problem_shape(0), 1);
|
||||
L = std::max(problem_shapes.groups(), L);
|
||||
|
||||
@@ -1034,6 +1357,26 @@ struct HostCollectiveEpilogue {
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
if constexpr (IsBlockScaleSupported) {
|
||||
for (int32_t i = 0; i < L; ++i) {
|
||||
auto [M, N, K, _] = cute::append<4>(problem_shapes.get_host_problem_shape(i), 1);
|
||||
// If block scaled output is supported we always have at least 1 SFD
|
||||
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 = [&] () {
|
||||
return cutlass::make_Coord(m_blks * Blk_MN{}, n_blks * Blk_SF{});
|
||||
}();
|
||||
tensors_SFD.push_back(cutlass::HostTensor<ElementSFD, LayoutTagD>(sfd_coord, cutlass::layout::Affine2Layout_Factory<LayoutTagD>::layout_factory(sfd_coord, stride_factor_D)));
|
||||
references_SFD.push_back(cutlass::HostTensor<ElementSFD, LayoutTagD>(sfd_coord, cutlass::layout::Affine2Layout_Factory<LayoutTagD>::layout_factory(sfd_coord, stride_factor_D), false));
|
||||
tensors_SFD[i].sync_device();
|
||||
}
|
||||
norm_constant.resize(scalar_coord, true);
|
||||
EXPECT_TRUE(initialize_tensor(norm_constant.host_view(), init_scale, seed + 2023));
|
||||
norm_constant.sync_device();
|
||||
}
|
||||
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
@@ -1116,6 +1459,17 @@ struct HostCollectiveEpilogue {
|
||||
passed &= tmp;
|
||||
}
|
||||
}
|
||||
|
||||
if constexpr (IsBlockScaleSupported) {
|
||||
tensors_SFD[batch].sync_host();
|
||||
bool passed_sf = equality_check(references_SFD[batch].host_view(), tensors_SFD[batch].host_view());
|
||||
if(!passed_sf) {
|
||||
std::cout<<"SF is incorrect"<<std::endl;
|
||||
}
|
||||
passed &= passed_sf;
|
||||
}
|
||||
|
||||
|
||||
return passed;
|
||||
}
|
||||
|
||||
@@ -1308,6 +1662,19 @@ struct HostCollectiveEpilogue {
|
||||
fusion_args.amax_aux_ptr = abs_max_Aux.device_data();
|
||||
}
|
||||
}
|
||||
|
||||
if constexpr (IsBlockScaleSupported) {
|
||||
std::vector<ElementSFD *> ptr_SFD_host(L);
|
||||
for (int32_t i = 0; i < L; ++i) {
|
||||
ptr_SFD_host.at(i) = tensors_SFD[i].device_data();
|
||||
}
|
||||
device_tensors_SFD.reset(L);
|
||||
device_tensors_SFD.copy_from_host(ptr_SFD_host.data());
|
||||
|
||||
arguments.thread.block_scale_factor_ptr = device_tensors_SFD.get();
|
||||
arguments.thread.norm_constant_ptr = norm_constant.device_data();
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
return arguments;
|
||||
@@ -1341,6 +1708,20 @@ struct HostCollectiveEpilogue {
|
||||
cute::make_layout(cute::make_shape(M, N, cute::_1{}), cute::make_stride(cute::_1{}, cute::_0{}, M)));
|
||||
auto Vbeta = cute::make_tensor(detail::make_iterator(beta.host_data()),
|
||||
cute::make_layout(cute::make_shape(M, N, cute::_1{}), cute::make_stride(cute::_1{}, cute::_0{}, N)));
|
||||
|
||||
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));
|
||||
return tensor;
|
||||
}
|
||||
else {
|
||||
// Reference kernel has a logic to ignore scalefactor computation if we pass the tensor type same as output D tensor.
|
||||
return D;
|
||||
}
|
||||
}();
|
||||
|
||||
|
||||
cutlass::reference::host::GettEpilogueParams<
|
||||
ElementScalar,
|
||||
ElementScalar,
|
||||
@@ -1353,8 +1734,11 @@ struct HostCollectiveEpilogue {
|
||||
decltype(Valpha),
|
||||
decltype(Vbeta),
|
||||
ActivationFunctor
|
||||
, decltype(SfD)
|
||||
, Int<SFD_VectorSize>
|
||||
, cutlass::plus<ElementCompute>
|
||||
, false
|
||||
, SfGenStrategy
|
||||
> epilogue_params{};
|
||||
|
||||
epilogue_params.C = C;
|
||||
@@ -1397,6 +1781,12 @@ struct HostCollectiveEpilogue {
|
||||
epilogue_params.Vbeta = Vbeta;
|
||||
}
|
||||
}
|
||||
|
||||
if constexpr (IsBlockScaleSupported) {
|
||||
epilogue_params.SfD = SfD;
|
||||
epilogue_params.st = norm_constant.at(coord_0);
|
||||
}
|
||||
|
||||
return epilogue_params;
|
||||
}
|
||||
};
|
||||
@@ -1812,8 +2202,24 @@ bool TestSmall(double alpha = 1.0, double beta = 1.0,
|
||||
using ElementB = typename Gemm::GemmKernel::ElementB;
|
||||
using TiledMma = typename Gemm::GemmKernel::TiledMma;
|
||||
int alignment_bits = 128;
|
||||
|
||||
static constexpr bool IsF8F6F4 = cutlass::gemm::collective::detail::is_sm100_mma_f8f6f4<TiledMma, ElementA, ElementB>();
|
||||
alignment_bits = cutlass::detail::get_input_alignment_bits<ElementA, IsF8F6F4>();
|
||||
// For fp4 and fp6 mx kernels, the min alignment_input is 128 elements, so we don't need to add alignment_input in test problem sizes.
|
||||
|
||||
int alignment_input = (alignment_bits / cute::sizeof_bits<ElementA>::value == 128) ? 0 : (alignment_bits / cute::sizeof_bits<ElementA>::value);
|
||||
|
||||
|
||||
if constexpr (apply_alignment_offset) {
|
||||
// If BlockScaled, then min alignment is SFVecSize
|
||||
static constexpr bool IsBlockScaleSupported = Gemm::EpilogueOutputOp::IsBlockScaleSupported;
|
||||
static constexpr int SFVecSize = Gemm::GemmKernel::CollectiveMainloop::SFVecSize;
|
||||
if constexpr (IsBlockScaleSupported) {
|
||||
alignment_input = cutlass::round_up(alignment_input, SFVecSize);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
using CtaShape_MNK = typename Gemm::GemmKernel::CollectiveMainloop::CtaShape_MNK;
|
||||
using DispatchPolicy = typename Gemm::GemmKernel::CollectiveMainloop::DispatchPolicy;
|
||||
CtaShape_MNK cta_shape;
|
||||
|
||||
Reference in New Issue
Block a user