co-authored by
Aniket Shivam
parent
9b8166e3f0
commit
d572cc1aab
@@ -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
|
||||
+102
-60
@@ -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;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
+28
-15
@@ -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
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
+558
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user