co-authored by
Aniket Shivam
parent
ca23ff7924
commit
b72cbf957d
@@ -65,6 +65,8 @@
|
||||
#include "cutlass/epilogue/threadblock/default_epilogue_simt.h"
|
||||
#include "cutlass/transform/threadblock/predicated_tile_iterator.h"
|
||||
|
||||
#include "cutlass/layout/permute.h"
|
||||
|
||||
#if defined(CUTLASS_ARCH_WMMA_ENABLED)
|
||||
#include "cutlass/epilogue/threadblock/default_epilogue_wmma_tensor_op.h"
|
||||
#endif //CUTLASS_ARCH_WMMA_ENABLED
|
||||
@@ -125,6 +127,8 @@ template <
|
||||
bool GatherB = false,
|
||||
/// Scatter result D by using an index array
|
||||
bool ScatterD = false,
|
||||
/// Permute result D
|
||||
typename PermuteDLayout = layout::NoPermute,
|
||||
///
|
||||
typename Enable = void
|
||||
>
|
||||
@@ -177,13 +181,15 @@ template <
|
||||
/// Gather operand B by using an index array
|
||||
bool GatherB,
|
||||
/// Scatter result D by using an index array
|
||||
bool ScatterD
|
||||
bool ScatterD,
|
||||
/// Permute result D
|
||||
typename PermuteDLayout
|
||||
>
|
||||
struct DefaultGemm<ElementA, LayoutA, kAlignmentA, ElementB, LayoutB, kAlignmentB, ElementC,
|
||||
LayoutC, ElementAccumulator, arch::OpClassTensorOp,
|
||||
arch::Sm80, ThreadblockShape, WarpShape, InstructionShape,
|
||||
EpilogueOutputOp, ThreadblockSwizzle, Stages, SplitKSerial,
|
||||
Operator, SharedMemoryClear, GatherA, GatherB, ScatterD> {
|
||||
Operator, SharedMemoryClear, GatherA, GatherB, ScatterD, PermuteDLayout> {
|
||||
|
||||
static_assert(platform::is_same<LayoutC, layout::RowMajor>::value
|
||||
|| platform::is_same<LayoutC, layout::AffineRankN<2>>::value,
|
||||
@@ -202,14 +208,14 @@ struct DefaultGemm<ElementA, LayoutA, kAlignmentA, ElementB, LayoutB, kAlignment
|
||||
using RegularEpilogue =
|
||||
typename cutlass::epilogue::threadblock::DefaultEpilogueTensorOp<
|
||||
ThreadblockShape, typename Mma::Operator, kPartitionsK, EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount, ScatterD>::Epilogue;
|
||||
EpilogueOutputOp::kCount, ScatterD, PermuteDLayout>::Epilogue;
|
||||
|
||||
using Affine2Epilogue =
|
||||
typename cutlass::epilogue::threadblock::DefaultEpilogueTensorOpAffineRankN<
|
||||
2, ThreadblockShape, typename Mma::Operator, kPartitionsK, EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount>::Epilogue;
|
||||
|
||||
using Epilogue = typename cutlass::platform::conditional<platform::is_same<LayoutC, layout::RowMajor>::value,
|
||||
using Epilogue = typename platform::conditional<platform::is_same<LayoutC, layout::RowMajor>::value,
|
||||
RegularEpilogue,
|
||||
Affine2Epilogue>::type;
|
||||
|
||||
@@ -258,7 +264,9 @@ template <
|
||||
/// Gather operand B by using an index array
|
||||
bool GatherB,
|
||||
/// Scatter result D by using an index array
|
||||
bool ScatterD
|
||||
bool ScatterD,
|
||||
/// Permute result D
|
||||
typename PermuteDLayout
|
||||
>
|
||||
struct DefaultGemm<
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
@@ -278,7 +286,8 @@ struct DefaultGemm<
|
||||
SharedMemoryClear,
|
||||
GatherA,
|
||||
GatherB,
|
||||
ScatterD
|
||||
ScatterD,
|
||||
PermuteDLayout
|
||||
> {
|
||||
|
||||
/// Define the threadblock-scoped matrix multiply-accumulate
|
||||
@@ -313,7 +322,8 @@ struct DefaultGemm<
|
||||
kPartitionsK,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount,
|
||||
ScatterD
|
||||
ScatterD,
|
||||
PermuteDLayout
|
||||
>::Epilogue;
|
||||
|
||||
/// Define the kernel-level GEMM operator.
|
||||
@@ -493,7 +503,9 @@ template <
|
||||
/// Gather operand B by using an index array
|
||||
bool GatherB,
|
||||
/// Scatter result D by using an index array
|
||||
bool ScatterD
|
||||
bool ScatterD,
|
||||
/// Permute result D
|
||||
typename PermuteDLayout
|
||||
>
|
||||
struct DefaultGemm<
|
||||
ElementA, LayoutA, kAlignmentA,
|
||||
@@ -513,7 +525,8 @@ struct DefaultGemm<
|
||||
SharedMemoryClear,
|
||||
GatherA,
|
||||
GatherB,
|
||||
ScatterD
|
||||
ScatterD,
|
||||
PermuteDLayout
|
||||
> {
|
||||
|
||||
/// Define the threadblock-scoped matrix multiply-accumulate
|
||||
@@ -548,7 +561,8 @@ struct DefaultGemm<
|
||||
kPartitionsK,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount,
|
||||
ScatterD
|
||||
ScatterD,
|
||||
PermuteDLayout
|
||||
>::Epilogue;
|
||||
|
||||
/// Define the kernel-level GEMM operator.
|
||||
@@ -598,7 +612,9 @@ template <
|
||||
/// Gather operand B by using an index array
|
||||
bool GatherB,
|
||||
/// Scatter result D by using an index array
|
||||
bool ScatterD
|
||||
bool ScatterD,
|
||||
/// Permute result D
|
||||
typename PermuteDLayout
|
||||
>
|
||||
struct DefaultGemm<
|
||||
ElementA,
|
||||
@@ -624,6 +640,7 @@ struct DefaultGemm<
|
||||
GatherA,
|
||||
GatherB,
|
||||
ScatterD,
|
||||
PermuteDLayout,
|
||||
typename platform::enable_if< ! platform::is_same<ArchTag, arch::Sm80>::value >::type > {
|
||||
|
||||
static_assert(platform::is_same<LayoutC, layout::RowMajor>::value
|
||||
@@ -661,7 +678,8 @@ struct DefaultGemm<
|
||||
typename Mma::Operator,
|
||||
EpilogueOutputOp,
|
||||
kEpilogueElementsPerAccess,
|
||||
ScatterD
|
||||
ScatterD,
|
||||
PermuteDLayout
|
||||
>::Epilogue;
|
||||
|
||||
using Affine2Epilogue = typename cutlass::epilogue::threadblock::DefaultEpilogueSimtAffineRankN<
|
||||
@@ -672,7 +690,7 @@ struct DefaultGemm<
|
||||
kEpilogueElementsPerAccess
|
||||
>::Epilogue;
|
||||
|
||||
using Epilogue = typename cutlass::platform::conditional<platform::is_same<LayoutC, layout::RowMajor>::value,
|
||||
using Epilogue = typename platform::conditional<platform::is_same<LayoutC, layout::RowMajor>::value,
|
||||
RegularEpilogue,
|
||||
Affine2Epilogue>::type;
|
||||
|
||||
@@ -723,7 +741,9 @@ template <
|
||||
/// Gather operand B by using an index array
|
||||
bool GatherB,
|
||||
/// Scatter result D by using an index array
|
||||
bool ScatterD
|
||||
bool ScatterD,
|
||||
/// Permute result D
|
||||
typename PermuteDLayout
|
||||
>
|
||||
struct DefaultGemm<ElementA,
|
||||
LayoutA,
|
||||
@@ -747,7 +767,8 @@ struct DefaultGemm<ElementA,
|
||||
SharedMemoryClear,
|
||||
GatherA,
|
||||
GatherB,
|
||||
ScatterD> {
|
||||
ScatterD,
|
||||
PermuteDLayout> {
|
||||
|
||||
static_assert(platform::is_same<LayoutC, layout::RowMajor>::value
|
||||
|| platform::is_same<LayoutC, layout::AffineRankN<2>>::value,
|
||||
@@ -769,7 +790,8 @@ struct DefaultGemm<ElementA,
|
||||
typename Mma::Operator,
|
||||
EpilogueOutputOp,
|
||||
kEpilogueElementsPerAccess,
|
||||
ScatterD
|
||||
ScatterD,
|
||||
PermuteDLayout
|
||||
>::Epilogue;
|
||||
|
||||
using Affine2Epilogue = typename cutlass::epilogue::threadblock::DefaultEpilogueSimtAffineRankN<
|
||||
@@ -780,7 +802,7 @@ struct DefaultGemm<ElementA,
|
||||
kEpilogueElementsPerAccess
|
||||
>::Epilogue;
|
||||
|
||||
using Epilogue = typename cutlass::platform::conditional<platform::is_same<LayoutC, layout::RowMajor>::value,
|
||||
using Epilogue = typename platform::conditional<platform::is_same<LayoutC, layout::RowMajor>::value,
|
||||
RegularEpilogue,
|
||||
Affine2Epilogue>::type;
|
||||
|
||||
|
||||
@@ -54,6 +54,8 @@
|
||||
#include "cutlass/gemm/kernel/default_gemm_complex.h"
|
||||
#include "cutlass/gemm/device/default_gemm_configuration.h"
|
||||
|
||||
#include "cutlass/layout/permute.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
@@ -101,12 +103,16 @@ template <
|
||||
typename ThreadblockSwizzle,
|
||||
/// Number of stages used in the pipelined mainloop
|
||||
int Stages,
|
||||
/// Whether the schedule of problems to visit has been precomputed
|
||||
GroupScheduleMode GroupScheduleMode_ = GroupScheduleMode::kDeviceOnly,
|
||||
/// Operation performed by GEMM
|
||||
typename Operator = typename device::DefaultGemmConfiguration<
|
||||
OperatorClass, ArchTag, ElementA_, ElementB_, ElementC_,
|
||||
ElementAccumulator>::Operator,
|
||||
/// Use zfill or predicate for out-of-bound cp.async
|
||||
SharedMemoryClearOption SharedMemoryClear = SharedMemoryClearOption::kNone,
|
||||
/// Permute result D
|
||||
typename PermuteDLayout = layout::NoPermute,
|
||||
///
|
||||
typename Enable = void
|
||||
>
|
||||
@@ -152,10 +158,14 @@ template <
|
||||
typename ThreadblockSwizzle,
|
||||
/// Number of stages used in the pipelined mainloop
|
||||
int Stages,
|
||||
/// Whether the schedule of problems to visit has been precomputed
|
||||
GroupScheduleMode GroupScheduleMode_,
|
||||
/// Operation performed by GEMM
|
||||
typename Operator,
|
||||
/// Use zfill or predicate for out-of-bound cp.async
|
||||
SharedMemoryClearOption SharedMemoryClear
|
||||
SharedMemoryClearOption SharedMemoryClear,
|
||||
/// Permute result D
|
||||
typename PermuteDLayout
|
||||
>
|
||||
struct DefaultGemmGrouped<
|
||||
ElementA,
|
||||
@@ -177,9 +187,11 @@ struct DefaultGemmGrouped<
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
Stages,
|
||||
GroupScheduleMode_,
|
||||
Operator,
|
||||
SharedMemoryClear,
|
||||
typename std::enable_if< ! cutlass::is_complex<ElementAccumulator>::value>::type
|
||||
PermuteDLayout,
|
||||
typename platform::enable_if< ! cutlass::is_complex<ElementAccumulator>::value>::type
|
||||
> {
|
||||
|
||||
// If true, we must construct a 'transposed-and-exchanged' Mma operator.
|
||||
@@ -219,7 +231,11 @@ struct DefaultGemmGrouped<
|
||||
Stages,
|
||||
true,
|
||||
Operator,
|
||||
SharedMemoryClear
|
||||
SharedMemoryClear,
|
||||
false, /*GatherA*/
|
||||
false, /*GatherB*/
|
||||
false, /*ScatterD*/
|
||||
PermuteDLayout
|
||||
>::GemmKernel;
|
||||
|
||||
/// Define the kernel in terms of the default kernel
|
||||
@@ -227,6 +243,7 @@ struct DefaultGemmGrouped<
|
||||
typename DefaultGemmKernel::Mma,
|
||||
typename DefaultGemmKernel::Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
GroupScheduleMode_,
|
||||
kInternalTranspose
|
||||
>;
|
||||
};
|
||||
@@ -276,6 +293,8 @@ template <
|
||||
typename ThreadblockSwizzle,
|
||||
/// Number of stages used in the pipelined mainloop
|
||||
int Stages,
|
||||
/// Whether the schedule of problems to visit has been precomputed
|
||||
GroupScheduleMode GroupScheduleMode_,
|
||||
/// Operation performed by GEMM
|
||||
typename Operator,
|
||||
/// Use zfill or predicate for out-of-bound cp.async
|
||||
@@ -301,9 +320,11 @@ struct DefaultGemmGrouped<
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
Stages,
|
||||
GroupScheduleMode_,
|
||||
Operator,
|
||||
SharedMemoryClear,
|
||||
typename std::enable_if<cutlass::is_complex<ElementAccumulator>::value>::type
|
||||
layout::NoPermute, /*PermuteDLayout*/
|
||||
typename platform::enable_if<cutlass::is_complex<ElementAccumulator>::value>::type
|
||||
> {
|
||||
|
||||
// If true, we must construct a 'transposed-and-exchanged' Mma operator.
|
||||
@@ -349,6 +370,7 @@ struct DefaultGemmGrouped<
|
||||
typename DefaultGemmKernel::Mma,
|
||||
typename DefaultGemmKernel::Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
GroupScheduleMode_,
|
||||
kInternalTranspose
|
||||
>;
|
||||
};
|
||||
|
||||
@@ -0,0 +1,164 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017 - 2022 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
|
||||
Default kernel-level softmax-grouped-GEMM
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
|
||||
#include "cutlass/complex.h"
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/kernel/gemm_grouped_softmax_mainloop_fusion.h"
|
||||
#include "cutlass/gemm/kernel/gemm_transpose_operands.h"
|
||||
#include "cutlass/gemm/kernel/default_gemm.h"
|
||||
#include "cutlass/gemm/kernel/default_gemm_complex.h"
|
||||
#include "cutlass/gemm/device/default_gemm_configuration.h"
|
||||
#include "cutlass/gemm/threadblock/default_mma_softmax_mainloop_fusion.h"
|
||||
|
||||
#include "cutlass/layout/permute.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace kernel {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
/// Element type for A matrix operand
|
||||
typename ElementA_,
|
||||
/// Layout type for A matrix operand
|
||||
typename LayoutA_,
|
||||
/// Complex elementwise transformation on A operand
|
||||
ComplexTransform TransformA,
|
||||
/// Access granularity of A matrix in units of elements
|
||||
int kAlignmentA,
|
||||
/// Element type for B matrix operand
|
||||
typename ElementB_,
|
||||
/// Layout type for B matrix operand
|
||||
typename LayoutB_,
|
||||
/// Complex elementwise transformation on B operand
|
||||
ComplexTransform TransformB,
|
||||
/// Access granularity of B matrix in units of elements
|
||||
int kAlignmentB,
|
||||
/// Element type for Scale/Bias vectors
|
||||
typename ElementScaleBias_,
|
||||
/// Layout type for Scale/Bias vectors
|
||||
typename LayoutScaleBias_,
|
||||
/// Element type for C and D matrix operands
|
||||
typename ElementC_,
|
||||
/// Layout type for C and D matrix operands
|
||||
typename LayoutC_,
|
||||
/// Element type for internal accumulation
|
||||
typename ElementAccumulator,
|
||||
/// Operator class tag
|
||||
typename OperatorClass,
|
||||
/// Tag indicating architecture to tune for
|
||||
typename ArchTag,
|
||||
/// Threadblock-level tile size (concept: GemmShape)
|
||||
typename ThreadblockShape,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename WarpShape,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename InstructionShape,
|
||||
/// Epilogue output operator
|
||||
typename EpilogueOutputOp,
|
||||
/// Threadblock-level swizzling operator
|
||||
typename ThreadblockSwizzle,
|
||||
/// Number of stages used in the pipelined mainloop
|
||||
int Stages,
|
||||
/// Whether the schedule of problems to visit has been precomputed
|
||||
GroupScheduleMode GroupScheduleMode_ = GroupScheduleMode::kDeviceOnly,
|
||||
/// Operation performed by GEMM
|
||||
typename Operator = typename device::DefaultGemmConfiguration<
|
||||
OperatorClass, ArchTag, ElementA_, ElementB_, ElementC_,
|
||||
ElementAccumulator>::Operator,
|
||||
/// Use zfill or predicate for out-of-bound cp.async
|
||||
SharedMemoryClearOption SharedMemoryClear = SharedMemoryClearOption::kNone
|
||||
>
|
||||
struct DefaultGemmGroupedSoftmaxMainloopFusion {
|
||||
// If true, we must construct a 'transposed-and-exchanged' Mma operator.
|
||||
static bool const kInternalTranspose = platform::is_same<LayoutC_, layout::ColumnMajor>::value;
|
||||
|
||||
using MapArguments = kernel::detail::MapArguments<
|
||||
ElementA_,
|
||||
LayoutA_,
|
||||
ComplexTransform::kNone,
|
||||
kAlignmentA,
|
||||
ElementB_,
|
||||
LayoutB_,
|
||||
ComplexTransform::kNone,
|
||||
kAlignmentB,
|
||||
LayoutC_,
|
||||
kInternalTranspose
|
||||
>;
|
||||
|
||||
private:
|
||||
/// Define the threadblock-scoped matrix multiply-accumulate
|
||||
using Mma = typename cutlass::gemm::threadblock::DefaultMmaSoftmaxMainloopFusion<
|
||||
typename MapArguments::ElementA, typename MapArguments::LayoutA, MapArguments::kAlignmentA,
|
||||
typename MapArguments::ElementB, typename MapArguments::LayoutB, MapArguments::kAlignmentB,
|
||||
ElementScaleBias_, LayoutScaleBias_, ElementAccumulator, layout::RowMajor, OperatorClass, ArchTag,
|
||||
ThreadblockShape, WarpShape, InstructionShape, Stages, kInternalTranspose,
|
||||
Operator, false, SharedMemoryClear>::ThreadblockMma;
|
||||
|
||||
static const int kPartitionsK = ThreadblockShape::kK / WarpShape::kK;
|
||||
|
||||
/// Define the epilogue
|
||||
using Epilogue =
|
||||
typename cutlass::epilogue::threadblock::DefaultEpilogueTensorOp<
|
||||
ThreadblockShape, typename Mma::Operator, kPartitionsK, EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount>::Epilogue;
|
||||
|
||||
public:
|
||||
using GemmKernel = kernel::GemmGroupedSoftmaxMainloopFusion<
|
||||
Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
GroupScheduleMode_,
|
||||
kInternalTranspose
|
||||
>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace kernel
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,137 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017 - 2022 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
|
||||
Default kernel-level GEMM definitions combine threadblock-scoped matrix multiply-add with
|
||||
the appropriate threadblock-scoped epilogue.
|
||||
|
||||
Note, CUTLASS epilogues universally target row-major outputs. Column-major outputs are
|
||||
accommodated by exchanging A and B operands and assuming transposed layouts. Partial
|
||||
specializations here choose 'device::GemmTransposed' to implement this functionality.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/arch/wmma.h"
|
||||
|
||||
#include "cutlass/epilogue/threadblock/epilogue.h"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/gemm/kernel/gemm_layernorm_mainloop_fusion.h"
|
||||
#include "cutlass/gemm/threadblock/default_mma_layernorm_mainloop_fusion.h"
|
||||
#include "cutlass/gemm/threadblock/threadblock_swizzle.h"
|
||||
|
||||
#include "cutlass/epilogue/threadblock/default_epilogue_tensor_op.h"
|
||||
#include "cutlass/transform/threadblock/predicated_tile_iterator.h"
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace kernel {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
/// Element type for A matrix operand
|
||||
typename ElementA,
|
||||
/// Layout type for A matrix operand
|
||||
typename LayoutA,
|
||||
/// Access granularity of A matrix in units of elements
|
||||
int kAlignmentA,
|
||||
/// Element type for B matrix operand
|
||||
typename ElementB,
|
||||
/// Layout type for B matrix operand
|
||||
typename LayoutB,
|
||||
/// Access granularity of B matrix in units of elements
|
||||
int kAlignmentB,
|
||||
/// Element type for Scale/Bias vectors
|
||||
typename ElementScaleBias,
|
||||
/// Layout type for Scale/Bias vectors
|
||||
typename LayoutScaleBias,
|
||||
/// Element type for C and D matrix operands
|
||||
typename ElementC,
|
||||
/// Layout type for C and D matrix operands
|
||||
typename LayoutC,
|
||||
/// Element type for internal accumulation
|
||||
typename ElementAccumulator,
|
||||
/// Operator class tag
|
||||
typename OperatorClass,
|
||||
/// Tag indicating architecture to tune for
|
||||
typename ArchTag,
|
||||
/// Threadblock-level tile size (concept: GemmShape)
|
||||
typename ThreadblockShape,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename WarpShape,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename InstructionShape,
|
||||
/// Epilogue output operator
|
||||
typename EpilogueOutputOp,
|
||||
/// Threadblock-level swizzling operator
|
||||
typename ThreadblockSwizzle,
|
||||
/// Number of stages used in the pipelined mainloop
|
||||
int Stages,
|
||||
/// Operation performed by GEMM
|
||||
typename Operator,
|
||||
/// Use zfill or predicate for out-of-bound cp.async
|
||||
SharedMemoryClearOption SharedMemoryClear = SharedMemoryClearOption::kNone>
|
||||
struct DefaultGemmLayernormMainloopFusion {
|
||||
|
||||
/// Define the threadblock-scoped matrix multiply-accumulate
|
||||
using Mma = typename cutlass::gemm::threadblock::DefaultMmaLayernormMainloopFusion<
|
||||
ElementA, LayoutA, kAlignmentA, ElementB, LayoutB, kAlignmentB,
|
||||
ElementScaleBias, LayoutScaleBias, ElementAccumulator, layout::RowMajor, arch::OpClassTensorOp, arch::Sm80,
|
||||
ThreadblockShape, WarpShape, InstructionShape, Stages,
|
||||
Operator, false, SharedMemoryClear>::ThreadblockMma;
|
||||
|
||||
static const int kPartitionsK = ThreadblockShape::kK / WarpShape::kK;
|
||||
|
||||
/// Define the epilogue
|
||||
using Epilogue =
|
||||
typename cutlass::epilogue::threadblock::DefaultEpilogueTensorOp<
|
||||
ThreadblockShape, typename Mma::Operator, kPartitionsK, EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount>::Epilogue;
|
||||
|
||||
/// Define the kernel-level GEMM operator.
|
||||
using GemmKernel = kernel::GemmLayernormMainloopFusion<Mma, Epilogue, ThreadblockSwizzle>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace kernel
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -52,6 +52,8 @@
|
||||
#include "cutlass/gemm/kernel/default_gemm.h"
|
||||
#include "cutlass/gemm/kernel/default_gemm_complex.h"
|
||||
|
||||
#include "cutlass/layout/permute.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
@@ -109,6 +111,8 @@ template <
|
||||
bool GatherB = false,
|
||||
/// Scatter result D by using an index array
|
||||
bool ScatterD = false,
|
||||
/// Permute result D
|
||||
typename PermuteDLayout = layout::NoPermute,
|
||||
///
|
||||
typename Enable = void
|
||||
>
|
||||
@@ -163,7 +167,9 @@ template <
|
||||
/// Gather operand B by using an index array
|
||||
bool GatherB,
|
||||
/// Scatter result D by using an index array
|
||||
bool ScatterD
|
||||
bool ScatterD,
|
||||
/// Permute result D
|
||||
typename PermuteDLayout
|
||||
>
|
||||
struct DefaultGemmUniversal<
|
||||
ElementA,
|
||||
@@ -190,6 +196,7 @@ struct DefaultGemmUniversal<
|
||||
GatherA,
|
||||
GatherB,
|
||||
ScatterD,
|
||||
PermuteDLayout,
|
||||
typename platform::enable_if< ! cutlass::is_complex<ElementAccumulator>::value>::type
|
||||
> {
|
||||
|
||||
@@ -216,7 +223,8 @@ struct DefaultGemmUniversal<
|
||||
SharedMemoryClear,
|
||||
GatherA,
|
||||
GatherB,
|
||||
ScatterD
|
||||
ScatterD,
|
||||
PermuteDLayout
|
||||
>::GemmKernel;
|
||||
|
||||
/// Define the kernel in terms of the default kernel
|
||||
@@ -302,6 +310,7 @@ struct DefaultGemmUniversal<
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
layout::NoPermute,
|
||||
typename platform::enable_if<cutlass::is_complex<ElementAccumulator>::value>::type
|
||||
> {
|
||||
|
||||
|
||||
@@ -0,0 +1,355 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017 - 2022 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
|
||||
Default kernel-level grouped Rank2K.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
|
||||
#include "cutlass/complex.h"
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#include "cutlass/gemm/kernel/rank_2k_transpose_operands.h"
|
||||
#include "cutlass/gemm/kernel/default_rank_2k.h"
|
||||
#include "cutlass/gemm/kernel/default_rank_2k_complex.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace kernel {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
/// Element type for A matrix operand
|
||||
typename ElementA,
|
||||
/// Layout type for A matrix operand
|
||||
typename LayoutA,
|
||||
/// Complex elementwise transformation on A operand
|
||||
ComplexTransform TransformA,
|
||||
/// Access granularity of A matrix in units of elements
|
||||
int kAlignmentA,
|
||||
/// Element type for B matrix operand
|
||||
typename ElementB,
|
||||
/// Layout type for B matrix operand
|
||||
typename LayoutB,
|
||||
/// Complex elementwise transformation on B operand
|
||||
ComplexTransform TransformB,
|
||||
/// Access granularity of B matrix in units of elements
|
||||
int kAlignmentB,
|
||||
/// Element type for C and D matrix operands
|
||||
typename ElementC,
|
||||
/// Layout type for C and D matrix operands
|
||||
typename LayoutC,
|
||||
/// Fill Mode for C (kLower or kUpper)
|
||||
FillMode FillModeC,
|
||||
/// Element type for internal accumulation
|
||||
typename ElementAccumulator,
|
||||
/// Operator class tag
|
||||
typename OperatorClass,
|
||||
/// Tag indicating architecture to tune for
|
||||
typename ArchTag,
|
||||
/// Threadblock-level tile size (concept: GemmShape)
|
||||
typename ThreadblockShape,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename WarpShape,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename InstructionShape,
|
||||
/// Epilogue output operator
|
||||
typename EpilogueOutputOp,
|
||||
/// Threadblock-level swizzling operator
|
||||
typename ThreadblockSwizzle,
|
||||
/// Number of stages used in the pipelined mainloop
|
||||
int Stages,
|
||||
/// Operation performed by GEMM
|
||||
typename Operator,
|
||||
/// Blas3 computation mode
|
||||
BlasMode BlasMode_ = BlasMode::kSymmetric,
|
||||
/// Whether the schedule of problems to visit has been precomputed
|
||||
GroupScheduleMode GroupScheduleMode_ = GroupScheduleMode::kDeviceOnly,
|
||||
///
|
||||
typename Enable = void
|
||||
>
|
||||
struct DefaultRank2KGrouped;
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
//
|
||||
// Real-valued grouped Rank2K
|
||||
//
|
||||
|
||||
template <
|
||||
/// Element type for A matrix operand
|
||||
typename ElementA,
|
||||
/// Layout type for A matrix operand
|
||||
typename LayoutA,
|
||||
/// Complex elementwise transformation on A operand
|
||||
ComplexTransform TransformA,
|
||||
/// Access granularity of A matrix in units of elements
|
||||
int kAlignmentA,
|
||||
/// Element type for B matrix operand
|
||||
typename ElementB,
|
||||
/// Layout type for B matrix operand
|
||||
typename LayoutB,
|
||||
/// Complex elementwise transformation on B operand
|
||||
ComplexTransform TransformB,
|
||||
/// Access granularity of B matrix in units of elements
|
||||
int kAlignmentB,
|
||||
/// Element type for C and D matrix operands
|
||||
typename ElementC,
|
||||
/// Layout type for C and D matrix operands
|
||||
typename LayoutC,
|
||||
/// Fill Mode for C (kLower or kUpper)
|
||||
FillMode FillModeC,
|
||||
/// Element type for internal accumulation
|
||||
typename ElementAccumulator,
|
||||
/// Operator class tag
|
||||
typename OperatorClass,
|
||||
/// Tag indicating architecture to tune for
|
||||
typename ArchTag,
|
||||
/// Threadblock-level tile size (concept: GemmShape)
|
||||
typename ThreadblockShape,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename WarpShape,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename InstructionShape,
|
||||
/// Epilogue output operator
|
||||
typename EpilogueOutputOp,
|
||||
/// Threadblock-level swizzling operator
|
||||
typename ThreadblockSwizzle,
|
||||
/// Number of stages used in the pipelined mainloop
|
||||
int Stages,
|
||||
/// Operation performed by GEMM
|
||||
typename Operator,
|
||||
/// Blas3 computation mode
|
||||
BlasMode BlasMode_,
|
||||
/// Whether the schedule of problems to visit has been precomputed
|
||||
GroupScheduleMode GroupScheduleMode_
|
||||
>
|
||||
struct DefaultRank2KGrouped<ElementA, LayoutA, TransformA, kAlignmentA,
|
||||
ElementB, LayoutB, TransformB, kAlignmentB,
|
||||
ElementC, LayoutC,
|
||||
FillModeC, ElementAccumulator, OperatorClass, ArchTag, ThreadblockShape,
|
||||
WarpShape, InstructionShape, EpilogueOutputOp,
|
||||
ThreadblockSwizzle, Stages, Operator, BlasMode_, GroupScheduleMode_,
|
||||
typename std::enable_if< ! cutlass::is_complex<ElementAccumulator>::value>::type
|
||||
> {
|
||||
// If true, we must construct a 'transposed-and-exchanged' Rank2K operator.
|
||||
static bool const kInternalTranspose = platform::is_same<LayoutC, layout::ColumnMajor>::value;
|
||||
|
||||
using MapArguments = kernel::detail::Rank2KMapArguments<
|
||||
ElementA,
|
||||
LayoutA,
|
||||
TransformA,
|
||||
kAlignmentA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
TransformB,
|
||||
kAlignmentB,
|
||||
LayoutC,
|
||||
FillModeC,
|
||||
kInternalTranspose
|
||||
>;
|
||||
|
||||
// Define the default grouped Rank2K kernel
|
||||
using DefaultRank2Kkernel = typename kernel::DefaultRank2K<
|
||||
typename MapArguments::ElementA,
|
||||
typename MapArguments::LayoutA,
|
||||
MapArguments::kAlignmentA,
|
||||
typename MapArguments::ElementB,
|
||||
typename MapArguments::LayoutB,
|
||||
MapArguments::kAlignmentB,
|
||||
ElementC,
|
||||
typename MapArguments::LayoutC,
|
||||
MapArguments::kFillModeC,
|
||||
ElementAccumulator,
|
||||
OperatorClass,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
Stages,
|
||||
false, // SplitKSerial
|
||||
Operator,
|
||||
BlasMode_
|
||||
>::Rank2Kkernel;
|
||||
|
||||
/// Define the kernel in terms of the default kernel
|
||||
using Rank2Kkernel = kernel::Rank2KGrouped<
|
||||
typename DefaultRank2Kkernel::Mma1,
|
||||
typename DefaultRank2Kkernel::Mma2,
|
||||
typename DefaultRank2Kkernel::Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
TransformA,
|
||||
TransformB,
|
||||
DefaultRank2Kkernel::kFillModeC,
|
||||
DefaultRank2Kkernel::kBlasMode,
|
||||
GroupScheduleMode_,
|
||||
kInternalTranspose
|
||||
>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
//
|
||||
// Complex-valued grouped Rank2K
|
||||
//
|
||||
|
||||
template <
|
||||
/// Element type for A matrix operand
|
||||
typename ElementA,
|
||||
/// Layout type for A matrix operand
|
||||
typename LayoutA,
|
||||
/// Complex elementwise transformation on A operand
|
||||
ComplexTransform TransformA,
|
||||
/// Access granularity of A matrix in units of elements
|
||||
int kAlignmentA,
|
||||
/// Element type for B matrix operand
|
||||
typename ElementB,
|
||||
/// Layout type for B matrix operand
|
||||
typename LayoutB,
|
||||
/// Complex elementwise transformation on B operand
|
||||
ComplexTransform TransformB,
|
||||
/// Access granularity of B matrix in units of elements
|
||||
int kAlignmentB,
|
||||
/// Element type for C and D matrix operands
|
||||
typename ElementC,
|
||||
/// Layout type for C and D matrix operands
|
||||
typename LayoutC,
|
||||
/// Fill Mode for C (kLower or kUpper)
|
||||
FillMode FillModeC,
|
||||
/// Element type for internal accumulation
|
||||
typename ElementAccumulator,
|
||||
/// Operator class tag
|
||||
typename OperatorClass,
|
||||
/// Tag indicating architecture to tune for
|
||||
typename ArchTag,
|
||||
/// Threadblock-level tile size (concept: GemmShape)
|
||||
typename ThreadblockShape,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename WarpShape,
|
||||
/// Warp-level tile size (concept: GemmShape)
|
||||
typename InstructionShape,
|
||||
/// Epilogue output operator
|
||||
typename EpilogueOutputOp,
|
||||
/// Threadblock-level swizzling operator
|
||||
typename ThreadblockSwizzle,
|
||||
/// Number of stages used in the pipelined mainloop
|
||||
int Stages,
|
||||
/// Operation performed by GEMM
|
||||
typename Operator,
|
||||
/// Blas3 computation mode
|
||||
BlasMode BlasMode_,
|
||||
/// Whether the schedule of problems to visit has been precomputed
|
||||
GroupScheduleMode GroupScheduleMode_
|
||||
>
|
||||
struct DefaultRank2KGrouped<ElementA, LayoutA, TransformA, kAlignmentA,
|
||||
ElementB, LayoutB, TransformB, kAlignmentB,
|
||||
ElementC, LayoutC,
|
||||
FillModeC, ElementAccumulator, OperatorClass, ArchTag, ThreadblockShape,
|
||||
WarpShape, InstructionShape, EpilogueOutputOp,
|
||||
ThreadblockSwizzle, Stages, Operator, BlasMode_, GroupScheduleMode_,
|
||||
typename std::enable_if<cutlass::is_complex<ElementAccumulator>::value>::type
|
||||
> {
|
||||
// If true, we must construct a 'transposed-and-exchanged' Rank2K operator.
|
||||
static bool const kInternalTranspose = platform::is_same<LayoutC, layout::ColumnMajor>::value;
|
||||
|
||||
using MapArguments = kernel::detail::Rank2KMapArguments<
|
||||
ElementA,
|
||||
LayoutA,
|
||||
TransformA,
|
||||
kAlignmentA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
TransformB,
|
||||
kAlignmentB,
|
||||
LayoutC,
|
||||
FillModeC,
|
||||
kInternalTranspose
|
||||
>;
|
||||
|
||||
// Define the default grouped Rank2K kernel
|
||||
using DefaultRank2Kkernel = typename kernel::DefaultRank2KComplex<
|
||||
typename MapArguments::ElementA,
|
||||
typename MapArguments::LayoutA,
|
||||
typename MapArguments::ElementB,
|
||||
typename MapArguments::LayoutB,
|
||||
ElementC,
|
||||
typename MapArguments::LayoutC,
|
||||
MapArguments::kFillModeC,
|
||||
ElementAccumulator,
|
||||
OperatorClass,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
Stages,
|
||||
MapArguments::kTransformA,
|
||||
MapArguments::kTransformB,
|
||||
Operator,
|
||||
false, // SplitKSerial
|
||||
BlasMode_
|
||||
>::Rank2Kkernel;
|
||||
|
||||
/// Define the kernel in terms of the default kernel
|
||||
/// Pass through the user-provided TransformA and TransformB so as to
|
||||
/// correctly set public-facing TransformA and TransformB in kernel::Rank2KGrouped.
|
||||
/// This is needed because kernel::DefaultRank2KComplex may change TransformA and
|
||||
/// TransformB that become template arguments to Mma1 and Mma2.
|
||||
using Rank2Kkernel = kernel::Rank2KGrouped<
|
||||
typename DefaultRank2Kkernel::Mma1,
|
||||
typename DefaultRank2Kkernel::Mma2,
|
||||
typename DefaultRank2Kkernel::Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
TransformA,
|
||||
TransformB,
|
||||
DefaultRank2Kkernel::kFillModeC,
|
||||
DefaultRank2Kkernel::kBlasMode,
|
||||
GroupScheduleMode_,
|
||||
kInternalTranspose
|
||||
>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace kernel
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -30,7 +30,7 @@
|
||||
**************************************************************************************************/
|
||||
|
||||
/*! \file
|
||||
\brief
|
||||
\brief Problem visitor for grouped GEMMs
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
@@ -45,6 +45,7 @@
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/trace.h"
|
||||
#include "cutlass/gemm/kernel/gemm_transpose_operands.h"
|
||||
#include "cutlass/gemm/kernel/gemm_grouped_problem_visitor.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -54,168 +55,11 @@ namespace kernel {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Visitor class to abstract away the algorithm for iterating over tiles
|
||||
template <bool Transposed = false>
|
||||
struct GemmGroupedProblemVisitor {
|
||||
|
||||
static bool const kTransposed = Transposed;
|
||||
|
||||
struct Params {
|
||||
cutlass::gemm::GemmCoord const *problem_sizes;
|
||||
int32_t problem_count;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
/// Ctor
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(): problem_sizes(nullptr), problem_count(0) { }
|
||||
|
||||
/// Ctor
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(
|
||||
cutlass::gemm::GemmCoord const *problem_sizes,
|
||||
int32_t problem_count
|
||||
):
|
||||
problem_sizes(problem_sizes),
|
||||
problem_count(problem_count)
|
||||
{}
|
||||
|
||||
};
|
||||
|
||||
struct SharedStorage {
|
||||
//
|
||||
// Nothing for now. As an optimization step, we could consider parallel
|
||||
// argmin or prefix sums across the block.
|
||||
//
|
||||
};
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
Params const ¶ms;
|
||||
SharedStorage &shared_storage;
|
||||
cutlass::MatrixCoord threadblock_shape;
|
||||
|
||||
int64_t tile_idx;
|
||||
int64_t tile_count_sum;
|
||||
int64_t problem_tile_start;
|
||||
int32_t problem_idx;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
CUTLASS_DEVICE
|
||||
GemmGroupedProblemVisitor(
|
||||
Params const ¶ms_,
|
||||
SharedStorage &shared_storage_,
|
||||
cutlass::MatrixCoord threadblock_shape_,
|
||||
int32_t block_idx
|
||||
):
|
||||
shared_storage(shared_storage_),
|
||||
params(params_),
|
||||
threadblock_shape(threadblock_shape_),
|
||||
tile_idx(block_idx),
|
||||
tile_count_sum(0),
|
||||
problem_idx(0)
|
||||
{
|
||||
|
||||
cutlass::gemm::GemmCoord problem = problem_size();
|
||||
cutlass::gemm::GemmCoord grid = grid_shape(problem);
|
||||
|
||||
problem_tile_start = 0;
|
||||
tile_count_sum = grid.m() * grid.n();
|
||||
}
|
||||
|
||||
/// Get the grid shape
|
||||
CUTLASS_HOST_DEVICE
|
||||
static cutlass::gemm::GemmCoord grid_shape(
|
||||
cutlass::gemm::GemmCoord problem,
|
||||
cutlass::MatrixCoord const & block_shape) {
|
||||
|
||||
return cutlass::gemm::GemmCoord(
|
||||
((problem.m() - 1 + block_shape.row()) / block_shape.row()),
|
||||
((problem.n() - 1 + block_shape.column()) / block_shape.column()),
|
||||
1);
|
||||
}
|
||||
|
||||
/// Get the grid shape
|
||||
CUTLASS_DEVICE
|
||||
cutlass::gemm::GemmCoord grid_shape(cutlass::gemm::GemmCoord const &problem) const {
|
||||
return grid_shape(problem, threadblock_shape);
|
||||
}
|
||||
|
||||
/// Returns true if there is a tile to compute
|
||||
CUTLASS_DEVICE
|
||||
bool next_tile() {
|
||||
|
||||
if (tile_idx < tile_count_sum) {
|
||||
return true;
|
||||
}
|
||||
|
||||
do {
|
||||
++problem_idx;
|
||||
|
||||
if (problem_idx >= params.problem_count) {
|
||||
return false;
|
||||
}
|
||||
|
||||
cutlass::gemm::GemmCoord problem = problem_size();
|
||||
cutlass::gemm::GemmCoord grid = grid_shape(problem);
|
||||
|
||||
int64_t tile_count = grid.m() * grid.n();
|
||||
|
||||
problem_tile_start = tile_count_sum;
|
||||
tile_count_sum += tile_count;
|
||||
|
||||
} while (tile_count_sum <= tile_idx);
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
/// Gets the global tile index
|
||||
CUTLASS_HOST_DEVICE
|
||||
int64_t tile_index() const {
|
||||
return tile_idx;
|
||||
}
|
||||
|
||||
/// Gets the index of the problem
|
||||
CUTLASS_HOST_DEVICE
|
||||
int32_t problem_index() const {
|
||||
return problem_idx;
|
||||
}
|
||||
|
||||
/// Returns the problem size for the current problem
|
||||
CUTLASS_HOST_DEVICE
|
||||
cutlass::gemm::GemmCoord problem_size() const {
|
||||
GemmCoord problem = params.problem_sizes[problem_idx];
|
||||
|
||||
if (kTransposed) {
|
||||
swap(problem.m(), problem.n());
|
||||
}
|
||||
|
||||
return problem;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
int64_t threadblock_index() const {
|
||||
return tile_idx - problem_tile_start;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void advance(int32_t grid_size) {
|
||||
tile_idx += grid_size;
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename Mma_, ///! Threadblock-scoped matrix multiply-accumulate
|
||||
typename Epilogue_, ///! Epilogue
|
||||
typename ThreadblockSwizzle_, ///! Threadblock swizzling function
|
||||
typename Mma_, ///! Threadblock-scoped matrix multiply-accumulate
|
||||
typename Epilogue_, ///! Epilogue
|
||||
typename ThreadblockSwizzle_, ///! Threadblock swizzling function
|
||||
GroupScheduleMode GroupScheduleMode_, ///! Type of scheduling to perform
|
||||
bool Transposed = false
|
||||
>
|
||||
struct GemmGrouped {
|
||||
@@ -225,6 +69,7 @@ public:
|
||||
using Epilogue = Epilogue_;
|
||||
using EpilogueOutputOp = typename Epilogue::OutputOp;
|
||||
using ThreadblockSwizzle = ThreadblockSwizzle_;
|
||||
static GroupScheduleMode const kGroupScheduleMode = GroupScheduleMode_;
|
||||
static bool const kTransposed = Transposed;
|
||||
|
||||
// Optional transpose
|
||||
@@ -270,6 +115,13 @@ public:
|
||||
using WarpCount = typename Mma::WarpCount;
|
||||
static int const kThreadCount = 32 * WarpCount::kCount;
|
||||
|
||||
using ProblemVisitor = GemmGroupedProblemVisitor<
|
||||
ThreadblockShape,
|
||||
kGroupScheduleMode,
|
||||
kThreadCount,
|
||||
kThreadCount,
|
||||
kTransposed>;
|
||||
|
||||
//
|
||||
// Structures
|
||||
//
|
||||
@@ -290,13 +142,16 @@ public:
|
||||
ElementA ** ptr_A;
|
||||
ElementB ** ptr_B;
|
||||
ElementC ** ptr_C;
|
||||
ElementC ** ptr_D;
|
||||
ElementC ** ptr_D;
|
||||
|
||||
typename LayoutA::Stride::LongIndex *lda;
|
||||
typename LayoutB::Stride::LongIndex *ldb;
|
||||
typename LayoutC::Stride::LongIndex *ldc;
|
||||
typename LayoutC::Stride::LongIndex *ldd;
|
||||
|
||||
// Only used by device-level operator
|
||||
GemmCoord *host_problem_sizes;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
@@ -304,7 +159,7 @@ public:
|
||||
/// Default ctor
|
||||
CUTLASS_HOST_DEVICE
|
||||
Arguments():
|
||||
problem_count(0),
|
||||
problem_count(0),
|
||||
threadblock_count(0),
|
||||
ptr_A(nullptr),
|
||||
ptr_B(nullptr),
|
||||
@@ -313,7 +168,8 @@ public:
|
||||
lda(nullptr),
|
||||
ldb(nullptr),
|
||||
ldc(nullptr),
|
||||
ldd(nullptr)
|
||||
ldd(nullptr),
|
||||
host_problem_sizes(nullptr)
|
||||
{
|
||||
|
||||
}
|
||||
@@ -328,11 +184,12 @@ public:
|
||||
ElementA ** ptr_A,
|
||||
ElementB ** ptr_B,
|
||||
ElementC ** ptr_C,
|
||||
ElementC ** ptr_D,
|
||||
ElementC ** ptr_D,
|
||||
typename LayoutA::Stride::LongIndex *lda,
|
||||
typename LayoutB::Stride::LongIndex *ldb,
|
||||
typename LayoutC::Stride::LongIndex *ldc,
|
||||
typename LayoutC::Stride::LongIndex *ldd
|
||||
typename LayoutC::Stride::LongIndex *ldd,
|
||||
GemmCoord *host_problem_sizes=nullptr
|
||||
):
|
||||
problem_sizes(problem_sizes),
|
||||
problem_count(problem_count),
|
||||
@@ -345,7 +202,8 @@ public:
|
||||
lda(lda),
|
||||
ldb(ldb),
|
||||
ldc(ldc),
|
||||
ldd(ldd)
|
||||
ldd(ldd),
|
||||
host_problem_sizes(host_problem_sizes)
|
||||
{
|
||||
|
||||
}
|
||||
@@ -358,7 +216,7 @@ public:
|
||||
/// Parameters structure
|
||||
struct Params {
|
||||
|
||||
typename GemmGroupedProblemVisitor<kTransposed>::Params problem_visitor;
|
||||
typename ProblemVisitor::Params problem_visitor;
|
||||
int threadblock_count;
|
||||
|
||||
typename EpilogueOutputOp::Params output_op;
|
||||
@@ -373,7 +231,6 @@ public:
|
||||
typename LayoutC::Stride::LongIndex *ldc;
|
||||
typename LayoutC::Stride::LongIndex *ldd;
|
||||
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
@@ -391,8 +248,10 @@ public:
|
||||
{ }
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(Arguments const &args, void *workspace = nullptr):
|
||||
problem_visitor(args.problem_sizes, args.problem_count),
|
||||
Params(Arguments const &args,
|
||||
void *workspace = nullptr,
|
||||
int tile_count = 0):
|
||||
problem_visitor(args.problem_sizes, args.problem_count, workspace, tile_count),
|
||||
threadblock_count(args.threadblock_count),
|
||||
output_op(args.output_op),
|
||||
ptr_A(args.ptr_A),
|
||||
@@ -410,9 +269,11 @@ public:
|
||||
CUTLASS_HOST_DEVICE
|
||||
void update(
|
||||
Arguments const &args,
|
||||
void *workspace = nullptr) {
|
||||
void *workspace = nullptr,
|
||||
int tile_count = 0) {
|
||||
|
||||
problem_visitor = typename GemmGroupedProblemVisitor<kTransposed>::Params(args.problem_sizes, args.problem_count);
|
||||
problem_visitor = typename ProblemVisitor::Params(args.problem_sizes, args.problem_count,
|
||||
workspace, tile_count);
|
||||
threadblock_count = args.threadblock_count;
|
||||
output_op = args.output_op;
|
||||
ptr_A = args.ptr_A;
|
||||
@@ -427,10 +288,14 @@ public:
|
||||
};
|
||||
|
||||
/// Shared memory storage structure
|
||||
union SharedStorage {
|
||||
typename GemmGroupedProblemVisitor<kTransposed>::SharedStorage problem_visitor;
|
||||
typename Mma::SharedStorage main_loop;
|
||||
typename Epilogue::SharedStorage epilogue;
|
||||
struct SharedStorage {
|
||||
union {
|
||||
typename Mma::SharedStorage main_loop;
|
||||
typename Epilogue::SharedStorage epilogue;
|
||||
} kernel;
|
||||
|
||||
// ProblemVisitor shared storage can't be overlapped with others
|
||||
typename ProblemVisitor::SharedStorage problem_visitor;
|
||||
};
|
||||
|
||||
public:
|
||||
@@ -476,24 +341,23 @@ public:
|
||||
//
|
||||
// Problem visitor.
|
||||
//
|
||||
GemmGroupedProblemVisitor<kTransposed> problem_visitor(
|
||||
params.problem_visitor,
|
||||
shared_storage.problem_visitor,
|
||||
{Mma::Shape::kM, Mma::Shape::kN},
|
||||
ProblemVisitor problem_visitor(
|
||||
params.problem_visitor,
|
||||
shared_storage.problem_visitor,
|
||||
blockIdx.x);
|
||||
|
||||
// Outer 'persistent' loop to iterate over tiles
|
||||
while (problem_visitor.next_tile()) {
|
||||
|
||||
GemmCoord problem_size = problem_visitor.problem_size();
|
||||
int32_t problem_idx = problem_visitor.problem_index();
|
||||
int32_t cta_idx = int32_t(problem_visitor.threadblock_index());
|
||||
GemmCoord problem_size = problem_visitor.problem_size();
|
||||
int32_t problem_idx = problem_visitor.problem_index();
|
||||
int32_t threadblock_idx = int32_t(problem_visitor.threadblock_idx());
|
||||
|
||||
GemmCoord grid_shape = problem_visitor.grid_shape(problem_size);
|
||||
|
||||
cutlass::gemm::GemmCoord threadblock_offset(
|
||||
int(cta_idx / grid_shape.n()) * Mma::Shape::kM,
|
||||
int(cta_idx % grid_shape.n()) * Mma::Shape::kN,
|
||||
int(threadblock_idx / grid_shape.n()) * Mma::Shape::kM,
|
||||
int(threadblock_idx % grid_shape.n()) * Mma::Shape::kN,
|
||||
0);
|
||||
|
||||
// Load element pointers. Exchange pointers and strides if working on the transpose
|
||||
@@ -547,7 +411,7 @@ public:
|
||||
//
|
||||
|
||||
// Construct thread-scoped matrix multiply
|
||||
Mma mma(shared_storage.main_loop, thread_idx, warp_idx, lane_idx);
|
||||
Mma mma(shared_storage.kernel.main_loop, thread_idx, warp_idx, lane_idx);
|
||||
|
||||
// Compute threadblock-scoped matrix multiply-add
|
||||
int gemm_k_iterations = (problem_size.k() + Mma::Shape::kK - 1) / Mma::Shape::kK;
|
||||
@@ -597,7 +461,7 @@ public:
|
||||
);
|
||||
|
||||
Epilogue epilogue(
|
||||
shared_storage.epilogue,
|
||||
shared_storage.kernel.epilogue,
|
||||
thread_idx,
|
||||
warp_idx,
|
||||
lane_idx);
|
||||
|
||||
@@ -0,0 +1,111 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017 - 2022 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 Scheduler for grouped GEMM
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/matrix_coord.h"
|
||||
#include "cutlass/gemm/kernel/grouped_problem_visitor.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace kernel {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace detail {
|
||||
// Helper for correctly representing problem sizes in grouped kernels
|
||||
template <bool Transposed>
|
||||
struct GemmGroupedProblemSizeHelper {
|
||||
|
||||
static bool const kTransposed = Transposed;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static void possibly_transpose_problem(cutlass::gemm::GemmCoord& problem) {
|
||||
if (kTransposed) {
|
||||
swap(problem.m(), problem.n());
|
||||
}
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static int32_t tile_count(const cutlass::gemm::GemmCoord& grid) {
|
||||
return grid.m() * grid.n();
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace detail
|
||||
|
||||
/// Visitor class to abstract away the algorithm for iterating over tiles
|
||||
template <typename ThreadblockShape,
|
||||
GroupScheduleMode GroupScheduleMode_,
|
||||
int PrefetchTileCount,
|
||||
int ThreadCount,
|
||||
bool Transposed = false>
|
||||
struct GemmGroupedProblemVisitor : public GroupedProblemVisitor<
|
||||
detail::GemmGroupedProblemSizeHelper<Transposed>,
|
||||
ThreadblockShape,
|
||||
GroupScheduleMode_,
|
||||
PrefetchTileCount,
|
||||
ThreadCount> {
|
||||
|
||||
static bool const kTransposed = Transposed;
|
||||
|
||||
using ProblemSizeHelper = detail::GemmGroupedProblemSizeHelper<Transposed>;
|
||||
using Base = GroupedProblemVisitor<ProblemSizeHelper, ThreadblockShape, GroupScheduleMode_, PrefetchTileCount, ThreadCount>;
|
||||
using Params = typename Base::Params;
|
||||
using SharedStorage = typename Base::SharedStorage;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
CUTLASS_DEVICE
|
||||
GemmGroupedProblemVisitor(
|
||||
Params const ¶ms_,
|
||||
SharedStorage &shared_storage_,
|
||||
int32_t block_idx
|
||||
): Base (params_, shared_storage_, block_idx)
|
||||
{}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace kernel
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,517 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017 - 2022 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 Problem visitor for grouped GEMMs with a softmax fused beforehand
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/fast_math.h"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/matrix_coord.h"
|
||||
#include "cutlass/complex.h"
|
||||
#include "cutlass/semaphore.h"
|
||||
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/trace.h"
|
||||
#include "cutlass/gemm/kernel/gemm_transpose_operands.h"
|
||||
#include "cutlass/gemm/kernel/gemm_grouped_problem_visitor.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace kernel {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename Mma_, ///! Threadblock-scoped matrix multiply-accumulate
|
||||
typename Epilogue_, ///! Epilogue
|
||||
typename ThreadblockSwizzle_, ///! Threadblock swizzling function
|
||||
GroupScheduleMode GroupScheduleMode_, ///! Type of scheduling to perform
|
||||
bool Transposed = false
|
||||
>
|
||||
struct GemmGroupedSoftmaxMainloopFusion {
|
||||
public:
|
||||
|
||||
using Mma = Mma_;
|
||||
using Epilogue = Epilogue_;
|
||||
using EpilogueOutputOp = typename Epilogue::OutputOp;
|
||||
using ThreadblockSwizzle = ThreadblockSwizzle_;
|
||||
static GroupScheduleMode const kGroupScheduleMode = GroupScheduleMode_;
|
||||
static bool const kTransposed = Transposed;
|
||||
|
||||
// Optional transpose
|
||||
using MapArguments = kernel::detail::MapArguments<
|
||||
typename Mma::IteratorA::Element,
|
||||
typename Mma::IteratorA::Layout,
|
||||
Mma::kTransformA,
|
||||
Mma::IteratorA::AccessType::kElements,
|
||||
typename Mma::IteratorB::Element,
|
||||
typename Mma::IteratorB::Layout,
|
||||
Mma::kTransformB,
|
||||
Mma::IteratorB::AccessType::kElements,
|
||||
typename Mma::LayoutC,
|
||||
kTransposed
|
||||
>;
|
||||
|
||||
// Public-facing type definitions related to operand element type, layout, and complex conjugate
|
||||
// operation. Must interact with the 'kTransposed' notion.
|
||||
using ElementA = typename MapArguments::ElementA;
|
||||
using LayoutA = typename MapArguments::LayoutA;
|
||||
using ElementB = typename MapArguments::ElementB;
|
||||
using LayoutB = typename MapArguments::LayoutB;
|
||||
using ElementC = typename Epilogue::OutputTileIterator::Element;
|
||||
using LayoutC = typename MapArguments::LayoutC;
|
||||
|
||||
using ElementScaleBias = typename Mma::IteratorNormSum::Element;
|
||||
|
||||
static ComplexTransform const kTransformA = MapArguments::kTransformA;
|
||||
static ComplexTransform const kTransformB = MapArguments::kTransformB;
|
||||
|
||||
// Type definitions about the mainloop.
|
||||
using Operator = typename Mma::Operator;
|
||||
using OperatorClass = typename Mma::Operator::OperatorClass;
|
||||
using ThreadblockShape = typename Mma::Shape;
|
||||
using WarpShape = typename Mma::Operator::Shape;
|
||||
using InstructionShape = typename Mma::Policy::Operator::InstructionShape;
|
||||
using ArchTag = typename Mma::ArchTag;
|
||||
|
||||
static int const kStages = Mma::kStages;
|
||||
static int const kAlignmentA = MapArguments::kAlignmentA;
|
||||
static int const kAlignmentB = MapArguments::kAlignmentB;
|
||||
static int const kAlignmentC = Epilogue::OutputTileIterator::kElementsPerAccess;
|
||||
|
||||
/// Warp count (concept: GemmShape)
|
||||
using WarpCount = typename Mma::WarpCount;
|
||||
static int const kThreadCount = 32 * WarpCount::kCount;
|
||||
|
||||
using ProblemVisitor = GemmGroupedProblemVisitor<
|
||||
ThreadblockShape,
|
||||
kGroupScheduleMode,
|
||||
kThreadCount,
|
||||
kThreadCount,
|
||||
kTransposed>;
|
||||
|
||||
//
|
||||
// Structures
|
||||
//
|
||||
|
||||
/// Argument structure
|
||||
struct Arguments {
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
GemmCoord *problem_sizes;
|
||||
int problem_count;
|
||||
int threadblock_count;
|
||||
|
||||
typename EpilogueOutputOp::Params output_op;
|
||||
|
||||
ElementA ** ptr_A;
|
||||
ElementB ** ptr_B;
|
||||
ElementC ** ptr_C;
|
||||
ElementC ** ptr_D;
|
||||
void ** ptr_norm;
|
||||
void ** ptr_sum;
|
||||
|
||||
typename LayoutA::Stride::LongIndex *lda;
|
||||
typename LayoutB::Stride::LongIndex *ldb;
|
||||
typename LayoutC::Stride::LongIndex *ldc;
|
||||
typename LayoutC::Stride::LongIndex *ldd;
|
||||
|
||||
// Only used by device-level operator
|
||||
GemmCoord *host_problem_sizes;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
/// Default ctor
|
||||
CUTLASS_HOST_DEVICE
|
||||
Arguments():
|
||||
problem_count(0),
|
||||
threadblock_count(0),
|
||||
ptr_A(nullptr),
|
||||
ptr_B(nullptr),
|
||||
ptr_C(nullptr),
|
||||
ptr_D(nullptr),
|
||||
ptr_norm(nullptr),
|
||||
ptr_sum(nullptr),
|
||||
lda(nullptr),
|
||||
ldb(nullptr),
|
||||
ldc(nullptr),
|
||||
ldd(nullptr),
|
||||
host_problem_sizes(nullptr)
|
||||
{
|
||||
|
||||
}
|
||||
|
||||
/// Ctor
|
||||
CUTLASS_HOST_DEVICE
|
||||
Arguments(
|
||||
GemmCoord *problem_sizes,
|
||||
int problem_count,
|
||||
int threadblock_count,
|
||||
typename EpilogueOutputOp::Params output_op,
|
||||
ElementA ** ptr_A,
|
||||
ElementB ** ptr_B,
|
||||
ElementC ** ptr_C,
|
||||
ElementC ** ptr_D,
|
||||
void ** ptr_norm,
|
||||
void ** ptr_sum,
|
||||
typename LayoutA::Stride::LongIndex *lda,
|
||||
typename LayoutB::Stride::LongIndex *ldb,
|
||||
typename LayoutC::Stride::LongIndex *ldc,
|
||||
typename LayoutC::Stride::LongIndex *ldd,
|
||||
GemmCoord *host_problem_sizes=nullptr
|
||||
):
|
||||
problem_sizes(problem_sizes),
|
||||
problem_count(problem_count),
|
||||
threadblock_count(threadblock_count),
|
||||
output_op(output_op),
|
||||
ptr_A(ptr_A),
|
||||
ptr_B(ptr_B),
|
||||
ptr_C(ptr_C),
|
||||
ptr_D(ptr_D),
|
||||
ptr_norm(ptr_norm),
|
||||
ptr_sum(ptr_sum),
|
||||
lda(lda),
|
||||
ldb(ldb),
|
||||
ldc(ldc),
|
||||
ldd(ldd),
|
||||
host_problem_sizes(host_problem_sizes)
|
||||
{
|
||||
|
||||
}
|
||||
};
|
||||
|
||||
//
|
||||
// Structure for precomputing values in host memory and passing to kernels
|
||||
//
|
||||
|
||||
/// Parameters structure
|
||||
struct Params {
|
||||
|
||||
typename ProblemVisitor::Params problem_visitor;
|
||||
int threadblock_count;
|
||||
|
||||
typename EpilogueOutputOp::Params output_op;
|
||||
|
||||
ElementA ** ptr_A;
|
||||
ElementB ** ptr_B;
|
||||
ElementC ** ptr_C;
|
||||
ElementC ** ptr_D;
|
||||
|
||||
void ** ptr_norm;
|
||||
void ** ptr_sum;
|
||||
|
||||
typename LayoutA::Stride::LongIndex *lda;
|
||||
typename LayoutB::Stride::LongIndex *ldb;
|
||||
typename LayoutC::Stride::LongIndex *ldc;
|
||||
typename LayoutC::Stride::LongIndex *ldd;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params():
|
||||
ptr_A(nullptr),
|
||||
ptr_B(nullptr),
|
||||
ptr_C(nullptr),
|
||||
ptr_D(nullptr),
|
||||
ptr_norm(nullptr),
|
||||
ptr_sum(nullptr),
|
||||
lda(nullptr),
|
||||
ldb(nullptr),
|
||||
ldc(nullptr),
|
||||
ldd(nullptr)
|
||||
{ }
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(Arguments const &args,
|
||||
void *workspace = nullptr,
|
||||
int tile_count = 0):
|
||||
problem_visitor(args.problem_sizes, args.problem_count, workspace, tile_count),
|
||||
threadblock_count(args.threadblock_count),
|
||||
output_op(args.output_op),
|
||||
ptr_A(args.ptr_A),
|
||||
ptr_B(args.ptr_B),
|
||||
ptr_C(args.ptr_C),
|
||||
ptr_D(args.ptr_D),
|
||||
ptr_norm(args.ptr_norm),
|
||||
ptr_sum(args.ptr_sum),
|
||||
lda(args.lda),
|
||||
ldb(args.ldb),
|
||||
ldc(args.ldc),
|
||||
ldd(args.ldd)
|
||||
{
|
||||
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
void update(
|
||||
Arguments const &args,
|
||||
void *workspace = nullptr,
|
||||
int tile_count = 0) {
|
||||
|
||||
problem_visitor = typename ProblemVisitor::Params(args.problem_sizes, args.problem_count,
|
||||
workspace, tile_count);
|
||||
threadblock_count = args.threadblock_count;
|
||||
output_op = args.output_op;
|
||||
ptr_A = args.ptr_A;
|
||||
ptr_B = args.ptr_B;
|
||||
ptr_C = args.ptr_C;
|
||||
ptr_D = args.ptr_D;
|
||||
ptr_norm = args.ptr_norm;
|
||||
ptr_sum = args.ptr_sum;
|
||||
lda = args.lda;
|
||||
ldb = args.ldb;
|
||||
ldc = args.ldc;
|
||||
ldd = args.ldd;
|
||||
}
|
||||
};
|
||||
|
||||
/// Shared memory storage structure
|
||||
struct SharedStorage {
|
||||
union {
|
||||
typename Mma::SharedStorage main_loop;
|
||||
typename Epilogue::SharedStorage epilogue;
|
||||
} kernel;
|
||||
|
||||
// ProblemVisitor shared storage can't be overlapped with others
|
||||
typename ProblemVisitor::SharedStorage problem_visitor;
|
||||
};
|
||||
|
||||
public:
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
CUTLASS_DEVICE
|
||||
GemmGroupedSoftmaxMainloopFusion() { }
|
||||
|
||||
/// Determines whether kernel satisfies alignment
|
||||
static Status can_implement(cutlass::gemm::GemmCoord const & problem_size) {
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
static Status can_implement(Arguments const &args) {
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
static size_t get_extra_workspace_size(
|
||||
Arguments const &args,
|
||||
cutlass::gemm::GemmCoord const &grid_tiled_shape) {
|
||||
|
||||
return 0;
|
||||
}
|
||||
|
||||
/// Executes one GEMM
|
||||
CUTLASS_DEVICE
|
||||
void operator()(Params const ¶ms, SharedStorage &shared_storage) {
|
||||
|
||||
//
|
||||
// These types shadow the type-level definitions and support the ability to implement
|
||||
// a 'transposed' GEMM that computes the transposed problems.
|
||||
//
|
||||
using ElementA = typename Mma::IteratorA::Element;
|
||||
using LayoutA = typename Mma::IteratorA::Layout;
|
||||
using ElementB = typename Mma::IteratorB::Element;
|
||||
using LayoutB = typename Mma::IteratorB::Layout;
|
||||
using ElementC = typename Epilogue::OutputTileIterator::Element;
|
||||
using LayoutC = typename Epilogue::OutputTileIterator::Layout;
|
||||
|
||||
//
|
||||
// Problem visitor.
|
||||
//
|
||||
ProblemVisitor problem_visitor(
|
||||
params.problem_visitor,
|
||||
shared_storage.problem_visitor,
|
||||
blockIdx.x);
|
||||
|
||||
// Outer 'persistent' loop to iterate over tiles
|
||||
while (problem_visitor.next_tile()) {
|
||||
|
||||
GemmCoord problem_size = problem_visitor.problem_size();
|
||||
int32_t problem_idx = problem_visitor.problem_index();
|
||||
int32_t threadblock_idx = int32_t(problem_visitor.threadblock_idx());
|
||||
|
||||
GemmCoord grid_shape = problem_visitor.grid_shape(problem_size);
|
||||
|
||||
cutlass::gemm::GemmCoord threadblock_offset(
|
||||
int(threadblock_idx / grid_shape.n()) * Mma::Shape::kM,
|
||||
int(threadblock_idx % grid_shape.n()) * Mma::Shape::kN,
|
||||
0);
|
||||
|
||||
// Load element pointers. Exchange pointers and strides if working on the transpose
|
||||
ElementA *ptr_A = reinterpret_cast<ElementA *>((kTransposed ? params.ptr_B[problem_idx] : params.ptr_A[problem_idx]));
|
||||
typename LayoutA::LongIndex ldm_A = (kTransposed ? params.ldb[problem_idx] : params.lda[problem_idx]);
|
||||
|
||||
ElementB *ptr_B = reinterpret_cast<ElementB *>((kTransposed ? params.ptr_A[problem_idx] : params.ptr_B[problem_idx]));
|
||||
typename LayoutB::LongIndex ldm_B = (kTransposed ? params.lda[problem_idx] : params.ldb[problem_idx]);
|
||||
|
||||
// Compute initial location in logical coordinates
|
||||
cutlass::MatrixCoord tb_offset_A{
|
||||
threadblock_offset.m(),
|
||||
0,
|
||||
};
|
||||
|
||||
cutlass::MatrixCoord tb_offset_B{
|
||||
0,
|
||||
threadblock_offset.n()
|
||||
};
|
||||
|
||||
// Compute position within threadblock
|
||||
int thread_idx = threadIdx.x;
|
||||
|
||||
// Construct iterators to A and B operands
|
||||
typename Mma::IteratorA iterator_A(
|
||||
LayoutA(ldm_A),
|
||||
ptr_A,
|
||||
{problem_size.m(), problem_size.k()},
|
||||
thread_idx,
|
||||
tb_offset_A);
|
||||
|
||||
typename Mma::IteratorB iterator_B(
|
||||
LayoutB(ldm_B),
|
||||
ptr_B,
|
||||
{problem_size.k(), problem_size.n()},
|
||||
thread_idx,
|
||||
tb_offset_B);
|
||||
|
||||
// Construct iterator to the softmax norm/sum vector
|
||||
typename Mma::IteratorNormSum iterator_norm_sum(
|
||||
problem_size.m(),
|
||||
static_cast<ElementScaleBias const *>(params.ptr_norm[problem_idx]),
|
||||
static_cast<ElementScaleBias const *>(params.ptr_sum[problem_idx]),
|
||||
thread_idx,
|
||||
MatrixCoord(0, threadblock_offset.m())
|
||||
);
|
||||
|
||||
typename Mma::FragmentC accumulators;
|
||||
|
||||
accumulators.clear();
|
||||
|
||||
// Broadcast the warp_id computed by lane 0 to ensure dependent code
|
||||
// is compiled as warp-uniform.
|
||||
int warp_idx = __shfl_sync(0xffffffff, threadIdx.x / 32, 0);
|
||||
|
||||
int lane_idx = threadIdx.x % 32;
|
||||
|
||||
//
|
||||
// Matrix multiply phase
|
||||
//
|
||||
|
||||
// Construct thread-scoped matrix multiply
|
||||
Mma mma(shared_storage.kernel.main_loop, thread_idx, warp_idx, lane_idx);
|
||||
|
||||
// Compute threadblock-scoped matrix multiply-add
|
||||
int gemm_k_iterations = (problem_size.k() + Mma::Shape::kK - 1) / Mma::Shape::kK;
|
||||
|
||||
// Wait for all threads to finish their epilogue phases from the previous tile.
|
||||
__syncthreads();
|
||||
|
||||
// Compute threadblock-scoped matrix multiply-add
|
||||
mma(
|
||||
gemm_k_iterations,
|
||||
accumulators,
|
||||
iterator_A,
|
||||
iterator_B,
|
||||
iterator_norm_sum,
|
||||
accumulators);
|
||||
|
||||
//
|
||||
// Epilogue
|
||||
//
|
||||
|
||||
EpilogueOutputOp output_op(params.output_op);
|
||||
|
||||
ElementC *ptr_C = params.ptr_C[problem_idx];
|
||||
ElementC *ptr_D = params.ptr_D[problem_idx];
|
||||
|
||||
LayoutC layout_C(params.ldc[problem_idx]);
|
||||
LayoutC layout_D(params.ldd[problem_idx]);
|
||||
|
||||
typename Epilogue::OutputTileIterator::Params params_C(layout_C);
|
||||
typename Epilogue::OutputTileIterator::Params params_D(layout_D);
|
||||
|
||||
// Tile iterator loading from source tensor.
|
||||
typename Epilogue::OutputTileIterator iterator_C(
|
||||
params_C,
|
||||
ptr_C,
|
||||
problem_size.mn(),
|
||||
thread_idx,
|
||||
threadblock_offset.mn()
|
||||
);
|
||||
|
||||
// Tile iterator writing to destination tensor.
|
||||
typename Epilogue::OutputTileIterator iterator_D(
|
||||
params_D,
|
||||
ptr_D,
|
||||
problem_size.mn(),
|
||||
thread_idx,
|
||||
threadblock_offset.mn()
|
||||
);
|
||||
|
||||
Epilogue epilogue(
|
||||
shared_storage.kernel.epilogue,
|
||||
thread_idx,
|
||||
warp_idx,
|
||||
lane_idx);
|
||||
|
||||
// Execute the epilogue operator to update the destination tensor.
|
||||
epilogue(
|
||||
output_op,
|
||||
iterator_D,
|
||||
accumulators,
|
||||
iterator_C);
|
||||
|
||||
// Next tile
|
||||
problem_visitor.advance(gridDim.x);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace kernel
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,818 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017 - 2022 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 Template for a multistage GEMM kernel with layernorm operations fused in mainloop.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/fast_math.h"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/matrix_coord.h"
|
||||
#include "cutlass/complex.h"
|
||||
#include "cutlass/semaphore.h"
|
||||
|
||||
#include "cutlass/layout/matrix.h"
|
||||
|
||||
#include "cutlass/trace.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace kernel {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename Mma_, ///! Threadblock-scoped matrix multiply-accumulate
|
||||
typename Epilogue_, ///! Epilogue
|
||||
typename ThreadblockSwizzle_ ///! Threadblock swizzling function
|
||||
>
|
||||
struct GemmLayernormMainloopFusion {
|
||||
public:
|
||||
|
||||
using Mma = Mma_;
|
||||
using Epilogue = Epilogue_;
|
||||
using EpilogueOutputOp = typename Epilogue::OutputOp;
|
||||
using ThreadblockSwizzle = ThreadblockSwizzle_;
|
||||
|
||||
using ElementA = typename Mma::IteratorA::Element;
|
||||
using LayoutA = typename Mma::IteratorA::Layout;
|
||||
using ElementB = typename Mma::IteratorB::Element;
|
||||
using LayoutB = typename Mma::IteratorB::Layout;
|
||||
using ElementC = typename Epilogue::OutputTileIterator::Element;
|
||||
using LayoutC = typename Epilogue::OutputTileIterator::Layout;
|
||||
|
||||
using ElementScaleBias = typename Mma::IteratorVarMean::Element;
|
||||
using LayoutScaleBias = typename Mma::IteratorVarMean::Layout;
|
||||
|
||||
static ComplexTransform const kTransformA = Mma::kTransformA;
|
||||
static ComplexTransform const kTransformB = Mma::kTransformB;
|
||||
using Operator = typename Mma::Operator;
|
||||
|
||||
using OperatorClass = typename Mma::Operator::OperatorClass;
|
||||
using ThreadblockShape = typename Mma::Shape;
|
||||
using WarpShape = typename Mma::Operator::Shape;
|
||||
using InstructionShape = typename Mma::Policy::Operator::InstructionShape;
|
||||
using ArchTag = typename Mma::ArchTag;
|
||||
|
||||
static int const kStages = Mma::kStages;
|
||||
static int const kAlignmentA = Mma::IteratorA::AccessType::kElements;
|
||||
static int const kAlignmentB = Mma::IteratorB::AccessType::kElements;
|
||||
static int const kAlignmentC = Epilogue::OutputTileIterator::kElementsPerAccess;
|
||||
|
||||
/// Warp count (concept: GemmShape)
|
||||
using WarpCount = typename Mma::WarpCount;
|
||||
static int const kThreadCount = 32 * WarpCount::kCount;
|
||||
|
||||
/// Split-K preserves splits that are 128b aligned
|
||||
static int const kSplitKAlignment = const_max(128 / sizeof_bits<ElementA>::value, 128 / sizeof_bits<ElementB>::value);
|
||||
|
||||
//
|
||||
// Structures
|
||||
//
|
||||
|
||||
/// Argument structure
|
||||
struct Arguments {
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
GemmUniversalMode mode;
|
||||
GemmCoord problem_size;
|
||||
int batch_count;
|
||||
|
||||
typename EpilogueOutputOp::Params epilogue;
|
||||
|
||||
void const * ptr_A;
|
||||
void const * ptr_B;
|
||||
void const * ptr_var;
|
||||
void const * ptr_mean;
|
||||
void const * ptr_gamma;
|
||||
void const * ptr_beta;
|
||||
void const * ptr_C;
|
||||
void * ptr_D;
|
||||
|
||||
int64_t batch_stride_A;
|
||||
int64_t batch_stride_B;
|
||||
int64_t batch_stride_var;
|
||||
int64_t batch_stride_mean;
|
||||
int64_t batch_stride_gamma;
|
||||
int64_t batch_stride_beta;
|
||||
int64_t batch_stride_C;
|
||||
int64_t batch_stride_D;
|
||||
|
||||
typename LayoutA::Stride stride_a;
|
||||
typename LayoutB::Stride stride_b;
|
||||
typename LayoutScaleBias::Stride stride_var;
|
||||
typename LayoutScaleBias::Stride stride_mean;
|
||||
typename LayoutScaleBias::Stride stride_gamma;
|
||||
typename LayoutScaleBias::Stride stride_beta;
|
||||
typename LayoutC::Stride stride_c;
|
||||
typename LayoutC::Stride stride_d;
|
||||
|
||||
typename LayoutA::Stride::LongIndex lda;
|
||||
typename LayoutB::Stride::LongIndex ldb;
|
||||
typename LayoutScaleBias::Stride::LongIndex ld_var;
|
||||
typename LayoutScaleBias::Stride::LongIndex ld_mean;
|
||||
typename LayoutScaleBias::Stride::LongIndex ld_gamma;
|
||||
typename LayoutScaleBias::Stride::LongIndex ld_beta;
|
||||
typename LayoutC::Stride::LongIndex ldc;
|
||||
typename LayoutC::Stride::LongIndex ldd;
|
||||
|
||||
int const * ptr_gather_A_indices;
|
||||
int const * ptr_gather_B_indices;
|
||||
int const * ptr_scatter_D_indices;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
Arguments():
|
||||
mode(GemmUniversalMode::kGemm),
|
||||
batch_count(1),
|
||||
ptr_A(nullptr), ptr_B(nullptr), ptr_C(nullptr), ptr_D(nullptr),
|
||||
ptr_var(nullptr), ptr_mean(nullptr),
|
||||
ptr_gamma(nullptr), ptr_beta(nullptr),
|
||||
ptr_gather_A_indices(nullptr),
|
||||
ptr_gather_B_indices(nullptr),
|
||||
ptr_scatter_D_indices(nullptr) {}
|
||||
|
||||
/// constructs an arguments structure
|
||||
Arguments(
|
||||
GemmUniversalMode mode,
|
||||
GemmCoord problem_size,
|
||||
int batch_count,
|
||||
typename EpilogueOutputOp::Params epilogue,
|
||||
void const * ptr_A,
|
||||
void const * ptr_B,
|
||||
void const * ptr_var,
|
||||
void const * ptr_mean,
|
||||
void const * ptr_gamma,
|
||||
void const * ptr_beta,
|
||||
void const * ptr_C,
|
||||
void * ptr_D,
|
||||
int64_t batch_stride_A,
|
||||
int64_t batch_stride_B,
|
||||
int64_t batch_stride_var,
|
||||
int64_t batch_stride_mean,
|
||||
int64_t batch_stride_gamma,
|
||||
int64_t batch_stride_beta,
|
||||
int64_t batch_stride_C,
|
||||
int64_t batch_stride_D,
|
||||
typename LayoutA::Stride stride_a,
|
||||
typename LayoutB::Stride stride_b,
|
||||
typename LayoutScaleBias::Stride stride_var,
|
||||
typename LayoutScaleBias::Stride stride_mean,
|
||||
typename LayoutScaleBias::Stride stride_gamma,
|
||||
typename LayoutScaleBias::Stride stride_beta,
|
||||
typename LayoutC::Stride stride_c,
|
||||
typename LayoutC::Stride stride_d,
|
||||
int const *ptr_gather_A_indices = nullptr,
|
||||
int const *ptr_gather_B_indices = nullptr,
|
||||
int const *ptr_scatter_D_indices = nullptr
|
||||
):
|
||||
mode(mode),
|
||||
problem_size(problem_size),
|
||||
batch_count(batch_count),
|
||||
epilogue(epilogue),
|
||||
ptr_A(ptr_A), ptr_B(ptr_B), ptr_C(ptr_C), ptr_D(ptr_D),
|
||||
ptr_var(ptr_var), ptr_mean(ptr_mean),
|
||||
ptr_gamma(ptr_gamma), ptr_beta(ptr_beta),
|
||||
batch_stride_A(batch_stride_A), batch_stride_B(batch_stride_B), batch_stride_C(batch_stride_C), batch_stride_D(batch_stride_D),
|
||||
stride_a(stride_a), stride_b(stride_b), stride_c(stride_c), stride_d(stride_d),
|
||||
stride_var(stride_var), stride_mean(stride_mean),
|
||||
stride_gamma(stride_gamma), stride_beta(stride_beta),
|
||||
ptr_gather_A_indices(ptr_gather_A_indices), ptr_gather_B_indices(ptr_gather_B_indices),
|
||||
ptr_scatter_D_indices(ptr_scatter_D_indices) {
|
||||
lda = 0;
|
||||
ldb = 0;
|
||||
ldc = 0;
|
||||
ldd = 0;
|
||||
ld_var = 0;
|
||||
ld_mean = 0;
|
||||
ld_gamma = 0;
|
||||
ld_beta = 0;
|
||||
CUTLASS_TRACE_HOST("GemmUniversal::Arguments::Arguments() - problem_size: " << problem_size);
|
||||
}
|
||||
|
||||
/// constructs an arguments structure
|
||||
Arguments(
|
||||
GemmUniversalMode mode,
|
||||
GemmCoord problem_size,
|
||||
int batch_count,
|
||||
typename EpilogueOutputOp::Params epilogue,
|
||||
void const * ptr_A,
|
||||
void const * ptr_B,
|
||||
void const * ptr_var,
|
||||
void const * ptr_mean,
|
||||
void const * ptr_gamma,
|
||||
void const * ptr_beta,
|
||||
void const * ptr_C,
|
||||
void * ptr_D,
|
||||
int64_t batch_stride_A,
|
||||
int64_t batch_stride_B,
|
||||
int64_t batch_stride_var,
|
||||
int64_t batch_stride_mean,
|
||||
int64_t batch_stride_gamma,
|
||||
int64_t batch_stride_beta,
|
||||
int64_t batch_stride_C,
|
||||
int64_t batch_stride_D,
|
||||
typename LayoutA::Stride::LongIndex lda,
|
||||
typename LayoutB::Stride::LongIndex ldb,
|
||||
typename LayoutScaleBias::Stride::LongIndex ld_var,
|
||||
typename LayoutScaleBias::Stride::LongIndex ld_mean,
|
||||
typename LayoutScaleBias::Stride::LongIndex ld_gamma,
|
||||
typename LayoutScaleBias::Stride::LongIndex ld_beta,
|
||||
typename LayoutC::Stride::LongIndex ldc,
|
||||
typename LayoutC::Stride::LongIndex ldd,
|
||||
int const *ptr_gather_A_indices = nullptr,
|
||||
int const *ptr_gather_B_indices = nullptr,
|
||||
int const *ptr_scatter_D_indices = nullptr
|
||||
):
|
||||
mode(mode),
|
||||
problem_size(problem_size),
|
||||
batch_count(batch_count),
|
||||
epilogue(epilogue),
|
||||
ptr_A(ptr_A), ptr_B(ptr_B), ptr_C(ptr_C), ptr_D(ptr_D),
|
||||
ptr_var(ptr_var), ptr_mean(ptr_mean),
|
||||
ptr_gamma(ptr_gamma), ptr_beta(ptr_beta),
|
||||
batch_stride_A(batch_stride_A), batch_stride_B(batch_stride_B), batch_stride_C(batch_stride_C), batch_stride_D(batch_stride_D),
|
||||
batch_stride_var(batch_stride_var), batch_stride_mean(batch_stride_mean),
|
||||
batch_stride_gamma(batch_stride_gamma), batch_stride_beta(batch_stride_beta),
|
||||
lda(lda), ldb(ldb), ldc(ldc), ldd(ldd),
|
||||
ld_var(ld_var), ld_mean(ld_mean),
|
||||
ld_gamma(ld_gamma), ld_beta(ld_beta),
|
||||
ptr_gather_A_indices(ptr_gather_A_indices), ptr_gather_B_indices(ptr_gather_B_indices),
|
||||
ptr_scatter_D_indices(ptr_scatter_D_indices) {
|
||||
stride_a = make_Coord(lda);
|
||||
stride_b = make_Coord(ldb);
|
||||
stride_c = make_Coord(ldc);
|
||||
stride_d = make_Coord(ldd);
|
||||
stride_var = make_Coord(ld_var);
|
||||
stride_mean = make_Coord(ld_mean);
|
||||
stride_gamma = make_Coord(ld_gamma);
|
||||
stride_beta = make_Coord(ld_beta);
|
||||
CUTLASS_TRACE_HOST("GemmUniversal::Arguments::Arguments() - problem_size: " << problem_size);
|
||||
}
|
||||
|
||||
/// Returns arguments for the transposed problem
|
||||
Arguments transposed_problem() const {
|
||||
Arguments args(*this);
|
||||
|
||||
std::swap(args.problem_size.m(), args.problem_size.n());
|
||||
std::swap(args.ptr_A, args.ptr_B);
|
||||
std::swap(args.lda, args.ldb);
|
||||
std::swap(args.stride_a, args.stride_b);
|
||||
std::swap(args.batch_stride_A, args.batch_stride_B);
|
||||
std::swap(args.ptr_gather_A_indices, args.ptr_gather_B_indices);
|
||||
|
||||
return args;
|
||||
}
|
||||
};
|
||||
|
||||
//
|
||||
// Structure for precomputing values in host memory and passing to kernels
|
||||
//
|
||||
|
||||
/// Parameters structure
|
||||
struct Params {
|
||||
|
||||
cutlass::gemm::GemmCoord problem_size;
|
||||
cutlass::gemm::GemmCoord grid_tiled_shape;
|
||||
int swizzle_log_tile;
|
||||
|
||||
typename Mma::IteratorA::Params params_A;
|
||||
typename Mma::IteratorB::Params params_B;
|
||||
typename Epilogue::OutputTileIterator::Params params_C;
|
||||
typename Epilogue::OutputTileIterator::Params params_D;
|
||||
|
||||
typename EpilogueOutputOp::Params output_op;
|
||||
|
||||
GemmUniversalMode mode;
|
||||
int batch_count;
|
||||
int gemm_k_size;
|
||||
|
||||
void * ptr_A;
|
||||
void * ptr_B;
|
||||
void * ptr_var;
|
||||
void * ptr_mean;
|
||||
void * ptr_gamma;
|
||||
void * ptr_beta;
|
||||
void * ptr_C;
|
||||
void * ptr_D;
|
||||
|
||||
int64_t batch_stride_A;
|
||||
int64_t batch_stride_B;
|
||||
int64_t batch_stride_var;
|
||||
int64_t batch_stride_mean;
|
||||
int64_t batch_stride_gamma;
|
||||
int64_t batch_stride_beta;
|
||||
int64_t batch_stride_C;
|
||||
int64_t batch_stride_D;
|
||||
|
||||
int * ptr_gather_A_indices;
|
||||
int * ptr_gather_B_indices;
|
||||
int * ptr_scatter_D_indices;
|
||||
|
||||
int *semaphore;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params():
|
||||
swizzle_log_tile(0),
|
||||
params_A(0),
|
||||
params_B(0),
|
||||
params_C(0),
|
||||
params_D(0),
|
||||
batch_count(0),
|
||||
gemm_k_size(0),
|
||||
mode(cutlass::gemm::GemmUniversalMode::kGemm),
|
||||
ptr_A(nullptr),
|
||||
ptr_B(nullptr),
|
||||
ptr_var(nullptr),
|
||||
ptr_mean(nullptr),
|
||||
ptr_gamma(nullptr),
|
||||
ptr_beta(nullptr),
|
||||
ptr_C(nullptr),
|
||||
ptr_D(nullptr),
|
||||
batch_stride_A(0),
|
||||
batch_stride_B(0),
|
||||
batch_stride_var(0),
|
||||
batch_stride_mean(0),
|
||||
batch_stride_C(0),
|
||||
batch_stride_D(0),
|
||||
ptr_gather_A_indices(nullptr),
|
||||
ptr_gather_B_indices(nullptr),
|
||||
ptr_scatter_D_indices(nullptr),
|
||||
semaphore(nullptr) { }
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(
|
||||
Arguments const &args,
|
||||
cutlass::gemm::GemmCoord const & grid_tiled_shape,
|
||||
int gemm_k_size,
|
||||
void *workspace = nullptr
|
||||
):
|
||||
problem_size(args.problem_size),
|
||||
grid_tiled_shape(grid_tiled_shape),
|
||||
swizzle_log_tile(ThreadblockSwizzle().get_log_tile(grid_tiled_shape)),
|
||||
params_A(args.lda ? make_Coord_with_padding<LayoutA::kStrideRank>(args.lda) : args.stride_a),
|
||||
params_B(args.ldb ? make_Coord_with_padding<LayoutB::kStrideRank>(args.ldb) : args.stride_b),
|
||||
params_C(args.ldc ? make_Coord_with_padding<LayoutC::kStrideRank>(args.ldc) : args.stride_c),
|
||||
params_D(args.ldd ? make_Coord_with_padding<LayoutC::kStrideRank>(args.ldd) : args.stride_d),
|
||||
output_op(args.epilogue),
|
||||
mode(args.mode),
|
||||
batch_count(args.batch_count),
|
||||
gemm_k_size(gemm_k_size),
|
||||
ptr_A(const_cast<void *>(args.ptr_A)),
|
||||
ptr_B(const_cast<void *>(args.ptr_B)),
|
||||
ptr_var(const_cast<void *>(args.ptr_var)),
|
||||
ptr_mean(const_cast<void *>(args.ptr_mean)),
|
||||
ptr_gamma(const_cast<void *>(args.ptr_gamma)),
|
||||
ptr_beta(const_cast<void *>(args.ptr_beta)),
|
||||
ptr_C(const_cast<void *>(args.ptr_C)),
|
||||
ptr_D(args.ptr_D),
|
||||
batch_stride_A(args.batch_stride_A),
|
||||
batch_stride_B(args.batch_stride_B),
|
||||
batch_stride_var(args.batch_stride_var),
|
||||
batch_stride_mean(args.batch_stride_mean),
|
||||
batch_stride_gamma(args.batch_stride_gamma),
|
||||
batch_stride_beta(args.batch_stride_beta),
|
||||
batch_stride_C(args.batch_stride_C),
|
||||
batch_stride_D(args.batch_stride_D),
|
||||
ptr_gather_A_indices(const_cast<int *>(args.ptr_gather_A_indices)),
|
||||
ptr_gather_B_indices(const_cast<int *>(args.ptr_gather_B_indices)),
|
||||
ptr_scatter_D_indices(const_cast<int *>(args.ptr_scatter_D_indices)),
|
||||
semaphore(static_cast<int *>(workspace)) {
|
||||
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
void update(
|
||||
Arguments const &args,
|
||||
void *workspace = nullptr) {
|
||||
|
||||
ptr_A = const_cast<void *>(args.ptr_A);
|
||||
ptr_B = const_cast<void *>(args.ptr_B);
|
||||
ptr_var = const_cast<void *>(args.ptr_var);
|
||||
ptr_mean = const_cast<void *>(args.ptr_mean);
|
||||
ptr_gamma = const_cast<void *>(args.ptr_gamma);
|
||||
ptr_beta = const_cast<void *>(args.ptr_beta);
|
||||
ptr_C = const_cast<void *>(args.ptr_C);
|
||||
ptr_D = args.ptr_D;
|
||||
|
||||
ptr_gather_A_indices = const_cast<int *>(args.ptr_gather_A_indices);
|
||||
ptr_gather_B_indices = const_cast<int *>(args.ptr_gather_B_indices);
|
||||
ptr_scatter_D_indices = const_cast<int *>(args.ptr_scatter_D_indices);
|
||||
|
||||
batch_stride_A = args.batch_stride_A;
|
||||
batch_stride_B = args.batch_stride_B;
|
||||
batch_stride_var = args.batch_stride_var;
|
||||
batch_stride_mean = args.batch_stride_mean;
|
||||
batch_stride_gamma = args.batch_stride_gamma;
|
||||
batch_stride_beta = args.batch_stride_beta;
|
||||
batch_stride_C = args.batch_stride_C;
|
||||
batch_stride_D = args.batch_stride_D;
|
||||
|
||||
output_op = args.epilogue;
|
||||
|
||||
semaphore = static_cast<int *>(workspace);
|
||||
CUTLASS_TRACE_HOST("GemmUniversal::Params::update()");
|
||||
}
|
||||
};
|
||||
|
||||
/// Shared memory storage structure
|
||||
union SharedStorage {
|
||||
typename Mma::SharedStorage main_loop;
|
||||
typename Epilogue::SharedStorage epilogue;
|
||||
};
|
||||
|
||||
public:
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
CUTLASS_DEVICE
|
||||
GemmLayernormMainloopFusion() { }
|
||||
|
||||
/// Determines whether kernel satisfies alignment
|
||||
static Status can_implement(
|
||||
cutlass::gemm::GemmCoord const & problem_size) {
|
||||
|
||||
CUTLASS_TRACE_HOST("GemmUniversal::can_implement()");
|
||||
|
||||
static int const kAlignmentA = (platform::is_same<LayoutA,
|
||||
layout::ColumnMajorInterleaved<32>>::value)
|
||||
? 32
|
||||
: (platform::is_same<LayoutA,
|
||||
layout::ColumnMajorInterleaved<64>>::value)
|
||||
? 64
|
||||
: Mma::IteratorA::AccessType::kElements;
|
||||
static int const kAlignmentB = (platform::is_same<LayoutB,
|
||||
layout::RowMajorInterleaved<32>>::value)
|
||||
? 32
|
||||
: (platform::is_same<LayoutB,
|
||||
layout::RowMajorInterleaved<64>>::value)
|
||||
? 64
|
||||
: Mma::IteratorB::AccessType::kElements;
|
||||
static int const kAlignmentC = (platform::is_same<LayoutC,
|
||||
layout::ColumnMajorInterleaved<32>>::value)
|
||||
? 32
|
||||
: (platform::is_same<LayoutC,
|
||||
layout::ColumnMajorInterleaved<64>>::value)
|
||||
? 64
|
||||
: Epilogue::OutputTileIterator::kElementsPerAccess;
|
||||
|
||||
bool isAMisaligned = false;
|
||||
bool isBMisaligned = false;
|
||||
bool isCMisaligned = false;
|
||||
|
||||
if (platform::is_same<LayoutA, layout::RowMajor>::value) {
|
||||
isAMisaligned = problem_size.k() % kAlignmentA;
|
||||
} else if (platform::is_same<LayoutA, layout::ColumnMajor>::value) {
|
||||
isAMisaligned = problem_size.m() % kAlignmentA;
|
||||
} else if (platform::is_same<LayoutA, layout::ColumnMajorInterleaved<32>>::value
|
||||
|| platform::is_same<LayoutA, layout::ColumnMajorInterleaved<64>>::value) {
|
||||
isAMisaligned = problem_size.k() % kAlignmentA;
|
||||
}
|
||||
|
||||
if (platform::is_same<LayoutB, layout::RowMajor>::value) {
|
||||
isBMisaligned = problem_size.n() % kAlignmentB;
|
||||
} else if (platform::is_same<LayoutB, layout::ColumnMajor>::value) {
|
||||
isBMisaligned = problem_size.k() % kAlignmentB;
|
||||
} else if (platform::is_same<LayoutB, layout::RowMajorInterleaved<32>>::value
|
||||
|| platform::is_same<LayoutB, layout::RowMajorInterleaved<64>>::value) {
|
||||
isBMisaligned = problem_size.k() % kAlignmentB;
|
||||
}
|
||||
|
||||
if (platform::is_same<LayoutC, layout::RowMajor>::value) {
|
||||
isCMisaligned = problem_size.n() % kAlignmentC;
|
||||
} else if (platform::is_same<LayoutC, layout::ColumnMajor>::value) {
|
||||
isCMisaligned = problem_size.m() % kAlignmentC;
|
||||
} else if (platform::is_same<LayoutC, layout::ColumnMajorInterleaved<32>>::value
|
||||
|| platform::is_same<LayoutC, layout::ColumnMajorInterleaved<64>>::value) {
|
||||
isCMisaligned = problem_size.n() % kAlignmentC;
|
||||
}
|
||||
|
||||
if (isAMisaligned) {
|
||||
CUTLASS_TRACE_HOST(" returning kErrorMisalignedOperand for A operand");
|
||||
return Status::kErrorMisalignedOperand;
|
||||
}
|
||||
|
||||
if (isBMisaligned) {
|
||||
CUTLASS_TRACE_HOST(" returning kErrorMisalignedOperand for B operand");
|
||||
return Status::kErrorMisalignedOperand;
|
||||
}
|
||||
|
||||
if (isCMisaligned) {
|
||||
CUTLASS_TRACE_HOST(" returning kErrorMisalignedOperand for C operand");
|
||||
return Status::kErrorMisalignedOperand;
|
||||
}
|
||||
|
||||
CUTLASS_TRACE_HOST(" returning kSuccess");
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
static Status can_implement(Arguments const &args) {
|
||||
return can_implement(args.problem_size);
|
||||
}
|
||||
|
||||
static size_t get_extra_workspace_size(Arguments const &args,
|
||||
cutlass::gemm::GemmCoord const &grid_tiled_shape) {
|
||||
|
||||
return 0;
|
||||
}
|
||||
|
||||
/// Executes one GEMM
|
||||
CUTLASS_DEVICE
|
||||
void operator()(Params const ¶ms, SharedStorage &shared_storage) {
|
||||
|
||||
// Compute threadblock location
|
||||
ThreadblockSwizzle threadblock_swizzle;
|
||||
|
||||
cutlass::gemm::GemmCoord threadblock_tile_offset =
|
||||
threadblock_swizzle.get_tile_offset(params.swizzle_log_tile);
|
||||
|
||||
// Early exit if CTA is out of range
|
||||
if (params.grid_tiled_shape.m() <= threadblock_tile_offset.m() ||
|
||||
params.grid_tiled_shape.n() <= threadblock_tile_offset.n()) {
|
||||
|
||||
return;
|
||||
}
|
||||
|
||||
int offset_k = 0;
|
||||
int problem_size_k = params.problem_size.k();
|
||||
|
||||
ElementA *ptr_A = static_cast<ElementA *>(params.ptr_A);
|
||||
ElementB *ptr_B = static_cast<ElementB *>(params.ptr_B);
|
||||
|
||||
//
|
||||
// Fetch pointers based on mode.
|
||||
//
|
||||
if (params.mode == GemmUniversalMode::kGemm ||
|
||||
params.mode == GemmUniversalMode::kGemmSplitKParallel) {
|
||||
|
||||
if (threadblock_tile_offset.k() + 1 < params.grid_tiled_shape.k()) {
|
||||
|
||||
problem_size_k = (threadblock_tile_offset.k() + 1) * params.gemm_k_size;
|
||||
}
|
||||
|
||||
offset_k = threadblock_tile_offset.k() * params.gemm_k_size;
|
||||
}
|
||||
else if (params.mode == GemmUniversalMode::kBatched) {
|
||||
ptr_A += threadblock_tile_offset.k() * params.batch_stride_A;
|
||||
ptr_B += threadblock_tile_offset.k() * params.batch_stride_B;
|
||||
}
|
||||
else if (params.mode == GemmUniversalMode::kArray) {
|
||||
ptr_A = static_cast<ElementA * const *>(params.ptr_A)[threadblock_tile_offset.k()];
|
||||
ptr_B = static_cast<ElementB * const *>(params.ptr_B)[threadblock_tile_offset.k()];
|
||||
}
|
||||
|
||||
__syncthreads();
|
||||
|
||||
// Compute initial location in logical coordinates
|
||||
cutlass::MatrixCoord tb_offset_A{
|
||||
threadblock_tile_offset.m() * Mma::Shape::kM,
|
||||
offset_k,
|
||||
};
|
||||
|
||||
cutlass::MatrixCoord tb_offset_B{
|
||||
offset_k,
|
||||
threadblock_tile_offset.n() * Mma::Shape::kN
|
||||
};
|
||||
|
||||
// Compute position within threadblock
|
||||
int thread_idx = threadIdx.x;
|
||||
|
||||
// Construct iterators to A and B operands
|
||||
typename Mma::IteratorA iterator_A(
|
||||
params.params_A,
|
||||
ptr_A,
|
||||
{params.problem_size.m(), problem_size_k},
|
||||
thread_idx,
|
||||
tb_offset_A,
|
||||
params.ptr_gather_A_indices);
|
||||
|
||||
typename Mma::IteratorB iterator_B(
|
||||
params.params_B,
|
||||
ptr_B,
|
||||
{problem_size_k, params.problem_size.n()},
|
||||
thread_idx,
|
||||
tb_offset_B,
|
||||
params.ptr_gather_B_indices);
|
||||
|
||||
// Construct iterators to A var/mean vector
|
||||
typename Mma::IteratorVarMean iterator_var_mean(
|
||||
params.problem_size.m(),
|
||||
static_cast<ElementScaleBias const *>(params.ptr_var),
|
||||
static_cast<ElementScaleBias const *>(params.ptr_mean),
|
||||
thread_idx,
|
||||
MatrixCoord(0, (threadblock_tile_offset.m() * Mma::Shape::kM))
|
||||
);
|
||||
|
||||
// Construct iterators to A scale/bias vector
|
||||
typename Mma::IteratorGammaBeta iterator_gamma_beta(
|
||||
problem_size_k,
|
||||
static_cast<ElementScaleBias const *>(params.ptr_gamma),
|
||||
static_cast<ElementScaleBias const *>(params.ptr_beta),
|
||||
thread_idx,
|
||||
MatrixCoord(
|
||||
0, (threadblock_tile_offset.k() * Mma::Shape::kK)
|
||||
)
|
||||
);
|
||||
|
||||
// Broadcast the warp_id computed by lane 0 to ensure dependent code
|
||||
// is compiled as warp-uniform.
|
||||
int warp_idx = __shfl_sync(0xffffffff, threadIdx.x / 32, 0);
|
||||
|
||||
int lane_idx = threadIdx.x % 32;
|
||||
|
||||
//
|
||||
// Main loop
|
||||
//
|
||||
|
||||
// Construct thread-scoped matrix multiply
|
||||
Mma mma(shared_storage.main_loop, thread_idx, warp_idx, lane_idx);
|
||||
|
||||
typename Mma::FragmentC accumulators;
|
||||
|
||||
accumulators.clear();
|
||||
|
||||
// Compute threadblock-scoped matrix multiply-add
|
||||
int gemm_k_iterations = (problem_size_k - offset_k + Mma::Shape::kK - 1) / Mma::Shape::kK;
|
||||
|
||||
// Compute threadblock-scoped matrix multiply-add
|
||||
mma(
|
||||
gemm_k_iterations,
|
||||
accumulators,
|
||||
iterator_A,
|
||||
iterator_B,
|
||||
iterator_var_mean,
|
||||
iterator_gamma_beta,
|
||||
accumulators);
|
||||
|
||||
//
|
||||
// Epilogue
|
||||
//
|
||||
|
||||
EpilogueOutputOp output_op(params.output_op);
|
||||
|
||||
//
|
||||
// Masked tile iterators constructed from members
|
||||
//
|
||||
|
||||
threadblock_tile_offset = threadblock_swizzle.get_tile_offset(params.swizzle_log_tile);
|
||||
|
||||
//assume identity swizzle
|
||||
MatrixCoord threadblock_offset(
|
||||
threadblock_tile_offset.m() * Mma::Shape::kM,
|
||||
threadblock_tile_offset.n() * Mma::Shape::kN
|
||||
);
|
||||
|
||||
int block_idx = threadblock_tile_offset.m() + threadblock_tile_offset.n() * params.grid_tiled_shape.m();
|
||||
|
||||
ElementC *ptr_C = static_cast<ElementC *>(params.ptr_C);
|
||||
ElementC *ptr_D = static_cast<ElementC *>(params.ptr_D);
|
||||
|
||||
//
|
||||
// Fetch pointers based on mode.
|
||||
//
|
||||
|
||||
// Construct the semaphore.
|
||||
Semaphore semaphore(params.semaphore + block_idx, thread_idx);
|
||||
|
||||
if (params.mode == GemmUniversalMode::kGemm) {
|
||||
|
||||
// If performing a reduction via split-K, fetch the initial synchronization
|
||||
if (params.grid_tiled_shape.k() > 1) {
|
||||
|
||||
// Fetch the synchronization lock initially but do not block.
|
||||
semaphore.fetch();
|
||||
|
||||
// Indicate which position in a serial reduction the output operator is currently updating
|
||||
output_op.set_k_partition(threadblock_tile_offset.k(), params.grid_tiled_shape.k());
|
||||
}
|
||||
}
|
||||
else if (params.mode == GemmUniversalMode::kGemmSplitKParallel) {
|
||||
ptr_D += threadblock_tile_offset.k() * params.batch_stride_D;
|
||||
}
|
||||
else if (params.mode == GemmUniversalMode::kBatched) {
|
||||
ptr_C += threadblock_tile_offset.k() * params.batch_stride_C;
|
||||
ptr_D += threadblock_tile_offset.k() * params.batch_stride_D;
|
||||
}
|
||||
else if (params.mode == GemmUniversalMode::kArray) {
|
||||
ptr_C = static_cast<ElementC * const *>(params.ptr_C)[threadblock_tile_offset.k()];
|
||||
ptr_D = static_cast<ElementC * const *>(params.ptr_D)[threadblock_tile_offset.k()];
|
||||
}
|
||||
|
||||
// Tile iterator loading from source tensor.
|
||||
typename Epilogue::OutputTileIterator iterator_C(
|
||||
params.params_C,
|
||||
ptr_C,
|
||||
params.problem_size.mn(),
|
||||
thread_idx,
|
||||
threadblock_offset,
|
||||
params.ptr_scatter_D_indices
|
||||
);
|
||||
|
||||
// Tile iterator writing to destination tensor.
|
||||
typename Epilogue::OutputTileIterator iterator_D(
|
||||
params.params_D,
|
||||
ptr_D,
|
||||
params.problem_size.mn(),
|
||||
thread_idx,
|
||||
threadblock_offset,
|
||||
params.ptr_scatter_D_indices
|
||||
);
|
||||
|
||||
Epilogue epilogue(
|
||||
shared_storage.epilogue,
|
||||
thread_idx,
|
||||
warp_idx,
|
||||
lane_idx);
|
||||
|
||||
// Wait on the semaphore - this latency may have been covered by iterator construction
|
||||
if (params.mode == GemmUniversalMode::kGemm && params.grid_tiled_shape.k() > 1) {
|
||||
|
||||
// For subsequent threadblocks, the source matrix is held in the 'D' tensor.
|
||||
if (threadblock_tile_offset.k()) {
|
||||
iterator_C = iterator_D;
|
||||
}
|
||||
|
||||
semaphore.wait(threadblock_tile_offset.k());
|
||||
}
|
||||
|
||||
// Execute the epilogue operator to update the destination tensor.
|
||||
epilogue(
|
||||
output_op,
|
||||
iterator_D,
|
||||
accumulators,
|
||||
iterator_C);
|
||||
|
||||
//
|
||||
// Release the semaphore
|
||||
//
|
||||
|
||||
if (params.mode == GemmUniversalMode::kGemm && params.grid_tiled_shape.k() > 1) {
|
||||
|
||||
int lock = 0;
|
||||
if (params.grid_tiled_shape.k() == threadblock_tile_offset.k() + 1) {
|
||||
|
||||
// The final threadblock resets the semaphore for subsequent grids.
|
||||
lock = 0;
|
||||
}
|
||||
else {
|
||||
// Otherwise, the semaphore is incremented
|
||||
lock = threadblock_tile_offset.k() + 1;
|
||||
}
|
||||
|
||||
semaphore.release(lock);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace kernel
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,468 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017 - 2022 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 Base scheduler for grouped problems
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/matrix_coord.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace kernel {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Enumerated type describing the type of scheduling to perform for the ProblemVisitor
|
||||
enum class GroupScheduleMode {
|
||||
// Perform all scheduling on device
|
||||
kDeviceOnly,
|
||||
// Precompute on the host the full sequence of problems to access
|
||||
kHostPrecompute
|
||||
};
|
||||
|
||||
/// Visitor class to abstract away the algorithm for iterating over tiles
|
||||
template <typename ProblemSizeHelper,
|
||||
typename ThreadblockShape_>
|
||||
struct BaseGroupedProblemVisitor {
|
||||
using ThreadblockShape = ThreadblockShape_;
|
||||
|
||||
struct ProblemInfo {
|
||||
static int32_t const kNoPrefetchEntry = -1;
|
||||
int32_t problem_idx;
|
||||
int32_t problem_start;
|
||||
|
||||
CUTLASS_DEVICE
|
||||
ProblemInfo() : problem_idx(kNoPrefetchEntry), problem_start(kNoPrefetchEntry) {}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
ProblemInfo(int32_t problem_idx_, int32_t problem_start_) :
|
||||
problem_idx(problem_idx_), problem_start(problem_start_) {}
|
||||
};
|
||||
|
||||
struct Params {
|
||||
cutlass::gemm::GemmCoord const *problem_sizes;
|
||||
int32_t problem_count;
|
||||
void const *workspace;
|
||||
int32_t tile_count;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
/// Ctor
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(): problem_sizes(nullptr), problem_count(0), workspace(nullptr), tile_count(0) { }
|
||||
|
||||
/// Ctor
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(
|
||||
cutlass::gemm::GemmCoord const *problem_sizes,
|
||||
int32_t problem_count,
|
||||
void const *workspace = nullptr,
|
||||
int32_t tile_count = 0
|
||||
):
|
||||
problem_sizes(problem_sizes),
|
||||
problem_count(problem_count),
|
||||
workspace(workspace),
|
||||
tile_count(tile_count)
|
||||
{}
|
||||
|
||||
};
|
||||
|
||||
Params const ¶ms;
|
||||
int32_t tile_idx;
|
||||
int32_t problem_tile_start;
|
||||
int32_t problem_idx;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
CUTLASS_DEVICE
|
||||
BaseGroupedProblemVisitor(
|
||||
Params const ¶ms_,
|
||||
int32_t block_idx
|
||||
):
|
||||
params(params_),
|
||||
tile_idx(block_idx),
|
||||
problem_tile_start(0),
|
||||
problem_idx(0)
|
||||
{}
|
||||
|
||||
/// Get the grid shape
|
||||
CUTLASS_HOST_DEVICE
|
||||
static cutlass::gemm::GemmCoord grid_shape(const cutlass::gemm::GemmCoord& problem) {
|
||||
|
||||
return cutlass::gemm::GemmCoord(
|
||||
((problem.m() - 1 + ThreadblockShape::kM) / ThreadblockShape::kM),
|
||||
((problem.n() - 1 + ThreadblockShape::kN) / ThreadblockShape::kN),
|
||||
1);
|
||||
}
|
||||
|
||||
/// Gets the global tile index
|
||||
CUTLASS_HOST_DEVICE
|
||||
int32_t tile_index() const {
|
||||
return tile_idx;
|
||||
}
|
||||
|
||||
/// Gets the index of the problem
|
||||
CUTLASS_HOST_DEVICE
|
||||
int32_t problem_index() const {
|
||||
return problem_idx;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
int32_t threadblock_idx() const {
|
||||
return tile_idx - problem_tile_start;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void advance(int32_t grid_size) {
|
||||
tile_idx += grid_size;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static void possibly_transpose_problem(cutlass::gemm::GemmCoord& problem) {
|
||||
ProblemSizeHelper::possibly_transpose_problem(problem);
|
||||
}
|
||||
|
||||
/// Returns the problem size for the current problem
|
||||
CUTLASS_HOST_DEVICE
|
||||
cutlass::gemm::GemmCoord problem_size() const {
|
||||
GemmCoord problem = params.problem_sizes[problem_idx];
|
||||
ProblemSizeHelper::possibly_transpose_problem(problem);
|
||||
return problem;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static int32_t tile_count(const cutlass::gemm::GemmCoord& grid) {
|
||||
return ProblemSizeHelper::tile_count(grid);
|
||||
}
|
||||
|
||||
static int32_t group_tile_count(const cutlass::gemm::GemmCoord* host_problem_sizes_ptr, int32_t problem_count) {
|
||||
int32_t total_tiles = 0;
|
||||
for (int32_t i = 0; i < problem_count; ++i) {
|
||||
auto problem = host_problem_sizes_ptr[i];
|
||||
possibly_transpose_problem(problem);
|
||||
auto grid = grid_shape(problem);
|
||||
total_tiles += tile_count(grid);
|
||||
}
|
||||
|
||||
return total_tiles;
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename ProblemSizeHelper,
|
||||
typename ThreadblockShape,
|
||||
GroupScheduleMode GroupScheduleMode_,
|
||||
int PrefetchTileCount,
|
||||
int ThreadCount
|
||||
>
|
||||
struct GroupedProblemVisitor;
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
// ProblemVisitor that performs all scheduling on device
|
||||
//
|
||||
template <typename ProblemSizeHelper,
|
||||
typename ThreadblockShape,
|
||||
int PrefetchTileCount,
|
||||
int ThreadCount>
|
||||
struct GroupedProblemVisitor<ProblemSizeHelper,
|
||||
ThreadblockShape,
|
||||
GroupScheduleMode::kDeviceOnly,
|
||||
PrefetchTileCount,
|
||||
ThreadCount>: public BaseGroupedProblemVisitor<ProblemSizeHelper, ThreadblockShape> {
|
||||
using Base = BaseGroupedProblemVisitor<ProblemSizeHelper, ThreadblockShape>;
|
||||
using Params = typename Base::Params;
|
||||
static int const kThreadCount = ThreadCount;
|
||||
static bool const kRequiresPrecomputation = false;
|
||||
static int const kThreadsPerWarp = 32;
|
||||
|
||||
struct SharedStorage {};
|
||||
|
||||
// Final tile of the problem loaded by this thread. Each thread will hold
|
||||
// a separate value.
|
||||
int32_t problem_ending_tile;
|
||||
|
||||
SharedStorage &shared_storage;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
CUTLASS_DEVICE
|
||||
GroupedProblemVisitor(
|
||||
Params const ¶ms_,
|
||||
SharedStorage &shared_storage_,
|
||||
int32_t block_idx
|
||||
): Base(params_, block_idx),
|
||||
problem_ending_tile(0),
|
||||
shared_storage(shared_storage_)
|
||||
{
|
||||
this->problem_idx = -1 * kThreadsPerWarp;
|
||||
this->problem_tile_start = 0;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
bool next_tile() {
|
||||
// Check whether the tile to compute is within the range of the current problem.
|
||||
int32_t problem_tile_end = __shfl_sync(0xffffffff, problem_ending_tile, this->problem_idx % kThreadsPerWarp);
|
||||
if (this->tile_idx < problem_tile_end) {
|
||||
return true;
|
||||
}
|
||||
|
||||
// Check whether the tile to compute is within the current group of problems fetched by the warp.
|
||||
// The last tile for this group is the final tile of the problem held by the final thread in the warp.
|
||||
int32_t group_tile_end = __shfl_sync(0xffffffff, problem_ending_tile, kThreadsPerWarp-1);
|
||||
|
||||
// Keep the starting problem for this group in `problem_idx`. This is done to reduce
|
||||
// register pressure. The starting problem for this group is simply the first problem
|
||||
// in the group most recently fetched by the warp.
|
||||
int32_t &group_problem_start = this->problem_idx;
|
||||
group_problem_start = (this->problem_idx / kThreadsPerWarp) * kThreadsPerWarp;
|
||||
|
||||
// Keep the starting tile for this group in `problem_tile_start`. This is done to reduce
|
||||
// register pressure.
|
||||
int32_t &group_tile_start = this->problem_tile_start;
|
||||
|
||||
// Each thread in the warp processes a separate problem to advance until
|
||||
// reaching a problem whose starting tile is less less than tile_idx.
|
||||
while (group_tile_end <= this->tile_idx) {
|
||||
group_problem_start += kThreadsPerWarp;
|
||||
if (group_problem_start > this->params.problem_count) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// Since `group_tile_start` is a reference to `this->problem_tile_start`, this
|
||||
// also sets `this->problem_tile_start`. The fact that `this->problem_tile_start`
|
||||
// is also set here is used later in `next_tile`.
|
||||
group_tile_start = group_tile_end;
|
||||
|
||||
int lane_idx = threadIdx.x % kThreadsPerWarp;
|
||||
int32_t lane_problem = group_problem_start + lane_idx;
|
||||
|
||||
// Compute the number of tiles in the problem assigned to each thread.
|
||||
problem_ending_tile = 0;
|
||||
if (lane_problem < this->params.problem_count) {
|
||||
cutlass::gemm::GemmCoord problem = this->params.problem_sizes[lane_problem];
|
||||
this->possibly_transpose_problem(problem);
|
||||
cutlass::gemm::GemmCoord grid = this->grid_shape(problem);
|
||||
problem_ending_tile = this->tile_count(grid);
|
||||
}
|
||||
|
||||
// Compute a warp-wide inclusive prefix sum to compute the ending tile index of
|
||||
// each thread's problem.
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 1; i < kThreadsPerWarp; i <<= 1) {
|
||||
int32_t val = __shfl_up_sync(0xffffffff, problem_ending_tile, i);
|
||||
if (lane_idx >= i) {
|
||||
problem_ending_tile += val;
|
||||
}
|
||||
}
|
||||
|
||||
// The total tile count for this group is now in the final position of the prefix sum
|
||||
int32_t tiles_in_group = __shfl_sync(0xffffffff, problem_ending_tile, kThreadsPerWarp-1);
|
||||
|
||||
problem_ending_tile += group_tile_start;
|
||||
group_tile_end += tiles_in_group;
|
||||
}
|
||||
|
||||
// The next problem to process is the first one that does not have ending tile position
|
||||
// that is greater than or equal to tile index.
|
||||
int32_t problem_idx_in_group =
|
||||
__popc(__ballot_sync(0xffffffff, problem_ending_tile <= this->tile_idx));
|
||||
|
||||
this->problem_idx = group_problem_start + problem_idx_in_group;
|
||||
|
||||
// The starting tile for this problem is the ending tile of the previous problem. In cases
|
||||
// where `problem_idx_in_group` is the first problem in the group, we do not need to reset
|
||||
// `problem_tile_start`, because it is set to the previous group's ending tile in the while
|
||||
// loop above.
|
||||
if (problem_idx_in_group > 0) {
|
||||
this->problem_tile_start = __shfl_sync(0xffffffff, problem_ending_tile, problem_idx_in_group - 1);
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
static size_t get_workspace_size(const cutlass::gemm::GemmCoord* host_problem_sizes_ptr,
|
||||
int32_t problem_count,
|
||||
int32_t block_count) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
static void host_precompute(const cutlass::gemm::GemmCoord* host_problem_sizes_ptr,
|
||||
int32_t problem_count,
|
||||
int32_t block_count,
|
||||
void* host_workspace_ptr) {}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
// Precomputes schedule on host and prefetches into shared memory
|
||||
//
|
||||
template <typename ProblemSizeHelper,
|
||||
typename ThreadblockShape,
|
||||
int PrefetchTileCount,
|
||||
int ThreadCount>
|
||||
struct GroupedProblemVisitor<ProblemSizeHelper,
|
||||
ThreadblockShape,
|
||||
GroupScheduleMode::kHostPrecompute,
|
||||
PrefetchTileCount,
|
||||
ThreadCount> : public BaseGroupedProblemVisitor<ProblemSizeHelper, ThreadblockShape> {
|
||||
static_assert(PrefetchTileCount > 0,
|
||||
"GroupedProblemVisitor with GroupScheduleMode `kHost` currently requires prefetching to shared memory");
|
||||
|
||||
using Base = BaseGroupedProblemVisitor<ProblemSizeHelper, ThreadblockShape>;
|
||||
using Params = typename Base::Params;
|
||||
using ProblemInfo = typename Base::ProblemInfo;
|
||||
static bool const kRequiresPrecomputation = true;
|
||||
|
||||
static int const kPrefetchTileCount = PrefetchTileCount;
|
||||
static int const kThreadCount = ThreadCount;
|
||||
|
||||
struct SharedStorage {
|
||||
// Sequence of problem IDs and starting tiles to compute
|
||||
cutlass::Array<ProblemInfo, kPrefetchTileCount> prefetched_problems;
|
||||
};
|
||||
|
||||
int32_t tiles_computed;
|
||||
int32_t iterations_per_block;
|
||||
int32_t block_load_start;
|
||||
SharedStorage &shared_storage;
|
||||
ProblemInfo const *problem_info_ptr;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
CUTLASS_DEVICE
|
||||
GroupedProblemVisitor(
|
||||
Params const ¶ms_,
|
||||
SharedStorage &shared_storage_,
|
||||
int32_t block_idx
|
||||
): Base(params_, block_idx),
|
||||
tiles_computed(0),
|
||||
shared_storage(shared_storage_),
|
||||
problem_info_ptr(reinterpret_cast<ProblemInfo const*>(params_.workspace))
|
||||
{
|
||||
iterations_per_block = (params_.tile_count - 1 + gridDim.x) / gridDim.x;
|
||||
block_load_start = iterations_per_block * block_idx;
|
||||
// Start prefetching the first set of tiles to compute
|
||||
prefetch_tiles();
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
bool next_tile() {
|
||||
if (this->tile_idx >= this->params.tile_count) {
|
||||
return false;
|
||||
}
|
||||
|
||||
int32_t prefetch_idx = (tiles_computed % kPrefetchTileCount);
|
||||
if (prefetch_idx == 0) {
|
||||
// Ensure all previous stores to shared memory have been completed
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
auto problem_info = shared_storage.prefetched_problems[prefetch_idx];
|
||||
++tiles_computed;
|
||||
|
||||
if ((tiles_computed % kPrefetchTileCount) == 0) {
|
||||
// Begin prefetching next set of tiles. Synchronize first to ensure that
|
||||
// we don't overwrite the current buffer while someone else is using it.
|
||||
__syncthreads();
|
||||
prefetch_tiles();
|
||||
}
|
||||
|
||||
this->problem_idx = problem_info.problem_idx;
|
||||
this->problem_tile_start = problem_info.problem_start;
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
static size_t get_workspace_size(const cutlass::gemm::GemmCoord* host_problem_sizes_ptr,
|
||||
int32_t problem_count,
|
||||
int32_t block_count) {
|
||||
int32_t total_tiles = Base::group_tile_count(host_problem_sizes_ptr, problem_count);
|
||||
int32_t entries_per_block = ((total_tiles - 1 + block_count) / block_count);
|
||||
return sizeof(ProblemInfo) * entries_per_block * block_count;
|
||||
}
|
||||
#if !defined(__CUDACC_RTC__)
|
||||
static void host_precompute(const cutlass::gemm::GemmCoord* host_problem_sizes_ptr,
|
||||
int32_t problem_count,
|
||||
int32_t block_count,
|
||||
void* host_workspace_ptr) {
|
||||
ProblemInfo* host_problem_info_ptr = reinterpret_cast<ProblemInfo*>(host_workspace_ptr);
|
||||
int32_t total_tiles = Base::group_tile_count(host_problem_sizes_ptr, problem_count);
|
||||
int32_t entries_per_block = (total_tiles - 1 + block_count) / block_count;
|
||||
|
||||
int tile = 0;
|
||||
int start_tile = 0;
|
||||
for (int p_idx = 0; p_idx < problem_count; ++p_idx) {
|
||||
auto problem = host_problem_sizes_ptr[p_idx];
|
||||
Base::possibly_transpose_problem(problem);
|
||||
auto grid = Base::grid_shape(problem);
|
||||
int tiles = Base::tile_count(grid);
|
||||
ProblemInfo problem_info(p_idx, start_tile);
|
||||
for (int i = 0; i < tiles; ++i, ++tile) {
|
||||
host_problem_info_ptr[(entries_per_block * (tile % block_count)) + (tile / block_count)] = problem_info;
|
||||
}
|
||||
start_tile += tiles;
|
||||
}
|
||||
}
|
||||
#endif
|
||||
private:
|
||||
CUTLASS_DEVICE
|
||||
void prefetch_tiles() {
|
||||
// TODO: Consider changing to use async copies from global to shared mem
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int32_t i = 0; i < kPrefetchTileCount; i += kThreadCount) {
|
||||
int32_t offset = threadIdx.x + i;
|
||||
if (offset < kPrefetchTileCount && (tiles_computed + offset < iterations_per_block)) {
|
||||
shared_storage.prefetched_problems[offset] = problem_info_ptr[block_load_start + tiles_computed + offset];
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace kernel
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,711 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017 - 2022 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 Grouped Rank2K kernel.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/blas3.h"
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/fast_math.h"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/matrix_coord.h"
|
||||
#include "cutlass/complex.h"
|
||||
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/trace.h"
|
||||
#include "cutlass/gemm/kernel/rank_2k_transpose_operands.h"
|
||||
#include "cutlass/gemm/kernel/rank_2k_grouped_problem_visitor.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace kernel {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename Mma1_, ///! Threadblock-scoped matrix multiply-accumulate (A*B^T)
|
||||
typename Mma2_, ///! Threadblock-scoped matrix multiply-accumulate (B*A^T)
|
||||
typename Epilogue_, ///! Epilogue
|
||||
typename ThreadblockSwizzle_, ///! Threadblock swizzling function
|
||||
ComplexTransform OriginalTransformA_, ///! Public-facing transformation on A
|
||||
ComplexTransform OriginalTransformB_, ///! Public-facing transformation on B
|
||||
FillMode FillModeC_, ///! Fill Mode for C (kLower or kUpper)
|
||||
BlasMode BlasMode_, ///! Blas3 computation mode
|
||||
GroupScheduleMode GroupScheduleMode_, ///! Type of scheduling to perform
|
||||
bool Transposed = false
|
||||
>
|
||||
struct Rank2KGrouped {
|
||||
public:
|
||||
|
||||
using Mma1 = Mma1_;
|
||||
using Mma2 = Mma2_;
|
||||
|
||||
static_assert(platform::is_same<typename Mma1::LayoutC, cutlass::layout::RowMajor>::value &&
|
||||
platform::is_same<typename Mma2::LayoutC, cutlass::layout::RowMajor>::value,
|
||||
"Kernel-level grouped Rank2K requires that LayoutC be row major.");
|
||||
|
||||
// Define generic Mma for usecases that use Kernel::Mma
|
||||
using Mma = Mma1_;
|
||||
|
||||
using Epilogue = Epilogue_;
|
||||
using EpilogueOutputOp = typename Epilogue::OutputOp;
|
||||
using ThreadblockSwizzle = ThreadblockSwizzle_;
|
||||
static GroupScheduleMode const kGroupScheduleMode = GroupScheduleMode_;
|
||||
static bool const kTransposed = Transposed;
|
||||
|
||||
// Public-facing type definitions related to operand element type, layout, and complex conjugate
|
||||
// operation. Must interact with the 'kTransposed' notion to reflect the original layout,
|
||||
// fill mode, etc. passed in.
|
||||
//
|
||||
// Recall that a Rank2K operation performs (A x BT) + (B x AT)
|
||||
// This is performed via:
|
||||
// Mma1 = (A x BT)
|
||||
// Mma2 = (B x AT)
|
||||
//
|
||||
// However, if C needs to be transposed, then this is changed to the following:
|
||||
// Mma1 = (B x AT)
|
||||
// Mma2 = (A x BT)
|
||||
//
|
||||
// The transformation above is achieved by swapping the Layouts/Elements/Transforms/etc.
|
||||
// of A and B as they are passed into the instantiations of Mma1 and Mma2.
|
||||
//
|
||||
// Now, given access to only Mma1 and Mma2, as well as whether a transposition has occurred,
|
||||
// we wish to retrieve the original Layouts/Elements/etc. for A and B that were passed into
|
||||
// the device-level call.
|
||||
//
|
||||
// The logic to do this (which is made clearer by referencing the above instantiations) is as follows:
|
||||
// LayoutA = kTransposed ? Mma2::LayoutA : Mma1::LayoutA
|
||||
// LayoutB = kTransposed ? Mma1::LayoutA : Mma2::LayoutA
|
||||
//
|
||||
// We achieve this swapping by passing Mma1::*A and Mma2::*B to Rank2KMapArguments:
|
||||
using MapArgumentsA = kernel::detail::Rank2KMapArguments<
|
||||
typename Mma1::IteratorA::Element,
|
||||
typename Mma1::IteratorA::Layout,
|
||||
Mma1::kTransformA,
|
||||
Mma1::IteratorA::AccessType::kElements,
|
||||
typename Mma2::IteratorA::Element,
|
||||
typename Mma2::IteratorA::Layout,
|
||||
Mma2::kTransformA,
|
||||
Mma2::IteratorA::AccessType::kElements,
|
||||
typename Mma1::LayoutC,
|
||||
FillModeC_,
|
||||
kTransposed
|
||||
>;
|
||||
|
||||
using ElementA = typename MapArgumentsA::ElementA;
|
||||
using LayoutA = typename MapArgumentsA::LayoutA;
|
||||
static int const kAlignmentA = MapArgumentsA::kAlignmentA;
|
||||
|
||||
using MapArgumentsB = kernel::detail::Rank2KMapArguments<
|
||||
typename Mma2::IteratorA::Element,
|
||||
typename Mma2::IteratorA::Layout,
|
||||
Mma2::kTransformA,
|
||||
Mma2::IteratorA::AccessType::kElements,
|
||||
typename Mma1::IteratorA::Element,
|
||||
typename Mma1::IteratorA::Layout,
|
||||
Mma1::kTransformA,
|
||||
Mma1::IteratorA::AccessType::kElements,
|
||||
typename Mma2::LayoutC,
|
||||
FillModeC_,
|
||||
kTransposed
|
||||
>;
|
||||
|
||||
using ElementB = typename MapArgumentsB::ElementA;
|
||||
using LayoutB = typename MapArgumentsB::LayoutA;
|
||||
static int const kAlignmentB = MapArgumentsB::kAlignmentA;
|
||||
|
||||
// Use the user-provided TransformA and TransformB, rather than those
|
||||
// resulting from MapArguments, because Mma1 and Mma2 may have different
|
||||
// complex transforms than those passed in by the user.
|
||||
// (See kernel/rank_2k_complex.h for an example of this)
|
||||
static cutlass::ComplexTransform const kTransformA = OriginalTransformA_;
|
||||
static cutlass::ComplexTransform const kTransformB = OriginalTransformB_;
|
||||
|
||||
using ElementC = typename Epilogue::OutputTileIterator::Element;
|
||||
using LayoutC = typename MapArgumentsA::LayoutC;
|
||||
static int const kAlignmentC = Epilogue::OutputTileIterator::kElementsPerAccess;
|
||||
static FillMode const kFillModeC = MapArgumentsA::kFillModeC;
|
||||
|
||||
// Common type definitions for Mma1 and Mma2
|
||||
using Operator = typename Mma1::Operator;
|
||||
using OperatorClass = typename Mma1::Operator::OperatorClass;
|
||||
using ThreadblockShape = typename Mma1::Shape;
|
||||
using WarpShape = typename Mma1::Operator::Shape;
|
||||
using InstructionShape = typename Mma1::Policy::Operator::InstructionShape;
|
||||
using ArchTag = typename Mma1::ArchTag;
|
||||
|
||||
static int const kStages = Mma1::kStages;
|
||||
static BlasMode const kBlasMode = BlasMode_;
|
||||
|
||||
private:
|
||||
static FillMode const kInternalFillModeC = FillModeC_;
|
||||
|
||||
public:
|
||||
|
||||
/// Warp count (concept: GemmShape)
|
||||
using WarpCount = typename Mma1::WarpCount;
|
||||
static int const kThreadCount = 32 * WarpCount::kCount;
|
||||
|
||||
using ProblemVisitor = Rank2KGroupedProblemVisitor<
|
||||
ThreadblockShape,
|
||||
kGroupScheduleMode,
|
||||
kThreadCount,
|
||||
kThreadCount,
|
||||
kInternalFillModeC>;
|
||||
|
||||
//
|
||||
// Structures
|
||||
//
|
||||
|
||||
/// Argument structure
|
||||
struct Arguments {
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
GemmUniversalMode mode;
|
||||
GemmCoord *problem_sizes;
|
||||
int problem_count;
|
||||
int threadblock_count;
|
||||
|
||||
typename EpilogueOutputOp::Params epilogue;
|
||||
|
||||
ElementA ** ptr_A;
|
||||
ElementB ** ptr_B;
|
||||
ElementC ** ptr_C;
|
||||
ElementC ** ptr_D;
|
||||
|
||||
typename LayoutA::Stride::LongIndex *lda;
|
||||
typename LayoutB::Stride::LongIndex *ldb;
|
||||
typename LayoutC::Stride::LongIndex *ldc;
|
||||
typename LayoutC::Stride::LongIndex *ldd;
|
||||
|
||||
// Only used by device-level operator
|
||||
GemmCoord *host_problem_sizes;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
/// Default ctor
|
||||
CUTLASS_HOST_DEVICE
|
||||
Arguments():
|
||||
mode(GemmUniversalMode::kGemm),
|
||||
problem_count(0),
|
||||
threadblock_count(0),
|
||||
ptr_A(nullptr),
|
||||
ptr_B(nullptr),
|
||||
ptr_C(nullptr),
|
||||
ptr_D(nullptr),
|
||||
lda(nullptr),
|
||||
ldb(nullptr),
|
||||
ldc(nullptr),
|
||||
ldd(nullptr),
|
||||
host_problem_sizes(nullptr)
|
||||
{
|
||||
|
||||
}
|
||||
|
||||
/// Ctor
|
||||
CUTLASS_HOST_DEVICE
|
||||
Arguments(
|
||||
GemmUniversalMode mode,
|
||||
GemmCoord *problem_sizes,
|
||||
int problem_count,
|
||||
int threadblock_count,
|
||||
typename EpilogueOutputOp::Params epilogue,
|
||||
ElementA ** ptr_A,
|
||||
ElementB ** ptr_B,
|
||||
ElementC ** ptr_C,
|
||||
ElementC ** ptr_D,
|
||||
typename LayoutA::Stride::LongIndex *lda,
|
||||
typename LayoutB::Stride::LongIndex *ldb,
|
||||
typename LayoutC::Stride::LongIndex *ldc,
|
||||
typename LayoutC::Stride::LongIndex *ldd,
|
||||
GemmCoord *host_problem_sizes=nullptr
|
||||
):
|
||||
mode(mode),
|
||||
problem_sizes(problem_sizes),
|
||||
problem_count(problem_count),
|
||||
threadblock_count(threadblock_count),
|
||||
epilogue(epilogue),
|
||||
ptr_A(ptr_A),
|
||||
ptr_B(ptr_B),
|
||||
ptr_C(ptr_C),
|
||||
ptr_D(ptr_D),
|
||||
lda(lda),
|
||||
ldb(ldb),
|
||||
ldc(ldc),
|
||||
ldd(ldd),
|
||||
host_problem_sizes(host_problem_sizes)
|
||||
{
|
||||
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
//
|
||||
// Structure for precomputing values in host memory and passing to kernels
|
||||
//
|
||||
|
||||
/// Parameters structure
|
||||
struct Params {
|
||||
|
||||
typename ProblemVisitor::Params problem_visitor;
|
||||
int threadblock_count;
|
||||
|
||||
typename EpilogueOutputOp::Params output_op;
|
||||
|
||||
GemmUniversalMode mode;
|
||||
int batch_count;
|
||||
|
||||
ElementA ** ptr_A;
|
||||
ElementB ** ptr_B;
|
||||
ElementC ** ptr_C;
|
||||
ElementC ** ptr_D;
|
||||
|
||||
typename LayoutA::Stride::LongIndex *lda;
|
||||
typename LayoutB::Stride::LongIndex *ldb;
|
||||
typename LayoutC::Stride::LongIndex *ldc;
|
||||
typename LayoutC::Stride::LongIndex *ldd;
|
||||
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params():
|
||||
mode(cutlass::gemm::GemmUniversalMode::kGemm),
|
||||
ptr_A(nullptr),
|
||||
ptr_B(nullptr),
|
||||
ptr_C(nullptr),
|
||||
ptr_D(nullptr),
|
||||
lda(nullptr),
|
||||
ldb(nullptr),
|
||||
ldc(nullptr),
|
||||
ldd(nullptr)
|
||||
{ }
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(Arguments const &args, void *workspace = nullptr, int tile_count = 0):
|
||||
problem_visitor(args.problem_sizes, args.problem_count, workspace, tile_count),
|
||||
threadblock_count(args.threadblock_count),
|
||||
output_op(args.epilogue),
|
||||
ptr_A(args.ptr_A),
|
||||
ptr_B(args.ptr_B),
|
||||
ptr_C(args.ptr_C),
|
||||
ptr_D(args.ptr_D),
|
||||
lda(args.lda),
|
||||
ldb(args.ldb),
|
||||
ldc(args.ldc),
|
||||
ldd(args.ldd)
|
||||
{
|
||||
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
void update(
|
||||
Arguments const &args,
|
||||
void *workspace = nullptr,
|
||||
int tile_count = 0) {
|
||||
|
||||
problem_visitor = typename ProblemVisitor::Params(args.problem_sizes, args.problem_count, workspace, tile_count);
|
||||
threadblock_count = args.threadblock_count;
|
||||
output_op = args.output_op;
|
||||
ptr_A = args.ptr_A;
|
||||
ptr_B = args.ptr_B;
|
||||
ptr_C = args.ptr_C;
|
||||
ptr_D = args.ptr_D;
|
||||
}
|
||||
};
|
||||
|
||||
/// Shared memory storage structure
|
||||
struct SharedStorage {
|
||||
union {
|
||||
typename Mma1::SharedStorage mma1_main_loop;
|
||||
typename Mma2::SharedStorage mma2_main_loop;
|
||||
typename Epilogue::SharedStorage epilogue;
|
||||
} kernel;
|
||||
|
||||
// ProblemVisitor shared storage can't be overlapped with others
|
||||
typename ProblemVisitor::SharedStorage problem_visitor;
|
||||
};
|
||||
|
||||
public:
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
CUTLASS_DEVICE
|
||||
Rank2KGrouped() { }
|
||||
|
||||
/// Determines whether kernel satisfies alignment
|
||||
static Status can_implement(cutlass::gemm::GemmCoord const & problem_size) {
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
static Status can_implement(Arguments const &args) {
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
static size_t get_extra_workspace_size(
|
||||
Arguments const &args,
|
||||
cutlass::gemm::GemmCoord const &grid_tiled_shape) {
|
||||
|
||||
return 0;
|
||||
}
|
||||
|
||||
/// Executes one GEMM
|
||||
CUTLASS_DEVICE
|
||||
void operator()(Params const ¶ms, SharedStorage &shared_storage) {
|
||||
|
||||
//
|
||||
// Problem visitor.
|
||||
//
|
||||
|
||||
ProblemVisitor problem_visitor(
|
||||
params.problem_visitor,
|
||||
shared_storage.problem_visitor,
|
||||
blockIdx.x);
|
||||
|
||||
// Outer 'persistent' loop to iterate over tiles
|
||||
while (problem_visitor.next_tile()) {
|
||||
|
||||
GemmCoord problem_size = problem_visitor.problem_size();
|
||||
int32_t problem_idx = problem_visitor.problem_index();
|
||||
int32_t threadblock_idx = int32_t(problem_visitor.threadblock_idx());
|
||||
|
||||
GemmCoord grid_shape = problem_visitor.grid_shape(problem_size);
|
||||
|
||||
cutlass::gemm::GemmCoord threadblock_tile_offset = problem_visitor.threadblock_offset(threadblock_idx);
|
||||
|
||||
//
|
||||
// Perform checks to determine whether the results of this threadblock will be needed.
|
||||
// An example of an unneeded threadblock is one that is assigned to compute in the upper
|
||||
// portion of a Rank2K kernel filled with mode kLower.
|
||||
//
|
||||
// TODO: Consider pushing these checks into ProblemVisitor to avoid spuriously
|
||||
// returning from `next_tile()`.
|
||||
//
|
||||
|
||||
// Early exit if threadblock is out of range
|
||||
if (grid_shape.m() <= threadblock_tile_offset.m() ||
|
||||
grid_shape.n() <= threadblock_tile_offset.n()) {
|
||||
// Next tile
|
||||
problem_visitor.advance(gridDim.x);
|
||||
continue;
|
||||
}
|
||||
|
||||
// Skip this tile if Fill Mode is Lower and
|
||||
// if the entire tile is above the main diagonal (bottom-left corner is at or above the diagonal)
|
||||
if (kInternalFillModeC == cutlass::FillMode::kLower &&
|
||||
(threadblock_tile_offset.m() + 1) * Mma1::Shape::kM <= threadblock_tile_offset.n() * Mma1::Shape::kN) {
|
||||
// Next tile
|
||||
problem_visitor.advance(gridDim.x);
|
||||
continue;
|
||||
}
|
||||
|
||||
// Skip this tile if Fill Mode is Upper and
|
||||
// if the entire tile is below the main diagonal (top-right corner is at or below the diagonal)
|
||||
if (kInternalFillModeC == cutlass::FillMode::kUpper &&
|
||||
threadblock_tile_offset.m() * Mma1::Shape::kM >= (threadblock_tile_offset.n() + 1) * Mma1::Shape::kN) {
|
||||
// Next tile
|
||||
problem_visitor.advance(gridDim.x);
|
||||
continue;
|
||||
}
|
||||
|
||||
bool tile_on_diagonal = false;
|
||||
// Mark tiles that are being crossed by the main diagonal
|
||||
// (top-right and bottom-left corners are on either side of the diagonal)
|
||||
if ((threadblock_tile_offset.m() + 1) * Mma1::Shape::kM > threadblock_tile_offset.n() * Mma1::Shape::kN
|
||||
&& threadblock_tile_offset.m() * Mma1::Shape::kM < (threadblock_tile_offset.n() + 1) * Mma1::Shape::kN) {
|
||||
tile_on_diagonal = true;
|
||||
}
|
||||
|
||||
int offset_k = 0;
|
||||
int problem_size_k = problem_size.k();
|
||||
|
||||
//
|
||||
// Fetch pointers based on mode.
|
||||
//
|
||||
if (params.mode == GemmUniversalMode::kGemm ||
|
||||
params.mode == GemmUniversalMode::kGemmSplitKParallel) {
|
||||
|
||||
if (threadblock_tile_offset.k() + 1 < grid_shape.k()) {
|
||||
problem_size_k = (threadblock_tile_offset.k() + 1) * problem_size.k();
|
||||
}
|
||||
|
||||
offset_k = threadblock_tile_offset.k() * problem_size.k();
|
||||
}
|
||||
|
||||
ElementA *ptr_A = reinterpret_cast<ElementA *>((kTransposed ? params.ptr_B[problem_idx] : params.ptr_A[problem_idx]));
|
||||
typename LayoutA::Stride::LongIndex ldm_A = (kTransposed ? params.ldb[problem_idx] : params.lda[problem_idx]);
|
||||
|
||||
ElementB *ptr_B = reinterpret_cast<ElementB *>((kTransposed ? params.ptr_A[problem_idx] : params.ptr_B[problem_idx]));
|
||||
typename LayoutB::Stride::LongIndex ldm_B = (kTransposed ? params.lda[problem_idx] : params.ldb[problem_idx]);
|
||||
|
||||
// Compute initial location in logical coordinates
|
||||
cutlass::MatrixCoord tb_offset_MxK{
|
||||
threadblock_tile_offset.m() * Mma1::Shape::kM,
|
||||
offset_k,
|
||||
};
|
||||
|
||||
cutlass::MatrixCoord tb_offset_KxN{
|
||||
offset_k,
|
||||
threadblock_tile_offset.n() * Mma1::Shape::kN
|
||||
};
|
||||
|
||||
// Assume identity swizzle
|
||||
MatrixCoord tb_offset(
|
||||
threadblock_tile_offset.m() * Mma1::Shape::kM,
|
||||
threadblock_tile_offset.n() * Mma1::Shape::kN
|
||||
);
|
||||
|
||||
// Compute position within threadblock
|
||||
int thread_idx = threadIdx.x;
|
||||
|
||||
// Construct iterators to A and B operands for Mma1
|
||||
typename Mma1::IteratorA iterator_A(
|
||||
Mma1::IteratorA::Params(ldm_A),
|
||||
ptr_A,
|
||||
{problem_size.m(), problem_size_k},
|
||||
thread_idx,
|
||||
tb_offset_MxK);
|
||||
|
||||
typename Mma1::IteratorB iterator_BT(
|
||||
Mma1::IteratorB::Params(ldm_B),
|
||||
ptr_B,
|
||||
{problem_size_k, problem_size.n()},
|
||||
thread_idx,
|
||||
tb_offset_KxN);
|
||||
|
||||
// Construct iterators to A and B operands for Mma2
|
||||
typename Mma2::IteratorA iterator_B(
|
||||
Mma2::IteratorA::Params(ldm_B),
|
||||
ptr_B,
|
||||
{problem_size.m(), problem_size_k},
|
||||
thread_idx,
|
||||
tb_offset_MxK);
|
||||
|
||||
typename Mma2::IteratorB iterator_AT(
|
||||
Mma2::IteratorB::Params(ldm_A),
|
||||
ptr_A,
|
||||
{problem_size_k, problem_size.n()},
|
||||
thread_idx,
|
||||
tb_offset_KxN);
|
||||
|
||||
// Broadcast the warp_id computed by lane 0 to ensure dependent code
|
||||
// is compiled as warp-uniform.
|
||||
int warp_idx = __shfl_sync(0xffffffff, threadIdx.x / 32, 0);
|
||||
|
||||
int lane_idx = threadIdx.x % 32;
|
||||
|
||||
//
|
||||
// Main loop
|
||||
//
|
||||
|
||||
// Construct thread-scoped matrix multiply for Mma1 (A x BT)
|
||||
Mma1 mma1(shared_storage.kernel.mma1_main_loop, thread_idx, warp_idx, lane_idx);
|
||||
|
||||
// Construct thread-scoped matrix multiply for Mma2 (B x AT)
|
||||
Mma2 mma2(shared_storage.kernel.mma2_main_loop, thread_idx, warp_idx, lane_idx);
|
||||
|
||||
typename Mma1::FragmentC accumulators;
|
||||
|
||||
accumulators.clear();
|
||||
|
||||
// Compute threadblock-scoped matrix multiply-add
|
||||
int gemm_k_iterations = (problem_size_k - offset_k + Mma1::Shape::kK - 1) / Mma1::Shape::kK;
|
||||
|
||||
// Wait for all threads to finish their epilogue phases from the previous tile.
|
||||
__syncthreads();
|
||||
|
||||
// Compute threadblock-scoped matrix multiply-add (A x BT)
|
||||
mma1(
|
||||
gemm_k_iterations,
|
||||
accumulators,
|
||||
iterator_A,
|
||||
iterator_BT,
|
||||
accumulators);
|
||||
|
||||
// HER2K kernel needs Alpha to be complex and is conj(Alpha) is applied to the second HERK.
|
||||
if (kBlasMode == BlasMode::kHermitian) {
|
||||
|
||||
//
|
||||
// Epilogue
|
||||
//
|
||||
|
||||
EpilogueOutputOp output_op(params.output_op);
|
||||
|
||||
int block_idx = threadblock_tile_offset.m() + threadblock_tile_offset.n() * grid_shape.m();
|
||||
|
||||
ElementC *ptr_C = static_cast<ElementC *>(params.ptr_C[problem_idx]);
|
||||
ElementC *ptr_D = static_cast<ElementC *>(params.ptr_D[problem_idx]);
|
||||
|
||||
// If TB not on diagonal, FillMode doesn't apply.
|
||||
FillMode kFillModeTB = tile_on_diagonal ? kInternalFillModeC : FillMode::kNone;
|
||||
|
||||
// Tile iterator loading from source tensor.
|
||||
typename Epilogue::OutputTileIterator iterator_C(
|
||||
Epilogue::OutputTileIterator::Params(params.ldc[problem_idx]),
|
||||
ptr_C,
|
||||
problem_size.mn(),
|
||||
thread_idx,
|
||||
tb_offset,
|
||||
kFillModeTB
|
||||
);
|
||||
|
||||
// Tile iterator writing to destination tensor.
|
||||
typename Epilogue::OutputTileIterator iterator_D(
|
||||
Epilogue::OutputTileIterator::Params(params.ldd[problem_idx]),
|
||||
ptr_D,
|
||||
problem_size.mn(),
|
||||
thread_idx,
|
||||
tb_offset,
|
||||
kFillModeTB
|
||||
);
|
||||
|
||||
Epilogue epilogue(
|
||||
shared_storage.kernel.epilogue,
|
||||
thread_idx,
|
||||
warp_idx,
|
||||
lane_idx);
|
||||
|
||||
// Execute the epilogue operator to update the destination tensor.
|
||||
epilogue(
|
||||
output_op,
|
||||
iterator_D,
|
||||
accumulators,
|
||||
iterator_C);
|
||||
|
||||
__syncthreads();
|
||||
|
||||
accumulators.clear();
|
||||
}
|
||||
|
||||
// Compute threadblock-scoped matrix multiply-add (B x AT)
|
||||
mma2(
|
||||
gemm_k_iterations,
|
||||
accumulators,
|
||||
iterator_B,
|
||||
iterator_AT,
|
||||
accumulators);
|
||||
|
||||
//
|
||||
// Epilogue
|
||||
//
|
||||
|
||||
EpilogueOutputOp output_op(params.output_op);
|
||||
|
||||
/* Needed for HER2K where the second HERK is multiplied by conj(alpha) */
|
||||
typename EpilogueOutputOp::Params second_her2k_params(conj(params.output_op.alpha), 1);
|
||||
EpilogueOutputOp output_op_her2k(second_her2k_params);
|
||||
|
||||
//
|
||||
// Masked tile iterators constructed from members
|
||||
//
|
||||
|
||||
int block_idx = threadblock_tile_offset.m() + threadblock_tile_offset.n() * grid_shape.m();
|
||||
|
||||
ElementC *ptr_C = static_cast<ElementC *>(params.ptr_C[problem_idx]);
|
||||
|
||||
// HER2K kernel needs Alpha to be complex and is conj(Alpha) is applied to the second HERK.
|
||||
if (kBlasMode == BlasMode::kHermitian) {
|
||||
ptr_C = static_cast<ElementC *>(params.ptr_D[problem_idx]);
|
||||
}
|
||||
|
||||
ElementC *ptr_D = static_cast<ElementC *>(params.ptr_D[problem_idx]);
|
||||
|
||||
// If TB not on diagonal, FillMode doesn't apply.
|
||||
FillMode kFillModeTB = tile_on_diagonal ? kInternalFillModeC : FillMode::kNone;
|
||||
|
||||
// Tile iterator loading from source tensor.
|
||||
typename Epilogue::OutputTileIterator iterator_C(
|
||||
Epilogue::OutputTileIterator::Params(params.ldc[problem_idx]),
|
||||
ptr_C,
|
||||
problem_size.mn(),
|
||||
thread_idx,
|
||||
tb_offset,
|
||||
kFillModeTB
|
||||
);
|
||||
|
||||
// Tile iterator writing to destination tensor.
|
||||
typename Epilogue::OutputTileIterator iterator_D(
|
||||
Epilogue::OutputTileIterator::Params(params.ldd[problem_idx]),
|
||||
ptr_D,
|
||||
problem_size.mn(),
|
||||
thread_idx,
|
||||
tb_offset,
|
||||
kFillModeTB
|
||||
);
|
||||
|
||||
Epilogue epilogue(
|
||||
shared_storage.kernel.epilogue,
|
||||
thread_idx,
|
||||
warp_idx,
|
||||
lane_idx);
|
||||
|
||||
// Execute the epilogue operator to update the destination tensor.
|
||||
if (kBlasMode == BlasMode::kSymmetric) {
|
||||
epilogue(
|
||||
output_op,
|
||||
iterator_D,
|
||||
accumulators,
|
||||
iterator_C);
|
||||
} else {
|
||||
epilogue(
|
||||
output_op_her2k,
|
||||
iterator_D,
|
||||
accumulators,
|
||||
iterator_C);
|
||||
}
|
||||
|
||||
// Next tile
|
||||
problem_visitor.advance(gridDim.x);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace kernel
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,368 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017 - 2022 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 Problem visitor for grouped Rank2K operations.
|
||||
|
||||
This problem visitor is specialized for Rank2K operations, for which matrix C is upper/lower
|
||||
triangular. Using a problem visitor designed for GEMMs for Rank2K problems is inefficient
|
||||
because threadblocks will be frequently assigned to tiles that exit early (e.g., due to
|
||||
being assigned to a tile in the upper-triangular portion of a lower-triangular problem).
|
||||
This can lead to load imbalance among threadblocks, as the GEMM-based scheduler
|
||||
assigns all threadblocks to nearly the same number of tiles, regardless of whether
|
||||
those tiles exit early.
|
||||
|
||||
Consider an example of a group of four Rank2Ks with matrix C consisting of a grid of 2x2 tiles.
|
||||
Consider a grid of 8 threadblocks. The default GEMM scheduler will assign threadblocks to
|
||||
tiles in the following order:
|
||||
Rank2K 0 Rank2K 1 Rank2K 2 Rank2K 3
|
||||
0 1 4 5 0 1 4 5
|
||||
2 3 6 7 2 3 6 7
|
||||
Assuming that the problems are lower triangular, blocks 1 and 5 are continuously assigned
|
||||
to inactive tiles.
|
||||
|
||||
This problem visitor aims to assign threadblocks to only those tiles which are in the
|
||||
upper/lower triangular portion of a given problem. Using the example above, the resulting
|
||||
assignment would be:
|
||||
Rank2K 0 Rank2K 1 Rank2K 2 Rank2K 3
|
||||
0 - 3 - 6 - 1 -
|
||||
1 2 4 5 7 0 2 3
|
||||
|
||||
Achieving the schedule above requires a mapping from threadblock ID to tile coordinates (i, j).
|
||||
We will illustrate this by mapping on a lower-triangular matrix with a 3x3 grid. We first
|
||||
calculate row and column indices assuming one-indexed rows, tiles, and threadblock IDs, and
|
||||
then subtract one to convert to zero-indexed.
|
||||
Col 1 Col 2 Col 3
|
||||
----------------------
|
||||
Row 1 | 1 - -
|
||||
Row 2 | 2 3 -
|
||||
Row 3 | 4 5 6
|
||||
|
||||
We next outline this mapping, borrowing from: https://stackoverflow.com/a/40954159
|
||||
|
||||
Calculating row i given threadblock ID t
|
||||
----------------------------------------
|
||||
For a given row i, all threadblock IDs t in that row satisfy the following:
|
||||
t <= 1 + 2 + 3 + ... + (i-1) + i
|
||||
|
||||
The closed-form equation for the right-hand side is: i(i+1)/2.
|
||||
Using this, we can solve for i given t:
|
||||
t <= i(i+1)/2
|
||||
2t <= i^2 + i
|
||||
2t <= i^2 + i + 0.25 - 0.25
|
||||
2t + 0.25 <= i^2 + i + 0.25
|
||||
2t + 0.25 <= (i + 0.5)^2
|
||||
sqrt(2t + 0.25) - 0.5 <= i
|
||||
|
||||
To account for fractional values, we set:
|
||||
i = ceil(sqrt(2t + 0.25) - 0.5)
|
||||
|
||||
To turn this into a zero-indexed row and work with zero-indexed t, we perform:
|
||||
i = ceil(sqrt(2(t+1) + 0.25) - 0.5) - 1
|
||||
= ceil(sqrt(2t + 2.25) - 0.5) - 1
|
||||
|
||||
Calculating column j given threadblock ID t and row i
|
||||
-----------------------------------------------------
|
||||
For a given row i, all threadblock IDs t in that row also satisfy the following:
|
||||
t > 1 + 2 + 3 + ... + (i-2) + (i-1)
|
||||
--> t > i(i-1)/2
|
||||
|
||||
Threadblock IDs within a given row are sequential, so the one-indexed column ID
|
||||
for one-indexed threadblock ID t and row i is:
|
||||
j = t - (i(i-1)/2)
|
||||
|
||||
The zero-indexed version becomes:
|
||||
j = (t+1) - (i(i+1)/2) -1
|
||||
= t - (i(i+1)/2)
|
||||
|
||||
Accounting for non-square grids
|
||||
-------------------------------
|
||||
Though the overall output problem size for Rank2K problems is guranteed to be square, the
|
||||
grids used in computing may not be square due to using non-square threadblock shapes. For
|
||||
example, a threadblock shape of 64x32 operating on a problem of output size 128x128 would
|
||||
result in a grid of 2x4 tiles.
|
||||
|
||||
This case can be handled by noting that the output resembles a square grid of 2x2 "macro tiles"
|
||||
each of which contains 2 "true tiles." We can thus first map a threadblock ID to its "macro tile"
|
||||
using the equations above, and then map it to the "true tile" within its "macro tile." In the example
|
||||
of a 2x4 grid, this mapping would look as follows:
|
||||
"Macro grid" "True grid"
|
||||
{0, 1} - 0 1 - -
|
||||
{2, 3} {4, 5} 2 3 4 5
|
||||
|
||||
A zero-indexed threadblock ID t is mapped to its "macro tile ID" t_macro as:
|
||||
t_macro = t // r
|
||||
Where r is the ratio of the maximum dimension of the grid to the minimum dimension of the grid
|
||||
(i.e., r = 4 / 2 = 2 in the previous example).
|
||||
|
||||
One uses t_macro and the calculations above to find the row and column in the square matrix to
|
||||
obtain i_macro and j_macro (zero-indexed). The mapping from (i_macro, j_macro) --> (i, j)
|
||||
is simply the following:
|
||||
if (ThreadblockShape::M > ThreadblockShape::N):
|
||||
r = ThreadblockShape::M / ThreadblockShape::N
|
||||
i = i_macro
|
||||
j = (j_macro * r) + (t % r)
|
||||
elif (ThreadblockShape::M < ThreadblockShape::N):
|
||||
r = ThreadblockShape::N / ThreadblockShape::M
|
||||
i = (i_macro * r) + (t % r)
|
||||
j = j_macro
|
||||
else:
|
||||
i = i_macro
|
||||
j = j_macro
|
||||
|
||||
Handling cases with grid dimensions that aren't multiples of eachother
|
||||
----------------------------------------------------------------------
|
||||
Even though threadblock shapes M and N are typically multiples of one another, the grid
|
||||
for a given problem may not have dimensions of the same ratio as that of the threadblock.
|
||||
For example, a problem of size 132x132 using a threadblock of shape 64x32 will result
|
||||
in a grid of 3x5 tiles. In this case, there is not an integer number of "true tiles"
|
||||
per "macro tile."
|
||||
|
||||
When this scenario arises, we simply pad the larger dimension of the grid such that
|
||||
there are an integer number of "true tiles" per "macro tile." Thus, the 3x5 grid in
|
||||
the example above will be treated as a 3x6 grid. Row and column positions for each
|
||||
tile are calculated as above. Any threadblocks that map to tiles that are outside the
|
||||
problem range or upper/lower triangular portion (e.g., (2, 5)) will exit early from
|
||||
this problem and may proceed to the next problem in the group.
|
||||
|
||||
Handling upper-triangular matrices
|
||||
----------------------------------
|
||||
The only modification needed for upper-triangular matrices is to swap i_macro and j_macro
|
||||
in the calculations above.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/blas3.h"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/matrix_coord.h"
|
||||
|
||||
#include "cutlass/gemm/kernel/grouped_problem_visitor.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace kernel {
|
||||
|
||||
namespace detail {
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
//
|
||||
// Helpers for calculating offsets for Rank2K problem visitor. These helpers specifically pertain
|
||||
// to the conversion from "macro tiles" to "true tiles" in the description above.
|
||||
//
|
||||
template <
|
||||
typename ThreadblockShape,
|
||||
typename Enable = void
|
||||
>
|
||||
struct Rank2KGroupedProblemVisitorOffsetHelper;
|
||||
|
||||
// Partial specialization for the case where threadblock shape M > threadblock shape N
|
||||
template <
|
||||
typename ThreadblockShape
|
||||
>
|
||||
struct Rank2KGroupedProblemVisitorOffsetHelper<
|
||||
ThreadblockShape,
|
||||
typename platform::enable_if< (ThreadblockShape::kM > ThreadblockShape::kN) >::type
|
||||
> {
|
||||
static_assert(ThreadblockShape::kM % ThreadblockShape::kN == 0,
|
||||
"Rank2KGroupedProblemVisitor with threadblock shape M > threadblock shape N "
|
||||
"requires that threadblock shape M be a multiple of threadblock shape N.");
|
||||
|
||||
static int32_t const kThreadblockSkewRatio = ThreadblockShape::kM / ThreadblockShape::kN;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static int32_t min_dim(cutlass::gemm::GemmCoord grid) {
|
||||
return grid.m();
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static int32_t macro_row_to_row(int32_t row, int32_t threadblock_id) {
|
||||
return row;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static int32_t macro_col_to_col(int32_t col, int32_t threadblock_id) {
|
||||
return (col * kThreadblockSkewRatio) + (threadblock_id % kThreadblockSkewRatio);
|
||||
}
|
||||
};
|
||||
|
||||
// Partial specialization for the case where threadblock shape M < threadblock shape N
|
||||
template <
|
||||
typename ThreadblockShape
|
||||
>
|
||||
struct Rank2KGroupedProblemVisitorOffsetHelper<
|
||||
ThreadblockShape,
|
||||
typename platform::enable_if< (ThreadblockShape::kM < ThreadblockShape::kN) >::type
|
||||
> {
|
||||
|
||||
static_assert(ThreadblockShape::kN % ThreadblockShape::kM == 0,
|
||||
"Rank2KGroupedProblemVisitor with threadblock shape M < threadblock shape N "
|
||||
"requires that threadblock shape N be a multiple of threadblock shape M.");
|
||||
|
||||
static int32_t const kThreadblockSkewRatio = ThreadblockShape::kN / ThreadblockShape::kM;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static int32_t min_dim(cutlass::gemm::GemmCoord grid) {
|
||||
return grid.n();
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static int32_t macro_row_to_row(int32_t row, int32_t threadblock_id) {
|
||||
return (row * kThreadblockSkewRatio) + (threadblock_id % kThreadblockSkewRatio);
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static int32_t macro_col_to_col(int32_t col, int32_t threadblock_id) {
|
||||
return col;
|
||||
}
|
||||
};
|
||||
|
||||
// Partial specialization for the case where threadblock shape M == threadblock shape N
|
||||
// In this case, macro tiles are equivalent to true tiles, so the conversions are
|
||||
// identity functions.
|
||||
template <
|
||||
typename ThreadblockShape
|
||||
>
|
||||
struct Rank2KGroupedProblemVisitorOffsetHelper<
|
||||
ThreadblockShape,
|
||||
typename platform::enable_if< (ThreadblockShape::kM == ThreadblockShape::kN) >::type
|
||||
> {
|
||||
|
||||
static int32_t const kThreadblockSkewRatio = 1;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static int32_t min_dim(cutlass::gemm::GemmCoord grid) {
|
||||
return grid.m();
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static int32_t macro_row_to_row(int32_t row, int32_t threadblock_id) {
|
||||
return row;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static int32_t macro_col_to_col(int32_t col, int32_t threadblock_id) {
|
||||
return col;
|
||||
}
|
||||
};
|
||||
|
||||
// Helper for correctly representing problem sizes in grouped kernels
|
||||
template <typename ThreadblockShape>
|
||||
struct Rank2KGroupedProblemSizeHelper {
|
||||
using OffsetHelper = Rank2KGroupedProblemVisitorOffsetHelper<ThreadblockShape>;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static int32_t tile_count(const cutlass::gemm::GemmCoord& grid) {
|
||||
// Return the number of tiles at or below the diagonal (or at and above
|
||||
// for mode kUpper). We do this by first calculating this value assuming
|
||||
// we have a square matrix of tiles of size `dim x dim` where `dim` is the
|
||||
// minimum among {grid.m(), grid.n()}. We then multiply the resulting value
|
||||
// by OffsetHelper::kThreadblockSkewRatio to account for cases in which there
|
||||
// are more tiles in one dimension than the other.
|
||||
int32_t dim = OffsetHelper::min_dim(grid);
|
||||
int32_t tiles_on_diagonal = dim;
|
||||
int32_t tiles_below_diagonal = ((dim * (dim - 1)) / 2);
|
||||
return (tiles_on_diagonal + tiles_below_diagonal) * OffsetHelper::kThreadblockSkewRatio;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static void possibly_transpose_problem(cutlass::gemm::GemmCoord& problem) {}
|
||||
};
|
||||
|
||||
} // namespace detail
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
//
|
||||
// Default problem visitor for fill modes kUpper and kLower.
|
||||
//
|
||||
template <typename ThreadblockShape,
|
||||
GroupScheduleMode GroupScheduleMode_,
|
||||
int PrefetchTileCount,
|
||||
int ThreadCount,
|
||||
cutlass::FillMode FillModeC>
|
||||
struct Rank2KGroupedProblemVisitor : public GroupedProblemVisitor<
|
||||
detail::Rank2KGroupedProblemSizeHelper<ThreadblockShape>,
|
||||
ThreadblockShape,
|
||||
GroupScheduleMode_,
|
||||
PrefetchTileCount,
|
||||
ThreadCount> {
|
||||
|
||||
static cutlass::FillMode const kFillModeC = FillModeC;
|
||||
|
||||
static_assert(kFillModeC == cutlass::FillMode::kLower || kFillModeC == cutlass::FillMode::kUpper,
|
||||
"Default Rank2KGroupedProblemVisitor requires fill mode of kLower or kUpper.");
|
||||
|
||||
using ProblemSizeHelper = detail::Rank2KGroupedProblemSizeHelper<ThreadblockShape>;
|
||||
using Base = GroupedProblemVisitor<ProblemSizeHelper,
|
||||
ThreadblockShape,
|
||||
GroupScheduleMode_,
|
||||
PrefetchTileCount,
|
||||
ThreadCount>;
|
||||
using OffsetHelper = typename ProblemSizeHelper::OffsetHelper;
|
||||
using Params = typename Base::Params;
|
||||
using SharedStorage = typename Base::SharedStorage;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
CUTLASS_DEVICE
|
||||
Rank2KGroupedProblemVisitor(
|
||||
Params const ¶ms_,
|
||||
SharedStorage &shared_storage_,
|
||||
int32_t block_idx
|
||||
): Base(params_, shared_storage_, block_idx)
|
||||
{}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
cutlass::gemm::GemmCoord threadblock_offset(int32_t threadblock_id) const {
|
||||
int32_t macro_id = threadblock_id / OffsetHelper::kThreadblockSkewRatio;
|
||||
int32_t macro_row = ceil(cutlass::fast_sqrt((2*macro_id) + 2.25) - 0.5) - 1;
|
||||
int32_t macro_col = macro_id - (((macro_row+1) * macro_row)/2);
|
||||
|
||||
if (kFillModeC == cutlass::FillMode::kUpper) {
|
||||
swap(macro_row, macro_col);
|
||||
}
|
||||
|
||||
int32_t row = OffsetHelper::macro_row_to_row(macro_row, threadblock_id);
|
||||
int32_t col = OffsetHelper::macro_col_to_col(macro_col, threadblock_id);
|
||||
|
||||
return cutlass::gemm::GemmCoord(row, col, 0);
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace kernel
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,129 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017 - 2022 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 Transpositions for Rank2K problems.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/blas3.h"
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace kernel {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace detail {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename ElementA_,
|
||||
typename LayoutA_,
|
||||
ComplexTransform TransformA,
|
||||
int AlignmentA,
|
||||
typename ElementB_,
|
||||
typename LayoutB_,
|
||||
ComplexTransform TransformB,
|
||||
int AlignmentB,
|
||||
typename LayoutC_,
|
||||
FillMode FillModeC_,
|
||||
bool Transpose
|
||||
>
|
||||
struct Rank2KMapArguments {
|
||||
using ElementA = ElementA_;
|
||||
using LayoutA = LayoutA_;
|
||||
static ComplexTransform const kTransformA = TransformA;
|
||||
static int const kAlignmentA = AlignmentA;
|
||||
using ElementB = ElementB_;
|
||||
using LayoutB = LayoutB_;
|
||||
static ComplexTransform const kTransformB = TransformB;
|
||||
static int const kAlignmentB = AlignmentB;
|
||||
using LayoutC = LayoutC_;
|
||||
static FillMode const kFillModeC = FillModeC_;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename ElementA_,
|
||||
typename LayoutA_,
|
||||
ComplexTransform TransformA,
|
||||
int AlignmentA,
|
||||
typename ElementB_,
|
||||
typename LayoutB_,
|
||||
ComplexTransform TransformB,
|
||||
int AlignmentB,
|
||||
typename LayoutC_,
|
||||
FillMode FillModeC_
|
||||
>
|
||||
struct Rank2KMapArguments<
|
||||
ElementA_,
|
||||
LayoutA_,
|
||||
TransformA,
|
||||
AlignmentA,
|
||||
ElementB_,
|
||||
LayoutB_,
|
||||
TransformB,
|
||||
AlignmentB,
|
||||
LayoutC_,
|
||||
FillModeC_,
|
||||
true
|
||||
> {
|
||||
using ElementA = ElementB_;
|
||||
using LayoutA = LayoutB_;
|
||||
static ComplexTransform const kTransformA = TransformB;
|
||||
static int const kAlignmentA = AlignmentB;
|
||||
using ElementB = ElementA_;
|
||||
using LayoutB = LayoutA_;
|
||||
static ComplexTransform const kTransformB = TransformA;
|
||||
static int const kAlignmentB = AlignmentA;
|
||||
using LayoutC = typename layout::LayoutTranspose<LayoutC_>::type;
|
||||
static FillMode const kFillModeC = InvertFillMode<FillModeC_>::mode;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
Reference in New Issue
Block a user