update 3.8 v2 (#2112)

* update 3.8 v2

* update 3.8

---------

Co-authored-by: yuzhai <yuzhai@nvidia.com>
This commit is contained in:
Yujia Zhai
2025-02-19 22:03:14 -05:00
committed by GitHub
co-authored by yuzhai
parent e9627ce55b
commit b84e9802d8
166 changed files with 3986 additions and 4037 deletions
@@ -567,7 +567,7 @@ sm100_make_trivial_fastFP32_tiled_mma() {
}
/**
* @brief Check for U4_UNPACK_U8, U6_UNPACK_U8 alignment requirement
* @brief Check for F8F6F4 alignment requirement
*
* @tparam TileShape_MNK (MmaAtomShape_M, MmaAtomShape_N, TileShape_K)
* @tparam ClusterShape_MNK (cluster_M, cluster_N, cluster_K)
@@ -85,7 +85,7 @@ compute_stage_count_or_override(StageCountAutoCarveout<carveout_bytes_> stage_co
}
// Returns the maximum number of smem tiles that can be used with a given smem capacity in gemm of blockwise/groupwise scale.
template<int capacity_bytes_, class ElementA, class ElementB, class ElementBlockScale, class TileShapeMNK, int ScaleMsPerTile, int carveout_bytes_, int alignment = 128>
template<int capacity_bytes_, class ElementA, class ElementB, class ElementBlockScale, class TileShapeMNK, int ScaleMsPerTile, int ScaleNsPerTile, int carveout_bytes_, int alignment = 128>
constexpr int
compute_stage_count_with_blockwise_scale(StageCountAutoCarveout<carveout_bytes_> stage_count) {
constexpr auto mainloop_pipeline_bytes = sizeof(typename cutlass::PipelineTmaAsync<1>::SharedStorage);
@@ -96,7 +96,7 @@ compute_stage_count_with_blockwise_scale(StageCountAutoCarveout<carveout_bytes_>
cutlass::bits_to_bytes(a_bits * size<0>(TileShapeMNK{}) * size<2>(TileShapeMNK{})) +
cutlass::bits_to_bytes(b_bits * size<1>(TileShapeMNK{}) * size<2>(TileShapeMNK{})) +
cutlass::bits_to_bytes(scale_bits * ScaleMsPerTile) + // scale of tensor A
cutlass::bits_to_bytes(scale_bits * 1); // scale of tensor B
cutlass::bits_to_bytes(scale_bits * ScaleNsPerTile); // scale of tensor B
constexpr int stage_bytes = cutlass::round_up(stage_bytes_, alignment) +
static_cast<int>(mainloop_pipeline_bytes);
@@ -1043,7 +1043,8 @@ template <
class TileShape_MNK,
class ClusterShape_MNK,
class StageCountType,
int ScaleGranularityM_
int ScaleGranularityM_,
int ScaleGranularityN_
>
struct CollectiveBuilder<
arch::Sm90,
@@ -1058,11 +1059,11 @@ struct CollectiveBuilder<
TileShape_MNK,
ClusterShape_MNK,
StageCountType,
KernelTmaWarpSpecializedCooperativeFP8BlockScaledAccum<ScaleGranularityM_>,
KernelTmaWarpSpecializedCooperativeFP8BlockScaledAccum<ScaleGranularityM_, ScaleGranularityN_>,
cute::enable_if_t<
not detail::is_use_rmem_A<ElementA, GmemLayoutATag, ElementB, GmemLayoutBTag>()>
> {
using KernelScheduleType = KernelTmaWarpSpecializedCooperativeFP8BlockScaledAccum<ScaleGranularityM_>;
using KernelScheduleType = KernelTmaWarpSpecializedCooperativeFP8BlockScaledAccum<ScaleGranularityM_, ScaleGranularityN_>;
static_assert(is_static<TileShape_MNK>::value);
static_assert(is_static<ClusterShape_MNK>::value);
@@ -1090,7 +1091,7 @@ struct CollectiveBuilder<
static constexpr bool IsCooperative = cute::is_any_of_v<KernelScheduleType,
KernelTmaWarpSpecializedCooperative,
KernelPtrArrayTmaWarpSpecializedCooperative,
KernelTmaWarpSpecializedCooperativeFP8BlockScaledAccum<ScaleGranularityM_>>;
KernelTmaWarpSpecializedCooperativeFP8BlockScaledAccum<ScaleGranularityM_, ScaleGranularityN_>>;
using AtomLayoutMNK = cute::conditional_t<IsCooperative,
Layout<Shape<_2,_1,_1>>, Layout<Shape<_1,_1,_1>>>;
@@ -1109,12 +1110,15 @@ struct CollectiveBuilder<
static constexpr int KernelSmemCarveout = static_cast<int>(TensorMapStorage);
static constexpr int ScaleGranularityM = ScaleGranularityM_ == 0 ? size<0>(TileShape_MNK{}) : ScaleGranularityM_;
static constexpr int ScaleGranularityN = ScaleGranularityN_ == 0 ? size<1>(TileShape_MNK{}) : ScaleGranularityN_;
static constexpr int ScaleMsPerTile = size<0>(TileShape_MNK{}) / ScaleGranularityM;
static constexpr int ScaleNsPerTile = size<1>(TileShape_MNK{}) / ScaleGranularityN;
static_assert((size<0>(TileShape_MNK{}) % ScaleGranularityM) == 0, "FP8 scaling granularity must evenly divide tile shape along M.");
static_assert((size<1>(TileShape_MNK{}) % ScaleGranularityN) == 0, "FP8 scaling granularity must evenly divide tile shape along N.");
static constexpr int PipelineStages = detail::compute_stage_count_with_blockwise_scale<detail::sm90_smem_capacity_bytes - KernelSmemCarveout,
ElementAMma, ElementBMma, ElementBlockScale, TileShape_MNK, ScaleMsPerTile>(StageCountType{});
using DispatchPolicy = MainloopSm90TmaGmmaWarpSpecializedBlockScalingFP8<PipelineStages, ClusterShape_MNK, KernelScheduleType, ScaleGranularityM_>;
ElementAMma, ElementBMma, ElementBlockScale, TileShape_MNK, ScaleMsPerTile, ScaleNsPerTile>(StageCountType{});
using DispatchPolicy = MainloopSm90TmaGmmaWarpSpecializedBlockScalingFP8<PipelineStages, ClusterShape_MNK, KernelScheduleType, ScaleGranularityM_, ScaleGranularityN_>;
using SmemCopyAtomA = void;
using SmemCopyAtomB = void;
@@ -75,6 +75,15 @@ private:
}
// `multiply` scale the partial accumulators and `add` to main accumulator (FFMA).
CUTLASS_DEVICE
void scale_core(ElementAccumulator const &scale) {
warpgroup_wait<0>();
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < size(accum_); ++i) {
accum_(i) += accum_temp_(i) * scale;
}
}
template <
class EngineScale,
class LayoutScale>
@@ -94,6 +103,31 @@ private:
}
}
template <
class EngineScaleA,
class LayoutScaleA,
class EngineScaleB,
class LayoutScaleB>
CUTLASS_DEVICE
void scale_core(const cute::Tensor<EngineScaleA, LayoutScaleA> &scaleA, const cute::Tensor<EngineScaleB, LayoutScaleB> &scaleB) {
using TensorScaleA = cute::Tensor<EngineScaleA, LayoutScaleA>;
using TensorScaleB = cute::Tensor<EngineScaleB, LayoutScaleB>;
static_assert(is_static<LayoutScaleA>::value, "ScaleA Layout should be static");
static_assert(is_static<LayoutScaleB>::value, "ScaleB Layout should be static");
static_assert(is_rmem<TensorScaleA>::value, "ScaleA tensor must be rmem resident.");
static_assert(is_rmem<TensorScaleB>::value, "ScaleB tensor must be rmem resident.");
static_assert(LayoutAccum{}.shape() == LayoutScaleA{}.shape(), "Accumulator and scaleA must have same shape.");
static_assert(LayoutAccum{}.shape() == LayoutScaleB{}.shape(), "Accumulator and scaleB must have same shape.");
warpgroup_wait<0>();
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < size(accum_); ++i) {
accum_(i) += accum_temp_(i) * scaleA(i) * scaleB(i);
}
}
public:
CUTLASS_DEVICE
GmmaFP8Accumulation(
@@ -152,6 +186,16 @@ public:
//
/// scale (multiply_add) the results from the MMA accumulators to main accumulator if needed.
CUTLASS_DEVICE
void scale_if_needed(ElementAccumulator const &scale) {
mma_count_ += mma_count_per_mainloop_iteration_;
reset_accum_flag_ = __shfl_sync(0xffffffff, mma_count_ == accum_promotion_interval_, 0);
if (reset_accum_flag_) {
scale_core(scale);
mma_count_ = 0;
}
}
template <
class EngineScale,
class LayoutScale>
@@ -165,7 +209,29 @@ public:
}
}
template <
class EngineScaleA,
class LayoutScaleA,
class EngineScaleB,
class LayoutScaleB>
CUTLASS_DEVICE
void scale_if_needed(const cute::Tensor<EngineScaleA, LayoutScaleA> &scaleA, const cute::Tensor<EngineScaleB, LayoutScaleB> &scaleB) {
mma_count_ += mma_count_per_mainloop_iteration_;
reset_accum_flag_ = __shfl_sync(0xffffffff, mma_count_ == accum_promotion_interval_, 0);
if (reset_accum_flag_) {
scale_core(scaleA, scaleB);
mma_count_ = 0;
}
}
/// scale (multiply_add) the residue results from the MMA accumulators to main accumulator if needed.
CUTLASS_DEVICE
void scale_residue_if_needed(ElementAccumulator const &scale) {
if (__shfl_sync(0xffffffff, mma_count_ > 0, 0)) {
scale_core(scale);
}
}
template <
class EngineScale,
class LayoutScale>
@@ -175,6 +241,18 @@ public:
scale_core(scale);
}
}
template <
class EngineScaleA,
class LayoutScaleA,
class EngineScaleB,
class LayoutScaleB>
CUTLASS_DEVICE
void scale_residue_if_needed(const cute::Tensor<EngineScaleA, LayoutScaleA> &scaleA, const cute::Tensor<EngineScaleB, LayoutScaleB> &scaleB) {
if (__shfl_sync(0xffffffff, mma_count_ > 0, 0)) {
scale_core(scaleA, scaleB);
}
}
};
} // namespace cutlass::gemm::collective
@@ -30,8 +30,6 @@
**************************************************************************************************/
#pragma once
#include "cutlass/cutlass.h"
@@ -288,23 +286,23 @@ struct CollectiveMma<
using TensorStorage = typename SharedStorage::TensorStorage;
using PipelineStorage = typename SharedStorage::PipelineStorage;
// Only one thread issues the TMA and updates the barriers in a 2SM MMA, adjust bytes accordingly
static constexpr uint32_t SFTransactionBytes =
cutlass::bits_to_bytes(size(AtomThrShapeMNK{}) * cosize(take<0,3>(SmemLayoutSFA{})) * cute::sizeof_bits_v<ElementSF>) +
cutlass::bits_to_bytes(size(AtomThrShapeMNK{}) * cosize(take<0,3>(SmemLayoutSFB{})) * cute::sizeof_bits_v<ElementSF>);
// Only one thread issues the TMA and updates the barriers in a 2SM MMA, adjust bytes accordingly
static constexpr uint32_t ABTmaTransactionBytes =
cutlass::bits_to_bytes(size(AtomThrShapeMNK{}) * cosize(take<0,3>(SmemLayoutA{})) * cute::sizeof_bits_v<ElementA>) +
cutlass::bits_to_bytes(size(AtomThrShapeMNK{}) * cosize(take<0,3>(SmemLayoutB{})) * cute::sizeof_bits_v<ElementB>);
static constexpr uint32_t TmaTransactionBytes = ABTmaTransactionBytes + SFTransactionBytes;
template<class AccTensor, class SfaTensor, class SfbTensor>
template <class AccTensor, class SfaTensor, class SfbTensor>
struct TmemStorage {
AccTensor accumulators;
SfaTensor tCtSFA;
SfbTensor tCtSFB;
};
template<
template <
class KTileCount,
class GTensorPartitionedA, class GTensorPartitionedB,
class STensorA, class STensorB,
@@ -348,7 +346,8 @@ struct CollectiveMma<
, mcast_mask_sfa(mcast_mask_sfa_), mcast_mask_sfb(mcast_mask_sfb_) {}
};
template<
template <
class TiledMma,
class FragmentA, class FragmentB,
class FragmentSFA, class FragmentSFB,
class SFATiledCopy, class SmemFrgSFA, class TmemFrgSFA,
@@ -496,6 +495,7 @@ struct CollectiveMma<
Tensor tensor_a = make_tensor(ptr_A, make_layout(make_shape(M,K,L), args.dA));
Tensor tensor_b = make_tensor(ptr_B, make_layout(make_shape(N,K,L), args.dB));
auto cluster_shape = cutlass::detail::select_cluster_shape(ClusterShape{}, hw_info.cluster_shape);
// Cluster layout for TMA construction
auto cluster_layout_vmnk = tiled_divide(make_layout(cluster_shape), make_tile(typename TiledMma::AtomThrID{}));
auto cluster_shape_fallback = cutlass::detail::select_cluster_shape(ClusterShape{}, hw_info.cluster_shape_fallback);
@@ -505,7 +505,7 @@ struct CollectiveMma<
// Cluster layout for TMA construction of SFB
auto cluster_layout_sfb_vmnk = tiled_divide(make_layout(cluster_shape), make_tile(typename TiledMMA_SF::AtomThrID{}));
auto cluster_layout_sfb_vmnk_fallback = tiled_divide(make_layout(cluster_shape_fallback), make_tile(typename TiledMMA_SF::AtomThrID{}));
auto cluster_layout_sfb_vmnk_fallback = tiled_divide(make_layout(cluster_shape_fallback), make_tile(typename TiledMMA_SF::AtomThrID{}));
typename Params::TMA_A tma_load_a = make_tma_atom_A_sm100<TmaInternalElementA>(
GmemTiledCopyA{},
@@ -649,7 +649,7 @@ struct CollectiveMma<
return cute::make_tuple(tmem_storage.accumulators(_,_,_,stage));
}
template<class EpilogueTile, bool IsOverlappingAccum = false>
template <class EpilogueTile, bool IsOverlappingAccum = false>
CUTLASS_DEVICE static
auto
init_tmem_tensors(EpilogueTile epi_tile) {
@@ -660,7 +660,7 @@ struct CollectiveMma<
tiled_mma, acc_shape, EpilogueTile{});
Tensor tCtSFA = make_tensor<typename TiledMma::FrgTypeSFA>(shape(SmemLayoutAtomSFA{}));
Tensor tCtSFB = make_tensor<typename TiledMma::FrgTypeSFB>(shape(SmemLayoutAtomSFB{}));
TmemStorage<decltype(accumulators), decltype(tCtSFA), decltype(tCtSFB)> tmem_storage;
tmem_storage.accumulators = accumulators;
tmem_storage.tCtSFA = tCtSFA;
@@ -669,10 +669,10 @@ struct CollectiveMma<
return tmem_storage;
}
template<class AccTensor, class SfaTensor, class SfbTensor>
template <class TmemStorage>
CUTLASS_DEVICE static
void
set_tmem_offsets(TmemStorage<AccTensor, SfaTensor, SfbTensor>& tmem_storage, uint32_t tmem_base_addr) {
set_tmem_offsets(TmemStorage& tmem_storage, uint32_t tmem_base_addr) {
tmem_storage.accumulators.data() = tmem_base_addr;
tmem_storage.tCtSFA.data() = tmem_storage.accumulators.data().get() + cutlass::detail::find_tmem_tensor_col_offset(tmem_storage.accumulators);
tmem_storage.tCtSFB.data() = tmem_storage.tCtSFA.data().get() + cutlass::detail::find_tmem_tensor_col_offset(tmem_storage.tCtSFA);
@@ -751,7 +751,6 @@ struct CollectiveMma<
Tensor sSFB = make_tensor(make_smem_ptr(shared_tensors.smem_SFB.begin()), SmemLayoutSFB{});
// Define the CTA-in-cluster Layout and Coord
Layout cta_layout_mnk = make_layout(cluster_shape_);
Layout cta_layout_vmnk = tiled_divide(cta_layout_mnk, make_tile(typename TiledMma::AtomThrID{}));
auto cta_coord_vmnk = cta_layout_vmnk.get_flat_coord(block_rank_in_cluster_);
@@ -785,13 +784,11 @@ struct CollectiveMma<
uint16_t mcast_mask_sfa = create_tma_multicast_mask<2>(cta_layout_vmnk, cta_coord_vmnk);
uint16_t mcast_mask_sfb = create_tma_multicast_mask<1>(cta_layout_sfb_vmnk, cta_coord_sfb_vmnk);
LoadParams load_params {
return LoadParams{
size<3>(gA_mkl), // for scheduler
tAgA_mkl, tBgB_nkl, tAsA, tBsB, // for input tensor values
tAgSFA_mkl, tBgSFB_nkl, tAsSFA, tBsSFB, // for input scale factor tensor values
mcast_mask_a, mcast_mask_b, mcast_mask_sfa, mcast_mask_sfb // multicast masks
};
return load_params;
mcast_mask_a, mcast_mask_b, mcast_mask_sfa, mcast_mask_sfb}; // multicast masks
}
/// Set up the data needed by this collective for mma compute.
@@ -802,8 +799,8 @@ struct CollectiveMma<
TensorStorage& shared_tensors) const {
// Allocate "fragments/descriptors" for A and B matrices
Tensor sA = make_tensor(make_smem_ptr(shared_tensors.smem_A.begin()), SmemLayoutA{}); // (BLK_M,BLK_K,PIPE)
Tensor sB = make_tensor(make_smem_ptr(shared_tensors.smem_B.begin()), SmemLayoutB{}); // (BLK_N,BLK_K,PIPE)
Tensor sA = make_tensor(make_smem_ptr(shared_tensors.smem_A.begin()), SmemLayoutA{}); // (BLK_M,BLK_K,PIPE)
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)
@@ -854,17 +851,12 @@ struct CollectiveMma<
tiled_mma.idesc_.a_format_ = uint8_t(runtime_data_type_a_) & 0b111;
tiled_mma.idesc_.b_format_ = uint8_t(runtime_data_type_b_) & 0b111;
}
MmaParams<
decltype(tCrA), decltype(tCrB), decltype(tCtSFA), decltype(tCtSFB),
decltype(tiled_copy_s2t_SFA), decltype(thr_tCsSFA_compact_s2t), decltype(thr_tCtSFA_compact_s2t),
decltype(tiled_copy_s2t_SFB), decltype(thr_tCsSFB_compact_s2t), decltype(thr_tCtSFB_compact_s2t)
> mma_params {
return MmaParams{
tiled_mma,
tCrA, tCrB, tCtSFA, tCtSFB,
tiled_copy_s2t_SFA, thr_tCsSFA_compact_s2t, thr_tCtSFA_compact_s2t,
tiled_copy_s2t_SFB, thr_tCsSFB_compact_s2t, thr_tCtSFB_compact_s2t
};
return mma_params;
tiled_copy_s2t_SFB, thr_tCsSFB_compact_s2t, thr_tCtSFB_compact_s2t};
}
/// Perform a collective-scoped matrix multiply-accumulate
@@ -983,52 +975,12 @@ 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 (k_tile_count > 0) { // first iteraion
// 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);
}
if constexpr (IsOverlappingAccum) {
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);
}
CUTLASS_PRAGMA_NO_UNROLL
while (k_tile_count > 0) {
// WAIT on mainloop_pipe_consumer_state until its data are available
@@ -1052,6 +1004,13 @@ struct CollectiveMma<
copy(tiled_copy_s2t_SFB, thr_tCsSFB_s2t(_,_,_,_,read_stage), thr_tCtSFB_s2t);
}
if constexpr (IsOverlappingAccum) {
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) {
@@ -1064,6 +1023,7 @@ struct CollectiveMma<
accumulators);
tiled_mma.accumulate_ = UMMA::ScaleOut::One;
}
mainloop_pipeline.consumer_release(curr_mainloop_pipe_consumer_state);
}
@@ -31,7 +31,6 @@
#pragma once
#include "cutlass/cutlass.h"
@@ -239,12 +238,12 @@ struct CollectiveMma<
cutlass::bits_to_bytes(size(AtomThrShapeMNK{}) * cosize(take<0,3>(SmemLayoutA{})) * cute::sizeof_bits_v<ElementA>) +
cutlass::bits_to_bytes(size(AtomThrShapeMNK{}) * cosize(take<0,3>(SmemLayoutB{})) * cute::sizeof_bits_v<ElementB>);
template<class AccTensor>
template <class AccTensor>
struct TmemStorage {
AccTensor accumulators;
};
template<
template <
class KTileCount,
class GTensorPartitionedA, class GTensorPartitionedB,
class STensorA, class STensorB
@@ -273,7 +272,10 @@ struct CollectiveMma<
, mcast_mask_a(mcast_mask_a_), mcast_mask_b(mcast_mask_b_) {}
};
template<class FragmentA, class FragmentB>
template <
class TiledMma,
class FragmentA, class FragmentB
>
struct MmaParams {
TiledMma tiled_mma;
FragmentA tCrA;
@@ -336,7 +338,7 @@ struct CollectiveMma<
, runtime_data_type_a_(params.runtime_data_type_a)
, runtime_data_type_b_(params.runtime_data_type_b) {
if constexpr (IsDynamicCluster) {
const bool is_fallback_cluster = (cute::size<0>(cluster_shape_) == params.cluster_shape_fallback.x &&
const bool is_fallback_cluster = (cute::size<0>(cluster_shape_) == params.cluster_shape_fallback.x &&
cute::size<1>(cluster_shape_) == params.cluster_shape_fallback.y);
observed_tma_load_a_ = is_fallback_cluster ? &params.tma_load_a_fallback : &params.tma_load_a;
observed_tma_load_b_ = is_fallback_cluster ? &params.tma_load_b_fallback : &params.tma_load_b;
@@ -461,7 +463,7 @@ struct CollectiveMma<
return cute::make_tuple(tmem_storage.accumulators(_,_,_,stage));
}
template<class EpilogueTile, bool IsOverlappingAccum = false>
template <class EpilogueTile, bool IsOverlappingAccum = false>
CUTLASS_DEVICE static
auto
init_tmem_tensors(EpilogueTile epi_tile) {
@@ -475,10 +477,10 @@ struct CollectiveMma<
return tmem_storage;
}
template<class AccTensor>
template <class TmemStorage>
CUTLASS_DEVICE static
void
set_tmem_offsets(TmemStorage<AccTensor>& tmem_storage, uint32_t tmem_base_addr) {
set_tmem_offsets(TmemStorage& tmem_storage, uint32_t tmem_base_addr) {
tmem_storage.accumulators.data() = tmem_base_addr;
}
@@ -535,21 +537,21 @@ struct CollectiveMma<
// TMA Multicast Masks
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);
LoadParams load_params {
return LoadParams{
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;
mcast_mask_a, mcast_mask_b}; // multicast masks
}
/// Set up the data needed by this collective for mma compute.
template <class AccTensor>
template <class TmemStorage>
CUTLASS_DEVICE auto
mma_init(
[[maybe_unused]] TmemStorage<AccTensor> tmem_tensors,
TensorStorage& shared_tensors) const {
[[maybe_unused]] TmemStorage tmem_storage,
TensorStorage& shared_tensors) const {
// Allocate "fragments/descriptors" for A and B matrices
Tensor sA = make_tensor(make_smem_ptr(shared_tensors.smem_A.begin()), SmemLayoutA{}); // (BLK_M,BLK_K,PIPE)
Tensor sB = make_tensor(make_smem_ptr(shared_tensors.smem_B.begin()), SmemLayoutB{}); // (BLK_N,BLK_K,PIPE)
@@ -558,7 +560,7 @@ struct CollectiveMma<
Tensor tCrB = TiledMma::make_fragment_B(sB); // (MMA,MMA_N,MMA_K,PIPE)
CUTE_STATIC_ASSERT_V(Int<DispatchPolicy::Stages>{} == size<3>(sA)); // PIPE
CUTE_STATIC_ASSERT_V(Int<DispatchPolicy::Stages>{} == size<3>(sB));
CUTE_STATIC_ASSERT_V(Int<DispatchPolicy::Stages>{} == size<3>(sB)); // PIPE
TiledMma tiled_mma;
@@ -568,11 +570,10 @@ struct CollectiveMma<
tiled_mma.idesc_.a_format_ = uint8_t(runtime_data_type_a_) & 0b111;
tiled_mma.idesc_.b_format_ = uint8_t(runtime_data_type_b_) & 0b111;
}
MmaParams<decltype(tCrA), decltype(tCrB)> mma_params {
return MmaParams{
tiled_mma,
tCrA, tCrB
};
return mma_params;
tCrA, tCrB};
}
/// Perform a collective-scoped matrix multiply-accumulate
@@ -657,6 +658,7 @@ struct CollectiveMma<
) {
static_assert(is_tmem<FrgEngine>::value, "Accumulator must be tmem resident.");
static_assert(rank(FrgLayout{}) == 3, "Accumulator must be MMA-partitioned: (MMA, MMA_M, MMA_N)");
auto accumulators = get<0>(accumulators_pair);
auto [tiled_mma, tCrA, tCrB] = mma_inputs;
@@ -58,6 +58,7 @@ template <
class ClusterShape,
class KernelSchedule,
int ScaleGranularityM_,
int ScaleGranularityN_,
class TileShape_,
class ElementA_,
class StrideA_,
@@ -73,7 +74,7 @@ template <
class SmemCopyAtomB_,
class TransformB_>
struct CollectiveMma<
MainloopSm90TmaGmmaWarpSpecializedBlockScalingFP8<Stages, ClusterShape, KernelSchedule, ScaleGranularityM_>,
MainloopSm90TmaGmmaWarpSpecializedBlockScalingFP8<Stages, ClusterShape, KernelSchedule, ScaleGranularityM_, ScaleGranularityN_>,
TileShape_,
ElementA_,
StrideA_,
@@ -92,7 +93,7 @@ struct CollectiveMma<
//
// Type Aliases
//
using DispatchPolicy = MainloopSm90TmaGmmaWarpSpecializedBlockScalingFP8<Stages, ClusterShape, KernelSchedule, ScaleGranularityM_>;
using DispatchPolicy = MainloopSm90TmaGmmaWarpSpecializedBlockScalingFP8<Stages, ClusterShape, KernelSchedule, ScaleGranularityM_, ScaleGranularityN_>;
using TileShape = TileShape_;
using ElementA = ElementA_;
using StrideA = StrideA_;
@@ -120,7 +121,9 @@ struct CollectiveMma<
static constexpr int NumProducerThreadEvents = 2;
static constexpr int ScaleGranularityM = ScaleGranularityM_ == 0 ? size<0>(TileShape{}) : ScaleGranularityM_;
static constexpr int ScaleGranularityN = ScaleGranularityN_ == 0 ? size<1>(TileShape{}) : ScaleGranularityN_;
static constexpr int ScaleMsPerTile = size<0>(TileShape{}) / ScaleGranularityM;
static constexpr int ScaleNsPerTile = size<1>(TileShape{}) / ScaleGranularityN;
static_assert(cute::rank(SmemLayoutAtomA{}) == 2, "SmemLayoutAtom must be rank 2 (M/N, K)");
static_assert((size<0>(TileShape{}) % size<0>(SmemLayoutAtomA{})) == 0, "SmemLayoutAtom must evenly divide tile shape.");
@@ -131,6 +134,7 @@ struct CollectiveMma<
static_assert((size<2>(TileShape{}) % size<1>(SmemLayoutAtomB{})) == 0, "SmemLayoutAtom must evenly divide tile shape.");
static_assert((size<0>(TileShape{}) % ScaleGranularityM) == 0, "FP8 scaling granularity must evenly divide tile shape along M.");
static_assert((size<1>(TileShape{}) % ScaleGranularityN) == 0, "FP8 scaling granularity must evenly divide tile shape along N.");
// Tile along modes in a way that maximizes the TMA box size.
using SmemLayoutA = decltype(tile_to_shape(
@@ -144,12 +148,13 @@ struct CollectiveMma<
// Block scaling gmem-to-smem copy atom
using BlockScaleCopyTypeA = cute::uint_byte_t<cute::min(static_cast<int>(sizeof(ElementBlockScale)) * ScaleMsPerTile, 16)>;
using BlockScaleCopyTypeB = cute::uint_byte_t<cute::min(static_cast<int>(sizeof(ElementBlockScale)) * ScaleNsPerTile, 16)>;
using SmemBlockScalingCopyAtomA = Copy_Atom<SM80_CP_ASYNC_CACHEALWAYS<BlockScaleCopyTypeA>, ElementBlockScale>;
using SmemBlockScalingCopyAtomB = Copy_Atom<SM80_CP_ASYNC_CACHEALWAYS<ElementBlockScale>, ElementBlockScale>;
using SmemBlockScalingCopyAtomB = Copy_Atom<SM80_CP_ASYNC_CACHEALWAYS<BlockScaleCopyTypeB>, ElementBlockScale>;
// Block scaling smem layout
using SmemLayoutScaleA = Layout<Shape<Int<ScaleMsPerTile>, Int<DispatchPolicy::Stages>>>;
using SmemLayoutScaleB = Layout<Shape<Int<DispatchPolicy::Stages>>, Stride<_1>>; // `ScaleNsPerTile` is always 1.
using SmemLayoutScaleB = Layout<Shape<Int<ScaleNsPerTile>, Int<DispatchPolicy::Stages>>>;
static_assert(DispatchPolicy::Stages >= 2, "Specialization requires Stages set to value 1 or more.");
static_assert(cute::is_base_of<cute::GMMA::DescriptorIterator, typename TiledMma::FrgTypeA>::value &&
@@ -168,7 +173,7 @@ struct CollectiveMma<
cute::array_aligned<typename TiledMma::ValTypeA, cute::cosize_v<SmemLayoutA>> smem_A; // mxk
cute::array_aligned<typename TiledMma::ValTypeB, cute::cosize_v<SmemLayoutB>> smem_B; // nxk
cute::array_aligned<ElementBlockScale, cute::cosize_v<SmemLayoutScaleA>> smem_scale_A; // ScaleMsPerTile x k
cute::array_aligned<ElementBlockScale, cute::cosize_v<SmemLayoutScaleB>> smem_scale_B; // 1xk
cute::array_aligned<ElementBlockScale, cute::cosize_v<SmemLayoutScaleB>> smem_scale_B; // ScaleNsPerTile x k
} tensors;
using PipelineStorage = typename MainloopPipeline::SharedStorage;
@@ -322,17 +327,17 @@ struct CollectiveMma<
Tensor gB_nkl = local_tile(mB_nkl, TileShape{}, make_coord(_,_,_), Step< X,_1,_1>{}); // (BLK_N,BLK_K,n,k,l)
// Make the tiled views of scale tensors
auto scaleA_shape = make_shape(get<2>(gA_mkl.shape()), Int<ScaleMsPerTile>{}, get<3>(gA_mkl.shape()), get<4>(gA_mkl.shape())); // (m,ScaleMsPerTile,k,l)
auto scale_dA = make_stride(get<3>(gA_mkl.shape()) * Int<ScaleMsPerTile>{}, Int<1>{}, Int<ScaleMsPerTile>{}, get<2>(gA_mkl.shape()) * get<3>(gA_mkl.shape()) * Int<ScaleMsPerTile>{});
auto scaleA_shape = make_shape(shape<2>(gA_mkl), Int<ScaleMsPerTile>{}, shape<3>(gA_mkl), shape<4>(gA_mkl)); // (m,ScaleMsPerTile,k,l)
auto scaleB_shape = make_shape(shape<2>(gB_nkl), Int<ScaleNsPerTile>{}, shape<3>(gB_nkl), shape<4>(gB_nkl)); // (n,ScaleNsPerTile,k,l)
auto scale_dA = compact_order(scaleA_shape, Step<_2,_0,_1,_3>{});
auto scale_dB = compact_order(scaleB_shape, Step<_2,_0,_1,_3>{});
auto scaleA_layout = make_layout(scaleA_shape, scale_dA);
auto scaleB_shape = make_shape(get<2>(gB_nkl.shape()), get<3>(gB_nkl.shape()), get<4>(gB_nkl.shape())); // (n,k,l)
auto scale_dB = make_stride(get<3>(gB_nkl.shape()), Int<1>{}, get<2>(gB_nkl.shape()) * get<3>(gB_nkl.shape()));
auto scaleB_layout = make_layout(scaleB_shape, scale_dB);
// Note that mScaleA_mkl and mScaleB_nkl are already blocked tiled in the `m` host and
// Note that mScaleA_mkl and mScaleB_nkl are already blocked tiled in the `m` host and
// gScaleA_mkl and gScaleB_nkl in `g` global memory are same as mScaleA_mkl and mScaleB_nkl.
Tensor mScaleA_mkl = make_tensor(make_gmem_ptr(mainloop_params.ptr_scale_A), scaleA_layout); // (m,ScaleMsPerTile,k,l)
Tensor mScaleB_nkl = make_tensor(make_gmem_ptr(mainloop_params.ptr_scale_B), scaleB_layout); // (n,k,l)
Tensor mScaleB_nkl = make_tensor(make_gmem_ptr(mainloop_params.ptr_scale_B), scaleB_layout); // (n,ScaleNsPerTile,k,l)
return cute::make_tuple(gA_mkl, gB_nkl, mScaleA_mkl, mScaleB_nkl);
}
@@ -356,13 +361,13 @@ struct CollectiveMma<
uint32_t block_rank_in_cluster,
TensorStorage& shared_tensors) {
int lane_predicate = cute::elect_one_sync();
// Blockscaling: Tma loads for load_input and CpAsync for load_scale
if (lane_predicate) {
Tensor sA = make_tensor(make_smem_ptr(shared_tensors.smem_A.data()), SmemLayoutA{}); // (BLK_M,BLK_K,PIPE)
Tensor sB = make_tensor(make_smem_ptr(shared_tensors.smem_B.data()), SmemLayoutB{}); // (BLK_N,BLK_K,PIPE)
Tensor sScaleA = make_tensor(cute::make_smem_ptr(shared_tensors.smem_scale_A.data()), SmemLayoutScaleA{}); // (ScaleMsPerTile,k)
Tensor sScaleB = make_tensor(cute::make_smem_ptr(shared_tensors.smem_scale_B.data()), SmemLayoutScaleB{}); // (k)
Tensor sScaleB = make_tensor(cute::make_smem_ptr(shared_tensors.smem_scale_B.data()), SmemLayoutScaleB{}); // (ScaleNsPerTile,k)
//
// Prepare the TMA loads for A and B
@@ -388,10 +393,10 @@ struct CollectiveMma<
Tensor mScaleB_nkl = get<3>(load_inputs);
Tensor gScaleA = mScaleA_mkl(m_coord,_,_,l_coord); // (1,ScaleMsPerTile,k,1)
Tensor gScaleB = mScaleB_nkl(n_coord,_,l_coord); // (1,k,1)
Tensor gScaleB = mScaleB_nkl(n_coord,_,_,l_coord); // (1,ScaleNsPerTile,k,1)
TiledCopy scale_copy_a = make_tiled_copy(SmemBlockScalingCopyAtomA{}, Layout<Shape<_1>>{}, Layout<Shape<Int<ScaleMsPerTile>>>{}); // (1,ScaleMsPerTile,1)
TiledCopy scale_copy_b = make_tiled_copy(SmemBlockScalingCopyAtomB{}, Layout<Shape<_1>>{}, Layout<Shape<_1>>{}); // (1,1,1)
TiledCopy scale_copy_b = make_tiled_copy(SmemBlockScalingCopyAtomB{}, Layout<Shape<_1>>{}, Layout<Shape<Int<ScaleNsPerTile>>>{}); // (1,ScaleNsPerTile,1)
ThrCopy thr_scale_copy_a = scale_copy_a.get_slice(threadIdx.x);
ThrCopy thr_scale_copy_b = scale_copy_b.get_slice(threadIdx.x);
@@ -446,7 +451,7 @@ struct CollectiveMma<
// Copy scale tensors from global memory to shared memory
copy(scale_copy_a, tAgA_ScaleA(_,_,*k_tile_iter), tAsA_ScaleA(_,_,write_stage));
copy(scale_copy_b, tBgB_ScaleB(_,*k_tile_iter), tBsB_ScaleB(_,write_stage));
copy(scale_copy_b, tBgB_ScaleB(_,_,*k_tile_iter), tBsB_ScaleB(_,_,write_stage));
pipeline.producer_commit(smem_pipe_write, cutlass::arch::cpasync_barrier_arrive_noinc);
++k_tile_iter;
@@ -508,7 +513,11 @@ struct CollectiveMma<
Shape<Shape<Int<ScaleGranularityM>, Int<ScaleMsPerTile>>, cute::tuple_element_t<1, TileShape>, Int<DispatchPolicy::Stages>>,
Stride<Stride<_0, _1>, _0, Int<ScaleMsPerTile>>
>{}); // ((ScaleGranularityM,ScaleMsPerTile),n,k)
Tensor sScaleB = make_tensor(cute::make_smem_ptr(shared_tensors.smem_scale_B.data()), SmemLayoutScaleB{}); // (k)
Tensor sScaleBViewAsC = make_tensor(cute::make_smem_ptr(shared_tensors.smem_scale_B.data()),
Layout<
Shape<cute::tuple_element_t<0, TileShape>, Shape<Int<ScaleGranularityN>, Int<ScaleNsPerTile>>, Int<DispatchPolicy::Stages>>,
Stride<_0, Stride<_0, _1>, Int<ScaleNsPerTile>>
>{}); // (m,(ScaleGranularityN,ScaleNsPerTile),k)
//
// Define C accumulators and A/B partitioning
@@ -531,7 +540,8 @@ struct CollectiveMma<
TiledMma tiled_mma;
auto thread_mma = tiled_mma.get_slice(warp_group_thread_layout(warp_group_idx));
Tensor tCsScaleAViewAsC = tiled_mma.get_slice(thread_idx).partition_C(sScaleAViewAsC); // (MMA,MMA_M,MMA_N,PIPE), `thread_mma` above is correct when partitioning A and B, but it is not correct when partitioning C.
Tensor tCsScaleAViewAsC = tiled_mma.get_slice(thread_idx).partition_C(sScaleAViewAsC); // (MMA,MMA_M,MMA_N,PIPE), `thread_mma` above is correct when partitioning A and B, but it is not correct when partitioning C.
Tensor tCsScaleBViewAsC = tiled_mma.get_slice(thread_idx).partition_C(sScaleBViewAsC); // (MMA,MMA_M,MMA_N,PIPE), `thread_mma` above is correct when partitioning A and B, but it is not correct when partitioning C.
Tensor tCsA = thread_mma.partition_A(sA); // (MMA,MMA_M,MMA_K,PIPE)
Tensor tCsB = thread_mma.partition_B(sB); // (MMA,MMA_N,MMA_K,PIPE)
@@ -557,11 +567,8 @@ struct CollectiveMma<
PipelineState smem_pipe_release = smem_pipe_read;
// Per block scale values for operand A and B
using RegLayoutScaleAViewAsC = decltype(make_layout_like(tCsScaleAViewAsC(_, _, _, 0).layout())); // `make_layout_like` makes a compact layout.
using RegLayoutScaleAEssential = decltype(filter_zeros(RegLayoutScaleAViewAsC{}.stride(), RegLayoutScaleAViewAsC{}.shape())); // an interface to traverse the underlying storage for the compact layout mentioned above
Tensor tCrScaleAViewAsC = make_tensor<ElementBlockScale>(RegLayoutScaleAViewAsC{}); // (MMA,MMA_M,MMA_N)
ElementBlockScale scale_b;
Tensor tCrScaleAViewAsC = make_tensor_like<ElementBlockScale>(tCsScaleAViewAsC(_, _, _, 0)); // (MMA,MMA_M,MMA_N)
Tensor tCrScaleBViewAsC = make_tensor_like<ElementBlockScale>(tCsScaleBViewAsC(_, _, _, 0)); // (MMA,MMA_M,MMA_N)
// Prologue GMMAs
int prologue_mma_count = min(K_PIPE_MMAS, k_tile_count);
@@ -583,21 +590,26 @@ struct CollectiveMma<
int read_stage = smem_pipe_read.index();
// Load per block scale values from shared memory to registers.
scale_b = sScaleB[read_stage];
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < size(RegLayoutScaleAEssential{}); i++) {
tCrScaleAViewAsC.data()[i] = tCsScaleAViewAsC(_, _, _, read_stage)(idx2crd(i, RegLayoutScaleAEssential{}));
// Load per block scale values from shared memory to registers
copy(tCsScaleAViewAsC(_, _, _, read_stage), tCrScaleAViewAsC);
copy(tCsScaleBViewAsC(_, _, _, read_stage), tCrScaleBViewAsC);
if constexpr (ScaleMsPerTile == 1 && ScaleNsPerTile == 1) {
tCrScaleAViewAsC.data()[0] = tCrScaleAViewAsC.data()[0] * tCrScaleBViewAsC.data()[0];
}
if constexpr (ScaleMsPerTile == 1) {
static_assert(size(RegLayoutScaleAEssential{}) == 1);
tCrScaleAViewAsC.data()[0] = __shfl_sync(0xffffffff, tCrScaleAViewAsC.data()[0] * scale_b, 0); // `tCrScaleAViewAsC.data()[0]` are all same in a warp group when `ScaleMsPerTile == 1`.
} else {
if constexpr (ScaleMsPerTile > 1 && ScaleNsPerTile == 1) {
ElementBlockScale scale_b = tCrScaleBViewAsC.data()[0];
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < size(RegLayoutScaleAEssential{}); i++) {
for (int i = 0; i < size(tCrScaleAViewAsC); i++) {
tCrScaleAViewAsC.data()[i] = tCrScaleAViewAsC.data()[i] * scale_b;
}
}
if constexpr (ScaleMsPerTile == 1 && ScaleNsPerTile > 1) {
ElementBlockScale scale_a = tCrScaleAViewAsC.data()[0];
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < size(tCrScaleBViewAsC); i++) {
tCrScaleBViewAsC.data()[i] = tCrScaleBViewAsC.data()[i] * scale_a;
}
}
warpgroup_arrive();
// Unroll the K mode manually to set scale D to 1
@@ -609,8 +621,20 @@ struct CollectiveMma<
}
warpgroup_commit_batch();
// Block scale the accumulators with reg tensor `tCrScaleAViewAsC`
accumulation.scale_if_needed(tCrScaleAViewAsC);
// Block scale the accumulators with reg tensor `tCrScaleAViewAsC` and `tCrScaleBViewAsC`
if constexpr (ScaleMsPerTile == 1 && ScaleNsPerTile == 1) {
ElementBlockScale scale_ab = tCrScaleAViewAsC.data()[0];
accumulation.scale_if_needed(scale_ab);
}
if constexpr (ScaleMsPerTile > 1 && ScaleNsPerTile == 1) {
accumulation.scale_if_needed(tCrScaleAViewAsC);
}
if constexpr (ScaleMsPerTile == 1 && ScaleNsPerTile > 1) {
accumulation.scale_if_needed(tCrScaleBViewAsC);
}
if constexpr (ScaleMsPerTile > 1 && ScaleNsPerTile > 1) {
accumulation.scale_if_needed(tCrScaleAViewAsC, tCrScaleBViewAsC);
}
++smem_pipe_read;
}
@@ -632,21 +656,26 @@ struct CollectiveMma<
int read_stage = smem_pipe_read.index();
// Load per block scale values from shared memory to registers (at most twice per block along M and exactly once per block along N)
scale_b = sScaleB[read_stage];
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < size(RegLayoutScaleAEssential{}); i++) {
tCrScaleAViewAsC.data()[i] = tCsScaleAViewAsC(_, _, _, read_stage)(idx2crd(i, RegLayoutScaleAEssential{}));
// Load per block scale values from shared memory to registers (at most twice per block along M and/or N)
copy(tCsScaleAViewAsC(_, _, _, read_stage), tCrScaleAViewAsC);
copy(tCsScaleBViewAsC(_, _, _, read_stage), tCrScaleBViewAsC);
if constexpr (ScaleMsPerTile == 1 && ScaleNsPerTile == 1) {
tCrScaleAViewAsC.data()[0] = tCrScaleAViewAsC.data()[0] * tCrScaleBViewAsC.data()[0];
}
if constexpr (ScaleMsPerTile == 1) {
static_assert(size(RegLayoutScaleAEssential{}) == 1);
tCrScaleAViewAsC.data()[0] = __shfl_sync(0xffffffff, tCrScaleAViewAsC.data()[0] * scale_b, 0); // `tCrScaleAViewAsC.data()[0]` are all same in a warp group when `ScaleMsPerTile == 1`.
} else {
if constexpr (ScaleMsPerTile > 1 && ScaleNsPerTile == 1) {
ElementBlockScale scale_b = tCrScaleBViewAsC.data()[0];
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < size(RegLayoutScaleAEssential{}); i++) {
for (int i = 0; i < size(tCrScaleAViewAsC); i++) {
tCrScaleAViewAsC.data()[i] = tCrScaleAViewAsC.data()[i] * scale_b;
}
}
if constexpr (ScaleMsPerTile == 1 && ScaleNsPerTile > 1) {
ElementBlockScale scale_a = tCrScaleAViewAsC.data()[0];
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < size(tCrScaleBViewAsC); i++) {
tCrScaleBViewAsC.data()[i] = tCrScaleBViewAsC.data()[i] * scale_a;
}
}
if (accumulation.prepare_if_needed()) {
tiled_mma.accumulate_ = GMMA::ScaleOut::Zero;
@@ -667,8 +696,20 @@ struct CollectiveMma<
warpgroup_wait<K_PIPE_MMAS>();
warpgroup_fence_operand(accumulation());
// Block scale the accumulators with reg tensor `tCrScaleAViewAsC`
accumulation.scale_if_needed(tCrScaleAViewAsC);
// Block scale the accumulators with reg tensor `tCrScaleAViewAsC` and `tCrScaleBViewAsC`
if constexpr (ScaleMsPerTile == 1 && ScaleNsPerTile == 1) {
ElementBlockScale scale_ab = tCrScaleAViewAsC.data()[0];
accumulation.scale_if_needed(scale_ab);
}
if constexpr (ScaleMsPerTile > 1 && ScaleNsPerTile == 1) {
accumulation.scale_if_needed(tCrScaleAViewAsC);
}
if constexpr (ScaleMsPerTile == 1 && ScaleNsPerTile > 1) {
accumulation.scale_if_needed(tCrScaleBViewAsC);
}
if constexpr (ScaleMsPerTile > 1 && ScaleNsPerTile > 1) {
accumulation.scale_if_needed(tCrScaleAViewAsC, tCrScaleBViewAsC);
}
pipeline.consumer_release(smem_pipe_release); // UNLOCK smem_pipe_release, done _computing_ on it
@@ -677,7 +718,19 @@ struct CollectiveMma<
++smem_pipe_release;
}
accumulation.scale_residue_if_needed(tCrScaleAViewAsC);
if constexpr (ScaleMsPerTile == 1 && ScaleNsPerTile == 1) {
ElementBlockScale scale_ab = tCrScaleAViewAsC.data()[0];
accumulation.scale_residue_if_needed(scale_ab);
}
if constexpr (ScaleMsPerTile > 1 && ScaleNsPerTile == 1) {
accumulation.scale_residue_if_needed(tCrScaleAViewAsC);
}
if constexpr (ScaleMsPerTile == 1 && ScaleNsPerTile > 1) {
accumulation.scale_residue_if_needed(tCrScaleBViewAsC);
}
if constexpr (ScaleMsPerTile > 1 && ScaleNsPerTile > 1) {
accumulation.scale_residue_if_needed(tCrScaleAViewAsC, tCrScaleBViewAsC);
}
warpgroup_fence_operand(accumulation());
}
+11 -3
View File
@@ -117,7 +117,11 @@ struct KernelPtrArrayTmaWarpSpecializedPingpong { };
// FP8 related policies (including Blocked Scaled Accumulation)
template<
int ScaleGranularityM = 0 // `ScaleGranularityM` specifies scaling granularity along M, while zero-value `ScaleGranularityM` indicates that scaling granularity is `size<0>(TileShape_MNK{})` along M.
// `ScaleGranularityM`/`ScaleGranularityN` specifies scaling granularity along M/N, while zero-value
// `ScaleGranularityM`/`ScaleGranularityN` indicates that scaling granularity is
// `size<0>(TileShape_MNK{})`/`size<1>(TileShape_MNK{})` along M/N.
int ScaleGranularityM = 0,
int ScaleGranularityN = 0
>
struct KernelTmaWarpSpecializedCooperativeFP8BlockScaledAccum: KernelTmaWarpSpecializedCooperative { };
@@ -302,12 +306,16 @@ template<
int Stages_,
class ClusterShape_ = Shape<_1,_1,_1>,
class KernelSchedule = KernelTmaWarpSpecialized,
int ScaleGranularityM = 0 // `ScaleGranularityM` specifies scaling granularity along M, while zero-value `ScaleGranularityM` indicates that scaling granularity is `size<0>(TileShape_MNK{})` along M.
// `ScaleGranularityM`/`ScaleGranularityN` specifies scaling granularity along M/N, while zero-value
// `ScaleGranularityM`/`ScaleGranularityN` indicates that scaling granularity is
// `size<0>(TileShape_MNK{})`/`size<1>(TileShape_MNK{})` along M/N.
int ScaleGranularityM = 0,
int ScaleGranularityN = 0
>
struct MainloopSm90TmaGmmaWarpSpecializedBlockScalingFP8
: MainloopSm90TmaGmmaWarpSpecialized<Stages_, ClusterShape_, KernelSchedule> {
static_assert(
cute::is_same_v<KernelSchedule, KernelTmaWarpSpecializedCooperativeFP8BlockScaledAccum<ScaleGranularityM>>,
cute::is_same_v<KernelSchedule, KernelTmaWarpSpecializedCooperativeFP8BlockScaledAccum<ScaleGranularityM, ScaleGranularityN>>,
"KernelSchedule must be one of the warp specialized policies");
};
@@ -397,8 +397,6 @@ public:
// An example of an unneeded threadblock is one that is assigned to compute in the upper
// portion of a Rank2K kernel filled with mode kLower.
//
// TODO: Consider pushing these checks into ProblemVisitor to avoid spuriously
// returning from `next_tile()`.
//
// Early exit if threadblock is out of range
@@ -1131,6 +1131,10 @@ public:
}
}
else {
// Register reconfiguration
arch::warpgroup_reg_dealloc<GenericRegisterRequirement>();
}
}
};
@@ -29,8 +29,6 @@
*
**************************************************************************************************/
#pragma once
#include "cutlass/cutlass.h"
@@ -564,20 +562,21 @@ public:
// Sync deallocation status between MMA warps of peer CTAs
arch::ClusterBarrier& tmem_deallocation_result_barrier = shared_storage.pipelines.tmem_dealloc;
[[maybe_unused]] uint32_t dealloc_barrier_phase = 0;
if constexpr(!IsOverlappingAccum) {
if (WarpCategory::MMA == warp_category && has_mma_peer_cta && lane_predicate) {
tmem_deallocation_result_barrier.init(NumMMAThreads);
if (WarpCategory::MMA == warp_category) {
if constexpr(!IsOverlappingAccum) {
if (has_mma_peer_cta && lane_predicate) {
tmem_deallocation_result_barrier.init(NumMMAThreads);
}
}
else {
if (has_mma_peer_cta && lane_predicate) {
tmem_deallocation_result_barrier.init(NumEpilogueThreads*2);
}
else if (lane_predicate) {
tmem_deallocation_result_barrier.init(NumEpilogueThreads);
}
}
}
else {
if (WarpCategory::MMA == warp_category && has_mma_peer_cta && lane_predicate) {
tmem_deallocation_result_barrier.init(NumEpilogueThreads*2);
}
else if (WarpCategory::MMA == warp_category && lane_predicate) {
tmem_deallocation_result_barrier.init(NumEpilogueThreads);
}
}
// Initialize smem barrier for prologue throttling. Epilogue warps are stalled until the prologue finishes.
arch::ClusterBarrier& epilogue_throttle_barrier = shared_storage.pipelines.epilogue_throttle;
@@ -699,7 +698,6 @@ public:
epilogue_throttle_barrier.arrive();
if constexpr (IsSchedDynamicPersistent) {
// Whether a new CLC query must be performed.
// See comment below where this variable is updated for a description of
// why this variable is needed.
@@ -738,7 +736,6 @@ public:
work_tile_info = next_work_tile_info;
} while (work_tile_info.is_valid());
clc_pipeline.producer_tail(clc_pipe_producer_state);
}
}
@@ -963,7 +960,6 @@ public:
epi_load_pipe_consumer_state = load_state_next;
epi_store_pipe_producer_state = store_state_next;
accumulator_pipe_consumer_state = acc_state_next;
do_tail_store = true;
}
work_tile_info = next_work_tile_info;
@@ -1057,6 +1057,10 @@ public:
}
}
else {
// Register reconfiguration
arch::warpgroup_reg_dealloc<GenericRegisterRequirement>();
}
}
};
@@ -783,7 +783,6 @@ private:
int L_idx, Split_idx;
params_.sk_params_.divmod_splits_(L_idx, Split_idx, work_tile_info.L_idx);
// TODO: Modularize the SM90 scheduler to pull out and reuse this redundant code
int additional_k_tiles = 0;
int split_start_offset = params_.sk_params_.big_units_;
@@ -455,8 +455,9 @@ public:
auto blk_shape = TileShape{}; // (BLK_M,BLK_N,BLK_K)
TileScheduler scheduler{params.scheduler};
auto work_tile_info = scheduler.initial_work_tile_info(ClusterShape{});
// Declare work_tile_info, then define it in each of warps that use it.
typename TileScheduler::WorkTileInfo work_tile_info;
// In a warp specialized kernel, collectives expose data movement and compute operations separately
CollectiveMainloop collective_mainloop;
@@ -474,6 +475,7 @@ public:
cluster_wait_fn();
if (warp_group_role == WarpGroupRole::Producer) {
work_tile_info = scheduler.initial_work_tile_info(ClusterShape{});
cutlass::arch::warpgroup_reg_dealloc<LoadRegisterRequirement>();
// Mainloop Producer Warp
@@ -578,6 +580,7 @@ public:
} // Producer Warp Group End
else if (warp_group_role == WarpGroupRole::Consumer0 || warp_group_role == WarpGroupRole::Consumer1) {
work_tile_info = scheduler.initial_work_tile_info(ClusterShape{});
cutlass::arch::warpgroup_reg_alloc<MmaRegisterRequirement>();
CollectiveEpilogue collective_epilogue(params.epilogue, shared_storage.tensors.epilogue);
@@ -265,7 +265,7 @@ struct PersistentTileSchedulerSm90Params {
}
// In case the maximum number of clusters that could co-exist on the target device is
// already calculated using cudaOccupancyMaxActiveClusters
else if (max_active_clusters != 0) {
else if (max_active_clusters != 0 && max_active_clusters * cluster_size <= sm_count) {
if (raster_order == RasterOrder::AlongN) {
launch_grid.y = max_active_clusters * cluster_shape.n();
}
@@ -1204,6 +1204,7 @@ struct PersistentTileSchedulerSm90StreamKParams {
KernelHardwareInfo new_hw_info;
new_hw_info.device_id = hw_info.device_id;
new_hw_info.sm_count = hw_info.sm_count;
new_hw_info.max_active_clusters = hw_info.max_active_clusters;
if (new_hw_info.sm_count <= 0) {
CUTLASS_TRACE_HOST(" WARNING: Arguments do not include a valid SM count.\n"
" For optimal performance, populate the arguments KernelHardwareInfo struct with the SM count.");
@@ -1787,7 +1788,7 @@ struct PersistentTileSchedulerSm90GroupParams {
}
// In case the maximum number of clusters that could co-exist on the target device is
// already calculated using cudaOccupancyMaxActiveClusters
else if (max_active_clusters != 0) {
else if (max_active_clusters != 0 && max_active_clusters * cluster_size <= sm_count) {
if (raster_order == RasterOrder::AlongN) {
launch_grid.y = max_active_clusters * cluster_shape.n();
}
@@ -2499,6 +2500,7 @@ struct PersistentTileSchedulerSm100GroupParams {
bool is_static_cluster_shape = false) {
int const sm_count = hw_info.sm_count;
int const max_active_clusters = hw_info.max_active_clusters;
// Round up to nearest multiple of swizzle_size along each mode
auto log_swizzle_size = get_log_swizzle_size(problem_blocks.x, problem_blocks.y, max_swizzle_size);
@@ -2542,6 +2544,18 @@ struct PersistentTileSchedulerSm100GroupParams {
launch_grid.x = possibly_truncate(sm_count, problem_blocks_total);
}
}
// In case the maximum number of clusters that could co-exist on the target device is
// already calculated using cudaOccupancyMaxActiveClusters
else if (max_active_clusters != 0 && max_active_clusters * cluster_size <= sm_count) {
if (raster_order == RasterOrder::AlongN) {
launch_grid.y = max_active_clusters * cluster_shape.n();
}
else {
launch_grid.x = max_active_clusters * cluster_shape.m();
}
CUTLASS_TRACE_HOST("get_grid_shape(): Proposed GridDims by the scheduler using cudaOccupancyMaxActiveClusters = "
"(" << launch_grid.x << ", " << launch_grid.y << ", " << launch_grid.z << ")\n");
}
else {
constexpr int max_sm_per_gpc = 20;
int cta_per_device = get_max_cta_occupancy(max_sm_per_gpc, cluster_shape, sm_count);
@@ -142,7 +142,6 @@ class MmaVoltaTensorOpMultiplicandTileIterator<
"Shape of warp-level Mma must be divisible by operator shape.");
// Shape of one individual LDS.128
// TODO: 32 and 4 are hardcoded, 32-by-4 is logical shape
using LdsShape = layout::PitchLinearShape<
32,
4
@@ -458,7 +457,6 @@ class MmaVoltaTensorOpMultiplicandTileIterator<
"Shape of warp-level Mma must be divisible by operator shape.");
// Shape of one individual LDS
// TODO: remove hardcoded 32 and 4
using LdsShape = layout::PitchLinearShape<
32,
4
@@ -995,7 +995,6 @@ public:
CUTLASS_DEVICE
MmaTensorOpMultiplicandTileIterator &add_tile_offset_negative(TensorCoord const &tile_offset) {
// TODO: fix this if it becomes an issue during warp it reset
add_tile_offset(tile_offset);
return *this;