CUTLASS 2.10 (#615)

Co-authored-by: Aniket Shivam <ashivam@nvidia.com>
This commit is contained in:
ANIKET SHIVAM
2022-09-03 18:48:46 -04:00
committed by GitHub
co-authored by Aniket Shivam
parent ca23ff7924
commit b72cbf957d
289 changed files with 43708 additions and 2513 deletions
@@ -0,0 +1,140 @@
/***************************************************************************************************
* 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 Templates implementing warp-level per channel scale+bias+relu before
matrix multiply-accumulate operations targeting Tensor Cores.
*/
#pragma once
#include "cutlass/cutlass.h"
#include "cutlass/array.h"
#include "cutlass/platform/platform.h"
#include "cutlass/numeric_conversion.h"
#include "cutlass/numeric_types.h"
#include "cutlass/matrix_shape.h"
#include "cutlass/arch/memory_sm75.h"
#include "cutlass/arch/mma_sm75.h"
#include "cutlass/arch/mma_sm80.h"
#include "cutlass/gemm/gemm.h"
#include "cutlass/gemm/warp/mma.h"
#include "cutlass/gemm/warp/mma_tensor_op_policy.h"
#include "cutlass/gemm/warp/mma_tensor_op_tile_iterator.h"
#include "cutlass/gemm/warp/mma_tensor_op_tile_iterator_sm80.h"
/////////////////////////////////////////////////////////////////////////////////////////////////
namespace cutlass {
namespace gemm {
namespace warp {
/////////////////////////////////////////////////////////////////////////////////////////////////
template <typename FragmentActivations, typename FragmentVarMean, typename FragmentGammaBeta>
struct LayernormScaleBiasTransform {
using T = typename FragmentActivations::Element;
static int const NumActivations = FragmentActivations::kElements;
static int const NumVarMean = FragmentVarMean::kElements;
static int const NumGammaBeta = FragmentGammaBeta::kElements;
static int const MmaElements = 2;
// One element has one scale and one bias
static int const MmaScaleBiasPair = 2;
// 16816 has 2 columns and 2 rows
static int const MmaCols = 2;
static int const MmaRows = 2;
using MmaOperand = Array<T, MmaElements>;
using VarMeanOperand = Array<__half2, MmaScaleBiasPair>;
using GammaBetaOperand = Array<T, MmaElements * MmaScaleBiasPair>;
CUTLASS_DEVICE
void transform(MmaOperand &activations,
VarMeanOperand const &var_mean,
GammaBetaOperand const &gamma_beta) {
#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 800))
uint32_t *ptr_activations = reinterpret_cast<uint32_t *>(&activations);
uint32_t const *ptr_var_mean = reinterpret_cast<uint32_t const *>(&var_mean);
uint32_t const *ptr_gamma_beta = reinterpret_cast<uint32_t const *>(&gamma_beta);
// Apply per channel scale+bias+relu if the data is not a special NaN
// (0x7eff). If it is a special NaN (0x7eff), hard code the output to 0.
// We assumes the pair of FP16 are either both inbound or both out-of-bound.
// It requires C to be an even number.
asm volatile(
"{\n\t"
" fma.rn.f16x2 %0, %1, %2, %3;\n"
" fma.rn.f16x2 %0, %4, %0, %5;\n"
"}\n"
: "=r"(ptr_activations[0])
: "r"(ptr_var_mean[0]), "r"(ptr_activations[0]),
"r"(ptr_var_mean[1]),
"r"(ptr_gamma_beta[0]), "r"(ptr_gamma_beta[1]));
#else
// TODO: write emulation code
assert(0);
#endif
}
CUTLASS_DEVICE
void operator()(FragmentActivations &activations,
FragmentVarMean const &var_mean,
FragmentGammaBeta const &gamma_beta) {
MmaOperand *ptr_activations = reinterpret_cast<MmaOperand *>(&activations);
VarMeanOperand const *ptr_var_mean =
reinterpret_cast<VarMeanOperand const *>(&var_mean);
GammaBetaOperand const *ptr_gamma_beta =
reinterpret_cast<GammaBetaOperand const *>(&gamma_beta);
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < (NumActivations / MmaElements); ++i) {
transform(ptr_activations[i],
ptr_var_mean[i / (MmaCols * MmaRows) * MmaRows + i % MmaRows],
ptr_gamma_beta[(i / MmaScaleBiasPair) % MmaCols]);
}
}
};
/////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace warp
} // namespace gemm
} // namespace cutlass
/////////////////////////////////////////////////////////////////////////////////////////////////
@@ -46,7 +46,6 @@
#include "cutlass/arch/memory_sm75.h"
#include "cutlass/arch/mma_sm75.h"
#include "cutlass/arch/mma_sm80.h"
#include "cutlass/gemm/gemm.h"
#include "cutlass/gemm/warp/mma.h"
@@ -251,6 +250,8 @@ template <
ComplexTransform TransformA = ComplexTransform::kNone,
/// Complex transform on B operand
ComplexTransform TransformB = ComplexTransform::kNone,
/// Do source operands need more than one elements
bool GeneralizedOperatorElements = false,
/// Used for partial specialization
typename Enable = bool
>
@@ -279,9 +280,7 @@ template <
/// Complex transform on A operand
ComplexTransform TransformA,
/// Complex transform on B operand
ComplexTransform TransformB,
/// Used for partial specialization
typename Enable
ComplexTransform TransformB
>
class MmaComplexTensorOp<
Shape_,
@@ -293,8 +292,7 @@ class MmaComplexTensorOp<
LayoutC_,
Policy_,
TransformA,
TransformB,
Enable> {
TransformB> {
public:
/// Shape of warp-level matrix operation (concept: GemmShape)
using Shape = Shape_;
@@ -565,9 +563,7 @@ template <
/// Complex transform on A operand
ComplexTransform TransformA,
/// Complex transform on B operand
ComplexTransform TransformB,
/// Used for partial specialization
typename Enable
ComplexTransform TransformB
>
class MmaComplexTensorOp<
Shape_,
@@ -579,8 +575,7 @@ class MmaComplexTensorOp<
LayoutC_,
Policy_,
TransformA,
TransformB,
Enable> {
TransformB> {
public:
/// Shape of warp-level matrix operation (concept: GemmShape)
using Shape = Shape_;
@@ -618,7 +618,7 @@ public:
/// Fragment object holding a thread's part of a tile
using Fragment = Array<Element, ThreadShape::kCount>;
private:
protected:
/// Internal reference
cutlass::TensorRef<Array<Element, Policy::LaneMmaShape::kN>, layout::RowMajor> ref_;
+16
View File
@@ -295,6 +295,14 @@ public:
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ < 800)
// Serpentine visitation order maximizing reuse of Rb
// The visitation order is like
// _
// | | | |
// | | | |
// |_| |_|
//
// Down Up Down Up
CUTLASS_PRAGMA_UNROLL
for (int n = 0; n < MmaIterations::kColumn; ++n) {
@@ -320,6 +328,14 @@ public:
}
#elif defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 800)
// Serpentine visitation order maximizing reuse of Ra
// The visitation order is like
// _________
// _________|
// |_________
// __________|
//
// Right Left Right Left
CUTLASS_PRAGMA_UNROLL
for (int m = 0; m < MmaIterations::kRow; ++m) {
@@ -146,6 +146,21 @@ public:
static bool const kReduceKForA = ReduceKForA_;
static_assert(platform::is_same<ElementA, cutlass::half_t>::value ||
platform::is_same<ElementA, cutlass::bfloat16_t>::value,
"ElementA needs to be fp16 or bf16.");
static_assert(platform::is_same<ElementB, cutlass::half_t>::value ||
platform::is_same<ElementB, cutlass::bfloat16_t>::value,
"ElementB needs to be fp16 or bf16.");
static_assert(platform::is_same<InstructionShape,
cutlass::gemm::GemmShape<16, 8, 16>>::value,
"Only supports 16x8x16 tensor core instruction.");
static_assert(!AccumulatorsInRowMajor,
"Only calls tensor core instructions in column major.");
public:
/// Iterates over the A operand in memory
@@ -226,30 +241,7 @@ public:
MmaOperandC *ptr_D = reinterpret_cast<MmaOperandC *>(&D);
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ < 800)
// Serpentine visitation order maximizing reuse of Rb
CUTLASS_PRAGMA_UNROLL
for (int n = 0; n < MmaIterations::kColumn; ++n) {
CUTLASS_PRAGMA_UNROLL
for (int m = 0; m < MmaIterations::kRow; ++m) {
int m_serpentine = ((n % 2) ? (MmaIterations::kRow - 1 - m) : m);
if (AccumulatorsInRowMajor) { // matrix B is reordered
mma(
ptr_D[n + m_serpentine * MmaIterations::kColumn],
ptr_A[m_serpentine],
ptr_B[n],
ptr_D[n + m_serpentine * MmaIterations::kColumn]);
} else {
mma(
ptr_D[m_serpentine + n * MmaIterations::kRow],
ptr_A[m_serpentine],
ptr_B[n],
ptr_D[m_serpentine + n * MmaIterations::kRow]);
}
}
}
assert(0);
#elif defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 800)
// Serpentine visitation order maximizing reuse of Ra
CUTLASS_PRAGMA_UNROLL
@@ -260,25 +252,21 @@ public:
int n_serpentine = ((m % 2) ? (MmaIterations::kColumn - 1 - n) : n);
if (AccumulatorsInRowMajor) { // matrix B is reordered
mma(
ptr_D[n_serpentine + m * MmaIterations::kColumn],
mma(ptr_D[m + n_serpentine * MmaIterations::kRow],
ptr_A[m],
ptr_B[n_serpentine],
ptr_D[n_serpentine + m * MmaIterations::kColumn]);
} else {
mma(ptr_D[m + n_serpentine * MmaIterations::kRow],
ptr_A[m],
ptr_B[n_serpentine],
ptr_D[m + n_serpentine * MmaIterations::kRow]);
ptr_D[m + n_serpentine * MmaIterations::kRow]);
if (!kReduceKForA && m == 0) {
// gemm_k_reduction[n_serpentine] += float(B[n_serpentine * 4]);
// gemm_k_reduction[n_serpentine] += float(B[n_serpentine * 4 + 1]);
// gemm_k_reduction[n_serpentine] += float(B[n_serpentine * 4 + 2]);
// gemm_k_reduction[n_serpentine] += float(B[n_serpentine * 4 + 3]);
if (!kReduceKForA && m == 0) {
#if 0
gemm_k_reduction[n_serpentine] += float(B[n_serpentine * 4]);
gemm_k_reduction[n_serpentine] += float(B[n_serpentine * 4 + 1]);
gemm_k_reduction[n_serpentine] += float(B[n_serpentine * 4 + 2]);
gemm_k_reduction[n_serpentine] += float(B[n_serpentine * 4 + 3]);
#else
uint32_t const *tmp = reinterpret_cast<uint32_t const *>(&B);
uint32_t const *tmp = reinterpret_cast<uint32_t const *>(&B);
if (platform::is_same<ElementB, cutlass::half_t>::value) {
asm volatile(
"{\n\t"
" .reg .f16 low, high;\n\t"
@@ -296,48 +284,99 @@ public:
"}\n\t"
: "+f"(gemm_k_reduction[n_serpentine])
: "r"(tmp[n_serpentine * 2]), "r"(tmp[n_serpentine * 2 + 1]));
} else if (platform::is_same<ElementB, cutlass::bfloat16_t>::value) {
asm volatile(
"{\n\t"
" .reg .f32 tmp;\n\t"
" shl.b32 tmp, %1, 16;\n\t"
" add.f32 %0, tmp, %0;\n\t"
" and.b32 tmp, %1, 0xffff0000;\n\t"
" add.f32 %0, tmp, %0;\n\t"
" shl.b32 tmp, %2, 16;\n\t"
" add.f32 %0, tmp, %0;\n\t"
" and.b32 tmp, %2, 0xffff0000;\n\t"
" add.f32 %0, tmp, %0;\n\t"
"}\n\t"
: "+f"(gemm_k_reduction[n_serpentine])
: "r"(tmp[n_serpentine * 2]), "r"(tmp[n_serpentine * 2 + 1]));
} else {
assert(0);
}
#endif
}
if (kReduceKForA && (n == 0)) {
// gemm_k_reduction[m * 2] += float(A[m * 8]);
// gemm_k_reduction[m * 2] += float(A[m * 8 + 1]);
// gemm_k_reduction[m * 2] += float(A[m * 8 + 4]);
// gemm_k_reduction[m * 2] += float(A[m * 8 + 5]);
//
// gemm_k_reduction[m * 2 + 1] += float(A[m * 8 + 2]);
// gemm_k_reduction[m * 2 + 1] += float(A[m * 8 + 3]);
// gemm_k_reduction[m * 2 + 1] += float(A[m * 8 + 6]);
// gemm_k_reduction[m * 2 + 1] += float(A[m * 8 + 7]);
#if 0
gemm_k_reduction[m * 2] += float(A[m * 8]);
gemm_k_reduction[m * 2] += float(A[m * 8 + 1]);
gemm_k_reduction[m * 2] += float(A[m * 8 + 4]);
gemm_k_reduction[m * 2] += float(A[m * 8 + 5]);
gemm_k_reduction[m * 2 + 1] += float(A[m * 8 + 2]);
gemm_k_reduction[m * 2 + 1] += float(A[m * 8 + 3]);
gemm_k_reduction[m * 2 + 1] += float(A[m * 8 + 6]);
gemm_k_reduction[m * 2 + 1] += float(A[m * 8 + 7]);
#else
uint32_t const *tmp = reinterpret_cast<uint32_t const *>(&A);
asm volatile(
"{\n\t"
" .reg .f16 low, high;\n\t"
" .reg .f32 tmp;\n\t"
" mov.b32 {low, high}, %2;\n\t"
" cvt.f32.f16 tmp, low;\n\t"
" add.f32 %0, tmp, %0;\n\t"
" cvt.f32.f16 tmp, high;\n\t"
" add.f32 %0, tmp, %0;\n\t"
" mov.b32 {low, high}, %3;\n\t"
" cvt.f32.f16 tmp, low;\n\t"
" add.f32 %1, tmp, %1;\n\t"
" cvt.f32.f16 tmp, high;\n\t"
" add.f32 %1, tmp, %1;\n\t"
" mov.b32 {low, high}, %4;\n\t"
" cvt.f32.f16 tmp, low;\n\t"
" add.f32 %0, tmp, %0;\n\t"
" cvt.f32.f16 tmp, high;\n\t"
" add.f32 %0, tmp, %0;\n\t"
" mov.b32 {low, high}, %5;\n\t"
" cvt.f32.f16 tmp, low;\n\t"
" add.f32 %1, tmp, %1;\n\t"
" cvt.f32.f16 tmp, high;\n\t"
" add.f32 %1, tmp, %1;\n\t"
"}\n\t"
: "+f"(gemm_k_reduction[m * 2]), "+f"(gemm_k_reduction[m * 2 + 1])
: "r"(tmp[m * 4]), "r"(tmp[m * 4 + 1]),"r"(tmp[m * 4 + 2]), "r"(tmp[m * 4 + 3]));
if (platform::is_same<ElementA, cutlass::half_t>::value) {
asm volatile(
"{\n\t"
" .reg .f16 low, high;\n\t"
" .reg .f32 tmp;\n\t"
" mov.b32 {low, high}, %2;\n\t"
" cvt.f32.f16 tmp, low;\n\t"
" add.f32 %0, tmp, %0;\n\t"
" cvt.f32.f16 tmp, high;\n\t"
" add.f32 %0, tmp, %0;\n\t"
" mov.b32 {low, high}, %3;\n\t"
" cvt.f32.f16 tmp, low;\n\t"
" add.f32 %1, tmp, %1;\n\t"
" cvt.f32.f16 tmp, high;\n\t"
" add.f32 %1, tmp, %1;\n\t"
" mov.b32 {low, high}, %4;\n\t"
" cvt.f32.f16 tmp, low;\n\t"
" add.f32 %0, tmp, %0;\n\t"
" cvt.f32.f16 tmp, high;\n\t"
" add.f32 %0, tmp, %0;\n\t"
" mov.b32 {low, high}, %5;\n\t"
" cvt.f32.f16 tmp, low;\n\t"
" add.f32 %1, tmp, %1;\n\t"
" cvt.f32.f16 tmp, high;\n\t"
" add.f32 %1, tmp, %1;\n\t"
"}\n\t"
: "+f"(gemm_k_reduction[m * 2]), "+f"(gemm_k_reduction[m * 2 + 1])
: "r"(tmp[m * 4]), "r"(tmp[m * 4 + 1]),"r"(tmp[m * 4 + 2]), "r"(tmp[m * 4 + 3]));
} else if (platform::is_same<ElementA, cutlass::bfloat16_t>::value) {
asm volatile(
"{\n\t"
" .reg .f32 tmp;\n\t"
" shl.b32 tmp, %2, 16;\n\t"
" add.f32 %0, tmp, %0;\n\t"
" and.b32 tmp, %2, 0xffff0000;\n\t"
" add.f32 %0, tmp, %0;\n\t"
" shl.b32 tmp, %3, 16;\n\t"
" add.f32 %1, tmp, %1;\n\t"
" and.b32 tmp, %3, 0xffff0000;\n\t"
" add.f32 %1, tmp, %1;\n\t"
" shl.b32 tmp, %4, 16;\n\t"
" add.f32 %0, tmp, %0;\n\t"
" and.b32 tmp, %4, 0xffff0000;\n\t"
" add.f32 %0, tmp, %0;\n\t"
" shl.b32 tmp, %5, 16;\n\t"
" add.f32 %1, tmp, %1;\n\t"
" and.b32 tmp, %5, 0xffff0000;\n\t"
" add.f32 %1, tmp, %1;\n\t"
"}\n\t"
: "+f"(gemm_k_reduction[m * 2]), "+f"(gemm_k_reduction[m * 2 + 1])
: "r"(tmp[m * 4]), "r"(tmp[m * 4 + 1]),"r"(tmp[m * 4 + 2]), "r"(tmp[m * 4 + 3]));
} else {
assert(0);
}
#endif
}
}
}
@@ -0,0 +1,574 @@
/***************************************************************************************************
* 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 Defines iterators used by warp-level loading scale and bias vectors.
Every scale/bias data only needs to be loaded once for every channel.
*/
#pragma once
#include "cutlass/cutlass.h"
#include "cutlass/array.h"
#include "cutlass/numeric_types.h"
#include "cutlass/tensor_ref.h"
#include "cutlass/matrix_shape.h"
#include "cutlass/arch/memory_sm75.h"
#include "cutlass/gemm/gemm.h"
#include "cutlass/layout/matrix.h"
#include "cutlass/layout/tensor.h"
#include "cutlass/layout/pitch_linear.h"
#include "cutlass/layout/tensor_op_multiplicand_sm75.h"
#include "cutlass/platform/platform.h"
#include "cutlass/fast_math.h"
////////////////////////////////////////////////////////////////////////////////
namespace cutlass {
namespace gemm {
namespace warp {
////////////////////////////////////////////////////////////////////////////////
template <
/// Size of the matrix to load (concept: MatrixShape)
typename Shape_,
/// Data type of A elements
typename Element_,
/// Layout of operand
typename Layout_,
/// Shape of one matrix production operation (concept: GemmShape)
typename InstructionShape_,
/// Policy of the details of LDSM shape and iterations
typename Policy_,
/// Number of threads participating in one matrix operation
int Threads,
/// Number of partitions along K dimension
int PartitionsK_ = 1>
class ScaleBiasTileIterator;
////////////////////////////////////////////////////////////////////////////////
/// This tile iterator is specialized for 32-thread TensorOps. It uses LDSM to
/// load from shared memory and therefore must be initialized with a TensorRef
/// to shared memory.
///
/// Satisfies:
/// ReadableRandomAccessContiguousTileIteratorConcept
///
template <
/// Size of the matrix to load (concept: PitchLinearShape)
typename Shape_,
/// Data type of elements
typename Element_,
/// Shape of one matrix product operation (concept: PitchLinearShape)
typename InstructionShape_,
/// Policy of the details of LDSM shape and iterations
typename Policy_,
/// Number of partitions along K dimension
int PartitionsK_>
class ScaleBiasTileIterator<Shape_, Element_, cutlass::layout::PitchLinear,
InstructionShape_, Policy_, 32, PartitionsK_> {
public:
/// Shape of tile to load (concept: PitchLinearShape)
using Shape = Shape_;
/// Element type
using Element = Element_;
/// Layout of source tile
using Layout = cutlass::layout::PitchLinear;
/// Shape of one matrix product operation (concept: GemmShape)
using InstructionShape = InstructionShape_;
/// Number of participating threads
static int const kThreads = 32;
/// Number of partitions along K dimension
static int const kPartitionsK = PartitionsK_;
/// Number of partitions along K dimension
static int const kElementsPerAccess = 128 / sizeof_bits<Element>::value;
/// TensorRef type for loading element from a tensor
using TensorRef = TensorRef<Element, Layout>;
/// Index type
using Index = typename TensorRef::Index;
/// Long Index type
using LongIndex = typename TensorRef::LongIndex;
/// Coordinate for an element in the tensor
using TensorCoord = typename TensorRef::TensorCoord;
/// Internal structure of iterator - made public to enable introspection
using Policy = Policy_;
private:
/// Pointer type used for accesses
using AccessType = Array<Element, kElementsPerAccess>;
public:
//
// Derived quantities
//
/// Fragment object holding a thread's part of a tile
using Fragment = Array<Element, 2 * Policy::kLdsmOpInner *
InstructionShape::kContiguous / kThreads>;
private:
/// Shared memory base pointers - not advanced
AccessType const *pointer_;
/// Byte offset incremented as iterator advances
Index byte_offset_;
/// Internal counter used to determine when to increment byte offset and when
/// to XOR it
int k_group_idx_;
public:
/// Default ctor constructs null iterator
CUTLASS_HOST_DEVICE
ScaleBiasTileIterator()
: pointer_(nullptr),
byte_offset_(0),
k_group_idx_(0) {}
/// Constructor from TensorRef
CUTLASS_DEVICE
ScaleBiasTileIterator(TensorRef const &ref_scale_bias,
int lane_id)
: byte_offset_(0), k_group_idx_(0) {
/// 16816 only
pointer_ = reinterpret_cast<AccessType const *>(ref_scale_bias.data()) +
((lane_id >> 3) & 1) * Shape::kContiguous / kElementsPerAccess +
(lane_id >> 4);
}
/// Adds a pointer offset to internal pointer(s) to advance through memory
CUTLASS_DEVICE
ScaleBiasTileIterator &add_pointer_offset(LongIndex offset) {
byte_offset_ += offset * sizeof_bits<Element>::value / 8;
return *this;
}
/// Advances an iterator along logical dimensions of matrix in units of whole
/// tiles
CUTLASS_DEVICE
ScaleBiasTileIterator &add_tile_offset(
TensorCoord const &tile_offset) {
int whole_tiles = tile_offset.contiguous() / Policy::kGroupsPerTile;
int k_groups_delta = tile_offset.contiguous() % Policy::kGroupsPerTile;
byte_offset_ += k_groups_delta * sizeof_bits<Element>::value *
kElementsPerAccess * Policy::LdsmShape::kContiguous / 8;
// Multiply by 2 because scale and bias belonging to the same stage are next
// to each other in the shared memory.
pointer_ += (2 * whole_tiles * Shape::kContiguous / kElementsPerAccess);
return *this;
}
/// Advances the iterator along the advance dimension
CUTLASS_DEVICE
ScaleBiasTileIterator &operator++() {
byte_offset_ += Policy::LdsmShape::kContiguous *
sizeof_bits<Element>::value * kElementsPerAccess / 8;
k_group_idx_++;
if (k_group_idx_ == (Policy::kGroupsPerTile / kPartitionsK)) {
k_group_idx_ = 0;
byte_offset_ -= (Policy::kGroupsPerTile / kPartitionsK) *
Policy::LdsmShape::kContiguous *
sizeof_bits<Element>::value * kElementsPerAccess / 8;
add_tile_offset({Policy::kGroupsPerTile, 0});
}
return *this;
}
/// Advances the iterator along the advance dimension
CUTLASS_HOST_DEVICE
ScaleBiasTileIterator &operator--() { assert(0); }
///< advances in units of whole tiles along the logical coordinate space of
///< the tensor
CUTLASS_DEVICE
ScaleBiasTileIterator &operator+=(
TensorCoord const &tile_offset) {
add_tile_offset(tile_offset);
return *this;
}
///< advances in units of whole tiles along the logical coordinate space of
///< the tensor
CUTLASS_DEVICE
ScaleBiasTileIterator &operator-=(
TensorCoord const &tile_offset) {
add_tile_offset(-tile_offset);
return *this;
}
/// Loads a fragment from memory at the location pointed to by the iterator.
CUTLASS_HOST_DEVICE
void load(Fragment &frag) const { load_with_byte_offset(frag, 0); }
/// Loads a fragment from memory with additional logical offset
CUTLASS_DEVICE
void load_with_byte_offset(
/// fragment to load from the tensor
Fragment &frag,
/// loads a tile with a linear offset in units of bytes
Index byte_offset) const {
Array<unsigned, 4> *fetch_ptr =
reinterpret_cast<Array<unsigned, 4> *>(&frag);
CUTLASS_PRAGMA_UNROLL
for (int s = 0; s < 1; ++s) {
CUTLASS_PRAGMA_UNROLL
for (int c = 0; c < Policy::LdsmIterations::kContiguous; ++c) {
int access_idx = c + s * Policy::LdsmIterations::kContiguous;
AccessType const *source_ptr =
pointer_ + Policy::LdsmShape::kContiguous * c;
char const *source_byte_ptr =
reinterpret_cast<char const *>(source_ptr) + byte_offset +
byte_offset_;
cutlass::arch::ldsm<layout::RowMajor, 4>(
fetch_ptr[access_idx], source_byte_ptr);
}
}
}
/// Loads a fragment from memory with additional logical offset
CUTLASS_DEVICE
void load_with_pointer_offset(
/// fragment to load from the tensor
Fragment &frag,
/// loads a tile with a linear offset
Index pointer_offset) const {
load_with_byte_offset(frag, pointer_offset * sizeof(Element));
}
/// Loads a fragment from memory with logical offset in units of whole tiles.
CUTLASS_DEVICE
void load(
/// fragment to load from the tensor
Fragment &frag,
/// loads a tile with a logical offset in units of whole tiles
TensorCoord const &tile_offset) const {
load_with_byte_offset(frag, tile_offset, 0);
}
/// Loads a fragment from memory with logical offset in units of whole tiles.
CUTLASS_DEVICE
void load(
/// fragment to load from the tensor
Fragment &frag,
/// loads a tile with a logical offset in units of whole tiles
TensorCoord const &tile_offset,
/// loads a tile with a logical offset AND a pointer offset
Index pointer_offset) const {
load_with_byte_offset(frag, tile_offset, pointer_offset * sizeof(Element));
}
/// Loads a fragment from memory with logical offset in units of whole tiles.
CUTLASS_DEVICE
void load_with_byte_offset(
/// fragment to load from the tensor
Fragment &frag,
/// loads a tile with a logical offset in units of whole tiles
TensorCoord const &tile_offset,
/// loads a tile with a logical offset AND a pointer offset
Index byte_offset) const {
Index pointer_offset = tile_offset.contiguous() *
InstructionShape::kContiguous /
kElementsPerAccess;
byte_offset += sizeof_bits<AccessType>::value * pointer_offset / 8;
load_with_byte_offset(frag, byte_offset);
}
/// Notify the iterator which k-group it is currently pointing to.
///
/// This does not advance the iterator. Rather, it overrides its internal
/// tracking with constant-valued k-group index to enable the compiler to
/// fold constants and achieve more efficient code.
///
/// This is used by some nontrivial permuted layouts.
CUTLASS_DEVICE
void set_kgroup_index(int k_group) {
k_group_idx_ = k_group % (Policy::kGroupsPerTile / kPartitionsK);
}
};
////////////////////////////////////////////////////////////////////////////////
/// This tile iterator is specialized for 32-thread TensorOps. It uses LDSM to
/// load from shared memory and therefore must be initialized with a TensorRef
/// to shared memory.
///
/// Satisfies:
/// ReadableRandomAccessContiguousTileIteratorConcept
///
template <
/// Size of the matrix to load (concept: MatrixShape)
typename Shape_,
/// Data type of elements
typename Element_,
/// Shape of one matrix product operation (concept: MatrixShape)
typename InstructionShape_,
/// Policy of the details of LDSM shape and iterations
typename Policy_,
/// Number of partitions along K dimension
int PartitionsK_>
class ScaleBiasTileIterator<Shape_, Element_, cutlass::layout::RowMajor,
InstructionShape_, Policy_, 32, PartitionsK_> {
public:
/// Shape of tile to load (concept: PitchLinearShape)
using Shape = Shape_;
/// 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_;
/// Number of participating threads
static int const kThreads = 32;
/// TensorRef type for loading element from a tensor
using TensorRef = TensorRef<Element, Layout>;
/// Index type
using Index = typename TensorRef::Index;
/// Long Index type
using LongIndex = typename TensorRef::LongIndex;
/// Coordinate for an element in the tensor
using TensorCoord = typename TensorRef::TensorCoord;
/// Internal structure of iterator - made public to enable introspection
using Policy = Policy_;
/// Underlying tile iterator implementation
using Base = ScaleBiasTileIterator<
layout::PitchLinearShape<Shape::kColumn, Shape::kRow>, Element,
layout::PitchLinear,
layout::PitchLinearShape<InstructionShape::kColumn,
InstructionShape::kRow>,
Policy, kThreads, PartitionsK_>;
public:
//
// Derived quantities
//
/// Fragment object holding a thread's part of a tile
using Fragment = typename Base::Fragment;
private:
/// Underlying tile iterator
Base iterator_;
public:
/// Default ctor constructs null iterator
CUTLASS_HOST_DEVICE
ScaleBiasTileIterator() {}
/// Constructor from TensorRef
CUTLASS_HOST_DEVICE
ScaleBiasTileIterator(TensorRef const &ref_scale_bias, int lane_id)
: iterator_({ref_scale_bias.data(), ref_scale_bias.stride()}, lane_id) {}
/// Adds a pointer offset to internal pointer(s) to advance through memory
CUTLASS_HOST_DEVICE
ScaleBiasTileIterator &add_pointer_offset(LongIndex offset) {
iterator_.add_pointer_offset(offset);
return *this;
}
/// Advances an iterator along logical dimensions of matrix in units of whole
/// tiles
CUTLASS_HOST_DEVICE
ScaleBiasTileIterator &add_tile_offset(
TensorCoord const &tile_offset) {
iterator_.add_tile_offset({tile_offset.column(), tile_offset.row()});
return *this;
}
/// Advances an iterator along logical dimensions of matrix in units of whole
/// tiles
CUTLASS_DEVICE
ScaleBiasTileIterator &add_tile_offset_negative(
TensorCoord const &tile_offset) {
iterator_.add_tile_offset_negative({tile_offset.column(), tile_offset.row()});
return *this;
}
/// Advances the iterator along the advance dimension
CUTLASS_HOST_DEVICE
ScaleBiasTileIterator &operator++() {
++iterator_;
return *this;
}
/// Advances the iterator along the advance dimension
CUTLASS_HOST_DEVICE
ScaleBiasTileIterator &operator--() {
--iterator_;
return *this;
}
///< advances in units of whole tiles along the logical coordinate space of
///< the tensor
CUTLASS_DEVICE
ScaleBiasTileIterator &operator+=(
TensorCoord const &tile_offset) {
add_tile_offset(PitchLinearCoord(tile_offset.column(), tile_offset.row()));
return *this;
}
///< advances in units of whole tiles along the logical coordinate space of
///< the tensor
CUTLASS_DEVICE
ScaleBiasTileIterator &operator-=(
TensorCoord const &tile_offset) {
add_tile_offset(-PitchLinearCoord(tile_offset.column(), tile_offset.row()));
return *this;
}
/// Loads a fragment from memory at the location pointed to by the iterator.
CUTLASS_HOST_DEVICE
void load(Fragment &frag) const { iterator_.load(frag); }
/// Loads a fragment from memory with additional logical offset
CUTLASS_DEVICE
void load_with_pointer_offset(
/// fragment to load from the tensor
Fragment &frag,
/// loads a tile with a linear offset
Index pointer_offset) const {
iterator_.load_with_pointer_offset(frag, pointer_offset);
}
/// Loads a fragment from memory with additional logical offset
CUTLASS_DEVICE
void load_with_byte_offset(
/// fragment to load from the tensor
Fragment &frag,
/// loads a tile with a linear offset
Index byte_offset) const {
iterator_.load_with_byte_offset(frag, byte_offset);
}
/// Loads a fragment from memory with logical offset in units of whole tiles.
CUTLASS_DEVICE
void load(
/// fragment to load from the tensor
Fragment &frag,
/// loads a tile with a logical offset in units of whole tiles
TensorCoord const &tile_offset) const {
// TODO
assert(0);
}
/// Loads a fragment from memory with logical offset in units of whole tiles.
CUTLASS_DEVICE
void load(
/// fragment to load from the tensor
Fragment &frag,
/// loads a tile with a logical offset in units of whole tiles
TensorCoord const &tile_offset,
/// loads a tile with a logical offset AND a pointer offset
Index pointer_offset) const {
// TODO
assert(0);
}
/// Loads a fragment from memory with logical offset in units of whole tiles.
CUTLASS_DEVICE
void load_with_byte_offset(
/// fragment to load from the tensor
Fragment &frag,
/// loads a tile with a logical offset in units of whole tiles
TensorCoord const &tile_offset,
/// loads a tile with a logical offset AND a pointer offset
Index byte_offset) const {
iterator_.load_with_byte_offset(
frag, {tile_offset.strided(), tile_offset.contiguous()}, byte_offset);
}
/// Notify the iterator which k-group it is currently pointing to.
///
/// This does not advance the iterator. Rather, it overrides its internal
/// tracking with constant-valued k-group index to enable the compiler to
/// fold constants and achieve more efficient code.
///
/// This is used by some nontrivial permuted layouts.
CUTLASS_DEVICE
void set_kgroup_index(int k_group) {
iterator_.set_kgroup_index(k_group);
}
};
////////////////////////////////////////////////////////////////////////////////
} // namespace warp
} // namespace gemm
} // namespace cutlass
////////////////////////////////////////////////////////////////////////////////
@@ -0,0 +1,117 @@
/***************************************************************************************************
* 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 Templates implementing warp-level per-channel softmax before
matrix multiply-accumulate operations targeting Tensor Cores.
*/
#pragma once
#include "cutlass/cutlass.h"
#include "cutlass/array.h"
#include "cutlass/platform/platform.h"
#include "cutlass/numeric_conversion.h"
#include "cutlass/numeric_types.h"
#include "cutlass/matrix_shape.h"
#include "cutlass/arch/memory_sm75.h"
#include "cutlass/arch/mma_sm75.h"
#include "cutlass/arch/mma_sm80.h"
#include "cutlass/gemm/gemm.h"
#include "cutlass/gemm/warp/mma.h"
#include "cutlass/gemm/warp/mma_tensor_op_policy.h"
#include "cutlass/gemm/warp/mma_tensor_op_tile_iterator.h"
#include "cutlass/gemm/warp/mma_tensor_op_tile_iterator_sm80.h"
/////////////////////////////////////////////////////////////////////////////////////////////////
namespace cutlass {
namespace gemm {
namespace warp {
/////////////////////////////////////////////////////////////////////////////////////////////////
template <typename FragmentActivations, typename FragmentNormSum>
struct SoftmaxScaleBiasTransform {
using T = typename FragmentActivations::Element;
static int const NumActivations = FragmentActivations::kElements;
static int const NumNormSum = FragmentNormSum::kElements;
static int const MmaElements = 2;
// One element has one scale and one bias
static int const MmaScaleBiasPair = 2;
// 16816 has 2 columns and 2 rows
static int const MmaCols = 2;
static int const MmaRows = 2;
using MmaOperand = Array<T, MmaElements>;
using NormSumOperand = Array<__half2, MmaScaleBiasPair>;
CUTLASS_DEVICE
void transform(MmaOperand &activations,
NormSumOperand const &norm_sum) {
__half2* packed_activations = reinterpret_cast<__half2*>(&activations);
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < MmaElements / 2; ++i) {
__half2 out = ::h2exp(__hsub2(packed_activations[i], norm_sum[2*i]));
packed_activations[i] = __hmul2(out, norm_sum[2*i + 1]);
}
}
CUTLASS_DEVICE
void operator()(FragmentActivations &activations,
FragmentNormSum const &norm_sum) {
MmaOperand *ptr_activations = reinterpret_cast<MmaOperand *>(&activations);
NormSumOperand const *ptr_norm_sum =
reinterpret_cast<NormSumOperand const *>(&norm_sum);
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < (NumActivations / MmaElements); ++i) {
transform(ptr_activations[i],
ptr_norm_sum[i / (MmaCols * MmaRows) * MmaRows + i % MmaRows]);
}
}
};
/////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace warp
} // namespace gemm
} // namespace cutlass
/////////////////////////////////////////////////////////////////////////////////////////////////