v3.9 (#2185)
* v3.8 update x * fix blackwell gg * doc change * doc change * doc change --------- Co-authored-by: yuzhai <yuzhai@nvidia.com> Co-authored-by: Haicheng Wu <haichengw@nvidia.com> Co-authored-by: Haicheng Wu <57973641+hwu36@users.noreply.github.com>
This commit is contained in:
co-authored by
yuzhai
Haicheng Wu
Haicheng Wu
parent
8c4d1dc47d
commit
62750a2b75
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user