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:
@@ -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(¶ms.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() {}
|
||||
|
||||
|
||||
@@ -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(¶ms) { }
|
||||
|
||||
|
||||
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(¶ms) { }
|
||||
|
||||
|
||||
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(¶ms) { }
|
||||
|
||||
|
||||
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(¶ms),
|
||||
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(¶ms) { }
|
||||
|
||||
|
||||
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>(
|
||||
|
||||
Reference in New Issue
Block a user