CUTLASS 3.3.0 (#1167)

* Release 3.3.0

Adds support for mixed precision GEMMs On Hopper and Ampere
Adds support for < 16B aligned GEMMs on Hopper
Enhancements to EVT
Enhancements to Python interface
Enhancements to Sub-byte type handling in CuTe
Several other bug-fixes and performance improvements.

* minor doc update
This commit is contained in:
Pradeep Ramani
2023-11-02 11:09:05 -04:00
committed by GitHub
parent 922fb5108b
commit c008b4aea8
263 changed files with 16214 additions and 5008 deletions
@@ -98,13 +98,18 @@ sm90_get_epilogue_smem_swizzle_layout_atom() {
}
// Attempts to compute a reasonable epilogue tile based on block tile shape or allows the user to provide one.
template <class ElementD, class EpilogueTileType, class Schedule>
template <class ElementD, class EpilogueTileType, class Schedule, class TileShape_MNK>
constexpr auto
sm90_compute_tile_shape_or_override() {
if constexpr (cute::is_same_v<EpilogueTileType, EpilogueTileAuto>) {
if constexpr (detail::sm90_is_cooperative_v<Schedule>) {
return Shape<_128,_32>{};
if constexpr (size<0>(TileShape_MNK{}) >= 128) {
return Shape<_128,_32>{};
}
else {
return Shape<_64,_32>{};
}
}
else if constexpr (detail::sm90_is_warp_specialized_v<Schedule>) {
if constexpr (sizeof_bits_v<ElementD> == 8) {
@@ -191,13 +196,17 @@ struct CallbacksBuilder<
TileShape_MNK,
EpilogueTile_MN,
ElementAccumulator,
enable_if_t<FusionOp::IsAuxOutSupported>
enable_if_t<(FusionOp::IsAuxOutSupported ^ FusionOp::IsAuxInSupported) // only one aux tensor
&& not is_subbyte_v<typename FusionOp::ElementAux>>
> {
using GmemStrideTypeAux = gemm::TagToStrideC_t<typename FusionOp::GmemLayoutTagAux>;
using SmemLayoutAtomAux = decltype(detail::sm90_get_epilogue_smem_swizzle_layout_atom<
GmemStrideTypeAux, typename FusionOp::ElementAux, EpilogueTile_MN>());
using SmemCopyOpAux = decltype(detail::sm90_get_smem_store_op_for_accumulator<
using CopyOpR2S = decltype(detail::sm90_get_smem_store_op_for_accumulator<
GmemStrideTypeAux, typename FusionOp::ElementAux>());
using CopyOpS2R = decltype(detail::sm90_get_smem_load_op_for_source<
GmemStrideTypeAux, typename FusionOp::ElementAux>());
using SmemCopyOpAux = conditional_t<FusionOp::IsAuxOutSupported, CopyOpR2S, CopyOpS2R>;
using Callbacks = fusion::FusionCallbacks<
Sm90TmaWarpSpecialized<StagesC, StagesD, FragmentSize, ReuseSmemC>,
@@ -206,6 +215,32 @@ struct CallbacksBuilder<
>;
};
template <
int StagesC,
int StagesD,
int FragmentSize,
bool ReuseSmemC,
class FusionOp,
class TileShape_MNK,
class EpilogueTile_MN,
class ElementAccumulator
>
struct CallbacksBuilder<
Sm90TmaWarpSpecialized<StagesC, StagesD, FragmentSize, ReuseSmemC>,
FusionOp,
TileShape_MNK,
EpilogueTile_MN,
ElementAccumulator,
enable_if_t<(FusionOp::IsAuxOutSupported ^ FusionOp::IsAuxInSupported) // only one aux tensor
&& sizeof_bits_v<typename FusionOp::ElementAux> == 1>
> {
using Callbacks = fusion::FusionCallbacks<
Sm90TmaWarpSpecialized<StagesC, StagesD, FragmentSize, ReuseSmemC>,
FusionOp, TileShape_MNK, EpilogueTile_MN,
Layout<_1,_0>, DefaultCopy // aux bit tensor doesn't use smem
>;
};
// Helper for building TMA warp-specialized collective epilogues, specialized by
// the fusion operation performed and the dispatch policy to use.
template <
@@ -279,7 +314,7 @@ struct EpilogueDescriptor {
using EpilogueTile =
decltype(
detail::sm90_compute_tile_shape_or_override<
ElementD, EpilogueTileType, Schedule
ElementD, EpilogueTileType, Schedule, TileShape_MNK
>()
);
using DispatchPolicy =
@@ -443,7 +478,7 @@ struct CollectiveBuilder<
cute::is_same_v<Schedule, TmaWarpSpecializedCooperative> >> {
private:
using EpilogueTile_MN =
decltype(detail::sm90_compute_tile_shape_or_override<ElementD, EpilogueTileType, Schedule>());
decltype(detail::sm90_compute_tile_shape_or_override<ElementD, EpilogueTileType, Schedule, TileShape_MNK>());
using DispatchPolicy =
decltype(detail::sm90_get_tma_dispatch_policy<TileShape_MNK,EpilogueTile_MN,ElementC,ElementD,Schedule>());
@@ -623,7 +658,7 @@ CollectiveBuilder<
cute::is_base_of_v<TmaWarpSpecializedCooperativeBiasElementwiseBase, Schedule> >> {
private:
using EpilogueTile_MN = decltype(detail::sm90_compute_tile_shape_or_override<
ElementD, EpilogueTileType, Schedule>());
ElementD, EpilogueTileType, Schedule, TileShape_MNK>());
// MSVC doesn't seem to be able to deduce DispatchPolicy correctly if it's
// defined as decltype of a detail::sm90_get_tma_dispatch_policy call.
// Instead, we paste in the contents of that function. A natural refactoring
@@ -111,6 +111,18 @@ public:
return args;
}
template <class ProblemShape>
static size_t
get_workspace_size(ProblemShape const& problem_shape, Arguments const& args) {
return 0;
}
template <class ProblemShape>
static cutlass::Status
initialize_workspace(ProblemShape const& problem_shape, Arguments const& args, void* workspace, cudaStream_t stream) {
return cutlass::Status::kSuccess;
}
template<class ProblemShape>
CUTLASS_HOST_DEVICE static bool
can_implement(
@@ -87,7 +87,7 @@ CUTLASS_HOST_DEVICE
auto get_epilogue_stride(Stride stride){
if constexpr (cute::is_base_of_v<cutlass::gemm::EpilogueTransposed, EpilogueSchedule>) {
return cute::make_stride(cute::get<1>(stride), cute::get<0>(stride), cute::get<2>(stride));
}
}
else {
return stride;
}
@@ -130,7 +130,19 @@ public:
return args;
}
template<class ProblemShape>
template <class ProblemShape>
static size_t
get_workspace_size(ProblemShape const& problem_shape, Arguments const& args) {
return 0;
}
template <class ProblemShape>
static cutlass::Status
initialize_workspace(ProblemShape const& problem_shape, Arguments const& args, void* workspace, cudaStream_t stream) {
return cutlass::Status::kSuccess;
}
template <class ProblemShape>
CUTLASS_HOST_DEVICE static bool
can_implement(
[[maybe_unused]] ProblemShape const& problem_shape,
@@ -52,7 +52,7 @@ namespace collective {
/// Ways to generalize this:
/// - CTA tile shape
/// - vectorization requirements (GMEM)
/// - vectoriz(able) transform()
/// - vectoriz(able) transform()
///
template <
class StrideC_,
@@ -120,7 +120,19 @@ public:
return args;
}
template<class ProblemShape>
template <class ProblemShape>
static size_t
get_workspace_size(ProblemShape const& problem_shape, Arguments const& args) {
return 0;
}
template <class ProblemShape>
static cutlass::Status
initialize_workspace(ProblemShape const& problem_shape, Arguments const& args, void* workspace, cudaStream_t stream) {
return cutlass::Status::kSuccess;
}
template <class ProblemShape>
CUTLASS_HOST_DEVICE static bool
can_implement(
[[maybe_unused]] ProblemShape const& problem_shape,
@@ -200,8 +212,8 @@ public:
// Tile gD and gC by the shape of SmemLayout first
auto tile = make_shape(size<0>(sC), size<1>(sC));
Tensor gCt = local_tile(gC, tile, _); // (SMEM_M,SMEM_N,TILE_M,TILE_N)
Tensor gDt = local_tile(gD, tile, _); // (SMEM_M,SMEM_N,TILE_M,TILE_N)
Tensor gCt = flat_divide(gC, tile); // (SMEM_M,SMEM_N,TILE_M,TILE_N)
Tensor gDt = flat_divide(gD, tile); // (SMEM_M,SMEM_N,TILE_M,TILE_N)
// Partition sC, gC, and gD for the output
auto tiled_s2r = TiledCopyS2R{};
@@ -216,7 +228,7 @@ public:
// Repeat the D-partitioning for coordinates and predication
Tensor cD = make_identity_tensor(make_shape(size<0>(gD),size<1>(gD))); // (BLK_M,BLK_N) -> (blk_m,blk_n)
Tensor cDt = local_tile(cD, tile, _); // (SMEM_M,SMEM_N,TILE_M,TILE_N)
Tensor cDt = flat_divide(cD, tile); // (SMEM_M,SMEM_N,TILE_M,TILE_N)
Tensor tDcD = tD.partition_D(cDt); // ((Atom,AtomNum),ATOM_M,ATOM_N,TILE_M,TILE_N)
CUTE_STATIC_ASSERT(size<1>(tCaC) % size<3>(tDgC) == 0); // TILE_M divides MMA_M
@@ -258,7 +270,7 @@ public:
for (int pipe_n = 0; pipe_n < size<2>(tCsC); ++pipe_n) {
int mma_m = step_m * size<1>(tCsC) + pipe_m;
int mma_n = step_n * size<2>(tCsC) + pipe_n;
copy(tiled_r2s, tCaC(_,mma_m,mma_n), tCsC(_,pipe_m,pipe_n));
}
}
@@ -279,14 +291,14 @@ public:
// source is needed
Tensor tDgCmn = tDgC(_,_,_,step_m,step_n);
CUTLASS_PRAGMA_UNROLL
for (int m = 0; m < size<1>(tDgDmn); ++m)
for (int m = 0; m < size<1>(tDgDmn); ++m)
{
CUTLASS_PRAGMA_UNROLL
for (int n = 0; n < size<2>(tDgDmn); ++n)
for (int n = 0; n < size<2>(tDgDmn); ++n)
{
// Predication
if (get<0>(tDcDmn(0,m,n)) < get<0>(residue_mnk) &&
get<1>(tDcDmn(0,m,n)) < get<1>(residue_mnk))
get<1>(tDcDmn(0,m,n)) < get<1>(residue_mnk))
{
// Step 5. Elementwise operation with conversion
CUTLASS_PRAGMA_UNROLL
@@ -309,14 +321,14 @@ public:
}
CUTLASS_PRAGMA_UNROLL
for (int m = 0; m < size<1>(tDgDmn); ++m)
for (int m = 0; m < size<1>(tDgDmn); ++m)
{
CUTLASS_PRAGMA_UNROLL
for (int n = 0; n < size<2>(tDgDmn); ++n)
for (int n = 0; n < size<2>(tDgDmn); ++n)
{
// Predication
if (get<0>(tDcDmn(0,m,n)) < get<0>(residue_mnk) &&
get<1>(tDcDmn(0,m,n)) < get<1>(residue_mnk))
get<1>(tDcDmn(0,m,n)) < get<1>(residue_mnk))
{
// Step 6. Copy to GMEM
copy(CopyAtomR2G{}, tDrD(_,m,n), tDgDmn(_,m,n));
@@ -40,6 +40,7 @@
#include "cutlass/epilogue/collective/detail.hpp"
#include "cutlass/epilogue/thread/scale_type.h"
#include "cutlass/epilogue/fusion/callbacks.hpp"
#include "cutlass/epilogue/fusion/sm90_callbacks_tma_warpspecialized.hpp"
#include "cutlass/detail/layout.hpp"
#include "cutlass/trace.h"
@@ -165,7 +166,7 @@ public:
using LoadPipeline = cutlass::PipelineTransactionAsync<StagesC>;
using LoadPipelineState = cutlass::PipelineState<StagesC>;
constexpr static uint32_t TmaTransactionBytes =
size(take<0,2>(SmemLayoutC{})) * static_cast<uint32_t>(sizeof(SmemElementC));
(size(take<0,2>(SmemLayoutC{})) * static_cast<uint32_t>(sizeof_bits<SmemElementC>::value)) / 8;
// TMA pipeline for storing D
using StorePipeline = cute::conditional_t<ReuseSmemC,
@@ -244,7 +245,19 @@ public:
};
}
template<class ProblemShape>
template <class ProblemShape>
static size_t
get_workspace_size(ProblemShape const& problem_shape, Arguments const& args) {
return FusionCallbacks::get_workspace_size(problem_shape, args.thread);
}
template <class ProblemShape>
static cutlass::Status
initialize_workspace(ProblemShape const& problem_shape, Arguments const& args, void* workspace, cudaStream_t stream) {
return FusionCallbacks::initialize_workspace(problem_shape, args.thread, workspace, stream);
}
template <class ProblemShape>
CUTLASS_HOST_DEVICE static bool
can_implement(
ProblemShape const& problem_shape,
@@ -252,7 +265,7 @@ public:
constexpr int tma_alignment_bits = 128;
auto problem_shape_MNKL = append<4>(problem_shape, 1);
auto [M,N,K,L] = problem_shape_MNKL;
constexpr int min_tma_aligned_elements_D = tma_alignment_bits / cutlass::sizeof_bits<ElementD>::value;
bool implementable = cutlass::detail::check_alignment<min_tma_aligned_elements_D>(cute::make_shape(M,N,L), StrideD{});
@@ -275,7 +288,7 @@ public:
// Compute number of epilogue subtiles
constexpr int epi_m = size<0>(tile_shape_MNK) / size<0>(EpilogueTile{});
constexpr int epi_n = size<1>(tile_shape_MNK) / size<1>(EpilogueTile{});
return epi_m * epi_n;
}
@@ -326,6 +339,15 @@ public:
auto [M, N, K, L] = problem_shape_mnkl;
auto [m_coord, n_coord, k_coord, l_coord] = tile_coord_mnkl;
// Tile residue
auto m_max_coord = unwrap(cute::transform(make_seq<rank<0>(tile_shape_MNK)>{}, [&](auto i) {
return get<0,i>(problem_shape_mnkl) - get<0,i>(tile_shape_MNK) * get<0,i>(tile_coord_mnkl);
}));
auto n_max_coord = unwrap(cute::transform(make_seq<rank<1>(tile_shape_MNK)>{}, [&](auto i) {
return get<1,i>(problem_shape_mnkl) - get<1,i>(tile_shape_MNK) * get<1,i>(tile_coord_mnkl);
}));
auto residue_mn = make_coord(m_max_coord, n_max_coord);
// Represent the full source tensor, slice to get the tile this CTA is currently responsible for
Tensor mC = params.tma_load_c.get_tma_tensor(make_shape(M,N,L)); // (M,N,L)
Tensor gC = local_tile(mC, take<0,2>(CtaTileMNK{}), make_coord(m_coord,n_coord,l_coord)); // (CTA_M,CTA_N)
@@ -335,7 +357,7 @@ public:
if constexpr (not ReuseSmemC and is_source_supported) {
ptr_sC = shared_tensors.smem_C.data();
}
Tensor gC_epi = local_tile(gC, EpilogueTile{}, _); // (EPI_TILE_M,EPI_TILE_N,EPI_M,EPI_N)
Tensor gC_epi = flat_divide(gC, EpilogueTile{}); // (EPI_TILE_M,EPI_TILE_N,EPI_M,EPI_N)
Tensor sC_epi = make_tensor(make_smem_ptr(ptr_sC), SmemLayoutC{}); // (EPI_TILE_M,EPI_TILE_N,PIPE_C)
// Prepare the thread(b)lock's (G)mem to (S)mem TMA tiled copy (bGS_)
@@ -344,12 +366,15 @@ public:
Tensor bGS_sC = thrblk_g2s.partition_D(sC_epi); // (G2S,G2S_M,G2S_N,PIPE_C)
// Get the fusion callbacks for the producer load warp
auto pld_callbacks = fusion_callbacks.get_producer_load_callbacks(
problem_shape_mnkl,
CtaTileMNK{},
tile_coord_mnkl,
EpilogueTile{},
thread_idx);
auto pld_args = cutlass::epilogue::fusion::detail::ProducerLoadArgs{
problem_shape_mnkl,
CtaTileMNK{},
tile_coord_mnkl,
residue_mn,
EpilogueTile{},
thread_idx
};
auto pld_callbacks = fusion_callbacks.get_producer_load_callbacks(pld_args);
bool is_C_load_needed = is_source_supported && fusion_callbacks.is_C_load_needed();
// Predication for TMA load (one thread issues TMA load)
@@ -445,12 +470,13 @@ public:
auto epi_tile_m = size<0>(EpilogueTile{});
auto epi_tile_n = size<1>(EpilogueTile{});
// Represent the full output tensor, slice to get the tile this CTA is responsible for
Tensor mD = params.tma_store_d.get_tma_tensor(make_shape(M,N,L)); // (M,N,L)
Tensor gD = local_tile(mD, take<0,2>(CtaTileMNK{}), make_coord(m_coord,n_coord,l_coord)); // (CTA_M,CTA_N)
// Apply epilogue subtiling
Tensor gD_epi = local_tile(gD, EpilogueTile{}, _); // (EPI_TILE_M,EPI_TILE_N,EPI_M,EPI_N)
Tensor gD_epi = flat_divide(gD, EpilogueTile{}); // (EPI_TILE_M,EPI_TILE_N,EPI_M,EPI_N)
// Construct the corresponding pipelined smem tensors
SmemElementC* ptr_sC = reinterpret_cast<SmemElementC*>(shared_tensors.smem_D.data());
@@ -503,19 +529,37 @@ public:
Tensor bSG_sD = thrblk_s2g.partition_S(sD_epi); // (S2G,S2G_M,S2G_N,PIPE_D)
Tensor bSG_gD = thrblk_s2g.partition_D(gD_epi); // (S2G,S2G_M,S2G_N,EPI_M,EPI_N)
// Coordinate tensors and residue for tile quantization
auto m_max_coord = unwrap(cute::transform(make_seq<rank<0>(CtaTileMNK{})>{}, [&](auto i) {
auto c_m = get<0,i>(problem_shape_mnkl) - get<0,i>(CtaTileMNK{}) * get<0,i>(tile_coord_mnkl);
return cute::max(0, c_m);
}));
auto n_max_coord = unwrap(cute::transform(make_seq<rank<1>(CtaTileMNK{})>{}, [&](auto i) {
auto c_n = get<1,i>(problem_shape_mnkl) - get<1,i>(CtaTileMNK{}) * get<1,i>(tile_coord_mnkl);
return cute::max(0, c_n);
}));
auto residue_mn = make_coord(m_max_coord, n_max_coord);
Tensor cD = make_identity_tensor(take<0,2>(CtaTileMNK{}));
Tensor tRS_cD = thread_r2s.partition_S(flat_divide(cD, EpilogueTile{}));
CUTE_STATIC_ASSERT(mma_tile_m == epi_tile_m, "EPI_TILE_M must equal MMA_TILE_M");
CUTE_STATIC_ASSERT(mma_tile_n % epi_tile_n == 0, "EPI_TILE_N must divide MMA_TILE_N");
// Get the fusion callbacks for the consumer store warps
constexpr bool RefSrc = true; // Register tensors reference R2S copy src layout
auto cst_callbacks = fusion_callbacks.get_consumer_store_callbacks<RefSrc>(
problem_shape_mnkl,
CtaTileMNK{},
tile_coord_mnkl,
EpilogueTile{},
tiled_copy_C_atom,
thread_idx,
tRS_rC);
auto cst_args = cutlass::epilogue::fusion::detail::ConsumerStoreArgs{
problem_shape_mnkl,
CtaTileMNK{},
tile_coord_mnkl,
residue_mn,
EpilogueTile{},
tiled_copy_C_atom,
thread_idx,
cD,
tRS_cD,
tRS_rC
};
auto cst_callbacks = fusion_callbacks.get_consumer_store_callbacks<RefSrc>(cst_args);
bool is_producer_load_needed = fusion_callbacks.is_producer_load_needed();
bool is_C_load_needed = is_source_supported && fusion_callbacks.is_C_load_needed();
@@ -548,27 +592,13 @@ public:
for (int epi_n = 0; epi_n < size<3>(gD_epi); ++epi_n) {
CUTLASS_PRAGMA_UNROLL
for (int epi_m = 0; epi_m < size<2>(gD_epi); ++epi_m) {
bool is_last_iteration = epi_m == size<2>(gD_epi)-1 && epi_n == size<3>(gD_epi)-1;
// The current tile in accumulator
int mma_m = epi_m;
int mma_n = (epi_n * epi_tile_n) / mma_tile_n;
Tensor tRS_rAcc_frg_mn = tRS_rAcc_frg(_,mma_m,mma_n);
// Wait for a smem buffer to be available
if (issue_tma_store) {
store_pipeline.producer_acquire(store_pipe_producer_state);
}
synchronize();
if constexpr (ReuseSmemC) {
// Let dma warp know smem buffer is consumed and empty after StagesD producer commits
if (issued_stores >= StagesD) {
if (is_producer_load_needed) {
load_pipeline.consumer_release(load_pipe_consumer_state);
}
++load_pipe_consumer_state;
}
}
if (is_producer_load_needed) {
// Wait for the producer load to fill smem
load_pipeline.consumer_wait(load_wait_state);
@@ -580,7 +610,7 @@ public:
}
// First loop fusion callback entry point
cst_callbacks.step_begin(epi_m, epi_n, load_wait_state.count(), is_producer_load_needed);
cst_callbacks.previsit(epi_m, epi_n, load_wait_state.count(), is_producer_load_needed);
if (is_producer_load_needed) {
if constexpr (not ReuseSmemC) {
@@ -602,9 +632,9 @@ public:
// Copy tile from register to smem
copy(tiled_r2s, tRS_rD, tRS_sD(_,_,_,store_pipe_producer_state.index()));
// Next loop fusion callback entry point
// Post visit, pre async fence callback entry point
constexpr bool issue_smem_store = true; // No smem store predication
cst_callbacks.step_next(epi_m, epi_n, store_pipe_producer_state.count(), issue_smem_store);
cst_callbacks.postvisit(epi_m, epi_n, store_pipe_producer_state.count(), issue_smem_store);
// Write the tile from smem to gmem with TMA
cutlass::arch::fence_view_async_shared(); // ensure smem writes are visible to TMA
@@ -613,8 +643,8 @@ public:
copy(params.tma_store_d, bSG_sD(_,_,_,store_pipe_producer_state.index()), bSG_gD(_,_,_,epi_m,epi_n));
}
// Last loop fusion callback entry point
cst_callbacks.step_end(epi_m, epi_n, store_pipe_producer_state.count(), issue_tma_store);
// Post async fence, pre TMA commit callback entry point
cst_callbacks.step(epi_m, epi_n, store_pipe_producer_state.count(), issue_tma_store);
// Commit the TMA stores for this stage
if (issue_tma_store) {
@@ -622,6 +652,42 @@ public:
}
++store_pipe_producer_state;
++issued_stores;
// Wait for the next smem buffer to be available
if (issue_tma_store) {
store_pipeline.producer_acquire(store_pipe_producer_state);
}
synchronize();
if constexpr (ReuseSmemC) {
// producer_acquire returns when at most StagesD-1 committed stores are pending
bool store_finished = issued_stores > StorePipeline::UnacquiredStages;
// Free an smem buffer for reduction if necessary
if (cst_callbacks.is_reduction_buffer_needed(epi_m, epi_n, is_last_iteration) && not store_finished) {
if (issue_tma_store) {
store_pipeline.producer_tail(store_pipe_producer_state); // wait for all TMA stores to finish
}
synchronize();
}
// Smem reduction callback entry point using least recently acquired load buffer for workspace
cst_callbacks.reduce(sC_epi(_,_,load_pipe_consumer_state.index()),
synchronize, epi_m, epi_n, is_last_iteration);
// Let dma warp know earliest smem buffer is consumed and empty after StagesD producer commits
if (store_finished) {
if (is_producer_load_needed) {
load_pipeline.consumer_release(load_pipe_consumer_state);
}
++load_pipe_consumer_state;
}
}
else {
// Smem reduction callback entry point using most recently acquired store buffer for workspace
cst_callbacks.reduce(sD_epi(_,_,store_pipe_producer_state.index()),
synchronize, epi_m, epi_n, is_last_iteration);
}
} // for epi_m
} // for epi_n
@@ -644,9 +710,8 @@ public:
if constexpr (ReuseSmemC) {
if (fusion_callbacks.is_producer_load_needed()) {
// Issue releases on up to StagesD previously issued TMA stores
constexpr int release_stages =
cute::min(StagesD, get_load_pipe_increment(CtaTileMNK{}));
// Issue releases on up to StagesD-1 previously issued TMA stores
constexpr int release_stages = cute::min(StorePipeline::UnacquiredStages, get_load_pipe_increment(CtaTileMNK{}));
CUTLASS_PRAGMA_UNROLL
for (int stage = 0; stage < release_stages; ++stage) {
load_pipeline.consumer_release(load_pipe_consumer_state);
@@ -60,15 +60,21 @@ struct FusionOperation {
using ElementBias = void;
static constexpr int AlignmentBias = 0;
static constexpr bool IsPerRowBiasSupported = false;
static constexpr bool IsDePerRowBiasSupported = false;
using ActivationFn = void;
static constexpr bool IsEltActSupported = false;
static constexpr bool IsDeEltActSupported = false;
using ElementAux = void;
using GmemLayoutTagAux = void;
static constexpr int AlignmentAux = 0;
static constexpr bool IsAuxOutSupported = false;
static constexpr bool IsAuxInSupported = false;
using ElementAmax = void;
static constexpr bool IsAbsMaxSupported = false;
};
// D = alpha * acc
@@ -242,6 +248,51 @@ struct ScaledLinCombPerRowBiasEltActAmaxAux
static constexpr bool IsAuxOutSupported = true;
};
// Z = Aux
// dY = alpha * acc + beta * C
// D = d_activation(dY, Z)
template<
class GmemLayoutTagAux_,
template <class> class ActivationFn_,
class ElementOutput_,
class ElementCompute_,
class ElementAux_ = ElementOutput_,
class ElementScalar_ = ElementCompute_,
int AlignmentAux_ = 128 / sizeof_bits_v<ElementAux_>,
FloatRoundStyle RoundStyle_ = FloatRoundStyle::round_to_nearest
>
struct LinCombDeEltAct
: LinCombEltAct<ActivationFn_, ElementOutput_, ElementCompute_, ElementScalar_, RoundStyle_> {
using ElementAux = ElementAux_;
using GmemLayoutTagAux = GmemLayoutTagAux_;
static constexpr int AlignmentAux = AlignmentAux_;
static constexpr bool IsAuxInSupported = true;
};
// Z = Aux
// dY = alpha * acc + beta * C
// D = d_activation(dY, Z)
// dBias = sum of columns of D
template<
class GmemLayoutTagAux_,
template <class> class ActivationFn_,
class ElementOutput_,
class ElementCompute_,
class ElementAux_ = ElementOutput_,
class ElementBias_ = ElementCompute_,
class ElementScalar_ = ElementCompute_,
int AlignmentAux_ = 128 / sizeof_bits_v<ElementAux_>,
int AlignmentBias_ = 128 / sizeof_bits_v<ElementBias_>,
FloatRoundStyle RoundStyle_ = FloatRoundStyle::round_to_nearest
>
struct LinCombDeEltActDePerRowBias
: LinCombDeEltAct<GmemLayoutTagAux_, ActivationFn_, ElementOutput_, ElementCompute_,
ElementAux_, ElementScalar_, AlignmentAux_, RoundStyle_> {
using ElementBias = ElementBias_;
static constexpr int AlignmentBias = AlignmentBias_;
static constexpr bool IsDePerRowBiasSupported = true;
};
/////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace cutlass::epilogue::fusion
@@ -117,7 +117,7 @@ template<
FloatRoundStyle RoundStyle = FloatRoundStyle::round_to_nearest
>
using Sm90LinearCombination =
Sm90EVT<Sm90Compute<multiply_add, ElementOutput, ElementCompute, RoundStyle>, // beta * C + (alpha * acc)
Sm90EVT<Sm90Compute<homogeneous_multiply_add, ElementOutput, ElementCompute, RoundStyle>, // beta * C + (alpha * acc)
Sm90ScalarBroadcast<ElementScalar>, // beta
Sm90SrcFetch, // C
Sm90EVT<Sm90Compute<multiplies, ElementCompute, ElementCompute, RoundStyle>, // alpha * acc
@@ -143,9 +143,9 @@ struct FusionCallbacks<
fusion::LinearCombination<ElementOutput, ElementCompute, ElementScalar, RoundStyle>,
CtaTileShapeMNK,
EpilogueTile
> : Sm90LinearCombination<ElementOutput, ElementCompute, ElementScalar, RoundStyle> {
> : Sm90LinearCombination<typename cutlass::detail::get_unpacked_element_type<ElementOutput>::type, ElementCompute, ElementScalar, RoundStyle> {
using Impl = Sm90LinearCombination<ElementOutput, ElementCompute, ElementScalar, RoundStyle>;
using Impl = Sm90LinearCombination<typename cutlass::detail::get_unpacked_element_type<ElementOutput>::type, ElementCompute, ElementScalar, RoundStyle>;
using Operation = fusion::LinearCombination<ElementOutput, ElementCompute, ElementScalar, RoundStyle>;
struct Arguments {
@@ -208,7 +208,7 @@ struct FusionCallbacks<
EpilogueTile
> : Sm90LinCombEltAct<ActivationFn, ElementOutput, ElementCompute, ElementScalar, RoundStyle> {
using Impl = Sm90LinCombEltAct<ActivationFn, ElementOutput, ElementCompute, ElementScalar, RoundStyle>;
using Impl = Sm90LinCombEltAct<ActivationFn, typename cutlass::detail::get_unpacked_element_type<ElementOutput>::type, ElementCompute, ElementScalar, RoundStyle>;
using Operation = fusion::LinCombEltAct<ActivationFn, ElementOutput, ElementCompute, ElementScalar, RoundStyle>;
struct Arguments {
@@ -255,10 +255,10 @@ template<
FloatRoundStyle RoundStyle = FloatRoundStyle::round_to_nearest
>
using Sm90LinCombPerRowBias =
Sm90EVT<Sm90Compute<multiply_add, ElementOutput, ElementCompute, RoundStyle>, // beta * C + (alpha * acc + bias)
Sm90EVT<Sm90Compute<homogeneous_multiply_add, ElementOutput, ElementCompute, RoundStyle>, // beta * C + (alpha * acc + bias)
Sm90ScalarBroadcast<ElementScalar>, // beta
Sm90SrcFetch, // C
Sm90EVT<Sm90Compute<multiply_add, ElementCompute, ElementCompute, RoundStyle>, // alpha * acc + bias
Sm90EVT<Sm90Compute<homogeneous_multiply_add, ElementCompute, ElementCompute, RoundStyle>, // alpha * acc + bias
Sm90ScalarBroadcast<ElementScalar>, // alpha
Sm90AccFetch, // acc
Sm90ColBroadcast<0, CtaTileShapeMNK, ElementBias, Stride<_1,_0,int>, AlignmentBias> // bias
@@ -541,10 +541,10 @@ template<
FloatRoundStyle RoundStyle = FloatRoundStyle::round_to_nearest
>
using Sm90PerRowLinCombPerRowBias =
Sm90EVT<Sm90Compute<multiply_add, ElementOutput, ElementCompute, RoundStyle>, // beta * C + (alpha * acc + bias)
Sm90EVT<Sm90Compute<homogeneous_multiply_add, ElementOutput, ElementCompute, RoundStyle>, // beta * C + (alpha * acc + bias)
Sm90ColBroadcast<0, CtaTileShapeMNK, ElementScalar, Stride<_1,_0,_0>, AlignmentScalar>, // beta
Sm90SrcFetch, // C
Sm90EVT<Sm90Compute<multiply_add, ElementCompute, ElementCompute, RoundStyle>, // alpha * acc + bias
Sm90EVT<Sm90Compute<homogeneous_multiply_add, ElementCompute, ElementCompute, RoundStyle>, // alpha * acc + bias
Sm90ColBroadcast<0, CtaTileShapeMNK, ElementScalar, Stride<_1,_0,_0>, AlignmentScalar>, // alpha
Sm90AccFetch, // acc
Sm90ColBroadcast<0, CtaTileShapeMNK, ElementBias, Stride<_1,_0,int>, AlignmentBias> // bias
@@ -669,10 +669,10 @@ template<
FloatRoundStyle RoundStyle = FloatRoundStyle::round_to_nearest
>
using Sm90ScaledLinCombPerRowBias =
Sm90EVT<Sm90Compute<multiply_add, ElementOutput, ElementCompute, RoundStyle>, // beta * C + (alpha * acc + bias)
Sm90EVT<Sm90Compute<homogeneous_multiply_add, ElementOutput, ElementCompute, RoundStyle>, // beta * C + (alpha * acc + bias)
Sm90ScalarBroadcast<ElementScalar, Stride<_0,_0,_0>, 2>, // scale_c * beta
Sm90SrcFetch, // C
Sm90EVT<Sm90Compute<multiply_add, ElementCompute, ElementCompute, RoundStyle>, // alpha * acc + bias
Sm90EVT<Sm90Compute<homogeneous_multiply_add, ElementCompute, ElementCompute, RoundStyle>, // alpha * acc + bias
Sm90ScalarBroadcast<ElementScalar, Stride<_0,_0,_0>, 3>, // scale_a * scale_b * alpha
Sm90AccFetch, // acc
Sm90ColBroadcast<0, CtaTileShapeMNK, ElementBias, Stride<_1,_0,int>, AlignmentBias> // bias
@@ -1003,6 +1003,229 @@ struct FusionCallbacks<
/////////////////////////////////////////////////////////////////////////////////////////////////
template<
class CtaTileShapeMNK,
class EpilogueTile,
int Stages,
class StrideAux,
class SmemLayoutAtom,
class CopyOpS2R,
template <class> class ActivationFn,
class ElementOutput,
class ElementCompute,
class ElementAux = ElementOutput,
class ElementScalar = ElementCompute,
int AlignmentAux = 128 / sizeof_bits_v<ElementAux>,
FloatRoundStyle RoundStyle = FloatRoundStyle::round_to_nearest
>
using Sm90LinCombDeEltAct =
Sm90EVT<Sm90Compute<ActivationFn, ElementOutput, ElementCompute, RoundStyle>, // activation(beta * C + (alpha * acc), aux)
Sm90LinearCombination<ElementCompute, ElementCompute, ElementScalar, RoundStyle>, // beta * C + (alpha * acc)
Sm90AuxLoad<Stages, EpilogueTile, ElementAux, StrideAux, SmemLayoutAtom, CopyOpS2R, AlignmentAux>, // aux
>;
template <
int StagesC,
int StagesD,
int FragmentSize,
bool ReuseSmemC,
class GmemLayoutTagAux,
template <class> class ActivationFn,
class ElementOutput,
class ElementCompute,
class ElementAux,
class ElementScalar,
int AlignmentAux,
FloatRoundStyle RoundStyle,
class CtaTileShapeMNK,
class EpilogueTile,
class SmemLayoutAtom,
class CopyOpS2R
>
struct FusionCallbacks<
epilogue::Sm90TmaWarpSpecialized<StagesC, StagesD, FragmentSize, ReuseSmemC>,
fusion::LinCombDeEltAct<
GmemLayoutTagAux, ActivationFn, ElementOutput, ElementCompute,
ElementAux, ElementScalar, AlignmentAux, RoundStyle
>,
CtaTileShapeMNK,
EpilogueTile,
SmemLayoutAtom,
CopyOpS2R
> : Sm90LinCombDeEltAct<
CtaTileShapeMNK, EpilogueTile, StagesC, cutlass::gemm::TagToStrideC_t<GmemLayoutTagAux>, SmemLayoutAtom, CopyOpS2R, ActivationFn,
ElementOutput, ElementCompute, ElementAux, ElementScalar, AlignmentAux, RoundStyle
> {
using Impl =
Sm90LinCombDeEltAct<
CtaTileShapeMNK, EpilogueTile, StagesC, cutlass::gemm::TagToStrideC_t<GmemLayoutTagAux>, SmemLayoutAtom, CopyOpS2R, ActivationFn,
ElementOutput, ElementCompute, ElementAux, ElementScalar, AlignmentAux, RoundStyle
>;
using Operation =
fusion::LinCombDeEltAct<
GmemLayoutTagAux, ActivationFn, ElementOutput, ElementCompute,
ElementAux, ElementScalar, AlignmentAux, RoundStyle
>;
struct Arguments {
ElementScalar alpha = ElementScalar(1);
ElementScalar beta = ElementScalar(0);
ElementScalar const* alpha_ptr = nullptr;
ElementScalar const* beta_ptr = nullptr;
using ActivationArguments = typename Sm90Compute<ActivationFn, ElementOutput, ElementCompute, RoundStyle>::Arguments;
ActivationArguments activation = ActivationArguments();
using StrideAux = cutlass::gemm::TagToStrideC_t<GmemLayoutTagAux>;
ElementAux const* aux_ptr = nullptr;
StrideAux dAux = {};
operator typename Impl::Arguments() const {
return
{ // binary op : activation(beta * C + (alpha * acc), aux)
{ // ternary op : beta * C + (alpha * acc)
{{beta}, {beta_ptr}}, // leaf args : beta
{}, // leaf args : C
{ // binary op : alpha * acc
{{alpha}, {alpha_ptr}}, // leaf args : alpha
{}, // leaf args : acc
{} // binary args : multiplies
}, // end binary op
{} // ternary args : multiply_add
}, // end ternary op
{aux_ptr, ElementAux(0), dAux}, // leaf args : aux
activation // binary args : activation
}; // end binary op
}
};
// Ctor inheritance
using Impl::Impl;
};
/////////////////////////////////////////////////////////////////////////////////////////////////
template<
class CtaTileShapeMNK,
class EpilogueTile,
int Stages,
class StrideAux,
class SmemLayoutAtom,
class CopyOpS2R,
template <class> class ActivationFn,
class ElementOutput,
class ElementCompute,
class ElementAux = ElementOutput,
class ElementBias = ElementOutput,
class ElementScalar = ElementCompute,
int AlignmentAux = 128 / sizeof_bits_v<ElementAux>,
int AlignmentBias = 128 / sizeof_bits_v<ElementBias>,
FloatRoundStyle RoundStyle = FloatRoundStyle::round_to_nearest
>
using Sm90LinCombDeEltActDePerRowBias =
Sm90EVT<Sm90Compute<cutlass::epilogue::thread::Identity, ElementOutput, ElementCompute, RoundStyle>, // Identity for final conversion
Sm90EVT<Sm90ColReduction<plus, plus, 0, CtaTileShapeMNK,
ElementBias, ElementCompute, RoundStyle, Stride<_1,_0,int>, AlignmentBias>,
Sm90LinCombDeEltAct<CtaTileShapeMNK, EpilogueTile, Stages, StrideAux, SmemLayoutAtom, CopyOpS2R, ActivationFn,
ElementCompute, ElementCompute, ElementAux, ElementScalar, AlignmentAux, RoundStyle>
>
>;
template <
int StagesC,
int StagesD,
int FragmentSize,
bool ReuseSmemC,
class GmemLayoutTagAux,
template <class> class ActivationFn,
class ElementOutput,
class ElementCompute,
class ElementAux,
class ElementBias,
class ElementScalar,
int AlignmentAux,
int AlignmentBias,
FloatRoundStyle RoundStyle,
class CtaTileShapeMNK,
class EpilogueTile,
class SmemLayoutAtom,
class CopyOpS2R
>
struct FusionCallbacks<
epilogue::Sm90TmaWarpSpecialized<StagesC, StagesD, FragmentSize, ReuseSmemC>,
fusion::LinCombDeEltActDePerRowBias<
GmemLayoutTagAux, ActivationFn, ElementOutput, ElementCompute,
ElementAux, ElementBias, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
>,
CtaTileShapeMNK,
EpilogueTile,
SmemLayoutAtom,
CopyOpS2R
> : Sm90LinCombDeEltActDePerRowBias<
CtaTileShapeMNK, EpilogueTile, StagesC, cutlass::gemm::TagToStrideC_t<GmemLayoutTagAux>, SmemLayoutAtom, CopyOpS2R, ActivationFn,
ElementOutput, ElementCompute, ElementAux, ElementBias, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
> {
using Impl =
Sm90LinCombDeEltActDePerRowBias<
CtaTileShapeMNK, EpilogueTile, StagesC, cutlass::gemm::TagToStrideC_t<GmemLayoutTagAux>, SmemLayoutAtom, CopyOpS2R, ActivationFn,
ElementOutput, ElementCompute, ElementAux, ElementBias, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
>;
using Operation =
fusion::LinCombDeEltActDePerRowBias<
GmemLayoutTagAux, ActivationFn, ElementOutput, ElementCompute,
ElementAux, ElementBias, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
>;
struct Arguments {
ElementScalar alpha = ElementScalar(1);
ElementScalar beta = ElementScalar(0);
ElementScalar const* alpha_ptr = nullptr;
ElementScalar const* beta_ptr = nullptr;
using ActivationArguments = typename Sm90Compute<ActivationFn, ElementOutput, ElementCompute, RoundStyle>::Arguments;
ActivationArguments activation = ActivationArguments();
using StrideAux = cutlass::gemm::TagToStrideC_t<GmemLayoutTagAux>;
ElementAux const* aux_ptr = nullptr;
StrideAux dAux = {};
using StrideBias = Stride<_1,_0,int>;
ElementBias* dbias_ptr = nullptr;
StrideBias dDbias = {};
operator typename Impl::Arguments() const {
return
{ // unary op : identity/convert
{ // unary op : reduce(activation(beta * C + (alpha * acc), aux))
{ // binary op : activation(beta * C + (alpha * acc), aux)
{ // ternary op : beta * C + (alpha * acc)
{{beta}, {beta_ptr}}, // leaf args : beta
{}, // leaf args : C
{ // binary op : alpha * acc
{{alpha}, {alpha_ptr}}, // leaf args : alpha
{}, // leaf args : acc
{} // binary args : multiplies
}, // end binary op
{} // ternary args : multiply_add
}, // end ternary op
{aux_ptr, ElementAux(0), dAux}, // leaf args : aux
activation // binary args : activation
}, // end binary op
{dbias_ptr, ElementCompute(0), dDbias} // unary args : reduce
}, // end unary op
{} // unary args : identity/convert
}; // end unary op
}
};
// Ctor inheritance
using Impl::Impl;
};
/////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace cutlass::epilogue::fusion
/////////////////////////////////////////////////////////////////////////////////////////////////
@@ -38,6 +38,7 @@
#include "cutlass/cutlass.h"
#include "cutlass/array.h"
#include "cutlass/numeric_conversion.h"
#include "cutlass/epilogue/thread/activation.h"
#include "cute/tensor.hpp"
@@ -60,6 +61,25 @@ using namespace detail;
//
/////////////////////////////////////////////////////////////////////////////////////////////////
// The template argument provided for ComputeFn must be able to accept
// exactly one template parameter. In Standard C++, it's OK for
// ComputeFn to have other template parameters, as long as those have
// defaults. For example, the following struct Foo would work.
//
// template<class A, class B = A>
// struct Foo {
// CUTLASS_HOST_DEVICE auto operator() (A a, B b);
// };
//
// However, some compilers, such as Clang, require that the argument
// take _exactly_ one template parameter. This is nonstandard C++
// behavior. One work-around for this case is to create a subclass
// with exactly one template parameter, and then use that subclass as
// the template argument.
//
// template<class A>
// struct FooHomogeneous : public Foo<A, B> {};
//
template<
template <class> class ComputeFn,
class ElementOutput,
@@ -67,77 +87,25 @@ template<
FloatRoundStyle RoundStyle,
class = void
>
struct Sm90Compute : Sm90VisitorImpl<> {
using Sm90VisitorImpl<>::Sm90VisitorImpl;
struct ConsumerStoreCallbacks : EmptyConsumerStoreCallbacks {
template <typename ElementAccumulator, typename... ElementInputs, int FragmentSize>
CUTLASS_DEVICE Array<ElementOutput, FragmentSize>
visit(Array<ElementAccumulator, FragmentSize> const& frg_acc, int epi_v, int epi_m, int epi_n,
Array<ElementInputs, FragmentSize> const&... frg_inputs) {
return transform_apply(cute::make_tuple(frg_inputs...),
[&] (auto&& frg_input) {
using ElementInput = typename cute::remove_cvref_t<decltype(frg_input)>::Element;
using ConvertInput = NumericArrayConverter<ElementCompute, ElementInput, FragmentSize, RoundStyle>;
ConvertInput convert_input{};
return convert_input(frg_input);
},
[&] (auto&&... cvt_frg_inputs) {
using ComputeOutput = ComputeFn<Array<ElementCompute, FragmentSize>>;
using ConvertOutput = NumericArrayConverter<ElementOutput, ElementCompute, FragmentSize, RoundStyle>;
ComputeOutput compute_output{};
ConvertOutput convert_output{};
return convert_output(compute_output(cvt_frg_inputs...));
}
);
}
struct Sm90Compute {
private:
using EmptyArguments = typename Sm90VisitorImpl<>::Arguments;
template <class Fn, class = void>
struct ComputeArguments {
using type = EmptyArguments;
};
template <
bool ReferenceSrc, // do register tensors reference the src or dst layout of the tiled copy
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class EpilogueTile,
class TiledCopy,
class SrcTensor
>
CUTLASS_DEVICE auto
get_consumer_store_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
TiledCopy tiled_copy,
int thread_idx,
SrcTensor const& tCrC) {
return ConsumerStoreCallbacks();
}
};
// partial specialization for compute fns that define an Arguments member, e.g. activation hyperparameters
template<
template <class> class ComputeFn,
class ElementOutput,
class ElementCompute,
FloatRoundStyle RoundStyle
>
struct Sm90Compute<
ComputeFn,
ElementOutput,
ElementCompute,
RoundStyle,
cute::void_t<typename ComputeFn<ElementCompute>::Arguments>
> {
// partial specialization for compute fns that define an Arguments member, e.g. activation hyperparameters
template <class Fn>
struct ComputeArguments<Fn, platform::void_t<typename Fn::Arguments>> {
using type = typename Fn::Arguments;
};
public:
struct SharedStorage { };
using Arguments = typename ComputeFn<ElementCompute>::Arguments;
using Arguments = typename ComputeArguments<ComputeFn<ElementCompute>>::type;
using Params = Arguments;
@@ -147,6 +115,18 @@ struct Sm90Compute<
return args;
}
template <class ProblemShape>
static size_t
get_workspace_size(ProblemShape const& problem_shape, Arguments const& args) {
return 0;
}
template <class ProblemShape>
static cutlass::Status
initialize_workspace(ProblemShape const& problem_shape, Arguments const& args, void* workspace, cudaStream_t stream) {
return cutlass::Status::kSuccess;
}
CUTLASS_DEVICE bool
is_producer_load_needed() const {
return false;
@@ -166,19 +146,9 @@ struct Sm90Compute<
Params const params;
template <
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class EpilogueTile
>
template <class... Args>
CUTLASS_DEVICE auto
get_producer_load_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
int thread_idx) {
get_producer_load_callbacks(ProducerLoadArgs<Args...> const& args) {
return EmptyProducerLoadCallbacks{};
}
@@ -207,7 +177,12 @@ struct Sm90Compute<
ComputeOutput compute_output{};
ConvertOutput convert_output{};
return convert_output(compute_output(cvt_frg_inputs..., params));
if constexpr (is_same_v<Arguments, EmptyArguments>) {
return convert_output(compute_output(cvt_frg_inputs...));
}
else {
return convert_output(compute_output(cvt_frg_inputs..., params));
}
}
);
}
@@ -216,22 +191,10 @@ struct Sm90Compute<
template <
bool ReferenceSrc, // do register tensors reference the src or dst layout of the tiled copy
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class EpilogueTile,
class TiledCopy,
class SrcTensor
class... Args
>
CUTLASS_DEVICE auto
get_consumer_store_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
TiledCopy tiled_copy,
int thread_idx,
SrcTensor const& tCrC) {
get_consumer_store_callbacks(ConsumerStoreArgs<Args...> const& args) {
return ConsumerStoreCallbacks(params);
}
@@ -255,7 +218,7 @@ template <
class InputAddOp // Z
>
struct Sm90TreeVisitor<
Sm90Compute<multiply_add, ElementOutput, ElementCompute, RoundStyle>,
Sm90Compute<homogeneous_multiply_add, ElementOutput, ElementCompute, RoundStyle>,
Sm90ScalarBroadcast<ElementScalar, StrideScalar, ScalarCount, ScalarReduceFn>,
Sm90SrcFetch,
InputAddOp
@@ -263,7 +226,7 @@ struct Sm90TreeVisitor<
Sm90ScalarBroadcast<ElementScalar, StrideScalar, ScalarCount, ScalarReduceFn>,
Sm90SrcFetch,
InputAddOp,
Sm90Compute<multiply_add, ElementOutput, ElementCompute, RoundStyle>
Sm90Compute<homogeneous_multiply_add, ElementOutput, ElementCompute, RoundStyle>
>
{
using Impl =
@@ -271,7 +234,7 @@ struct Sm90TreeVisitor<
Sm90ScalarBroadcast<ElementScalar, StrideScalar, ScalarCount, ScalarReduceFn>,
Sm90SrcFetch,
InputAddOp,
Sm90Compute<multiply_add, ElementOutput, ElementCompute, RoundStyle>
Sm90Compute<homogeneous_multiply_add, ElementOutput, ElementCompute, RoundStyle>
>;
CUTLASS_DEVICE bool
@@ -334,37 +297,482 @@ struct Sm90TreeVisitor<
template <
bool ReferenceSrc, // do register tensors reference the src or dst layout of the tiled copy
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class EpilogueTile,
class TiledCopy,
class SrcTensor
class... Args
>
CUTLASS_DEVICE auto
get_consumer_store_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
TiledCopy tiled_copy,
int thread_idx,
SrcTensor const& tCrC) {
get_consumer_store_callbacks(ConsumerStoreArgs<Args...> const& args) {
return ConsumerStoreCallbacks(
is_C_load_needed(),
Impl::get_consumer_store_callbacks<ReferenceSrc>(
problem_shape_mnkl,
tile_shape_mnk,
tile_coord_mnkl,
epi_tile,
tiled_copy,
thread_idx,
tCrC
)
Impl::get_consumer_store_callbacks<ReferenceSrc>(args)
);
}
};
// ReLU with aux bit tensor dReLU/dZ
// Aux(i) = Z(i) >= 0 ? 1 : 0
namespace detail {
template <
class ElementOutput,
class ElementCompute,
FloatRoundStyle RoundStyle,
class StrideMNL,
int Alignment,
bool EnableNullptr
>
struct Sm90ReLUAuxStore {
static_assert(Alignment % 128 == 0, "sub-16B alignment not supported yet");
struct SharedStorage {};
struct Arguments {
cutlass::uint1b_t* ptr_aux = nullptr;
StrideMNL dAux = {};
};
using Params = Arguments;
template <class ProblemShape>
static constexpr Params
to_underlying_arguments(ProblemShape const& problem_shape, Arguments const& args, void* workspace) {
return args;
}
template <class ProblemShape>
static size_t
get_workspace_size(ProblemShape const& problem_shape, Arguments const& args) {
return 0;
}
template <class ProblemShape>
static cutlass::Status
initialize_workspace(ProblemShape const& problem_shape, Arguments const& args, void* workspace, cudaStream_t stream) {
return cutlass::Status::kSuccess;
}
CUTLASS_HOST_DEVICE
Sm90ReLUAuxStore() { }
CUTLASS_HOST_DEVICE
Sm90ReLUAuxStore(Params const& params, SharedStorage const& shared_storage)
: params(params) { }
Params const params;
CUTLASS_DEVICE bool
is_producer_load_needed() const {
return false;
}
CUTLASS_DEVICE bool
is_C_load_needed() const {
return false;
}
template <class... Args>
CUTLASS_DEVICE auto
get_producer_load_callbacks(ProducerLoadArgs<Args...> const& args) {
return EmptyProducerLoadCallbacks{};
}
template <class RTensor, class GTensor, class CTensor, class ResidueMN>
struct ConsumerStoreCallbacks : EmptyConsumerStoreCallbacks {
CUTLASS_DEVICE
ConsumerStoreCallbacks(
RTensor&& tC_rAux,
GTensor&& tC_gAux,
CTensor tC_cAux,
ResidueMN residue_mn,
Params const& params)
: tC_rAux(cute::forward<RTensor>(tC_rAux)),
tC_gAux(cute::forward<GTensor>(tC_gAux)),
tC_cAux(tC_cAux),
residue_mn(residue_mn),
params(params) {}
RTensor tC_rAux; // (CPY,CPY_M,CPY_N,EPI_M,EPI_N)
GTensor tC_gAux; // (CPY,CPY_M,CPY_N,EPI_M,EPI_N)
CTensor tC_cAux; // (CPY,CPY_M,CPY_N,EPI_M,EPI_N)
ResidueMN residue_mn;
Params const& params;
template <typename ElementAccumulator, typename ElementInput, int FragmentSize>
CUTLASS_DEVICE auto
visit(Array<ElementAccumulator, FragmentSize> const& frg_acc, int epi_v, int epi_m, int epi_n,
Array<ElementInput, FragmentSize> const& frg_input) {
using ConvertInput = NumericArrayConverter<ElementCompute, ElementInput, FragmentSize, RoundStyle>;
using ConvertAux = PackPredicates<FragmentSize>;
using ComputeOutput = cutlass::epilogue::thread::ReLu<ElementCompute>;
using ConvertOutput = NumericArrayConverter<ElementOutput, ElementCompute, FragmentSize, RoundStyle>;
ConvertInput convert_input{};
ComputeOutput relu{};
ConvertAux convert_aux{};
ConvertOutput convert_output{};
Array frg_compute = convert_input(frg_input);
bool frg_aux[FragmentSize];
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < FragmentSize; ++i) {
ElementCompute pre_relu = frg_compute[i];
frg_compute[i] = relu(frg_compute[i]);
frg_aux[i] = frg_compute[i] == pre_relu;
}
static_assert(FragmentSize % 8 == 0, "Predicate vector must be byte-aligned");
Tensor tC_rAux_frg = recast<typename ConvertAux::result_type>(coalesce(tC_rAux(_,_,_,epi_m,epi_n))); // (EPI_V)
tC_rAux_frg(epi_v) = convert_aux(frg_aux);
return convert_output(frg_compute);
}
CUTLASS_DEVICE void
end() {
if constexpr (EnableNullptr) {
if (params.ptr_aux == nullptr) {
return;
}
}
// Copy vectorizes into byte-aligned stores
constexpr int V = cute::min(Alignment, decltype(max_common_vector(tC_rAux, tC_gAux))::value);
if constexpr (V > 0 && V % 8 == 0) {
using VecType = uint_bit_t<V>;
Tensor tC_rAux_vec = recast<VecType>(tC_rAux);
Tensor tC_gAux_vec = recast<VecType>(tC_gAux);
Tensor tC_cAux_vec = tC_cAux.compose(make_layout(Int<size(tC_rAux_vec)>{}, Int<V>{}));
auto predicate_fn = [&] (auto&&... coords) { return elem_less(tC_cAux_vec(coords...), residue_mn); };
copy_if(FunctionPredTensor(predicate_fn), tC_rAux_vec, tC_gAux_vec);
}
// sub-byte vectorization, must serialize threads
else {
// Assumes no inter-warp sharing of bytes (most copy layouts should satisfy this)
int lane_idx = canonical_lane_idx();
auto predicate_fn = [&] (auto&&... coords) { return elem_less(tC_cAux(coords...), residue_mn); };
CUTLASS_PRAGMA_NO_UNROLL
for (int i = 0; i < NumThreadsPerWarp; ++i) {
if (lane_idx == i) {
copy_if(FunctionPredTensor(predicate_fn), tC_rAux, tC_gAux);
}
__syncwarp();
}
}
}
};
template <
bool ReferenceSrc, // do register tensors reference the src or dst layout of the tiled copy
class... Args
>
CUTLASS_DEVICE auto
get_consumer_store_callbacks(ConsumerStoreArgs<Args...> const& args) {
auto [M, N, K, L] = args.problem_shape_mnkl;
auto [m, n, k, l] = args.tile_coord_mnkl;
gmem_ptr ptr_aux = make_gmem_ptr(subbyte_iterator<cutlass::uint1b_t>(params.ptr_aux));
Tensor mAux = make_tensor(ptr_aux, make_layout(make_shape(M,N,L), params.dAux)); // (M,N,L)
Tensor gAux = local_tile(mAux, take<0,2>(args.tile_shape_mnk), make_coord(m,n,l)); // (CTA_M,CTA_N)
Tensor tC_gAux = sm90_partition_for_epilogue<ReferenceSrc>( // (CPY,CPY_M,CPY_N,EPI_M,EPI_N)
gAux, args.epi_tile, args.tiled_copy, args.thread_idx);
Tensor tC_rAux = make_tensor<cutlass::uint1b_t>(shape(tC_gAux)); // (CPY,CPY_M,CPY_N,EPI_M,EPI_N)
return ConsumerStoreCallbacks(cute::move(tC_rAux), cute::move(tC_gAux), args.tCcD, args.residue_mn, params);
}
};
} // namespace detail
// Specialization on the generic compute+aux EVT
template <
// Compute node
template <class> class Activation,
class ElementOutput,
class ElementCompute,
FloatRoundStyle RoundStyle,
// Aux node
int Stages,
class EpilogueTile,
class StrideMNL,
class SmemLayoutAtom,
class CopyOpR2S,
int Alignment,
bool EnableNullptr,
// Input node
class InputOp
>
struct Sm90TreeVisitor<
Sm90Compute<Activation, ElementOutput, ElementCompute, RoundStyle,
enable_if_t<is_same_v<Activation<ElementCompute>, cutlass::epilogue::thread::ReLu<ElementCompute>>, void>>,
Sm90TreeVisitor<
Sm90AuxStore<
Stages,
EpilogueTile,
cutlass::uint1b_t,
RoundStyle,
StrideMNL,
SmemLayoutAtom,
CopyOpR2S,
Alignment,
EnableNullptr
>,
InputOp
>
> : Sm90VisitorImpl<
Sm90VisitorImpl<
InputOp,
detail::Sm90ReLUAuxStore<ElementOutput, ElementCompute, RoundStyle, StrideMNL, Alignment, EnableNullptr>
>,
Sm90Compute<Activation, ElementOutput, ElementCompute, RoundStyle>
>
{
using Impl =
Sm90VisitorImpl<
Sm90VisitorImpl<
InputOp,
detail::Sm90ReLUAuxStore<ElementOutput, ElementCompute, RoundStyle, StrideMNL, Alignment, EnableNullptr>
>,
Sm90Compute<Activation, ElementOutput, ElementCompute, RoundStyle>
>;
using Impl::Sm90VisitorImpl;
template <class CallbacksImpl>
struct ConsumerStoreCallbacks : CallbacksImpl {
CUTLASS_DEVICE
ConsumerStoreCallbacks(CallbacksImpl&& impl)
: CallbacksImpl(cute::forward<CallbacksImpl>(impl)) { }
template <typename ElementAccumulator, int FragmentSize>
CUTLASS_DEVICE Array<ElementOutput, FragmentSize>
visit(Array<ElementAccumulator, FragmentSize> const& frg_acc, int epi_v, int epi_m, int epi_n) {
auto& [callbacks_input, callbacks_relu_aux] = get<0>(CallbacksImpl::callbacks_tuple).callbacks_tuple;
Array frg_input = callbacks_input.visit(frg_acc, epi_v, epi_m, epi_n);
return callbacks_relu_aux.visit(frg_acc, epi_v, epi_m, epi_n, frg_input);
}
};
template <
bool ReferenceSrc, // do register tensors reference the src or dst layout of the tiled copy
class... Args
>
CUTLASS_DEVICE auto
get_consumer_store_callbacks(ConsumerStoreArgs<Args...> const& args) {
return ConsumerStoreCallbacks(
Impl::get_consumer_store_callbacks<ReferenceSrc>(args)
);
}
};
// Aux load for uint1b_t
template <
int Stages,
class EpilogueTile,
class StrideMNL,
class SmemLayoutAtom,
class CopyOpS2R,
int Alignment,
bool EnableNullptr
>
struct Sm90AuxLoad<
Stages,
EpilogueTile,
cutlass::uint1b_t,
StrideMNL,
SmemLayoutAtom,
CopyOpS2R,
Alignment,
EnableNullptr
> {
static_assert(Alignment % 128 == 0, "sub-16B alignment not supported yet");
struct SharedStorage {};
struct Arguments {
cutlass::uint1b_t const* ptr_aux = nullptr;
cutlass::uint1b_t null_default = cutlass::uint1b_t(0);
StrideMNL dAux = {};
};
using Params = Arguments;
template <class ProblemShape>
static constexpr Params
to_underlying_arguments(ProblemShape const& problem_shape, Arguments const& args, void* workspace) {
return args;
}
template <class ProblemShape>
static size_t
get_workspace_size(ProblemShape const& problem_shape, Arguments const& args) {
return 0;
}
template <class ProblemShape>
static cutlass::Status
initialize_workspace(ProblemShape const& problem_shape, Arguments const& args, void* workspace, cudaStream_t stream) {
return cutlass::Status::kSuccess;
}
CUTLASS_HOST_DEVICE
Sm90AuxLoad() { }
CUTLASS_HOST_DEVICE
Sm90AuxLoad(Params const& params, SharedStorage const&)
: params(params) { }
Params const params;
CUTLASS_DEVICE bool
is_producer_load_needed() const {
return false;
}
CUTLASS_DEVICE bool
is_C_load_needed() const {
return false;
}
template <class... Args>
CUTLASS_DEVICE auto
get_producer_load_callbacks(ProducerLoadArgs<Args...> const& args) {
return EmptyProducerLoadCallbacks{};
}
template <class RTensor, class GTensor, class ResidueMN>
struct ConsumerStoreCallbacks : EmptyConsumerStoreCallbacks {
CUTLASS_DEVICE
ConsumerStoreCallbacks(RTensor&& tC_rAux_, GTensor&& tC_gAux_, ResidueMN residue_mn_, Params const& params_)
: tC_rAux(cute::forward<RTensor>(tC_rAux_)),
tC_gAux(cute::forward<GTensor>(tC_gAux_)),
residue_mn(residue_mn_),
params(params_) {}
RTensor tC_rAux; // (CPY,CPY_M,CPY_N,{EPI_M,EPI_N})
GTensor tC_gAux; // (CPY,CPY_M,CPY_N,EPI_M,EPI_N)
ResidueMN residue_mn;
Params const& params;
CUTLASS_DEVICE void
begin() {
if constexpr (decltype(rank(tC_rAux))::value == 5) {
if constexpr (EnableNullptr) {
if (params.ptr_aux == nullptr) {
return;
}
}
if (elem_less(repeat_like(residue_mn, _0{}), residue_mn)) { // (partially) in-bounds CTA tile
copy(tC_gAux, tC_rAux);
}
}
}
CUTLASS_DEVICE void
previsit(int epi_m, int epi_n, int load_iteration, bool is_producer_load_needed) {
if constexpr (decltype(rank(tC_rAux))::value == 3) {
if constexpr (EnableNullptr) {
if (params.ptr_aux == nullptr) {
return;
}
}
if (elem_less(repeat_like(residue_mn, _0{}), residue_mn)) {
copy(tC_gAux(_,_,_,epi_m,epi_n), tC_rAux);
}
}
}
template <typename ElementAccumulator, int FragmentSize>
CUTLASS_DEVICE auto
visit(Array<ElementAccumulator, FragmentSize> const& frg_acc, int epi_v, int epi_m, int epi_n) {
using ElementRegister = typename remove_cvref_t<RTensor>::value_type;
if constexpr (decltype(rank(tC_rAux))::value == 3) {
return recast<Array<ElementRegister, FragmentSize>>(coalesce(tC_rAux))(epi_v);
}
else {
return recast<Array<ElementRegister, FragmentSize>>(coalesce(tC_rAux(_,_,_,epi_m,epi_n)))(epi_v);
}
}
};
template <
bool ReferenceSrc, // do register tensors reference the src or dst layout of the tiled copy
class... Args
>
CUTLASS_DEVICE auto
get_consumer_store_callbacks(ConsumerStoreArgs<Args...> const& args) {
auto [M, N, K, L] = args.problem_shape_mnkl;
auto [m, n, k, l] = args.tile_coord_mnkl;
gmem_ptr ptr_aux = make_gmem_ptr(subbyte_iterator<cutlass::uint1b_t const>(params.ptr_aux));
Tensor mAux = make_tensor(ptr_aux, make_layout(make_shape(M,N,L), params.dAux)); // (M,N,L)
Tensor gAux = local_tile(mAux, take<0,2>(args.tile_shape_mnk), make_coord(m,n,l)); // (CTA_M,CTA_N)
Tensor tC_gAux = sm90_partition_for_epilogue<ReferenceSrc>( // (CPY,CPY_M,CPY_N,EPI_M,EPI_N)
gAux, args.epi_tile, args.tiled_copy, args.thread_idx);
// If byte-unaligned vectorization, store in registers as uint32_t to reduce redundant pack+unpack instruction sequences
constexpr int V = decltype(max_common_vector(tC_gAux.layout(), make_layout(tC_gAux.shape())))::value;
Tensor tC_rAux = [&] () {
if constexpr (V % 8 != 0) {
return make_tensor<uint32_t>(take<0,3>(shape(tC_gAux))); // (CPY,CPY_M,CPY_N)
} else {
return make_tensor<cutlass::uint1b_t>(shape(tC_gAux)); // (CPY,CPY_M,CPY_N,EPI_M,EPI_N)
}
}();
if constexpr (EnableNullptr) {
if (params.ptr_aux == nullptr) {
fill(tC_rAux, params.null_default);
}
}
return ConsumerStoreCallbacks(cute::move(tC_rAux), cute::move(tC_gAux), args.residue_mn, params);
}
};
// dReLU specialization
template<
class ElementOutput,
class ElementCompute,
FloatRoundStyle RoundStyle
>
struct Sm90Compute<
cutlass::epilogue::thread::dReLU,
ElementOutput,
ElementCompute,
RoundStyle
> : Sm90VisitorImpl<> {
using Sm90VisitorImpl<>::Sm90VisitorImpl;
struct ConsumerStoreCallbacks : EmptyConsumerStoreCallbacks {
template <typename ElementAccumulator, typename ElementInput, typename ElementAux, int FragmentSize>
CUTLASS_DEVICE Array<ElementOutput, FragmentSize>
visit(Array<ElementAccumulator, FragmentSize> const& frg_acc, int epi_v, int epi_m, int epi_n,
Array<ElementInput , FragmentSize> const& frg_input,
Array<ElementAux , FragmentSize> const& frg_aux) {
using ConvertInput = NumericArrayConverter<ElementCompute, ElementInput, FragmentSize, RoundStyle>;
using ComputeOutput = cutlass::epilogue::thread::dReLU<Array<ElementCompute, FragmentSize>>;
using ConvertOutput = NumericArrayConverter<ElementOutput, ElementCompute, FragmentSize, RoundStyle>;
ConvertInput convert_input{};
ComputeOutput compute_output{};
ConvertOutput convert_output{};
return convert_output(compute_output(convert_input(frg_input), frg_aux)); // don't convert frg_aux for dReLU
}
};
template <
bool ReferenceSrc, // do register tensors reference the src or dst layout of the tiled copy
class... Args
>
CUTLASS_DEVICE auto
get_consumer_store_callbacks(ConsumerStoreArgs<Args...> const& args) {
return ConsumerStoreCallbacks();
}
};
/////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace cutlass::epilogue::fusion
@@ -71,23 +71,10 @@ struct Sm90AccFetch : Sm90VisitorImpl<> {
template <
bool ReferenceSrc, // do register tensors reference the src or dst layout of the tiled copy
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class EpilogueTile,
class TiledCopy,
class SrcTensor
class... Args
>
CUTLASS_DEVICE auto
get_consumer_store_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
TiledCopy tiled_copy,
int thread_idx,
SrcTensor const& tCrC) {
get_consumer_store_callbacks(ConsumerStoreArgs<Args...> const& args) {
return ConsumerStoreCallbacks{};
}
};
@@ -131,24 +118,12 @@ struct Sm90SrcFetch : Sm90VisitorImpl<> {
template <
bool ReferenceSrc, // do register tensors reference the src or dst layout of the tiled copy
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class EpilogueTile,
class TiledCopy,
class SrcTensor
class... Args
>
CUTLASS_DEVICE auto
get_consumer_store_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
TiledCopy tiled_copy,
int thread_idx,
SrcTensor const& tCrC) {
get_consumer_store_callbacks(ConsumerStoreArgs<Args...> const& args) {
return ConsumerStoreCallbacks(tCrC);
return ConsumerStoreCallbacks(args.tCrC);
}
};
@@ -223,6 +198,18 @@ struct Sm90AuxLoad {
return Params{tma_load_aux, args.null_default, use_default};
}
template <class ProblemShape>
static size_t
get_workspace_size(ProblemShape const& problem_shape, Arguments const& args) {
return 0;
}
template <class ProblemShape>
static cutlass::Status
initialize_workspace(ProblemShape const& problem_shape, Arguments const& args, void* workspace, cudaStream_t stream) {
return cutlass::Status::kSuccess;
}
CUTLASS_HOST_DEVICE
Sm90AuxLoad() { }
@@ -277,33 +264,23 @@ struct Sm90AuxLoad {
}
};
template <
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL
>
template <class... Args>
CUTLASS_DEVICE auto
get_producer_load_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
int thread_idx) {
get_producer_load_callbacks(ProducerLoadArgs<Args...> const& args) {
auto [M, N, K, L] = problem_shape_mnkl;
auto [m, n, k, l] = tile_coord_mnkl;
auto [M, N, K, L] = args.problem_shape_mnkl;
auto [m, n, k, l] = args.tile_coord_mnkl;
Tensor mAux = params_ptr->tma_load_aux.get_tma_tensor(make_shape(M,N,L)); // (M,N,L)
Tensor gAux = local_tile(mAux, take<0,2>(tile_shape_mnk), make_coord(m,n,l)); // (CTA_M,CTA_N)
Tensor gAux = local_tile(mAux, take<0,2>(args.tile_shape_mnk), make_coord(m,n,l)); // (CTA_M,CTA_N)
Tensor gAux_epi = local_tile(gAux, epi_tile, _); // (EPI_TILE_M,EPI_TILE_N,EPI_M,EPI_N)
Tensor gAux_epi = flat_divide(gAux, args.epi_tile); // (EPI_TILE_M,EPI_TILE_N,EPI_M,EPI_N)
Tensor sAux_epi = make_tensor(make_smem_ptr(smem_aux), SmemLayout{}); // (EPI_TILE_M,EPI_TILE_N,PIPE)
ThrCopy thrblk_g2s = params_ptr->tma_load_aux.get_slice(_0{});
Tensor bGS_gAux = thrblk_g2s.partition_S(gAux_epi); // (TMA,TMA_M,TMA_N,EPI_M,EPI_N)
Tensor bGS_sAux = thrblk_g2s.partition_D(sAux_epi); // (TMA,TMA_M,TMA_N,PIPE)
return ProducerLoadCallbacks(
cute::move(bGS_gAux), cute::move(bGS_sAux), params_ptr);
return ProducerLoadCallbacks(cute::move(bGS_gAux), cute::move(bGS_sAux), params_ptr);
}
template <class RTensor, class TiledS2R, class STensorS2R>
@@ -321,7 +298,7 @@ struct Sm90AuxLoad {
Params const* params_ptr;
CUTLASS_DEVICE void
step_begin(int epi_m, int epi_n, int load_iteration, bool is_producer_load_needed) {
previsit(int epi_m, int epi_n, int load_iteration, bool is_producer_load_needed) {
if constexpr (EnableNullptr) {
if (params_ptr->use_default) {
fill(tC_rAux, params_ptr->null_default);
@@ -347,35 +324,24 @@ struct Sm90AuxLoad {
template <
bool ReferenceSrc, // do register tensors reference the src or dst layout of the tiled copy
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class TiledCopy,
class SrcTensor
class... Args
>
CUTLASS_DEVICE auto
get_consumer_store_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
TiledCopy tiled_copy,
int thread_idx,
SrcTensor const& tCrC) {
get_consumer_store_callbacks(ConsumerStoreArgs<Args...> const& args) {
auto [M, N, K, L] = problem_shape_mnkl;
auto [M, N, K, L] = args.problem_shape_mnkl;
Tensor mAux = params_ptr->tma_load_aux.get_tma_tensor(make_shape(M,N,L)); // (M,N,L)
Tensor tC_gAux = sm90_partition_for_epilogue<ReferenceSrc>( // (CPY,CPY_M,CPY_N,EPI_M,EPI_N)
mAux, tile_shape_mnk, tile_coord_mnkl, epi_tile, tiled_copy, thread_idx);
mAux, args.tile_shape_mnk, args.tile_coord_mnkl, args.epi_tile, args.tiled_copy, args.thread_idx);
Tensor tC_rAux = make_tensor<Element>(take<0,3>(shape(tC_gAux))); // (CPY,CPY_M,CPY_N)
auto tiled_s2r = conditional_return<ReferenceSrc>(
make_tiled_copy_S(Copy_Atom<CopyOpS2R,Element>{}, tiled_copy),
make_tiled_copy_D(Copy_Atom<CopyOpS2R,Element>{}, tiled_copy)
make_tiled_copy_S(Copy_Atom<CopyOpS2R,Element>{}, args.tiled_copy),
make_tiled_copy_D(Copy_Atom<CopyOpS2R,Element>{}, args.tiled_copy)
);
Tensor sAux_epi = cute::as_position_independent_swizzle_tensor(
make_tensor(make_smem_ptr(smem_aux), SmemLayout{})); // (EPI_TILE_M,EPI_TILE_N,PIPE)
auto tSR_sAux = tiled_s2r.get_slice(thread_idx).partition_S(sAux_epi); // (S2R,S2R_M,S2R_N,PIPE)
auto tSR_sAux = tiled_s2r.get_slice(args.thread_idx).partition_S(sAux_epi); // (S2R,S2R_M,S2R_N,PIPE)
return ConsumerStoreCallbacks(cute::move(tC_rAux), tiled_s2r, cute::move(tSR_sAux), params_ptr);
@@ -400,7 +366,7 @@ struct Sm90ScalarBroadcast {
static_assert(
(cute::is_same_v<StrideMNL, Stride<_0,_0, _0>>) || // scalar broadcast, e.g. alpha
(cute::is_same_v<StrideMNL, Stride<_0,_0, _1>>) || // batched scalar broadcast, e.g. per-batch alpha
(cute::is_same_v<StrideMNL, Stride<_0,_0,int>>));
(cute::is_same_v<StrideMNL, Stride<_0,_0,int>>));
struct SharedStorage { };
@@ -418,6 +384,18 @@ struct Sm90ScalarBroadcast {
return args;
}
template <class ProblemShape>
static size_t
get_workspace_size(ProblemShape const& problem_shape, Arguments const& args) {
return 0;
}
template <class ProblemShape>
static cutlass::Status
initialize_workspace(ProblemShape const& problem_shape, Arguments const& args, void* workspace, cudaStream_t stream) {
return cutlass::Status::kSuccess;
}
CUTLASS_DEVICE bool
is_producer_load_needed() const {
return false;
@@ -443,24 +421,14 @@ struct Sm90ScalarBroadcast {
Element scalar;
Params const* params_ptr;
template <
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class EpilogueTile
>
template <class... Args>
CUTLASS_DEVICE auto
get_producer_load_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
int thread_idx) {
get_producer_load_callbacks(ProducerLoadArgs<Args...> const& args) {
// Get the scalar for batched broadcast
if constexpr (
cute::is_same_v<StrideMNL, Stride<_0,_0,_1>> ||
cute::is_same_v<StrideMNL, Stride<_0,_0,_1>> ||
cute::is_same_v<StrideMNL, Stride<_0,_0,int>>) {
auto [m_coord, n_coord, k_coord, l_coord] = tile_coord_mnkl;
auto [m_coord, n_coord, k_coord, l_coord] = args.tile_coord_mnkl;
update_scalar(l_coord);
}
@@ -487,28 +455,16 @@ struct Sm90ScalarBroadcast {
template <
bool ReferenceSrc, // do register tensors reference the src or dst layout of the tiled copy
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class EpilogueTile,
class TiledCopy,
class SrcTensor
class... Args
>
CUTLASS_DEVICE auto
get_consumer_store_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
TiledCopy tiled_copy,
int thread_idx,
SrcTensor const& tCrC) {
get_consumer_store_callbacks(ConsumerStoreArgs<Args...> const& args) {
// Get the scalar for batched broadcast
if constexpr (
cute::is_same_v<StrideMNL, Stride<_0,_0,_1>> ||
cute::is_same_v<StrideMNL, Stride<_0,_0,_1>> ||
cute::is_same_v<StrideMNL, Stride<_0,_0,int>>) {
auto [m_coord, n_coord, k_coord, l_coord] = tile_coord_mnkl;
auto [m_coord, n_coord, k_coord, l_coord] = args.tile_coord_mnkl;
update_scalar(l_coord);
}
@@ -579,6 +535,18 @@ struct Sm90RowBroadcast {
return args;
}
template <class ProblemShape>
static size_t
get_workspace_size(ProblemShape const& problem_shape, Arguments const& args) {
return 0;
}
template <class ProblemShape>
static cutlass::Status
initialize_workspace(ProblemShape const& problem_shape, Arguments const& args, void* workspace, cudaStream_t stream) {
return cutlass::Status::kSuccess;
}
CUTLASS_HOST_DEVICE
Sm90RowBroadcast() { }
@@ -633,29 +601,19 @@ struct Sm90RowBroadcast {
}
};
template <
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class EpilogueTile
>
template <class... Args>
CUTLASS_DEVICE auto
get_producer_load_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
int thread_idx) {
get_producer_load_callbacks(ProducerLoadArgs<Args...> const& args) {
auto [M, N, K, L] = problem_shape_mnkl;
auto [m, n, k, l] = tile_coord_mnkl;
auto [M, N, K, L] = args.problem_shape_mnkl;
auto [m, n, k, l] = args.tile_coord_mnkl;
Tensor mRow = make_tensor(make_gmem_ptr(params.ptr_row), make_shape(M,N,L), params.dRow);
Tensor gRow = local_tile(mRow, take<0,2>(tile_shape_mnk), make_coord(m,n,l)); // (CTA_M,CTA_N)
Tensor gRow = local_tile(mRow, take<0,2>(args.tile_shape_mnk), make_coord(m,n,l)); // (CTA_M,CTA_N)
Tensor sRow = make_tensor(make_smem_ptr(smem_row), // (CTA_M,CTA_N,PIPE)
make_shape(size<0>(CtaTileShapeMNK{}), size<1>(CtaTileShapeMNK{}), Stages),
make_stride(_0{},_1{},size<1>(CtaTileShapeMNK{})));
constexpr int EpiTiles = size(shape_div(take<0,2>(tile_shape_mnk), epi_tile));
constexpr int EpiTiles = decltype(size(shape_div(take<0,2>(args.tile_shape_mnk), args.epi_tile)))::value;
return ProducerLoadCallbacks<EpiTiles, decltype(gRow), decltype(sRow)>(
cute::move(gRow), cute::move(sRow), params);
}
@@ -673,7 +631,7 @@ struct Sm90RowBroadcast {
Params const& params;
CUTLASS_DEVICE void
step_begin(int epi_m, int epi_n, int load_iteration, bool is_producer_load_needed) {
previsit(int epi_m, int epi_n, int load_iteration, bool is_producer_load_needed) {
if constexpr (EnableNullptr) {
if (params.ptr_row == nullptr) {
fill(tCrRow, params.null_default);
@@ -704,31 +662,19 @@ struct Sm90RowBroadcast {
template <
bool ReferenceSrc, // do register tensors reference the src or dst layout of the tiled copy
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class EpilogueTile,
class TiledCopy,
class SrcTensor
class... Args
>
CUTLASS_DEVICE auto
get_consumer_store_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
TiledCopy tiled_copy,
int thread_idx,
SrcTensor const& tCrC) {
get_consumer_store_callbacks(ConsumerStoreArgs<Args...> const& args) {
Tensor sRow = make_tensor(make_smem_ptr(smem_row), // (CTA_M,CTA_N,PIPE)
make_shape(size<0>(CtaTileShapeMNK{}), size<1>(CtaTileShapeMNK{}), Stages),
make_stride(_0{},_1{},size<1>(CtaTileShapeMNK{})));
Tensor tCsRow = sm90_partition_for_epilogue<ReferenceSrc>( // (CPY,CPY_M,CPY_N,EPI_M,EPI_N,PIPE)
sRow, epi_tile, tiled_copy, thread_idx);
sRow, args.epi_tile, args.tiled_copy, args.thread_idx);
Tensor tCrRow = make_tensor_like(take<0,3>(tCsRow)); // (CPY,CPY_M,CPY_N)
constexpr int EpiTiles = size(shape_div(take<0,2>(tile_shape_mnk), epi_tile));
constexpr int EpiTiles = decltype(size(shape_div(take<0,2>(args.tile_shape_mnk), args.epi_tile)))::value;
return ConsumerStoreCallbacks<EpiTiles, decltype(tCrRow), decltype(tCsRow)>(
cute::move(tCrRow), cute::move(tCsRow), params);
}
@@ -769,6 +715,18 @@ struct Sm90ColBroadcast {
return args;
}
template <class ProblemShape>
static size_t
get_workspace_size(ProblemShape const& problem_shape, Arguments const& args) {
return 0;
}
template <class ProblemShape>
static cutlass::Status
initialize_workspace(ProblemShape const& problem_shape, Arguments const& args, void* workspace, cudaStream_t stream) {
return cutlass::Status::kSuccess;
}
CUTLASS_DEVICE bool
is_producer_load_needed() const {
return false;
@@ -788,19 +746,9 @@ struct Sm90ColBroadcast {
Params params;
template <
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class EpilogueTile
>
template <class... Args>
CUTLASS_DEVICE auto
get_producer_load_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
int thread_idx) {
get_producer_load_callbacks(ProducerLoadArgs<Args...> const& args) {
return EmptyProducerLoadCallbacks{};
}
@@ -847,27 +795,15 @@ struct Sm90ColBroadcast {
template <
bool ReferenceSrc, // do register tensors reference the src or dst layout of the tiled copy
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class EpilogueTile,
class TiledCopy,
class SrcTensor
class... Args
>
CUTLASS_DEVICE auto
get_consumer_store_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
TiledCopy tiled_copy,
int thread_idx,
SrcTensor const& tCrC) {
get_consumer_store_callbacks(ConsumerStoreArgs<Args...> const& args) {
auto [M, N, K, L] = problem_shape_mnkl;
auto [M, N, K, L] = args.problem_shape_mnkl;
Tensor mCol = make_tensor(make_gmem_ptr(params.ptr_col), make_shape(M,N,L), params.dCol);
Tensor tCgCol = sm90_partition_for_epilogue<ReferenceSrc>( // (CPY,CPY_M,CPY_N,EPI_M,EPI_N)
mCol, tile_shape_mnk, tile_coord_mnkl, epi_tile, tiled_copy, thread_idx);
mCol, args.tile_shape_mnk, args.tile_coord_mnkl, args.epi_tile, args.tiled_copy, args.thread_idx);
Tensor tCrCol = make_tensor_like(tCgCol); // (CPY,CPY_M,CPY_N,EPI_M,EPI_N)
return ConsumerStoreCallbacks<decltype(tCgCol), decltype(tCrCol)>(
@@ -36,6 +36,7 @@
#pragma once
#include "cutlass/cutlass.h"
#include "cutlass/workspace.h"
#include "cute/tensor.hpp"
#include "sm90_visitor_tma_warpspecialized.hpp"
@@ -122,6 +123,18 @@ struct Sm90AuxStore {
return {tma_store_aux, is_nullptr};
}
template <class ProblemShape>
static size_t
get_workspace_size(ProblemShape const& problem_shape, Arguments const& args) {
return 0;
}
template <class ProblemShape>
static cutlass::Status
initialize_workspace(ProblemShape const& problem_shape, Arguments const& args, void* workspace, cudaStream_t stream) {
return cutlass::Status::kSuccess;
}
CUTLASS_HOST_DEVICE
Sm90AuxStore() { }
@@ -143,18 +156,9 @@ struct Sm90AuxStore {
return false;
}
template <
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL
>
template <class... Args>
CUTLASS_DEVICE auto
get_producer_load_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
int thread_idx) {
get_producer_load_callbacks(ProducerLoadArgs<Args...> const& args) {
return EmptyProducerLoadCallbacks{};
}
@@ -202,7 +206,7 @@ struct Sm90AuxStore {
}
CUTLASS_DEVICE void
step_next(int epi_m, int epi_n, int store_iteration, bool issue_smem_store) {
postvisit(int epi_m, int epi_n, int store_iteration, bool issue_smem_store) {
if constexpr (EnableNullptr) {
if (params_ptr->is_nullptr) {
return;
@@ -219,7 +223,7 @@ struct Sm90AuxStore {
}
CUTLASS_DEVICE void
step_end(int epi_m, int epi_n, int store_iteration, bool issue_tma_store) {
step(int epi_m, int epi_n, int store_iteration, bool issue_tma_store) {
if constexpr (EnableNullptr) {
if (params_ptr->is_nullptr) {
return;
@@ -236,40 +240,29 @@ struct Sm90AuxStore {
template <
bool ReferenceSrc, // do register tensors reference the src or dst layout of the tiled copy
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class TiledCopy,
class SrcTensor
class... Args
>
CUTLASS_DEVICE auto
get_consumer_store_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
TiledCopy tiled_copy,
int thread_idx,
SrcTensor const& tCrC) {
get_consumer_store_callbacks(ConsumerStoreArgs<Args...> const& args) {
auto [M, N, K, L] = problem_shape_mnkl;
auto [m, n, k, l] = tile_coord_mnkl;
auto [M, N, K, L] = args.problem_shape_mnkl;
auto [m, n, k, l] = args.tile_coord_mnkl;
Tensor mAux = params_ptr->tma_store_aux.get_tma_tensor(make_shape(M,N,L)); // (M,N,L)
Tensor gAux = local_tile(mAux, take<0,2>(tile_shape_mnk), make_coord(m,n,l)); // (CTA_M,CTA_N)
Tensor gAux = local_tile(mAux, take<0,2>(args.tile_shape_mnk), make_coord(m,n,l)); // (CTA_M,CTA_N)
Tensor tC_gAux = sm90_partition_for_epilogue<ReferenceSrc>( // (CPY,CPY_M,CPY_N,EPI_M,EPI_N)
gAux, epi_tile, tiled_copy, thread_idx);
gAux, args.epi_tile, args.tiled_copy, args.thread_idx);
Tensor tC_rAux = make_tensor<Element>(take<0,3>(shape(tC_gAux))); // (CPY,CPY_M,CPY_N)
Tensor sAux_epi = cute::as_position_independent_swizzle_tensor(
make_tensor(make_smem_ptr(smem_aux), SmemLayout{})); // (EPI_TILE_M,EPI_TILE_N,PIPE)
Tensor gAux_epi = local_tile(gAux, epi_tile, _); // (EPI_TILE_M,EPI_TILE_N,EPI_M,EPI_N)
Tensor gAux_epi = flat_divide(gAux, args.epi_tile); // (EPI_TILE_M,EPI_TILE_N,EPI_M,EPI_N)
auto tiled_r2s = conditional_return<ReferenceSrc>(
make_tiled_copy_S(Copy_Atom<CopyOpR2S,Element>{}, tiled_copy),
make_tiled_copy_D(Copy_Atom<CopyOpR2S,Element>{}, tiled_copy)
make_tiled_copy_S(Copy_Atom<CopyOpR2S,Element>{}, args.tiled_copy),
make_tiled_copy_D(Copy_Atom<CopyOpR2S,Element>{}, args.tiled_copy)
);
auto tRS_sAux = tiled_r2s.get_slice(thread_idx).partition_D(sAux_epi); // (R2S,R2S_M,R2S_N,PIPE)
auto tRS_sAux = tiled_r2s.get_slice(args.thread_idx).partition_D(sAux_epi); // (R2S,R2S_M,R2S_N,PIPE)
ThrCopy thrblk_s2g = params_ptr->tma_store_aux.get_slice(_0{});
Tensor bSG_sAux = thrblk_s2g.partition_S(sAux_epi); // (TMA,TMA_M,TMA_N,PIPE)
@@ -294,7 +287,7 @@ struct Sm90AuxStore {
// Scalar reduction
template <
template <class> class RegReduceFn,
template <class> class AtomicReduceFn,
template <class> class GmemReduceFn,
class ElementOutput,
class ElementCompute,
FloatRoundStyle RoundStyle,
@@ -302,10 +295,15 @@ template <
bool EnableNullptr = true // Noop on nullptr params
>
struct Sm90ScalarReduction {
private:
static_assert(
(cute::is_same_v<StrideMNL, Stride<_0,_0, _0>>) || // scalar reduction, e.g. tensor max element
(cute::is_same_v<StrideMNL, Stride<_0,_0, _1>>) || // batched scalar reduction, e.g. per-batch max element
(cute::is_same_v<StrideMNL, Stride<_0,_0,int>>));
(cute::is_same_v<StrideMNL, Stride<_0,_0,int>>));
static constexpr bool IsAtomic = is_atomic<GmemReduceFn<ElementCompute>>::value;
static_assert(IsAtomic, "non-atomic scalar reduction not supported yet");
public:
struct SharedStorage { };
struct Arguments {
@@ -322,6 +320,26 @@ struct Sm90ScalarReduction {
return args;
}
template <class ProblemShape>
static size_t
get_workspace_size(ProblemShape const& problem_shape, Arguments const& args) {
return 0;
}
template <class ProblemShape>
static cutlass::Status
initialize_workspace(ProblemShape const& problem_shape, Arguments const& args, void* workspace, cudaStream_t stream) {
if constexpr (IsAtomic) {
auto [M, N, K, L] = problem_shape;
Layout mScalar_layout = make_layout(make_shape(M,N,L), args.dScalar);
if (args.ptr_scalar != nullptr) {
return fill_workspace(args.ptr_scalar, ElementOutput(args.reduction_identity), cosize(mScalar_layout), stream);
}
}
return cutlass::Status::kSuccess;
}
CUTLASS_DEVICE bool
is_producer_load_needed() const {
return false;
@@ -341,19 +359,9 @@ struct Sm90ScalarReduction {
Params const params;
template <
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class EpilogueTile
>
template <class... Args>
CUTLASS_DEVICE auto
get_producer_load_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
int thread_idx) {
get_producer_load_callbacks(ProducerLoadArgs<Args...> const& args) {
return EmptyProducerLoadCallbacks{};
}
@@ -362,18 +370,18 @@ struct Sm90ScalarReduction {
CUTLASS_DEVICE
ConsumerStoreCallbacks(
int l_coord,
CTensor&& tCcScalar,
CTensor tCcScalar,
ResidueMN residue_mn,
Params const& params)
: scalar(params.reduction_identity),
l_coord(l_coord),
tCcScalar(cute::forward<CTensor>(tCcScalar)),
tCcScalar(tCcScalar),
residue_mn(residue_mn),
params(params) {}
ElementCompute scalar;
int l_coord;
CTensor tCcScalar;
CTensor tCcScalar; // (CPY,CPY_M,CPY_N,EPI_M,EPI_N)
ResidueMN residue_mn;
Params params;
@@ -393,10 +401,11 @@ struct Sm90ScalarReduction {
ReduceInput reduce_input{};
Array frg_I = convert_input(frg_input);
Tensor tCcScalar_mn = tCcScalar(_,_,_,epi_m,epi_n);
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < FragmentSize; ++i) {
if (elem_less(tCcScalar(epi_v * FragmentSize + i), residue_mn)) {
if (elem_less(tCcScalar_mn(epi_v * FragmentSize + i), residue_mn)) {
scalar = reduce_input(scalar, frg_I[i]);
}
}
@@ -413,7 +422,7 @@ struct Sm90ScalarReduction {
}
using ConvertI = NumericConverter<ElementOutput, ElementCompute, RoundStyle>;
using ReduceInput = AtomicReduceFn<ElementOutput>;
using ReduceInput = GmemReduceFn<ElementOutput>;
ConvertI convert_I{};
ReduceInput reduce_input{};
@@ -426,36 +435,12 @@ struct Sm90ScalarReduction {
template <
bool ReferenceSrc, // do register tensors reference the src or dst layout of the tiled copy
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class EpilogueTile,
class TiledCopy,
class SrcTensor
class... Args
>
CUTLASS_DEVICE auto
get_consumer_store_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
TiledCopy tiled_copy,
int thread_idx,
SrcTensor const& tCrC) {
int l_coord = static_cast<int>(get<3>(tile_coord_mnkl));
// Compute tile residues and coordinate tensors for predication
auto [M, N, K, L] = problem_shape_mnkl;
auto [m, n, k, l] = tile_coord_mnkl;
auto residue_mn = make_coord(
M - static_cast<int>(m) * size<0>(tile_shape_mnk),
N - static_cast<int>(n) * size<1>(tile_shape_mnk)
);
Tensor cScalar = make_identity_tensor(take<0,2>(tile_shape_mnk));
Tensor tCcScalar = sm90_partition_for_epilogue<ReferenceSrc>(cScalar, epi_tile, tiled_copy, thread_idx);
return ConsumerStoreCallbacks(l_coord, cute::move(tCcScalar), residue_mn, params);
get_consumer_store_callbacks(ConsumerStoreArgs<Args...> const& args) {
return ConsumerStoreCallbacks<decltype(args.tCcD), decltype(args.residue_mn)>(
get<3>(args.tile_coord_mnkl), args.tCcD, args.residue_mn, params);
}
};
@@ -466,7 +451,7 @@ struct Sm90ScalarReduction {
// Row vector reduction
template <
template <class> class RegReduceFn,
template <class> class AtomicReduceFn,
template <class> class GmemReduceFn,
int Stages,
class CtaTileShapeMNK,
class ElementOutput,
@@ -477,12 +462,16 @@ template <
bool EnableNullptr = true // Noop on nullptr params
>
struct Sm90RowReduction {
private:
static_assert(Stages == 0, "Smem usage not supported yet");
static_assert(Alignment * sizeof_bits_v<ElementOutput> % 128 == 0, "sub-16B alignment not supported yet");
static_assert(
(cute::is_same_v<StrideMNL, Stride<_0,_1, _0>>) || // row vector reduction, e.g. per-col sum over all batches
(cute::is_same_v<StrideMNL, Stride<_0,_1,int>>)); // batched row vector reduction, e.g. per-col sum per batch
static constexpr bool IsAtomic = is_atomic<GmemReduceFn<ElementCompute>>::value;
static_assert(IsAtomic, "non-atomic row reduction not supported yet");
public:
struct SharedStorage { };
struct Arguments {
@@ -499,6 +488,26 @@ struct Sm90RowReduction {
return args;
}
template <class ProblemShape>
static size_t
get_workspace_size(ProblemShape const& problem_shape, Arguments const& args) {
return 0;
}
template <class ProblemShape>
static cutlass::Status
initialize_workspace(ProblemShape const& problem_shape, Arguments const& args, void* workspace, cudaStream_t stream) {
if constexpr (IsAtomic) {
auto [M, N, K, L] = problem_shape;
Layout mRow_layout = make_layout(make_shape(M,N,L), args.dRow);
if (args.ptr_row != nullptr) {
return fill_workspace(args.ptr_row, ElementOutput(args.reduction_identity), cosize(mRow_layout), stream);
}
}
return cutlass::Status::kSuccess;
}
CUTLASS_DEVICE bool
is_producer_load_needed() const {
return false;
@@ -518,19 +527,9 @@ struct Sm90RowReduction {
Params params;
template <
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class EpilogueTile
>
template <class... Args>
CUTLASS_DEVICE auto
get_producer_load_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
int thread_idx) {
get_producer_load_callbacks(ProducerLoadArgs<Args...> const& args) {
return EmptyProducerLoadCallbacks{};
}
@@ -540,12 +539,12 @@ struct Sm90RowReduction {
ConsumerStoreCallbacks(
RTensor&& tCrRow,
GTensor&& tCgRow,
CTensor&& tCcRow,
CTensor tCcRow,
ResidueMN residue_mn,
Params const& params)
: tCrRow(cute::forward<RTensor>(tCrRow)),
tCgRow(cute::forward<GTensor>(tCgRow)),
tCcRow(cute::forward<CTensor>(tCcRow)),
tCcRow(tCcRow),
residue_mn(residue_mn),
params(params) {}
@@ -580,7 +579,7 @@ struct Sm90RowReduction {
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < FragmentSize; ++i) {
if (elem_less(tCcRow_mn(i), residue_mn)) {
if (elem_less(tCcRow_mn(epi_v * FragmentSize + i), residue_mn)) {
ElementCompute& tCrRow_vmn = tCrRow_mn(epi_v * FragmentSize + i);
tCrRow_vmn = reduce_input(tCrRow_vmn, frg_I[i]);
}
@@ -590,7 +589,7 @@ struct Sm90RowReduction {
}
CUTLASS_DEVICE void
step_end(int epi_m, int epi_n, int store_iteration, bool issue_tma_store) {
step(int epi_m, int epi_n, int store_iteration, bool issue_tma_store) {
if constexpr (EnableNullptr) {
if (params.ptr_row == nullptr) {
return;
@@ -599,7 +598,7 @@ struct Sm90RowReduction {
if (epi_m == size<3>(tCrRow)-1) { // assumes M-major subtile loop
using ConvertI = NumericConverter<ElementOutput, ElementCompute, RoundStyle>;
using ReduceInput = AtomicReduceFn<ElementOutput>;
using ReduceInput = GmemReduceFn<ElementOutput>;
ConvertI convert_I{};
ReduceInput reduce_input{};
@@ -616,7 +615,7 @@ struct Sm90RowReduction {
for (int i = 0; i < size(tCrRow_flt); ++i) {
// partially OOB in M must still issue gmem reduction, so only consider residue_n
// in case last epi tile in column is fully OOB in M and CTA tile is partially OOB in M
if (residue_n > get<1>(tCcRow_flt(i)) &&
if (residue_n > get<1>(tCcRow_flt(i)) &&
// fully OOB in M does not need to issue gmem reduction, skip
residue_m > 0) {
reduce_input(&tCgRow_flt(i), convert_I(tCrRow_flt(i)));
@@ -632,41 +631,20 @@ struct Sm90RowReduction {
template <
bool ReferenceSrc, // do register tensors reference the src or dst layout of the tiled copy
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class EpilogueTile,
class TiledCopy,
class SrcTensor
class... Args
>
CUTLASS_DEVICE auto
get_consumer_store_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
TiledCopy tiled_copy,
int thread_idx,
SrcTensor const& tCrC) {
get_consumer_store_callbacks(ConsumerStoreArgs<Args...> const& args) {
auto [M, N, K, L] = problem_shape_mnkl;
auto [M, N, K, L] = args.problem_shape_mnkl;
Tensor mRow = make_tensor(make_gmem_ptr(params.ptr_row), make_shape(M,N,L), params.dRow);
Tensor tCgRow = sm90_partition_for_epilogue<ReferenceSrc>( // (CPY,CPY_M,CPY_N,EPI_M,EPI_N)
mRow, tile_shape_mnk, tile_coord_mnkl, epi_tile, tiled_copy, thread_idx);
mRow, args.tile_shape_mnk, args.tile_coord_mnkl, args.epi_tile, args.tiled_copy, args.thread_idx);
Tensor tCrRow = make_tensor_like<ElementCompute>(tCgRow(_,_,_,_,_0{})); // (CPY,CPY_M,CPY_N,EPI_M)
fill(tCrRow, params.reduction_identity);
// Compute tile residues and coordinate tensors for predication
auto [m, n, k, l] = tile_coord_mnkl;
auto residue_mn = make_coord(
M - static_cast<int>(m) * size<0>(tile_shape_mnk),
N - static_cast<int>(n) * size<1>(tile_shape_mnk)
);
Tensor cRow = make_identity_tensor(take<0,2>(tile_shape_mnk));
Tensor tCcRow = sm90_partition_for_epilogue<ReferenceSrc>(cRow, epi_tile, tiled_copy, thread_idx);
return ConsumerStoreCallbacks<decltype(tCrRow),decltype(tCgRow),decltype(tCcRow),decltype(residue_mn)>(
cute::move(tCrRow), cute::move(tCgRow), cute::move(tCcRow), residue_mn, params);
return ConsumerStoreCallbacks<decltype(tCrRow),decltype(tCgRow),decltype(args.tCcD),decltype(args.residue_mn)>(
cute::move(tCrRow), cute::move(tCgRow), args.tCcD, args.residue_mn, params);
}
};
@@ -675,7 +653,7 @@ struct Sm90RowReduction {
// Col vector reduction
template <
template <class> class RegReduceFn,
template <class> class AtomicReduceFn,
template <class> class GmemReduceFn,
int Stages,
class CtaTileShapeMNK,
class ElementOutput,
@@ -683,29 +661,110 @@ template <
FloatRoundStyle RoundStyle,
class StrideMNL = Stride<_1,_0,_0>,
int Alignment = 128 / sizeof_bits_v<ElementOutput>,
bool EnableNullptr = true // Noop on nullptr params
bool EnableNullptr = true, // Noop on nullptr params
// If this is false, ptr_col is assumed to point to a compact m-major (round_nearest(M,CTA_M), ceil_div(N,CTA_N), L)
// tensor of ElementCompute. It is the user's responsibility to reduce this to a (M, L) tensor of ElementOutput
bool FinalReduction = true
>
struct Sm90ColReduction {
private:
static_assert(Stages == 0, "Smem usage not supported yet");
static_assert(Alignment * sizeof_bits_v<ElementOutput> % 128 == 0, "sub-16B alignment not supported yet");
static_assert(
(cute::is_same_v<StrideMNL, Stride<_1,_0, _0>>) || // col vector reduction, e.g. per-row sum over all batches
(cute::is_same_v<StrideMNL, Stride<_1,_0,int>>)); // batched col vector reduction, e.g. per-row sum per batch
static constexpr bool IsAtomic = is_atomic<GmemReduceFn<ElementCompute>>::value;
static_assert(not (IsAtomic && not FinalReduction), "atomic reduction must be final");
public:
struct SharedStorage { };
struct Arguments {
ElementOutput* ptr_col = nullptr;
void* ptr_col = nullptr; // ElementOutput* if FinalReduction, else ElementCompute*
ElementCompute reduction_identity = 0;
StrideMNL dCol = {};
};
using Params = Arguments;
struct Params {
void* ptr_col = nullptr;
ElementCompute reduction_identity = 0;
StrideMNL dCol = {};
ElementCompute* reduction_buffer = nullptr;
int* tile_counters = nullptr;
};
template <class ProblemShape>
static constexpr Params
to_underlying_arguments(ProblemShape const& problem_shape, Arguments const& args, void* workspace) {
return args;
ElementCompute* reduction_buffer;
int* tile_counters;
if constexpr (IsAtomic) {
reduction_buffer = nullptr;
tile_counters = nullptr;
}
else if constexpr (not FinalReduction) {
reduction_buffer = reinterpret_cast<ElementCompute*>(args.ptr_col);
tile_counters = nullptr;
}
else {
auto [M, N, K, L] = problem_shape;
auto [tile_M, tile_N, tile_K] = CtaTileShapeMNK{};
size_t tile_counters_offset = product(ceil_div(make_shape(M,N,L), make_shape(tile_M, tile_N))) * tile_M * sizeof(ElementCompute);
tile_counters_offset = round_nearest(tile_counters_offset, sizeof(int));
reduction_buffer = reinterpret_cast<ElementCompute*>(workspace);
tile_counters = reinterpret_cast<int*>(reinterpret_cast<uint8_t*>(workspace) + tile_counters_offset);
}
return {
args.ptr_col,
args.reduction_identity,
args.dCol,
reduction_buffer,
tile_counters
};
}
template <class ProblemShape>
static size_t
get_workspace_size(ProblemShape const& problem_shape, Arguments const& args) {
if constexpr (IsAtomic || not FinalReduction) {
return 0;
}
size_t workspace_size = 0;
auto [M, N, K, L] = problem_shape;
auto [tile_M, tile_N, tile_K] = CtaTileShapeMNK{};
// Increment by size of reduction buffer
workspace_size += product(ceil_div(make_shape(M,N,L), make_shape(tile_M, tile_N))) * tile_M * sizeof(ElementCompute);
// Align and increment by size of tile counters
workspace_size = round_nearest(workspace_size, sizeof(int));
workspace_size += cute::ceil_div(M, tile_M) * sizeof(int);
return workspace_size;
}
template <class ProblemShape>
static cutlass::Status
initialize_workspace(ProblemShape const& problem_shape, Arguments const& args, void* workspace, cudaStream_t stream) {
if constexpr (IsAtomic) {
auto [M, N, K, L] = problem_shape;
Layout mCol_layout = make_layout(make_shape(M,N,L), args.dCol);
if (args.ptr_col != nullptr) {
return fill_workspace(args.ptr_col, ElementOutput(args.reduction_identity), cosize(mCol_layout), stream);
}
return Status::kSuccess;
}
auto [M, N, K, L] = problem_shape;
auto [tile_M, tile_N, tile_K] = CtaTileShapeMNK{};
size_t tile_counters_offset = product(ceil_div(make_shape(M,N,L), make_shape(tile_M, tile_N))) * tile_M * sizeof(ElementCompute);
tile_counters_offset = round_nearest(tile_counters_offset, sizeof(int));
int* tile_counters = reinterpret_cast<int*>(reinterpret_cast<uint8_t*>(workspace) + tile_counters_offset);
size_t tile_counters_size = cute::ceil_div(M, tile_M) * sizeof(int);
return zero_workspace(tile_counters, tile_counters_size, stream);
}
CUTLASS_DEVICE bool
@@ -727,66 +786,48 @@ struct Sm90ColReduction {
Params params;
template <
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class EpilogueTile
>
template <class... Args>
CUTLASS_DEVICE auto
get_producer_load_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
int thread_idx) {
get_producer_load_callbacks(ProducerLoadArgs<Args...> const& args) {
return EmptyProducerLoadCallbacks{};
}
template<class RTensor, class GTensor, class CTensor, class ResidueMN>
template<class ArgsTuple>
struct ConsumerStoreCallbacks : EmptyConsumerStoreCallbacks {
CUTLASS_DEVICE
ConsumerStoreCallbacks(
RTensor&& tCrCol,
GTensor&& tCgCol,
CTensor&& tCcCol,
ResidueMN residue_mn,
Params const& params)
: tCrCol(cute::forward<RTensor>(tCrCol)),
tCgCol(cute::forward<GTensor>(tCgCol)),
tCcCol(cute::forward<CTensor>(tCcCol)),
residue_mn(residue_mn),
ConsumerStoreCallbacks(ArgsTuple&& args_tuple, Params const& params)
: args_tuple(cute::forward<ArgsTuple>(args_tuple)),
params(params) {}
RTensor tCrCol; // (CPY,CPY_M,CPY_N,EPI_M,EPI_N)
GTensor tCgCol; // (CPY,CPY_M,CPY_N,EPI_M,EPI_N)
CTensor tCcCol; // (CPY,CPY_M,CPY_N,EPI_M,EPI_N)
ResidueMN residue_mn;
ArgsTuple args_tuple;
Params const& params;
bool do_final_reduction = false;
template <typename ElementAccumulator, typename ElementInput, int FragmentSize>
CUTLASS_DEVICE auto
visit(Array<ElementAccumulator, FragmentSize> const& frg_acc, int epi_v, int epi_m, int epi_n,
Array<ElementInput, FragmentSize> const& frg_input) {
if constexpr (EnableNullptr) {
if (params.ptr_col == nullptr) {
return frg_input;
}
}
auto& [ref_src, tCrCol, tCcCol, gCol_l, cCol, gBuf_nl, sBuf_layout,
lane_layout_MN, lane_mn, warp_layout_MN, warp_mn,
tile_coord_mnkl, residue_mn, epi_tile, tiled_copy, thread_idx] = args_tuple;
Tensor tCrCol_mn = tCrCol(_,_,_,epi_m,epi_n);
Tensor tCcCol_mn = tCcCol(_,_,_,epi_m,epi_n);
using ConvertInput = NumericArrayConverter<ElementCompute, ElementInput, FragmentSize, RoundStyle>;
using ReduceInput = RegReduceFn<ElementCompute>;
ConvertInput convert_input{};
ReduceInput reduce_input{};
Array frg_I = convert_input(frg_input);
Tensor tCrCol_mn = tCrCol(_,_,_,epi_m,epi_n);
Tensor tCcCol_mn = tCcCol(_,_,_,epi_m,epi_n);
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < FragmentSize; ++i) {
if (elem_less(tCcCol_mn(i), residue_mn)) {
if (elem_less(tCcCol_mn(epi_v * FragmentSize + i), residue_mn)) {
ElementCompute& tCrCol_vmn = tCrCol_mn(epi_v * FragmentSize + i);
tCrCol_vmn = reduce_input(tCrCol_vmn, frg_I[i]);
}
@@ -795,71 +836,283 @@ struct Sm90ColReduction {
return frg_input;
}
template <class STensor, class SyncFn>
CUTLASS_DEVICE void
end() {
reduce(STensor&& smem_buffer, SyncFn const& sync_fn, int epi_m, int epi_n, bool is_last_iteration) {
if (not is_last_iteration) {
return;
}
auto& [ref_src, tCrCol, tCcCol, gCol_l, cCol, gBuf_nl, sBuf_layout,
lane_layout_MN, lane_mn, warp_layout_MN, warp_mn,
tile_coord_mnkl, residue_mn, epi_tile, tiled_copy, thread_idx] = args_tuple;
auto [m, n, k, l] = tile_coord_mnkl;
constexpr bool ReferenceSrc = decltype(ref_src)::value;
// Runtime nullptr is noop
if constexpr (EnableNullptr) {
if (params.ptr_col == nullptr) {
return;
}
}
using ConvertI = NumericConverter<ElementOutput, ElementCompute, RoundStyle>;
using ReduceInput = AtomicReduceFn<ElementOutput>;
ConvertI convert_I{};
ReduceInput reduce_input{};
// Filter so we don't issue redunant copies over stride-0 modes
Tensor tCrCol_flt = filter_zeros(tCrCol);
Tensor tCgCol_flt = filter_zeros(tCgCol);
Tensor tCcCol_flt = make_tensor(tCcCol.data(), make_layout(tCgCol_flt.shape(), tCcCol.stride()));
// fully OOB CTA in partially OOB cluster
if (not elem_less(cCol(_0{},_0{}), residue_mn)) {
return;
}
//
// 1. Warp shuffle reduction
//
using FragmentShuffle = Array<ElementCompute, sizeof(uint64_t) / sizeof(ElementCompute)>;
using ReduceShuffle = RegReduceFn<FragmentShuffle>;
ReduceShuffle reduce_shuffle{};
Tensor tCrCol_frg = recast<FragmentShuffle>(filter(tCrCol));
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < size(tCrCol_flt); ++i) {
if (elem_less(tCcCol_flt(i), residue_mn)) {
reduce_input(&tCgCol_flt(i), convert_I(tCrCol_flt(i)));
for (int reduction_cols = size<1>(lane_layout_MN) / 2; reduction_cols > 0; reduction_cols /= 2) {
CUTLASS_PRAGMA_UNROLL
for (int frg_idx = 0; frg_idx < size(tCrCol_frg); ++frg_idx) {
uint64_t frg_shfl = reinterpret_cast<uint64_t&>(tCrCol_frg(frg_idx));
frg_shfl = __shfl_down_sync(0xFFFFFFFF, frg_shfl, lane_layout_MN(_0{},reduction_cols));
tCrCol_frg(frg_idx) = reduce_shuffle(tCrCol_frg(frg_idx), reinterpret_cast<FragmentShuffle&>(frg_shfl));
}
}
bool is_reduced_lane = get<1>(lane_mn) == 0;
//
// 2. Atomic reduction
//
if constexpr (IsAtomic) {
// Filter so we don't issue redunant copies over stride-0 modes
Tensor tCrCol_flt = filter_zeros(tCrCol);
Tensor tCcCol_flt = make_tensor(tCcCol.data(), make_layout(tCrCol_flt.shape(), tCcCol.stride()));
Tensor tCgCol = sm90_partition_for_epilogue<ReferenceSrc>(gCol_l(_,_,l), epi_tile, tiled_copy, thread_idx);
Tensor tCgCol_flt = filter_zeros(tCgCol);
// NOTE: atomic reduction is performed in the output type
using ConvertOutput = NumericConverter<ElementOutput, ElementCompute, RoundStyle>;
using ReduceOutput = GmemReduceFn<ElementOutput>;
ConvertOutput convert_output{};
ReduceOutput reduce_output{};
if (is_reduced_lane) {
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < size(tCrCol_flt); ++i) {
if (elem_less(tCcCol_flt(i), residue_mn)) {
reduce_output(&tCgCol_flt(i), convert_output(tCrCol_flt(i)));
}
}
}
sync_fn();
}
//
// 2. One warp in N, skip threadblock smem reduction
//
else if constexpr (decltype(size<1>(warp_layout_MN))::value <= 1) {
// Dump warp reduction to gmem workspace
using ElementGmem = conditional_t<FinalReduction, ElementCompute volatile, ElementCompute>;
Tensor tCgBuf = sm90_partition_for_epilogue<ReferenceSrc>(gBuf_nl(_,_,n,l), epi_tile, tiled_copy, thread_idx);
if (is_reduced_lane) {
// Filter so we don't issue redunant copies over stride-0 modes
copy(filter(tCrCol), recast<ElementGmem>(filter(tCgBuf)));
}
sync_fn();
}
//
// 2. Multiple warps in N, do threadblock smem reduction
//
else {
Tensor sBuf = make_tensor(make_smem_ptr<ElementCompute>(raw_pointer_cast(smem_buffer.data())), sBuf_layout);
static_assert(decltype(cosize(sBuf.layout()))::value * sizeof(ElementCompute) <=
decltype(cosize(smem_buffer.layout()))::value * sizeof(typename remove_cvref_t<STensor>::value_type),
"smem reduction buffer not large enough, use a larger epilogue tile");
// Dump warp reduction to smem workspace
Tensor tCsBuf = sm90_partition_for_epilogue<ReferenceSrc>(sBuf(_,_,get<1>(warp_mn)), epi_tile, tiled_copy, thread_idx);
if (is_reduced_lane) {
// Filter so we don't issue redunant copies over stride-0 modes
copy(filter(tCrCol), filter(tCsBuf));
}
sync_fn();
constexpr int SmemFragSize = cute::max(1, sizeof(uint32_t) / sizeof(ElementCompute));
using FragmentSmem = Array<ElementCompute, SmemFragSize>;
using VectorSmem = uint_bit_t<sizeof_bits_v<FragmentSmem>>;
using ReduceSmem = GmemReduceFn<FragmentSmem>;
ReduceSmem reduce_smem{};
Tensor sBuf_frg = recast<FragmentSmem>(filter_zeros(sBuf));
Tensor sBuf_vec = recast<VectorSmem>(filter_zeros(sBuf));
constexpr int FragsPerCol = decltype(size<0>(sBuf_frg))::value;
// Do the threadblock smem reduction
CUTLASS_PRAGMA_UNROLL
for (int reduction_cols = size<1>(warp_layout_MN) / 2; reduction_cols > 1; reduction_cols /= 2) {
int FragsPerReduction = reduction_cols * FragsPerCol;
CUTLASS_PRAGMA_NO_UNROLL
for (int frg_idx = thread_idx; frg_idx < FragsPerReduction; frg_idx += size(tiled_copy)) {
FragmentSmem frg_smem = reduce_smem(sBuf_frg(frg_idx), sBuf_frg(frg_idx + FragsPerReduction));
sBuf_vec(frg_idx) = reinterpret_cast<VectorSmem&>(frg_smem);
}
sync_fn();
}
// Do final smem reduction and dump to gmem workspace
using VectorGmem = conditional_t<FinalReduction, VectorSmem volatile, VectorSmem>;
Tensor gBuf_vec = recast<VectorGmem>(filter(gBuf_nl(_,_,n,l)));
CUTLASS_PRAGMA_NO_UNROLL
for (int frg_idx = thread_idx; frg_idx < FragsPerCol; frg_idx += size(tiled_copy)) {
FragmentSmem frg_smem = reduce_smem(sBuf_frg(frg_idx), sBuf_frg(frg_idx + FragsPerCol));
gBuf_vec(frg_idx) = reinterpret_cast<VectorSmem&>(frg_smem);
}
sync_fn();
}
//
// 3. Increment atomic counters to signal final gmem reduction
//
if constexpr (not IsAtomic && FinalReduction) {
// Ensure gmem writes are visible to other threads before incrementing counter
__threadfence();
sync_fn();
// Collective thread 0 increments atomic tile counter and copies value to smem
int* prev_tile_count = reinterpret_cast<int*>(raw_pointer_cast(smem_buffer.data()));
if (thread_idx == 0) {
*prev_tile_count = atomicAdd(&params.tile_counters[m], 1);
}
sync_fn();
// Broadcast tile count to other threads in CTA and determine final reduction status
do_final_reduction = *prev_tile_count == size<2>(gBuf_nl) * size<3>(gBuf_nl) - 1;
sync_fn();
}
}
CUTLASS_DEVICE void
end() {
//
// 4. Do final gmem reduction if necessary
//
if constexpr (not IsAtomic && FinalReduction) {
if (not do_final_reduction) {
return;
}
auto& [ref_src, tCrCol, tCcCol, gCol_l, cCol, gBuf_nl, sBuf_layout,
lane_layout_MN, lane_mn, warp_layout_MN, warp_mn,
tile_coord_mnkl, residue_mn, epi_tile, tiled_copy, thread_idx] = args_tuple;
using ReduceOutput = GmemReduceFn<ElementCompute>;
using ConvertOutput = NumericConverter<ElementOutput, ElementCompute, RoundStyle>;
ReduceOutput reduce_output{};
ConvertOutput convert_output{};
// Reduction over batches
if (size<2>(stride(gCol_l)) == 0) {
CUTLASS_PRAGMA_NO_UNROLL
for (int m = thread_idx; m < size<0>(gBuf_nl); m += size(tiled_copy)) {
Tensor tRgBuf_nl = gBuf_nl(m,_0{},_,_);
ElementCompute output = tRgBuf_nl(_0{});
CUTLASS_PRAGMA_NO_UNROLL
for (int nl = 1; nl < size(tRgBuf_nl); ++nl) {
output = reduce_output(output, tRgBuf_nl(nl));
}
if (elem_less(cCol(m,_0{}), residue_mn)) {
gCol_l(m,_0{},_0{}) = convert_output(output);
}
}
}
// No reduction over batches
else {
CUTLASS_PRAGMA_NO_UNROLL
for (int m = thread_idx; m < size<0>(gBuf_nl); m += size(tiled_copy)) {
bool do_store = elem_less(cCol(m,_0{}), residue_mn);
CUTLASS_PRAGMA_NO_UNROLL
for (int l = 0; l < size<3>(gBuf_nl); ++l) {
Tensor tRgBuf_n = gBuf_nl(m,_0{},_,l);
ElementCompute output = tRgBuf_n(_0{});
CUTLASS_PRAGMA_NO_UNROLL
for (int n = 1; n < size(tRgBuf_n); ++n) {
output = reduce_output(output, tRgBuf_n(n));
}
if (do_store) {
gCol_l(m,_0{},l) = convert_output(output);
}
}
}
}
}
}
CUTLASS_DEVICE bool
is_reduction_buffer_needed(int epi_m, int epi_n, bool is_last_iteration) const {
auto const& [ref_src, tCrCol, tCcCol, gCol_l, cCol, gBuf_nl, sBuf_layout,
lane_layout_MN, lane_mn, warp_layout_MN, warp_mn,
tile_coord_mnkl, residue_mn, epi_tile, tiled_copy, thread_idx] = args_tuple;
return (not IsAtomic && // atomic reduction doesn't use smem
is_last_iteration && // smem reduction happens after epilogue loop
(decltype(size<1>(warp_layout_MN))::value > 1 || // smem reduction happens when multiple warps are in N
FinalReduction)); // smem is used to broadcast tile counters for final reduction
}
};
template <
bool ReferenceSrc, // do register tensors reference the src or dst layout of the tiled copy
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class EpilogueTile,
class TiledCopy,
class SrcTensor
class... Args
>
CUTLASS_DEVICE auto
get_consumer_store_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
TiledCopy tiled_copy,
int thread_idx,
SrcTensor const& tCrC) {
get_consumer_store_callbacks(ConsumerStoreArgs<Args...> const& args) {
Layout ref_layout_MN = [&] () {
if constexpr (ReferenceSrc) { return get<0>(args.tiled_copy.get_layoutS_MN()); }
else { return get<0>(args.tiled_copy.get_layoutD_MN()); }
}(); // tile_mn -> tv_idx
auto [M, N, K, L] = problem_shape_mnkl;
Tensor mCol = make_tensor(make_gmem_ptr(params.ptr_col), make_shape(M,N,L), params.dCol);
// Get the MN layout + coord of lanes to determine shuffle reduction iterations
using _W = Int<decltype(args.tiled_copy)::TiledNumThr::value / NumThreadsPerWarp>;
Layout tv2lane = Layout<Shape<Int<NumThreadsPerWarp>,_W,_1>,Stride<_1,_0,_0>>{}; // tv_idx -> lane_idx
Layout ref2lane = composition(tv2lane, ref_layout_MN); // tile_mn -> lane_idx
Layout lane_layout_MN = make_layout(filter(get<0>(ref2lane)), filter(get<1>(ref2lane))); // lane_mn -> lane_idx
Layout inv_lane_layout_MN = right_inverse(lane_layout_MN); // lane_idx -> lane_mn
int lane_idx = canonical_lane_idx();
auto lane_mn = idx2crd(inv_lane_layout_MN(lane_idx), shape(lane_layout_MN));
// Get the MN layout + coord of warps to determine smem reduction iterations
Layout tv2warp = Layout<Shape<Int<NumThreadsPerWarp>,_W,_1>,Stride<_0,_1,_0>>{}; // tv_idx -> warp_idx
Layout ref2warp = composition(tv2warp, ref_layout_MN); // tile_mn -> warp_idx
Layout warp_layout_MN = make_layout(filter(get<0>(ref2warp)), filter(get<1>(ref2warp))); // warp_mn -> warp_idx
Layout inv_warp_layout_MN = right_inverse(warp_layout_MN); // warp_idx -> warp_mn
int warp_idx = args.thread_idx / NumThreadsPerWarp;
auto warp_mn = idx2crd(inv_warp_layout_MN(warp_idx), shape(warp_layout_MN));
// Partition output gmem and register tensors
auto [tile_M, tile_N, tile_K] = args.tile_shape_mnk;
auto [M, N, K, L] = args.problem_shape_mnkl;
auto [m, n, k, l] = args.tile_coord_mnkl;
Tensor mCol = make_tensor(make_gmem_ptr<ElementOutput>(params.ptr_col), make_shape(M,N,L), params.dCol); // (M,N,L)
Tensor gCol_l = local_tile(mCol, take<0,2>(args.tile_shape_mnk), make_coord(m,n,_)); // (CTA_M,CTA_N,L)
Tensor tCgCol = sm90_partition_for_epilogue<ReferenceSrc>( // (CPY,CPY_M,CPY_N,EPI_M,EPI_N)
mCol, tile_shape_mnk, tile_coord_mnkl, epi_tile, tiled_copy, thread_idx);
gCol_l(_,_,l), args.epi_tile, args.tiled_copy, args.thread_idx);
Tensor tCrCol = make_tensor_like<ElementCompute>(tCgCol); // (CPY,CPY_M,CPY_N,EPI_M,EPI_N)
fill(tCrCol, params.reduction_identity);
// Compute tile residues and coordinate tensors for predication
auto [m, n, k, l] = tile_coord_mnkl;
auto residue_mn = make_coord(
M - static_cast<int>(m) * size<0>(tile_shape_mnk),
N - static_cast<int>(n) * size<1>(tile_shape_mnk)
);
Tensor cCol = make_identity_tensor(take<0,2>(tile_shape_mnk));
Tensor tCcCol = sm90_partition_for_epilogue<ReferenceSrc>(cCol, epi_tile, tiled_copy, thread_idx);
// Partition gmem+smem reduction buffer tensors
Layout gBuf_layout = make_layout(take<0,2>(args.tile_shape_mnk), make_stride(_1{}, _0{}));
Layout mBuf_layout = blocked_product(gBuf_layout, make_layout(ceil_div(make_shape(M,N,L), shape(gBuf_layout))));
Tensor mBuf = make_tensor(make_gmem_ptr(params.reduction_buffer), mBuf_layout); // (ceil_M,ceil_N,L)
Tensor gBuf_nl = local_tile(mBuf, take<0,2>(args.tile_shape_mnk), make_coord(m,_,_)); // (CTA_M,CTA_N,REST_N,L)
Layout sBuf_layout = blocked_product(gBuf_layout,make_layout(make_shape(_1{},_1{},size<1>(warp_layout_MN)))); // (CTA_M,CTA_N,WARPS_N)
return ConsumerStoreCallbacks(cute::move(tCrCol), cute::move(tCgCol), cute::move(tCcCol), residue_mn, params);
return ConsumerStoreCallbacks(
make_tuple(bool_constant<ReferenceSrc>{}, cute::move(tCrCol), args.tCcD, gCol_l, args.cD, gBuf_nl, sBuf_layout,
lane_layout_MN, lane_mn, warp_layout_MN, warp_mn,
args.tile_coord_mnkl, args.residue_mn, args.epi_tile, args.tiled_copy, args.thread_idx),
params
);
}
};
@@ -37,6 +37,7 @@
#pragma once
#include "cutlass/cutlass.h"
#include "cutlass/workspace.h"
#include "cute/tensor.hpp"
@@ -71,7 +72,7 @@ sm90_partition_for_epilogue(
TiledCopy tiled_copy,
int thread_idx) {
ThrCopy thread_copy = tiled_copy.get_thread_slice(thread_idx);
Tensor cT_epi = local_tile(cT, epi_tile, _); // (EPI_TILE_M,EPI_TILE_N,EPI_M,EPI_N,...)
Tensor cT_epi = flat_divide(cT, epi_tile); // (EPI_TILE_M,EPI_TILE_N,EPI_M,EPI_N,...)
if constexpr (ReferenceSrc) {
return thread_copy.partition_S(cT_epi); // (CPY,CPY_M,CPY_N,EPI_M,EPI_N,...)
}
@@ -111,6 +112,84 @@ sm90_partition_for_epilogue(
//
/////////////////////////////////////////////////////////////////////////////////////////////////
template<
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class ResidueMN,
class EpilogueTile
>
struct ProducerLoadArgs {
ProblemShapeMNKL problem_shape_mnkl;
TileShapeMNK tile_shape_mnk;
TileCoordMNKL tile_coord_mnkl;
ResidueMN residue_mn;
EpilogueTile epi_tile;
int thread_idx;
CUTLASS_DEVICE
ProducerLoadArgs(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
ResidueMN residue_mn,
EpilogueTile epi_tile,
int thread_idx)
: problem_shape_mnkl(problem_shape_mnkl),
tile_shape_mnk(tile_shape_mnk),
tile_coord_mnkl(tile_coord_mnkl),
residue_mn(residue_mn),
epi_tile(epi_tile),
thread_idx(thread_idx) {}
};
template<
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class ResidueMN,
class EpilogueTile,
class TiledCopy,
class CoordTensor,
class ThrCoordTensor,
class ThrSrcTensor
>
struct ConsumerStoreArgs {
ProblemShapeMNKL problem_shape_mnkl;
TileShapeMNK tile_shape_mnk;
TileCoordMNKL tile_coord_mnkl;
ResidueMN residue_mn;
EpilogueTile epi_tile;
TiledCopy tiled_copy;
int thread_idx;
CoordTensor cD;
ThrCoordTensor tCcD;
ThrSrcTensor const& tCrC;
CUTLASS_DEVICE
ConsumerStoreArgs(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
ResidueMN residue_mn,
EpilogueTile epi_tile,
TiledCopy tiled_copy,
int thread_idx,
CoordTensor cD,
ThrCoordTensor tCcD,
ThrSrcTensor const& tCrC)
: problem_shape_mnkl(problem_shape_mnkl),
tile_shape_mnk(tile_shape_mnk),
tile_coord_mnkl(tile_coord_mnkl),
residue_mn(residue_mn),
epi_tile(epi_tile),
tiled_copy(tiled_copy),
thread_idx(thread_idx),
cD(cD),
tCcD(tCcD),
tCrC(tCrC) {}
};
template <class... Ops>
struct Sm90VisitorImplBase {
// Shared memory allocation
@@ -132,6 +211,46 @@ struct Sm90VisitorImplBase {
);
}
template <class ProblemShape>
static size_t
get_workspace_size(ProblemShape const& problem_shape, Arguments const& args) {
return transform_apply(tuple<Ops...>{}, args,
[&] (auto&& op, auto const& op_args) {
using Op = cute::remove_cvref_t<decltype(op)>;
size_t op_workspace_size = Op::get_workspace_size(problem_shape, op_args);
return round_nearest(op_workspace_size, MinWorkspaceAlignment);
},
[&] (auto&&... op_workspace_size) {
return (0 + ... + op_workspace_size);
}
);
}
template <class ProblemShape>
static cutlass::Status
initialize_workspace(ProblemShape const& problem_shape, Arguments const& args, void* workspace, cudaStream_t stream) {
Status status = Status::kSuccess;
uint8_t* op_workspace = reinterpret_cast<uint8_t*>(workspace);
return transform_apply(tuple<Ops...>{}, args,
// Initialize each operation's workspace, stopping at the first error
[&] (auto&& op, auto const& op_args) {
if (status != Status::kSuccess) {
return status;
}
using Op = cute::remove_cvref_t<decltype(op)>;
status = Op::initialize_workspace(problem_shape, op_args, op_workspace, stream);
if (op_workspace != nullptr) {
size_t op_workspace_size = Op::get_workspace_size(problem_shape, op_args);
op_workspace += round_nearest(op_workspace_size, MinWorkspaceAlignment);
}
return status;
},
// Return the final status
[&] (auto const&...) { return status; }
);
}
CUTLASS_HOST_DEVICE
Sm90VisitorImplBase() {}
@@ -167,13 +286,11 @@ struct Sm90VisitorImpl : Sm90VisitorImplBase<Ops...> {
// e.g. for batched beta this must always be true regardless of current batch idx
CUTLASS_DEVICE bool
is_producer_load_needed() const {
bool needed = false;
for_each(ops,
[&] (auto const& op) {
needed |= op.is_producer_load_needed();
return apply(ops,
[] (auto const&... op) {
return (false || ... || op.is_producer_load_needed());
}
);
return needed;
}
// Is a producer TMA load specifically for C needed
@@ -183,13 +300,11 @@ struct Sm90VisitorImpl : Sm90VisitorImplBase<Ops...> {
// e.g. for batched beta this can be false depending on current batch idx
CUTLASS_DEVICE bool
is_C_load_needed() const {
bool needed = false;
for_each(ops,
[&] (auto const& op) {
needed |= op.is_C_load_needed();
return apply(ops,
[] (auto const&... op) {
return (false || ... || op.is_C_load_needed());
}
);
return needed;
}
//
@@ -241,28 +356,12 @@ struct Sm90VisitorImpl : Sm90VisitorImplBase<Ops...> {
// Producer load callbacks factory
// All operations must redefine this, but most can just dispatch to the base impl
template <
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class EpilogueTile
>
template <class... Args>
CUTLASS_DEVICE auto
get_producer_load_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
int thread_idx) {
get_producer_load_callbacks(ProducerLoadArgs<Args...> const& args) {
return transform_apply(ops,
[&] (auto& op) {
return op.get_producer_load_callbacks(
problem_shape_mnkl,
tile_shape_mnk,
tile_coord_mnkl,
epi_tile,
thread_idx
);
return op.get_producer_load_callbacks(args);
},
[] (auto&&... callbacks) {
auto callbacks_tuple = cute::make_tuple(callbacks...);
@@ -293,10 +392,10 @@ struct Sm90VisitorImpl : Sm90VisitorImplBase<Ops...> {
// Start of subtile store iteration. Smem broadcasts usually performed here.
// Upon entry, all producer loads for this subtile are completed and visible.
CUTLASS_DEVICE void
step_begin(int epi_m, int epi_n, int load_iteration, bool is_producer_load_needed) {
previsit(int epi_m, int epi_n, int load_iteration, bool is_producer_load_needed) {
for_each(callbacks_tuple,
[&] (auto& callbacks) {
callbacks.step_begin(epi_m, epi_n, load_iteration, is_producer_load_needed);
callbacks.previsit(epi_m, epi_n, load_iteration, is_producer_load_needed);
}
);
}
@@ -308,24 +407,49 @@ struct Sm90VisitorImpl : Sm90VisitorImplBase<Ops...> {
Array<ElementInputs, FragmentSize> const&... frg_inputs) // depends on the N-naryness of the op
= delete; // Must be implemented for each operation
// After D smem store, before smem async fence. Smem reductions usually performed here.
// After visit call, before smem async fence. Smem stores usually performed here.
// Upon exit, all smem stores for TMA must have been issued
CUTLASS_DEVICE void
step_next(int epi_m, int epi_n, int store_iteration, bool issue_smem_store) {
postvisit(int epi_m, int epi_n, int store_iteration, bool issue_smem_store) {
for_each(callbacks_tuple,
[&] (auto& callbacks) {
callbacks.step_next(epi_m, epi_n, store_iteration, issue_smem_store);
callbacks.postvisit(epi_m, epi_n, store_iteration, issue_smem_store);
}
);
}
// End of subtile iteration, before TMA store commit. Aux stores usually performed here
// After async fence, before TMA store commit. Aux stores usually performed here
// Upon exit, all TMA stores for this subtile must have been issued
CUTLASS_DEVICE void
step_end(int epi_m, int epi_n, int store_iteration, bool issue_tma_store) {
step(int epi_m, int epi_n, int store_iteration, bool issue_tma_store) {
for_each(callbacks_tuple,
[&] (auto& callbacks) {
callbacks.step_end(epi_m, epi_n, store_iteration, issue_tma_store);
callbacks.step(epi_m, epi_n, store_iteration, issue_tma_store);
}
);
}
// After TMA store commit. Smem reductions usually performed here
// reduction_buffer is an arbitrary smem tensor that can be used for workspace
// It is each nodes reponsibility to assert that this buffer is sufficiently sized
// and to ensure that this buffer is no longer needed upon callback exit
// i.e. results are synchronized and no longer in the reduction buffer
template <class STensor, class SyncFn>
CUTLASS_DEVICE void
reduce(STensor&& reduction_buffer, SyncFn const& sync_fn, int epi_m, int epi_n, bool is_last_iteration) {
for_each(callbacks_tuple,
[&] (auto& callbacks) {
callbacks.reduce(reduction_buffer, sync_fn, epi_m, epi_n, is_last_iteration);
}
);
}
// Collective can query this to determine whether a buffer needs to be freed for reduction
CUTLASS_DEVICE bool
is_reduction_buffer_needed(int epi_m, int epi_n, bool is_last_iteration) const {
return apply(callbacks_tuple,
[&] (auto const&... callbacks) {
return (false || ... || callbacks.is_reduction_buffer_needed(epi_m, epi_n, is_last_iteration));
}
);
}
@@ -345,33 +469,13 @@ struct Sm90VisitorImpl : Sm90VisitorImplBase<Ops...> {
// All operations must redefine this
template <
bool ReferenceSrc, // do register tensors reference the src or dst layout of the tiled copy
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class EpilogueTile,
class TiledCopy,
class SrcTensor
class... Args
>
CUTLASS_DEVICE auto
get_consumer_store_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
TiledCopy tiled_copy,
int thread_idx,
SrcTensor const& tCrC) {
get_consumer_store_callbacks(ConsumerStoreArgs<Args...> const& args) {
return transform_apply(ops,
[&] (auto& op) {
return op.template get_consumer_store_callbacks<ReferenceSrc>(
problem_shape_mnkl,
tile_shape_mnk,
tile_coord_mnkl,
epi_tile,
tiled_copy,
thread_idx,
tCrC
);
return op.template get_consumer_store_callbacks<ReferenceSrc>(args);
},
[] (auto&&... callbacks) {
auto callbacks_tuple = cute::make_tuple(callbacks...);
@@ -430,33 +534,13 @@ struct Sm90TreeVisitor : Sm90VisitorImpl<ChildOps..., NodeOp> {
template <
bool ReferenceSrc, // do register tensors reference the src or dst layout of the tiled copy
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class EpilogueTile,
class TiledCopy,
class SrcTensor
class... Args
>
CUTLASS_DEVICE auto
get_consumer_store_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
TiledCopy tiled_copy,
int thread_idx,
SrcTensor const& tCrC) {
get_consumer_store_callbacks(ConsumerStoreArgs<Args...> const& args) {
return ConsumerStoreCallbacks(
Sm90VisitorImpl<ChildOps..., NodeOp>::
get_consumer_store_callbacks<ReferenceSrc>(
problem_shape_mnkl,
tile_shape_mnk,
tile_coord_mnkl,
epi_tile,
tiled_copy,
thread_idx,
tCrC
)
get_consumer_store_callbacks<ReferenceSrc>(args)
);
}
@@ -502,33 +586,13 @@ struct Sm90SplitTreeVisitor : Sm90VisitorImpl<InputTree, AuxOutTrees..., OutputT
template <
bool ReferenceSrc, // do register tensors reference the src or dst layout of the tiled copy
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class EpilogueTile,
class TiledCopy,
class SrcTensor
class... Args
>
CUTLASS_DEVICE auto
get_consumer_store_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
TiledCopy tiled_copy,
int thread_idx,
SrcTensor const& tCrC) {
get_consumer_store_callbacks(ConsumerStoreArgs<Args...> const& args) {
return ConsumerStoreCallbacks(
Sm90VisitorImpl<InputTree, AuxOutTrees..., OutputTree>::
get_consumer_store_callbacks<ReferenceSrc>(
problem_shape_mnkl,
tile_shape_mnk,
tile_coord_mnkl,
epi_tile,
tiled_copy,
thread_idx,
tCrC
)
get_consumer_store_callbacks<ReferenceSrc>(args)
);
}
@@ -601,33 +665,13 @@ struct Sm90TopologicalVisitor : Sm90VisitorImpl<Ops...> {
template <
bool ReferenceSrc, // do register tensors reference the src or dst layout of the tiled copy
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class EpilogueTile,
class TiledCopy,
class SrcTensor
class... Args
>
CUTLASS_DEVICE auto
get_consumer_store_callbacks(
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_mnk,
TileCoordMNKL tile_coord_mnkl,
EpilogueTile epi_tile,
TiledCopy tiled_copy,
int thread_idx,
SrcTensor const& tCrC) {
get_consumer_store_callbacks(ConsumerStoreArgs<Args...> const& args) {
return ConsumerStoreCallbacks(
Sm90VisitorImpl<Ops...>::
get_consumer_store_callbacks<ReferenceSrc>(
problem_shape_mnkl,
tile_shape_mnk,
tile_coord_mnkl,
epi_tile,
tiled_copy,
thread_idx,
tCrC
)
get_consumer_store_callbacks<ReferenceSrc>(args)
);
}
@@ -663,6 +707,33 @@ struct Sm90VisitorImplBase<Op0> {
};
}
template <class ProblemShape>
static size_t
get_workspace_size(ProblemShape const& problem_shape, Arguments const& args) {
size_t workspace_size = 0;
workspace_size += Op0::get_workspace_size(problem_shape, args.op_0);
workspace_size = round_nearest(workspace_size, MinWorkspaceAlignment);
return workspace_size;
}
template <class ProblemShape>
static cutlass::Status
initialize_workspace(ProblemShape const& problem_shape, Arguments const& args, void* workspace, cudaStream_t stream) {
Status status = Status::kSuccess;
uint8_t* workspace_ptr = reinterpret_cast<uint8_t*>(workspace);
size_t workspace_offset = 0;
status = Op0::initialize_workspace(problem_shape, args.op_0, workspace_ptr + workspace_offset, stream);
workspace_offset += Op0::get_workspace_size(problem_shape, args.op_0);
workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment);
if (status != Status::kSuccess) {
return status;
}
return status;
}
CUTLASS_HOST_DEVICE
Sm90VisitorImplBase() {}
@@ -702,6 +773,43 @@ struct Sm90VisitorImplBase<Op0, Op1> {
};
}
template <class ProblemShape>
static size_t
get_workspace_size(ProblemShape const& problem_shape, Arguments const& args) {
size_t workspace_size = 0;
workspace_size += Op0::get_workspace_size(problem_shape, args.op_0);
workspace_size = round_nearest(workspace_size, MinWorkspaceAlignment);
workspace_size += Op1::get_workspace_size(problem_shape, args.op_1);
workspace_size = round_nearest(workspace_size, MinWorkspaceAlignment);
return workspace_size;
}
template <class ProblemShape>
static cutlass::Status
initialize_workspace(ProblemShape const& problem_shape, Arguments const& args, void* workspace, cudaStream_t stream) {
Status status = Status::kSuccess;
uint8_t* workspace_ptr = reinterpret_cast<uint8_t*>(workspace);
size_t workspace_offset = 0;
status = Op0::initialize_workspace(problem_shape, args.op_0, workspace_ptr + workspace_offset, stream);
workspace_offset += Op0::get_workspace_size(problem_shape, args.op_0);
workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment);
if (status != Status::kSuccess) {
return status;
}
status = Op1::initialize_workspace(problem_shape, args.op_1, workspace_ptr + workspace_offset, stream);
workspace_offset += Op1::get_workspace_size(problem_shape, args.op_1);
workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment);
if (status != Status::kSuccess) {
return status;
}
return status;
}
CUTLASS_HOST_DEVICE
Sm90VisitorImplBase() {}
@@ -746,6 +854,53 @@ struct Sm90VisitorImplBase<Op0, Op1, Op2> {
};
}
template <class ProblemShape>
static size_t
get_workspace_size(ProblemShape const& problem_shape, Arguments const& args) {
size_t workspace_size = 0;
workspace_size += Op0::get_workspace_size(problem_shape, args.op_0);
workspace_size = round_nearest(workspace_size, MinWorkspaceAlignment);
workspace_size += Op1::get_workspace_size(problem_shape, args.op_1);
workspace_size = round_nearest(workspace_size, MinWorkspaceAlignment);
workspace_size += Op2::get_workspace_size(problem_shape, args.op_2);
workspace_size = round_nearest(workspace_size, MinWorkspaceAlignment);
return workspace_size;
}
template <class ProblemShape>
static cutlass::Status
initialize_workspace(ProblemShape const& problem_shape, Arguments const& args, void* workspace, cudaStream_t stream) {
Status status = Status::kSuccess;
uint8_t* workspace_ptr = reinterpret_cast<uint8_t*>(workspace);
size_t workspace_offset = 0;
status = Op0::initialize_workspace(problem_shape, args.op_0, workspace_ptr + workspace_offset, stream);
workspace_offset += Op0::get_workspace_size(problem_shape, args.op_0);
workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment);
if (status != Status::kSuccess) {
return status;
}
status = Op1::initialize_workspace(problem_shape, args.op_1, workspace_ptr + workspace_offset, stream);
workspace_offset += Op1::get_workspace_size(problem_shape, args.op_1);
workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment);
if (status != Status::kSuccess) {
return status;
}
status = Op2::initialize_workspace(problem_shape, args.op_2, workspace_ptr + workspace_offset, stream);
workspace_offset += Op2::get_workspace_size(problem_shape, args.op_2);
workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment);
if (status != Status::kSuccess) {
return status;
}
return status;
}
CUTLASS_HOST_DEVICE
Sm90VisitorImplBase() {}
@@ -795,6 +950,63 @@ struct Sm90VisitorImplBase<Op0, Op1, Op2, Op3> {
};
}
template <class ProblemShape>
static size_t
get_workspace_size(ProblemShape const& problem_shape, Arguments const& args) {
size_t workspace_size = 0;
workspace_size += Op0::get_workspace_size(problem_shape, args.op_0);
workspace_size = round_nearest(workspace_size, MinWorkspaceAlignment);
workspace_size += Op1::get_workspace_size(problem_shape, args.op_1);
workspace_size = round_nearest(workspace_size, MinWorkspaceAlignment);
workspace_size += Op2::get_workspace_size(problem_shape, args.op_2);
workspace_size = round_nearest(workspace_size, MinWorkspaceAlignment);
workspace_size += Op3::get_workspace_size(problem_shape, args.op_3);
workspace_size = round_nearest(workspace_size, MinWorkspaceAlignment);
return workspace_size;
}
template <class ProblemShape>
static cutlass::Status
initialize_workspace(ProblemShape const& problem_shape, Arguments const& args, void* workspace, cudaStream_t stream) {
Status status = Status::kSuccess;
uint8_t* workspace_ptr = reinterpret_cast<uint8_t*>(workspace);
size_t workspace_offset = 0;
status = Op0::initialize_workspace(problem_shape, args.op_0, workspace_ptr + workspace_offset, stream);
workspace_offset += Op0::get_workspace_size(problem_shape, args.op_0);
workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment);
if (status != Status::kSuccess) {
return status;
}
status = Op1::initialize_workspace(problem_shape, args.op_1, workspace_ptr + workspace_offset, stream);
workspace_offset += Op1::get_workspace_size(problem_shape, args.op_1);
workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment);
if (status != Status::kSuccess) {
return status;
}
status = Op2::initialize_workspace(problem_shape, args.op_2, workspace_ptr + workspace_offset, stream);
workspace_offset += Op2::get_workspace_size(problem_shape, args.op_2);
workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment);
if (status != Status::kSuccess) {
return status;
}
status = Op3::initialize_workspace(problem_shape, args.op_3, workspace_ptr + workspace_offset, stream);
workspace_offset += Op3::get_workspace_size(problem_shape, args.op_3);
workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment);
if (status != Status::kSuccess) {
return status;
}
return status;
}
CUTLASS_HOST_DEVICE
Sm90VisitorImplBase() {}
+19 -4
View File
@@ -171,8 +171,8 @@ struct ReLu<Array<T, N>> {
template <typename T>
struct Clamp {
struct Arguments {
T lower_bound = cutlass::platform::numeric_limits<T>::min();
T upper_bound = cutlass::platform::numeric_limits<T>::max();
T lower_bound = CUTLASS_STL_NAMESPACE::numeric_limits<T>::min();
T upper_bound = CUTLASS_STL_NAMESPACE::numeric_limits<T>::max();
};
CUTLASS_HOST_DEVICE
@@ -615,12 +615,13 @@ struct dGELU<Array<T, N> > {
template <typename T>
struct dReLU {
CUTLASS_HOST_DEVICE
T operator()(T const& d_t, bool d_relu) const {
T operator()(T d_t, bool d_relu) const {
return d_relu ? d_t : T(0);
}
template <typename U>
CUTLASS_HOST_DEVICE
T operator()(T const& d_t, uint1b_t d_relu) const {
T operator()(T d_t, U d_relu) const {
return operator()(d_t, static_cast<bool>(d_relu));
}
};
@@ -649,6 +650,20 @@ struct dReLU<Array<T, N>> {
return operator()(d_t, preds);
}
template <typename U>
CUTLASS_HOST_DEVICE
Array<T, N> operator()(Array<T, N> const& d_t, Array<U, N> const& d_relu) const {
Array<T, N> y;
dReLU<T> relu_op;
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < N; ++i) {
y[i] = relu_op(d_t[i], d_relu[i]);
}
return y;
}
};
/////////////////////////////////////////////////////////////////////////////////////////////////
@@ -44,6 +44,14 @@
namespace cutlass {
namespace epilogue {
namespace threadblock {
namespace detail {
struct EVT2xBase { };
template <class T>
static constexpr bool is_2x_evt_v = platform::is_base_of<EVT2xBase, T>::value;
} // namespace detail
////////////////////////////////////////////////////////////////////////////////////////////////////
@@ -67,7 +75,8 @@ class EpilogueWithVisitorCallbacks :
typename DefaultEpilogue::Shape,
DefaultEpilogue::kPartitionsK,
typename DefaultEpilogue::WarpMmaOperator,
typename DefaultEpilogue::AccumulatorFragmentIterator>
typename DefaultEpilogue::AccumulatorFragmentIterator>,
public detail::EVT2xBase
{
public:
@@ -96,7 +96,7 @@ struct VisitorScalarBroadcast {
(cute::is_same_v<StrideMNL, Stride<_0,_0,_0>>) || // scalar broadcast, e.g. alpha
(cute::is_same_v<StrideMNL, Stride<_0,_0,_1>>) ||
(cute::is_same_v<StrideMNL, Stride<_0,_0,int>>)); // batched scalar broadcast, e.g. per-batch alpha
struct SharedStorage { };
struct Arguments {
@@ -132,12 +132,12 @@ struct VisitorScalarBroadcast {
CUTLASS_DEVICE
Callbacks(Element scalar)
: scalar(scalar) {}
Element scalar;
template <class ElementAccumulator, int FragmentSize>
CUTLASS_DEVICE auto // returns an Array
visit(int iter_idx, int row_idx, int column_idx, int frg_idx,
visit(int iter_idx, int row_idx, int column_idx, int frg_idx,
Array<ElementAccumulator, FragmentSize> const& frg_acc) {
Array<Element, FragmentSize> frg_scalar;
frg_scalar.fill(scalar);
@@ -224,7 +224,7 @@ struct VisitorAuxLoad{
// Global load type
static int constexpr vec_bits = ThreadMap::kElementsPerAccess * sizeof_bits<Element>::value;
using VecType = uint_bit_t<cute::min(128, vec_bits)>;
static int constexpr VecLength = sizeof(VecType) / sizeof(Element);
static int constexpr VecLength = sizeof(VecType) / sizeof(Element);
CUTLASS_HOST_DEVICE
VisitorAuxLoad() { }
@@ -272,7 +272,7 @@ struct VisitorAuxLoad{
template <class ElementAccumulator, int FragmentSize>
CUTLASS_DEVICE auto // returns an Array
visit(int iter_idx, int row_idx, int column_idx, int frg_idx,
visit(int iter_idx, int row_idx, int column_idx, int frg_idx,
Array<ElementAccumulator, FragmentSize> const& frg_acc) {
Tensor tC_rAux_frg = recast<Array<Element, FragmentSize>>(coalesce(tC_rAux(_,_,_,iter_idx%Stages)));
return tC_rAux_frg(frg_idx);
@@ -285,9 +285,9 @@ struct VisitorAuxLoad{
gemm::GemmCoord threadblock_tile_offset,
int thread_idx,
ProblemShape problem_shape
) {
) {
Tensor mAux = make_tensor(
make_gmem_ptr(params_ptr->ptr_aux),
make_gmem_ptr(params_ptr->ptr_aux),
problem_shape,
params_ptr->dAux); // (M,N,L)
// VECTOR, FRAGMENT_COLUMN, FRAGMENT_ROW, ITERATION_ROW, ITERATION_GROUP, ITERATION_CLUSTER
@@ -299,14 +299,14 @@ struct VisitorAuxLoad{
// Generate the pred tensor
Tensor cAux = make_identity_tensor(mAux.shape());
Tensor tC_cAux = local_partition(
Tensor tC_cAux = outer_partition(
group_modes<3,6>(ThreadMap::partition(cAux, thread_idx, threadblock_tile_offset)),
Shape<Int<VecLength>>{},
(_0{})
);
return Callbacks<
decltype(tC_gAux), decltype(tC_rAux),
decltype(tC_gAux), decltype(tC_rAux),
decltype(tC_cAux), ProblemShape>(
cute::move(tC_gAux),
cute::move(tC_rAux),
@@ -354,7 +354,7 @@ struct VisitorRowBroadcast {
CUTLASS_HOST_DEVICE
VisitorRowBroadcast(Params const& params, SharedStorage const& shared_storage)
: params_ptr(&params) { }
Params const* params_ptr;
template <class GTensor, class RTensor, class CTensor, class ProblemShape>
@@ -372,7 +372,7 @@ struct VisitorRowBroadcast {
tC_cRow(cute::forward<CTensor>(tC_cRow)),
n(get<1>(problem_shape)),
params_ptr(params_ptr) { }
GTensor tC_gRow;
RTensor tC_rRow;
CTensor tC_cRow;
@@ -394,7 +394,7 @@ struct VisitorRowBroadcast {
template <class ElementAccumulator, int FragmentSize>
CUTLASS_DEVICE auto // returns an Array
visit(int iter_idx, int row_idx, int column_idx, int frg_idx,
visit(int iter_idx, int row_idx, int column_idx, int frg_idx,
Array<ElementAccumulator, FragmentSize> const& frg_acc) {
Tensor rRow_frg = recast<Array<Element, FragmentSize>>(coalesce(tC_rRow));
return rRow_frg(column_idx);
@@ -409,10 +409,10 @@ struct VisitorRowBroadcast {
ProblemShape problem_shape
) {
Tensor mRow = make_tensor(
make_gmem_ptr(params_ptr->ptr_row),
make_gmem_ptr(params_ptr->ptr_row),
problem_shape,
params_ptr->dRow);
// VECTOR, FRAGMENT_COLUMN
Tensor tC_gRow = recast<VecType>(
ThreadMap::partition(mRow, thread_idx, threadblock_tile_offset)
@@ -421,14 +421,14 @@ struct VisitorRowBroadcast {
// Generate the pred tensor
Tensor cRow = make_identity_tensor(mRow.shape());
Tensor tC_cRow = local_partition(
Tensor tC_cRow = outer_partition(
ThreadMap::partition(cRow, thread_idx, threadblock_tile_offset)(_,_,_0{},_0{},_0{},_0{}),
Shape<Int<VecLength>>{},
(_0{})
);
return Callbacks<
decltype(tC_gRow), decltype(tC_rRow),
decltype(tC_gRow), decltype(tC_rRow),
decltype(tC_cRow), ProblemShape>(
cute::move(tC_gRow),
cute::move(tC_rRow),
@@ -472,7 +472,7 @@ struct VisitorColBroadcast {
CUTLASS_HOST_DEVICE
VisitorColBroadcast(Params const& params, SharedStorage const& shared_storage)
: params_ptr(&params) { }
Params const* params_ptr;
template <class GTensor, class RTensor, class CTensor, class ProblemShape>
@@ -490,7 +490,7 @@ struct VisitorColBroadcast {
tC_cCol(cute::forward<CTensor>(tC_cCol)),
m(get<0>(problem_shape)),
params_ptr(params_ptr) { }
GTensor tC_gCol;
RTensor tC_rCol;
CTensor tC_cCol;
@@ -510,7 +510,7 @@ struct VisitorColBroadcast {
template <class ElementAccumulator, int FragmentSize>
CUTLASS_DEVICE auto // returns an Array
visit(int iter_idx, int row_idx, int column_idx, int frg_idx,
visit(int iter_idx, int row_idx, int column_idx, int frg_idx,
Array<ElementAccumulator, FragmentSize> const& frg_acc) {
Array<Element, FragmentSize> frg_col;
frg_col.fill(tC_rCol(row_idx,iter_idx));
@@ -529,7 +529,7 @@ struct VisitorColBroadcast {
make_gmem_ptr(params_ptr->ptr_col),
problem_shape,
params_ptr->dCol);
// VECTOR, FRAGMENT_COLUMN, FRAGMENT_ROW, ITERATION_ROW, ITERATION_GROUP, ITERATION_CLUSTER
Tensor tC_gCol = group_modes<1,4>(
ThreadMap::partition(mCol, thread_idx, threadblock_tile_offset)(_0{},_0{},_,_,_,_));
@@ -118,7 +118,7 @@ struct VisitorAuxStore{
template <class ElementAccumulator, class ElementInput, int FragmentSize>
CUTLASS_DEVICE auto // returns an Array
visit(int iter_idx, int row_idx, int column_idx, int frg_idx,
visit(int iter_idx, int row_idx, int column_idx, int frg_idx,
Array<ElementAccumulator, FragmentSize> const& frg_acc,
Array<ElementInput, FragmentSize> const& frg_input) {
using ConvertInput = NumericArrayConverter<Element, ElementInput, FragmentSize, RoundStyle>;
@@ -152,8 +152,8 @@ struct VisitorAuxStore{
ProblemShape problem_shape
) {
Tensor mAux = make_tensor(
make_gmem_ptr(params_ptr->ptr_aux),
problem_shape,
make_gmem_ptr(params_ptr->ptr_aux),
problem_shape,
params_ptr->dAux); // (M,N,L)
// VECTOR, FRAGMENT_COLUMN, FRAGMENT_ROW, ITERATION_ROW, ITERATION_GROUP, ITERATION_CLUSTER
Tensor tC_gAux = recast<VecType>(group_modes<3,6>(ThreadMap::partition(mAux, thread_idx, threadblock_tile_offset)));
@@ -161,14 +161,14 @@ struct VisitorAuxStore{
// Generate the pred tensor
Tensor cAux = make_identity_tensor(mAux.shape());
Tensor tC_cAux = local_partition(
Tensor tC_cAux = outer_partition(
group_modes<3,6>(ThreadMap::partition(cAux, thread_idx, threadblock_tile_offset)),
Shape<Int<VecLength>>{},
(_0{})
);
return Callbacks<
decltype(tC_gAux), decltype(tC_rAux),
decltype(tC_gAux), decltype(tC_rAux),
decltype(tC_cAux), ProblemShape>(
cute::move(tC_gAux),
cute::move(tC_rAux),
@@ -186,7 +186,7 @@ struct VisitorAuxStore{
/////////////////////////////////////////////////////////////////////////////////////////////////
// Helper functions
template <
template <class> class ReduceFn,
template <class> class ReduceFn,
int kThreads, class T>
CUTLASS_DEVICE
void intra_warp_row_reduce(T& value) {
@@ -266,7 +266,7 @@ struct VisitorColReduction {
CUTLASS_HOST_DEVICE
VisitorColReduction(Params const& params, SharedStorage const& shared_storage)
: params_ptr(&params) { }
Params const* params_ptr;
template <class GTensor, class CTensor, class ProblemShape>
@@ -307,7 +307,7 @@ struct VisitorColReduction {
template <class ElementAccumulator, class ElementInput, int FragmentSize>
CUTLASS_DEVICE auto // returns an Array
visit(int iter_idx, int row_idx, int column_idx, int frg_idx,
visit(int iter_idx, int row_idx, int column_idx, int frg_idx,
Array<ElementAccumulator, FragmentSize> const& frg_acc,
Array<ElementInput, FragmentSize> const& frg_input) {
@@ -414,13 +414,13 @@ struct VisitorRowReduction {
VisitorRowReduction(Params const& params, SharedStorage const& shared_storage)
: params_ptr(&params),
smem_reduce(const_cast<ElementCompute*>(shared_storage.reduction.data())) { }
Params const* params_ptr;
ElementCompute* smem_reduce;
template <
class RTensorR2S, class STensorR2S, class CTensorR2S,
class STensorS2R, class RTensorS2R, class CTensorS2R,
class STensorS2R, class RTensorS2R, class CTensorS2R,
class GTensor, class CTensor, class ProblemShape>
struct Callbacks : EmptyCallbacks {
CUTLASS_DEVICE
@@ -465,7 +465,7 @@ struct VisitorRowReduction {
// R->G
GTensor tC_gRow;
CTensor tC_cRow;
Params const* params_ptr;
int n;
int m;
@@ -477,10 +477,10 @@ struct VisitorRowReduction {
template <class ElementAccumulator, class ElementInput, int FragmentSize>
CUTLASS_DEVICE auto // returns an Array
visit(int iter_idx, int row_idx, int column_idx, int frg_idx,
visit(int iter_idx, int row_idx, int column_idx, int frg_idx,
Array<ElementAccumulator, FragmentSize> const& frg_acc,
Array<ElementInput, FragmentSize> const& frg_input) {
using ConvertInput = NumericArrayConverter<ElementCompute, ElementInput, FragmentSize, RoundStyle>;
ConvertInput convert_input{};
Tensor tRS_rRow_frg = recast<Array<ElementCompute, FragmentSize>>(coalesce(tRS_rSrc));
@@ -528,13 +528,13 @@ struct VisitorRowReduction {
atomic_reduce<AtomicReduceFn, RoundStyle>(&tC_gRow(j), tSR_rRows(j));
}
}
}
}
private:
template <int FragmentSize>
CUTLASS_DEVICE ElementCompute
CUTLASS_DEVICE ElementCompute
reduction(Array<ElementCompute, FragmentSize>& reduce_buffer, Array<ElementCompute, FragmentSize> const& result) {
using ReduceInput = RegReduceFn<ElementCompute>;
ReduceInput reduce_input{};
@@ -556,7 +556,7 @@ struct VisitorRowReduction {
make_gmem_ptr(params_ptr->ptr_row),
problem_shape,
params_ptr->dRow);
//
// Step 1: reduce fragment input (Src) into tRS_rSrc
//
@@ -567,7 +567,7 @@ struct VisitorRowReduction {
Tensor cSrc = make_identity_tensor(mRow.shape());
// FRAGMENT_COLUMN, FRAGMENT_ROW, (ITERATION_ROW, ITERATION_GROUP, ITERATION_CLUSTER)
Tensor tRS_cSrc = group_modes<2,5>(ThreadMap::partition(cSrc, thread_idx, threadblock_tile_offset)(_0{},_,_,_,_,_));
//
// Step 2: copy the partial results in tRS_rSrc to sRows in shared memory
//
@@ -587,18 +587,18 @@ struct VisitorRowReduction {
// VECTOR*ACCESS_WIDTH*FRAGMENT_COL,ACCESS_ROWS*WARPS_PER_ROW*GROUPS*CLUSTERS
Tensor sRows_nm = coalesce(group_modes<1,5>(group_modes<0,3>(sRows)), Shape<_1,_1>{});
// SMEM_ROW/THREADS,ACCESS_ROWS*WARPS_PER_ROW*GROUPS*CLUSTERS
Tensor tSR_sRows = local_partition(sRows_nm, Shape<Int<ThreadMap::kThreads>,_1>{}, thread_idx);
Tensor tSR_sRows = outer_partition(sRows_nm, Shape<Int<ThreadMap::kThreads>,_1>{}, thread_idx);
// SMEM_ROW/THREADS
Tensor tSR_rRows = make_tensor_like(tSR_sRows(_,_0{}));
// Coord
Tensor cRows_nm = make_identity_tensor(sRows_nm.shape());
Tensor tSR_cRows = local_partition(cRows_nm, Shape<Int<ThreadMap::kThreads>,_1>{}, thread_idx)(_,_0{});
Tensor tSR_cRows = outer_partition(cRows_nm, Shape<Int<ThreadMap::kThreads>,_1>{}, thread_idx)(_,_0{});
//
// Step 4: atomically reduce the results to global memory
//
Tensor tC_gRow = local_partition(
Tensor tC_gRow = outer_partition(
// Cta tile
local_tile(
mRow, typename ThreadMap::CtaShapeMNL{}, make_coord(_,_,_),Step<_1,_1, X>{}
@@ -608,7 +608,7 @@ struct VisitorRowReduction {
)(_0{},_);
Tensor cRow = make_identity_tensor(mRow.shape());
Tensor tC_cRow = local_partition(
Tensor tC_cRow = outer_partition(
// Cta tile
local_tile(
cRow, typename ThreadMap::CtaShapeMNL{}, make_coord(_,_,_), Step<_1,_1, X>{}
@@ -680,7 +680,7 @@ struct VisitorScalarReduction {
CUTLASS_HOST_DEVICE
VisitorScalarReduction(Params const& params, SharedStorage const& shared_storage)
: params_ptr(&params) { }
Params const* params_ptr;
template <class CTensor, class GTensor, class ProblemShape>
@@ -760,7 +760,7 @@ struct VisitorScalarReduction {
);
Tensor tC_gScalar = mScalar(_,_,threadblock_tile_offset.k());
return Callbacks<
decltype(tC_cSrc), decltype(tC_gScalar),
ProblemShape>(