releaase 2.11 (#703)
This commit is contained in:
@@ -0,0 +1,154 @@
|
||||
/***************************************************************************************************
|
||||
* 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 Epilogue for threadblock scoped GEMMs using Tensor Ops.
|
||||
|
||||
The epilogue rearranges the result of a matrix product through shared memory to match canonical
|
||||
tensor layouts in global memory. Epilogues support conversion and reduction operations.
|
||||
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/array.h"
|
||||
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/epilogue/thread/linear_combination_clamp.h"
|
||||
#include "cutlass/epilogue/thread/conversion_op.h"
|
||||
#include "cutlass/epilogue/thread/reduction_op.h"
|
||||
|
||||
#include "cutlass/transform/threadblock/regular_tile_iterator_pitch_linear.h"
|
||||
|
||||
#include "cutlass/epilogue/warp/fragment_iterator_tensor_op.h"
|
||||
#include "cutlass/epilogue/warp/fragment_iterator_complex_tensor_op.h"
|
||||
#include "cutlass/epilogue/warp/tile_iterator_tensor_op.h"
|
||||
#include "cutlass/epilogue/warp/tile_iterator_tensor_op_mixed.h"
|
||||
#include "cutlass/epilogue/threadblock/default_thread_map_tensor_op.h"
|
||||
#include "cutlass/epilogue/threadblock/predicated_tile_iterator.h"
|
||||
#include "cutlass/epilogue/threadblock/shared_load_iterator.h"
|
||||
#include "cutlass/epilogue/threadblock/shared_load_iterator_mixed.h"
|
||||
|
||||
// #include "cutlass/epilogue/threadblock/epilogue.h"
|
||||
#include "cutlass/epilogue/threadblock/interleaved_epilogue.h"
|
||||
|
||||
#include "fused_bias_act_epilogue.h"
|
||||
#include "../warp/fused_bias_act_fragment_iterator_tensor_op.h"
|
||||
#include "output_tile_thread_map_for_fused_bias.h"
|
||||
#include "default_thread_map_tensor_op_for_fused_bias.h"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace epilogue {
|
||||
namespace threadblock {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Defines sensible defaults for epilogues for TensorOps.
|
||||
template <
|
||||
typename Shape_,
|
||||
typename WarpMmaTensorOp_,
|
||||
int PartitionsK,
|
||||
typename OutputOp_,
|
||||
int ElementsPerAccess
|
||||
>
|
||||
struct DefaultFusedBiasActEpilogueTensorOp {
|
||||
|
||||
using Shape = Shape_;
|
||||
using WarpMmaTensorOp = WarpMmaTensorOp_;
|
||||
static int const kPartitionsK = PartitionsK;
|
||||
using OutputOp = OutputOp_;
|
||||
static int const kElementsPerAccess = ElementsPerAccess;
|
||||
using ElementOutput = typename OutputOp::ElementOutput;
|
||||
using LayoutC = typename WarpMmaTensorOp::LayoutC;
|
||||
using ElementAccumulator = typename WarpMmaTensorOp::ElementC;
|
||||
|
||||
//
|
||||
// Thread map
|
||||
//
|
||||
|
||||
using OutputTileThreadMap = typename cutlass::epilogue::threadblock::DefaultThreadMapTensorOpForFusedBias<
|
||||
Shape,
|
||||
typename WarpMmaTensorOp::Shape,
|
||||
kPartitionsK,
|
||||
ElementOutput,
|
||||
kElementsPerAccess
|
||||
>::Type;
|
||||
|
||||
using OutputTileIterator = cutlass::epilogue::threadblock::PredicatedTileIterator<
|
||||
OutputTileThreadMap,
|
||||
ElementOutput
|
||||
>;
|
||||
|
||||
using AccumulatorFragmentIterator = typename std::conditional<is_complex<ElementOutput>::value,
|
||||
cutlass::epilogue::warp::FragmentIteratorComplexTensorOp<
|
||||
typename WarpMmaTensorOp::Shape,
|
||||
typename WarpMmaTensorOp::Policy::Operator::Shape,
|
||||
typename WarpMmaTensorOp::Policy::Operator::ElementC,
|
||||
typename WarpMmaTensorOp::Policy::Operator::FragmentC,
|
||||
LayoutC>,
|
||||
cutlass::epilogue::warp::FusedBiasActFragmentIteratorTensorOp<
|
||||
typename WarpMmaTensorOp::Shape,
|
||||
typename WarpMmaTensorOp::Policy::Operator::Shape,
|
||||
typename WarpMmaTensorOp::Policy::Operator::ElementC,
|
||||
typename WarpMmaTensorOp::Policy::Operator::FragmentC,
|
||||
LayoutC> >::type;
|
||||
|
||||
//
|
||||
// Define the epilogue
|
||||
//
|
||||
using Epilogue = cutlass::epilogue::threadblock::FusedBiasActEpilogue<
|
||||
Shape,
|
||||
WarpMmaTensorOp,
|
||||
kPartitionsK,
|
||||
OutputTileIterator,
|
||||
AccumulatorFragmentIterator,
|
||||
OutputOp
|
||||
>;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace epilogue
|
||||
} // namespace cutlass
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,113 @@
|
||||
/***************************************************************************************************
|
||||
* 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
|
||||
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/epilogue/threadblock/predicated_tile_iterator.h"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/layout/pitch_linear.h"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace epilogue {
|
||||
namespace threadblock {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Defines the optimal thread map for TensorOp accumulator layouts
|
||||
template <
|
||||
typename ThreadblockShape_,
|
||||
typename WarpShape_,
|
||||
int PartitionsK,
|
||||
typename Element_,
|
||||
int ElementsPerAccess
|
||||
>
|
||||
struct DefaultThreadMapTensorOpForFusedBias {
|
||||
|
||||
using ThreadblockShape = ThreadblockShape_;
|
||||
using WarpShape = WarpShape_;
|
||||
static int const kPartitionsK = PartitionsK;
|
||||
using Element = Element_;
|
||||
static int const kElementsPerAccess = ElementsPerAccess;
|
||||
|
||||
//
|
||||
// Definitions
|
||||
//
|
||||
|
||||
struct Detail {
|
||||
|
||||
/// Tensor Operations fundamentally perform operations on 8 rows
|
||||
static int const kTensorOpRows = 8;
|
||||
static int const kWarpSize = 32;
|
||||
|
||||
static_assert(
|
||||
!(ThreadblockShape::kM % WarpShape::kM) &&
|
||||
!(ThreadblockShape::kM % WarpShape::kM), "Divisibility");
|
||||
|
||||
/// Number of warps
|
||||
using WarpCount = gemm::GemmShape<
|
||||
ThreadblockShape::kM / WarpShape::kM,
|
||||
ThreadblockShape::kN / WarpShape::kN,
|
||||
kPartitionsK
|
||||
>;
|
||||
|
||||
/// Number of participating threads
|
||||
static int const kThreads = WarpCount::kCount * kWarpSize;
|
||||
};
|
||||
|
||||
//
|
||||
// ThreadMap
|
||||
//
|
||||
|
||||
/// ThreadMap to be used by epilogue::PredicatedTileIterator satisfying concept OutputTileThreadMap
|
||||
using Type = OutputTileOptimalThreadMapBiasAct <
|
||||
OutputTileShape<ThreadblockShape::kN, Detail::kTensorOpRows, Detail::WarpCount::kM, 1, 1>,
|
||||
OutputTileShape<1, WarpShape::kM / Detail::kTensorOpRows, 1, 1, WarpShape::kM / Detail::kTensorOpRows>,
|
||||
Detail::kThreads,
|
||||
kElementsPerAccess,
|
||||
sizeof_bits<Element>::value
|
||||
>;
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace epilogue
|
||||
} // namespace cutlass
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,222 @@
|
||||
/***************************************************************************************************
|
||||
* 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 Epilogue for threadblock scoped GEMMs using Tensor Ops.
|
||||
|
||||
The epilogue rearranges the result of a matrix product through shared memory to match canonical
|
||||
tensor layouts in global memory. Epilogues support conversion and reduction operations.
|
||||
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#if defined(__CUDACC_RTC__)
|
||||
#include <cuda/std/cassert>
|
||||
#else
|
||||
#include <assert.h>
|
||||
#endif
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/layout/vector.h"
|
||||
#include "cutlass/layout/tensor.h"
|
||||
#include "cutlass/tensor_coord.h"
|
||||
#include "cutlass/aligned_buffer.h"
|
||||
#include "cutlass/functional.h"
|
||||
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
|
||||
#include "cutlass/transform/pitch_linear_thread_map.h"
|
||||
#include "cutlass/transform/threadblock/regular_tile_iterator.h"
|
||||
|
||||
#include "cutlass/epilogue/threadblock/epilogue_base.h"
|
||||
#include "cutlass/epilogue/threadblock/predicated_tile_iterator.h"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace epilogue {
|
||||
namespace threadblock {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Epilogue operator without splitk
|
||||
template <
|
||||
typename Shape_, ///< Shape of threadblock tile (concept: GemmShape)
|
||||
typename WarpMmaOperator_, ///< Warp-level MMA operator (concept: gemm::warp::MmaTensorOp)
|
||||
int PartitionsK, ///< Number of partitions of the K dimension
|
||||
typename OutputTileIterator_, ///< Tile iterator reading and writing output tensors
|
||||
typename AccumulatorFragmentIterator_, ///< Fragment iterator selecting accumulators
|
||||
typename OutputOp_ ///< Output operator
|
||||
>
|
||||
class FusedBiasActEpilogue {
|
||||
|
||||
public:
|
||||
|
||||
using Shape = Shape_;
|
||||
using WarpMmaOperator = WarpMmaOperator_;
|
||||
static int const kPartitionsK = PartitionsK;
|
||||
using OutputTileIterator = OutputTileIterator_;
|
||||
using AccumulatorFragmentIterator = AccumulatorFragmentIterator_;
|
||||
using OutputOp = OutputOp_;
|
||||
|
||||
/// Output layout is always row-major
|
||||
using Layout = layout::RowMajor;
|
||||
using LongIndex = typename Layout::LongIndex;
|
||||
|
||||
/// The complete warp-level accumulator tile
|
||||
using AccumulatorTile = typename AccumulatorFragmentIterator::AccumulatorTile;
|
||||
|
||||
/// Output element
|
||||
using ElementOutput = typename OutputTileIterator::Element;
|
||||
|
||||
/// Output access size
|
||||
static int const kElementsPerAccess = OutputTileIterator::kElementsPerAccess;
|
||||
|
||||
|
||||
public:
|
||||
|
||||
|
||||
static_assert(OutputTileIterator::kElementsPerAccess, "OutputTileIterator::kElementsPerAccess must not be zero.");
|
||||
|
||||
static_assert(!(OutputTileIterator::Fragment::kElements % OutputTileIterator::kElementsPerAccess),
|
||||
"Divisibility");
|
||||
|
||||
public:
|
||||
|
||||
/// Constructor
|
||||
CUTLASS_DEVICE
|
||||
FusedBiasActEpilogue(
|
||||
){ }
|
||||
|
||||
/// Streams the result to global memory
|
||||
CUTLASS_DEVICE
|
||||
void operator()(
|
||||
OutputOp const &output_op, ///< Output operator
|
||||
AccumulatorTile &accumulators, ///< Complete warp-level accumulator tile
|
||||
AccumulatorTile & fused_bias_act_accumlators,
|
||||
OutputTileIterator source_iterator) { ///< Threadblock tile coordinate in GEMM (in units of threadblock tiles)
|
||||
|
||||
bool need_bias = output_op.is_source_needed();
|
||||
|
||||
if (need_bias)
|
||||
compute_source_needed_(output_op, accumulators, fused_bias_act_accumlators, source_iterator);
|
||||
else
|
||||
compute_source_no_needed_(output_op, accumulators, fused_bias_act_accumlators);
|
||||
|
||||
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void operator()(
|
||||
OutputOp const &output_op, ///< Output operator
|
||||
AccumulatorTile &accumulators, ///< Complete warp-level accumulator tile
|
||||
AccumulatorTile & fused_bias_act_accumlators) { ///< Threadblock tile coordinate in GEMM (in units of threadblock tiles)
|
||||
|
||||
compute_source_no_needed_(output_op, accumulators, fused_bias_act_accumlators);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void compute_source_needed_(
|
||||
OutputOp const &output_op, ///< Output operator
|
||||
AccumulatorTile &accumulators, ///< Complete warp-level accumulator tile
|
||||
AccumulatorTile & fused_bias_act_accumlators,
|
||||
OutputTileIterator source_iterator) { ///< Threadblock tile coordinate in GEMM (in units of threadblock tiles)
|
||||
|
||||
typename OutputTileIterator::Fragment source_fragment;
|
||||
|
||||
|
||||
source_fragment.clear();
|
||||
|
||||
AccumulatorFragmentIterator accum_fragment_iterator(accumulators);
|
||||
AccumulatorFragmentIterator fused_bias_act_fragment_iterator(fused_bias_act_accumlators);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int iter = 0; iter < OutputTileIterator::kIterations; ++iter) {
|
||||
|
||||
source_iterator.load(source_fragment);
|
||||
++source_iterator;
|
||||
|
||||
typename AccumulatorFragmentIterator::Fragment accum_fragment;
|
||||
|
||||
accum_fragment_iterator.load(accum_fragment);
|
||||
++accum_fragment_iterator;
|
||||
|
||||
typename AccumulatorFragmentIterator::Fragment fused_bias_act_fragment;
|
||||
fused_bias_act_fragment = output_op(accum_fragment, source_fragment);
|
||||
|
||||
fused_bias_act_fragment_iterator.store(fused_bias_act_fragment);
|
||||
++fused_bias_act_fragment_iterator;
|
||||
}
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void compute_source_no_needed_(
|
||||
OutputOp const &output_op, ///< Output operator
|
||||
AccumulatorTile &accumulators, ///< Complete warp-level accumulator tile
|
||||
AccumulatorTile & fused_bias_act_accumlators) { ///< Threadblock tile coordinate in GEMM (in units of threadblock tiles)
|
||||
|
||||
|
||||
AccumulatorFragmentIterator accum_fragment_iterator(accumulators);
|
||||
AccumulatorFragmentIterator fused_bias_act_fragment_iterator(fused_bias_act_accumlators);
|
||||
|
||||
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int iter = 0; iter < AccumulatorFragmentIterator::kIterations; ++iter) {
|
||||
|
||||
typename AccumulatorFragmentIterator::Fragment accum_fragment;
|
||||
|
||||
accum_fragment_iterator.load(accum_fragment);
|
||||
++accum_fragment_iterator;
|
||||
|
||||
typename AccumulatorFragmentIterator::Fragment fused_bias_act_fragment;
|
||||
fused_bias_act_fragment = output_op(accum_fragment);
|
||||
|
||||
fused_bias_act_fragment_iterator.store(fused_bias_act_fragment);
|
||||
++fused_bias_act_fragment_iterator;
|
||||
}
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
|
||||
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace epilogue
|
||||
} // namespace cutlass
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,311 @@
|
||||
/***************************************************************************************************
|
||||
* 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 Metaprogram for determining the mapping of output elements to threads for epilogue tiles.
|
||||
|
||||
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/matrix_shape.h"
|
||||
#include "cutlass/tensor_ref.h"
|
||||
#include "cutlass/fast_math.h"
|
||||
|
||||
#include "cutlass/epilogue/threadblock/output_tile_thread_map.h"
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace epilogue {
|
||||
namespace threadblock {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace detail {
|
||||
|
||||
/// RowArrangement determines how one or more warps cover a region of consecutive rows.
|
||||
template <
|
||||
typename Shape,
|
||||
int WarpsRemaining,
|
||||
int ElementsPerAccess,
|
||||
int ElementSize,
|
||||
bool Is2dTile
|
||||
>
|
||||
struct RowArrangementBiasAct;
|
||||
|
||||
/// RowArrangement in which each warp's access is a 1D tiled arrangement.
|
||||
template <
|
||||
typename Shape,
|
||||
int WarpsRemaining,
|
||||
int ElementsPerAccess,
|
||||
int ElementSize
|
||||
>
|
||||
struct RowArrangementBiasAct<Shape, WarpsRemaining, ElementsPerAccess, ElementSize, false> {
|
||||
static int const kWarpSize = 32;
|
||||
static int const kElementsPerAccess = ElementsPerAccess;
|
||||
static int const kElementSize = ElementSize;
|
||||
|
||||
static int const kIterationsRow = 1;
|
||||
static int const kDeltaRow = 1;
|
||||
static int const kIterationsColumn = Shape::kColumn / kElementsPerAccess / kWarpSize;
|
||||
static int const kDeltaColumn = kWarpSize * kElementsPerAccess;
|
||||
|
||||
static int const kAccessWidth = kWarpSize;
|
||||
static int const kAccessRows = 1;
|
||||
static int const kWarpPartitionsRow = 1;
|
||||
static int const kWarpPartitionsColumn = WarpsRemaining;
|
||||
};
|
||||
|
||||
/// RowArrangement in which each warp's access is a 2D tiled arrangement.
|
||||
template <
|
||||
typename Shape,
|
||||
int WarpsRemaining,
|
||||
int ElementsPerAccess,
|
||||
int ElementSize
|
||||
>
|
||||
struct RowArrangementBiasAct<Shape, WarpsRemaining, ElementsPerAccess, ElementSize, true> {
|
||||
|
||||
static int const kMemoryAccessSize = 4;//128;
|
||||
static int const kWarpSize = 32;
|
||||
|
||||
static int const kElementsPerAccess = ElementsPerAccess;
|
||||
static int const kElementSize = ElementSize;
|
||||
|
||||
struct Detail {
|
||||
static int const kShapeRow = Shape::kRow / WarpsRemaining;
|
||||
static int const kShapeWidth = Shape::kColumn / kElementsPerAccess;
|
||||
|
||||
static int const kTargetMemoryAccessWidth =
|
||||
kMemoryAccessSize / (kElementsPerAccess * kElementSize / 8);
|
||||
|
||||
static int const kTargetAccessRows = kWarpSize / kTargetMemoryAccessWidth;
|
||||
};
|
||||
|
||||
static int const kAccessWidth =
|
||||
(Detail::kTargetAccessRows > Detail::kShapeRow ?
|
||||
kWarpSize / Detail::kShapeRow
|
||||
: const_min(
|
||||
Detail::kShapeWidth,
|
||||
const_min(kWarpSize, kMemoryAccessSize / (kElementsPerAccess * kElementSize / 8))
|
||||
));
|
||||
|
||||
static int const kAccessRows =
|
||||
(Detail::kTargetAccessRows > Detail::kShapeRow ?
|
||||
Detail::kShapeRow
|
||||
: const_min(Shape::kRow, kWarpSize / kAccessWidth));
|
||||
|
||||
static int const kIterationsRow = Detail::kShapeRow / kAccessRows;
|
||||
static int const kDeltaRow = kAccessRows;
|
||||
|
||||
static int const kIterationsColumn = Detail::kShapeWidth / kAccessWidth;
|
||||
static int const kDeltaColumn = kAccessWidth * kElementsPerAccess;
|
||||
|
||||
static_assert( kAccessWidth * kElementsPerAccess <= Shape::kColumn, "Accessing too many elements per access");
|
||||
static_assert( kIterationsColumn > 0, "Iteration Count Column must be > 0" );
|
||||
static_assert( kIterationsRow > 0, "Iteration Count Row must be > 0" );
|
||||
|
||||
static int const kWarpPartitionsRow = 1;
|
||||
static int const kWarpPartitionsColumn = 1;
|
||||
};
|
||||
|
||||
}
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Template metaprogram for partitioning a 4D space across warps to achieve several performance
|
||||
/// objectives:
|
||||
///
|
||||
/// - coalesced memory accesses in units of 16 Byte lines
|
||||
/// - minimal address arithmetic
|
||||
/// - minimal predicate calculations
|
||||
///
|
||||
template <
|
||||
typename Shape_,
|
||||
typename Count_,
|
||||
int Threads,
|
||||
int ElementsPerAccess,
|
||||
int ElementSize
|
||||
>
|
||||
struct OutputTileOptimalThreadMapBiasAct {
|
||||
|
||||
using Shape = Shape_;
|
||||
using Count = Count_;
|
||||
|
||||
static int const kWarpSize = 32;
|
||||
static int const kThreads = Threads;
|
||||
static int const kWarpCount = kThreads / kWarpSize;
|
||||
|
||||
static int const kElementsPerAccess = ElementsPerAccess;
|
||||
static int const kElementSize = ElementSize;
|
||||
|
||||
//
|
||||
// Metaprogram computation
|
||||
//
|
||||
|
||||
struct Detail {
|
||||
|
||||
// Clusters
|
||||
static int const kIterationsCluster =
|
||||
((Shape::kCluster > kWarpCount) ?
|
||||
Shape::kCluster / kWarpCount
|
||||
: 1);
|
||||
|
||||
static int const kDeltaCluster =
|
||||
((Shape::kCluster > kWarpCount) ?
|
||||
Shape::kRow * Count::kRow * Shape::kGroup * Count::kGroup * Shape::kCluster / kIterationsCluster
|
||||
: 1);
|
||||
|
||||
static int const kCompactedDeltaCluster =
|
||||
((Shape::kCluster > kWarpCount) ?
|
||||
Shape::kRow * Shape::kGroup * Shape::kCluster / kIterationsCluster
|
||||
: 1);
|
||||
|
||||
static int const kWarpPartitionsCluster =
|
||||
((Shape::kCluster > kWarpCount) ?
|
||||
kWarpCount
|
||||
: kWarpCount / Shape::kCluster);
|
||||
|
||||
static int const kWarpsRemainingForGroups =
|
||||
((Shape::kCluster > kWarpCount) ? 1 : kWarpCount / Shape::kCluster);
|
||||
|
||||
// Groups
|
||||
static int const kIterationsGroup =
|
||||
((Shape::kGroup > kWarpsRemainingForGroups) ?
|
||||
Shape::kGroup / kWarpsRemainingForGroups
|
||||
: 1);
|
||||
|
||||
static int const kDeltaGroup =
|
||||
((Shape::kGroup > kWarpsRemainingForGroups) ?
|
||||
Shape::kRow * Count::kRow * Shape::kGroup / kIterationsGroup
|
||||
: 1);
|
||||
|
||||
static int const kCompactedDeltaGroup =
|
||||
((Shape::kGroup > kWarpsRemainingForGroups) ?
|
||||
Shape::kRow * Shape::kGroup / kIterationsGroup
|
||||
: 1);
|
||||
|
||||
static int const kWarpPartitionsGroup =
|
||||
((Shape::kGroup > kWarpsRemainingForGroups) ?
|
||||
1
|
||||
: kWarpsRemainingForGroups / Shape::kGroup);
|
||||
|
||||
static int const kWarpsRemainingForRows =
|
||||
((Shape::kGroup > kWarpsRemainingForGroups) ?
|
||||
1
|
||||
: kWarpsRemainingForGroups / Shape::kGroup);
|
||||
|
||||
// Rows
|
||||
using RowArrangement = detail::RowArrangementBiasAct<
|
||||
Shape,
|
||||
kWarpsRemainingForRows,
|
||||
kElementsPerAccess,
|
||||
kElementSize,
|
||||
(Shape::kRow > kWarpsRemainingForRows)
|
||||
>;
|
||||
|
||||
// Warp partitions
|
||||
using WarpPartitions = OutputTileShape<
|
||||
RowArrangement::kWarpPartitionsColumn,
|
||||
RowArrangement::kWarpPartitionsRow,
|
||||
kWarpPartitionsGroup,
|
||||
kWarpPartitionsCluster,
|
||||
1>;
|
||||
|
||||
static int const kAccessWidth = RowArrangement::kAccessWidth;
|
||||
static int const kAccessRows = RowArrangement::kAccessRows;
|
||||
};
|
||||
|
||||
//
|
||||
// Output
|
||||
//
|
||||
|
||||
using Iterations = OutputTileShape<
|
||||
Detail::RowArrangement::kIterationsColumn,
|
||||
Detail::RowArrangement::kIterationsRow,
|
||||
Detail::kIterationsGroup,
|
||||
Detail::kIterationsCluster,
|
||||
1>;
|
||||
|
||||
using Delta = OutputTileShape<
|
||||
Detail::RowArrangement::kDeltaColumn,
|
||||
Detail::RowArrangement::kDeltaRow,
|
||||
Detail::kDeltaGroup,
|
||||
Detail::kDeltaCluster,
|
||||
1>;
|
||||
|
||||
/// Initial offset function
|
||||
CUTLASS_HOST_DEVICE
|
||||
static MatrixCoord initial_offset(int thread_idx) {
|
||||
|
||||
int warp_idx = thread_idx / kWarpSize;
|
||||
int lane_idx = thread_idx % kWarpSize;
|
||||
|
||||
// Compute warp location
|
||||
int cluster_idx = warp_idx / Detail::WarpPartitions::kCluster;
|
||||
int residual_cluster = warp_idx % Detail::WarpPartitions::kCluster;
|
||||
|
||||
int group_idx = residual_cluster / Detail::WarpPartitions::kGroup;
|
||||
int residual_group = residual_cluster % Detail::WarpPartitions::kGroup;
|
||||
|
||||
int row_idx = residual_group / Detail::WarpPartitions::kRow;
|
||||
int col_idx = residual_group % Detail::WarpPartitions::kRow;
|
||||
|
||||
// Compute per-lane offset
|
||||
int lane_row_offset = lane_idx / Detail::kAccessWidth;
|
||||
int lane_col_offset = lane_idx % Detail::kAccessWidth;
|
||||
|
||||
// Compute coordinate in output space
|
||||
int cluster_offset = cluster_idx * Shape::kRow * Count::kRow * Shape::kGroup * Count::kGroup;
|
||||
int group_offset = group_idx * Shape::kRow * Count::kRow;
|
||||
int row_offset = row_idx * Iterations::kRow * Detail::kAccessRows;
|
||||
int column_offset = col_idx * Iterations::kColumn * Detail::kAccessWidth * kElementsPerAccess;
|
||||
|
||||
return MatrixCoord(
|
||||
cluster_offset + group_offset + row_offset + lane_row_offset,
|
||||
(column_offset + lane_col_offset) * kElementsPerAccess
|
||||
);
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace epilogue
|
||||
} // namespace cutlass
|
||||
@@ -0,0 +1,189 @@
|
||||
/***************************************************************************************************
|
||||
* 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 This defines a "fragment" iterator for visiting the fragments of an accumulator tile
|
||||
that participate in one warp-level store operation.
|
||||
|
||||
Typically, the accumulator tile is the largest single block of register-backed storage
|
||||
within the kernel. Storing it to memory is best accomplished by partitioning it into
|
||||
smaller tiles and storing these sequentially.
|
||||
|
||||
Round trips through shared memory during the Epilogue phase require partitioning, as
|
||||
shared memory capacity is typically insufficient for a threadblock's total accumulator
|
||||
size.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/layout/matrix.h"
|
||||
|
||||
#include "cutlass/epilogue/warp/tensor_op_policy.h"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace epilogue {
|
||||
namespace warp {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
///
|
||||
template <
|
||||
typename WarpShape, ///< shape of warp-level GEMM (concept: MatrixShape)
|
||||
typename OperatorShape, ///< matrix multiply operation shape (concept: gemm::GemmShape)
|
||||
typename OperatorElementC, ///< matrix multiply operation data type (concept: data type)
|
||||
typename OperatorFragmentC, ///< matrix multiply operation fragment (concept: Array)
|
||||
typename Layout ///< target shared memory layout
|
||||
>
|
||||
class FusedBiasActFragmentIteratorTensorOp;
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Partial specialization for row-major shared memory
|
||||
template <
|
||||
typename WarpShape_, ///< shape of the warp-level GEMM tile
|
||||
typename OperatorShape_, ///< matrix multiply operation shape (concept: gemm::GemmShape)
|
||||
typename OperatorElementC_, ///< matrix multiply operation data type (concept: data type)
|
||||
typename OperatorFragmentC_ ///< matrix multiply operation fragment (concept: Array)
|
||||
>
|
||||
class FusedBiasActFragmentIteratorTensorOp<WarpShape_, OperatorShape_, OperatorElementC_, OperatorFragmentC_, layout::RowMajor> {
|
||||
public:
|
||||
|
||||
using WarpShape = WarpShape_;
|
||||
using OperatorShape = OperatorShape_;
|
||||
using OperatorElementC = OperatorElementC_;
|
||||
using OperatorFragmentC = OperatorFragmentC_;
|
||||
using Layout = layout::RowMajor;
|
||||
|
||||
using Policy = TensorOpPolicy<WarpShape, OperatorShape, Layout>;
|
||||
|
||||
/// This is the fragment size produced by one access of the iterator.
|
||||
using Fragment = Array<
|
||||
OperatorElementC,
|
||||
Policy::OperatorCount::kColumn * Policy::kElementsPerAccess>;
|
||||
|
||||
/// This is the complete warp-level accumulator tile.
|
||||
using AccumulatorTile = Array<
|
||||
OperatorElementC,
|
||||
OperatorFragmentC::kElements * Policy::OperatorCount::kRow * Policy::OperatorCount::kColumn>;
|
||||
|
||||
using OutputAccumulatorTile = AccumulatorTile;
|
||||
|
||||
/// Number of times this iterator can be incremented
|
||||
static int const kIterations = Policy::kIterations;
|
||||
|
||||
private:
|
||||
|
||||
/// Internal access type
|
||||
using AccessType = Array<OperatorElementC, Policy::kElementsPerAccess>;
|
||||
|
||||
private:
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Accumulator tile
|
||||
AccessType *accumulators_;
|
||||
|
||||
/// Internal index
|
||||
int index_;
|
||||
|
||||
public:
|
||||
|
||||
/// Constructs an iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
FusedBiasActFragmentIteratorTensorOp(AccumulatorTile &accum):
|
||||
accumulators_(reinterpret_cast<AccessType *>(&accum)),
|
||||
index_(0) {
|
||||
}
|
||||
|
||||
/// Increments
|
||||
CUTLASS_HOST_DEVICE
|
||||
FusedBiasActFragmentIteratorTensorOp &operator++() {
|
||||
++index_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Decrements
|
||||
CUTLASS_HOST_DEVICE
|
||||
FusedBiasActFragmentIteratorTensorOp &operator--() {
|
||||
--index_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Loads a fragment from the referenced part of the accumulator tile
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load(Fragment &frag, int index_offset = 0) const {
|
||||
|
||||
int index = index_ + index_offset;
|
||||
|
||||
AccessType *frag_ptr = reinterpret_cast<AccessType *>(&frag);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int n = 0; n < Policy::OperatorCount::kColumn; ++n) {
|
||||
|
||||
int accumulator_access_offset =
|
||||
index + n * Policy::kAccumulatorColumnStride / Policy::kElementsPerAccess;
|
||||
|
||||
frag_ptr[n] = accumulators_[accumulator_access_offset];
|
||||
}
|
||||
}
|
||||
/// Stores a fragment from the referenced part of the accumulator tile
|
||||
CUTLASS_HOST_DEVICE
|
||||
void store(Fragment &frag, int index_offset = 0) const {
|
||||
|
||||
int index = index_ + index_offset;
|
||||
|
||||
AccessType *frag_ptr = reinterpret_cast<AccessType *>(&frag);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int n = 0; n < Policy::OperatorCount::kColumn; ++n) {
|
||||
|
||||
int accumulator_access_offset =
|
||||
index + n * Policy::kAccumulatorColumnStride / Policy::kElementsPerAccess;
|
||||
|
||||
accumulators_[accumulator_access_offset] = frag_ptr[n];
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace warp
|
||||
} // namespace epilogue
|
||||
} // namespace cutlass
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,427 @@
|
||||
/***************************************************************************************************
|
||||
* 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/matrix_shape.h"
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/layout/tensor.h"
|
||||
#include "cutlass/numeric_conversion.h"
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace warp {
|
||||
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
/// Size of the matrix to load (concept: MatrixShape)
|
||||
typename Shape_,
|
||||
/// Size of the accumulation tile shape (concept: MatrixShape)
|
||||
typename AccumulatorShape_,
|
||||
/// KBlocks columns to compute residual
|
||||
int KBlocksColumn_,
|
||||
/// Accumulator Element type
|
||||
typename ElementAccumulator_,
|
||||
/// Element type
|
||||
typename Element_,
|
||||
/// Layout of operand in memory
|
||||
typename Layout_,
|
||||
/// Shape of one matrix product operation (concept: MatrixShape)
|
||||
typename InstructionShape_,
|
||||
/// Whether beta is zero
|
||||
bool IsBetaZero_ >
|
||||
class MmaTensorOpPureFragmentIterator;
|
||||
|
||||
|
||||
// Partial specialization for col-major accumulator tile
|
||||
// And Element type is the same as Accumulator Element type
|
||||
|
||||
template <
|
||||
/// Shape of warp tile to load (concept: MatrixShape)
|
||||
typename Shape_,
|
||||
/// Shape of the warp accumulation tile (concept: MatrixShape)
|
||||
typename AccumulatorShape_,
|
||||
/// KBlocks columns to compute residual
|
||||
int KBlocksColumn_,
|
||||
/// Element type
|
||||
typename Element_,
|
||||
/// Shape of one matrix product operation (concept: MatrixShape)
|
||||
typename InstructionShape_>
|
||||
class MmaTensorOpPureFragmentIterator<Shape_, AccumulatorShape_, KBlocksColumn_, Element_, Element_,
|
||||
cutlass::layout::ColumnMajor,
|
||||
InstructionShape_, true> {
|
||||
public:
|
||||
|
||||
/// Shape of warp tile to load (concept: MatrixShape)
|
||||
using Shape = Shape_;
|
||||
|
||||
/// Shape of the warp accumulation tile (concept: MatrixShape)
|
||||
using AccumulatorShape = AccumulatorShape_;
|
||||
|
||||
/// KBlocks columns to compute residual
|
||||
static int const kKBlockColumn = KBlocksColumn_;
|
||||
|
||||
/// Element type
|
||||
using Element = Element_;
|
||||
|
||||
/// Layout of source tile
|
||||
using Layout = cutlass::layout::ColumnMajor;
|
||||
|
||||
/// Shape of one matrix product operation (concept: MatrixShape)
|
||||
using InstructionShape = InstructionShape_;
|
||||
|
||||
/// Whether beta is zero
|
||||
static bool const IsBetaZero = true;
|
||||
|
||||
/// Number of participating threads
|
||||
static int const kThreads = 32;
|
||||
|
||||
/// Internal structure of iterator - made public to enable introspection
|
||||
struct Policy {
|
||||
static_assert(
|
||||
!(Shape::kRow % InstructionShape::kM) &&
|
||||
!(Shape::kColumn % InstructionShape::kN),
|
||||
"Shape of warp-level Mma must be divisible by operator shape.");
|
||||
static_assert(
|
||||
!(AccumulatorShape::kRow % Shape::kRow) &&
|
||||
!(AccumulatorShape::kColumn % Shape::kColumn),
|
||||
"Shape of Warp Accumulator must be divisible by warp shape.");
|
||||
static_assert(
|
||||
!(kKBlockColumn % Shape::kColumn),
|
||||
"KBlock size must be divisible by warp shape.");
|
||||
|
||||
/// Number of times this iterator can be incremented
|
||||
static int const kIterations = AccumulatorShape::kCount / Shape::kCount;
|
||||
};
|
||||
|
||||
private:
|
||||
|
||||
static int const kElementsPerAccess = InstructionShape::kM * InstructionShape::kN / kThreads;
|
||||
|
||||
/// Number of mma operations performed by a warp
|
||||
using MmaIterations = MatrixShape<Shape::kRow / InstructionShape::kM,
|
||||
Shape::kColumn / InstructionShape::kN>;
|
||||
/// Number of mma operations performed by the entire accumulator
|
||||
using AccumulatorIterations = MatrixShape<AccumulatorShape::kRow / InstructionShape::kM,
|
||||
AccumulatorShape::kColumn / InstructionShape::kN>;
|
||||
|
||||
/// Number of K iterations
|
||||
static int const kKBlockIterations = (AccumulatorShape::kColumn + kKBlockColumn - 1) / kKBlockColumn;
|
||||
static int const kResidualColumn = AccumulatorShape::kColumn - (kKBlockIterations - 1) * kKBlockColumn;
|
||||
static int const kKBlockColumnIterations = kKBlockColumn / Shape::kColumn
|
||||
* (AccumulatorShape::kRow / Shape::kRow);
|
||||
static int const kResidualIndex = kResidualColumn / Shape::kColumn
|
||||
* (AccumulatorShape::kRow / Shape::kRow);
|
||||
|
||||
public:
|
||||
|
||||
//
|
||||
// Derived quantities
|
||||
//
|
||||
|
||||
/// Fragment object holding a thread's part of a tile
|
||||
/// This is the fragment size produced by one access of the iterator.
|
||||
using Fragment = Array<Element, Shape::kCount / kThreads>;
|
||||
|
||||
/// Accumulator Fragment object
|
||||
using AccumulatorFragment = Array<Element, AccumulatorShape::kCount / kThreads>;
|
||||
|
||||
|
||||
private:
|
||||
|
||||
/// Internal access type
|
||||
using AccessType = Array<Element, kElementsPerAccess>;
|
||||
|
||||
private:
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Accumulator tile
|
||||
AccessType const *accumulators_;
|
||||
|
||||
/// Internal index
|
||||
int index_;
|
||||
|
||||
/// Used to access residual tile first
|
||||
bool is_residual_tile_;
|
||||
|
||||
public:
|
||||
/// Constructs an iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpPureFragmentIterator(AccumulatorFragment const &accum)
|
||||
: accumulators_(reinterpret_cast<AccessType const *>(&accum)),
|
||||
index_(0), is_residual_tile_(true) {}
|
||||
|
||||
/// Add offset
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_offset(int index_offset) {
|
||||
index_ += index_offset;
|
||||
if(is_residual_tile_ && index_ >= kKBlockColumnIterations) {
|
||||
index_ = index_ - kKBlockColumnIterations + kResidualIndex;
|
||||
is_residual_tile_ = false;
|
||||
}
|
||||
}
|
||||
|
||||
/// Increments
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpPureFragmentIterator &operator++() {
|
||||
add_offset(1);
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Decrements
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpPureFragmentIterator &operator--() {
|
||||
add_offset(-1);
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Loads a fragment from the referenced part of the accumulator tile
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load(Fragment &frag) const {
|
||||
|
||||
AccessType src_fragment;
|
||||
src_fragment.clear();
|
||||
|
||||
|
||||
AccessType *frag_ptr = reinterpret_cast<AccessType *>(&frag);
|
||||
|
||||
int index_m = (index_ * MmaIterations::kRow) % AccumulatorIterations::kRow;
|
||||
int index_n = (index_ * MmaIterations::kRow) / AccumulatorIterations::kRow
|
||||
* MmaIterations::kColumn;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int n = 0; n < MmaIterations::kColumn; n++) {
|
||||
for (int m = 0; m < MmaIterations::kRow; m++) {
|
||||
int accumulator_access_offset =
|
||||
(n + index_n) * AccumulatorIterations::kRow + m + index_m;
|
||||
|
||||
frag_ptr[n * MmaIterations::kRow + m].clear();
|
||||
if(!(is_residual_tile_ && index_ >= kResidualIndex))
|
||||
frag_ptr[n * MmaIterations::kRow + m] = accumulators_[accumulator_access_offset];
|
||||
// frag_ptr[n * MmaIterations::kRow + m] = output_op(accumulators_[accumulator_access_offset], src_fragment);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
// Partial specialization for row-major accumulator tile
|
||||
|
||||
template <
|
||||
/// Shape of warp tile to load (concept: MatrixShape)
|
||||
typename Shape_,
|
||||
/// Shape of the warp accumulation tile (concept: MatrixShape)
|
||||
typename AccumulatorShape_,
|
||||
/// KBlocks columns to compute residual
|
||||
int KBlocksColumn_,
|
||||
/// Accumulator Element type
|
||||
typename ElementAccumulator_,
|
||||
/// Element type
|
||||
typename Element_,
|
||||
/// Shape of one matrix product operation (concept: MatrixShape)
|
||||
typename InstructionShape_>
|
||||
class MmaTensorOpPureFragmentIterator<Shape_, AccumulatorShape_, KBlocksColumn_, ElementAccumulator_, Element_,
|
||||
cutlass::layout::RowMajor,
|
||||
InstructionShape_, true> {
|
||||
public:
|
||||
|
||||
/// Shape of warp tile to load (concept: MatrixShape)
|
||||
using Shape = Shape_;
|
||||
|
||||
/// Shape of the warp accumulation tile (concept: MatrixShape)
|
||||
using AccumulatorShape = AccumulatorShape_;
|
||||
|
||||
/// KBlocks columns to compute residual
|
||||
static int const kKBlockColumn = KBlocksColumn_;
|
||||
|
||||
/// Accumulator Element type
|
||||
using ElementAccumulator = ElementAccumulator_;
|
||||
|
||||
/// Element type
|
||||
using Element = Element_;
|
||||
|
||||
/// Layout of source tile
|
||||
using Layout = cutlass::layout::RowMajor;
|
||||
|
||||
/// Shape of one matrix product operation (concept: MatrixShape)
|
||||
using InstructionShape = InstructionShape_;
|
||||
|
||||
/// Whether beta is zero
|
||||
static bool const IsBetaZero = true;
|
||||
|
||||
/// Number of participating threads
|
||||
static int const kThreads = 32;
|
||||
|
||||
/// Internal structure of iterator - made public to enable introspection
|
||||
struct Policy {
|
||||
static_assert(
|
||||
!(Shape::kRow % InstructionShape::kM) &&
|
||||
!(Shape::kColumn % InstructionShape::kN),
|
||||
"Shape of warp-level Mma must be divisible by operator shape.");
|
||||
static_assert(
|
||||
!(AccumulatorShape::kRow % Shape::kRow) &&
|
||||
!(AccumulatorShape::kColumn % Shape::kColumn),
|
||||
"Shape of Warp Accumulator must be divisible by warp shape.");
|
||||
static_assert(
|
||||
!(kKBlockColumn % Shape::kColumn),
|
||||
"KBlock size must be divisible by warp shape.");
|
||||
|
||||
/// Number of times this iterator can be incremented
|
||||
static int const kIterations = AccumulatorShape::kCount / Shape::kCount;
|
||||
};
|
||||
|
||||
private:
|
||||
|
||||
static int const kElementsPerAccess = InstructionShape::kM * InstructionShape::kN / kThreads;
|
||||
|
||||
/// Number of mma operations performed by a warp
|
||||
using MmaIterations = MatrixShape<Shape::kRow / InstructionShape::kM,
|
||||
Shape::kColumn / InstructionShape::kN>;
|
||||
/// Number of mma operations performed by the entire accumulator
|
||||
using AccumulatorIterations = MatrixShape<AccumulatorShape::kRow / InstructionShape::kM,
|
||||
AccumulatorShape::kColumn / InstructionShape::kN>;
|
||||
|
||||
/// Number of K iterations
|
||||
static int const kKBlockIterations = (AccumulatorShape::kColumn + kKBlockColumn - 1) / kKBlockColumn;
|
||||
static int const kResidualColumn = AccumulatorShape::kColumn - (kKBlockIterations - 1) * kKBlockColumn;
|
||||
static int const kKBlockColumnIterations = kKBlockColumn / Shape::kColumn
|
||||
* (AccumulatorShape::kRow / Shape::kRow);
|
||||
static int const kResidualIndex = kResidualColumn / Shape::kColumn
|
||||
* (AccumulatorShape::kRow / Shape::kRow);
|
||||
|
||||
public:
|
||||
|
||||
//
|
||||
// Derived quantities
|
||||
//
|
||||
|
||||
/// Fragment object holding a thread's part of a tile
|
||||
/// This is the fragment size produced by one access of the iterator.
|
||||
using Fragment = Array<Element, Shape::kCount / kThreads>;
|
||||
|
||||
/// Accumulator Fragment object
|
||||
using AccumulatorFragment = Array<ElementAccumulator, AccumulatorShape::kCount / kThreads>;
|
||||
|
||||
|
||||
private:
|
||||
|
||||
/// Internal access type
|
||||
using AccessType = Array<ElementAccumulator, kElementsPerAccess>;
|
||||
using FragmentAccessType = Array<Element, kElementsPerAccess>;
|
||||
|
||||
private:
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Accumulator tile
|
||||
AccessType const *accumulators_;
|
||||
|
||||
/// Internal index
|
||||
int index_;
|
||||
|
||||
/// Used to access residual tile first
|
||||
bool is_residual_tile_;
|
||||
|
||||
public:
|
||||
/// Constructs an iterator
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpPureFragmentIterator(AccumulatorFragment const &accum)
|
||||
: accumulators_(reinterpret_cast<AccessType const *>(&accum)),
|
||||
index_(0), is_residual_tile_(true) {}
|
||||
|
||||
/// Add offset
|
||||
CUTLASS_HOST_DEVICE
|
||||
void add_offset(int index_offset) {
|
||||
index_ += index_offset;
|
||||
if(is_residual_tile_ && index_ >= kKBlockColumnIterations) {
|
||||
index_ = index_ - kKBlockColumnIterations + kResidualIndex;
|
||||
is_residual_tile_ = false;
|
||||
}
|
||||
}
|
||||
|
||||
/// Increments
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpPureFragmentIterator &operator++() {
|
||||
add_offset(1);
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Decrements
|
||||
CUTLASS_HOST_DEVICE
|
||||
MmaTensorOpPureFragmentIterator &operator--() {
|
||||
add_offset(-1);
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// Loads a fragment from the referenced part of the accumulator tile
|
||||
CUTLASS_HOST_DEVICE
|
||||
void load(Fragment &frag) const {
|
||||
|
||||
|
||||
FragmentAccessType src_fragment;
|
||||
src_fragment.clear();
|
||||
|
||||
FragmentAccessType *frag_ptr = reinterpret_cast<FragmentAccessType *>(&frag);
|
||||
|
||||
int index_m = (index_ * MmaIterations::kRow) % AccumulatorIterations::kRow;
|
||||
int index_n = (index_ * MmaIterations::kRow) / AccumulatorIterations::kRow
|
||||
* MmaIterations::kColumn;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int m = 0; m < MmaIterations::kRow; m++) {
|
||||
for (int n = 0; n < MmaIterations::kColumn; n++) {
|
||||
int accumulator_access_offset =
|
||||
(m + index_m) * AccumulatorIterations::kColumn + n + index_n;
|
||||
|
||||
frag_ptr[m * MmaIterations::kColumn + n].clear();
|
||||
if(!(is_residual_tile_ && index_ >= kResidualIndex))
|
||||
frag_ptr[m * MmaIterations::kColumn + n] = (accumulators_[accumulator_access_offset]);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace warp
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
Reference in New Issue
Block a user