v3.9 update (#2203)

* v3.9 update

* voidD

---------

Co-authored-by: yuzhai <yuzhai@nvidia.com>
This commit is contained in:
Yujia Zhai
2025-04-02 15:11:18 -04:00
committed by GitHub
co-authored by yuzhai
parent 62750a2b75
commit 6f4921858b
129 changed files with 7719 additions and 2036 deletions
@@ -113,6 +113,122 @@ sm100_compute_stage_count_or_override_blockwise(StageCountAutoCarveout<carveout_
return (CapacityBytes - carveout_bytes) / stage_bytes;
}
template<class Element, typename LayoutSFA, class CtaShape_MNK>
auto sm100_make_simt_gmem_tiled_copy_SFA() {
// we have at most a warp to perform the loads
constexpr int ScaleGranularityM = size<0,0>(LayoutSFA{});
constexpr int ScaleMsPerTile = size<0>(CtaShape_MNK{}) / ScaleGranularityM;
constexpr int ScaleGranularityK = size<1,0>(LayoutSFA{});
constexpr int ScaleKsPerTile = size<2>(CtaShape_MNK{}) / ScaleGranularityK;
if constexpr (size<0,1>(LayoutSFA{}.stride()) == 1) {
constexpr int LeadingScalesPerTileSFA = ScaleMsPerTile;
if constexpr (LeadingScalesPerTileSFA >= 32) {
constexpr int Alignment = cute::min(static_cast<int>(LeadingScalesPerTileSFA * sizeof(Element)) / 32, 16);
using ScaleCopyTypeA = cute::uint_byte_t<Alignment>;
using SmemScalingCopyAtomA = Copy_Atom<SM80_CP_ASYNC_CACHEALWAYS<ScaleCopyTypeA>, Element>;
constexpr int ElementsPerSFACopy = static_cast<int>(sizeof(ScaleCopyTypeA) / sizeof(Element));
return make_tiled_copy(SmemScalingCopyAtomA{}, Layout<Shape<_32>>{}, Layout<Shape<Int<ElementsPerSFACopy>>>{});
}
else {
using SmemScalingCopyAtomA = Copy_Atom<SM80_CP_ASYNC_CACHEALWAYS<Element>, Element>;
return make_tiled_copy(SmemScalingCopyAtomA{}, Layout<Shape<Int<LeadingScalesPerTileSFA>>>{}, Layout<Shape<_1>>{});
}
}
else {
// we expect scale Ks per tile to be small
constexpr int LeadingScalesPerTileSFA = ScaleKsPerTile;
using SmemScalingCopyAtomA = Copy_Atom<SM80_CP_ASYNC_CACHEALWAYS<Element>, Element>;
return make_tiled_copy(SmemScalingCopyAtomA{}, Layout<Shape<_1, Int<LeadingScalesPerTileSFA>>>{}, Layout<Shape<_1,_1>>{});
}
}
template<class Element, typename LayoutSFB, class CtaShape_MNK>
auto sm100_make_simt_gmem_tiled_copy_SFB() {
// we have at most a warp to perform the loads
constexpr int ScaleGranularityN = size<0,0>(LayoutSFB{});
constexpr int ScaleNsPerTile = size<1>(CtaShape_MNK{}) / ScaleGranularityN;
constexpr int ScaleGranularityK = size<1,0>(LayoutSFB{});
constexpr int ScaleKsPerTile = size<2>(CtaShape_MNK{}) / ScaleGranularityK;
if constexpr (size<0,1>(LayoutSFB{}.stride()) == 1) {
constexpr int LeadingScalesPerTileSFB = ScaleNsPerTile;
if constexpr (LeadingScalesPerTileSFB >= 32) {
constexpr int Alignment = cute::min(static_cast<int>(LeadingScalesPerTileSFB * sizeof(Element)) / 32, 16);
using ScaleCopyTypeB = cute::uint_byte_t<Alignment>;
using SmemScalingCopyAtomB = Copy_Atom<SM80_CP_ASYNC_CACHEALWAYS<ScaleCopyTypeB>, Element>;
constexpr int ElementsPerSFBCopy = static_cast<int>(sizeof(ScaleCopyTypeB) / sizeof(Element));
return make_tiled_copy(SmemScalingCopyAtomB{}, Layout<Shape<_32>>{}, Layout<Shape<Int<ElementsPerSFBCopy>>>{});
}
else {
using SmemScalingCopyAtomB = Copy_Atom<SM80_CP_ASYNC_CACHEALWAYS<Element>, Element>;
return make_tiled_copy(SmemScalingCopyAtomB{}, Layout<Shape<Int<LeadingScalesPerTileSFB>>>{}, Layout<Shape<_1>>{});
}
}
else {
// we expect scale Ks per tile to be small
constexpr int LeadingScalesPerTileSFB = ScaleKsPerTile;
using SmemScalingCopyAtomB = Copy_Atom<SM80_CP_ASYNC_CACHEALWAYS<Element>, Element>;
return make_tiled_copy(SmemScalingCopyAtomB{}, Layout<Shape<_1, Int<LeadingScalesPerTileSFB>>>{}, Layout<Shape<_1,_1>>{});
}
}
// For new MMA construction and partitioning that supports both dynamic and static cluster shape.
// Used in conjunction with make_tma_atom_(A|B)_sm100
// TileShape_MNK is always static and has shape (MmaAtomShapeM, MmaAtomShapeN, TileK)
// ClusterShape_MNK can be dynamic or static.
template<
class ElementAMma,
class ElementBMma,
class ElementAccumulator,
class TileShape_MNK,
class ClusterShape_MNK,
UMMA::Major UmmaMajorA,
UMMA::Major UmmaMajorB,
class BuilderScheduleTag,
UMMA::ScaleIn ANeg = UMMA::ScaleIn::One,
UMMA::ScaleIn BNeg = UMMA::ScaleIn::One
>
constexpr auto
sm100_make_trivial_tiled_mma_blockwise() {
// MMA_2SM requested
if constexpr (cute::is_base_of_v<KernelSchedule2Sm, BuilderScheduleTag> ) {
return sm100_make_2sm_trivial_tiled_mma<ElementAMma, ElementBMma, ElementAccumulator,
TileShape_MNK, ClusterShape_MNK, UmmaMajorA, UmmaMajorB, ANeg, BNeg>();
}
// MMA_1SM requested
else if constexpr (cute::is_base_of_v<KernelSchedule1Sm, BuilderScheduleTag> ) {
return sm100_make_1sm_trivial_tiled_mma<ElementAMma, ElementBMma, ElementAccumulator,
TileShape_MNK, ClusterShape_MNK, UmmaMajorA, UmmaMajorB, ANeg, BNeg>();
}
// Auto scheduling requested
else if constexpr (cute::is_same_v<BuilderScheduleTag, KernelScheduleSm100Blockwise>) {
// Static cluster
if constexpr (cute::is_static_v<ClusterShape_MNK>) {
// For MMA_2SM we need a cluster shape that is multiple of 2x1
// and only M=128 and M=256 are supported, otherwise, fall back to MMA_1SM
if constexpr (cute::size<0>(ClusterShape_MNK{}) % 2 == 0 &&
cute::size<0>(TileShape_MNK{}) % 128 == 0) {
return sm100_make_2sm_trivial_tiled_mma<ElementAMma, ElementBMma, ElementAccumulator,
TileShape_MNK, ClusterShape_MNK, UmmaMajorA, UmmaMajorB, ANeg, BNeg>();
}
else {
return sm100_make_1sm_trivial_tiled_mma<ElementAMma, ElementBMma, ElementAccumulator,
TileShape_MNK, ClusterShape_MNK, UmmaMajorA, UmmaMajorB, ANeg, BNeg>();
}
// Dynamic cluster shape means we cannot assume we can use 2SM MMA
}
else {
return sm100_make_1sm_trivial_tiled_mma<ElementAMma, ElementBMma, ElementAccumulator,
TileShape_MNK, ClusterShape_MNK, UmmaMajorA, UmmaMajorB, ANeg, BNeg>();
}
}
}
} // namespace detail
/////////////////////////////////////////////////////////////////////////////////////////////////
@@ -161,9 +277,11 @@ struct CollectiveBuilder<
using GmemLayoutBTag = cute::remove_cvref_t<decltype(get<0>(GmemLayoutBTagPair{}))>;
using GmemLayoutSFBTag = cute::remove_cvref_t<decltype(get<1>(GmemLayoutBTagPair{}))>;
static_assert(cute::depth(GmemLayoutSFATag{}) == 2 and cute::depth(GmemLayoutSFBTag{}) == 2,
static_assert(cute::depth(cute::remove_pointer_t<GmemLayoutSFATag>{}) == 2 and
cute::depth(cute::remove_pointer_t<GmemLayoutSFBTag>{}) == 2,
"Expect SFA and SFB layout to be depth of two with shape ((SFVecMN, restMN),(SFVecK, restK), L)");
static_assert(size<1,0>(GmemLayoutSFATag{}) == size<1, 0>(GmemLayoutSFBTag{}),
static_assert(size<1,0>(cute::remove_pointer_t<GmemLayoutSFATag>{}) ==
size<1,0>(cute::remove_pointer_t<GmemLayoutSFBTag>{}),
"SFA and SFB must have equivalent SF vector sizes along K");
static constexpr cute::UMMA::Major UmmaMajorA = cutlass::gemm::collective::detail::tag_to_umma_major_A<GmemLayoutATag>();
@@ -183,7 +301,7 @@ struct CollectiveBuilder<
TileShape_MNK, ClusterShape_MNK,
GmemLayoutATag, GmemLayoutBTag, false /*is_sparse*/, is_2sm>(),
"TileSize and MNK Major does not met with MMA Mix 8-bit TMA load requirement" );
using TiledMma = decltype(detail::sm100_make_trivial_tiled_mma<
using TiledMma = decltype(detail::sm100_make_trivial_tiled_mma_blockwise<
ElementAMma, ElementBMma, ElementAccumulator,
decltype(cute::product_each(TileShape_MNK{})), ClusterShape_MNK,
UmmaMajorA, UmmaMajorB, BuilderScheduleTag>());
@@ -238,12 +356,14 @@ struct CollectiveBuilder<
// SchedulerPipelineStageCount could be set to zero for Grouped GEMM, but we shouldn't define CLC Pipeline's barrier arrays of size zero.
static constexpr uint32_t SchedulerPipelineStageCount = cute::is_same_v<InternalStrideA, StrideA> ? (AccumulatorPipelineStageCount + 1) : 1;
static constexpr bool IsArrayOfPointersGemm = (cute::is_base_of_v<KernelScheduleSm100PtrArrayBlockwise, BuilderScheduleTag>);
static constexpr uint32_t KernelSmemCarveout = detail::Sm100DenseGemmTmaUmmaCarveout<
ClusterShape_MNK,
AccumulatorPipelineStageCount,
SchedulerPipelineStageCount,
detail::CLCResponseSize,
false
IsArrayOfPointersGemm
>::KernelSmemCarveout;
// Reduce SMEM capacity available for buffers considering barrier allocations.
static constexpr int Sm100ReducedSmemCapacityBytes = cutlass::gemm::collective::detail::sm100_smem_capacity_bytes - KernelSmemCarveout;
@@ -253,14 +373,23 @@ struct CollectiveBuilder<
using TransformLoadPipelineStorage = typename cutlass::PipelineAsync<1>::SharedStorage;
using TransformPipelineStorage = typename cutlass::PipelineUmmaAsync<1>::SharedStorage;
static constexpr int ScaleGranularityM = size<0,0>(GmemLayoutSFATag{});
static constexpr int ScaleGranularityN = size<0,0>(GmemLayoutSFBTag{});
static constexpr int ScaleGranularityK = size<1,0>(GmemLayoutSFBTag{});
static constexpr int ScaleGranularityM = size<0,0>(cute::remove_pointer_t<GmemLayoutSFATag>{});
static constexpr int ScaleGranularityN = size<0,0>(cute::remove_pointer_t<GmemLayoutSFBTag>{});
static constexpr int ScaleGranularityK = size<1,0>(cute::remove_pointer_t<GmemLayoutSFBTag>{});
static_assert(size<0>(CtaTileShape_MNK{}) >= ScaleGranularityM, "Scale Granularity must be smaller than or equal to the tile shape");
static_assert(size<1>(CtaTileShape_MNK{}) >= ScaleGranularityN, "Scale Granularity must be smaller than or equal to the tile shape");
static_assert(size<2>(CtaTileShape_MNK{}) >= ScaleGranularityK, "Scale Granularity must be smaller than or equal to the tile shape");
using GmemTiledCopySFA = decltype(detail::sm100_make_simt_gmem_tiled_copy_SFA<
ElementAccumulator,
cute::remove_pointer_t<GmemLayoutSFATag>,
CtaTileShape_MNK>());
using GmemTiledCopySFB = decltype(detail::sm100_make_simt_gmem_tiled_copy_SFB<
ElementAccumulator,
cute::remove_pointer_t<GmemLayoutSFBTag>,
CtaTileShape_MNK>());
using BlockTileScale_M = Int<size<0>(TileShape_MNK{}) / ScaleGranularityM>;
using BlockTileScale_N = Int<size<1>(TileShape_MNK{}) / ScaleGranularityN>;
using BlockTileScale_K = Int<size<2>(TileShape_MNK{}) / ScaleGranularityK>;
@@ -273,11 +402,18 @@ struct CollectiveBuilder<
TransformLoadPipelineStorage, TransformPipelineStorage>(StageCountType{});
static_assert(PipelineStages > 0, "Smem usage is too high. Can't create any SMEM buffers for A, B, and scales.");
using DispatchPolicy = cutlass::gemm::MainloopSm100TmaUmmaWarpSpecializedBlockwiseScaling<
using DispatchPolicy = cute::conditional_t<
IsArrayOfPointersGemm,
cutlass::gemm::MainloopSm100ArrayTmaUmmaWarpSpecializedBlockwiseScaling<
PipelineStages,
SchedulerPipelineStageCount,
AccumulatorPipelineStageCount,
ClusterShape_MNK>;
ClusterShape_MNK>,
cutlass::gemm::MainloopSm100TmaUmmaWarpSpecializedBlockwiseScaling<
PipelineStages,
SchedulerPipelineStageCount,
AccumulatorPipelineStageCount,
ClusterShape_MNK>>;
using CollectiveOp = cutlass::gemm::collective::CollectiveMma<
DispatchPolicy,
@@ -287,11 +423,11 @@ struct CollectiveBuilder<
ElementB,
cute::tuple<cutlass::gemm::TagToStrideB_t<GmemLayoutBTag>, cutlass::gemm::TagToStrideB_t<GmemLayoutSFBTag>>,
TiledMma,
GmemTiledCopyA,
cute::tuple<GmemTiledCopyA, GmemTiledCopySFA>,
SmemLayoutAtomA,
void,
cute::identity,
GmemTiledCopyB,
cute::tuple<GmemTiledCopyB, GmemTiledCopySFB>,
SmemLayoutAtomB,
void,
cute::identity
@@ -104,10 +104,10 @@ struct CollectiveBuilder<
UmmaMajorB,
BuilderScheduleTag>();
static constexpr bool UseMxf8f6f4 = Instr == detail::blockscaled::BlockScaledInstr::MXF4F6F8;
using PermTileM = decltype(cute::min(size<0>(TileShape_MNK{}), _128{}));
using PermTileN = decltype(detail::sm120_tile_n_permute_selector<SFVectorSize>());
using PermTileK = cute::conditional_t<UseMxf8f6f4, _32, _64>;
using PermTileK = cute::conditional_t<(UseMxf8f6f4
), _32, _64>;
static constexpr bool IsCooperative = !cute::is_base_of_v<KernelTmaWarpSpecializedPingpong, BuilderScheduleTag>;
// Data type used by MMA instruction
@@ -124,7 +124,13 @@ struct CollectiveBuilder<
Layout<Shape<_4,_2,_1>>, Layout<Shape<_2,_2,_1>>>;
using TiledMma = decltype(cute::make_tiled_mma(
cute::rr_blockscaled_op_selector_sm120<ElementA, ElementB, ElementAccumulator, ElementSF, SFVectorSize, UseMxf8f6f4>(),
cute::rr_blockscaled_op_selector_sm120<ElementA,
ElementB,
ElementAccumulator,
ElementSF,
SFVectorSize,
UseMxf8f6f4
>(),
AtomLayoutMNK{},
Tile<PermTileM, PermTileN, PermTileK>{}
));
@@ -150,8 +156,14 @@ struct CollectiveBuilder<
using SmemLayoutAtomA = decltype(detail::sm120_rr_smem_selector<SmemAllocTypeA, decltype(size<2>(TileShape_MNK{}))>());
using SmemLayoutAtomB = decltype(detail::sm120_rr_smem_selector<SmemAllocTypeB, decltype(size<2>(TileShape_MNK{}))>());
using SmemCopyAtomA = Copy_Atom<decltype(detail::sm120_rr_smem_copy_selector_A<ElementA, ElementB, UseMxf8f6f4>()), SmemAllocTypeA>;
using SmemCopyAtomB = Copy_Atom<decltype(detail::sm120_rr_smem_copy_selector_B<ElementA, ElementB, UseMxf8f6f4>()), SmemAllocTypeB>;
using SmemCopyAtomA = Copy_Atom<decltype(detail::sm120_rr_smem_copy_selector_A<ElementA,
ElementB,
UseMxf8f6f4
>()), SmemAllocTypeA>;
using SmemCopyAtomB = Copy_Atom<decltype(detail::sm120_rr_smem_copy_selector_B<ElementA,
ElementB,
UseMxf8f6f4
>()), SmemAllocTypeB>;
using SmemCopyAtomSF = Copy_Atom<UniversalCopy<SmemAllocTypeSF>, SmemAllocTypeSF>; // auto-vectorized LDS
using SmemCopyAtomSFA = SmemCopyAtomSF;
@@ -45,7 +45,11 @@ namespace cutlass::gemm::collective::detail {
constexpr int sm120_smem_capacity_bytes = cutlass::arch::sm120_smem_capacity_bytes;
// Helper for selecting the shared memory copy atom to use for operand A
template <class ElementA, class ElementB, bool UseF8f6f4>
template <
class ElementA,
class ElementB,
bool UseF8f6f4
>
CUTLASS_HOST_DEVICE constexpr
auto
sm120_rr_smem_copy_selector_A() {
@@ -66,7 +70,11 @@ sm120_rr_smem_copy_selector_A() {
}
// Helper for selecting the shared memory copy atom to use for operand B
template <class ElementA, class ElementB, bool UseF8f6f4>
template <
class ElementA,
class ElementB,
bool UseF8f6f4
>
CUTLASS_HOST_DEVICE constexpr
auto
sm120_rr_smem_copy_selector_B() {
@@ -467,6 +467,8 @@ check_input_datatypes() {
// SfVectorSize = 64 for blockscaled sparse gemm
static_assert(
((SfVectorSizeA == 32 && cute::is_same_v<KernelScheduleAuto, BuilderScheduleTag>)
|| (SfVectorSizeA == 32 && cute::is_same_v<KernelTmaWarpSpecializedPingpong, BuilderScheduleTag>)
|| (SfVectorSizeA == 32 && cute::is_same_v<KernelTmaWarpSpecializedCooperative, BuilderScheduleTag>)
|| (SfVectorSizeA == 32 && cute::is_base_of_v<KernelScheduleBlockScaledGemmSm100, BuilderScheduleTag>)
|| (SfVectorSizeA == 32 && cute::is_base_of_v<KernelSchedulePtrArrayBlockScaledGemmSm100, BuilderScheduleTag>)
|| (SfVectorSizeA == 64 && cute::is_base_of_v<KernelScheduleBlockScaledSparseGemmSm100, BuilderScheduleTag>)
@@ -645,6 +647,8 @@ select_instr() {
static_assert(
(SfVectorSize == 32 && cute::is_same_v<KernelScheduleAuto, BuilderScheduleTag>)
|| (SfVectorSize == 32 && cute::is_base_of_v<KernelScheduleBlockScaledGemmSm100, BuilderScheduleTag>)
|| (SfVectorSize == 32 && cute::is_base_of_v<KernelTmaWarpSpecializedPingpong, BuilderScheduleTag>)
|| (SfVectorSize == 32 && cute::is_base_of_v<KernelTmaWarpSpecializedCooperative, BuilderScheduleTag>)
|| (SfVectorSize == 32 && cute::is_base_of_v<KernelSchedulePtrArrayBlockScaledGemmSm100, BuilderScheduleTag>)
|| (SfVectorSize == 64 && cute::is_base_of_v<KernelScheduleBlockScaledSparseGemmSm100, BuilderScheduleTag>
|| (SfVectorSize == 32 && cute::is_base_of_v<KernelScheduleBlockScaledGemmSm120, BuilderScheduleTag>)
@@ -666,6 +670,8 @@ select_instr() {
else {
static_assert(
((SfVectorSize == 32 && cute::is_same_v<KernelScheduleAuto, BuilderScheduleTag>)
|| (SfVectorSize == 32 && cute::is_base_of_v<KernelTmaWarpSpecializedPingpong, BuilderScheduleTag>)
|| (SfVectorSize == 32 && cute::is_base_of_v<KernelTmaWarpSpecializedCooperative, BuilderScheduleTag>)
|| (SfVectorSize == 32 && cute::is_base_of_v<KernelScheduleBlockScaledGemmSm100, BuilderScheduleTag>)
|| (SfVectorSize == 32 && cute::is_base_of_v<KernelSchedulePtrArrayBlockScaledGemmSm100, BuilderScheduleTag>)
|| (SfVectorSize == 64 && cute::is_base_of_v<KernelScheduleBlockScaledSparseGemmSm100, BuilderScheduleTag>)
@@ -61,6 +61,7 @@
#include "cutlass/gemm/collective/sm100_blockscaled_mma_warpspecialized.hpp"
#include "cutlass/gemm/collective/sm100_blockscaled_mma_array_warpspecialized.hpp"
#include "cutlass/gemm/collective/sm100_mma_warpspecialized_blockwise_scaling.hpp"
#include "cutlass/gemm/collective/sm100_mma_array_warpspecialized_blockwise_scaling.hpp"
#include "cutlass/gemm/collective/sm120_mma_tma.hpp"
#include "cutlass/gemm/collective/sm120_blockscaled_mma_tma.hpp"
#include "cutlass/gemm/collective/sm120_sparse_mma_tma.hpp"
@@ -28,8 +28,6 @@
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include "cutlass/cutlass.h"
@@ -989,12 +987,59 @@ struct CollectiveMma<
uint32_t skip_wait = k_tile_count <= 0;
auto barrier_token = mainloop_pipeline.consumer_try_wait(mainloop_pipe_consumer_state, skip_wait);
bool is_first_iter = true;
//
// PIPELINED MAIN LOOP
//
tiled_mma.accumulate_ = UMMA::ScaleOut::Zero;
if constexpr (IsOverlappingAccum) {
// first iteration manual unroll for tmem overlap kernel
if (k_tile_count > 0) {
// WAIT on mainloop_pipe_consumer_state until its data are available
// (phase bit flips from mainloop_pipe_consumer_state.phase() value)
mainloop_pipeline.consumer_wait(mainloop_pipe_consumer_state, barrier_token);
// Compute on k_tile
int read_stage = mainloop_pipe_consumer_state.index();
// Save current mainlop pipeline read state
auto curr_mainloop_pipe_consumer_state = mainloop_pipe_consumer_state;
// Advance mainloop_pipe
++mainloop_pipe_consumer_state;
--k_tile_count;
skip_wait = k_tile_count <= 0;
// Peek at next iteration
barrier_token = mainloop_pipeline.consumer_try_wait(mainloop_pipe_consumer_state, skip_wait);
if (cute::elect_one_sync()) {
copy(tiled_copy_s2t_SFA, thr_tCsSFA_s2t(_,_,_,_,read_stage), thr_tCtSFA_s2t);
copy(tiled_copy_s2t_SFB, thr_tCsSFB_s2t(_,_,_,_,read_stage), thr_tCtSFB_s2t);
}
// Wait for tmem accumulator buffer to become empty with a flipped phase
accumulator_pipeline.producer_acquire(accumulator_pipe_producer_state);
// Unroll the K mode manually so we can set scale C to 1
CUTLASS_PRAGMA_UNROLL
for (int k_block = 0; k_block < size<2>(tCrA); ++k_block) {
// (V,M) x (V,N) => (V,M,N)
cute::gemm(tiled_mma.with(tiled_mma.accumulate_,
tCtSFA(_,_,k_block),
tCtSFB_mma(_,_,k_block)),
tCrA(_,_,k_block,read_stage),
tCrB(_,_,k_block,read_stage),
accumulators);
tiled_mma.accumulate_ = UMMA::ScaleOut::One;
}
mainloop_pipeline.consumer_release(curr_mainloop_pipe_consumer_state);
}
}
else {
// Wait for tmem accumulator buffer to become empty with a flipped phase
accumulator_pipeline.producer_acquire(accumulator_pipe_producer_state);
}
CUTLASS_PRAGMA_NO_UNROLL
while (k_tile_count > 0) {
// WAIT on mainloop_pipe_consumer_state until its data are available
@@ -1018,12 +1063,6 @@ struct CollectiveMma<
copy(tiled_copy_s2t_SFB, thr_tCsSFB_s2t(_,_,_,_,read_stage), thr_tCtSFB_s2t);
}
// Wait for tmem accumulator buffer to become empty with a flipped phase
if (is_first_iter) {
accumulator_pipeline.producer_acquire(accumulator_pipe_producer_state);
is_first_iter = false;
}
// Unroll the K mode manually so we can set scale C to 1
CUTLASS_PRAGMA_UNROLL
for (int k_block = 0; k_block < size<2>(tCrA); ++k_block) {
@@ -1036,6 +1075,7 @@ struct CollectiveMma<
accumulators);
tiled_mma.accumulate_ = UMMA::ScaleOut::One;
}
mainloop_pipeline.consumer_release(curr_mainloop_pipe_consumer_state);
}
@@ -1197,12 +1197,61 @@ struct CollectiveMma<
uint32_t skip_wait = k_tile_count <= 0;
auto barrier_token = mainloop_pipeline.consumer_try_wait(mainloop_pipe_consumer_state, skip_wait);
bool is_first_iter = true;
//
// PIPELINED MAIN LOOP
//
tiled_mma.accumulate_ = UMMA::ScaleOut::Zero;
if constexpr (IsOverlappingAccum) {
// first iteration manual unroll for tmem overlap kernel
if (k_tile_count > 0) {
// WAIT on mainloop_pipe_consumer_state until its data are available
// (phase bit flips from mainloop_pipe_consumer_state.phase() value)
mainloop_pipeline.consumer_wait(mainloop_pipe_consumer_state, barrier_token);
// Compute on k_tile
int read_stage = mainloop_pipe_consumer_state.index();
// Save current mainlop pipeline read state
auto curr_mainloop_pipe_consumer_state = mainloop_pipe_consumer_state;
// Advance mainloop_pipe
++mainloop_pipe_consumer_state;
--k_tile_count;
skip_wait = k_tile_count <= 0;
// Peek at next iteration
barrier_token = mainloop_pipeline.consumer_try_wait(mainloop_pipe_consumer_state, skip_wait);
if (cute::elect_one_sync()) {
copy(tiled_copy_s2t_E, thr_tCsE_s2t(_,_,_,_,read_stage), thr_tCtE_s2t);
copy(tiled_copy_s2t_SFA, thr_tCsSFA_s2t(_,_,_,_,read_stage), thr_tCtSFA_s2t);
copy(tiled_copy_s2t_SFB, thr_tCsSFB_s2t(_,_,_,_,read_stage), thr_tCtSFB_s2t);
}
// Wait for tmem accumulator buffer to become empty with a flipped phase
accumulator_pipeline.producer_acquire(accumulator_pipe_producer_state);
// Unroll the K mode manually so we can set scale C to 1
CUTLASS_PRAGMA_UNROLL
for (int k_block = 0; k_block < size<2>(tCrA); ++k_block) {
// (V,M) x (V,N) => (V,M,N)
cute::gemm(tiled_mma.with(tiled_mma.accumulate_,
tCtE(_,_,k_block),
tCtSFA(_,_,k_block),
tCtSFB_mma(_,_,k_block)),
tCrA(_,_,k_block,read_stage),
tCrB(_,_,k_block,read_stage),
accumulators);
tiled_mma.accumulate_ = UMMA::ScaleOut::One;
}
mainloop_pipeline.consumer_release(curr_mainloop_pipe_consumer_state);
}
}
else {
// Wait for tmem accumulator buffer to become empty with a flipped phase
accumulator_pipeline.producer_acquire(accumulator_pipe_producer_state);
}
CUTLASS_PRAGMA_NO_UNROLL
while (k_tile_count > 0) {
// WAIT on mainloop_pipe_consumer_state until its data are available
@@ -1227,12 +1276,6 @@ struct CollectiveMma<
copy(tiled_copy_s2t_SFB, thr_tCsSFB_s2t(_,_,_,_,read_stage), thr_tCtSFB_s2t);
}
// Wait for tmem accumulator buffer to become empty with a flipped phase
if (is_first_iter) {
accumulator_pipeline.producer_acquire(accumulator_pipe_producer_state);
is_first_iter = false;
}
// Unroll the K mode manually so we can set scale C to 1
CUTLASS_PRAGMA_UNROLL
for (int k_block = 0; k_block < size<2>(tCrA); ++k_block) {
@@ -667,12 +667,14 @@ struct CollectiveMma<
uint32_t skip_wait = k_tile_count <= 0;
auto barrier_token = mainloop_pipeline.consumer_try_wait(mainloop_pipe_consumer_state, skip_wait);
bool is_first_iter = true;
//
// PIPELINED MAIN LOOP
//
tiled_mma.accumulate_ = UMMA::ScaleOut::Zero;
// Wait for tmem accumulator buffer to become empty with a flipped phase
accumulator_pipeline.producer_acquire(accumulator_pipe_producer_state);
CUTLASS_PRAGMA_NO_UNROLL
while (k_tile_count > 0) {
// WAIT on mainloop_pipe_consumer_state until its data are available
@@ -690,11 +692,6 @@ struct CollectiveMma<
skip_wait = k_tile_count <= 0;
// Peek at next iteration
barrier_token = mainloop_pipeline.consumer_try_wait(mainloop_pipe_consumer_state, skip_wait);
// Wait for tmem accumulator buffer to become empty with a flipped phase
if (is_first_iter) {
accumulator_pipeline.producer_acquire(accumulator_pipe_producer_state);
is_first_iter = false;
}
// Unroll the K mode manually so we can set scale C to 1
CUTLASS_PRAGMA_UNROLL
@@ -70,11 +70,11 @@ template <
class ElementB_,
class StridePairB_,
class TiledMma_,
class GmemTiledCopyA_,
class GmemTiledCopyPairA_,
class SmemLayoutAtomA_,
class SmemCopyAtomA_,
class TransformA_,
class GmemTiledCopyB_,
class GmemTiledCopyPairB_,
class SmemLayoutAtomB_,
class SmemCopyAtomB_,
class TransformB_>
@@ -90,11 +90,11 @@ struct CollectiveMma<
ElementB_,
StridePairB_,
TiledMma_,
GmemTiledCopyA_,
GmemTiledCopyPairA_,
SmemLayoutAtomA_,
SmemCopyAtomA_,
TransformA_,
GmemTiledCopyB_,
GmemTiledCopyPairB_,
SmemLayoutAtomB_,
SmemCopyAtomB_,
TransformB_>
@@ -142,9 +142,6 @@ struct CollectiveMma<
static constexpr int K_BLOCK_MMAS_PER_SCALE_K = ScaleGranularityK / size<2>(typename TiledMma::AtomShape_MNK{});
static constexpr int TILE_M = size<0>(TileShape{});
static constexpr int TILE_N = size<1>(TileShape{});
using ScaleConfig = cutlass::detail::Sm100BlockwiseScaleConfig<ScaleGranularityM,
ScaleGranularityN,
ScaleGranularityK,
@@ -156,8 +153,6 @@ struct CollectiveMma<
using CtaShape_MNK = decltype(shape_div(TileShape{}, AtomThrShapeMNK{}));
static_assert(size<0>(AtomThrShapeMNK{}) == 1, "2SM MMA is not yet supported");
static_assert(size<0>(CtaShape_MNK{}) >= ScaleGranularityM, "Scale Granularity must be smaller than or equal to the tile shape");
static_assert(size<1>(CtaShape_MNK{}) >= ScaleGranularityN, "Scale Granularity must be smaller than or equal to the tile shape");
static_assert(size<2>(CtaShape_MNK{}) >= ScaleGranularityK, "Scale Granularity must be smaller than or equal to the tile shape");
@@ -180,8 +175,10 @@ struct CollectiveMma<
static constexpr bool IsRuntimeDataType = IsRuntimeDataTypeA && IsRuntimeDataTypeB;
using ElementAccumulator = typename TiledMma::ValTypeC;
using GmemTiledCopyA = GmemTiledCopyA_;
using GmemTiledCopyB = GmemTiledCopyB_;
using GmemTiledCopyA = cute::remove_cvref_t<decltype(get<0>(GmemTiledCopyPairA_{}))>;
using GmemTiledCopySFA = cute::remove_cvref_t<decltype(get<1>(GmemTiledCopyPairA_{}))>;
using GmemTiledCopyB = cute::remove_cvref_t<decltype(get<0>(GmemTiledCopyPairB_{}))>;
using GmemTiledCopySFB = cute::remove_cvref_t<decltype(get<1>(GmemTiledCopyPairB_{}))>;
using SmemLayoutAtomA = SmemLayoutAtomA_;
using SmemLayoutAtomB = SmemLayoutAtomB_;
using SmemCopyAtomA = SmemCopyAtomA_;
@@ -190,22 +187,22 @@ struct CollectiveMma<
using TransformB = TransformB_;
using ArchTag = typename DispatchPolicy::ArchTag;
using MainloopPipeline = cutlass::PipelineTmaUmmaAsync<
DispatchPolicy::Stages,
ClusterShape,
AtomThrShapeMNK>;
using MainloopPipelineState = typename MainloopPipeline::PipelineState;
using MainloopABPipeline = cutlass::PipelineTmaUmmaAsync<
DispatchPolicy::Stages,
ClusterShape,
AtomThrShapeMNK>;
using MainloopABPipelineState = typename MainloopABPipeline::PipelineState;
using Load2TransformPipeline = cutlass::PipelineAsync<DispatchPolicy::Stages>;
using Load2TransformPipelineState = typename Load2TransformPipeline::PipelineState;
using MainloopSFPipeline = cutlass::PipelineAsync<DispatchPolicy::Stages>;
using MainloopSFPipelineState = typename MainloopSFPipeline::PipelineState;
using Mma2TransformPipeline = cutlass::PipelineUmmaAsync<
using AccumulatorPipeline = cutlass::PipelineUmmaAsync<
AccumulatorPipelineStageCount,
AtomThrShapeMNK>;
using Mma2TransformPipelineState = typename Mma2TransformPipeline::PipelineState;
using AccumulatorPipelineState = typename AccumulatorPipeline::PipelineState;
// Two arrivals per CTA (1 arrival and 1 arrival through cp.async.mbarrier)
static constexpr int NumLoad2TransformProducerThreadEvents = 2;
// Two arrivals per thread in the warp (1 arrival and 1 arrival through cp.async.mbarrier)
static constexpr int NumMainloopSFProducerThreadEvents = 64;
static_assert(rank(SmemLayoutAtomA{}) == 2, "SmemLayoutAtomA must be rank 2 (M,K)");
static_assert(((size<0,0>(MmaShapeA_MK{}) * size<1>(MmaShapeA_MK{})) % size<0>(SmemLayoutAtomA{})) == 0,
@@ -277,43 +274,28 @@ struct CollectiveMma<
append(stride(SmemLayoutAtomSFB{}), size(filter_zeros(SmemLayoutAtomSFB{})))
));
// Scaling gmem-to-smem copy atom
static constexpr int LeadingScalesPerTileSFA = size<0,1>(LayoutSFA{}.stride()) == 1 ? ScaleMsPerTile : ScaleKsPerTile;
using ScaleCopyTypeA = cute::uint_byte_t<cute::min(static_cast<int>(sizeof(ElementAccumulator)) * LeadingScalesPerTileSFA, 16)>;
using SmemScalingCopyAtomA = Copy_Atom<SM80_CP_ASYNC_CACHEALWAYS<ScaleCopyTypeA>, ElementAccumulator>;
static constexpr int ElementsPerSFACopy = static_cast<int>(sizeof(ScaleCopyTypeA) / sizeof(ElementAccumulator));
static constexpr int LeadingScalesPerTileSFB = size<0,1>(LayoutSFB{}.stride()) == 1 ? ScaleNsPerTile : ScaleKsPerTile;
using ScaleCopyTypeB = cute::uint_byte_t<cute::min(static_cast<int>(sizeof(ElementAccumulator)) * LeadingScalesPerTileSFB, 16)>;
using SmemScalingCopyAtomB = Copy_Atom<SM80_CP_ASYNC_CACHEALWAYS<ScaleCopyTypeB>, ElementAccumulator>;
static constexpr int ElementsPerSFBCopy = static_cast<int>(sizeof(ScaleCopyTypeB) / sizeof(ElementAccumulator));
using TiledCopyScaleA = decltype(make_tiled_copy(SmemScalingCopyAtomA{}, Layout<Shape<_1>>{}, Layout<Shape<Int<ElementsPerSFACopy>>>{}));
using TiledCopyScaleB = decltype(make_tiled_copy(SmemScalingCopyAtomB{}, Layout<Shape<_1>>{}, Layout<Shape<Int<ElementsPerSFBCopy>>>{}));
struct SharedStorage {
struct TensorStorage : cute::aligned_struct<128, _0> {
cute::ArrayEngine<SmemAllocTypeA, cute::cosize_v<SmemLayoutA>> smem_A;
cute::ArrayEngine<SmemAllocTypeB, cute::cosize_v<SmemLayoutB>> smem_B;
cute::ArrayEngine<ElementAccumulator, cute::cosize_v<SmemLayoutScaleA>> smem_scale_A;
cute::ArrayEngine<ElementAccumulator, cute::cosize_v<SmemLayoutScaleB>> smem_scale_B;
cute::ArrayEngine<ElementAccumulator, cute::cosize_v<SmemLayoutScaleA>> smem_SFA;
cute::ArrayEngine<ElementAccumulator, cute::cosize_v<SmemLayoutScaleB>> smem_SFB;
} tensors;
using PipelineStorage = typename MainloopPipeline::SharedStorage;
PipelineStorage pipeline;
using PipelineABStorage = typename MainloopABPipeline::SharedStorage;
using PipelineSFStorage = typename MainloopSFPipeline::SharedStorage;
using AccumulatorPipelineStorage = typename AccumulatorPipeline::SharedStorage;
using Load2TransformPipelineStorage = typename Load2TransformPipeline::SharedStorage;
Load2TransformPipelineStorage transform2load_pipeline;
using Mma2TransformPipelineStorage = typename Mma2TransformPipeline::SharedStorage;
Mma2TransformPipelineStorage mma2transform_pipeline;
struct PipelineStorage {
alignas(16) PipelineABStorage pipeline_ab;
alignas(16) PipelineSFStorage pipeline_sf;
alignas(16) AccumulatorPipelineStorage pipeline_accum;
};
};
// Expose shared storage for tensors/pipelines separately to allow kernel layer to reorder them.
using TensorStorage = typename SharedStorage::TensorStorage;
using PipelineStorage = typename SharedStorage::PipelineStorage;
using Mma2TransformPipelineStorage = typename SharedStorage::Mma2TransformPipelineStorage;
using Load2TransformPipelineStorage = typename SharedStorage::Load2TransformPipelineStorage;
// Only one thread issues the TMA and updates the barriers in a 2SM MMA, adjust bytes accordingly
static constexpr uint32_t TmaTransactionBytes =
@@ -328,12 +310,9 @@ struct CollectiveMma<
template<
class KTileCount,
class GTensorPartitionedA, class GTensorPartitionedB,
class STensorA, class STensorB,
class GTensorPartitionedScaleA, class GTensorPartitionedScaleB,
class IdentTensorPartitionedScaleA, class IdentTensorPartitionedScaleB,
class STensorScaleA, class STensorScaleB
class STensorA, class STensorB
>
struct LoadParams {
struct LoadABParams {
// for scheduler
KTileCount k_tiles;
// for input tensor values
@@ -342,6 +321,32 @@ struct CollectiveMma<
STensorA tAsA;
STensorB tBsB;
// the TMA multicast masks
uint16_t mcast_mask_a;
uint16_t mcast_mask_b;
CUTLASS_DEVICE
LoadABParams (
KTileCount k_tiles_,
GTensorPartitionedA tAgA_mkl_, GTensorPartitionedB tBgB_nkl_,
STensorA tAsA_, STensorB tBsB_,
uint16_t mcast_mask_a_, uint16_t mcast_mask_b_)
: k_tiles(k_tiles_)
, tAgA_mkl(tAgA_mkl_), tBgB_nkl(tBgB_nkl_)
, tAsA(tAsA_), tBsB(tBsB_)
, mcast_mask_a(mcast_mask_a_), mcast_mask_b(mcast_mask_b_) {}
};
template<
class KTileCount,
class GTensorPartitionedScaleA, class GTensorPartitionedScaleB,
class IdentTensorPartitionedScaleA, class IdentTensorPartitionedScaleB,
class STensorScaleA, class STensorScaleB
>
struct LoadSFParams {
// for scheduler
KTileCount k_tiles;
GTensorPartitionedScaleA tSFAgSFA_mkl;
GTensorPartitionedScaleB tSFBgSFB_nkl;
IdentTensorPartitionedScaleA tSFAIdentSFA_mkl;
@@ -349,30 +354,20 @@ struct CollectiveMma<
STensorScaleA tSFAsSFA;
STensorScaleB tSFBsSFB;
// the TMA multicast masks
uint16_t mcast_mask_a;
uint16_t mcast_mask_b;
LayoutSFA layout_SFA;
LayoutSFB layout_SFB;
CUTLASS_DEVICE
LoadParams (
LoadSFParams (
KTileCount k_tiles_,
GTensorPartitionedA tAgA_mkl_, GTensorPartitionedB tBgB_nkl_,
STensorA tAsA_, STensorB tBsB_,
GTensorPartitionedScaleA tSFAgSFA_mkl_, GTensorPartitionedScaleB tSFBgSFB_nkl_,
IdentTensorPartitionedScaleA tSFAIdentSFA_mkl_, IdentTensorPartitionedScaleB tSFBIdentSFB_nkl_,
STensorScaleA tSFAsSFA_, STensorScaleB tSFBsSFB_,
uint16_t mcast_mask_a_, uint16_t mcast_mask_b_,
LayoutSFA layout_SFA_, LayoutSFB layout_SFB_)
: k_tiles(k_tiles_)
, tAgA_mkl(tAgA_mkl_), tBgB_nkl(tBgB_nkl_)
, tAsA(tAsA_), tBsB(tBsB_)
, tSFAgSFA_mkl(tSFAgSFA_mkl_), tSFBgSFB_nkl(tSFBgSFB_nkl_)
, tSFAIdentSFA_mkl(tSFAIdentSFA_mkl_), tSFBIdentSFB_nkl(tSFBIdentSFB_nkl_)
, tSFAsSFA(tSFAsSFA_), tSFBsSFB(tSFBsSFB_)
, mcast_mask_a(mcast_mask_a_), mcast_mask_b(mcast_mask_b_)
, layout_SFA(layout_SFA_), layout_SFB(layout_SFB_) {}
};
@@ -393,14 +388,14 @@ struct CollectiveMma<
template<
class STensorScaleA, class STensorScaleB
>
struct TransformParams {
struct AccumTransformParams {
// for scheduler
STensorScaleA sSFA;
STensorScaleB sSFB;
CUTLASS_DEVICE
TransformParams (
AccumTransformParams (
STensorScaleA sSFA_, STensorScaleB sSFB_)
: sSFA(sSFA_), sSFB(sSFB_) {}
};
@@ -412,9 +407,9 @@ struct CollectiveMma<
StrideA dA{};
ArrayElementB const* ptr_B{nullptr};
StrideB dB{};
ElementAccumulator const* ptr_scale_A{nullptr};
ElementAccumulator const* ptr_SFA{nullptr};
LayoutSFA layout_SFA{};
ElementAccumulator const* ptr_scale_B{nullptr};
ElementAccumulator const* ptr_SFB{nullptr};
LayoutSFB layout_SFB{};
RuntimeDataTypeA runtime_data_type_a{};
RuntimeDataTypeB runtime_data_type_b{};
@@ -451,9 +446,9 @@ struct CollectiveMma<
RuntimeDataTypeA runtime_data_type_a;
RuntimeDataTypeB runtime_data_type_b;
ElementAccumulator const* ptr_scale_A;
ElementAccumulator const* ptr_SFA;
LayoutSFA layout_SFA;
ElementAccumulator const* ptr_scale_B;
ElementAccumulator const* ptr_SFB;
LayoutSFB layout_SFB;
};
@@ -539,9 +534,9 @@ struct CollectiveMma<
hw_info.cluster_shape_fallback,
args.runtime_data_type_a,
args.runtime_data_type_b,
args.ptr_scale_A,
args.ptr_SFA,
args.layout_SFA,
args.ptr_scale_B,
args.ptr_SFB,
args.layout_SFB
};
}
@@ -568,8 +563,8 @@ struct CollectiveMma<
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Problem Size doesn't meet the minimum alignment requirements for TMA.\n");
}
bool implementable_sf = cutlass::detail::check_alignment<sizeof(ScaleCopyTypeA) / sizeof(ElementAccumulator)>(args.layout_SFA);
implementable_sf = implementable_sf && cutlass::detail::check_alignment<sizeof(ScaleCopyTypeB) / sizeof(ElementAccumulator)>(args.layout_SFB);
bool implementable_sf = cutlass::detail::check_alignment<sizeof(typename GmemTiledCopySFA::ValType) / sizeof(ElementAccumulator)>(args.layout_SFA);
implementable_sf = implementable_sf && cutlass::detail::check_alignment<sizeof(typename GmemTiledCopySFB::ValType) / sizeof(ElementAccumulator)>(args.layout_SFB);
if (!implementable_sf) {
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Problem Size doesn't meet the minimum alignment requirements for Scale Factors.\n");
@@ -628,20 +623,12 @@ struct CollectiveMma<
/// gB_nkl - The tiled tma tensor for input B
/// tAsA - partitioned smem tensor for A
/// tBsB - partitioned smem tensor for B
/// tSFAgSFA_mkl - partitioned gmem tensor for SFA
/// tSFBgSFB_nkl - partitioned gmem tensor for SFB
/// tSFAIdentSFA_mkl - partitioned identity tensor for SFA in gmem
/// tSFBIdentSFB_nkl - partitioned identity tensor for SFB in gmem
/// tSFAsSFA - partitioned smem tensor for SFA
/// tSFBsSFB - partitioned smem tensor for SFB
/// mcast_mask_a - tma multicast mask for A
/// mcast_mask_b - tma multicast mask for B
/// layout_SFA - layout of SFA in gmem
/// layout_SFB - layout of SFB in gmem
template <class ProblemShape_MNKL,
class MainloopParams>
CUTLASS_DEVICE auto
load_init(
load_ab_init(
ProblemShape_MNKL const& problem_shape_MNKL,
MainloopParams const& mainloop_params,
TensorStorage& shared_tensors) const {
@@ -686,10 +673,38 @@ struct CollectiveMma<
uint16_t mcast_mask_a = create_tma_multicast_mask<2>(cta_layout_vmnk, cta_coord_vmnk);
uint16_t mcast_mask_b = create_tma_multicast_mask<1>(cta_layout_vmnk, cta_coord_vmnk);
// Scales
LoadABParams load_params {
shape<3>(gA_mkl), // for scheduler
tAgA_mkl, tBgB_nkl, tAsA, tBsB, // for input tensor values
mcast_mask_a, mcast_mask_b, // multicast masks
};
return load_params;
}
Tensor mSFA_mkl = make_tensor(make_gmem_ptr(mainloop_params.ptr_scale_A), mainloop_params.layout_SFA); // (m,k,l)
Tensor mSFB_nkl = make_tensor(make_gmem_ptr(mainloop_params.ptr_scale_B), mainloop_params.layout_SFB); // (n,k,l)
/// Set up the data needed by this collective for load.
/// Return load params containing
/// tSFAgSFA_mkl - partitioned gmem tensor for SFA
/// tSFBgSFB_nkl - partitioned gmem tensor for SFB
/// tSFAIdentSFA_mkl - partitioned identity tensor for SFA in gmem
/// tSFBIdentSFB_nkl - partitioned identity tensor for SFB in gmem
/// tSFAsSFA - partitioned smem tensor for SFA
/// tSFBsSFB - partitioned smem tensor for SFB
/// layout_SFA - layout of SFA in gmem
/// layout_SFB - layout of SFB in gmem
template <class ProblemShape_MNKL,
class MainloopParams>
CUTLASS_DEVICE auto
load_sf_init(
ProblemShape_MNKL const& problem_shape_MNKL,
MainloopParams const& mainloop_params,
TensorStorage& shared_tensors) const {
using X = Underscore;
// Separate out problem shape for convenience
auto [M,N,K,L] = problem_shape_MNKL;
Tensor mSFA_mkl = make_tensor(make_gmem_ptr(mainloop_params.ptr_SFA), mainloop_params.layout_SFA); // (m,k,l)
Tensor mSFB_nkl = make_tensor(make_gmem_ptr(mainloop_params.ptr_SFB), mainloop_params.layout_SFB); // (n,k,l)
Tensor SFA_mkl_ident = make_identity_tensor(shape(mainloop_params.layout_SFA));
@@ -710,15 +725,15 @@ struct CollectiveMma<
static_assert(rank(decltype(gSFB_nkl){}) == 5);
// 1 thread copies entire set of scalar
TiledCopyScaleA scale_copy_a{};
TiledCopyScaleB scale_copy_b{};
GmemTiledCopySFA scale_copy_a{};
GmemTiledCopySFB scale_copy_b{};
ThrCopy thr_scale_copy_a = scale_copy_a.get_slice(_0{});
ThrCopy thr_scale_copy_b = scale_copy_b.get_slice(_0{});
ThrCopy thr_scale_copy_a = scale_copy_a.get_slice(threadIdx.x % size(scale_copy_a));
ThrCopy thr_scale_copy_b = scale_copy_b.get_slice(threadIdx.x % size(scale_copy_b));
Tensor sSFA = make_tensor(make_smem_ptr(shared_tensors.smem_scale_A.begin()),
Tensor sSFA = make_tensor(make_smem_ptr(shared_tensors.smem_SFA.begin()),
SmemLayoutScaleA{}); // (CTA_M,CTA_K,P)
Tensor sSFB = make_tensor(make_smem_ptr(shared_tensors.smem_scale_B.begin()),
Tensor sSFB = make_tensor(make_smem_ptr(shared_tensors.smem_SFB.begin()),
SmemLayoutScaleB{}); // (CTA_M,CTA_K,P)
Tensor tSFAgSFA_mkl = thr_scale_copy_a.partition_S(gSFA_mkl); // (CPY, BLK_M, BLK_K, m, k, l)
@@ -733,19 +748,18 @@ struct CollectiveMma<
static_assert(rank(decltype(tSFAgSFA_mkl){}) == 6);
static_assert(rank(decltype(tSFBgSFB_nkl){}) == 6);
LoadParams load_params {
shape<3>(gA_mkl), // for scheduler
tAgA_mkl, tBgB_nkl, tAsA, tBsB, // for input tensor values
LoadSFParams load_params {
size<3>(gSFA_mkl),
tSFAgSFA_mkl, tSFBgSFB_nkl, // for input scale tensor values
tSFAIdentSFA_mkl, tSFBIdentSFB_nkl, // for predicating scale tensor copies
tSFAsSFA, tSFBsSFB, // for scale tensor values
mcast_mask_a, mcast_mask_b, // multicast masks
mainloop_params.layout_SFA, // for predicating scale tensor copies
mainloop_params.layout_SFB // for predicating scale tensor copies
};
return load_params;
}
/// Set up the data needed by this collective for mma compute.
template <class AccTensor>
CUTLASS_DEVICE auto
@@ -756,8 +770,27 @@ struct CollectiveMma<
Tensor sB = make_tensor(make_smem_ptr(shared_tensors.smem_B.begin()), SmemLayoutB{}); // (BLK_N,BLK_K,PIPE)
// Allocate "fragments/descriptors" for A and B matrices
Tensor tCrA = TiledMma::make_fragment_A(sA); // (MMA,MMA_M,MMA_K,PIPE)
Tensor tCrB = TiledMma::make_fragment_B(sB); // (MMA,MMA_N,MMA_K,PIPE)
Tensor tCrA_ = TiledMma::make_fragment_A(sA); // (MMA,MMA_M,MMA_K,PIPE)
Tensor tCrB_ = TiledMma::make_fragment_B(sB); // (MMA,MMA_N,MMA_K,PIPE)
CUTE_STATIC_ASSERT_V(rank(tCrA_) == _4{});
auto mma_tile_shape_A = make_shape(get<0>(shape(tCrA_.layout())),
get<1>(shape(tCrA_.layout())),
Int<K_BLOCK_MMAS_PER_SCALE_K>{},
_1{});
auto mma_tile_shape_B = make_shape(get<0>(shape(tCrB_.layout())),
get<1>(shape(tCrB_.layout())),
Int<K_BLOCK_MMAS_PER_SCALE_K>{},
_1{});
Tensor tCrA = flat_divide(tCrA_,
mma_tile_shape_A)(_,_,_,_0{},_0{},_0{},_,_); // (MMA,MMA_M,MMA_K_PER_SCALE,MMA_K_REST,PIPE)
Tensor tCrB = flat_divide(tCrB_,
mma_tile_shape_B)(_,_,_,_0{},_0{},_0{},_,_); // (MMA,MMA_N,MMA_K_PER_SCALE,MMA_K_REST,PIPE)
CUTE_STATIC_ASSERT_V(Int<DispatchPolicy::Stages>{} == size<3>(sA)); // PIPE
CUTE_STATIC_ASSERT_V(Int<DispatchPolicy::Stages>{} == size<3>(sB));
@@ -780,7 +813,7 @@ struct CollectiveMma<
/// Set up the data needed by this collective for transform.
template <class ProblemShape_MNKL>
CUTLASS_DEVICE auto
transform_init(
accum_init(
ProblemShape_MNKL const& problem_shape_MNKL,
TensorStorage& shared_tensors) const {
using X = Underscore;
@@ -788,13 +821,13 @@ struct CollectiveMma<
// Separate out problem shape for convenience
auto [M,N,K,L] = problem_shape_MNKL;
Tensor sSFA = make_tensor(cute::make_smem_ptr(shared_tensors.smem_scale_A.begin()),
Tensor sSFA = make_tensor(cute::make_smem_ptr(shared_tensors.smem_SFA.begin()),
SmemLayoutScaleA{}); // (ScaleMsPerTile,ScakeKsPerTile,P)
Tensor sSFB = make_tensor(cute::make_smem_ptr(shared_tensors.smem_scale_B.begin()),
Tensor sSFB = make_tensor(cute::make_smem_ptr(shared_tensors.smem_SFB.begin()),
SmemLayoutScaleB{}); // (ScaleNsPerTile,ScaleKsPerTile,P)
TransformParams transform_params {
AccumTransformParams transform_params {
sSFA, sSFB // for input tensor values
};
return transform_params;
@@ -803,34 +836,92 @@ struct CollectiveMma<
/// Perform a collective-scoped matrix multiply-accumulate
/// Producer Perspective
template <
class LoadParams,
class LoadABParams,
class TileCoordMNKL,
class KTileIterator
>
CUTLASS_DEVICE auto
load(
MainloopPipeline mainloop_pipeline,
Load2TransformPipeline load2transform_pipeline,
MainloopPipelineState mainloop_pipe_producer_state,
Load2TransformPipelineState load2transform_pipe_producer_state,
LoadParams const& load_inputs,
load_ab(
MainloopABPipeline mainloop_pipeline,
MainloopABPipelineState mainloop_pipe_producer_state,
LoadABParams const& load_inputs,
TileCoordMNKL const& cta_coord_mnkl,
KTileIterator k_tile_iter, int k_tile_count) {
auto [unused_k_tiles,
tAgA_mkl, tBgB_nkl, tAsA, tBsB,
tSFAgSFA_mkl, tSFBgSFB_nkl,
tSFAIdentSFA_mkl, tSFBIdentSFB_nkl,
tSFAsSFA, tSFBsSFB,
mcast_mask_a, mcast_mask_b,
layout_SFA, layout_SFB] = load_inputs;
mcast_mask_a, mcast_mask_b] = load_inputs;
// slice out the work coord from partitioned tensors
Tensor tAgA = tAgA_mkl(_, get<0>(cta_coord_mnkl) / size(typename TiledMma::AtomThrID{}), _, get<3>(cta_coord_mnkl));
Tensor tBgB = tBgB_nkl(_, get<1>(cta_coord_mnkl), _, get<3>(cta_coord_mnkl));
TiledCopyScaleA scale_copy_a{};
TiledCopyScaleB scale_copy_b{};
auto barrier_token = mainloop_pipeline.producer_try_acquire(mainloop_pipe_producer_state);
// Issue the Mainloop loads
CUTLASS_PRAGMA_NO_UNROLL
while (k_tile_count > 0) {
// LOCK mainloop_pipe_producer_state for _writing_
mainloop_pipeline.producer_acquire(mainloop_pipe_producer_state, barrier_token);
using BarrierType = typename MainloopABPipeline::ProducerBarrierType;
BarrierType* tma_barrier = mainloop_pipeline.producer_get_barrier(mainloop_pipe_producer_state);
int write_stage = mainloop_pipe_producer_state.index();
auto curr_mainloop_pipe_producer_state = mainloop_pipe_producer_state;
++mainloop_pipe_producer_state;
barrier_token = mainloop_pipeline.producer_try_acquire(mainloop_pipe_producer_state);
if (cute::elect_one_sync()) {
copy(observed_tma_load_a_->with(*tma_barrier, mcast_mask_a), tAgA(_,*k_tile_iter), tAsA(_,write_stage));
copy(observed_tma_load_b_->with(*tma_barrier, mcast_mask_b), tBgB(_,*k_tile_iter), tBsB(_,write_stage));
}
--k_tile_count;
++k_tile_iter;
}
return cute::make_tuple(mainloop_pipe_producer_state, k_tile_iter);
}
/// Perform a Producer Epilogue to prevent early exit of ctas in a Cluster
CUTLASS_DEVICE void
load_ab_tail(
MainloopABPipeline mainloop_pipeline,
MainloopABPipelineState mainloop_pipe_producer_state) {
// Issue the epilogue waits
// This helps avoid early exit of ctas in Cluster
// Waits for all stages to either be released (all
// Consumer UNLOCKs), or if the stage was never used
// then would just be acquired since the phase was
// still inverted from make_producer_start_state
mainloop_pipeline.producer_tail(mainloop_pipe_producer_state);
}
/// Perform a collective-scoped transform
/// Load producer Perspective
template <
class LoadSFParams,
class TileCoordMNKL,
class KTileIterator
>
CUTLASS_DEVICE auto
load_sf(
MainloopSFPipeline mainloop_sf_pipeline,
MainloopSFPipelineState mainloop_sf_pipe_producer_state,
LoadSFParams const& load_inputs,
TileCoordMNKL const& cta_coord_mnkl,
KTileIterator k_tile_iter, int k_tile_count) {
auto [unused_k_tiles,
tSFAgSFA_mkl, tSFBgSFB_nkl,
tSFAIdentSFA_mkl, tSFBIdentSFB_nkl,
tSFAsSFA, tSFBsSFB,
layout_SFA, layout_SFB] = load_inputs;
// slice out the work coord from partitioned tensors
GmemTiledCopySFA scale_copy_a{};
GmemTiledCopySFB scale_copy_b{};
Tensor tSFAgSFA = tSFAgSFA_mkl(_, _, _, get<0>(cta_coord_mnkl), _, get<3>(cta_coord_mnkl));
@@ -842,69 +933,50 @@ struct CollectiveMma<
Tensor thr_tile_pSFB = make_tensor<bool>(shape(filter_zeros(thr_tile_SFB_k(_,_,_0{}), tSFBgSFB(_0{},_,_,_0{}).stride())));
auto barrier_token = mainloop_pipeline.producer_try_acquire(mainloop_pipe_producer_state);
// Issue the Mainloop loads
// Issue the loads
CUTLASS_PRAGMA_NO_UNROLL
while (k_tile_count > 0) {
// LOCK mainloop_pipe_producer_state for _writing_
mainloop_pipeline.producer_acquire(mainloop_pipe_producer_state, barrier_token);
load2transform_pipeline.producer_acquire(load2transform_pipe_producer_state);
using BarrierType = typename MainloopPipeline::ProducerBarrierType;
BarrierType* tma_barrier = mainloop_pipeline.producer_get_barrier(mainloop_pipe_producer_state);
int write_stage = mainloop_pipe_producer_state.index();
auto curr_mainloop_pipe_producer_state = mainloop_pipe_producer_state;
++mainloop_pipe_producer_state;
barrier_token = mainloop_pipeline.producer_try_acquire(mainloop_pipe_producer_state);
// LOCK pipe_producer_state for _writing_
mainloop_sf_pipeline.producer_acquire(mainloop_sf_pipe_producer_state);
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < size(thr_tile_pSFA); ++i) {
Tensor thr_tile_SFA = filter_zeros(thr_tile_SFA_k(_,_,*k_tile_iter), tSFAgSFA(_0{},_,_,_0{}).stride());
thr_tile_pSFA(i) = elem_less(thr_tile_SFA(i), shape(filter_zeros(layout_SFA)));
thr_tile_pSFA(i) = elem_less(thr_tile_SFA(i), shape(filter_zeros(layout_SFA))) && threadIdx.x % 32 < size(scale_copy_a);
}
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < size(thr_tile_pSFB); ++i) {
Tensor thr_tile_SFB = filter_zeros(thr_tile_SFB_k(_,_,*k_tile_iter), tSFBgSFB(_0{},_,_,_0{}).stride());
thr_tile_pSFB(i) = elem_less(thr_tile_SFB(i), shape(filter_zeros(layout_SFB)));
thr_tile_pSFB(i) = elem_less(thr_tile_SFB(i), shape(filter_zeros(layout_SFB))) && threadIdx.x % 32 < size(scale_copy_b);
}
if (cute::elect_one_sync()) {
copy(observed_tma_load_a_->with(*tma_barrier, mcast_mask_a), tAgA(_,*k_tile_iter), tAsA(_,write_stage));
copy(observed_tma_load_b_->with(*tma_barrier, mcast_mask_b), tBgB(_,*k_tile_iter), tBsB(_,write_stage));
copy_if(scale_copy_a, thr_tile_pSFA, filter_zeros(tSFAgSFA(_,_,_,*k_tile_iter)), filter_zeros(tSFAsSFA(_,_,_,load2transform_pipe_producer_state.index())));
copy_if(scale_copy_b, thr_tile_pSFB, filter_zeros(tSFBgSFB(_,_,_,*k_tile_iter)), filter_zeros(tSFBsSFB(_,_,_,load2transform_pipe_producer_state.index())));
load2transform_pipeline.producer_commit(load2transform_pipe_producer_state, cutlass::arch::cpasync_barrier_arrive_noinc);
}
copy_if(scale_copy_a, thr_tile_pSFA, filter_zeros(tSFAgSFA(_,_,_,*k_tile_iter)), filter_zeros(tSFAsSFA(_,_,_,mainloop_sf_pipe_producer_state.index())));
copy_if(scale_copy_b, thr_tile_pSFB, filter_zeros(tSFBgSFB(_,_,_,*k_tile_iter)), filter_zeros(tSFBsSFB(_,_,_,mainloop_sf_pipe_producer_state.index())));
mainloop_sf_pipeline.producer_commit(mainloop_sf_pipe_producer_state, cutlass::arch::cpasync_barrier_arrive_noinc);
__syncwarp();
++load2transform_pipe_producer_state;
++mainloop_sf_pipe_producer_state;
--k_tile_count;
++k_tile_iter;
}
return cute::make_tuple(mainloop_pipe_producer_state, load2transform_pipe_producer_state, k_tile_iter);
return cute::make_tuple(mainloop_sf_pipe_producer_state, k_tile_iter);
}
/// Perform a Producer Epilogue to prevent early exit of ctas in a Cluster
CUTLASS_DEVICE void
load_tail(
MainloopPipeline mainloop_pipeline,
Load2TransformPipeline load2transform_pipeline,
MainloopPipelineState mainloop_pipe_producer_state,
Load2TransformPipelineState load2transform_pipe_producer_state) {
load_sf_tail(
MainloopSFPipeline mainloop_sf_pipeline,
MainloopSFPipelineState mainloop_sf_pipe_producer_state) {
// Issue the epilogue waits
// This helps avoid early exit of ctas in Cluster
// Waits for all stages to either be released (all
// Consumer UNLOCKs), or if the stage was never used
// then would just be acquired since the phase was
// still inverted from make_producer_start_state
mainloop_pipeline.producer_tail(mainloop_pipe_producer_state);
load2transform_pipeline.producer_tail(load2transform_pipe_producer_state);
mainloop_sf_pipeline.producer_tail(mainloop_sf_pipe_producer_state);
}
/// Perform a collective-scoped matrix multiply-accumulate
@@ -916,10 +988,10 @@ struct CollectiveMma<
>
CUTLASS_DEVICE auto
mma(
cute::tuple<MainloopPipeline,
Mma2TransformPipeline> pipelines,
cute::tuple<MainloopPipelineState,
Mma2TransformPipelineState> pipeline_states,
cute::tuple<MainloopABPipeline,
AccumulatorPipeline> pipelines,
cute::tuple<MainloopABPipelineState,
AccumulatorPipelineState> pipeline_states,
TmemStorage tmem_storage,
MmaParams const& mma_inputs,
CtaTileCoord cta_tile_coord,
@@ -927,10 +999,10 @@ struct CollectiveMma<
auto [tiled_mma, tCrA, tCrB] = mma_inputs;
auto [mainloop_pipeline,
mma2transform_pipeline] = pipelines;
accumulator_pipeline] = pipelines;
auto [mainloop_pipe_consumer_state,
mma2transform_pipe_producer_state] = pipeline_states;
accumulator_pipe_producer_state] = pipeline_states;
uint32_t skip_wait = k_tile_count <= 0;
auto barrier_token = mainloop_pipeline.consumer_try_wait(mainloop_pipe_consumer_state, skip_wait);
@@ -958,54 +1030,50 @@ struct CollectiveMma<
// Peek at next iteration
barrier_token = mainloop_pipeline.consumer_try_wait(mainloop_pipe_consumer_state, skip_wait);
static_assert(size<2>(tCrA) / K_BLOCK_MMAS_PER_SCALE_K, "k blocks must be divisible by K_BLOCK_MMAS_PER_SCALE_K");
CUTLASS_PRAGMA_UNROLL
for (int scale_k_blocks = 0; scale_k_blocks < size<2>(tCrA) / K_BLOCK_MMAS_PER_SCALE_K; ++scale_k_blocks) {
mma2transform_pipeline.producer_acquire(mma2transform_pipe_producer_state);
for (int scale_k_iter = 0; scale_k_iter < size<3>(tCrA); ++scale_k_iter) {
accumulator_pipeline.producer_acquire(accumulator_pipe_producer_state);
auto acc = get<0>(slice_accumulator(tmem_storage, mma2transform_pipe_producer_state.index()));
auto acc = get<0>(slice_accumulator(tmem_storage, accumulator_pipe_producer_state.index()));
static_assert(is_tmem<remove_cvref_t<decltype(acc)>>::value, "Accumulator must be tmem resident.");
static_assert(rank(remove_cvref_t<decltype(acc)>{}) == 3, "Accumulator must be MMA-partitioned: (MMA, MMA_M, MMA_N)");
// for each set of scale_k_blocks we zero the accumulator
tiled_mma.accumulate_ = UMMA::ScaleOut::Zero;
int start_k_block = scale_k_blocks * size<2>(tCrA) / K_BLOCK_MMAS_PER_SCALE_K;
// Unroll the K mode manually so we can set scale C to 1
CUTLASS_PRAGMA_UNROLL
for (int k_block_offset = 0; k_block_offset < K_BLOCK_MMAS_PER_SCALE_K; ++k_block_offset) {
int k_block = start_k_block + k_block_offset;
for (int k_block = 0; k_block < size<2>(tCrA); ++k_block) {
// (V,M) x (V,N) => (V,M,N)
cute::gemm(tiled_mma,
tCrA(_,_,k_block,read_stage),
tCrB(_,_,k_block,read_stage),
tCrA(_,_,k_block,scale_k_iter,read_stage),
tCrB(_,_,k_block,scale_k_iter,read_stage),
acc);
tiled_mma.accumulate_ = UMMA::ScaleOut::One;
}
mma2transform_pipeline.producer_commit(mma2transform_pipe_producer_state);
++mma2transform_pipe_producer_state;
accumulator_pipeline.producer_commit(accumulator_pipe_producer_state);
++accumulator_pipe_producer_state;
}
mainloop_pipeline.consumer_release(curr_mainloop_pipe_consumer_state);
}
return make_tuple(mainloop_pipe_consumer_state, mma2transform_pipe_producer_state);
return make_tuple(mainloop_pipe_consumer_state, accumulator_pipe_producer_state);
}
/// Transform
template <
class TransformParams,
class AccumTransformParams,
class TmemStorage,
class CtaTileCoord,
class CopyOpT2R,
class EpilogueTile
>
CUTLASS_DEVICE auto
transform(
cute::tuple<Mma2TransformPipeline, Load2TransformPipeline> pipelines,
cute::tuple<Mma2TransformPipelineState, Load2TransformPipelineState> consumer_states,
accum(
cute::tuple<AccumulatorPipeline, MainloopSFPipeline> pipelines,
cute::tuple<AccumulatorPipelineState, MainloopSFPipelineState> consumer_states,
TmemStorage tmem_storage,
TransformParams const& transform_inputs,
AccumTransformParams const& transform_inputs,
CtaTileCoord cta_tile_coord,
CopyOpT2R,
EpilogueTile,
@@ -1076,14 +1144,14 @@ struct CollectiveMma<
// Zero our accumulator
clear(tTR_FullAcc);
auto [mma2transform_pipeline, load2transform_pipeline] = pipelines;
auto [mma2transform_pipe_state, load2transform_pipe_state] = consumer_states;
auto [accumulator_pipeline, mainloop_sf_pipeline] = pipelines;
auto [accumulator_pipe_state, mainloop_sf_pipe_state] = consumer_states;
CUTLASS_PRAGMA_NO_UNROLL
while (k_tile_count > 0) {
load2transform_pipeline.consumer_wait(load2transform_pipe_state);
int read_idx = load2transform_pipe_state.index();
mainloop_sf_pipeline.consumer_wait(mainloop_sf_pipe_state);
int read_idx = mainloop_sf_pipe_state.index();
copy(filter_zeros(tTR_sSFA_epi(_,_,_,_,_,_,read_idx)), tTR_rSFA_compact);
copy(filter_zeros(tTR_sSFB_epi(_,_,_,_,_,_,read_idx)), tTR_rSFB_compact);
@@ -1094,15 +1162,15 @@ struct CollectiveMma<
Tensor tTR_rSFA = make_tensor(tTR_rSFA_compact.data(), tTR_rSFA_layout);
Tensor tTR_rSFB = make_tensor(tTR_rSFB_compact.data(), tTR_rSFB_layout);
load2transform_pipeline.consumer_release(load2transform_pipe_state);
++load2transform_pipe_state;
mainloop_sf_pipeline.consumer_release(mainloop_sf_pipe_state);
++mainloop_sf_pipe_state;
CUTLASS_PRAGMA_UNROLL
for (int k_block = 0; k_block < ScaleKsPerTile; ++k_block) {
mma2transform_pipeline.consumer_wait(mma2transform_pipe_state);
accumulator_pipeline.consumer_wait(accumulator_pipe_state);
Tensor acc = get<0>(slice_accumulator(tmem_storage, mma2transform_pipe_state.index()));
Tensor acc = get<0>(slice_accumulator(tmem_storage, accumulator_pipe_state.index()));
Tensor tAcc = acc(make_coord(_,_),_0{},_0{});
Tensor tAcc_epi = flat_divide(tAcc, EpilogueTile{}); // (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N)
Tensor tTR_tAcc = thread_t2r_epi.partition_S(tAcc_epi); // (T2R, T2R_M, T2R_N, EPI_M, EPI_N)
@@ -1128,15 +1196,15 @@ struct CollectiveMma<
}
}
cutlass::arch::fence_view_async_tmem_load();
mma2transform_pipeline.consumer_release(mma2transform_pipe_state);
accumulator_pipeline.consumer_release(accumulator_pipe_state);
// release acc
++mma2transform_pipe_state;
++accumulator_pipe_state;
}
--k_tile_count;
}
return cute::make_tuple(tTR_FullAcc, tiled_t2r_epi, cute::make_tuple(mma2transform_pipe_state, load2transform_pipe_state));
return cute::make_tuple(tTR_FullAcc, tiled_t2r_epi, cute::make_tuple(accumulator_pipe_state, mainloop_sf_pipe_state));
}
protected:
@@ -866,6 +866,11 @@ struct CollectiveMma<
// PIPELINED MAIN LOOP
//
tiled_mma.accumulate_ = UMMA::ScaleOut::Zero;
if constexpr (not IsOverlappingAccum) {
// Wait for tmem accumulator buffer to become empty with a flipped phase
accumulator_pipeline.producer_acquire(accumulator_pipe_producer_state);
}
CUTLASS_PRAGMA_NO_UNROLL
while (k_tile_count > 0) {
// WAIT on mainloop_pipe_consumer_state until its data are available
@@ -884,15 +889,23 @@ struct CollectiveMma<
// Peek at next iteration
barrier_token = mainloop_pipeline.consumer_try_wait(mainloop_pipe_consumer_state, skip_wait);
if (iter % UtccpReuseCnt == 0) {
if constexpr (UtccpReuseCnt == 1) {
if (cute::elect_one_sync()) {
copy(tiled_copy_s2t_E, thr_tCsE_s2t(_,_,_,_,read_stage), thr_tCtE_s2t);
}
}
else {
if (not (iter & 1)) {
if (cute::elect_one_sync()) {
copy(tiled_copy_s2t_E, thr_tCsE_s2t(_,_,_,_,read_stage), thr_tCtE_s2t);
}
}
}
// Wait for tmem accumulator buffer to become empty with a flipped phase
if (iter == 0) {
accumulator_pipeline.producer_acquire(accumulator_pipe_producer_state);
if constexpr (IsOverlappingAccum) {
if (iter == 0) {
accumulator_pipeline.producer_acquire(accumulator_pipe_producer_state);
}
}
// Unroll the K mode manually so we can set scale C to 1
+31 -4
View File
@@ -475,6 +475,15 @@ struct KernelTmaWarpSpecializedMmaTransformSm100 final {
static constexpr int AccumulatorPipelineStageCount = AccumulatorPipelineStageCount_;
};
template<
int SchedulerPipelineStageCount_,
int AccumulatorPipelineStageCount_
>
struct KernelPtrArrayTmaWarpSpecializedMmaTransformSm100 final {
static constexpr int SchedulerPipelineStageCount = SchedulerPipelineStageCount_;
static constexpr int AccumulatorPipelineStageCount = AccumulatorPipelineStageCount_;
};
// Sparse Gemm
template<
int SchedulerPipelineStageCount_,
@@ -602,12 +611,16 @@ struct KernelScheduleSm100PtrArrayDenseGemm : KernelScheduleSm100DenseGemm {};
struct KernelPtrArrayTmaWarpSpecialized1SmSm100 final : KernelSchedule1Sm, KernelScheduleSm100PtrArrayDenseGemm {};
struct KernelPtrArrayTmaWarpSpecialized2SmSm100 final : KernelSchedule2Sm, KernelScheduleSm100PtrArrayDenseGemm {};
///////////////////////////////////////////////////////////////////////////////////////////////////////
// SM100 Blockwise GEMM Dispatch Policies
// SM100 Blockwise GEMM + Ptr-Array GEMM Dispatch Policies
///////////////////////////////////////////////////////////////////////////////////////////////////////
struct KernelScheduleSm100Blockwise : KernelScheduleSm100 {};
struct KernelTmaWarpSpecializedBlockwise1SmSm100 final : KernelSchedule1Sm, KernelScheduleSm100Blockwise {};
struct KernelTmaWarpSpecializedBlockwise2SmSm100 final : KernelSchedule2Sm, KernelScheduleSm100Blockwise {};
struct KernelScheduleSm100PtrArrayBlockwise : KernelScheduleSm100Blockwise {};
struct KernelPtrArrayTmaWarpSpecializedBlockwise1SmSm100 final : KernelSchedule1Sm, KernelScheduleSm100PtrArrayBlockwise {};
struct KernelPtrArrayTmaWarpSpecializedBlockwise2SmSm100 final : KernelSchedule2Sm, KernelScheduleSm100PtrArrayBlockwise {};
///////////////////////////////////////////////////////////////////////////////////////////////////////
// SM100 Planar Complex GEMM Dispatch Policies
@@ -728,14 +741,13 @@ struct KernelScheduleF8f6f4Sm120 final : KernelScheduleSm120DenseGemm {};
struct KernelScheduleBlockScaledGemmSm120 : KernelScheduleSm120 {};
struct KernelScheduleMxf8f6f4Sm120 : KernelScheduleBlockScaledGemmSm120 {};
struct KernelScheduleMxNvf4Sm120 : KernelScheduleBlockScaledGemmSm120 {};
// Block Scaled Sparse GEMM: Specialize for instruction type, scale factor vector size.
// Block Scaled GEMM: Specialize for instruction type, scale factor vector size.
struct KernelTmaWarpSpecializedNvf4Sm120 final : KernelScheduleMxNvf4Sm120, KernelTmaWarpSpecializedCooperative { };
struct KernelTmaWarpSpecializedPingpongNvf4Sm120 final : KernelScheduleMxNvf4Sm120, KernelTmaWarpSpecializedPingpong { };
struct KernelTmaWarpSpecializedMxf4Sm120 final : KernelScheduleMxNvf4Sm120, KernelTmaWarpSpecializedCooperative { };
struct KernelTmaWarpSpecializedPingpongMxf4Sm120 final : KernelScheduleMxNvf4Sm120, KernelTmaWarpSpecializedPingpong { };
struct KernelTmaWarpSpecializedMxf8f6f4Sm120 final : KernelScheduleMxf8f6f4Sm120, KernelTmaWarpSpecializedCooperative { };
struct KernelTmaWarpSpecializedPingpongMxf8f6f4Sm120 final : KernelScheduleMxf8f6f4Sm120, KernelTmaWarpSpecializedPingpong { };
///////////////////////////////////////////////////////////////////////////////////////////////////////
// SM120 Sparse GEMM Dispatch Policies
///////////////////////////////////////////////////////////////////////////////////////////////////////
@@ -786,6 +798,21 @@ struct MainloopSm100TmaUmmaWarpSpecializedBlockwiseScaling {
constexpr static bool IsOverlappingAccum = false;
};
// n-buffer in smem, pipelined with Blackwell UMMA and TMA, Warp specialized dynamic schedule
template<
int Stages_,
int SchedulerPipelineStageCount_,
int AccumulatorPipelineStageCount_,
class ClusterShape_ = Shape<_1,_1,_1>
>
struct MainloopSm100ArrayTmaUmmaWarpSpecializedBlockwiseScaling {
constexpr static int Stages = Stages_;
using ClusterShape = ClusterShape_;
using ArchTag = arch::Sm100;
using Schedule = KernelPtrArrayTmaWarpSpecializedMmaTransformSm100<SchedulerPipelineStageCount_, AccumulatorPipelineStageCount_>;
constexpr static bool IsOverlappingAccum = false;
};
// n-buffer in smem, pipelined with Blackwell UMMA and TMA, Warp specialized dynamic schedule
template<
int Stages_,
@@ -68,6 +68,7 @@ struct IsCutlass3ArrayKernel<ProblemShape, cute::void_t<typename ProblemShape::U
#include "cutlass/gemm/kernel/sm100_gemm_array_tma_warpspecialized.hpp"
#include "cutlass/gemm/kernel/sm100_gemm_tma_warpspecialized_input_transform.hpp"
#include "cutlass/gemm/kernel/sm100_gemm_array_tma_warpspecialized_input_transform.hpp"
#include "cutlass/gemm/kernel/sm100_gemm_array_tma_warpspecialized_mma_transform.hpp"
#include "cutlass/gemm/kernel/sm100_sparse_gemm_tma_warpspecialized.hpp"
#include "cutlass/gemm/kernel/sm120_gemm_tma_warpspecialized_cooperative_asymmetric_dma.hpp"
////////////////////////////////////////////////////////////////////////////////
@@ -131,16 +131,19 @@ public:
static constexpr bool IsGdcEnabled = cutlass::arch::IsGdcGloballyEnabled;
// Warp specialization thread count per threadblock
static constexpr uint32_t NumSchedThreads = NumThreadsPerWarp; // 1 warp
static constexpr uint32_t NumMMAThreads = NumThreadsPerWarp; // 1 warp
static constexpr uint32_t NumMainloopLoadThreads = NumThreadsPerWarp; // 1 warp
static constexpr uint32_t NumEpilogueLoadThreads = NumThreadsPerWarp; // 1 warp
static constexpr uint32_t NumEpilogueThreads = CollectiveEpilogue::ThreadCount;
static constexpr uint32_t NumEpilogueWarps = NumEpilogueThreads / NumThreadsPerWarp;
static constexpr uint32_t NumSchedThreads = NumThreadsPerWarp; // 1 warp
static constexpr uint32_t NumMMAThreads = NumThreadsPerWarp; // 1 warp
static constexpr uint32_t NumMainloopABLoadThreads = NumThreadsPerWarp; // 1 warp
static constexpr uint32_t NumEpilogueLoadThreads = NumThreadsPerWarp; // 1 warp
static constexpr uint32_t NumEpilogueThreads = CollectiveEpilogue::ThreadCount;
static constexpr uint32_t NumEpilogueWarps = NumEpilogueThreads / NumThreadsPerWarp;
static constexpr uint32_t NumMainloopSFLoadThreads = NumThreadsPerWarp; // 1 warp
static constexpr uint32_t MaxThreadsPerBlock = NumSchedThreads +
NumMainloopLoadThreads + NumMMAThreads +
NumEpilogueLoadThreads + NumEpilogueThreads;
static constexpr uint32_t MaxThreadsPerBlock = cute::round_up(NumSchedThreads +
NumMainloopABLoadThreads + NumMMAThreads +
NumEpilogueLoadThreads + NumEpilogueThreads +
NumMainloopSFLoadThreads, 128);
static constexpr uint32_t MinBlocksPerMultiprocessor = 1;
static constexpr uint32_t NumEpilogueSubTiles = CollectiveEpilogue::get_load_pipe_increment(CtaShape_MNK{});
@@ -152,8 +155,8 @@ public:
static constexpr uint32_t CLCResponseSize = sizeof(typename TileScheduler::CLCResponse);
// Pipeline and pipeline state types
using MainloopPipeline = typename CollectiveMainloop::MainloopPipeline;
using MainloopPipelineState = typename CollectiveMainloop::MainloopPipelineState;
using MainloopABPipeline = typename CollectiveMainloop::MainloopABPipeline;
using MainloopABPipelineState = typename CollectiveMainloop::MainloopABPipelineState;
using EpiLoadPipeline = typename CollectiveEpilogue::LoadPipeline;
using EpiLoadPipelineState = typename CollectiveEpilogue::LoadPipelineState;
@@ -163,11 +166,11 @@ public:
using LoadOrderBarrier = cutlass::OrderedSequenceBarrier<1,2>;
using Mma2TransformPipeline = typename CollectiveMainloop::Mma2TransformPipeline;
using Mma2TransformPipelineState = typename Mma2TransformPipeline::PipelineState;
using AccumulatorPipeline = typename CollectiveMainloop::AccumulatorPipeline;
using AccumulatorPipelineState = typename AccumulatorPipeline::PipelineState;
using Load2TransformPipeline = typename CollectiveMainloop::Load2TransformPipeline;
using Load2TransformPipelineState = typename Load2TransformPipeline::PipelineState;
using MainloopSFPipeline = typename CollectiveMainloop::MainloopSFPipeline;
using MainloopSFPipelineState = typename MainloopSFPipeline::PipelineState;
using CLCPipeline = cutlass::PipelineCLCFetchAsync<SchedulerPipelineStageCount, ClusterShape>;
using CLCPipelineState = typename CLCPipeline::PipelineState;
@@ -178,7 +181,7 @@ public:
using TmemAllocator = cute::conditional_t<cute::size(cute::shape<0>(typename TiledMma::ThrLayoutVMNK{})) == 1,
cute::TMEM::Allocator1Sm, cute::TMEM::Allocator2Sm>;
static constexpr uint32_t GenericRegisterRequirement = 104;
static constexpr uint32_t GenericRegisterRequirement = 48;
static constexpr uint32_t AccumRegisterRequirement = 256;
// Kernel level shared memory storage
@@ -186,19 +189,15 @@ public:
// Barriers should be allocated in lower 8KB of SMEM for SM100
struct PipelineStorage : cute::aligned_struct<16, _1> {
using MainloopPipelineStorage = typename CollectiveMainloop::PipelineStorage;
using Load2TransformPipelineStorage = typename CollectiveMainloop::Load2TransformPipelineStorage;
using EpiLoadPipelineStorage = typename CollectiveEpilogue::PipelineStorage;
using LoadOrderBarrierStorage = typename LoadOrderBarrier::SharedStorage;
using CLCPipelineStorage = typename CLCPipeline::SharedStorage;
using Mma2TransformPipelineStorage = typename CollectiveMainloop::Mma2TransformPipelineStorage;
using CLCThrottlePipelineStorage = typename CLCThrottlePipeline::SharedStorage;
alignas(16) MainloopPipelineStorage mainloop;
alignas(16) Load2TransformPipelineStorage load2transform;
alignas(16) EpiLoadPipelineStorage epi_load;
alignas(16) LoadOrderBarrierStorage load_order;
alignas(16) CLCPipelineStorage clc;
alignas(16) Mma2TransformPipelineStorage mma2transform;
alignas(16) CLCThrottlePipelineStorage clc_throttle;
alignas(16) arch::ClusterBarrier tmem_dealloc;
alignas(16) arch::ClusterBarrier epilogue_throttle;
@@ -240,19 +239,23 @@ public:
};
enum class WarpCategory : int32_t {
MMA = 0,
Sched = 1,
MainloopLoad = 2,
EpilogueLoad = 3,
Epilogue = 4
MMA = 0,
Sched = 1,
MainloopABLoad = 2,
EpilogueLoad = 3,
Epilogue = 4, // 4 warps
MainloopSFLoad = 8,
Unused = 9,
};
struct IsParticipant {
uint32_t mma = false;
uint32_t sched = false;
uint32_t main_load = false;
uint32_t epi_load = false;
uint32_t epilogue = false;
uint32_t mma = false;
uint32_t sched = false;
uint32_t main_ab_load = false;
uint32_t epi_load = false;
uint32_t epilogue = false;
uint32_t main_sf_load = false;
uint32_t unused = false;
};
//
@@ -407,8 +410,20 @@ public:
// Account for more than one epilogue warp
int warp_idx = canonical_warp_idx_sync();
WarpCategory warp_category = warp_idx < static_cast<int>(WarpCategory::Epilogue) ? WarpCategory(warp_idx)
: WarpCategory::Epilogue;
WarpCategory warp_category = [&] () CUTLASS_LAMBDA_FUNC_INLINE {
if (warp_idx < static_cast<int>(WarpCategory::Epilogue)) {
return WarpCategory(warp_idx);
}
else if (warp_idx < static_cast<int>(WarpCategory::MainloopSFLoad)) {
return WarpCategory::Epilogue;
}
else if (warp_idx == static_cast<int>(WarpCategory::MainloopSFLoad)) {
return WarpCategory::MainloopSFLoad;
}
else {
return WarpCategory::Unused;
}
}();
uint32_t lane_predicate = cute::elect_one_sync();
auto cluster_shape = cutlass::detail::select_cluster_shape(ClusterShape{});
@@ -440,41 +455,43 @@ public:
IsParticipant is_participant = {
(warp_category == WarpCategory::MMA), // mma
(warp_category == WarpCategory::Sched) && is_first_cta_in_cluster, // sched
(warp_category == WarpCategory::MainloopLoad), // main_load
(warp_category == WarpCategory::MainloopABLoad), // main_ab_load
(warp_category == WarpCategory::EpilogueLoad) && is_epi_load_needed, // epi_load
(warp_category == WarpCategory::Epilogue) // epilogue
(warp_category == WarpCategory::Epilogue), // epilogue
(warp_category == WarpCategory::MainloopSFLoad), // main_sf_load
(warp_category == WarpCategory::Unused) // unused
};
// Mainloop Load pipeline
typename MainloopPipeline::Params mainloop_pipeline_params;
if (WarpCategory::MainloopLoad == warp_category) {
mainloop_pipeline_params.role = MainloopPipeline::ThreadCategory::Producer;
typename MainloopABPipeline::Params mainloop_ab_pipeline_params;
if (WarpCategory::MainloopABLoad == warp_category) {
mainloop_ab_pipeline_params.role = MainloopABPipeline::ThreadCategory::Producer;
}
if (WarpCategory::MMA == warp_category) {
mainloop_pipeline_params.role = MainloopPipeline::ThreadCategory::Consumer;
mainloop_ab_pipeline_params.role = MainloopABPipeline::ThreadCategory::Consumer;
}
mainloop_pipeline_params.is_leader = lane_predicate && is_mma_leader_cta && is_participant.main_load;
mainloop_pipeline_params.transaction_bytes = CollectiveMainloop::TmaTransactionBytes;
mainloop_pipeline_params.initializing_warp = 0;
MainloopPipeline mainloop_pipeline(shared_storage.pipelines.mainloop,
mainloop_pipeline_params,
cluster_shape,
cute::true_type{}, // Perform barrier init
cute::false_type{}); // Delay mask calculation
mainloop_ab_pipeline_params.is_leader = lane_predicate && is_mma_leader_cta && is_participant.main_ab_load;
mainloop_ab_pipeline_params.transaction_bytes = CollectiveMainloop::TmaTransactionBytes;
mainloop_ab_pipeline_params.initializing_warp = 0;
MainloopABPipeline mainloop_ab_pipeline(shared_storage.pipelines.mainloop.pipeline_ab,
mainloop_ab_pipeline_params,
cluster_shape,
cute::true_type{}, // Perform barrier init
cute::false_type{}); // Delay mask calculation
typename Load2TransformPipeline::Params load2transform_pipeline_params;
if (WarpCategory::MainloopLoad == warp_category) {
load2transform_pipeline_params.role = Load2TransformPipeline::ThreadCategory::Producer;
typename MainloopSFPipeline::Params mainloop_sf_pipeline_params;
if (WarpCategory::MainloopSFLoad == warp_category) {
mainloop_sf_pipeline_params.role = MainloopSFPipeline::ThreadCategory::Producer;
}
if (WarpCategory::Epilogue == warp_category) {
load2transform_pipeline_params.role = Load2TransformPipeline::ThreadCategory::Consumer;
mainloop_sf_pipeline_params.role = MainloopSFPipeline::ThreadCategory::Consumer;
}
load2transform_pipeline_params.initializing_warp = 0;
load2transform_pipeline_params.producer_arv_count = CollectiveMainloop::NumLoad2TransformProducerThreadEvents;
load2transform_pipeline_params.consumer_arv_count = NumEpilogueThreads;
mainloop_sf_pipeline_params.initializing_warp = 8;
mainloop_sf_pipeline_params.producer_arv_count = CollectiveMainloop::NumMainloopSFProducerThreadEvents;
mainloop_sf_pipeline_params.consumer_arv_count = NumEpilogueThreads;
Load2TransformPipeline load2transform_pipeline(shared_storage.pipelines.load2transform,
load2transform_pipeline_params);
MainloopSFPipeline mainloop_sf_pipeline(shared_storage.pipelines.mainloop.pipeline_sf,
mainloop_sf_pipeline_params);
// Epilogue Load pipeline
typename EpiLoadPipeline::Params epi_load_pipeline_params;
@@ -498,8 +515,8 @@ public:
// Load order barrier
typename LoadOrderBarrier::Params load_order_barrier_params;
load_order_barrier_params.group_id = (warp_category == WarpCategory::MainloopLoad) ? 0 : 1;
load_order_barrier_params.group_size = NumMainloopLoadThreads;
load_order_barrier_params.group_id = (warp_category == WarpCategory::MainloopABLoad) ? 0 : 1;
load_order_barrier_params.group_size = NumMainloopABLoadThreads;
load_order_barrier_params.initializing_warp = 5;
LoadOrderBarrier load_order_barrier(shared_storage.pipelines.load_order, load_order_barrier_params);
@@ -514,7 +531,8 @@ public:
clc_pipeline_params.producer_blockid = 0;
clc_pipeline_params.producer_arv_count = 1;
clc_pipeline_params.consumer_arv_count = NumSchedThreads + cluster_size *
(NumMainloopLoadThreads + NumEpilogueThreads + NumMMAThreads);
(NumMainloopABLoadThreads + NumEpilogueThreads +
NumMMAThreads + NumMainloopSFLoadThreads);
if (is_epi_load_needed) {
clc_pipeline_params.consumer_arv_count += cluster_size * NumEpilogueLoadThreads;
}
@@ -523,30 +541,30 @@ public:
CLCPipeline clc_pipeline(shared_storage.pipelines.clc, clc_pipeline_params, cluster_shape);
// Mainloop-Epilogue pipeline
typename Mma2TransformPipeline::Params mma2transform_pipeline_params;
typename AccumulatorPipeline::Params accumulator_pipeline_params;
if (WarpCategory::MMA == warp_category) {
mma2transform_pipeline_params.role = Mma2TransformPipeline::ThreadCategory::Producer;
accumulator_pipeline_params.role = AccumulatorPipeline::ThreadCategory::Producer;
}
if (WarpCategory::Epilogue == warp_category) {
mma2transform_pipeline_params.role = Mma2TransformPipeline::ThreadCategory::Consumer;
accumulator_pipeline_params.role = AccumulatorPipeline::ThreadCategory::Consumer;
}
// Only one producer thread arrives on this barrier.
mma2transform_pipeline_params.producer_arv_count = 1;
mma2transform_pipeline_params.consumer_arv_count = size(AtomThrShapeMNK{}) * NumEpilogueThreads;
mma2transform_pipeline_params.initializing_warp = 2;
Mma2TransformPipeline mma2transform_pipeline(shared_storage.pipelines.mma2transform,
mma2transform_pipeline_params,
accumulator_pipeline_params.producer_arv_count = 1;
accumulator_pipeline_params.consumer_arv_count = size(AtomThrShapeMNK{}) * NumEpilogueThreads;
accumulator_pipeline_params.initializing_warp = 2;
AccumulatorPipeline accumulator_pipeline(shared_storage.pipelines.mainloop.pipeline_accum,
accumulator_pipeline_params,
cluster_shape);
// CLC throttle pipeline
typename CLCThrottlePipeline::Params clc_throttle_pipeline_params;
if (WarpCategory::MainloopLoad == warp_category) {
if (WarpCategory::MainloopABLoad == warp_category) {
clc_throttle_pipeline_params.role = CLCThrottlePipeline::ThreadCategory::Producer;
}
if (WarpCategory::Sched == warp_category) {
clc_throttle_pipeline_params.role = CLCThrottlePipeline::ThreadCategory::Consumer;
}
clc_throttle_pipeline_params.producer_arv_count = NumMainloopLoadThreads;
clc_throttle_pipeline_params.producer_arv_count = NumMainloopABLoadThreads;
clc_throttle_pipeline_params.consumer_arv_count = NumSchedThreads;
clc_throttle_pipeline_params.dst_blockid = 0;
clc_throttle_pipeline_params.initializing_warp = 3;
@@ -573,7 +591,7 @@ public:
if (WarpCategory::MMA == warp_category && lane_predicate) {
epilogue_throttle_barrier.init( NumMMAThreads +
(is_first_cta_in_cluster ? NumSchedThreads : 0) +
NumMainloopLoadThreads +
NumMainloopABLoadThreads +
(is_epi_load_needed ? NumEpilogueLoadThreads : 0));
}
@@ -581,11 +599,11 @@ public:
// To all producers and consumer threadblocks in the cluster
pipeline_init_arrive_relaxed(cluster_size);
auto load_inputs = collective_mainloop.load_init(
auto load_inputs = collective_mainloop.load_ab_init(
problem_shape_MNKL, params.mainloop, shared_storage.tensors.mainloop);
MainloopPipelineState mainloop_pipe_consumer_state;
MainloopPipelineState mainloop_pipe_producer_state = cutlass::make_producer_start_state<MainloopPipeline>();
MainloopABPipelineState mainloop_ab_pipe_consumer_state;
MainloopABPipelineState mainloop_ab_pipe_producer_state = cutlass::make_producer_start_state<MainloopABPipeline>();
EpiLoadPipelineState epi_load_pipe_consumer_state;
EpiLoadPipelineState epi_load_pipe_producer_state = cutlass::make_producer_start_state<EpiLoadPipeline>();
@@ -596,17 +614,17 @@ public:
CLCPipelineState clc_pipe_consumer_state;
CLCPipelineState clc_pipe_producer_state = cutlass::make_producer_start_state<CLCPipeline>();
Mma2TransformPipelineState mma2transform_pipe_consumer_state;
Mma2TransformPipelineState mma2transform_pipe_producer_state = cutlass::make_producer_start_state<Mma2TransformPipeline>();
AccumulatorPipelineState accumulator_pipe_consumer_state;
AccumulatorPipelineState accumulator_pipe_producer_state = cutlass::make_producer_start_state<AccumulatorPipeline>();
Load2TransformPipelineState load2transform_pipe_consumer_state;
Load2TransformPipelineState load2transform_pipe_producer_state = cutlass::make_producer_start_state<Load2TransformPipeline>();
MainloopSFPipelineState mainloop_sf_pipe_consumer_state;
MainloopSFPipelineState mainloop_sf_pipe_producer_state = cutlass::make_producer_start_state<MainloopSFPipeline>();
dim3 block_id_in_cluster = cute::block_id_in_cluster();
// Calculate mask after cluster barrier arrival
mainloop_pipeline.init_masks(cluster_shape, block_id_in_cluster);
mma2transform_pipeline.init_masks(cluster_shape, block_id_in_cluster);
mainloop_ab_pipeline.init_masks(cluster_shape, block_id_in_cluster);
accumulator_pipeline.init_masks(cluster_shape, block_id_in_cluster);
// TileID scheduler
TileScheduler scheduler(&shared_storage.clc_response[0], params.scheduler, block_id_in_cluster);
@@ -619,7 +637,7 @@ public:
pipeline_init_wait(cluster_size);
if (is_participant.main_load) {
if (is_participant.main_ab_load) {
// Register reconfiguration
arch::warpgroup_reg_dealloc<GenericRegisterRequirement>();
@@ -633,15 +651,12 @@ public:
epilogue_throttle_barrier.arrive();
bool requires_clc_query = true;
auto pipelines = cute::make_tuple(mainloop_pipeline, load2transform_pipeline);
auto states = cute::make_tuple(mainloop_pipe_producer_state, load2transform_pipe_producer_state);
do {
// Get the number of K tiles to compute for this work as well as the starting K tile offset of the work.
auto k_tile_iter = scheduler.get_k_tile_iterator(work_tile_info, problem_shape_MNKL, CtaShape_MNK{}, load_inputs.k_tiles);
auto k_tile_count = TileScheduler::get_work_k_tile_count(work_tile_info, problem_shape_MNKL, CtaShape_MNK{});
auto k_tile_prologue = min(MainloopPipeline::Stages, k_tile_count);
auto k_tile_prologue = min(MainloopABPipeline::Stages, k_tile_count);
if constexpr (IsSchedDynamicPersistent) {
if (is_first_cta_in_cluster && requires_clc_query) {
@@ -652,34 +667,28 @@ public:
}
// Start mainloop prologue loads, arrive on the epilogue residual load barrier, resume mainloop loads
auto [mainloop_producer_state_next, load2transform_producer_state_next, k_tile_iter_next] = collective_mainloop.load(
mainloop_pipeline,
load2transform_pipeline,
mainloop_pipe_producer_state,
load2transform_pipe_producer_state,
auto [mainloop_ab_producer_state_next, k_tile_iter_next] = collective_mainloop.load_ab(
mainloop_ab_pipeline,
mainloop_ab_pipe_producer_state,
load_inputs,
cta_coord_mnkl,
k_tile_iter, k_tile_prologue
);
mainloop_pipe_producer_state = mainloop_producer_state_next;
load2transform_pipe_producer_state = load2transform_producer_state_next;
mainloop_ab_pipe_producer_state = mainloop_ab_producer_state_next;
if (do_load_order_arrive) {
load_order_barrier.arrive();
do_load_order_arrive = false;
}
auto [mainloop_producer_state_next_, load2transform_producer_state_next_, unused_] = collective_mainloop.load(
mainloop_pipeline,
load2transform_pipeline,
mainloop_pipe_producer_state,
load2transform_pipe_producer_state,
auto [mainloop_ab_producer_state_next_, unused_] = collective_mainloop.load_ab(
mainloop_ab_pipeline,
mainloop_ab_pipe_producer_state,
load_inputs,
cta_coord_mnkl,
k_tile_iter_next, k_tile_count - k_tile_prologue
);
mainloop_pipe_producer_state = mainloop_producer_state_next_;
load2transform_pipe_producer_state = load2transform_producer_state_next_;
mainloop_ab_pipe_producer_state = mainloop_ab_producer_state_next_;
// Sync warp to prevent non-participating threads entering next wave early
__syncwarp();
@@ -697,11 +706,61 @@ public:
}
} while (work_tile_info.is_valid());
collective_mainloop.load_tail(
mainloop_pipeline,
load2transform_pipeline,
mainloop_pipe_producer_state,
load2transform_pipe_producer_state
collective_mainloop.load_ab_tail(
mainloop_ab_pipeline,
mainloop_ab_pipe_producer_state
);
}
else if (is_participant.main_sf_load) {
auto mainloop_sf_inputs = collective_mainloop.load_sf_init(
problem_shape_MNKL, params.mainloop, shared_storage.tensors.mainloop);
// Register reconfiguration
arch::warpgroup_reg_dealloc<GenericRegisterRequirement>();
// Ensure that the prefetched kernel does not touch
// unflushed global memory prior to this instruction
cutlass::arch::wait_on_dependent_grids();
bool requires_clc_query = true;
do {
// Get the number of K tiles to compute for this work as well as the starting K tile offset of the work.
auto k_tile_iter = scheduler.get_k_tile_iterator(work_tile_info, problem_shape_MNKL, CtaShape_MNK{}, mainloop_sf_inputs.k_tiles);
auto k_tile_count = TileScheduler::get_work_k_tile_count(work_tile_info, problem_shape_MNKL, CtaShape_MNK{});
// Start mainloop prologue loads, arrive on the epilogue residual load barrier, resume mainloop loads
auto [mainloop_sf_producer_state_next, k_tile_iter_next] = collective_mainloop.load_sf(
mainloop_sf_pipeline,
mainloop_sf_pipe_producer_state,
mainloop_sf_inputs,
cta_coord_mnkl,
k_tile_iter, k_tile_count
);
mainloop_sf_pipe_producer_state = mainloop_sf_producer_state_next;
// Sync warp to prevent non-participating threads entering next wave early
__syncwarp();
auto [next_work_tile_info, increment_pipe] = scheduler.fetch_next_work(
work_tile_info,
clc_pipeline,
clc_pipe_consumer_state
);
work_tile_info = next_work_tile_info;
cta_coord_mnkl = scheduler.work_tile_to_cta_coord(work_tile_info);
requires_clc_query = increment_pipe;
if (increment_pipe) {
++clc_pipe_consumer_state;
}
} while (work_tile_info.is_valid());
collective_mainloop.load_sf_tail(
mainloop_sf_pipeline,
mainloop_sf_pipe_producer_state
);
}
@@ -791,16 +850,16 @@ public:
}
if (is_mma_leader_cta) {
auto [mainloop_pipe_consumer_state_, mma2transform_pipe_producer_state_] = collective_mainloop.mma(
cute::make_tuple(mainloop_pipeline, mma2transform_pipeline),
cute::make_tuple(mainloop_pipe_consumer_state, mma2transform_pipe_producer_state),
auto [mainloop_ab_pipe_consumer_state_, accumulator_pipe_producer_state_] = collective_mainloop.mma(
cute::make_tuple(mainloop_ab_pipeline, accumulator_pipeline),
cute::make_tuple(mainloop_ab_pipe_consumer_state, accumulator_pipe_producer_state),
tmem_storage,
mma_inputs,
cta_coord_mnkl,
k_tile_count
);
mainloop_pipe_consumer_state = mainloop_pipe_consumer_state_;
mma2transform_pipe_producer_state = mma2transform_pipe_producer_state_;
mainloop_ab_pipe_consumer_state = mainloop_ab_pipe_consumer_state_;
accumulator_pipe_producer_state = accumulator_pipe_producer_state_;
}
work_tile_info = next_work_tile_info;
@@ -817,7 +876,7 @@ public:
// Leader MMA waits for leader + peer epilogues to release stage
if (is_mma_leader_cta) {
mma2transform_pipeline.producer_tail(mma2transform_pipe_producer_state);
accumulator_pipeline.producer_tail(accumulator_pipe_producer_state);
}
// Signal to peer MMA that entire tmem allocation can be deallocated
if constexpr (has_mma_peer_cta) {
@@ -912,13 +971,13 @@ public:
uint32_t tmem_base_ptr = shared_storage.tmem_base_ptr;
collective_mainloop.set_tmem_offsets(tmem_storage, tmem_base_ptr);
auto transform_inputs = collective_mainloop.transform_init(
auto accum_inputs = collective_mainloop.accum_init(
problem_shape_MNKL,
shared_storage.tensors.mainloop
);
auto pipelines = cute::make_tuple(mma2transform_pipeline, load2transform_pipeline);
auto states = cute::make_tuple(mma2transform_pipe_consumer_state, load2transform_pipe_consumer_state);
auto pipelines = cute::make_tuple(accumulator_pipeline, mainloop_sf_pipeline);
auto states = cute::make_tuple(accumulator_pipe_consumer_state, mainloop_sf_pipe_consumer_state);
bool do_tail_store = false;
do {
@@ -935,11 +994,11 @@ public:
++clc_pipe_consumer_state;
}
auto [accum, tiled_t2r, next_state] = collective_mainloop.transform(
auto [accum, tiled_t2r, next_state] = collective_mainloop.accum(
pipelines,
states,
tmem_storage,
transform_inputs,
accum_inputs,
cta_coord_mnkl,
typename CollectiveEpilogue::CopyOpT2R{},
typename CollectiveEpilogue::EpilogueTile{},
@@ -405,7 +405,7 @@ public:
return make_coord(m_coord, n_coord, _, l_coord);
}
CUTLASS_DEVICE
CUTLASS_HOST_DEVICE
static void
issue_clc_query(PipelineState<Stages> state, uint32_t mbarrier_addr, CLCResponse* clc_response_ptr) {
#if defined(CUTLASS_ARCH_CLC_ENABLED)
@@ -468,7 +468,7 @@ public:
// Kernel helper function to get next work tile
template <class TileSchedulerPipeline, class TileSchedulerPipelineState>
CUTLASS_DEVICE
CUTLASS_HOST_DEVICE
auto
fetch_next_work(
WorkTileInfo work_tile_info,
@@ -627,9 +627,10 @@ public:
store_query_response(state, make_invalid_response());
}
CUTLASS_DEVICE
CUTLASS_HOST_DEVICE
void
store_query_response(PipelineState<Stages> state, CLCResponse clc_response) {
#if defined(__CUDA_ARCH__)
uint32_t smem_ptr = cute::cast_smem_ptr_to_uint(&clc_response_ptr_[state.index()]);
asm volatile("st.shared.v4.b32 [%0], {%1, %2, %3, %4};\n"
: : "r"(smem_ptr)
@@ -638,6 +639,7 @@ public:
, "r"(clc_response.data[2])
, "r"(clc_response.data[3]));
cutlass::arch::fence_view_async_shared();
#endif
}
CUTLASS_DEVICE