CUTLASS 3.1 (#915)

Co-authored-by: Aniket Shivam <ashivam@nvidia.com>
This commit is contained in:
ANIKET SHIVAM
2023-04-14 23:19:34 -04:00
committed by GitHub
co-authored by Aniket Shivam
parent 9b8166e3f0
commit d572cc1aab
482 changed files with 37175 additions and 16410 deletions
@@ -0,0 +1,536 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include "cute/atom/mma_traits_sm90.hpp"
#include "cute/atom/mma_traits_sm90_gmma.hpp"
#include "cute/atom/copy_traits_sm90.hpp"
#include "cutlass/detail/dependent_false.hpp"
#include "cutlass/gemm/gemm.h"
#include "cutlass/gemm/collective/collective_builder.hpp"
#include "cutlass/epilogue/dispatch_policy.hpp"
#include "cutlass/epilogue/collective/collective_epilogue.hpp"
#include "cutlass/epilogue/thread/linear_combination.h"
#include "cutlass/epilogue/thread/linear_combination_generic.h"
#include "cutlass/epilogue/thread/linear_combination_bias_elementwise.h"
#if defined(__CUDACC_RTC__)
#include <cuda/std/type_traits>
#else
#include <type_traits>
#endif
///////////////////////////////////////////////////////////////////////////////
namespace cutlass::epilogue::collective {
///////////////////////////////////////////////////////////////////////////////
namespace detail {
// Returns the smem layout atom to be used for C or D matrix
template<class GmemStrideType, class Element, class EpilogueTile_MN>
constexpr auto
sm90_get_epilogue_smem_swizzle_layout_atom() {
using namespace cute;
// ColMajor C/D (M-major)
if constexpr (size<0>(GmemStrideType{}) == 1) {
return cutlass::gemm::collective::detail::ss_smem_selector<
cute::GMMA::Major::MN, Element, decltype(get<0>(EpilogueTile_MN{})), decltype(get<1>(EpilogueTile_MN{}))
>();
}
// RowMajor C/D (N-major)
else if constexpr (size<1>(GmemStrideType{}) == 1) {
return cutlass::gemm::collective::detail::ss_smem_selector<
cute::GMMA::Major::K , Element, decltype(get<0>(EpilogueTile_MN{})), decltype(get<1>(EpilogueTile_MN{}))
>();
}
else {
static_assert(cutlass::detail::dependent_false<GmemStrideType>, "Unsupported gmem layout.");
}
}
// Attempts to compute a reasonable epilogue tile based on block tile shape or allows the user to provide one.
template <class Element, class EpilogueTileType, class Schedule>
constexpr auto
sm90_compute_tile_shape_or_override() {
if constexpr (cute::is_same_v<EpilogueTileType, EpilogueTileAuto>) {
constexpr int SmemAlloc = 4096;
if constexpr (detail::sm90_is_cooperative_v<Schedule>) {
constexpr int M = 128;
constexpr int N = SmemAlloc / (M * sizeof(Element));
return make_shape(Int<M>{}, Int<N>{});
}
else if constexpr (detail::sm90_is_warp_specialized_v<Schedule>) {
constexpr int M = 64;
constexpr int N = SmemAlloc / (M * sizeof(Element));
return make_shape(Int<M>{}, Int<N>{});
}
else {
static_assert(cutlass::detail::dependent_false<Schedule>, "Unsupported schedule.");
}
}
else if constexpr (cute::is_tuple<EpilogueTileType>::value) {
EpilogueTileType epi_tile;
constexpr int M = size<0>(shape(epi_tile));
constexpr int N = size<1>(shape(epi_tile));
static_assert(!is_layout<EpilogueTileType>::value, "EpilogueTile must be a cute::Tile or cute::Shape");
static_assert(M == 64 && detail::sm90_is_warp_specialized_v<Schedule> ||
M == 128 && detail::sm90_is_cooperative_v<Schedule>, "Unsupported tile shape");
static_assert(N % 8 == 0, "Unsupported tile shape");
return epi_tile;
}
else {
static_assert(cutlass::detail::dependent_false<EpilogueTileType>, "Invalid type for EpilogueTileType.");
}
}
// Selects the largest vectorized smem store atom available
template <class GmemStrideTypeD, class ElementD>
constexpr auto
sm90_get_smem_store_op_for_accumulator() {
using namespace cute;
if constexpr (sizeof(ElementD) == 2 && size<0>(GmemStrideTypeD{}) == 1) {
return SM90_U16x8_STSM_T{};
}
else if constexpr (sizeof(ElementD) == 2 && size<1>(GmemStrideTypeD{}) == 1) {
return SM90_U32x4_STSM_N{};
}
else {
// auto-vectorizing store
return DefaultCopy{};
}
}
// Selects the largest vectorized smem load atom available
template <class GmemStrideTypeC, class ElementC>
constexpr auto
sm90_get_smem_load_op_for_source() {
using namespace cute;
// Reuse the logic from smem store selector
using SmemStoreOp = decltype(sm90_get_smem_store_op_for_accumulator<GmemStrideTypeC, ElementC>());
if constexpr (cute::is_same_v<SmemStoreOp, SM90_U16x8_STSM_T>) {
return SM75_U16x8_LDSM_T{};
}
else if constexpr (cute::is_same_v<SmemStoreOp, SM90_U32x4_STSM_N>) {
return SM75_U32x4_LDSM_N{};
}
else {
// auto-vectorizing load
return DefaultCopy{};
}
}
// Helper for building TMA warp-specialized collective epilogues, specialized by
// the thread-level epilogue operation performed and the dispatch policy to use.
template <
class TileShape_MNK,
class ClusterShape_MNK,
class EpilogueTileType,
class ElementAccumulator,
class ElementCompute,
class ElementC,
class GmemLayoutTagC,
int AlignmentC,
class ElementD,
class GmemLayoutTagD,
int AlignmentD,
class Schedule,
class ThreadOp,
class DispatchPolicy
>
struct TmaBuilderImpl {
using GmemStrideTypeC = gemm::TagToStrideC_t<GmemLayoutTagC>;
using GmemStrideTypeD = gemm::TagToStrideC_t<GmemLayoutTagD>;
using EpilogueTile_MN = decltype(detail::sm90_compute_tile_shape_or_override<
ElementD, EpilogueTileType, Schedule>());
using CollectiveOp = cutlass::epilogue::collective::CollectiveEpilogue<
DispatchPolicy,
TileShape_MNK,
EpilogueTile_MN,
ElementC,
GmemStrideTypeC,
ElementD,
GmemStrideTypeD,
ThreadOp,
SM90_TMA_LOAD,
decltype(detail::sm90_get_epilogue_smem_swizzle_layout_atom<GmemStrideTypeC, ElementC, TileShape_MNK>()),
decltype(detail::sm90_get_smem_load_op_for_source<GmemStrideTypeC, ElementC>()),
SM90_TMA_STORE,
decltype(detail::sm90_get_epilogue_smem_swizzle_layout_atom<GmemStrideTypeD, ElementD, EpilogueTile_MN>()),
decltype(detail::sm90_get_smem_store_op_for_accumulator<GmemStrideTypeD, ElementD>())
>;
};
} // namespace detail
///////////////////////////////////////////////////////////////////////////////
// No-smem builder
template <
class TileShape_MNK,
class ClusterShape_MNK,
class EpilogueTileType,
class ElementAccumulator,
class ElementCompute,
class ElementC_,
class GmemLayoutTagC_,
int AlignmentC,
class ElementD,
class GmemLayoutTagD,
int AlignmentD,
class Schedule
>
struct CollectiveBuilder<
arch::Sm90,
arch::OpClassTensorOp,
TileShape_MNK,
ClusterShape_MNK,
EpilogueTileType,
ElementAccumulator,
ElementCompute,
ElementC_,
GmemLayoutTagC_,
AlignmentC,
ElementD,
GmemLayoutTagD,
AlignmentD,
Schedule,
cute::enable_if_t<cute::is_same_v<Schedule, NoSmemWarpSpecialized>>> {
// Passing void C disables source load
using ElementC = cute::conditional_t<cute::is_void_v<ElementC_>,
ElementD, ElementC_>; // prevents cute breakages
using GmemLayoutTagC = cute::conditional_t<cute::is_void_v<ElementC_>,
GmemLayoutTagD, GmemLayoutTagC_>;
static constexpr thread::ScaleType::Kind ScaleType = cute::is_void_v<ElementC_> ?
thread::ScaleType::OnlyAlphaScaling : thread::ScaleType::Default;
using ThreadOp = thread::LinearCombination<
ElementD, 1, ElementAccumulator, ElementCompute,
ScaleType, FloatRoundStyle::round_to_nearest, ElementC>;
using CollectiveOp = cutlass::epilogue::collective::detail::Sm90TmaWarpSpecializedAdapter<
cutlass::epilogue::collective::DefaultEpilogue<
cutlass::gemm::TagToStrideC_t<GmemLayoutTagC>,
cutlass::gemm::TagToStrideC_t<GmemLayoutTagD>,
ThreadOp,
cutlass::gemm::EpilogueDefault>
>;
};
// Tma warp-specialized builder
template <
class TileShape_MNK,
class ClusterShape_MNK,
class EpilogueTileType,
class ElementAccumulator,
class ElementCompute,
class ElementC_,
class GmemLayoutTagC_,
int AlignmentC,
class ElementD,
class GmemLayoutTagD,
int AlignmentD,
class Schedule
>
struct CollectiveBuilder<
arch::Sm90,
arch::OpClassTensorOp,
TileShape_MNK,
ClusterShape_MNK,
EpilogueTileType,
ElementAccumulator,
ElementCompute,
ElementC_,
GmemLayoutTagC_,
AlignmentC,
ElementD,
GmemLayoutTagD,
AlignmentD,
Schedule,
cute::enable_if_t<cute::is_same_v<Schedule, TmaWarpSpecialized> ||
cute::is_same_v<Schedule, TmaWarpSpecializedCooperative> >> {
public:
// Passing void C disables source load
using ElementC = cute::conditional_t<cute::is_void_v<ElementC_>,
ElementD, ElementC_>; // prevents cute breakages
using GmemLayoutTagC = cute::conditional_t<cute::is_void_v<ElementC_>,
GmemLayoutTagD, GmemLayoutTagC_>;
static constexpr thread::ScaleType::Kind ScaleType = cute::is_void_v<ElementC_> ?
thread::ScaleType::OnlyAlphaScaling : thread::ScaleType::Default;
using ThreadOp = thread::LinearCombination<
ElementD, AlignmentD, ElementAccumulator, ElementCompute,
thread::ScaleType::Default, FloatRoundStyle::round_to_nearest, ElementC>;
private:
using Impl = detail::TmaBuilderImpl<
TileShape_MNK, ClusterShape_MNK, EpilogueTileType, ElementAccumulator, ElementCompute,
ElementC, GmemLayoutTagC, AlignmentC, ElementD, GmemLayoutTagD, AlignmentD,
Schedule, ThreadOp, cutlass::epilogue::Sm90TmaWarpSpecialized<1,2,true>>;
public:
using CollectiveOp = typename Impl::CollectiveOp;
};
// Auto builder
template <
class TileShape_MNK,
class ClusterShape_MNK,
class EpilogueTileType,
class ElementAccumulator,
class ElementCompute,
class ElementC,
class GmemLayoutTagC,
int AlignmentC,
class ElementD,
class GmemLayoutTagD,
int AlignmentD,
class Schedule
>
struct CollectiveBuilder<
arch::Sm90,
arch::OpClassTensorOp,
TileShape_MNK,
ClusterShape_MNK,
EpilogueTileType,
ElementAccumulator,
ElementCompute,
ElementC,
GmemLayoutTagC,
AlignmentC,
ElementD,
GmemLayoutTagD,
AlignmentD,
Schedule,
cute::enable_if_t<cute::is_same_v<Schedule, EpilogueScheduleAuto>>> {
private:
static constexpr bool IsTmaAligned = cutlass::gemm::collective::detail::is_aligned<
ElementC, AlignmentC, ElementD, AlignmentD, cutlass::gemm::collective::detail::tma_alignment_bytes>();
// Current TMA epilogues require sixteen-bit data types and epilogue tile M to be of size 64.
// Only dispatch to the TMA builder if these requirements are satisfied.
static constexpr bool IsSixteenBit = sizeof_bits<ElementC>::value == 16 && sizeof_bits<ElementD>::value == 16;
static constexpr bool IsEpiTileM64 = size<0>(shape(TileShape_MNK{})) == 64;
using _CollectiveBuilder = CollectiveBuilder<
arch::Sm90,
arch::OpClassTensorOp,
TileShape_MNK,
ClusterShape_MNK,
EpilogueTileType,
ElementAccumulator,
ElementCompute,
ElementC,
GmemLayoutTagC,
AlignmentC,
ElementD,
GmemLayoutTagD,
AlignmentD,
cute::conditional_t<IsTmaAligned && IsSixteenBit && IsEpiTileM64,
TmaWarpSpecialized, NoSmemWarpSpecialized>
>;
public:
using ThreadOp = typename _CollectiveBuilder::ThreadOp;
using CollectiveOp = typename _CollectiveBuilder::CollectiveOp;
};
// Tma warp-specialized builder for elementwise fusion
template <
class TileShape_MNK,
class ClusterShape_MNK,
class EpilogueTileType,
class ElementAccumulator,
class ElementCompute,
class ElementC,
class GmemLayoutTagC,
int AlignmentC,
class ElementD,
class GmemLayoutTagD,
int AlignmentD,
class Schedule
>
struct CollectiveBuilder<
arch::Sm90,
arch::OpClassTensorOp,
TileShape_MNK,
ClusterShape_MNK,
EpilogueTileType,
ElementAccumulator,
ElementCompute,
ElementC,
GmemLayoutTagC,
AlignmentC,
ElementD,
GmemLayoutTagD,
AlignmentD,
Schedule,
cute::enable_if_t<cute::is_base_of_v<TmaWarpSpecializedElementwiseBase, Schedule> ||
cute::is_base_of_v<TmaWarpSpecializedCooperativeElementwiseBase, Schedule> >> {
public:
using ThreadOp = thread::LinearCombinationGeneric<
Schedule::ActivationFunctor,
ElementD, AlignmentD,
ElementAccumulator, ElementCompute, Schedule::Scale,
Schedule::Round>;
private:
using Impl = detail::TmaBuilderImpl<
TileShape_MNK, ClusterShape_MNK, EpilogueTileType, ElementAccumulator, ElementCompute,
ElementC, GmemLayoutTagC, AlignmentC, ElementD, GmemLayoutTagD, AlignmentD,
Schedule, ThreadOp, cutlass::epilogue::Sm90TmaWarpSpecialized<1,2,true>>;
public:
using CollectiveOp = typename Impl::CollectiveOp;
};
// Tma warp-specialized builder for bias + elementwise fusion
template <
class TileShape_MNK,
class ClusterShape_MNK,
class EpilogueTileType,
class ElementAccumulator,
class ElementCompute,
class ElementC,
class GmemLayoutTagC,
int AlignmentC,
class ElementD,
class GmemLayoutTagD,
int AlignmentD,
class Schedule
>
struct CollectiveBuilder<
arch::Sm90,
arch::OpClassTensorOp,
TileShape_MNK,
ClusterShape_MNK,
EpilogueTileType,
ElementAccumulator,
ElementCompute,
ElementC,
GmemLayoutTagC,
AlignmentC,
ElementD,
GmemLayoutTagD,
AlignmentD,
Schedule,
cute::enable_if_t<cute::is_base_of_v<TmaWarpSpecializedBiasElementwiseBase, Schedule> ||
cute::is_base_of_v<TmaWarpSpecializedCooperativeBiasElementwiseBase, Schedule> >> {
public:
using ThreadOp = thread::LinearCombinationBiasElementwise<
ElementC, ElementAccumulator, ElementCompute, ElementD, typename Schedule::ElementT, AlignmentD,
typename Schedule::ActivationFunctor<ElementCompute>, typename Schedule::BiasOp<ElementCompute>,
Schedule::StoreT, typename Schedule::ElementBias>;
private:
using Impl = detail::TmaBuilderImpl<
TileShape_MNK, ClusterShape_MNK, EpilogueTileType, ElementAccumulator, ElementCompute,
ElementC, GmemLayoutTagC, AlignmentC, ElementD, GmemLayoutTagD, AlignmentD,
Schedule, ThreadOp, cutlass::epilogue::Sm90TmaWarpSpecializedBiasElementwise<1,2>>;
public:
using CollectiveOp = typename Impl::CollectiveOp;
};
// CollectiveBuilder that transposed epilogue below is used for sm90 gmma RS TT kernels
// since swapping NNN kernels input matrix and transposing its output at the same time then
// we can get TTN kernel.
template <
class TileShape_MNK,
class ClusterShape_MNK,
class EpilogueTileType,
class ElementAccumulator,
class ElementCompute,
class ElementC_,
class GmemLayoutTagC_,
int AlignmentC,
class ElementD,
class GmemLayoutTagD,
int AlignmentD,
class Schedule
>
struct CollectiveBuilder<
arch::Sm90,
arch::OpClassTensorOp,
TileShape_MNK,
ClusterShape_MNK,
EpilogueTileType,
ElementAccumulator,
ElementCompute,
ElementC_,
GmemLayoutTagC_,
AlignmentC,
ElementD,
GmemLayoutTagD,
AlignmentD,
Schedule,
cute::enable_if_t<cute::is_same_v<Schedule, cutlass::gemm::EpilogueTransposed>>> {
// Passing void C disables source load
using ElementC = cute::conditional_t<cute::is_void_v<ElementC_>,
ElementD, ElementC_>; // prevents cute breakages
using GmemLayoutTagC = cute::conditional_t<cute::is_void_v<ElementC_>,
GmemLayoutTagD, GmemLayoutTagC_>;
static constexpr thread::ScaleType::Kind ScaleType = cute::is_void_v<ElementC_> ?
thread::ScaleType::OnlyAlphaScaling : thread::ScaleType::Default;
using ThreadOp = thread::LinearCombination<
ElementD, 1, ElementAccumulator, ElementCompute,
ScaleType, FloatRoundStyle::round_to_nearest, ElementC>;
using CollectiveOp = cutlass::epilogue::collective::detail::Sm90TmaWarpSpecializedAdapter<
cutlass::epilogue::collective::DefaultEpilogue<
cutlass::gemm::TagToStrideC_t<GmemLayoutTagC>,
cutlass::gemm::TagToStrideC_t<GmemLayoutTagD>,
ThreadOp,
cutlass::gemm::EpilogueTransposed>
>;
};
///////////////////////////////////////////////////////////////////////////////
} // namespace cutlass::epilogue::collective
@@ -0,0 +1,77 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include "cutlass/detail/dependent_false.hpp"
/////////////////////////////////////////////////////////////////////////////////////////////////
namespace cutlass::epilogue::collective {
/////////////////////////////////////////////////////////////////////////////////////////////////
// Used to specify epilogue subtile shape or dispatch to automatic computation of subtile shape
struct EpilogueTileAuto {};
// Used to let the builder pick the epilogue schedule automatically.
// Can be overridden with kernel schedule tags in cutlass/gemm/dispatch_policy.hpp
struct EpilogueScheduleAuto {};
template <
class ArchTag,
class OpClass,
class TileShape_MNK,
class ClusterShape_MNK,
class EpilogueTileType,
class ElementAccumulator,
class ElementCompute,
class ElementC,
class GmemLayoutTagC,
int AlignmentC,
class ElementD,
class GmemLayoutTagD,
int AlignmentD,
class Schedule,
class Enable = void
>
struct CollectiveBuilder {
static_assert(cutlass::detail::dependent_false<ArchTag>,
"Could not build a collective epilogue for given parameters.");
};
/////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace cutlass::epilogue::collective
/////////////////////////////////////////////////////////////////////////////////////////////////
#include "builders/sm90_builder.inl"
/////////////////////////////////////////////////////////////////////////////////////////////////
@@ -24,6 +24,8 @@
**************************************************************************************************/
#pragma once
#include <cutlass/detail/dependent_false.hpp>
/////////////////////////////////////////////////////////////////////////////////////////////////
namespace cutlass::epilogue::collective {
@@ -34,8 +36,8 @@ template <
class DispatchPolicy,
class... Args
>
struct CollectiveEpilogue {
static_assert(std::is_void_v<DispatchPolicy>, "Could not find an epilogue specialization.");
class CollectiveEpilogue {
static_assert(cutlass::detail::dependent_false<DispatchPolicy>, "Could not find an epilogue specialization.");
};
/////////////////////////////////////////////////////////////////////////////////////////////////
@@ -44,6 +46,10 @@ struct CollectiveEpilogue {
/////////////////////////////////////////////////////////////////////////////////////////////////
#include "detail.hpp"
#include "default_epilogue.hpp"
#include "epilogue.hpp"
#include "epilogue_tensor_broadcast.hpp"
#include "sm70_epilogue_vectorized.hpp"
#include "sm90_epilogue_tma_warpspecialized.hpp"
#include "sm90_epilogue_tma_warpspecialized_bias_elementwise.hpp"
/////////////////////////////////////////////////////////////////////////////////////////////////
@@ -35,6 +35,8 @@
#pragma once
#include "cutlass/cutlass.h"
#include "cutlass/gemm/dispatch_policy.hpp"
#include "cutlass/epilogue/collective/detail.hpp"
#include "cute/tensor.hpp"
#include "cute/numeric/int.hpp"
@@ -52,13 +54,16 @@ namespace collective {
template <
class StrideC_,
class StrideD_,
class ThreadEpilogueOp_
class ThreadEpilogueOp_,
class EpilogueSchedule_
>
class DefaultEpilogue {
public:
//
// Type Aliases
//
using EpilogueSchedule = EpilogueSchedule_;
// derived types of output thread level operator
using ThreadEpilogueOp = ThreadEpilogueOp_;
using ElementOutput = typename ThreadEpilogueOp::ElementOutput;
@@ -78,28 +83,40 @@ public:
struct SharedStorage { };
// Params of epilogue::collective contain the epilogue::thread params
struct Params {
// Host side epilgoue arguments
struct Arguments {
typename ThreadEpilogueOp::Params thread{};
ElementC const* ptr_C = nullptr;
StrideC dC{};
ElementD* ptr_D = nullptr;
StrideD dD{};
typename ThreadEpilogueOp::Params thread_params{};
};
// Device side epilogue params
using Params = Arguments;
//
// Methods
//
template <class Args>
template <class ProblemShape>
static constexpr Params
to_underlying_arguments(Args const& args, void* workspace) {
(void) workspace;
return {args.epilogue_params};
to_underlying_arguments(
[[maybe_unused]] ProblemShape const& _,
Arguments const& args,
[[maybe_unused]] void* workspace) {
return args;
}
CUTLASS_HOST_DEVICE
DefaultEpilogue(Params const& params_) : params(params_) { }
DefaultEpilogue(Params const& params_)
: params(params_), epilogue_op(params_.thread) { }
CUTLASS_DEVICE
bool
is_source_needed() {
return epilogue_op.is_source_needed();
}
template<
class ProblemShapeMNKL,
@@ -118,7 +135,7 @@ public:
TiledMma tiled_mma,
ResidueMNK residue_mnk,
int thread_idx,
char* smem_buf)
[[maybe_unused]] char* smem_buf)
{
using namespace cute;
using X = Underscore;
@@ -128,17 +145,17 @@ public:
static_assert(rank(BlockShapeMNK{}) == 3, "BlockShapeMNK must be rank 3");
static_assert(rank(BlockCoordMNKL{}) == 4, "BlockCoordMNKL must be rank 3");
(void) smem_buf;
ThreadEpilogueOp epilogue_op{params.thread_params};
// Separate out problem shape for convenience
auto M = get<0>(problem_shape_mnkl);
auto N = get<1>(problem_shape_mnkl);
auto L = get<3>(problem_shape_mnkl);
auto stride_c = detail::get_epilogue_stride<EpilogueSchedule>(params.dC);
auto stride_d = detail::get_epilogue_stride<EpilogueSchedule>(params.dD);
// Represent the full output tensor
Tensor mC_mnl = make_tensor(make_gmem_ptr(params.ptr_C), make_shape(M,N,L), params.dC); // (m,n,l)
Tensor mD_mnl = make_tensor(make_gmem_ptr(params.ptr_D), make_shape(M,N,L), params.dD); // (m,n,l)
Tensor mC_mnl = make_tensor(make_gmem_ptr(params.ptr_C), make_shape(M,N,L), stride_c); // (m,n,l)
Tensor mD_mnl = make_tensor(make_gmem_ptr(params.ptr_D), make_shape(M,N,L), stride_d); // (m,n,l)
Tensor gC_mnl = local_tile(mC_mnl, blk_shape_MNK, make_coord(_,_,_), Step<_1,_1, X>{}); // (BLK_M,BLK_N,m,n,l)
Tensor gD_mnl = local_tile(mD_mnl, blk_shape_MNK, make_coord(_,_,_), Step<_1,_1, X>{}); // (BLK_M,BLK_N,m,n,l)
@@ -184,6 +201,7 @@ public:
private:
Params params;
ThreadEpilogueOp epilogue_op;
};
/////////////////////////////////////////////////////////////////////////////////////////////////
@@ -0,0 +1,211 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include "cutlass/cutlass.h"
#include "cutlass/pipeline/pipeline.hpp"
#include "cutlass/gemm/dispatch_policy.hpp"
#include "cutlass/epilogue/dispatch_policy.hpp"
#include "cute/tensor.hpp"
#include "cute/numeric/int.hpp"
/////////////////////////////////////////////////////////////////////////////////////////////////
namespace cutlass {
namespace epilogue {
namespace collective {
namespace detail {
template <class T>
static constexpr int elements_per_access_v = cutlass::sizeof_bits<uint32_t>::value / cutlass::sizeof_bits<T>::value;
template <class EpilogueSchedule>
static constexpr bool sm90_is_cooperative_v =
std::is_base_of_v<cutlass::epilogue::TmaWarpSpecializedCooperative, EpilogueSchedule>;
template <class EpilogueSchedule>
static constexpr bool sm90_is_warp_specialized_v =
std::is_base_of_v<cutlass::epilogue::TmaWarpSpecialized, EpilogueSchedule>;
template <class T>
struct EmptyStorage {
CUTLASS_HOST_DEVICE
T* data() { return nullptr; }
};
template<class EpilogueSchedule, class Stride>
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;
}
}
template <typename ThreadEpilogueOp, typename = void>
struct IsThreadEpilogueOpWithBias {
static constexpr bool value = false;
using type = typename ThreadEpilogueOp::ElementCompute;
};
template <typename ThreadEpilogueOp>
struct IsThreadEpilogueOpWithBias <ThreadEpilogueOp, cute::void_t<typename ThreadEpilogueOp::ElementBias>> {
static constexpr bool value = true;
using type = typename ThreadEpilogueOp::ElementBias;
};
// IF_EPILOGUE_USES_TMA<T>::value will be true only if:
// class T has member CopyOpS2G and T::CopyOpS2G is true
template <typename T, typename = void>
struct IF_EPILOGUE_USES_TMA { static constexpr bool value = false; };
template <typename T>
struct IF_EPILOGUE_USES_TMA <T, void_t<typename T::CopyOpS2G>>
{ static constexpr bool value = true; };
// Wrapper class to use operator-style epilogues in sm90 TMA warp-specialized kernels
template <class EpilogueOp>
class Sm90TmaWarpSpecializedAdapter : public EpilogueOp {
public:
using LoadPipeline = cutlass::PipelineTransactionAsync<0>; // 0 stage to disable smem alloc
using LoadPipelineState = cutlass::PipelineState<0>;
constexpr static uint32_t TmaTransactionBytes = 0;
using StorePipeline = cutlass::PipelineTmaStore<1>; // tma store pipe has no smem alloc
using StorePipelineState = cutlass::PipelineState<1>;
using TensorStorage = typename EpilogueOp::SharedStorage;
using PipelineStorage = typename LoadPipeline::SharedStorage;
template<class TileShapeMNK>
CUTLASS_HOST_DEVICE
static constexpr int
get_load_pipe_increment([[maybe_unused]] TileShapeMNK) {
return 1;
}
template<class TileShapeMNK>
CUTLASS_HOST_DEVICE
static constexpr int
get_store_pipe_increment([[maybe_unused]] TileShapeMNK) {
return 1;
}
CUTLASS_DEVICE
static void prefetch_tma_descriptors([[maybe_unused]] typename EpilogueOp::Params const&)
{
}
// ctor inheritance
using EpilogueOp::EpilogueOp;
template<
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class TiledMma
>
CUTLASS_DEVICE void
load(
[[maybe_unused]] LoadPipeline load_pipeline,
[[maybe_unused]] LoadPipelineState load_pipe_producer_state,
[[maybe_unused]] ProblemShapeMNKL problem_shape_mnkl,
[[maybe_unused]] TileShapeMNK tile_shape_MNK,
[[maybe_unused]] TileCoordMNKL tile_coord_mnkl,
[[maybe_unused]] TiledMma tiled_mma,
[[maybe_unused]] int thread_idx,
[[maybe_unused]] TensorStorage& shared_tensors)
{
// source load is performed in epilogue operator
}
CUTLASS_DEVICE void
load_tail(
[[maybe_unused]] LoadPipeline load_pipeline,
[[maybe_unused]] LoadPipelineState load_pipe_producer_state)
{
}
template<
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class AccEngine, class AccLayout,
class TiledMma
>
CUTLASS_DEVICE void
store(
[[maybe_unused]] LoadPipeline load_pipeline,
[[maybe_unused]] LoadPipelineState load_pipe_consumer_state,
[[maybe_unused]] StorePipeline store_pipeline,
[[maybe_unused]] StorePipelineState store_pipe_producer_state,
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_MNK,
TileCoordMNKL tile_coord_mnkl,
cute::Tensor<AccEngine,AccLayout> accumulators,
TiledMma tiled_mma,
int thread_idx,
TensorStorage& shared_tensors)
{
constexpr int BLK_M_RANK = rank<0>(tile_shape_MNK);
auto m_max_coord = unwrap(cute::transform(make_seq<BLK_M_RANK>{}, [&](auto i) {
return get<0,i>(problem_shape_mnkl) - get<0,i>(tile_shape_MNK) * get<0,i>(tile_coord_mnkl);
}));
constexpr int BLK_N_RANK = rank<1>(tile_shape_MNK);
auto n_max_coord = unwrap(cute::transform(make_seq<BLK_N_RANK>{}, [&](auto i) {
return get<1,i>(problem_shape_mnkl) - get<1,i>(tile_shape_MNK) * get<1,i>(tile_coord_mnkl);
}));
auto residue_mnk = make_tuple(m_max_coord, n_max_coord, Int<0>{});
(*this)(
problem_shape_mnkl,
tile_shape_MNK,
tile_coord_mnkl,
accumulators,
tiled_mma,
residue_mnk,
thread_idx,
reinterpret_cast<char*>(&shared_tensors));
}
};
} // namespace detail
} // namespace collective
} // namespace epilogue
} // namespace cutlass
@@ -28,83 +28,117 @@
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
/*! \file
\brief Functor performing elementwise operations used by epilogues.
\brief Functor for performing tensor-tensor broadacasts atop existing epilogues.
Concretely, the opeartion performed is the following:
UnaryOp(
BinaryOp1(
BinaryOp0(
Activation((alpha * A @ B) + bias),
beta * C0
),
beta * C1
)
)
where:
- C0 and C1 have the same extents as the output
- BinaryOp0 and BinaryOp1 perform elementwise binary operations
- UnaryOp is an elementwise operation
*/
#pragma once
#include "cutlass/cutlass.h"
#include "cutlass/epilogue/collective/detail.hpp"
#include "cute/tensor.hpp"
#include "cute/numeric/int.hpp"
/////////////////////////////////////////////////////////////////////////////////////////////////
namespace cutlass {
namespace epilogue {
namespace collective {
/////////////////////////////////////////////////////////////////////////////////////////////////
using namespace cute;
/////////////////////////////////////////////////////////////////////////////////////////////////
/// Applies an element wise operation to all elements within the fragment
/// and writes them out to destination storage.
/// Collective epilogue that applies elementwise tensor-tensor operations atop other epilogues
///
template <
class StrideC_,
class StrideD_,
class ThreadEpilogueOp_
class ThreadEpilogueOp_,
class EpilogueSchedule_
>
class DefaultTransposedEpilogue {
class EpilogueTensorBroadcast {
public:
//
// Type Aliases
//
using EpilogueSchedule = EpilogueSchedule_;
// derived types of output thread level operator
using ThreadEpilogueOp = ThreadEpilogueOp_;
using ElementOutput = typename ThreadEpilogueOp::ElementOutput;
using ElementAccumulator = typename ThreadEpilogueOp::ElementAccumulator;
using ElementCompute = typename ThreadEpilogueOp::ElementCompute;
using ElementScalar = ElementCompute;
using ElementBias = typename ThreadEpilogueOp::ElementBias;
using ElementC = typename ThreadEpilogueOp::ElementC;
using StrideC = StrideC_;
using ElementD = typename ThreadEpilogueOp::ElementD;
using StrideD = StrideD_;
static const int kOutputAlignment = ThreadEpilogueOp::kCount;
using AlignmentType = typename cute::uint_bit<sizeof_bits<ElementOutput>::value * kOutputAlignment>::type;
using ActivationFunctor = typename ThreadEpilogueOp::ActivationFunctor;
static_assert(rank(StrideC{}) == 3, "StrideCD must be rank-3: [M, N, L]");
static_assert(rank(StrideD{}) == 3, "StrideCD must be rank-3: [M, N, L]");
static constexpr int kOutputAlignment = ThreadEpilogueOp::kCount;
using AlignmentType = typename cute::uint_bit<sizeof_bits<ElementOutput>::value * kOutputAlignment>::type;
static constexpr bool IsBinaryOp0Enabled = ThreadEpilogueOp::IsBinaryOp0Enabled;
static constexpr bool IsBinaryOp1Enabled = ThreadEpilogueOp::IsBinaryOp1Enabled;
static constexpr bool IsUnaryOpEnabled = ThreadEpilogueOp::IsUnaryOpEnabled;
struct SharedStorage { };
// Params of epilogue::collective contain the epilogue::thread params
struct Params {
ElementC const* ptr_C = nullptr;
// Host side epilogue arguments
struct Arguments {
typename ThreadEpilogueOp::Params thread{};
StrideC dC{};
ElementD* ptr_D = nullptr;
StrideD dD{};
typename ThreadEpilogueOp::Params thread_params{};
ElementBias* ptr_Bias = nullptr;
ElementC* ptr_C0 = nullptr;
ElementC* ptr_C1 = nullptr;
};
// Device side epilogue params
using Params = Arguments;
//
// Methods
//
template <class Args>
template <class ProblemShape>
static constexpr Params
to_underlying_arguments(Args const& args, void* workspace) {
(void) workspace;
return {args.epilogue_params};
to_underlying_arguments(
[[maybe_unused]] ProblemShape const& _,
Arguments const& args,
[[maybe_unused]] void* workspace) {
return args;
}
CUTLASS_HOST_DEVICE
DefaultTransposedEpilogue(Params const& params_) : params(params_) { }
EpilogueTensorBroadcast(Params const& params_)
: params(params_), epilogue_op(params_.thread) { }
CUTLASS_DEVICE
bool
is_source_needed() {
return epilogue_op.is_source0_needed() || epilogue_op.is_source1_needed();
}
template<
class ProblemShapeMNKL,
@@ -123,7 +157,7 @@ public:
TiledMma tiled_mma,
ResidueMNK residue_mnk,
int thread_idx,
char* smem_buf)
[[maybe_unused]] char* smem_buf)
{
using namespace cute;
using X = Underscore;
@@ -131,67 +165,75 @@ public:
static_assert(rank(ProblemShapeMNKL{}) == 4, "ProblemShapeMNKL must be rank 4");
static_assert(is_static<BlockShapeMNK>::value, "ThreadBlock tile shape must be static");
static_assert(rank(BlockShapeMNK{}) == 3, "BlockShapeMNK must be rank 3");
static_assert(rank(BlockCoordMNKL{}) == 4, "BlockCoordMNKL must be rank 3");
(void) smem_buf;
ThreadEpilogueOp epilogue_op{params.thread_params};
static_assert(rank(BlockCoordMNKL{}) == 4, "BlockCoordMNKL must be rank 4");
// Separate out problem shape for convenience
auto M = get<0>(problem_shape_mnkl);
auto N = get<1>(problem_shape_mnkl);
auto L = get<3>(problem_shape_mnkl);
// Tranpose stride C/D.
auto stride_c = make_stride(get<1>(params.dC), get<0>(params.dC), get<2>(params.dC));
auto stride_d = make_stride(get<1>(params.dD), get<0>(params.dD), get<2>(params.dD));
auto stride_c = detail::get_epilogue_stride<EpilogueSchedule>(params.dC);
auto stride_d = detail::get_epilogue_stride<EpilogueSchedule>(params.dD);
auto stride_bias = detail::get_epilogue_stride<EpilogueSchedule>(Stride<_1, _0, _0>{});
// Represent the full output tensor
Tensor mC_mnl = make_tensor(make_gmem_ptr(params.ptr_C), make_shape(M,N,L), stride_c); // (m,n,l)
Tensor mD_mnl = make_tensor(make_gmem_ptr(params.ptr_D), make_shape(M,N,L), stride_d); // (m,n,l)
Tensor gC_mnl = local_tile(mC_mnl, blk_shape_MNK, make_coord(_,_,_), Step<_1,_1, X>{}); // (BLK_M,BLK_N,m,n,l)
Tensor gD_mnl = local_tile(mD_mnl, blk_shape_MNK, make_coord(_,_,_), Step<_1,_1, X>{}); // (BLK_M,BLK_N,m,n,l)
Tensor mC0_mnl = make_tensor(make_gmem_ptr(params.ptr_C0), make_shape(M,N,L), stride_c); // (m,n,l)
Tensor mC1_mnl = make_tensor(make_gmem_ptr(params.ptr_C1), make_shape(M,N,L), stride_c); // (m,n,l)
Tensor mD_mnl = make_tensor(make_gmem_ptr(params.ptr_D), make_shape(M,N,L), stride_d); // (m,n,l)
Tensor mBias_mnl = make_tensor(make_gmem_ptr(params.ptr_Bias), make_shape(M,N,L), stride_bias); // (m,n,l)
// Slice to get the tile this CTA is responsible for
Tensor gC0_mnl = local_tile(mC0_mnl, blk_shape_MNK, make_coord(_,_,_), Step<_1,_1, X>{}); // (BLK_M,BLK_N,m,n,l)
Tensor gC1_mnl = local_tile(mC1_mnl, blk_shape_MNK, make_coord(_,_,_), Step<_1,_1, X>{}); // (BLK_M,BLK_N,m,n,l)
Tensor gD_mnl = local_tile(mD_mnl, blk_shape_MNK, make_coord(_,_,_), Step<_1,_1, X>{}); // (BLK_M,BLK_N,m,n,l)
Tensor gBias_mnl = local_tile(mBias_mnl, blk_shape_MNK, make_coord(_,_,_), Step<_1,_1, X>{}); // (BLK_M,BLK_N,m,n,l)
// Slice to get the tile this thread block is responsible for
auto [m_coord, n_coord, k_coord, l_coord] = blk_coord_mnkl;
Tensor gC = gC_mnl(_,_,m_coord,n_coord,l_coord); // (BLK_M,BLK_N)
Tensor gD = gD_mnl(_,_,m_coord,n_coord,l_coord); // (BLK_M,BLK_N)
Tensor gC0 = gC0_mnl(_,_,m_coord,n_coord,l_coord); // (BLK_M,BLK_N)
Tensor gC1 = gC1_mnl(_,_,m_coord,n_coord,l_coord); // (BLK_M,BLK_N)
Tensor gD = gD_mnl(_,_,m_coord,n_coord,l_coord); // (BLK_M,BLK_N)
Tensor gBias = gBias_mnl(_,_,m_coord,n_coord,l_coord); // (BLK_M,BLK_N)
// Partition source and destination tiles to match the accumulator partitioning
auto thr_mma = tiled_mma.get_thread_slice(thread_idx);
Tensor tCgD = thr_mma.partition_C(gD); // (VEC,THR_M,THR_N)
Tensor tCgC = thr_mma.partition_C(gC); // (VEC,THR_M,THR_N)
Tensor tCgD = thr_mma.partition_C(gD); // (VEC,THR_M,THR_N)
Tensor tCgC0 = thr_mma.partition_C(gC0); // (VEC,THR_M,THR_N)
Tensor tCgC1 = thr_mma.partition_C(gC1); // (VEC,THR_M,THR_N)
Tensor tCgBias = thr_mma.partition_C(gBias); // (VEC,THR_M,THR_N)
static_assert(is_static<FrgLayout>::value, "Accumulator layout must be static");
CUTE_STATIC_ASSERT_V(size(tCgC) == size(tCgD),
static_assert(is_static<FrgLayout>::value,
"Accumulator layout must be static");
CUTE_STATIC_ASSERT_V(size(tCgC0) == size(tCgD),
"Source and destination must have the same number of elements.");
CUTE_STATIC_ASSERT_V(size(tCgC1) == size(tCgD),
"Source and destination must have the same number of elements.");
CUTE_STATIC_ASSERT_V(size(tCgD) == size(accumulators),
"Accumulator count must have the same destination element count.");
CUTE_STATIC_ASSERT_V(size(tCgBias) == size(accumulators),
"Accumulator count must have the same destination element count.");
auto cD = make_identity_tensor(make_shape(unwrap(shape<0>(gD)), unwrap(shape<1>(gD))));
auto cD = make_identity_tensor(make_shape(unwrap(shape<0>(gD)), unwrap(shape<1>(gD))));
Tensor tCcD = thr_mma.partition_C(cD);
// source is needed
if (epilogue_op.is_source_needed()) {
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < size(accumulators); ++i) {
if (elem_less(tCcD(i), make_coord(get<0>(residue_mnk), get<1>(residue_mnk)))) {
tCgD(i) = epilogue_op(accumulators(i), tCgC(i));
}
}
}
// source is not needed, avoid load
else {
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < size(accumulators); ++i) {
if (elem_less(tCcD(i), make_coord(get<0>(residue_mnk), get<1>(residue_mnk)))) {
tCgD(i) = epilogue_op(accumulators(i));
}
bool bias_needed = params.ptr_Bias != nullptr;
bool c0_needed = (params.ptr_C0 != nullptr) && epilogue_op.is_source0_needed();
bool c1_needed = (params.ptr_C1 != nullptr) && epilogue_op.is_source1_needed();
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < size(accumulators); ++i) {
if (elem_less(tCcD(i), make_coord(get<0>(residue_mnk), get<1>(residue_mnk)))) {
ElementBias bias = bias_needed ? tCgBias(i) : ElementBias(0);
ElementC c0 = c0_needed ? tCgC0(i) : ElementC(0);
ElementC c1 = c1_needed ? tCgC1(i) : ElementC(0);
tCgD(i) = epilogue_op(accumulators(i), c0, c1, bias);
}
}
}
private:
Params params;
ThreadEpilogueOp epilogue_op;
};
/////////////////////////////////////////////////////////////////////////////////////////////////
@@ -95,28 +95,40 @@ public:
cute::array_aligned<ElementAccumulator, cute::cosize_v<SmemLayout>> smem_epilogue;
};
// Params of epilogue::collective contain the epilogue::thread params
struct Params {
// Host side epilogue arguments
struct Arguments {
typename ThreadEpilogueOp::Params thread{};
ElementC const* ptr_C = nullptr;
StrideC dC{};
ElementD* ptr_D = nullptr;
StrideD dD{};
typename ThreadEpilogueOp::Params thread_params{};
};
// Device side epilogue params
using Params = Arguments;
//
// Methods
//
template <class Args>
template <class ProblemShape>
static constexpr Params
to_underlying_arguments(Args const& args, void* workspace) {
(void) workspace;
return {args.epilogue_params};
to_underlying_arguments(
[[maybe_unused]] ProblemShape const& _,
Arguments const& args,
[[maybe_unused]] void* workspace) {
return args;
}
CUTLASS_HOST_DEVICE
Epilogue(Params const& params_) : params(params_) { };
Epilogue(Params const& params_)
: params(params_), epilogue_op(params_.thread) { }
CUTLASS_DEVICE
bool
is_source_needed() {
return epilogue_op.is_source_needed();
}
template<
class ProblemShapeMNKL,
@@ -147,13 +159,11 @@ public:
// synchronizing function for smem reads/writes
#if CUDA_BARRIER_ENABLED
auto synchronize = [] () { NamedBarrier::sync(typename TiledCopyS2R::TiledNumThr{}, 0); };
auto synchronize = [] () { cutlass::arch::NamedBarrier::sync(typename TiledCopyS2R::TiledNumThr{}, 0); };
#else
auto synchronize = [] () { __syncthreads(); };
#endif
ThreadEpilogueOp epilogue_op{this->params.thread_params};
// Separate out problem shape for convenience
auto M = get<0>(problem_shape_mnkl);
auto N = get<1>(problem_shape_mnkl);
@@ -175,7 +185,8 @@ public:
Tensor sC = make_tensor(make_smem_ptr(storage.smem_epilogue.data()), SmemLayout{}); // (SMEM_M,SMEM_N)
// Partition sC to match the accumulator partitioning
auto tC = make_tiled_copy_C(CopyAtomR2S{}, tiled_mma).get_thread_slice(thread_idx);
auto tiled_r2s = make_tiled_copy_C(CopyAtomR2S{}, tiled_mma);
auto tC = tiled_r2s.get_thread_slice(thread_idx);
Tensor tCaC = tC.retile_S(accumulators); // ((Atom,AtomNum), MMA_M, MMA_N)
Tensor tCsC = tC.partition_D(sC); // ((Atom,AtomNum),PIPE_M,PIPE_N)
@@ -185,7 +196,8 @@ public:
Tensor gDt = local_tile(gD, tile, _); // (SMEM_M,SMEM_N,TILE_M,TILE_N)
// Partition sC, gC, and gD for the output
auto tD = TiledCopyS2R{}.get_thread_slice(thread_idx);
auto tiled_s2r = TiledCopyS2R{};
auto tD = tiled_s2r.get_thread_slice(thread_idx);
Tensor tDsC = tD.partition_S(sC); // ((Atom,AtomNum),ATOM_M,ATOM_N)
Tensor tDgC = tD.partition_D(gCt); // ((Atom,AtomNum),ATOM_M,ATOM_N,TILE_M,TILE_N)
Tensor tDgD = tD.partition_D(gDt); // ((Atom,AtomNum),ATOM_M,ATOM_N,TILE_M,TILE_N)
@@ -239,7 +251,7 @@ public:
int mma_m = step_m * size<1>(tCsC) + pipe_m;
int mma_n = step_n * size<2>(tCsC) + pipe_n;
copy(tC, tCaC(_,mma_m,mma_n), tCsC(_,pipe_m,pipe_n));
copy(tiled_r2s, tCaC(_,mma_m,mma_n), tCsC(_,pipe_m,pipe_n));
}
}
@@ -247,7 +259,7 @@ public:
synchronize();
// Step 3. Copy from SMEM into a fragment
copy(tD, tDsC, tDrC);
copy(tiled_s2r, tDsC, tDrC);
// Step 4. Wait for SMEM reads to complete
synchronize();
@@ -310,6 +322,7 @@ public:
private:
Params params;
ThreadEpilogueOp epilogue_op;
};
@@ -0,0 +1,582 @@
/***************************************************************************************************
* Copyright (c) 2023, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * Redistributions in binary form must reproduce the above copyright notice, this list of
* conditions and the following disclaimer in the documentation and/or other materials
* provided with the distribution.
* * Neither the name of the NVIDIA CORPORATION nor the names of its contributors may be used
* to endorse or promote products derived from this software without specific prior written
* permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND ANY EXPRESS OR
* IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND
* FITNESS FOR A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL NVIDIA CORPORATION BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING,
* BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS;
* OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT,
* STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
/*! \file
\brief Functor performing elementwise operations used by epilogues.
*/
#pragma once
#include "cutlass/cutlass.h"
#include "cutlass/arch/barrier.h"
#include "cutlass/epilogue/dispatch_policy.hpp"
#include "cutlass/epilogue/collective/detail.hpp"
#include "cutlass/epilogue/thread/scale_type.h"
#include "cute/tensor.hpp"
/////////////////////////////////////////////////////////////////////////////////////////////////
namespace cutlass {
namespace epilogue {
namespace collective {
/////////////////////////////////////////////////////////////////////////////////////////////////
template <
int StagesC_,
int StagesD_,
bool DisableSmemReuseC_,
class BlockTileShape_, // (BLK_M,BLK_N,BLK_K)
class EpilogueTile_, // (EPI_TILE_M,EPI_TILE_N) per-collective
class ElementC_,
class StrideC_,
class ElementD_,
class StrideD_,
class ThreadEpilogueOp_,
class CopyOpG2S_,
class SmemLayoutAtomC_,
class CopyOpS2R_,
class CopyOpS2G_,
class SmemLayoutAtomD_,
class CopyOpR2S_
>
class CollectiveEpilogue<
Sm90TmaWarpSpecialized<StagesC_,StagesD_,DisableSmemReuseC_>,
BlockTileShape_,
EpilogueTile_,
ElementC_,
StrideC_,
ElementD_,
StrideD_,
ThreadEpilogueOp_,
CopyOpG2S_,
SmemLayoutAtomC_,
CopyOpS2R_,
CopyOpS2G_,
SmemLayoutAtomD_,
CopyOpR2S_
> {
public:
//
// Type Aliases
//
// derived types of output thread level operator
using DispatchPolicy = Sm90TmaWarpSpecialized<StagesC_,StagesD_,DisableSmemReuseC_>;
using BlockTileShape = BlockTileShape_;
using EpilogueTile = EpilogueTile_;
using ThreadEpilogueOp = ThreadEpilogueOp_;
using ElementAccumulator = typename ThreadEpilogueOp::ElementAccumulator;
using ElementCompute = typename ThreadEpilogueOp::ElementCompute;
using ElementScalar = ElementCompute;
using ElementBias = typename detail::IsThreadEpilogueOpWithBias<ThreadEpilogueOp>::type;
using ElementOutput = typename ThreadEpilogueOp::ElementOutput;
using ElementC = ElementC_;
using StrideC = StrideC_;
using ElementD = ElementD_;
using StrideD = StrideD_;
using CopyOpG2S = CopyOpG2S_;
using SmemLayoutAtomC = SmemLayoutAtomC_;
using CopyOpS2R = CopyOpS2R_;
using CopyOpS2G = CopyOpS2G_;
using SmemLayoutAtomD = SmemLayoutAtomD_;
using CopyOpR2S = CopyOpR2S_;
constexpr static int kOutputAlignment = ThreadEpilogueOp::kCount;
constexpr static bool iskThreadEpilogueOpWithBias = detail::IsThreadEpilogueOpWithBias<ThreadEpilogueOp>::value;
using AlignmentType = typename uint_bit<sizeof_bits<ElementOutput>::value * kOutputAlignment>::type;
static_assert(sizeof(ElementC) == 2, "Only 16b source supported for now");
static_assert(sizeof(ElementD) == 2, "Only 16b output supported for now");
static_assert(!is_layout<EpilogueTile>::value && is_tuple<EpilogueTile>::value, "EpilogueTile must be a cute::Tile or cute::Shape");
static_assert(rank(BlockTileShape{}) == 3, "BlockTileShape must be rank-3: [BLK_M,BLK_N,BLK_K]");
static_assert(rank(EpilogueTile{}) == 2, "EpilogueTile must be rank-2: [EPI_TILE_M,EPI_TILE_N]");
static_assert(rank(StrideC{}) == 3, "StrideCD must be rank-3: [M, N, L]");
static_assert(rank(StrideD{}) == 3, "StrideCD must be rank-3: [M, N, L]");
private:
constexpr static int StagesC = StagesC_;
constexpr static int StagesD = StagesD_;
constexpr static bool is_source_supported = ThreadEpilogueOp::kScale == cutlass::epilogue::thread::ScaleType::Default ||
ThreadEpilogueOp::kScale == cutlass::epilogue::thread::ScaleType::NoBetaScaling;
// internal optimization to reuse C shared memory for storing D
using SmemLayoutAtomBitsC = decltype(downcast<sizeof_bits<ElementC>::value>(SmemLayoutAtomC{}));
using SmemLayoutAtomBitsD = decltype(downcast<sizeof_bits<ElementD>::value>(SmemLayoutAtomD{}));
constexpr static bool ReuseSmemC = not DispatchPolicy::DisableSmemReuseC &&
is_source_supported &&
sizeof(ElementC) == sizeof(ElementD) &&
StrideC{} == StrideD{} &&
cute::is_same_v<SmemLayoutAtomBitsC,SmemLayoutAtomBitsD>;
// Find the max contiguous layout usable by TMA (if EpilogueTile is a by-mode tiler)
using SmemLayoutTmaD = decltype(tile_to_shape(
SmemLayoutAtomD{},
make_shape(max_common_vector(make_layout(get<0>(EpilogueTile{})),make_layout(get<0>(EpilogueTile{}))),
max_common_vector(make_layout(get<1>(EpilogueTile{})),make_layout(get<1>(EpilogueTile{})))),
cute::conditional_t<get<0>(StrideD{}) == 1, Step<_2,_1>, Step<_1,_2>>{} ));
public:
using SmemLayoutC = decltype(tile_to_shape(
SmemLayoutAtomC{},
make_shape(size<0>(BlockTileShape{}), size<1>(BlockTileShape{}), Int<StagesC>{}),
cute::conditional_t<get<0>(StrideC{}) == 1, Step<_2,_1,_3>, Step<_1,_2,_3>>{} ));
using SmemLayoutD = decltype(tile_to_shape(
SmemLayoutTmaD{},
make_shape(size<0>(shape(EpilogueTile{})), size<1>(shape(EpilogueTile{})), Int<StagesD>{}),
cute::conditional_t<get<0>(StrideD{}) == 1, Step<_2,_1,_3>, Step<_1,_2,_3>>{} ));
// TMA pipeline for loading C
using LoadPipeline = cutlass::PipelineTransactionAsync<is_source_supported ? StagesC : 0>;
using LoadPipelineState = cutlass::PipelineState<is_source_supported ? StagesC : 0>;
constexpr static uint32_t TmaTransactionBytes =
size(take<0,2>(SmemLayoutC{})) * static_cast<uint32_t>(sizeof(ElementC));
// TMA pipeline for storing D
using StorePipeline = cutlass::PipelineTmaStore<ReuseSmemC ? StagesC : StagesD>;
using StorePipelineState = cutlass::PipelineState<ReuseSmemC ? StagesC : StagesD>;
struct SharedStorage {
struct TensorStorage : aligned_struct<128> {
cute::conditional_t<not is_source_supported,
detail::EmptyStorage<ElementC>,
array_aligned<ElementC, size(SmemLayoutC{})>> smem_C;
alignas(128) cute::conditional_t<ReuseSmemC,
detail::EmptyStorage<ElementD>,
array_aligned<ElementD, size(SmemLayoutD{})>> smem_D;
} tensors;
using PipelineStorage = typename LoadPipeline::SharedStorage;
PipelineStorage pipeline;
};
using TensorStorage = typename SharedStorage::TensorStorage;
using PipelineStorage = typename SharedStorage::PipelineStorage;
// Host side epilogue arguments
struct Arguments {
typename ThreadEpilogueOp::Params thread;
ElementC const* ptr_C;
StrideC dC;
ElementD const* ptr_D;
StrideD dD;
};
// Device side epilgoue params
struct Params {
using TMA_C = decltype(make_tma_copy(
CopyOpG2S{},
make_tensor(static_cast<ElementC const*>(nullptr),
repeat_like(StrideC{}, int32_t(0)), StrideC{}),
SmemLayoutC{}(_,_,0)));
using TMA_D = decltype(make_tma_copy(
CopyOpS2G{},
make_tensor(static_cast<ElementD const*>(nullptr),
repeat_like(StrideD{}, int32_t(0)), StrideD{}),
SmemLayoutTmaD{}));
typename ThreadEpilogueOp::Params thread{};
TMA_C tma_load_c;
TMA_D tma_store_d;
};
//
// Methods
//
template <class ProblemShape>
static constexpr Params
to_underlying_arguments(
ProblemShape const& problem_shape,
Arguments const& args,
[[maybe_unused]] void* workspace)
{
// Optionally append _1s until problem shape is rank-4 in case its is only rank-3 (MNK)
auto problem_shape_MNKL = append<4>(problem_shape, Int<1>{});
auto M = get<0>(problem_shape_MNKL);
auto N = get<1>(problem_shape_MNKL);
auto L = get<3>(problem_shape_MNKL);
Tensor tensor_c = make_tensor(args.ptr_C, make_layout(make_shape(M,N,L), args.dC));
Tensor tensor_d = make_tensor(args.ptr_D, make_layout(make_shape(M,N,L), args.dD));
typename Params::TMA_C tma_load_c = make_tma_copy(
CopyOpG2S{},
tensor_c,
SmemLayoutC{}(_,_,0));
typename Params::TMA_D tma_store_d = make_tma_copy(
CopyOpS2G{},
tensor_d,
SmemLayoutTmaD{});
return {
args.thread,
tma_load_c,
tma_store_d
};
}
template<class TileShapeMNK>
CUTLASS_HOST_DEVICE
static constexpr int
get_load_pipe_increment(TileShapeMNK tile_shape_MNK) {
// Compute number of C subtiles (currently always one)
constexpr int epi_m = size<0>(tile_shape_MNK) / size<0>(SmemLayoutC{});
constexpr int epi_n = size<1>(tile_shape_MNK) / size<1>(SmemLayoutC{});
return epi_m * epi_n;
}
template<class TileShapeMNK>
CUTLASS_HOST_DEVICE
static constexpr int
get_store_pipe_increment(TileShapeMNK tile_shape_MNK) {
if constexpr (ReuseSmemC) {
return get_load_pipe_increment(tile_shape_MNK);
}
// Compute number of D subtiles
constexpr int epi_m = size<0>(tile_shape_MNK) / size<0>(SmemLayoutD{});
constexpr int epi_n = size<1>(tile_shape_MNK) / size<1>(SmemLayoutD{});
return epi_m * epi_n;
}
CUTLASS_HOST_DEVICE
CollectiveEpilogue(Params const& params_)
: params(params_), epilogue_op(params_.thread) { }
CUTLASS_DEVICE
bool
is_source_needed() {
return epilogue_op.is_source_needed();
}
/// Issue Tma Descriptor Prefetch -- ideally from a single thread for best performance
CUTLASS_DEVICE
static void prefetch_tma_descriptors(Params const& epilogue_params) {
cute::prefetch_tma_descriptor(epilogue_params.tma_load_c.get_tma_descriptor());
cute::prefetch_tma_descriptor(epilogue_params.tma_store_d.get_tma_descriptor());
}
template<
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class TiledMma
>
CUTLASS_DEVICE void
load(
LoadPipeline load_pipeline,
LoadPipelineState load_pipe_producer_state,
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_MNK,
TileCoordMNKL tile_coord_mnkl,
TiledMma tiled_mma,
[[maybe_unused]] int thread_idx,
TensorStorage& shared_tensors) {
using namespace cute;
using X = Underscore;
int warp_idx = canonical_warp_idx();
int warp_idx_in_warp_group = warp_idx % 4;
int lane_predicate = cute::elect_one_sync();
auto M = get<0>(problem_shape_mnkl);
auto N = get<1>(problem_shape_mnkl);
auto L = get<3>(problem_shape_mnkl);
auto [m_coord, n_coord, k_coord, l_coord] = tile_coord_mnkl;
// Represent the full source tensor
Tensor mC_mnl = params.tma_load_c.get_tma_tensor(make_shape(M,N,L)); // (m,n,l)
Tensor gC_mnl = local_tile(mC_mnl, tile_shape_MNK, make_coord(_,_,_), Step<_1,_1, X>{}); // (TILE_M,TILE_N,m,n,l)
// Slice to get the gmem tile of C (gC) this CTA is currently responsible for
Tensor gC = gC_mnl(_,_,m_coord,n_coord,l_coord); // (TILE_M,TILE_N)
// Get the corresponding smem tile of C (sC)
Tensor sC = make_tensor(make_smem_ptr(shared_tensors.smem_C.data()), SmemLayoutC{}); // (TILE_M,TILE_N,PIPE)
// Prepare the thread(b)lock (G)mem to (S)mem TMA copy (bGS_)
ThrCopy thrblk_g2s = params.tma_load_c.get_slice(Int<0>{});
Tensor bGS_gC = thrblk_g2s.partition_S(gC); // (TMA,TMA_M,TMA_N)
Tensor bGS_sC = thrblk_g2s.partition_D(sC); // (TMA,TMA_M,TMA_N,PIPE)
auto* tma_barrier = load_pipeline.producer_get_barrier(load_pipe_producer_state);
uint16_t mcast_mask = 0;
// Execute the TMA load for C
if (warp_idx_in_warp_group == 0 and lane_predicate) {
load_pipeline.producer_acquire(load_pipe_producer_state);
copy(params.tma_load_c.with(*tma_barrier, mcast_mask), bGS_gC, bGS_sC(_,_,_,load_pipe_producer_state.index()));
load_pipeline.producer_commit(load_pipe_producer_state);
}
}
CUTLASS_DEVICE void
load_tail(
LoadPipeline load_pipeline,
LoadPipelineState load_pipe_producer_state) {
int warp_idx = canonical_warp_idx();
int warp_idx_in_warp_group = warp_idx % 4;
int lane_predicate = cute::elect_one_sync();
if (warp_idx_in_warp_group == 0 and lane_predicate) {
load_pipeline.producer_tail(load_pipe_producer_state);
}
}
template<
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class AccEngine, class AccLayout,
class TiledMma
>
CUTLASS_DEVICE void
store(
LoadPipeline load_pipeline,
LoadPipelineState load_pipe_consumer_state,
StorePipeline store_pipeline,
StorePipelineState store_pipe_producer_state,
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_MNK,
TileCoordMNKL tile_coord_mnkl,
cute::Tensor<AccEngine,AccLayout> accumulators,
TiledMma tiled_mma,
int thread_idx,
TensorStorage& shared_tensors) {
using namespace cute;
using X = Underscore;
static_assert(is_rmem<AccEngine>::value, "Accumulator must be RF resident.");
static_assert(rank(AccLayout{}) == 3, "Accumulator must be MMA-partitioned: (MMA,MMA_M,MMA_N)");
static_assert(rank(ProblemShapeMNKL{}) == 4, "ProblemShapeMNKL must be rank 4");
static_assert(is_static<TileShapeMNK>::value, "TileShapeMNK must be static");
static_assert(rank(TileShapeMNK{}) == 3, "TileShapeMNK must be rank 3");
static_assert(rank(TileCoordMNKL{}) == 4, "TileCoordMNKL must be rank 4");
// Separate out problem shape for convenience
auto M = get<0>(problem_shape_mnkl);
auto N = get<1>(problem_shape_mnkl);
auto L = get<3>(problem_shape_mnkl);
auto mma_tile_m = size<0>(typename TiledMma::TiledShape_MNK{});
auto mma_tile_n = size<1>(typename TiledMma::TiledShape_MNK{});
auto epi_tile_m = size<0>(shape(EpilogueTile{}));
auto epi_tile_n = size<1>(shape(EpilogueTile{}));
// Represent the full output tensor
Tensor mD_mnl = params.tma_store_d.get_tma_tensor(make_shape(M,N,L)); // (m,n,l)
Tensor gD_mnl = local_tile(mD_mnl, tile_shape_MNK, make_coord(_,_,_), Step<_1,_1, X>{}); // (TILE_M,TILE_N,m,n,l)
// Slice to get the tile this CTA is responsible for
auto [m_coord, n_coord, k_coord, l_coord] = tile_coord_mnkl;
Tensor gD = gD_mnl(_,_,m_coord,n_coord,l_coord); // (TILE_M,TILE_N)
// Construct the smem tensors for source (sC) and output (sD)
Tensor sC = make_tensor(make_smem_ptr(shared_tensors.smem_C.data()), // (TILE_M,TILE_N)
SmemLayoutC{})(_,_,load_pipe_consumer_state.index());
Tensor bEsD = make_tensor(make_smem_ptr(shared_tensors.smem_D.data()), // (EPI_TILE_M,EPI_TILE_N,PIPE)
SmemLayoutD{});
// Tile thread(b)lock tensors by (E)pilogue output tile shape (bE)
Tensor bEsC = local_tile(sC, EpilogueTile{}, _); // (EPI_TILE_M,EPI_TILE_N,EPI_M,EPI_N)
Tensor bEgD = local_tile(gD, EpilogueTile{}, _); // (EPI_TILE_M,EPI_TILE_N,EPI_M,EPI_N)
// Partition for register to smem copy (tRS_)
TiledCopy tiled_r2s = make_tiled_copy_C_atom(Copy_Atom<CopyOpR2S,ElementD>{}, tiled_mma);
ThrCopy thread_r2s = tiled_r2s.get_slice(thread_idx);
Tensor tRS_rAcc = thread_r2s.retile_S(accumulators); // ((R2S,R2S_V),MMA_M,MMA_N)
Tensor tRS_sD = conditional_return<ReuseSmemC>(
thread_r2s.partition_D(recast<ElementD>(bEsC)), // (R2S,R2S_M,R2S_N,EPI_M,EPI_N)
thread_r2s.partition_D(bEsD) ); // (R2S,R2S_M,R2S_N,PIPE)
// Allocate register tensors
auto tRS_rD_shape = take<0,3>(shape(thread_r2s.partition_S(bEsD))); // (R2S,R2S_M,R2S_N)
Tensor tRS_rC = make_tensor<ElementC>(tRS_rD_shape); // (R2S,R2S_M,R2S_N)
Tensor tRS_rD = make_tensor<ElementD>(tRS_rD_shape); // (R2S,R2S_M,R2S_N)
// Vectorized fragment view for thread epilogue op
Tensor tRS_rAcc_frg = recast<typename ThreadEpilogueOp::FragmentAccumulator>(tRS_rAcc);
Tensor tRS_rC_frg = recast<typename ThreadEpilogueOp::FragmentSource>(tRS_rC);
Tensor tRS_rD_frg = recast<typename ThreadEpilogueOp::FragmentOutput>(tRS_rD);
// Partition for smem to register copy (tSR_)
TiledCopy tiled_s2r = make_tiled_copy_S(Copy_Atom<CopyOpS2R,ElementC>{}, tiled_r2s);
ThrCopy thread_s2r = tiled_s2r.get_slice(thread_idx);
Tensor tSR_sC = thread_s2r.partition_S(bEsC); // (S2R,S2R_M,S2R_N,EPI_M,EPI_N)
Tensor tSR_rC = thread_s2r.retile_D(tRS_rC); // (S2R,S2R_M,S2R_N)
// Partition for smem to gmem copy (tSG_)
ThrCopy thrblk_s2g = params.tma_store_d.get_slice(Int<0>{});
Tensor tSG_sD = conditional_return<ReuseSmemC>(
thrblk_s2g.partition_S(recast<ElementD>(bEsC)), // (S2G,S2G_M,S2G_N,EPI_M,EPI_N)
thrblk_s2g.partition_S(bEsD) ); // (S2G,S2G_M,S2G_N,PIPE)
Tensor tSG_gD = thrblk_s2g.partition_D(bEgD); // (S2G,S2G_M,S2G_N,EPI_M,EPI_N)
CUTE_STATIC_ASSERT(size<0,0>(tRS_rAcc) % ThreadEpilogueOp::kCount == 0, "ThreadEpilogueOp does not vectorize properly");
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");
// Thread synchronizer for previously issued waits or fences
// to ensure visibility of smem reads/writes to threads or TMA unit
auto synchronize = [&] () { cutlass::arch::NamedBarrier::sync(size(TiledMma{}), 0); };
// Predication for TMA store (one warp issues TMA store)
bool issue_tma_store = (thread_idx / NumThreadsPerWarp) == 0;
if (epilogue_op.is_source_needed()) {
// Wait for epilogue load to fill smem buffer with C
load_pipeline.consumer_wait(load_pipe_consumer_state);
}
// Delay issue of TMA store by 1 iteration to achieve better instruction pipelining
PipelineState store_pipe_producer_state_prev = store_pipe_producer_state;
int epi_m_prev = 0, epi_n_prev = 0;
// For each output tile
CUTLASS_PRAGMA_UNROLL
for (int epi_n = 0; epi_n < size<3>(bEgD); ++epi_n) {
CUTLASS_PRAGMA_UNROLL
for (int epi_m = 0; epi_m < size<2>(bEgD); ++epi_m) {
// 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);
// Elementwise operation with conversion
int r2s_v = epi_n * size(tRS_rD_frg);
if (epilogue_op.is_source_needed()) {
// Copy source tile to register from smem
copy(tiled_s2r, tSR_sC(_,_,_,epi_m,epi_n), tSR_rC);
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < size(tRS_rD_frg); ++i) {
tRS_rD_frg(i) = epilogue_op(tRS_rAcc_frg_mn(r2s_v + i), tRS_rC_frg(i));
}
}
else {
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < size(tRS_rD_frg); ++i) {
tRS_rD_frg(i) = epilogue_op(tRS_rAcc_frg_mn(r2s_v + i));
}
}
if constexpr (ReuseSmemC) {
// Issue the TMA store of the previous iteration
if (not (epi_m == 0 && epi_n == 0)) {
// Make sure smem writes are visible to TMA
cutlass::arch::fence_view_async_shared();
synchronize(); // ensure all threads have issued their async fence
// Write the tile to gmem from smem with TMA
if (issue_tma_store) {
copy(params.tma_store_d, tSG_sD(_,_,_,epi_m_prev,epi_n_prev), tSG_gD(_,_,_,epi_m_prev,epi_n_prev));
}
}
// Copy output tile to smem from register
copy(tiled_r2s, tRS_rD, tRS_sD(_,_,_,epi_m,epi_n));
}
else {
// Issue the TMA store of the previous iteration
if (not (epi_m == 0 && epi_n == 0)) {
// Make sure smem writes are visible to TMA
cutlass::arch::fence_view_async_shared();
synchronize(); // ensure all threads have issued their async fence
// Write the tile to gmem from smem with TMA
if (issue_tma_store) {
copy(params.tma_store_d, tSG_sD(_,_,_,store_pipe_producer_state_prev.index()), tSG_gD(_,_,_,epi_m_prev,epi_n_prev));
store_pipeline.producer_commit(store_pipe_producer_state_prev);
}
}
// Wait for a smem buffer to be available
if (issue_tma_store) {
store_pipeline.producer_acquire(store_pipe_producer_state);
}
synchronize();
// Copy tile to smem from register
copy(tiled_r2s, tRS_rD, tRS_sD(_,_,_,store_pipe_producer_state.index()));
// Advance pipeline state
store_pipe_producer_state_prev = store_pipe_producer_state;
++store_pipe_producer_state;
}
epi_m_prev = epi_m;
epi_n_prev = epi_n;
}
}
if constexpr (ReuseSmemC) {
// Fence and issue the TMA store of the last iteration
cutlass::arch::fence_view_async_shared();
synchronize(); // ensure all threads have issued their async fence
if (issue_tma_store) {
copy(params.tma_store_d, tSG_sD(_,_,_,epi_m_prev,epi_n_prev), tSG_gD(_,_,_,epi_m_prev,epi_n_prev));
}
// Arrive and advance pipeline state
if (issue_tma_store) {
store_pipeline.producer_commit(store_pipe_producer_state);
}
++store_pipe_producer_state;
// Wait for a smem buffer to be available
if (issue_tma_store) {
store_pipeline.producer_acquire(store_pipe_producer_state);
}
synchronize();
// Let dma warp know smem buffer is consumed and empty
if (epilogue_op.is_source_needed()) {
load_pipeline.consumer_release(store_pipe_producer_state);
}
}
else {
// Fence and issue the TMA store of the last iteration
cutlass::arch::fence_view_async_shared();
synchronize(); // ensure all threads have issued their async fence
if (issue_tma_store) {
copy(params.tma_store_d, tSG_sD(_,_,_,store_pipe_producer_state_prev.index()), tSG_gD(_,_,_,epi_m_prev,epi_n_prev));
store_pipeline.producer_commit(store_pipe_producer_state_prev);
}
// Let dma warp know smem buffer is consumed and empty
if (epilogue_op.is_source_needed()) {
load_pipeline.consumer_release(load_pipe_consumer_state);
}
}
}
private:
Params const& params;
ThreadEpilogueOp epilogue_op;
};
/////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace collective
} // namespace epilogue
} // namespace cutlass
/////////////////////////////////////////////////////////////////////////////////////////////////
@@ -0,0 +1,558 @@
/***************************************************************************************************
* Copyright (c) 2023, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * Redistributions in binary form must reproduce the above copyright notice, this list of
* conditions and the following disclaimer in the documentation and/or other materials
* provided with the distribution.
* * Neither the name of the NVIDIA CORPORATION nor the names of its contributors may be used
* to endorse or promote products derived from this software without specific prior written
* permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND ANY EXPRESS OR
* IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND
* FITNESS FOR A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL NVIDIA CORPORATION BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING,
* BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS;
* OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT,
* STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
/*! \file
\brief Functor performing pipelined epilogues with bias add and elementwise activation functions.
*/
#pragma once
#include "cutlass/cutlass.h"
#include "cutlass/arch/barrier.h"
#include "cutlass/epilogue/dispatch_policy.hpp"
#include "cutlass/epilogue/collective/detail.hpp"
#include "cutlass/epilogue/thread/scale_type.h"
#include "cute/tensor.hpp"
/////////////////////////////////////////////////////////////////////////////////////////////////
namespace cutlass {
namespace epilogue {
namespace collective {
/////////////////////////////////////////////////////////////////////////////////////////////////
template <
int StagesC_,
int StagesD_,
class BlockTileShape_, // (BLK_M,BLK_N,BLK_K)
class EpilogueTile_, // (EPI_TILE_M,EPI_TILE_N) per-collective
class ElementC_,
class StrideC_,
class ElementD_,
class StrideD_,
class ThreadEpilogueOp_,
class CopyOpG2S_,
class SmemLayoutAtomC_,
class CopyOpS2R_,
class CopyOpS2G_,
class SmemLayoutAtomD_,
class CopyOpR2S_
>
class CollectiveEpilogue<
Sm90TmaWarpSpecializedBiasElementwise<StagesC_, StagesD_>,
BlockTileShape_,
EpilogueTile_,
ElementC_,
StrideC_,
ElementD_,
StrideD_,
ThreadEpilogueOp_,
CopyOpG2S_,
SmemLayoutAtomC_,
CopyOpS2R_,
CopyOpS2G_,
SmemLayoutAtomD_,
CopyOpR2S_
> {
public:
//
// Type Aliases
//
// derived types of output thread level operator
using DispatchPolicy = Sm90TmaWarpSpecializedBiasElementwise<StagesC_, StagesD_>;
using BlockTileShape = BlockTileShape_;
using EpilogueTile = EpilogueTile_;
using ThreadEpilogueOp = ThreadEpilogueOp_;
using ElementAccumulator = typename ThreadEpilogueOp::ElementAccumulator;
using ElementCompute = typename ThreadEpilogueOp::ElementCompute;
using ElementScalar = ElementCompute;
using ElementBias = typename detail::IsThreadEpilogueOpWithBias<ThreadEpilogueOp>::type;
using ElementT = typename ThreadEpilogueOp::ElementT;
using ElementOutput = typename ThreadEpilogueOp::ElementOutput;
using ElementC = ElementC_;
using StrideC = StrideC_;
using ElementD = ElementD_;
using StrideD = StrideD_;
using ActivationFunctor = typename ThreadEpilogueOp::ActivationFunctor;
using BinaryOp = typename ThreadEpilogueOp::BinaryOp;
using CopyOpG2S = CopyOpG2S_;
using SmemLayoutAtomC = SmemLayoutAtomC_;
using CopyOpS2R = CopyOpS2R_;
using CopyOpS2G = CopyOpS2G_;
using SmemLayoutAtomD = SmemLayoutAtomD_;
using CopyOpR2S = CopyOpR2S_;
constexpr static bool StoreT = ThreadEpilogueOp::kStoreT;
constexpr static int kOutputAlignment = ThreadEpilogueOp::kCount;
static_assert(detail::IsThreadEpilogueOpWithBias<ThreadEpilogueOp>::value,
"Epilogue dispatch policy Sm90TmaWarpSpecializedBiasElementwise requires the use of a thread-level epiogue that supports bias calculation");
constexpr static bool iskThreadEpilogueOpWithBias = true;
using AlignmentType = typename uint_bit<sizeof_bits<ElementOutput>::value * kOutputAlignment>::type;
static_assert(sizeof(ElementC) == 2, "Only 16b source supported for now");
static_assert(sizeof(ElementD) == 2, "Only 16b output supported for now");
static_assert(!is_layout<EpilogueTile>::value && is_tuple<EpilogueTile>::value, "EpilogueTile must be a cute::Tile or cute::Shape");
static_assert(rank(BlockTileShape{}) == 3, "BlockTileShape must be rank-3: [BLK_M,BLK_N,BLK_K]");
static_assert(rank(EpilogueTile{}) == 2, "EpilogueTile must be rank-2: [EPI_TILE_M,EPI_TILE_N]");
static_assert(rank(StrideC{}) == 3, "StrideCD must be rank-3: [M, N, L]");
static_assert(rank(StrideD{}) == 3, "StrideCD must be rank-3: [M, N, L]");
private:
constexpr static int StagesC = StagesC_;
constexpr static int StagesD = StagesD_;
constexpr static bool is_source_supported = ThreadEpilogueOp::kScale == cutlass::epilogue::thread::ScaleType::Default ||
ThreadEpilogueOp::kScale == cutlass::epilogue::thread::ScaleType::NoBetaScaling;
// Find the max contiguous layout usable by TMA (if EpilogueTile is a by-mode tiler)
using SmemLayoutTmaD = decltype(tile_to_shape(
SmemLayoutAtomD{},
make_shape(max_common_vector(make_layout(get<0>(EpilogueTile{})),make_layout(get<0>(EpilogueTile{}))),
max_common_vector(make_layout(get<1>(EpilogueTile{})),make_layout(get<1>(EpilogueTile{})))),
cute::conditional_t<get<0>(StrideD{}) == 1, Step<_2,_1>, Step<_1,_2>>{} ));
public:
using SmemLayoutC = decltype(tile_to_shape(
SmemLayoutAtomC{},
make_shape(size<0>(BlockTileShape{}), size<1>(BlockTileShape{}), Int<StagesC>{}),
cute::conditional_t<get<0>(StrideC{}) == 1, Step<_2,_1,_3>, Step<_1,_2,_3>>{} ));
using SmemLayoutD = decltype(tile_to_shape(
SmemLayoutTmaD{},
make_shape(size<0>(shape(EpilogueTile{})), size<1>(shape(EpilogueTile{})), Int<StagesD>{}),
cute::conditional_t<get<0>(StrideD{}) == 1, Step<_2,_1,_3>, Step<_1,_2,_3>>{} ));
// TMA pipeline for loading C
using LoadPipeline = cutlass::PipelineTransactionAsync<is_source_supported ? StagesC : 0>;
using LoadPipelineState = cutlass::PipelineState<is_source_supported ? StagesC : 0>;
constexpr static uint32_t TmaTransactionBytes =
size(take<0,2>(SmemLayoutC{})) * static_cast<uint32_t>(sizeof(ElementC));
// TMA pipeline for storing D and T
using StorePipeline = cutlass::PipelineTmaStore<StagesD>;
using StorePipelineState = cutlass::PipelineState<StagesD>;
struct SharedStorage {
struct TensorStorage : aligned_struct<128> {
cute::conditional_t<not is_source_supported,
detail::EmptyStorage<ElementC>,
array_aligned<ElementC, size(SmemLayoutC{})>> smem_C;
alignas(128) array_aligned<ElementD, size(SmemLayoutD{})> smem_D;
alignas(128) cute::conditional_t<not StoreT,
detail::EmptyStorage<ElementT>,
array_aligned<ElementT, size(SmemLayoutD{})>> smem_T;
} tensors;
using PipelineStorage = typename LoadPipeline::SharedStorage;
PipelineStorage pipeline;
};
using TensorStorage = typename SharedStorage::TensorStorage;
using PipelineStorage = typename SharedStorage::PipelineStorage;
// Host side epilogue arguments
struct Arguments {
typename ThreadEpilogueOp::Params thread{};
ElementC const* ptr_C = nullptr;
StrideC dC{};
ElementD* ptr_D = nullptr;
StrideD dD{};
ElementBias const* ptr_Bias = nullptr;
ElementT* ptr_T = nullptr;
};
// Device side epilogue params
struct Params {
using TMA_C = decltype(make_tma_copy(
CopyOpG2S{},
make_tensor(static_cast<ElementC const*>(nullptr), repeat_like(StrideC{}, int32_t(0)), StrideC{}),
SmemLayoutC{}(_,_,0)));
using TMA_D = decltype(make_tma_copy(
CopyOpS2G{},
make_tensor(static_cast<ElementD*>(nullptr), repeat_like(StrideD{}, int32_t(0)), StrideD_{}),
SmemLayoutTmaD{}));
using TMA_T = decltype(make_tma_copy(
CopyOpS2G{},
make_tensor(static_cast<ElementT*>(nullptr), repeat_like(StrideD{}, int32_t(0)), StrideD{}),
SmemLayoutTmaD{}));
typename ThreadEpilogueOp::Params thread{};
TMA_C tma_load_c;
TMA_D tma_store_d;
TMA_T tma_store_t;
ElementBias const* ptr_Bias = nullptr;
};
//
// Methods
//
template <class ProblemShape>
static constexpr Params
to_underlying_arguments(ProblemShape const& problem_shape, Arguments const& args, [[maybe_unused]] void* workspace) {
// Optionally append _1s until problem shape is rank-4 in case its is only rank-3 (MNK)
auto problem_shape_MNKL = append<4>(problem_shape, Int<1>{});
auto M = get<0>(problem_shape_MNKL);
auto N = get<1>(problem_shape_MNKL);
auto L = get<3>(problem_shape_MNKL);
Tensor tensor_c = make_tensor(args.ptr_C, make_layout(make_shape(M,N,L), args.dC));
Tensor tensor_d = make_tensor(args.ptr_D, make_layout(make_shape(M,N,L), args.dD));
typename Params::TMA_C tma_load_c = make_tma_copy(
CopyOpG2S{},
tensor_c,
SmemLayoutC{}(_,_,0));
typename Params::TMA_D tma_store_d = make_tma_copy(
CopyOpS2G{},
tensor_d,
SmemLayoutTmaD{});
typename Params::TMA_T tma_store_t = [&]() {
if constexpr (StoreT) {
Tensor tensor_t = make_tensor(args.ptr_T, make_layout(make_shape(M,N,L), args.dD));
return make_tma_copy(
CopyOpS2G{},
tensor_t,
SmemLayoutTmaD{});
}
else {
return typename Params::TMA_T{};
}
}();
return {
args.thread,
tma_load_c,
tma_store_d,
tma_store_t,
args.ptr_Bias
};
}
template<class TileShapeMNK>
CUTLASS_HOST_DEVICE
static constexpr int
get_load_pipe_increment(TileShapeMNK tile_shape_MNK) {
// Compute number of C subtiles (currently always one)
constexpr int epi_m = size<0>(tile_shape_MNK) / size<0>(SmemLayoutC{});
constexpr int epi_n = size<1>(tile_shape_MNK) / size<1>(SmemLayoutC{});
return epi_m * epi_n;
}
template<class TileShapeMNK>
CUTLASS_HOST_DEVICE
static constexpr int
get_store_pipe_increment(TileShapeMNK tile_shape_MNK) {
// Compute number of D subtiles
constexpr int epi_m = size<0>(tile_shape_MNK) / size<0>(SmemLayoutD{});
constexpr int epi_n = size<1>(tile_shape_MNK) / size<1>(SmemLayoutD{});
return epi_m * epi_n;
}
CUTLASS_HOST_DEVICE
CollectiveEpilogue(Params const& params_)
: params(params_), epilogue_op(params_.thread) { }
CUTLASS_DEVICE
bool
is_source_needed() {
return epilogue_op.is_source_needed();
}
/// Issue Tma Descriptor Prefetch -- ideally from a single thread for best performance
CUTLASS_DEVICE
static void prefetch_tma_descriptors(Params const& epilogue_params) {
cute::prefetch_tma_descriptor(epilogue_params.tma_load_c.get_tma_descriptor());
cute::prefetch_tma_descriptor(epilogue_params.tma_store_d.get_tma_descriptor());
if constexpr (StoreT) {
cute::prefetch_tma_descriptor(epilogue_params.tma_store_t.get_tma_descriptor());
}
}
template<
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class TiledMma
>
CUTLASS_DEVICE void
load(
LoadPipeline load_pipeline,
LoadPipelineState load_pipe_producer_state,
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_MNK,
TileCoordMNKL tile_coord_mnkl,
TiledMma tiled_mma,
[[maybe_unused]] int thread_idx,
TensorStorage& shared_tensors) {
using namespace cute;
using X = Underscore;
int warp_idx = canonical_warp_idx();
int warp_idx_in_warp_group = warp_idx % 4;
int lane_predicate = cute::elect_one_sync();
auto M = get<0>(problem_shape_mnkl);
auto N = get<1>(problem_shape_mnkl);
auto L = get<3>(problem_shape_mnkl);
auto [m_coord, n_coord, k_coord, l_coord] = tile_coord_mnkl;
// Represent the full source tensor
Tensor mC_mnl = params.tma_load_c.get_tma_tensor(make_shape(M,N,L)); // (m,n,l)
Tensor gC_mnl = local_tile(mC_mnl, tile_shape_MNK, make_coord(_,_,_), Step<_1,_1, X>{}); // (TILE_M,TILE_N,m,n,l)
// Slice to get the gmem tile of C (gC) this CTA is currently responsible for
Tensor gC = gC_mnl(_,_,m_coord,n_coord,l_coord); // (TILE_M,TILE_N)
// Get the corresponding smem tile of C (sC)
Tensor sC = make_tensor(make_smem_ptr(shared_tensors.smem_C.data()), SmemLayoutC{}); // (TILE_M,TILE_N,PIPE)
// Prepare the thread(b)lock (G)mem to (S)mem TMA copy (bGS_)
ThrCopy thrblk_g2s = params.tma_load_c.get_slice(Int<0>{});
Tensor bGS_gC = thrblk_g2s.partition_S(gC); // (TMA,TMA_M,TMA_N)
Tensor bGS_sC = thrblk_g2s.partition_D(sC); // (TMA,TMA_M,TMA_N,PIPE)
auto* tma_barrier = load_pipeline.producer_get_barrier(load_pipe_producer_state);
uint16_t mcast_mask = 0;
// Execute the TMA load for C
if (warp_idx_in_warp_group == 0 and lane_predicate) {
load_pipeline.producer_acquire(load_pipe_producer_state);
copy(params.tma_load_c.with(*tma_barrier, mcast_mask), bGS_gC, bGS_sC(_,_,_,load_pipe_producer_state.index()));
load_pipeline.producer_commit(load_pipe_producer_state);
}
}
CUTLASS_DEVICE void
load_tail(
LoadPipeline load_pipeline,
LoadPipelineState load_pipe_producer_state) {
int warp_idx = canonical_warp_idx();
int warp_idx_in_warp_group = warp_idx % 4;
int lane_predicate = cute::elect_one_sync();
if (warp_idx_in_warp_group == 0 and lane_predicate) {
load_pipeline.producer_tail(load_pipe_producer_state);
}
}
template<
class ProblemShapeMNKL,
class TileShapeMNK,
class TileCoordMNKL,
class AccEngine, class AccLayout,
class TiledMma
>
CUTLASS_DEVICE void
store(
LoadPipeline load_pipeline,
LoadPipelineState load_pipe_consumer_state,
StorePipeline store_pipeline,
StorePipelineState store_pipe_producer_state,
ProblemShapeMNKL problem_shape_mnkl,
TileShapeMNK tile_shape_MNK,
TileCoordMNKL tile_coord_mnkl,
cute::Tensor<AccEngine,AccLayout> accumulators,
TiledMma tiled_mma,
int thread_idx,
TensorStorage& shared_tensors) {
using namespace cute;
using X = Underscore;
static_assert(is_rmem<AccEngine>::value, "Accumulator must be RF resident.");
static_assert(rank(AccLayout{}) == 3, "Accumulator must be MMA-partitioned: (MMA,MMA_M,MMA_N)");
static_assert(rank(ProblemShapeMNKL{}) == 4, "ProblemShapeMNKL must be rank 4");
static_assert(is_static<TileShapeMNK>::value, "TileShapeMNK must be static");
static_assert(rank(TileShapeMNK{}) == 3, "TileShapeMNK must be rank 3");
static_assert(rank(TileCoordMNKL{}) == 4, "TileCoordMNKL must be rank 4");
// Separate out problem shape for convenience
auto M = get<0>(problem_shape_mnkl);
auto N = get<1>(problem_shape_mnkl);
auto L = get<3>(problem_shape_mnkl);
auto mma_tile_m = size<0>(typename TiledMma::TiledShape_MNK{});
auto mma_tile_n = size<1>(typename TiledMma::TiledShape_MNK{});
auto epi_tile_m = size<0>(shape(EpilogueTile{}));
auto epi_tile_n = size<1>(shape(EpilogueTile{}));
// Represent the full output tensor
Tensor mD_mnl = params.tma_store_d.get_tma_tensor(make_shape(M,N,L)); // (m,n,l)
Tensor gD_mnl = local_tile(mD_mnl, tile_shape_MNK, make_coord(_,_,_), Step<_1, _1, X>{}); // (TILE_M,TILE_N,m,n,l)
Tensor mT_mnl = params.tma_store_t.get_tma_tensor(make_shape(M,N,L)); // (m,n,l)
Tensor gT_mnl = local_tile(mT_mnl, tile_shape_MNK, make_coord(_,_,_), Step<_1, _1, X>{}); // (TILE_M,TILE_N,m,n,l)
Tensor mBias_mnl = make_tensor(make_gmem_ptr(params.ptr_Bias), make_shape(M,N,L), Stride<_1, _0, _0>{}); // (m,n,l)
Tensor gBias_mnl = local_tile(mBias_mnl, tile_shape_MNK, make_coord(_,_,_), Step<_1,_1,X>{}); // (TILE_M,TILE_N,m,n,l)
// Slice to get the tile this CTA is responsible for
auto [m_coord, n_coord, k_coord, l_coord] = tile_coord_mnkl;
Tensor gD = gD_mnl(_,_,m_coord,n_coord,l_coord); // (TILE_M,TILE_N)
Tensor gT = gT_mnl(_,_,m_coord,n_coord,l_coord); // (TILE_M,TILE_N)
Tensor gBias = gBias_mnl(_,_,m_coord,n_coord,l_coord); // (TILE_M,TILE_N)
// Construct the smem tensors for source (sC) and output (sD)
Tensor sC = make_tensor(make_smem_ptr(shared_tensors.smem_C.data()), // (TILE_M,TILE_N)
SmemLayoutC{})(_,_,load_pipe_consumer_state.index());
Tensor bEsD = make_tensor(make_smem_ptr(shared_tensors.smem_D.data()), // (EPI_TILE_M,EPI_TILE_N,PIPE)
SmemLayoutD{});
Tensor bEsT = make_tensor(make_smem_ptr(shared_tensors.smem_T.data()), // (EPI_TILE_M,EPI_TILE_N,PIPE)
SmemLayoutD{});
// Tile thread(b)lock tensors by (E)pilogue output tile shape (bE)
Tensor bEsC = local_tile(sC, EpilogueTile{}, _); // (EPI_TILE_M,EPI_TILE_N,EPI_M,EPI_N)
Tensor bEgD = local_tile(gD, EpilogueTile{}, _); // (EPI_TILE_M,EPI_TILE_N,EPI_M,EPI_N)
Tensor bEgT = local_tile(gT, EpilogueTile{}, _); // (EPI_TILE_M,EPI_TILE_N,EPI_M,EPI_N)
Tensor bEgBias = local_tile(gBias, EpilogueTile{}, _); // (EPI_TILE_M,EPI_TILE_N,EPI_M,EPI_N)
// Partition for register to smem copy (tRS_)
TiledCopy tiled_r2s = make_tiled_copy_C_atom(Copy_Atom<CopyOpR2S,ElementD>{}, tiled_mma);
ThrCopy thread_r2s = tiled_r2s.get_slice(thread_idx);
Tensor tRS_rAcc = thread_r2s.retile_S(accumulators); // ((R2S,R2S_V),MMA_M,MMA_N)
Tensor tRS_sD = thread_r2s.partition_D(bEsD); // (R2S,R2S_M,R2S_N,PIPE)
Tensor tRS_sT = thread_r2s.partition_D(bEsT); // (R2S,R2S_M,R2S_N,PIPE)
// Allocate register tensors
auto tRS_rD_shape = take<0,3>(shape(thread_r2s.partition_S(bEsD))); // (R2S,R2S_M,R2S_N)
Tensor tRS_rC = make_tensor<ElementC>(tRS_rD_shape); // (R2S,R2S_M,R2S_N)
Tensor tRS_rD = make_tensor<ElementD>(tRS_rD_shape); // (R2S,R2S_M,R2S_N)
Tensor tRS_rT = make_tensor<ElementT>(tRS_rD_shape); // (R2S,R2S_M,R2S_N)
Tensor tRS_gBias = thread_r2s.partition_S(bEgBias); // (R2S,R2S_M,R2S_N,EPI_M,EPI_N)
Tensor tRS_rBias = make_tensor<ElementBias>(take<0,3>(shape(tRS_gBias))); // (R2S,R2S_M,R2S_N)
// Vectorized fragment view for thread epilogue op
Tensor tRS_rAcc_frg = recast<typename ThreadEpilogueOp::FragmentAccumulator>(tRS_rAcc);
Tensor tRS_rC_frg = recast<typename ThreadEpilogueOp::FragmentSource>(tRS_rC);
Tensor tRS_rD_frg = recast<typename ThreadEpilogueOp::FragmentOutput>(tRS_rD);
Tensor tRS_rT_frg = recast<typename ThreadEpilogueOp::FragmentT>(tRS_rT);
Tensor tRS_rBias_frg = recast<typename ThreadEpilogueOp::FragmentBias>(tRS_rBias);
// Partition for smem to register copy (tSR_)
TiledCopy tiled_s2r = make_tiled_copy_S(Copy_Atom<CopyOpS2R,ElementC>{}, tiled_r2s);
ThrCopy thread_s2r = tiled_s2r.get_slice(thread_idx);
Tensor tSR_sC = thread_s2r.partition_S(bEsC); // (S2R,S2R_M,S2R_N,EPI_M,EPI_N)
Tensor tSR_rC = thread_s2r.retile_D(tRS_rC); // (S2R,S2R_M,S2R_N)
// Partition for smem to gmem copy (tSG_)
ThrCopy thrblk_s2g = params.tma_store_d.get_slice(Int<0>{});
Tensor tSG_sD = thrblk_s2g.partition_S(bEsD); // (S2G,S2G_M,S2G_N,PIPE)
Tensor tSG_gD = thrblk_s2g.partition_D(bEgD); // (S2G,S2G_M,S2G_N,EPI_M,EPI_N)
ThrCopy thrblk_s2g_t = params.tma_store_t.get_slice(Int<0>{});
Tensor tSG_sT = thrblk_s2g_t.partition_S(bEsT); // (S2G,S2G_M,S2G_N,PIPE)
Tensor tSG_gT = thrblk_s2g_t.partition_D(bEgT); // (S2G,S2G_M,S2G_N,EPI_M,EPI_N)
CUTE_STATIC_ASSERT(size<0,0>(tRS_rAcc) % ThreadEpilogueOp::kCount == 0, "ThreadEpilogueOp does not vectorize properly");
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");
// Thread synchronizer for previously issued waits or fences
// to ensure visibility of smem reads/writes to threads or TMA unit
auto synchronize = [&] () { cutlass::arch::NamedBarrier::sync(size(TiledMma{}), 0); };
// Predication for TMA store (one warp issues TMA store)
bool issue_tma_store = (thread_idx / NumThreadsPerWarp) == 0;
if (epilogue_op.is_source_needed()) {
// Wait for epilogue load to fill smem buffer with C
load_pipeline.consumer_wait(load_pipe_consumer_state);
}
// For each output tile
CUTLASS_PRAGMA_UNROLL
for (int epi_n = 0; epi_n < size<3>(bEgD); ++epi_n) {
CUTLASS_PRAGMA_UNROLL
for (int epi_m = 0; epi_m < size<2>(bEgD); ++epi_m) {
// 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);
// Copy bias to registers from gmem
copy(tRS_gBias(_,_,_,epi_m,epi_n), tRS_rBias);
// Elementwise operation with conversion
int r2s_v = epi_n * size(tRS_rD_frg);
if (epilogue_op.is_source_needed()) {
// Copy source tile to registers from smem
copy(tiled_s2r, tSR_sC(_,_,_,epi_m,epi_n), tSR_rC);
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < size(tRS_rD_frg); ++i) {
epilogue_op(tRS_rD_frg(i), tRS_rT_frg(i), tRS_rAcc_frg_mn(r2s_v + i), tRS_rC_frg(i), tRS_rBias_frg(i));
}
}
else {
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < size(tRS_rD_frg); ++i) {
epilogue_op(tRS_rD_frg(i), tRS_rT_frg(i), tRS_rAcc_frg_mn(r2s_v + i), tRS_rBias_frg(i));
}
}
// Wait for a smem buffer to be available
if (issue_tma_store) {
store_pipeline.producer_acquire(store_pipe_producer_state);
}
synchronize();
// Copy tile to smem from register
copy(tiled_r2s, tRS_rD, tRS_sD(_,_,_,store_pipe_producer_state.index()));
if constexpr (StoreT) {
copy(tiled_r2s, tRS_rT, tRS_sT(_,_,_,store_pipe_producer_state.index()));
}
// Make sure smem writes are visible to TMA
cutlass::arch::fence_view_async_shared();
synchronize(); // ensure all threads have issued their async fence
// Write the tile to gmem from smem with TMA
if (issue_tma_store) {
copy(params.tma_store_d, tSG_sD(_,_,_,store_pipe_producer_state.index()), tSG_gD(_,_,_,epi_m,epi_n));
if constexpr (StoreT) {
copy(params.tma_store_t, tSG_sT(_,_,_,store_pipe_producer_state.index()), tSG_gT(_,_,_,epi_m,epi_n));
}
store_pipeline.producer_commit(store_pipe_producer_state);
}
// Advance pipeline state
++store_pipe_producer_state;
}
}
// Let dma warp know smem buffer is consumed and empty
if (epilogue_op.is_source_needed()) {
load_pipeline.consumer_release(load_pipe_consumer_state);
}
}
private:
Params const& params;
ThreadEpilogueOp epilogue_op;
};
/////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace collective
} // namespace epilogue
} // namespace cutlass
/////////////////////////////////////////////////////////////////////////////////////////////////
@@ -24,16 +24,113 @@
**************************************************************************************************/
#pragma once
#include "cutlass/numeric_conversion.h"
#include "cutlass/epilogue/thread/scale_type.h"
//////////////////////////////////////////////////////////////////////////////
namespace cutlass::epilogue {
//////////////////////////////////////////////////////////////////////////////
// Epilogue schedule types that can be used for categorical dispatch
struct NoSmemWarpSpecialized {};
struct TmaWarpSpecialized {};
struct TmaWarpSpecializedCooperative {};
struct TmaWarpSpecializedElementwiseBase : public TmaWarpSpecialized {};
struct TmaWarpSpecializedCooperativeElementwiseBase : public TmaWarpSpecializedCooperative {};
template <
template <class T> class ActivationFunctor_,
thread::ScaleType::Kind Scale_ = thread::ScaleType::Default,
FloatRoundStyle Round_ = FloatRoundStyle::round_to_nearest
>
struct TmaWarpSpecializedElementwise : public TmaWarpSpecializedElementwiseBase {
template <class T>
using ActivationFunctor = ActivationFunctor_<T>;
static constexpr thread::ScaleType::Kind Scale = Scale_;
static constexpr FloatRoundStyle Round = Round_;
};
template <
template <class T> class ActivationFunctor_,
thread::ScaleType::Kind Scale_ = thread::ScaleType::Default,
FloatRoundStyle Round_ = FloatRoundStyle::round_to_nearest
>
struct TmaWarpSpecializedCooperativeElementwise : public TmaWarpSpecializedCooperativeElementwiseBase {
template <class T>
using ActivationFunctor = ActivationFunctor_<T>;
static constexpr thread::ScaleType::Kind Scale = Scale_;
static constexpr FloatRoundStyle Round = Round_;
};
struct TmaWarpSpecializedBiasElementwiseBase : public TmaWarpSpecialized{};
struct TmaWarpSpecializedCooperativeBiasElementwiseBase : public TmaWarpSpecializedCooperative {};
template <
template <class T> class ActivationFunctor_,
class ElementT_,
template <class T> class BiasOp_,
bool StoreT_,
class ElementBias_
>
struct TmaWarpSpecializedBiasElementwise : public TmaWarpSpecializedBiasElementwiseBase {
template <class T>
using ActivationFunctor = ActivationFunctor_<T>;
using ElementT = ElementT_;
template <class T>
using BiasOp = BiasOp_<T>;
static constexpr bool StoreT = StoreT_;
using ElementBias = ElementBias_;
};
template <
template <class T> class ActivationFunctor_,
class ElementT_,
template <class T> class BiasOp_,
bool StoreT_,
class ElementBias_
>
struct TmaWarpSpecializedCooperativeBiasElementwise : public TmaWarpSpecializedCooperativeBiasElementwiseBase {
template <class T>
using ActivationFunctor = ActivationFunctor_<T>;
using ElementT = ElementT_;
template <class T>
using BiasOp = BiasOp_<T>;
static constexpr bool StoreT = StoreT_;
using ElementBias = ElementBias_;
};
//
// Collective Epilogue Policies
//
template<
int StagesC_,
int StagesD_,
bool DisableSmemReuseC_
>
struct Sm90TmaWarpSpecialized {
constexpr static int StagesC = StagesC_;
constexpr static int StagesD = StagesD_;
constexpr static bool DisableSmemReuseC = DisableSmemReuseC_;
};
template<
int StagesC_,
int StagesD_
>
struct Sm90TmaWarpSpecializedBiasElementwise {
constexpr static int StagesC = StagesC_;
constexpr static int StagesD = StagesD_;
};
//////////////////////////////////////////////////////////////////////////////
} // namespace cutlass::epilogue
@@ -0,0 +1,52 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
/*! \file
\brief Utilities for thread-level epilogues
*/
#pragma once
/////////////////////////////////////////////////////////////////////////////////////////////////
namespace cutlass {
namespace epilogue {
namespace thread {
namespace detail {
/// Class used to identify cases in which no operation is performed
template <typename T_>
struct NoOp {};
} // namespace detail
} // namespace thread
} // namespace epilogue
} // namespace cutlass
@@ -52,7 +52,7 @@ namespace thread {
/// Applies a linear combination operator to an array of elements.
///
/// D = alpha * accumulator + beta * source + uniform
/// D = alpha * accumulator + beta * source
///
template <
typename ElementOutput_, ///< Data type used to load and store tensors
@@ -69,6 +69,7 @@ class LinearCombination {
public:
using ElementOutput = ElementOutput_;
using ElementSource = ElementSource_;
using ElementAccumulator = ElementAccumulator_;
using ElementCompute = ElementCompute_;
using ElementC = ElementSource_;
@@ -77,14 +78,15 @@ public:
static int const kCount = Count;
static const ScaleType::Kind kScale = Scale;
using FragmentOutput = Array<ElementOutput, kCount>;
using FragmentSource = Array<ElementSource, kCount>;
using FragmentAccumulator = Array<ElementAccumulator, kCount>;
using ComputeFragment = Array<ElementCompute, kCount>;
using FragmentCompute = Array<ElementCompute, kCount>;
using ParamsBase = LinearCombinationParams;
static FloatRoundStyle const kRound = Round;
/// Host-constructable parameters structure
struct Params : ParamsBase{
struct Params
{
ElementCompute alpha; ///< scales accumulators
ElementCompute beta; ///< scales source tensor
ElementCompute const *alpha_ptr; ///< pointer to accumulator scalar - if not null, loads it from memory
@@ -92,10 +94,6 @@ public:
CUTLASS_HOST_DEVICE
Params():
ParamsBase(
ElementCompute(1),
ElementCompute(0)
),
alpha(ElementCompute(1)),
beta(ElementCompute(0)),
alpha_ptr(nullptr),
@@ -106,14 +104,12 @@ public:
ElementCompute alpha,
ElementCompute beta
):
ParamsBase(alpha, beta),
alpha(alpha), beta(beta), alpha_ptr(nullptr), beta_ptr(nullptr) { }
CUTLASS_HOST_DEVICE
Params(
ElementCompute alpha
):
ParamsBase(alpha, ElementCompute(0)),
alpha(alpha), beta(0), alpha_ptr(nullptr), beta_ptr(nullptr) { }
CUTLASS_HOST_DEVICE
@@ -121,28 +117,13 @@ public:
ElementCompute const *alpha_ptr,
ElementCompute const *beta_ptr
):
ParamsBase(*alpha_ptr, *beta_ptr),
alpha(0), beta(0), alpha_ptr(alpha_ptr), beta_ptr(beta_ptr) { }
CUTLASS_HOST_DEVICE
Params(
ElementCompute const *alpha_ptr
):
ParamsBase(*alpha_ptr, ElementCompute(0)),
alpha(0), beta(0), alpha_ptr(alpha_ptr), beta_ptr(nullptr) { }
CUTLASS_HOST_DEVICE
Params(
ParamsBase const& base
): ParamsBase(base), alpha_ptr(nullptr), beta_ptr(nullptr) {
#if defined(__CUDA_ARCH__)
alpha = reinterpret_cast<ElementCompute const&>(base.alpha_data);
beta = reinterpret_cast<ElementCompute const&>(base.beta_data);
#else
memcpy( alpha, base.alpha_data, sizeof(ElementCompute) );
memcpy( beta, base.alpha_data, sizeof(ElementCompute) );
#endif
}
};
private:
@@ -183,30 +164,73 @@ public:
}
}
/// Computes linear scaling: D = alpha * accumulator + beta * source
/// Computes intermediate: X = beta * source
CUTLASS_HOST_DEVICE
FragmentCompute compute_intermediate(
FragmentSource const &source) const {
// Convert source to internal compute numeric type
NumericArrayConverter<ElementCompute, ElementSource, kCount, Round> source_converter;
FragmentCompute converted_source = source_converter(source);
if (Scale == ScaleType::NoBetaScaling) {
return converted_source;
}
else {
multiplies<FragmentCompute> mul_source;
return mul_source(beta_, converted_source);
}
}
/// Computes linear scaling with intermediate: D = alpha * accumulator + X
CUTLASS_HOST_DEVICE
FragmentOutput with_intermediate(
FragmentAccumulator const& accumulator,
FragmentCompute const& intermediate) const {
// Convert accumulator to internal compute numeric type
NumericArrayConverter<ElementCompute, ElementAccumulator, kCount, Round> accumulator_converter;
// Convert to destination numeric type
NumericArrayConverter<ElementOutput, ElementCompute, kCount, Round> destination_converter;
FragmentCompute converted_accumulator = accumulator_converter(accumulator);
if (Scale == ScaleType::Nothing) {
return destination_converter(converted_accumulator);
} else {
// Perform binary operations
multiply_add<FragmentCompute> mul_add_accumulator;
FragmentCompute computed_output = mul_add_accumulator(alpha_, converted_accumulator, intermediate);
return destination_converter(computed_output);
}
}
/// Computes linear scaling with source: D = alpha * accumulator + beta * source
CUTLASS_HOST_DEVICE
FragmentOutput operator()(
FragmentAccumulator const &accumulator,
FragmentOutput const &source) const {
FragmentAccumulator const &accumulator,
FragmentSource const &source) const {
// Convert source to interal compute numeric type
NumericArrayConverter<ElementCompute, ElementOutput, kCount, Round> source_converter;
// Convert source to internal compute numeric type
NumericArrayConverter<ElementCompute, ElementSource, kCount, Round> source_converter;
NumericArrayConverter<ElementCompute, ElementAccumulator, kCount, Round> accumulator_converter;
// Convert to destination numeric type
NumericArrayConverter<ElementOutput, ElementCompute, kCount, Round> destination_converter;
ComputeFragment converted_source = source_converter(source);
ComputeFragment converted_accumulator = accumulator_converter(accumulator);
FragmentCompute converted_source = source_converter(source);
FragmentCompute converted_accumulator = accumulator_converter(accumulator);
if (Scale == ScaleType::Nothing)
return destination_converter(converted_accumulator);
// Perform binary operations
ComputeFragment intermediate;
FragmentCompute intermediate;
multiplies<ComputeFragment> mul_add_source;
multiply_add<ComputeFragment> mul_add_accumulator;
multiplies<FragmentCompute> mul_add_source;
multiply_add<FragmentCompute> mul_add_accumulator;
if (Scale == ScaleType::NoBetaScaling)
intermediate = converted_source;
@@ -221,7 +245,7 @@ public:
/// Computes linear scaling: D = alpha * accumulator
CUTLASS_HOST_DEVICE
FragmentOutput operator()(
FragmentAccumulator const &accumulator) const {
FragmentAccumulator const &accumulator) const {
// Convert source to interal compute numeric type
NumericArrayConverter<ElementCompute, ElementAccumulator, kCount, Round> accumulator_converter;
@@ -229,14 +253,14 @@ public:
// Convert to destination numeric type
NumericArrayConverter<ElementOutput, ElementCompute, kCount, Round> destination_converter;
ComputeFragment converted_accumulator = accumulator_converter(accumulator);
FragmentCompute converted_accumulator = accumulator_converter(accumulator);
if (Scale == ScaleType::Nothing)
return destination_converter(converted_accumulator);
// Perform binary operations
ComputeFragment intermediate;
multiplies<ComputeFragment> mul_accumulator;
FragmentCompute intermediate;
multiplies<FragmentCompute> mul_accumulator;
intermediate = mul_accumulator(alpha_, converted_accumulator); // D = alpha * Accum
@@ -42,6 +42,7 @@
#include "cutlass/numeric_conversion.h"
#include "cutlass/epilogue/thread/activation.h"
#include "cutlass/epilogue/thread/scale_type.h"
/////////////////////////////////////////////////////////////////////////////////////////////////
@@ -90,7 +91,13 @@ public:
using FragmentZ = Array<ElementZ, kElementsPerAccess>;
using FragmentT = Array<ElementT, kElementsPerAccess>;
// Definitions needed for collective epilogue
using FragmentSource = FragmentC;
using FragmentOutput = FragmentZ;
using ElementBias = ElementVector;
using FragmentBias = FragmentCompute;
using ActivationFunctor = ElementwiseOp;
static const ScaleType::Kind kScale = ScaleType::Default;
static bool const kIsHeavy = ElementwiseOp::kIsHeavy;
@@ -196,8 +203,8 @@ public:
/// Applies the operation when is_source_needed() is true
CUTLASS_HOST_DEVICE
void operator()(
FragmentZ &frag_Z,
FragmentT &frag_T,
FragmentZ &frag_Z,
FragmentT &frag_T,
FragmentAccumulator const &AB,
FragmentC const &frag_C,
FragmentCompute const &V) const {
@@ -227,8 +234,8 @@ public:
/// Applies the operation when is_source_needed() is false
CUTLASS_HOST_DEVICE
void operator()(
FragmentZ &frag_Z,
FragmentT &frag_T,
FragmentZ &frag_Z,
FragmentT &frag_T,
FragmentAccumulator const &AB,
FragmentCompute const &V) const {
@@ -87,6 +87,7 @@ public:
using FragmentOutput = Array<ElementOutput, kCount>;
using FragmentAccumulator = Array<ElementAccumulator, kCount>;
using ComputeFragment = Array<ElementCompute, kCount>;
using FragmentSource = Array<ElementOutput, kCount>;
static FloatRoundStyle const kRound = Round;
@@ -35,7 +35,7 @@
#pragma once
#include <cutlass/half.h>
#include "cutlass/half.h"
#include "cutlass/cutlass.h"
#include "cutlass/numeric_types.h"
#include "cutlass/array.h"
@@ -34,7 +34,7 @@
#pragma once
#include <cutlass/half.h>
#include "cutlass/half.h"
#include "cutlass/cutlass.h"
#include "cutlass/numeric_types.h"
#include "cutlass/array.h"
@@ -78,6 +78,7 @@ public:
using FragmentOutput = Array<ElementOutput, kCount>;
using FragmentAccumulator = Array<ElementAccumulator, kCount>;
using FragmentSource = Array<ElementOutput, kCount>;
using FragmentCompute = Array<ElementCompute, kCount>;
static FloatRoundStyle const kRound = Round;
@@ -72,6 +72,7 @@ public:
using FragmentOutput = Array<ElementOutput, kCount>;
using FragmentAccumulator = Array<ElementAccumulator, kCount>;
using ComputeFragment = Array<ElementCompute, kCount>;
using FragmentSource = Array<ElementOutput, kCount>;
static FloatRoundStyle const kRound = Round;
@@ -56,13 +56,13 @@ struct LinearCombinationParams {
LinearCombinationParams(ElementCompute alpha, ElementCompute beta)
: alpha_data {0lu, 0lu}, beta_data {0lu, 0lu}
{
#if defined(__CUDA_ARCH__)
#if defined(__CUDA_ARCH__)
reinterpret_cast<ElementCompute&>(alpha_data) = alpha;
reinterpret_cast<ElementCompute&>(beta_data) = beta;
#else
#else
memcpy( alpha_data, &alpha, sizeof(ElementCompute) );
memcpy( beta_data, &beta, sizeof(ElementCompute) );
#endif
#endif
}
};
@@ -34,7 +34,7 @@
#pragma once
#include <cutlass/half.h>
#include "cutlass/half.h"
#include "cutlass/cutlass.h"
#include "cutlass/numeric_types.h"
#include "cutlass/array.h"
@@ -90,6 +90,7 @@ public:
using FragmentAccumulator = Array<ElementAccumulator, kCount>;
using FragmentCompute = Array<ElementCompute, kCount>;
using FragmentScaleBias = Array<ElementCompute, kCount>;
using FragmentSource = Array<ElementOutput, kCount>;
static FloatRoundStyle const kRound = Round;
@@ -321,6 +322,7 @@ public:
using FragmentAccumulator = Array<ElementAccumulator, kCount>;
using FragmentCompute = Array<ElementCompute, kCount>;
using FragmentScaleBias = Array<ElementCompute, kCount>;
using FragmentSource = Array<ElementOutput, kCount>;
static FloatRoundStyle const kRound = Round;
@@ -37,7 +37,7 @@
#pragma once
#include <cutlass/half.h>
#include "cutlass/half.h"
#include "cutlass/cutlass.h"
#include "cutlass/numeric_types.h"
#include "cutlass/array.h"
@@ -93,6 +93,7 @@ public:
using FragmentAccumulator = Array<ElementAccumulator, kCount>;
using FragmentCompute = Array<ElementCompute, kCount>;
using FragmentScaleBias = Array<ElementCompute, kCount>;
using FragmentSource = Array<ElementOutput, kCount>;
static FloatRoundStyle const kRound = Round;
@@ -308,6 +309,7 @@ public:
using FragmentAccumulator = Array<ElementAccumulator, kCount>;
using FragmentCompute = Array<ElementCompute, kCount>;
using FragmentScaleBias = Array<ElementCompute, kCount>;
using FragmentSource = Array<ElementOutput, kCount>;
static FloatRoundStyle const kRound = Round;
@@ -38,6 +38,7 @@
#include "cutlass/array.h"
#include "cutlass/functional.h"
#include "cutlass/numeric_conversion.h"
#include "cutlass/epilogue/thread/detail.hpp"
/////////////////////////////////////////////////////////////////////////////////////////////////
@@ -45,14 +46,6 @@ namespace cutlass {
namespace epilogue {
namespace thread {
namespace detail {
/// Dummy class used to designate that the second binary operator in the epilogue is unsued
template <typename T>
class NoOp {};
}
/// Models a residual block of the form: UnaryOp(BinaryOp(BinaryOp(ActivationOp(TensorOp(X) + bias), residual1), residual2))
template <typename ElementOutput_, typename ElementAccumulator_,
typename ElementCompute_, typename ElementC_, int ElementsPerAccess,
@@ -0,0 +1,251 @@
/***************************************************************************************************
* Copyright (c) 2017 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
/*! \file
\brief Functor performing linear combination operation, bias addition, and tensor-tensor
elementwise operations
*/
#pragma once
#include "cutlass/cutlass.h"
#include "cutlass/array.h"
#include "cutlass/functional.h"
#include "cutlass/numeric_conversion.h"
#include "cutlass/numeric_types.h"
#include "cutlass/epilogue/thread/activation.h"
#include "cutlass/epilogue/thread/detail.hpp"
#include "cutlass/epilogue/thread/scale_type.h"
/////////////////////////////////////////////////////////////////////////////////////////////////
namespace cutlass {
namespace epilogue {
namespace thread {
namespace detail {
/// Returns whether a source operand is needed for a combination of binary operation and scale
/// type. Simple specialized checks are made for cases in which 0 is an identity element of
/// the binary operation.
template <class BinaryOp, class ElementCompute, ScaleType::Kind Scale>
CUTLASS_HOST_DEVICE
bool is_binary_op_source_needed(ElementCompute scale) {
if constexpr (cute::is_same_v<BinaryOp, NoOp<ElementCompute>>) {
return false;
}
else if constexpr (cute::is_same_v<BinaryOp, plus<ElementCompute>> || cute::is_same_v<BinaryOp, minus<ElementCompute>>) {
// Cases for binary operators for which 0 is an identity element
if constexpr (Scale == ScaleType::NoBetaScaling) return true;
if constexpr (Scale == ScaleType::OnlyAlphaScaling) return false;
if constexpr (Scale == ScaleType::Nothing) return false;
return scale != ElementCompute(0);
}
return true;
}
} // namespace detail
/////////////////////////////////////////////////////////////////////////////////////////////////
/** Compute a tensor-tensor broadcast epilogue.
*
* @param ElementOutput_ Data type used to load and store tensors
* @param ElementAccumulator_ Accumulator data type
* @param ElementCompute_ Data type used to compute linear combination
* @param ElementBias_ Data type of Bias elements
* @param ActivationFunctor_ Fused Activation
* @param BinaryOp0_ Binary operation to perform on O0 and C0. detail::NoOp means no operation
* @param BinaryOp1_ Binary operation to perform on O1 and C1. detail::NoOp means no operation
* @param UnaryOp_ Unary operation to perform on final result
* @param Scale Controls the type of Alpha and Beta scaling to perform
* @param Round How values should be rounded in conversions
* @param ElementSource_ Data type used for source operands
*
* Computes the following:
* O0 = alpha * accumulator + bias
* O1 = BinaryOp0(O0, beta * C0)
* O2 = BinaryOp1(O1, beta * C1)
* D = UnaryOp(O2)
*/
template <
class ElementOutput_,
class ElementAccumulator_ = ElementOutput_,
class ElementCompute_ = ElementOutput_,
class ElementBias_ = ElementCompute_,
template <class T> class ActivationFunctor_ = Identity,
template <class T> class BinaryOp0_ = plus,
template <class T> class BinaryOp1_ = detail::NoOp,
template <class T> class UnaryOp_ = Identity,
ScaleType::Kind Scale = ScaleType::Default,
FloatRoundStyle Round = FloatRoundStyle::round_to_nearest,
class ElementSource_ = ElementOutput_
>
class LinearCombinationTensorBroadcast {
public:
using ElementOutput = ElementOutput_;
using ElementAccumulator = ElementAccumulator_;
using ElementCompute = ElementCompute_;
using ElementBias = ElementBias_;
using ElementC = ElementSource_;
using ElementD = ElementOutput_;
using ElementScalingFactor = ElementAccumulator_;
using UnaryOp = UnaryOp_<ElementCompute>;
using BinaryOp0 = BinaryOp0_<ElementCompute>;
using BinaryOp1 = BinaryOp1_<ElementCompute>;
using ActivationFunctor = ActivationFunctor_<ElementCompute>;
static constexpr int kCount = 1;
using FragmentOutput = Array<ElementOutput, kCount>;
using FragmentAccumulator = Array<ElementAccumulator, kCount>;
using ComputeFragment = Array<ElementCompute, kCount>;
using FragmentBias = Array<ElementBias, kCount>;
static constexpr FloatRoundStyle kRound = Round;
using NoOpType = detail::NoOp<ElementCompute>;
static constexpr bool IsBinaryOp0Enabled = !cute::is_same_v<BinaryOp0, NoOpType>;
static constexpr bool IsBinaryOp1Enabled = !cute::is_same_v<BinaryOp1, NoOpType>;
static constexpr bool IsUnaryOpEnabled = !cute::is_same_v<UnaryOp, NoOpType> && !cute::is_same_v<UnaryOp, Identity<ElementCompute>>;
/// Host-constructable parameters structure
struct Params {
ElementCompute alpha{}; ///< scales accumulators
ElementCompute beta{}; ///< scales source tensor
ElementCompute const* alpha_ptr = nullptr; ///< pointer to accumulator scalar - if not null, loads it from memory
ElementCompute const* beta_ptr = nullptr; ///< pointer to source scalar - if not null, loads it from memory
//
// Methods
//
Params() = default;
CUTLASS_HOST_DEVICE
Params(ElementCompute const* alpha_ptr, ElementCompute const* beta_ptr)
: alpha_ptr(alpha_ptr),
beta_ptr(beta_ptr) {}
CUTLASS_HOST_DEVICE
Params(ElementCompute const* alpha_ptr)
: alpha_ptr(alpha_ptr) {}
CUTLASS_HOST_DEVICE
Params(ElementCompute alpha,
ElementCompute beta)
: alpha(alpha),
beta(beta) {}
};
private:
//
// Data members
//
ElementCompute alpha_;
ElementCompute beta_;
public:
/// Constructs the function object, possibly loading from pointers in host memory
CUTLASS_HOST_DEVICE
LinearCombinationTensorBroadcast(Params const& params)
: alpha_(params.alpha_ptr ? *params.alpha_ptr : params.alpha),
beta_(params.beta_ptr ? *params.beta_ptr : params.beta) {}
/// Returns true if source 0 is needed
CUTLASS_HOST_DEVICE
bool is_source0_needed() const {
return detail::is_binary_op_source_needed<BinaryOp0, ElementCompute, Scale>(beta_);
}
/// Returns true if source 1 is needed
CUTLASS_HOST_DEVICE
bool is_source1_needed() const {
return detail::is_binary_op_source_needed<BinaryOp1, ElementCompute, Scale>(beta_);
}
//
// Specialization for scalar
//
CUTLASS_HOST_DEVICE
ElementD operator()(ElementAccumulator const accumulator, ElementC const source0, ElementC source1, ElementBias const bias) {
// Convert everything to Compute type, do compute, and then store to output type
NumericConverter<ElementCompute, ElementAccumulator, Round> accumulator_converter;
NumericConverter<ElementCompute, ElementBias, Round> bias_converter;
NumericConverter<ElementCompute, ElementC, Round> source_converter;
NumericConverter<ElementD, ElementCompute, Round> destination_converter;
ActivationFunctor act;
multiplies<ElementCompute> mul;
multiply_add<ElementCompute> madd;
ElementCompute intermediate = accumulator_converter(accumulator);
intermediate = madd(alpha_, intermediate, bias_converter(bias));
intermediate = act(intermediate);
// Apply BinaryOp0, if needed
if constexpr (IsBinaryOp0Enabled) {
BinaryOp0 bin0;
ElementCompute converted_source = source_converter(source0);
intermediate = bin0(intermediate, mul(beta_, converted_source));
}
// Apply BinaryOp1, if needed
if constexpr (IsBinaryOp1Enabled) {
BinaryOp1 bin1;
ElementCompute converted_source = source_converter(source1);
intermediate = bin1(intermediate, mul(beta_, converted_source));
}
// Apply UnaryOp, if needed
if constexpr (IsUnaryOpEnabled) {
UnaryOp unary;
intermediate = unary(intermediate);
}
return destination_converter(intermediate);
}
};
/////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace thread
} // namespace epilogue
} // namespace cutlass
/////////////////////////////////////////////////////////////////////////////////////////////////
@@ -35,7 +35,7 @@
#pragma once
#include <cutlass/half.h>
#include "cutlass/half.h"
#include "cutlass/cutlass.h"
#include "cutlass/numeric_types.h"
#include "cutlass/array.h"
@@ -71,7 +71,7 @@ template <
typename ThreadMap_, ///< Thread map (conept: OutputTileThreadMap)
typename Element_, ///< Element data type
bool ScatterD = false, ///< Scatter D operand or not
typename PermuteDLayout = layout::NoPermute, ///< Permute D operand or not
typename PermuteDLayout = layout::NoPermute, ///< Permute D operand or not
bool UseCUDAStore = false
>
class PredicatedTileIterator {
@@ -93,6 +93,8 @@ public:
static int const kThreads = ThreadMap::kThreads;
static int const kIterations = ThreadMap::Count::kTile;
static bool constexpr PermuteD = !layout::is_trivial_permute<PermuteDLayout>;
static_assert( ThreadMap::Iterations::kRow > 0,"ThreadMap::Iterations::kRow must be > 0");
static_assert( ThreadMap::Iterations::kGroup > 0,"ThreadMap::Iterations::kGroup must be > 0");
static_assert( ThreadMap::Iterations::kCluster > 0,"ThreadMap::Iterations::kCluster must be > 0");
@@ -202,11 +204,9 @@ private:
/// Scatter indices
int const *indices_;
/// Whether to perform Permute Op
bool PermuteD;
/// PermuteDLayout
mutable PermuteDLayout permute_layout_;
PermuteDLayout permute_layout_;
//
// Static asserts about internal strides
//
@@ -237,7 +237,8 @@ public:
TensorCoord threadblock_offset = TensorCoord(),
int const *indices = nullptr
):
params_(params), indices_(indices)
params_(params), indices_(indices),
permute_layout_(PitchLinearCoord(extent.column(), extent.row()), params_.stride * kElementsPerAccess / sizeof(AccessType))
{
TensorCoord thread_offset = ThreadMap::initial_offset(thread_idx) + threadblock_offset;
@@ -276,17 +277,7 @@ public:
}
// store_byte_pointer_ is set to be the same with byte_pointer_ unless PermuteD is used.
store_byte_pointer_ = byte_pointer_;
// Initialize PermuteD. If PermuteD is true, store_byte_pointer_ is initialized accordingly.
if (platform::is_same<PermuteDLayout, layout::NoPermute>::value) {
PermuteD = false;
}else{
PermuteD = true;
store_byte_pointer_ = reinterpret_cast<uint8_t *>(pointer);
permute_layout_ = PermuteDLayout(extent,
params_.stride * kElementsPerAccess / sizeof(AccessType));
}
store_byte_pointer_ = PermuteD ? reinterpret_cast<uint8_t *>(pointer) : byte_pointer_;
// Initialize internal state counter
state_[0] = state_[1] = state_[2] = 0;
@@ -411,18 +402,17 @@ public:
for (int column = 0; column < ThreadMap::Iterations::kColumn; ++column) {
bool guard = row_guard && mask_.predicates[column];
int col_offset = column * ThreadMap::Delta::kColumn;
if (PermuteD) {
int col_offset = column * ThreadMap::Delta::kColumn;
int col = col_offset + thread_start_column_;
int row = row_offset + thread_start_row_;
TensorCoord init_coord(row, col);
// Locate memory_pointer
memory_pointer = reinterpret_cast<AccessType *>(byte_pointer + byte_offset
+ permute_layout_(init_coord) * sizeof(AccessType) / kElementsPerAccess);
+ permute_layout_(PitchLinearCoord(col, row)) * sizeof(AccessType) / kElementsPerAccess);
}
if (UseCUDAStore) {
@@ -249,17 +249,21 @@ public:
/////////////////////////////////////////////////////////////////////////////////////////////////
/// Partial specialization for int32_t x 16 => int8_t/int4b_t x 16
/// Partial specialization for
/// int32_t x 16 => int8_t/int4b_t x 16 and
/// float x 16 => float_e4m3_t/float_e5m2_t x 16
template <
typename ThreadMap_, ///< Thread map (conept: OutputTileThreadMap)
typename ThreadMap_, ///< Thread map (concept: OutputTileThreadMap)
typename Element_,
int OutputSizeBits_ ///< Size of output element in bits
>
class SharedLoadIteratorMixed<ThreadMap_, int32_t, 32, OutputSizeBits_, 16, 8, true> {
class SharedLoadIteratorMixed<ThreadMap_, Element_, 32, OutputSizeBits_, 16, 8, true> {
public:
using ThreadMap = ThreadMap_;
using Shape = typename ThreadMap::Shape;
using Element = int32_t;
using Element = Element_;
static_assert(sizeof_bits<Element>::value == 32, "Element size in bits must be 32.");
using Layout = layout::RowMajor;
using TensorRef = TensorRef<Element, Layout>;
@@ -414,17 +418,21 @@ public:
/////////////////////////////////////////////////////////////////////////////////////////////////
/// Partial specialization for int32_t x 8 => int8_t/int4b_t x 8
/// Partial specialization for:
/// int32_t x 8 => int8_t/int4b_t x 8 and
/// float x 8 => float_e4m3_t/float_e5m2_t x 8
template <
typename ThreadMap_, ///< Thread map (conept: OutputTileThreadMap)
typename ThreadMap_, ///< Thread map (concept: OutputTileThreadMap)
typename Element_,
int OutputSizeBits_
>
class SharedLoadIteratorMixed<ThreadMap_, int32_t, 32, OutputSizeBits_, 8, 8, true> {
class SharedLoadIteratorMixed<ThreadMap_, Element_, 32, OutputSizeBits_, 8, 8, true> {
public:
using ThreadMap = ThreadMap_;
using Shape = typename ThreadMap::Shape;
using Element = int32_t;
using Element = Element_;
static_assert(sizeof_bits<Element>::value == 32, "Element size in bits must be 32.");
using Layout = layout::RowMajor;
using TensorRef = TensorRef<Element, Layout>;
@@ -716,7 +716,6 @@ public:
};
/////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace warp
} // namespace epilogue
} // namespace cutlass