co-authored by
Aniket Shivam
parent
ca23ff7924
commit
b72cbf957d
@@ -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_;
|
||||
|
||||
@@ -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
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
Reference in New Issue
Block a user