update 3.8 v2 (#2112)
* update 3.8 v2 * update 3.8 --------- Co-authored-by: yuzhai <yuzhai@nvidia.com>
This commit is contained in:
@@ -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 ? ¶ms.tma_load_a_fallback : ¶ms.tma_load_a;
|
||||
observed_tma_load_b_ = is_fallback_cluster ? ¶ms.tma_load_b_fallback : ¶ms.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;
|
||||
|
||||
|
||||
+101
-48
@@ -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());
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
Reference in New Issue
Block a user