@@ -54,12 +54,11 @@ using ElementInputA = cutlass::half_t; // <- data type of elements
|
||||
using ElementInputB = cutlass::half_t; // <- data type of elements in input matrix B
|
||||
using ElementOutput = float; // <- data type of elements in output matrix D
|
||||
|
||||
// The code section below describes matrix layout of input and output matrices.
|
||||
// Column Major for Matrix A, B and C.
|
||||
|
||||
// Note that if the output is column major, the bias has to be per row. i.e. every row has different bias.
|
||||
// If the output is row major, the bias has to be per column, i.e. every column has different bias.
|
||||
// Below list some other notices:
|
||||
//
|
||||
// Note this example only works for ColumnMajor output because
|
||||
// 1) we only have row major epilogue.
|
||||
// 2) we swap A and B if the output is column major then we can still use the
|
||||
// row major epilogue.
|
||||
|
||||
@@ -457,9 +457,13 @@ Result profile_convolution(Options const &options) {
|
||||
ElementInputB(-8),
|
||||
0);
|
||||
|
||||
// Fill tensor C on host with zeros
|
||||
cutlass::reference::host::TensorFill(
|
||||
tensor_c.host_view());
|
||||
// Fill tensor C on host with uniform-distribution random data
|
||||
cutlass::reference::host::TensorFillRandomUniform(
|
||||
tensor_c.host_view(),
|
||||
1,
|
||||
ElementOutput(7),
|
||||
ElementOutput(-8),
|
||||
0);
|
||||
|
||||
// Fill tensor D on host with zeros
|
||||
cutlass::reference::host::TensorFill(
|
||||
@@ -686,7 +690,7 @@ int main(int argc, char const **args) {
|
||||
cudaDeviceProp props;
|
||||
CUDA_CHECK(cudaGetDeviceProperties(&props, 0));
|
||||
|
||||
if (!(props.major > 8 || (props.major == 8 && props.minor >= 0))) {
|
||||
if (!(props.major >= 8)) {
|
||||
std::cerr << "Ampere Tensor Ops must be run on a machine with compute capability at least 80."
|
||||
<< std::endl;
|
||||
notSupported = true;
|
||||
|
||||
@@ -290,7 +290,7 @@ int main(int argc, char const **args) {
|
||||
cudaDeviceProp props;
|
||||
CUDA_CHECK(cudaGetDeviceProperties(&props, 0));
|
||||
|
||||
if (!(props.major > 8 || (props.major == 8 && props.minor >= 0))) {
|
||||
if (!(props.major >= 8)) {
|
||||
std::cerr << "Ampere Tensor Ops must be run on a machine with compute capability at least 80."
|
||||
<< std::endl;
|
||||
notSupported = true;
|
||||
|
||||
@@ -326,7 +326,7 @@ int main(int argc, char const **args) {
|
||||
cudaDeviceProp props;
|
||||
CUDA_CHECK(cudaGetDeviceProperties(&props, 0));
|
||||
|
||||
if (!(props.major > 8 || (props.major == 8 && props.minor >= 0))) {
|
||||
if (!(props.major >= 8)) {
|
||||
std::cerr << "Ampere Tensor Ops must be run on a machine with compute capability at least 80."
|
||||
<< std::endl;
|
||||
notSupported = true;
|
||||
|
||||
@@ -32,7 +32,7 @@
|
||||
/**
|
||||
The example demenstrates how to reduce one of the operands of the GEMM along the k-dimension when
|
||||
computing GEMM. So the output also contains either a Mx1 or 1XN vector. It only works with Ampere
|
||||
HMMA 16x8x16 FP16 tensor cores, though it is not difficult to apply to other Turing/Ampere tensor
|
||||
16x8x16 FP16/BF16 tensor cores, though it is not difficult to apply to other Turing/Ampere tensor
|
||||
core instructions.
|
||||
|
||||
Most of the reduction is done in gemm/warp level, see gemm/warp/mma_with_reduction_tensor_op.h
|
||||
@@ -67,9 +67,9 @@ epilogue/threadblock/epilogue_gemm_k_reduction.h
|
||||
// elements
|
||||
using ElementAccumulator = float; // Data type of accumulator
|
||||
using ElementComputeEpilogue = ElementAccumulator; // Data type of epilogue computation
|
||||
using ElementInputA = cutlass::half_t; // Data type of elements in input tensor
|
||||
using ElementInputB = cutlass::half_t; // Data type of elements in input tensor
|
||||
using ElementOutput = cutlass::half_t; // Data type of elements in output tensor
|
||||
using ElementInputA = cutlass::bfloat16_t; // Data type of elements in input tensor
|
||||
using ElementInputB = cutlass::bfloat16_t; // Data type of elements in input tensor
|
||||
using ElementOutput = cutlass::bfloat16_t; // Data type of elements in output tensor
|
||||
|
||||
using LayoutInputA = cutlass::layout::ColumnMajor;
|
||||
using LayoutInputB = cutlass::layout::RowMajor;
|
||||
@@ -369,22 +369,22 @@ Result profile(Options const &options) {
|
||||
cutlass::reference::host::TensorFillRandomUniform(
|
||||
tensor_a.host_view(),
|
||||
1,
|
||||
ElementInputA(4),
|
||||
ElementInputA(-4),
|
||||
ElementInputA(2),
|
||||
ElementInputA(-2),
|
||||
0); // <- Fill tensor A on host with uniform-distribution random data
|
||||
|
||||
cutlass::reference::host::TensorFillRandomUniform(
|
||||
tensor_b.host_view(),
|
||||
1,
|
||||
ElementInputB(4),
|
||||
ElementInputB(-4),
|
||||
ElementInputB(2),
|
||||
ElementInputB(-2),
|
||||
0); // <- Fill tensor B on host with uniform-distribution random data
|
||||
|
||||
cutlass::reference::host::TensorFillRandomUniform(
|
||||
tensor_c.host_view(),
|
||||
1,
|
||||
ElementOutput(4),
|
||||
ElementOutput(-4),
|
||||
ElementOutput(2),
|
||||
ElementOutput(-2),
|
||||
0); // <- Fill matrix C on host with uniform-distribution random data
|
||||
cutlass::reference::host::TensorFill(
|
||||
tensor_d.host_view()); // <- fill matrix D on host with zeros
|
||||
@@ -612,10 +612,10 @@ Result profile(Options const &options) {
|
||||
|
||||
if (options.reference_check) {
|
||||
output_workspace << "Reference D = \n" << tensor_ref_d.host_view() << "\n\n";
|
||||
output_workspace << "Reference reduction vector= \n" << tensor_ref_reduction.host_view() << "\n\n";
|
||||
output_workspace << "Reference reduction vector = \n" << tensor_ref_reduction.host_view() << "\n\n";
|
||||
}
|
||||
|
||||
output_workspace << "Computed = \n" << tensor_d.host_view() << std::endl;
|
||||
output_workspace << "Computed D = \n" << tensor_d.host_view() << std::endl;
|
||||
output_workspace << "Computed reduction vector = \n" << tensor_reduction.host_view() << std::endl;
|
||||
|
||||
std::cout << "Results written to '" << ss.str() << "'." << std::endl;
|
||||
@@ -699,7 +699,7 @@ int main(int argc, char const **args) {
|
||||
cudaDeviceProp props;
|
||||
CUDA_CHECK(cudaGetDeviceProperties(&props, 0));
|
||||
|
||||
if (!(props.major > 8 || (props.major == 8 && props.minor >= 0))) {
|
||||
if (!(props.major >= 8)) {
|
||||
std::cerr << "Ampere Tensor Ops must be run on a machine with compute capability at least 80."
|
||||
<< std::endl;
|
||||
notSupported = true;
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -34,3 +34,8 @@ cutlass_example_add_executable(
|
||||
ampere_fprop_mainloop_fusion.cu
|
||||
)
|
||||
|
||||
cutlass_example_add_executable(
|
||||
25_ampere_3d_fprop_mainloop_fusion
|
||||
ampere_3d_fprop_mainloop_fusion.cu
|
||||
)
|
||||
|
||||
|
||||
@@ -0,0 +1,776 @@
|
||||
/***************************************************************************************************
|
||||
* 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
/**
|
||||
|
||||
This example shows how to fuse per channel scale+bias+relu of the activations
|
||||
into the 3D fprop mainloop.
|
||||
|
||||
Compared with original 3D fprop kernel, this example has two more vectors, one for
|
||||
the scale and one for the bias. The length of the vectors is the same as the
|
||||
activation channel number. This kernel loads the vectors when the associated
|
||||
activation channels are loaded in the mainloop. Between reading the
|
||||
activations and scale/bias data from the shared memory and calling tensor core
|
||||
instructions, scale+bias+relu is computed in the register file.
|
||||
|
||||
This example is customized for Ampere 16816 fp16 tensor core instruction.
|
||||
Changing to different data types or different tensor core instruction require
|
||||
source code changing. See
|
||||
include/cutlass/conv/threadblock/implicit_gemm_fprop_fusion_multistage.h for more
|
||||
technical details.
|
||||
|
||||
This example is modified based on 25_ampere_fprop_mainloop_fusion. The command
|
||||
line is the same.
|
||||
*/
|
||||
|
||||
#include <iostream>
|
||||
#include <fstream>
|
||||
#include <sstream>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/gemm/device/gemm.h"
|
||||
#include "cutlass/conv/kernel/default_conv3d_fprop_fusion.h"
|
||||
#include "cutlass/conv/device/implicit_gemm_convolution_fusion.h"
|
||||
|
||||
#include "cutlass/util/command_line.h"
|
||||
#include "cutlass/util/host_tensor.h"
|
||||
#include "cutlass/util/tensor_view_io.h"
|
||||
#include "cutlass/util/reference/device/gemm.h"
|
||||
#include "cutlass/util/reference/host/tensor_compare.h"
|
||||
#include "cutlass/util/reference/host/tensor_copy.h"
|
||||
#include "cutlass/util/reference/host/tensor_fill.h"
|
||||
#include "cutlass/util/reference/device/convolution.h"
|
||||
#include "cutlass/util/tensor_view_io.h"
|
||||
|
||||
#include "helper.h"
|
||||
|
||||
// The code section below describes datatype for input, output tensors and computation between
|
||||
// elements
|
||||
using ElementAccumulator = float; // Data type of accumulator
|
||||
using ElementComputeEpilogue = float; // Data type of epilogue computation (alpha, beta)
|
||||
using ElementInputA = cutlass::half_t; // Data type of elements in input tensor
|
||||
using ElementInputB = cutlass::half_t; // Data type of elements in input tensor
|
||||
using ElementInputScaleBias = cutlass::half_t; // Data type of elements in input sclae and bias vectors
|
||||
using ElementOutput = float; // Data type of elements in output tensor
|
||||
|
||||
using LayoutInputA = cutlass::layout::TensorNDHWC;
|
||||
using LayoutInputB = cutlass::layout::TensorNDHWC;
|
||||
using LayoutInputScaleBias = cutlass::layout::RowMajor;
|
||||
using LayoutOutput = cutlass::layout::TensorNDHWC;
|
||||
|
||||
// This code section describes whether you want to use tensor cores or regular SIMT cores on GPU SM
|
||||
using MMAOp = cutlass::arch::OpClassTensorOp;
|
||||
|
||||
// This code section describes CUDA SM architecture number
|
||||
using SmArch = cutlass::arch::Sm80;
|
||||
|
||||
// This code section describes the tile size a thread block will compute
|
||||
using ThreadblockShape = cutlass::gemm::GemmShape<128, 128, 32>; // Threadblock tile shape
|
||||
|
||||
// This code section describes tile size a warp will compute
|
||||
using WarpShape = cutlass::gemm::GemmShape<64, 64, 32>; // Warp tile shape
|
||||
|
||||
// This code section describes the size of MMA op
|
||||
using InstructionShape = cutlass::gemm::GemmShape<16, 8, 16>; // TensorCore instruction shape
|
||||
|
||||
// This code section describes how threadblocks are scheduled on GPU
|
||||
using SwizzleThreadBlock = cutlass::gemm::threadblock::GemmIdentityThreadblockSwizzle<>;
|
||||
|
||||
// Number of pipelines you want to use
|
||||
constexpr int NumStages = 4;
|
||||
|
||||
// This code section describe iterator algorithm selected is Analytic or Optimized
|
||||
static cutlass::conv::IteratorAlgorithm const IteratorAlgorithm = cutlass::conv::IteratorAlgorithm::kOptimized;
|
||||
|
||||
// This code section describes the epilogue part of the kernel, we use default value
|
||||
using EpilogueOp = cutlass::epilogue::thread::LinearCombination<
|
||||
ElementOutput, // Data type of output matrix.
|
||||
128 / cutlass::sizeof_bits<ElementOutput>::value, // The number of elements per vectorized.
|
||||
// memory access. This becomes the vector width of
|
||||
// math instructions in the epilogue too.
|
||||
ElementAccumulator, // Data type of accumulator
|
||||
ElementComputeEpilogue>; // Data type for alpha/beta in linear combination
|
||||
|
||||
using Conv3dFpropFusionKernel = typename cutlass::conv::kernel::DefaultConv3dFpropFusion<
|
||||
ElementInputA, LayoutInputA,
|
||||
ElementInputB, LayoutInputB,
|
||||
ElementInputScaleBias, LayoutInputScaleBias,
|
||||
ElementOutput, LayoutOutput,
|
||||
ElementAccumulator,
|
||||
MMAOp,
|
||||
SmArch,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOp,
|
||||
SwizzleThreadBlock,
|
||||
NumStages,
|
||||
cutlass::arch::OpMultiplyAdd,
|
||||
IteratorAlgorithm
|
||||
>::Kernel;
|
||||
|
||||
using ImplicitGemmFusion = cutlass::conv::device::ImplicitGemmConvolutionFusion<Conv3dFpropFusionKernel>;
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// Command line options parsing
|
||||
struct Options {
|
||||
|
||||
bool help;
|
||||
cutlass::Tensor5DCoord input_size;
|
||||
cutlass::Tensor5DCoord filter_size;
|
||||
cutlass::Coord<3> padding;
|
||||
cutlass::Coord<3> conv_stride;
|
||||
cutlass::Coord<3> dilation;
|
||||
bool reference_check;
|
||||
bool measure_performance;
|
||||
int iterations;
|
||||
bool save_workspace;
|
||||
ElementComputeEpilogue alpha;
|
||||
ElementComputeEpilogue beta;
|
||||
bool benchmark;
|
||||
std::string tag;
|
||||
|
||||
Options():
|
||||
help(false),
|
||||
input_size(1, 32, 32, 32, 32),
|
||||
filter_size(32, 3, 3, 3, 32),
|
||||
padding(cutlass::make_Coord(1, 1, 1)),
|
||||
conv_stride(cutlass::make_Coord(1, 1, 1)),
|
||||
dilation(cutlass::make_Coord(1, 1, 1)),
|
||||
reference_check(true),
|
||||
measure_performance(false),
|
||||
iterations(20),
|
||||
save_workspace(false),
|
||||
alpha(1),
|
||||
beta(0),
|
||||
benchmark(false) { }
|
||||
|
||||
// Verify the problem size is compatible with the CUTLASS Convolution implementation.
|
||||
bool valid() {
|
||||
|
||||
//
|
||||
// CUTLASS attempts to load 128b vectors of cutlass::half_t (F16) elements. Consequently,
|
||||
// all pointers, strides, and tensor extents must be divisible by 8 elements.
|
||||
//
|
||||
int const kAlignment = 8;
|
||||
|
||||
if ((input_size.c() % kAlignment) ||
|
||||
(filter_size.n() % kAlignment)) {
|
||||
|
||||
// misaligned tensors
|
||||
return false;
|
||||
}
|
||||
|
||||
// Invalid padding
|
||||
if ((padding[0] != filter_size.d() / 2) ||
|
||||
(padding[1] != filter_size.h() / 2) ||
|
||||
(padding[2] != filter_size.w() / 2)) {
|
||||
|
||||
return false;
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
/// Updates input and filter sizes
|
||||
void update(
|
||||
cutlass::Tensor5DCoord input_size,
|
||||
cutlass::Tensor5DCoord filter_size,
|
||||
cutlass::Coord<3> stride) {
|
||||
|
||||
this->input_size = input_size;
|
||||
this->filter_size = filter_size;
|
||||
conv_stride = stride;
|
||||
|
||||
padding[0] = filter_size.d() / 2;
|
||||
padding[1] = filter_size.h() / 2;
|
||||
padding[2] = filter_size.w() / 2;
|
||||
}
|
||||
|
||||
// Parses the command line
|
||||
void parse(int argc, char const **args) {
|
||||
cutlass::CommandLine cmd(argc, args);
|
||||
|
||||
if (cmd.check_cmd_line_flag("help")) {
|
||||
help = true;
|
||||
}
|
||||
|
||||
if (cmd.check_cmd_line_flag("ref-check")) {
|
||||
reference_check = true;
|
||||
}
|
||||
|
||||
if (cmd.check_cmd_line_flag("perf-check")) {
|
||||
measure_performance = true;
|
||||
}
|
||||
|
||||
if (cmd.check_cmd_line_flag("save-workspace")) {
|
||||
save_workspace = true;
|
||||
}
|
||||
|
||||
if (cmd.check_cmd_line_flag("benchmark")) {
|
||||
benchmark = true;
|
||||
}
|
||||
|
||||
cmd.get_cmd_line_argument("n", input_size.n());
|
||||
cmd.get_cmd_line_argument("d", input_size.d());
|
||||
cmd.get_cmd_line_argument("h", input_size.h());
|
||||
cmd.get_cmd_line_argument("w", input_size.w());
|
||||
cmd.get_cmd_line_argument("c", input_size.c());
|
||||
|
||||
cmd.get_cmd_line_argument("k", filter_size.n());
|
||||
cmd.get_cmd_line_argument("t", filter_size.d());
|
||||
cmd.get_cmd_line_argument("r", filter_size.h());
|
||||
cmd.get_cmd_line_argument("s", filter_size.w());
|
||||
filter_size.c() = input_size.c();
|
||||
|
||||
cmd.get_cmd_line_argument("alpha", alpha);
|
||||
cmd.get_cmd_line_argument("beta", beta);
|
||||
|
||||
cmd.get_cmd_line_argument("iterations", iterations);
|
||||
cmd.get_cmd_line_argument("tag", tag);
|
||||
|
||||
if (filter_size.d() == 3 && filter_size.h() == 3 && filter_size.w() == 3) {
|
||||
padding = cutlass::make_Coord(1, 1, 1);
|
||||
}
|
||||
else {
|
||||
filter_size.d() = 1;
|
||||
filter_size.h() = 1;
|
||||
filter_size.w() = 1;
|
||||
padding = cutlass::make_Coord(0, 0, 0);
|
||||
}
|
||||
}
|
||||
|
||||
/// Prints the usage statement.
|
||||
std::ostream & print_usage(std::ostream &out) const {
|
||||
|
||||
out << "25_ampere_3d_fprop_mainloop_fusion example\n\n"
|
||||
<< " This example fuses scale+bias+relu of the activations into Ampere's\n"
|
||||
<< " Tensor Core operators on F16 data types to compute\n"
|
||||
<< " forward convolution on tensors of layout NDHWC.\n\n"
|
||||
<< "Options:\n\n"
|
||||
<< " --help If specified, displays this usage statement.\n\n"
|
||||
<< " --n <int> Input tensor extent N\n"
|
||||
<< " --d <int> Input tensor extent D\n"
|
||||
<< " --h <int> Input tensor extent H\n"
|
||||
<< " --w <int> Input tensor extent W\n"
|
||||
<< " --c <int> Input tensor extent C\n"
|
||||
<< " --k <int> Filter extent K\n"
|
||||
<< " --t <int> Filter extent T\n"
|
||||
<< " --r <int> Filter extent R\n"
|
||||
<< " --s <int> Filter extent S\n\n"
|
||||
<< " --alpha <float> Epilogue scalar alpha\n"
|
||||
<< " --beta <float> Epilogue scalar beta\n\n"
|
||||
<< " --ref-check If set (true), reference check on the host is computed\n"
|
||||
<< " --perf-check If set (true), performance is measured.\n"
|
||||
<< " --benchmark If set (true), performance benchmarking on several layers and batch-size.\n"
|
||||
<< " --iterations <int> Number of profiling iterations to perform.\n"
|
||||
<< " --save-workspace If set, workspace is written to a text file.\n"
|
||||
<< " --tag <string> String to replicate across the first column in the results table\n";
|
||||
|
||||
out << "\n\nExamples:\n\n"
|
||||
<< "$ ./25_ampere_3d_fprop_mainloop_fusion --n=32 --d=96 --h=96 --w=96 --c=64 --k=64 --t=1 --r=1 --s=1\n\n"
|
||||
<< "$ ./25_ampere_3d_fprop_mainloop_fusion --n=1 --d=224 --h=224 --w=224 --c=32 --k=32 --t=3 --r=3 --s=3 --ref-check\n\n"
|
||||
<< "$ ./25_ampere_3d_fprop_mainloop_fusion --n=19 --d=94 --h=96 --w=96 --c=128 --k=128 --t=1 --r=1 --s=1\n\n";
|
||||
|
||||
return out;
|
||||
}
|
||||
|
||||
/// Computes the output tensor size (NPQK)
|
||||
cutlass::Tensor5DCoord output_size() const {
|
||||
return cutlass::Tensor5DCoord(
|
||||
input_size.n(),
|
||||
(input_size.d() + padding[0] + padding[0] - filter_size.d()) / conv_stride[0] + 1,
|
||||
(input_size.h() + padding[1] + padding[1] - filter_size.h()) / conv_stride[1] + 1,
|
||||
(input_size.w() + padding[2] + padding[2] - filter_size.w()) / conv_stride[2] + 1,
|
||||
filter_size.n());
|
||||
}
|
||||
|
||||
/// Compute performance in GFLOP/s
|
||||
double gflops(double runtime_s) const {
|
||||
|
||||
// Number of multiply-adds = NPQK * CRS
|
||||
int64_t fmas = output_size().product() * int64_t(filter_size.d() * filter_size.h() * filter_size.w() * filter_size.c());
|
||||
|
||||
// Two flops per multiply-add
|
||||
return 2.0 * double(fmas) / double(1.0e9) / runtime_s;
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
struct Result {
|
||||
double runtime_ms;
|
||||
double gflops;
|
||||
cutlass::Status status;
|
||||
cutlass::Status reference_check;
|
||||
cudaError_t error;
|
||||
|
||||
Result():
|
||||
runtime_ms(0),
|
||||
gflops(0),
|
||||
status(cutlass::Status::kSuccess),
|
||||
reference_check(cutlass::Status::kInvalid),
|
||||
error(cudaSuccess) { }
|
||||
|
||||
static std::ostream & print_header(std::ostream &out, Options const &options) {
|
||||
|
||||
if (!options.tag.empty()) {
|
||||
out << "Name,";
|
||||
}
|
||||
|
||||
out << "Layer,N,D,H,W,C,K,T,R,S,Stride_D,Stride_H,Stride_W,Runtime,GFLOPs";
|
||||
|
||||
return out;
|
||||
}
|
||||
|
||||
std::ostream & print(std::ostream &out, int idx, Options const &options) {
|
||||
|
||||
if (!options.tag.empty()) {
|
||||
out << options.tag << ",";
|
||||
}
|
||||
|
||||
out
|
||||
<< "conv_" << idx << ","
|
||||
<< options.input_size.n() << ","
|
||||
<< options.input_size.d() << ","
|
||||
<< options.input_size.h() << ","
|
||||
<< options.input_size.w() << ","
|
||||
<< options.input_size.c() << ","
|
||||
<< options.filter_size.n() << ","
|
||||
<< options.filter_size.d() << ","
|
||||
<< options.filter_size.h() << ","
|
||||
<< options.filter_size.w() << ","
|
||||
<< options.conv_stride[0] << ","
|
||||
<< options.conv_stride[1] << ","
|
||||
<< options.conv_stride[2] << ","
|
||||
<< runtime_ms << ","
|
||||
<< gflops;
|
||||
|
||||
return out;
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Runs one benchmark
|
||||
Result profile_convolution(Options const &options) {
|
||||
|
||||
Result result;
|
||||
|
||||
//
|
||||
// Allocate host-device tensors using the CUTLASS Utilities.
|
||||
//
|
||||
|
||||
cutlass::HostTensor<ElementInputA, LayoutInputA> tensor_a(options.input_size);
|
||||
cutlass::HostTensor<ElementInputA, LayoutInputA> tensor_transformed_a(options.input_size);
|
||||
cutlass::HostTensor<ElementInputB, LayoutInputB> tensor_b(options.filter_size);
|
||||
cutlass::HostTensor<ElementInputScaleBias, LayoutInputScaleBias>
|
||||
tensor_a_scale({1, options.input_size.c()});
|
||||
cutlass::HostTensor<ElementInputScaleBias, LayoutInputScaleBias>
|
||||
tensor_a_bias({1, options.input_size.c()});
|
||||
cutlass::HostTensor<ElementOutput, LayoutOutput> tensor_c(options.output_size());
|
||||
cutlass::HostTensor<ElementOutput, LayoutOutput> tensor_d(options.output_size());
|
||||
cutlass::HostTensor<ElementOutput, LayoutOutput> tensor_ref_d(options.output_size());
|
||||
|
||||
//
|
||||
// Initialize tensors
|
||||
//
|
||||
|
||||
// Fill tensor A on host with uniform-distribution random data
|
||||
cutlass::reference::host::TensorFillRandomUniform(
|
||||
tensor_a.host_view(),
|
||||
1,
|
||||
ElementInputA(3),
|
||||
ElementInputA(-4),
|
||||
0);
|
||||
|
||||
// Fill scale vector for tensor A on host with uniform-distribution random
|
||||
// data
|
||||
cutlass::reference::host::TensorFillRandomUniform(
|
||||
tensor_a_scale.host_view(),
|
||||
1,
|
||||
ElementInputA(3),
|
||||
ElementInputA(-4),
|
||||
0);
|
||||
|
||||
// Fill bias vector for tensor A on host with uniform-distribution random
|
||||
// data
|
||||
cutlass::reference::host::TensorFillRandomUniform(
|
||||
tensor_a_bias.host_view(),
|
||||
1,
|
||||
ElementInputA(3),
|
||||
ElementInputA(-4),
|
||||
0);
|
||||
|
||||
// Fill tensor B on host with uniform-distribution random data
|
||||
cutlass::reference::host::TensorFillRandomUniform(
|
||||
tensor_b.host_view(),
|
||||
1,
|
||||
ElementInputB(7),
|
||||
ElementInputB(-8),
|
||||
0);
|
||||
|
||||
// Fill tensor C on host with uniform-distribution random data
|
||||
cutlass::reference::host::TensorFillRandomUniform(
|
||||
tensor_c.host_view(),
|
||||
1,
|
||||
ElementOutput(7),
|
||||
ElementOutput(-8),
|
||||
0);
|
||||
|
||||
// Fill tensor D for reference on host with zeros
|
||||
cutlass::reference::host::TensorFill(
|
||||
tensor_ref_d.host_view());
|
||||
|
||||
// Copy data from host to GPU
|
||||
tensor_a.sync_device();
|
||||
tensor_a_scale.sync_device();
|
||||
tensor_a_bias.sync_device();
|
||||
tensor_b.sync_device();
|
||||
tensor_c.sync_device();
|
||||
tensor_d.sync_device();
|
||||
tensor_ref_d.sync_device();
|
||||
|
||||
//
|
||||
// Define arguments for CUTLASS Convolution
|
||||
//
|
||||
|
||||
cutlass::conv::Mode mode = cutlass::conv::Mode::kCrossCorrelation;
|
||||
|
||||
// Split K dimension into 1 partitions
|
||||
int split_k_slices = 1;
|
||||
|
||||
// Construct Conv3dProblemSize with user defined output size
|
||||
cutlass::conv::Conv3dProblemSize problem_size(
|
||||
options.input_size,
|
||||
options.filter_size,
|
||||
options.padding,
|
||||
options.conv_stride,
|
||||
options.dilation,
|
||||
options.output_size(),
|
||||
mode,
|
||||
split_k_slices
|
||||
);
|
||||
|
||||
typename ImplicitGemmFusion::Arguments arguments{
|
||||
problem_size,
|
||||
tensor_a.device_ref(),
|
||||
tensor_b.device_ref(),
|
||||
tensor_a_scale.device_ref(),
|
||||
tensor_a_bias.device_ref(),
|
||||
tensor_c.device_ref(),
|
||||
tensor_d.device_ref(),
|
||||
{options.alpha, options.beta},
|
||||
};
|
||||
|
||||
//
|
||||
// Initialize CUTLASS Convolution
|
||||
//
|
||||
|
||||
ImplicitGemmFusion implicit_gemm_fusion_op;
|
||||
|
||||
size_t workspace_size = implicit_gemm_fusion_op.get_workspace_size(arguments);
|
||||
|
||||
// Allocate workspace memory
|
||||
cutlass::device_memory::allocation<uint8_t> workspace(workspace_size);
|
||||
|
||||
result.status = implicit_gemm_fusion_op.can_implement(arguments);
|
||||
CUTLASS_CHECK(result.status);
|
||||
|
||||
result.status = implicit_gemm_fusion_op.initialize(arguments, workspace.get());
|
||||
CUTLASS_CHECK(result.status);
|
||||
|
||||
//
|
||||
// Launch initialized CUTLASS kernel
|
||||
//
|
||||
result.status = implicit_gemm_fusion_op();
|
||||
|
||||
CUTLASS_CHECK(result.status);
|
||||
|
||||
//
|
||||
// Optional reference check
|
||||
//
|
||||
|
||||
if (options.reference_check) {
|
||||
std::cout << "Verification on device...\n";
|
||||
|
||||
// Compute scale + bias + relu in host code
|
||||
for (int n = 0; n < options.input_size.n(); ++n) {
|
||||
for (int d = 0; d < options.input_size.d(); ++d) {
|
||||
for (int h = 0; h < options.input_size.h(); ++h) {
|
||||
for (int w = 0; w < options.input_size.w(); ++w) {
|
||||
for (int c = 0; c < options.input_size.c(); ++c) {
|
||||
tensor_transformed_a.at({n, d, h, w, c}) = std::max(
|
||||
ElementOutput(0), ElementOutput(tensor_a.at({n, d, h, w, c}) *
|
||||
tensor_a_scale.at({0, c}) +
|
||||
tensor_a_bias.at({0, c})));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
tensor_transformed_a.sync_device();
|
||||
|
||||
// Compute with reference implementation
|
||||
cutlass::reference::device::Conv3dFprop<
|
||||
ElementInputA,
|
||||
LayoutInputA,
|
||||
ElementInputB,
|
||||
LayoutInputB,
|
||||
ElementOutput,
|
||||
LayoutOutput,
|
||||
ElementComputeEpilogue,
|
||||
ElementAccumulator,
|
||||
cutlass::NumericConverter<ElementOutput, ElementComputeEpilogue>
|
||||
>(
|
||||
problem_size,
|
||||
tensor_transformed_a.device_ref(),
|
||||
tensor_b.device_ref(),
|
||||
tensor_c.device_ref(),
|
||||
tensor_ref_d.device_ref(),
|
||||
options.alpha,
|
||||
options.beta
|
||||
);
|
||||
|
||||
// Check if output from CUTLASS kernel and reference kernel are equal or not
|
||||
tensor_d.sync_host();
|
||||
tensor_ref_d.sync_host();
|
||||
|
||||
bool passed = cutlass::reference::host::TensorEquals(
|
||||
tensor_d.host_view(),
|
||||
tensor_ref_d.host_view());
|
||||
|
||||
if (!passed) {
|
||||
result.reference_check = cutlass::Status::kErrorInternal;
|
||||
std::cout << "ERROR - results miscompared.\n";
|
||||
}
|
||||
else {
|
||||
result.reference_check = cutlass::Status::kSuccess;
|
||||
std::cout << "Passed.\n";
|
||||
}
|
||||
}
|
||||
else {
|
||||
result.reference_check = cutlass::Status::kInvalid;
|
||||
}
|
||||
|
||||
if (options.save_workspace) {
|
||||
|
||||
std::stringstream ss;
|
||||
|
||||
ss << "25_ampere_3d_fprop_mainloop_fusion"
|
||||
<< options.input_size.n() << "x" << options.input_size.h() << "x" << options.input_size.w() << "x" << options.input_size.c()
|
||||
<< "_"
|
||||
<< options.filter_size.n() << "x" << options.filter_size.h() << "x" << options.filter_size.w() << "x" << options.filter_size.c()
|
||||
<< ".dat";
|
||||
|
||||
std::ofstream output_workspace(ss.str());
|
||||
|
||||
output_workspace
|
||||
<< "Input = \n" << tensor_a.host_view() << "\n\n"
|
||||
<< "Filters = \n" << tensor_b.host_view() << "\n\n";
|
||||
|
||||
if (options.reference_check) {
|
||||
output_workspace << "Reference = \n" << tensor_ref_d.host_view() << "\n\n";
|
||||
}
|
||||
|
||||
output_workspace << "Computed = \n" << tensor_d.host_view() << std::endl;
|
||||
|
||||
std::cout << "Results written to '" << ss.str() << "'." << std::endl;
|
||||
}
|
||||
|
||||
//
|
||||
// Performance measurement
|
||||
//
|
||||
|
||||
if (options.measure_performance) {
|
||||
|
||||
cudaEvent_t events[2];
|
||||
|
||||
for (auto & event : events) {
|
||||
result.error = cudaEventCreate(&event);
|
||||
if (result.error != cudaSuccess) {
|
||||
std::cerr << "cudaEventCreate() failed: " << cudaGetErrorString(result.error) << std::endl;
|
||||
return result;
|
||||
}
|
||||
}
|
||||
|
||||
// Record an event at the start of a series of convolution operations.
|
||||
result.error = cudaEventRecord(events[0]);
|
||||
if (result.error != cudaSuccess) {
|
||||
std::cerr << "cudaEventRecord() failed: " << cudaGetErrorString(result.error) << std::endl;
|
||||
return result;
|
||||
}
|
||||
|
||||
// Launch a sequence of implicit GEMM operations on the device
|
||||
for (int iteration = 0; iteration < options.iterations; ++iteration) {
|
||||
result.status = implicit_gemm_fusion_op();
|
||||
CUTLASS_CHECK(result.status);
|
||||
}
|
||||
|
||||
// Record an event when the convolutions have been launched.
|
||||
result.error = cudaEventRecord(events[1]);
|
||||
if (result.error != cudaSuccess) {
|
||||
std::cerr << "cudaEventRecord() failed: " << cudaGetErrorString(result.error) << std::endl;
|
||||
return result;
|
||||
}
|
||||
|
||||
// Wait for work on the device to complete.
|
||||
result.error = cudaEventSynchronize(events[1]);
|
||||
if (result.error != cudaSuccess) {
|
||||
std::cerr << "cudaEventSynchronize() failed: " << cudaGetErrorString(result.error) << std::endl;
|
||||
return result;
|
||||
}
|
||||
|
||||
// Measure elapsed runtime
|
||||
float runtime_ms = 0;
|
||||
result.error = cudaEventElapsedTime(&runtime_ms, events[0], events[1]);
|
||||
if (result.error != cudaSuccess) {
|
||||
std::cerr << "cudaEventElapsed() failed: " << cudaGetErrorString(result.error) << std::endl;
|
||||
return result;
|
||||
}
|
||||
|
||||
// Print average runtime and GFLOPs.
|
||||
result.runtime_ms = double(runtime_ms) / double(options.iterations);
|
||||
result.gflops = options.gflops(result.runtime_ms / 1000.0);
|
||||
|
||||
// Cleanup
|
||||
for (auto event : events) {
|
||||
(void)cudaEventDestroy(event);
|
||||
}
|
||||
}
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
int main(int argc, char const **args) {
|
||||
|
||||
bool notSupported = false;
|
||||
|
||||
// Ampere Tensor Core operations exposed with mma.sync are first available in CUDA 11.0.
|
||||
//
|
||||
// CUTLASS must be compiled with CUDA 11 Toolkit to run Conv3dFprop examples.
|
||||
if (!(__CUDACC_VER_MAJOR__ > 11 || (__CUDACC_VER_MAJOR__ == 11 && __CUDACC_VER_MINOR__ >= 0))) {
|
||||
std::cerr << "Ampere Tensor Core operations must be compiled with CUDA 11.0 Toolkit or later." << std::endl;
|
||||
notSupported = true;
|
||||
}
|
||||
|
||||
cudaDeviceProp props;
|
||||
CUDA_CHECK(cudaGetDeviceProperties(&props, 0));
|
||||
|
||||
if (!(props.major >= 8)) {
|
||||
std::cerr << "This test must run on SM80 or above.\n";
|
||||
notSupported = true;
|
||||
}
|
||||
|
||||
if (notSupported) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
Options options;
|
||||
|
||||
options.parse(argc, args);
|
||||
|
||||
if (options.help) {
|
||||
options.print_usage(std::cout) << std::endl;
|
||||
return 0;
|
||||
}
|
||||
|
||||
if (options.benchmark) {
|
||||
// Benchmark several layers
|
||||
|
||||
int batch_sizes[] = {34, 18};
|
||||
|
||||
struct Benchmark {
|
||||
int d, h, w, c, k, t, r, s, stride_d, stride_h, stride_w;
|
||||
} layers[] = {
|
||||
{56, 56, 56, 64, 256, 1, 1, 1, 1, 1, 1},
|
||||
{56, 56, 56, 64, 64, 1, 1, 1, 1, 1, 1},
|
||||
{56, 56, 56, 64, 64, 3, 3, 3, 1, 1, 1},
|
||||
{56, 56, 56, 256, 64, 1, 1, 1, 1, 1, 1},
|
||||
{56, 56, 56, 256, 512, 1, 1, 1, 2, 2, 2},
|
||||
{56, 56, 56, 256, 128, 1, 1, 1, 1, 1, 1},
|
||||
{56, 56, 56, 128, 128, 3, 3, 3, 2, 2, 2},
|
||||
{28, 28, 28, 128, 512, 1, 1, 1, 1, 1, 1},
|
||||
{28, 28, 28, 512, 128, 1, 1, 1, 1, 1, 1},
|
||||
{28, 28, 28, 128, 128, 3, 3, 3, 1, 1, 1},
|
||||
{28, 28, 28, 512, 1024, 1, 1, 1, 2, 2, 2},
|
||||
{28, 28, 28, 512, 256, 1, 1, 1, 1, 1, 1},
|
||||
{28, 28, 28, 256, 256, 3, 3, 3, 2, 2, 2},
|
||||
{14, 14, 14, 256, 1024, 1, 1, 1, 1, 1, 1},
|
||||
{14, 14, 14, 1024, 256, 1, 1, 1, 1, 1, 1},
|
||||
{14, 14, 14, 256, 256, 3, 3, 3, 1, 1, 1},
|
||||
{14, 14, 14, 1024, 2048, 1, 1, 1, 2, 2, 2},
|
||||
{14, 14, 14, 1024, 512, 1, 1, 1, 1, 1, 1},
|
||||
{14, 14, 14, 512, 512, 3, 3, 3, 2, 2, 2},
|
||||
{ 7, 7, 7, 512, 2048, 1, 1, 1, 1, 1, 1},
|
||||
{ 7, 7, 7, 2048, 512, 1, 1, 1, 1, 1, 1},
|
||||
{ 7, 7, 7, 512, 512, 3, 3, 3, 1, 1, 1},
|
||||
};
|
||||
|
||||
Result::print_header(std::cout, options) << std::endl;
|
||||
|
||||
int idx = 1;
|
||||
|
||||
for (auto const &layer : layers) {
|
||||
for (auto N : batch_sizes) {
|
||||
options.update({N, layer.d, layer.h, layer.w, layer.c},
|
||||
{layer.k, layer.t, layer.r, layer.s, layer.c},
|
||||
cutlass::make_Coord(layer.stride_d, layer.stride_h, layer.stride_w));
|
||||
|
||||
Result result = profile_convolution(options);
|
||||
result.print(std::cout, idx, options) << std::endl;
|
||||
}
|
||||
|
||||
++idx;
|
||||
}
|
||||
}
|
||||
else {
|
||||
|
||||
// Execute one problem size
|
||||
if (!options.valid()) {
|
||||
std::cerr << "Invalid problem." << std::endl;
|
||||
return -1;
|
||||
}
|
||||
|
||||
Result result = profile_convolution(options);
|
||||
|
||||
Result::print_header(std::cout, options) << std::endl;
|
||||
result.print(std::cout, 1, options) << std::endl;
|
||||
}
|
||||
|
||||
return 0;
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -429,9 +429,13 @@ Result profile_convolution(Options const &options) {
|
||||
ElementInputB(-8),
|
||||
0);
|
||||
|
||||
// Fill tensor C on host with zeros
|
||||
cutlass::reference::host::TensorFill(
|
||||
tensor_c.host_view());
|
||||
// Fill tensor C on host with uniform-distribution random data
|
||||
cutlass::reference::host::TensorFillRandomUniform(
|
||||
tensor_c.host_view(),
|
||||
1,
|
||||
ElementOutput(7),
|
||||
ElementOutput(-8),
|
||||
0);
|
||||
|
||||
// Fill tensor D on host with zeros
|
||||
cutlass::reference::host::TensorFill(
|
||||
@@ -575,7 +579,7 @@ Result profile_convolution(Options const &options) {
|
||||
|
||||
std::stringstream ss;
|
||||
|
||||
ss << "25_ampere_fprop_mainloop_fusion_"
|
||||
ss << "25_ampere_fprop_mainloop_fusion"
|
||||
<< options.input_size.n() << "x" << options.input_size.h() << "x" << options.input_size.w() << "x" << options.input_size.c()
|
||||
<< "_"
|
||||
<< options.filter_size.n() << "x" << options.filter_size.h() << "x" << options.filter_size.w() << "x" << options.filter_size.c()
|
||||
@@ -677,8 +681,8 @@ int main(int argc, char const **args) {
|
||||
cudaDeviceProp props;
|
||||
CUDA_CHECK(cudaGetDeviceProperties(&props, 0));
|
||||
|
||||
if (!(props.major == 8 && props.minor == 0)) {
|
||||
std::cerr << "This test must run on SM80 A100.\n";
|
||||
if (!(props.major >= 8)) {
|
||||
std::cerr << "This test must run on SM80 or above.\n";
|
||||
notSupported = true;
|
||||
}
|
||||
|
||||
|
||||
@@ -266,8 +266,8 @@ struct Options {
|
||||
/// Prints the usage statement.
|
||||
std::ostream & print_usage(std::ostream &out) const {
|
||||
|
||||
out << "26_ampere_fused_wgrad_batch_normalization example\n\n"
|
||||
<< " This example fuses scale+bias+relu from batch norm into Ampere's\n"
|
||||
out << "26_ampere_wgrad_mainloop_fusion example\n\n"
|
||||
<< " This example fuses scale+bias+relu of the activation into Ampere's\n"
|
||||
<< " Tensor Core operators on F16 data types to compute\n"
|
||||
<< " backward convolution on tensors of layout NHWC.\n\n"
|
||||
<< "Options:\n\n"
|
||||
@@ -289,8 +289,8 @@ struct Options {
|
||||
<< " --tag=<string> String to replicate across the first column in the results table\n";
|
||||
|
||||
out << "\n\nExamples:\n\n"
|
||||
<< "$ ./examples/26_ampere_fused_fprop_batch_normalization/26_ampere_fused_wgrad_batch_normalization --n=32 --h=224 --w=224 --c=128 --k=256 --r=1 --s=1\n\n"
|
||||
<< "$ ./examples/26_ampere_fused_fprop_batch_normalization/26_ampere_fused_wgrad_batch_normalization --n=1 --h=224 --w=224 --c=32 --k=32 --r=3 --s=3 --ref-check\n\n";
|
||||
<< "$ ./examples/26_ampere_wgrad_mainloop_fusion/26_ampere_wgrad_mainloop_fusion --n=32 --h=224 --w=224 --c=128 --k=256 --r=1 --s=1\n\n"
|
||||
<< "$ ./examples/26_ampere_wgrad_mainloop_fusion/26_ampere_wgrad_mainloop_fusion --n=1 --h=224 --w=224 --c=32 --k=32 --r=3 --s=3 --ref-check\n\n";
|
||||
|
||||
return out;
|
||||
}
|
||||
@@ -427,9 +427,13 @@ Result profile_convolution(Options const &options) {
|
||||
ElementInputA(-4),
|
||||
0);
|
||||
|
||||
// Fill tensor C on host with zeros
|
||||
cutlass::reference::host::TensorFill(
|
||||
tensor_c.host_view());
|
||||
// Fill tensor C on host with uniform-distribution random data
|
||||
cutlass::reference::host::TensorFillRandomUniform(
|
||||
tensor_c.host_view(),
|
||||
1,
|
||||
ElementOutput(7),
|
||||
ElementOutput(-8),
|
||||
0);
|
||||
|
||||
// Fill tensor D on host with zeros
|
||||
cutlass::reference::host::TensorFill(
|
||||
|
||||
@@ -740,7 +740,7 @@ int main(int argc, char const **args) {
|
||||
cudaDeviceProp props;
|
||||
CUDA_CHECK(cudaGetDeviceProperties(&props, 0));
|
||||
|
||||
if (!(props.major > 8 || (props.major == 8 && props.minor >= 0))) {
|
||||
if (!(props.major >= 8)) {
|
||||
std::cerr << "Ampere Tensor Ops must be run on a machine with compute capability at least 80."
|
||||
<< std::endl;
|
||||
notSupported = true;
|
||||
|
||||
@@ -703,7 +703,7 @@ int main(int argc, char const **args) {
|
||||
cudaDeviceProp props;
|
||||
CUDA_CHECK(cudaGetDeviceProperties(&props, 0));
|
||||
|
||||
if (!(props.major > 8 || (props.major == 8 && props.minor >= 0))) {
|
||||
if (!(props.major >= 8)) {
|
||||
std::cerr << "Ampere Tensor Ops must be run on a machine with compute capability at least 80."
|
||||
<< std::endl;
|
||||
notSupported = true;
|
||||
|
||||
@@ -603,7 +603,7 @@ int main(int argc, char const **args) {
|
||||
cudaDeviceProp props;
|
||||
CUDA_CHECK(cudaGetDeviceProperties(&props, 0));
|
||||
|
||||
if (!(props.major > 8 || (props.major == 8 && props.minor >= 0))) {
|
||||
if (!(props.major >= 8)) {
|
||||
std::cerr << "Ampere Tensor Ops must be run on a machine with compute capability at least 80."
|
||||
<< std::endl;
|
||||
notSupported = true;
|
||||
|
||||
@@ -1,407 +0,0 @@
|
||||
/***************************************************************************************************
|
||||
* 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 Generic epilogue for implementing certain kinds of fused epilogue behavior.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/fast_math.h"
|
||||
#include "cutlass/matrix_coord.h"
|
||||
#include "cutlass/semaphore.h"
|
||||
#include "cutlass/epilogue/threadblock/epilogue_base.h"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace epilogue {
|
||||
namespace threadblock {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
class EpilogueFusedVisitorConcept {
|
||||
public:
|
||||
|
||||
static int const kIterations = 1;
|
||||
static int const kElementsPerAccess = 4;
|
||||
using ElementOutput = float;
|
||||
using ElementAccumulator = float;
|
||||
using AccumulatorFragment = Array<ElementAccumulator, kElementsPerAccess>;
|
||||
|
||||
/// Arguments structure
|
||||
struct Arguments { };
|
||||
|
||||
/// Params structure
|
||||
struct Params {
|
||||
|
||||
Params() { }
|
||||
Params(Arguments const &args) { }
|
||||
};
|
||||
|
||||
/// Shared storage
|
||||
struct SharedStorage { };
|
||||
|
||||
public:
|
||||
|
||||
CUTLASS_DEVICE
|
||||
EpilogueFusedVisitorConcept(
|
||||
Params const ¶ms, ///< Parameters routed to the epilogue
|
||||
SharedStorage &shared_storage, ///< Shared storage needed by the functors here
|
||||
MatrixCoord const &problem_size, ///< Problem size of the output
|
||||
int thread_idx, ///< Thread index within the threadblock
|
||||
int warp_idx, ///< Warp index within the threadblock
|
||||
int lane_idx, ///< Lane index within the warp
|
||||
MatrixCoord const &threadblock_offset = MatrixCoord(0, 0)) { ///< Coordinate
|
||||
|
||||
}
|
||||
|
||||
/// Helper to indicate split-K behavior
|
||||
CUTLASS_DEVICE
|
||||
void set_k_partition(
|
||||
int split_k_index, ///< Index of this threadblock within split-K partitioned scheme
|
||||
int split_k_slices) { ///< Total number of split-K slices
|
||||
|
||||
}
|
||||
|
||||
/// Called to set the batch index
|
||||
CUTLASS_DEVICE
|
||||
void set_batch_index(int batch_idx) {
|
||||
|
||||
}
|
||||
|
||||
/// Called at the start of the epilogue just before iterating over accumulator slices
|
||||
CUTLASS_DEVICE
|
||||
void begin_epilogue() {
|
||||
|
||||
}
|
||||
|
||||
/// Called at the start of one step before starting accumulator exchange
|
||||
CUTLASS_DEVICE
|
||||
void begin_step(int step_idx) {
|
||||
|
||||
}
|
||||
|
||||
/// Called at the start of a row
|
||||
CUTLASS_DEVICE
|
||||
void begin_row(int row_idx) {
|
||||
|
||||
}
|
||||
|
||||
/// Called after accumulators have been exchanged for each accumulator vector
|
||||
CUTLASS_DEVICE
|
||||
void visit(
|
||||
int row_idx,
|
||||
int column_idx,
|
||||
int frag_idx,
|
||||
AccumulatorFragment const &accum) {
|
||||
|
||||
}
|
||||
|
||||
/// Called at the start of a row
|
||||
CUTLASS_DEVICE
|
||||
void end_row(int row_idx) {
|
||||
|
||||
}
|
||||
|
||||
/// Called after all accumulator elements have been visited
|
||||
CUTLASS_DEVICE
|
||||
void end_step(int step_idx) {
|
||||
|
||||
}
|
||||
|
||||
/// Called after all steps have been completed
|
||||
CUTLASS_DEVICE
|
||||
void end_epilogue() {
|
||||
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Epilogue operator
|
||||
template <
|
||||
typename Visitor_, ///< Functor containing fused operations (satisfies EpilogueFusedVisitorConcept)
|
||||
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 AccumulatorFragmentIterator_, ///< Fragment iterator selecting accumulators
|
||||
typename WarpTileIterator_, ///< Warp-scoped tile iterator writing accumulators to SMEM
|
||||
typename SharedLoadIterator_, ///< Threadblock-scoped tile iterator loading from SMEM
|
||||
typename Padding_, ///< Padding added to SMEM allocation to avoid bank conflicts (concept: MatrixShape)
|
||||
int FragmentsPerPartition = 1, ///< Used to coarsten the epilogue granularity
|
||||
int IterationsUnroll = ///< Used to reduce binary size when epilogue op is large
|
||||
(true || !IsEpilogueFunctorHeavy<Visitor_>::value)
|
||||
>
|
||||
class EpilogueWithVisitor :
|
||||
public EpilogueBase<
|
||||
Shape_,
|
||||
typename WarpMmaOperator_::Shape,
|
||||
PartitionsK,
|
||||
AccumulatorFragmentIterator_,
|
||||
WarpTileIterator_,
|
||||
Padding_,
|
||||
FragmentsPerPartition> {
|
||||
|
||||
public:
|
||||
|
||||
using Visitor = Visitor_;
|
||||
|
||||
using Base = EpilogueBase<
|
||||
Shape_,
|
||||
typename WarpMmaOperator_::Shape,
|
||||
PartitionsK,
|
||||
AccumulatorFragmentIterator_,
|
||||
WarpTileIterator_,
|
||||
Padding_,
|
||||
FragmentsPerPartition>;
|
||||
|
||||
using Shape = Shape_;
|
||||
using WarpMmaOperator = WarpMmaOperator_;
|
||||
static int const kPartitionsK = PartitionsK;
|
||||
|
||||
using AccumulatorFragmentIterator = AccumulatorFragmentIterator_;
|
||||
using WarpTileIterator = WarpTileIterator_;
|
||||
using SharedLoadIterator = SharedLoadIterator_;
|
||||
using Padding = Padding_;
|
||||
|
||||
using Layout = layout::RowMajor;
|
||||
using LongIndex = typename Layout::LongIndex;
|
||||
|
||||
/// The complete warp-level accumulator tile
|
||||
using AccumulatorTile = typename Base::AccumulatorTile;
|
||||
|
||||
/// Accumulator element
|
||||
using ElementAccumulator = typename WarpTileIterator::Element;
|
||||
|
||||
/// Output access size
|
||||
static int const kElementsPerAccess = Visitor::kElementsPerAccess;
|
||||
|
||||
/// Tensor reference to sync tensor
|
||||
using SyncTensorRef = typename cutlass::TensorRef<int, cutlass::layout::PackedVectorLayout>;
|
||||
|
||||
/// Array type used by output functor
|
||||
using AccumulatorAccessType = Array<
|
||||
typename WarpTileIterator::Element, kElementsPerAccess>;
|
||||
|
||||
/// Number of warps
|
||||
using WarpCount = typename Base::WarpCount;
|
||||
|
||||
static int constexpr kSmemTiles = Base::kFragmentsPerIteration > 1 ? Base::kFragmentsPerIteration : kPartitionsK;
|
||||
static int constexpr kSmemPointerOffset = Base::SharedStorage::StorageShape::kCount / kSmemTiles;
|
||||
|
||||
using SharedStorage = typename Base::SharedStorage;
|
||||
|
||||
private:
|
||||
|
||||
/// Loads fragment from shared memory aligned with output tensor
|
||||
SharedLoadIterator shared_load_iterator_;
|
||||
|
||||
public:
|
||||
|
||||
/// Constructor
|
||||
CUTLASS_DEVICE
|
||||
EpilogueWithVisitor(
|
||||
SharedStorage &shared_storage, ///< Shared storage object
|
||||
int thread_idx, ///< ID of a thread within the threadblock
|
||||
int warp_idx, ///< ID of warp within threadblock
|
||||
int lane_idx ///< Id of thread within warp
|
||||
):
|
||||
Base(shared_storage, thread_idx, warp_idx, lane_idx),
|
||||
shared_load_iterator_(shared_storage.reference(), thread_idx)
|
||||
{
|
||||
|
||||
}
|
||||
|
||||
/// Streams the result to global memory
|
||||
CUTLASS_DEVICE
|
||||
void operator()(
|
||||
Visitor & visitor,
|
||||
AccumulatorTile const &accumulators) { ///< Threadblock tile coordinate in GEMM (in units of threadblock tiles)
|
||||
|
||||
visitor.begin_epilogue();
|
||||
|
||||
//
|
||||
// Iterator over warp-level accumulator fragment
|
||||
//
|
||||
|
||||
AccumulatorFragmentIterator accum_fragment_iterator(accumulators);
|
||||
|
||||
//
|
||||
// Iterate over accumulator tile
|
||||
//
|
||||
|
||||
#pragma unroll(IterationsUnroll ? Visitor::kIterations : 1)
|
||||
for (int iter_idx = 0; iter_idx < Visitor::kIterations; ++iter_idx) {
|
||||
|
||||
//
|
||||
// Load the source
|
||||
//
|
||||
|
||||
visitor.begin_step(iter_idx);
|
||||
|
||||
//
|
||||
// Convert and store fragment
|
||||
//
|
||||
|
||||
__syncthreads();
|
||||
|
||||
acc2smem_source_needed<cutlass::make_index_sequence<Visitor::kIterations>>::push(
|
||||
iter_idx, accum_fragment_iterator, this->warp_tile_iterator_);
|
||||
|
||||
__syncthreads();
|
||||
|
||||
//
|
||||
// Load fragments from shared memory
|
||||
//
|
||||
|
||||
typename SharedLoadIterator::Fragment aligned_accum_fragment[kPartitionsK];
|
||||
|
||||
shared_load_iterator_.load(aligned_accum_fragment[0]);
|
||||
|
||||
// If the number of k-slices is > 1 - perform a reduction amongst the k-slices
|
||||
if (kPartitionsK > 1) {
|
||||
|
||||
plus <typename SharedLoadIterator::Fragment> add_fragments;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for ( int i = 1; i < kPartitionsK; ++i) {
|
||||
shared_load_iterator_.add_pointer_offset(kSmemPointerOffset);
|
||||
shared_load_iterator_.load(aligned_accum_fragment[i]);
|
||||
aligned_accum_fragment[0] = add_fragments(aligned_accum_fragment[0], aligned_accum_fragment[i]);
|
||||
}
|
||||
|
||||
shared_load_iterator_.add_pointer_offset((1 - kPartitionsK) * kSmemPointerOffset);
|
||||
}
|
||||
|
||||
//
|
||||
// Iterate over output fragments
|
||||
//
|
||||
|
||||
AccumulatorAccessType const *accum_frag_ptr =
|
||||
reinterpret_cast<AccumulatorAccessType const *>(&aligned_accum_fragment[0]);
|
||||
|
||||
int const kAccumulatorFragmentCount = AccumulatorTile::kElements / (Visitor::kIterations * AccumulatorAccessType::kElements);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int idx = 0; idx < kAccumulatorFragmentCount; ++idx) {
|
||||
|
||||
int row_idx = idx / SharedLoadIterator::ThreadMap::Iterations::kColumn;
|
||||
int col_idx = idx % SharedLoadIterator::ThreadMap::Iterations::kColumn;
|
||||
|
||||
// Start a new row of the output fragment
|
||||
if (!col_idx) {
|
||||
visitor.begin_row(row_idx);
|
||||
}
|
||||
|
||||
visitor.visit(
|
||||
row_idx,
|
||||
col_idx,
|
||||
idx,
|
||||
accum_frag_ptr[idx]
|
||||
);
|
||||
|
||||
// End the row of the output fragment
|
||||
if (col_idx + 1 == SharedLoadIterator::ThreadMap::Iterations::kColumn) {
|
||||
visitor.end_row(row_idx);
|
||||
}
|
||||
}
|
||||
|
||||
//
|
||||
// Conclude the step
|
||||
//
|
||||
|
||||
visitor.end_step(iter_idx);
|
||||
}
|
||||
|
||||
visitor.end_epilogue();
|
||||
}
|
||||
|
||||
private:
|
||||
|
||||
|
||||
template<class Seq>
|
||||
struct acc2smem_source_needed;
|
||||
|
||||
template <size_t... Seq>
|
||||
struct acc2smem_source_needed<cutlass::index_sequence<Seq...>> {
|
||||
template<int Advance>
|
||||
CUTLASS_DEVICE
|
||||
static void helper(AccumulatorFragmentIterator accum_fragment_iterator,
|
||||
WarpTileIterator &warp_tile_iterator) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < Advance; i++) {
|
||||
++accum_fragment_iterator;
|
||||
}
|
||||
|
||||
typename AccumulatorFragmentIterator::Fragment accum_fragment;
|
||||
accum_fragment_iterator.load(accum_fragment);
|
||||
warp_tile_iterator.store(accum_fragment);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static void push(size_t pos,
|
||||
AccumulatorFragmentIterator const &iterator_begin,
|
||||
WarpTileIterator &warp_tile_iterator) {
|
||||
int dummy[] = {(pos == Seq) && (helper<Seq>(iterator_begin, warp_tile_iterator), 0)...};
|
||||
}
|
||||
};
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Helper to create an EpilogueWithVisitor from an existing epilogue
|
||||
template <typename Visitor_, typename Existing_, bool IterationsUnroll = true>
|
||||
struct EpilogueWithVisitorFromExistingEpilogue {
|
||||
|
||||
using Epilogue = EpilogueWithVisitor<
|
||||
Visitor_,
|
||||
typename Existing_::Shape,
|
||||
typename Existing_::WarpMmaOperator,
|
||||
Existing_::kPartitionsK,
|
||||
typename Existing_::AccumulatorFragmentIterator,
|
||||
typename Existing_::WarpTileIterator,
|
||||
typename Existing_::SharedLoadIterator,
|
||||
typename Existing_::Padding,
|
||||
Existing_::kFragmentsPerIteration,
|
||||
IterationsUnroll
|
||||
>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace threadblock
|
||||
} // namespace epilogue
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -47,14 +47,17 @@
|
||||
#include "cutlass/util/host_tensor.h"
|
||||
|
||||
#include "cutlass/util/reference/host/gemm_complex.h"
|
||||
#include "cutlass/util/reference/device/gemm_complex.h"
|
||||
#include "cutlass/util/reference/host/tensor_reduce.h"
|
||||
#include "cutlass/util/reference/host/tensor_compare.h"
|
||||
#include "cutlass/util/reference/host/tensor_norm.h"
|
||||
#include "cutlass/util/reference/host/tensor_copy.h"
|
||||
#include "cutlass/util/reference/device/tensor_fill.h"
|
||||
#include "cutlass/util/reference/host/tensor_fill.h"
|
||||
#include "cutlass/util/reference/host/error_metrics.h"
|
||||
#include "cutlass/util/tensor_view_io.h"
|
||||
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -85,18 +88,18 @@ struct Options {
|
||||
float alpha;
|
||||
float beta;
|
||||
bool verification_enabled;
|
||||
double tolerance;
|
||||
float tolerance;
|
||||
|
||||
Options():
|
||||
help(false),
|
||||
problem_size({16, 24, 64}),
|
||||
batch_count(1), // As a temporary limitation to the test bench, batch count must be 1. The kernels support arbitrary batching.
|
||||
batch_count(16),
|
||||
iterations(20),
|
||||
seed(2022),
|
||||
alpha(1),
|
||||
beta(),
|
||||
beta(0),
|
||||
verification_enabled(true),
|
||||
tolerance(0.01)
|
||||
tolerance(1e-5f)
|
||||
{ }
|
||||
|
||||
bool valid() {
|
||||
@@ -116,6 +119,8 @@ struct Options {
|
||||
cmd.get_cmd_line_argument("n", problem_size.n());
|
||||
cmd.get_cmd_line_argument("k", problem_size.k());
|
||||
|
||||
cmd.get_cmd_line_argument("batch_count", batch_count);
|
||||
|
||||
cmd.get_cmd_line_argument("alpha", alpha);
|
||||
cmd.get_cmd_line_argument("beta", beta);
|
||||
|
||||
@@ -135,6 +140,7 @@ struct Options {
|
||||
<< " --m=<int> GEMM M dimension\n"
|
||||
<< " --n=<int> GEMM N dimension\n"
|
||||
<< " --k=<int> GEMM K dimension\n"
|
||||
<< " --batch_count=<int> Batch number\n"
|
||||
<< " --alpha=<f32> Epilogue scalar alpha\n"
|
||||
<< " --beta=<f32> Epilogue scalar beta\n\n"
|
||||
<< " --seed=<int> Random number seed (1*)\n\n"
|
||||
@@ -198,13 +204,22 @@ struct Testbed {
|
||||
using ElementA = cutlass::half_t;
|
||||
using ElementB = cutlass::half_t;
|
||||
using ElementC = cutlass::half_t;
|
||||
using ElementD = cutlass::half_t;
|
||||
using ElementCompute = float;
|
||||
using ElementSoftmax = cutlass::half_t;
|
||||
using ElementD = ElementC;
|
||||
using ElementSoftmax = ElementC;
|
||||
|
||||
using LayoutA = cutlass::layout::RowMajor;
|
||||
using LayoutB = cutlass::layout::ColumnMajor;
|
||||
|
||||
using ThreadblockShape = cutlass::gemm::GemmShape<128, 128, 32>;
|
||||
using WarpShape = cutlass::gemm::GemmShape<64, 64, 32>;
|
||||
using InstructionShape = cutlass::gemm::GemmShape<16, 8, 16>;
|
||||
|
||||
using OperatorClass = cutlass::arch::OpClassTensorOp;
|
||||
using ArchTag = cutlass::arch::Sm80;
|
||||
|
||||
static int const kStages = 3;
|
||||
|
||||
/// Linear scaling operator
|
||||
using EpilogueFunctorOp = cutlass::epilogue::thread::LinearCombination<
|
||||
ElementC,
|
||||
@@ -218,12 +233,21 @@ struct Testbed {
|
||||
ElementB, LayoutB,
|
||||
ElementC,
|
||||
ElementCompute,
|
||||
EpilogueFunctorOp
|
||||
OperatorClass,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueFunctorOp,
|
||||
kStages
|
||||
>;
|
||||
|
||||
using ElementNorm = typename GemmSoftmax::ElementNorm;
|
||||
using ElementSum = typename GemmSoftmax::ElementSum;
|
||||
using LayoutC = typename GemmSoftmax::LayoutC;
|
||||
using LayoutN = typename GemmSoftmax::LayoutN;
|
||||
using LayoutS = typename GemmSoftmax::LayoutS;
|
||||
using MatrixCoord = typename LayoutC::TensorCoord;
|
||||
|
||||
//
|
||||
// Data members
|
||||
@@ -231,20 +255,42 @@ struct Testbed {
|
||||
|
||||
Options const &options;
|
||||
|
||||
cutlass::HostTensor<ElementA, LayoutA> tensor_A;
|
||||
cutlass::HostTensor<ElementB, LayoutB> tensor_B;
|
||||
cutlass::HostTensor<ElementC, LayoutC> tensor_C;
|
||||
cutlass::HostTensor<ElementD, LayoutC> tensor_D;
|
||||
cutlass::HostTensor<ElementNorm, LayoutC> tensor_N;
|
||||
cutlass::HostTensor<ElementSum, LayoutC> tensor_S;
|
||||
cutlass::HostTensor<ElementSoftmax, LayoutC> tensor_Softmax;
|
||||
|
||||
cutlass::HostTensor<ElementD, LayoutC> reference_D;
|
||||
cutlass::HostTensor<ElementNorm, LayoutC> reference_N;
|
||||
cutlass::HostTensor<ElementSoftmax, LayoutC> reference_Softmax;
|
||||
|
||||
cutlass::DeviceAllocation<ElementA> block_A;
|
||||
cutlass::DeviceAllocation<ElementB> block_B;
|
||||
cutlass::DeviceAllocation<ElementC> block_C;
|
||||
cutlass::DeviceAllocation<ElementD> block_D;
|
||||
cutlass::DeviceAllocation<ElementD> block_Ref;
|
||||
cutlass::DeviceAllocation<ElementSoftmax> block_Softmax;
|
||||
cutlass::DeviceAllocation<ElementNorm> block_Norm;
|
||||
cutlass::DeviceAllocation<ElementSum> block_Sum;
|
||||
|
||||
int block_num = (options.problem_size.n() + GemmSoftmax::ThreadblockShape::kN - 1) / GemmSoftmax::ThreadblockShape::kN;
|
||||
|
||||
cutlass::gemm::GemmCoord problem = options.problem_size;
|
||||
|
||||
int64_t lda = LayoutA::packed({problem.m(), problem.k()}).stride(0);
|
||||
int64_t ldb = LayoutB::packed({problem.k(), problem.n()}).stride(0);
|
||||
int64_t ldc = LayoutC::packed({problem.m(), problem.n()}).stride(0);
|
||||
|
||||
// fixed rowmajor for norm and sum
|
||||
int64_t ldn = problem.m();
|
||||
int64_t lds = ldn;
|
||||
|
||||
int64_t total_elements_A_per_batch = problem.m() * problem.k();
|
||||
int64_t total_elements_B_per_batch = problem.k() * problem.n();
|
||||
int64_t total_elements_C_per_batch = problem.m() * problem.n();
|
||||
int64_t total_elements_D_per_batch = problem.m() * problem.n();
|
||||
int64_t total_elements_partial_norm_per_batch = block_num * problem.m();
|
||||
|
||||
int64_t total_elements_A = total_elements_A_per_batch * options.batch_count;
|
||||
int64_t total_elements_B = total_elements_B_per_batch * options.batch_count;
|
||||
int64_t total_elements_C = total_elements_C_per_batch * options.batch_count;
|
||||
int64_t total_elements_D = total_elements_D_per_batch * options.batch_count;
|
||||
int64_t total_elements_partial_norm = total_elements_partial_norm_per_batch * options.batch_count;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
@@ -254,20 +300,7 @@ struct Testbed {
|
||||
):
|
||||
options(options_)
|
||||
{
|
||||
|
||||
tensor_A.reset({options.problem_size.m(), options.problem_size.k()});
|
||||
tensor_B.reset({options.problem_size.k(), options.problem_size.n()});
|
||||
|
||||
tensor_C.reset({options.problem_size.m(), options.problem_size.n()});
|
||||
tensor_D.reset({options.problem_size.m(), options.problem_size.n()});
|
||||
|
||||
tensor_N.reset({block_num, options.problem_size.m()});
|
||||
tensor_S.reset({block_num, options.problem_size.m()});
|
||||
tensor_Softmax.reset({options.problem_size.m(), options.problem_size.n()});
|
||||
|
||||
reference_D.reset({options.problem_size.m(), options.problem_size.n()}, false);
|
||||
reference_N.reset({options.problem_size.m(), 1}, false);
|
||||
reference_Softmax.reset({options.problem_size.m(), options.problem_size.n()}, false);
|
||||
}
|
||||
|
||||
/// Run
|
||||
@@ -300,11 +333,6 @@ struct Testbed {
|
||||
return disposition;
|
||||
}
|
||||
|
||||
//
|
||||
// Compute the reference
|
||||
//
|
||||
compute_reference();
|
||||
|
||||
//
|
||||
// Verify
|
||||
//
|
||||
@@ -334,43 +362,38 @@ struct Testbed {
|
||||
/// Random initialization
|
||||
void initialize() {
|
||||
|
||||
cutlass::reference::host::TensorFillRandomUniform(
|
||||
tensor_A.host_view(),
|
||||
options.seed,
|
||||
ElementD(5),
|
||||
ElementD(-5),
|
||||
0
|
||||
);
|
||||
block_A.reset(total_elements_A);
|
||||
block_B.reset(total_elements_B);
|
||||
block_C.reset(total_elements_C);
|
||||
block_D.reset(total_elements_D);
|
||||
block_Softmax.reset(total_elements_D);
|
||||
block_Ref.reset(total_elements_D_per_batch);
|
||||
block_Norm.reset(total_elements_partial_norm);
|
||||
block_Sum.reset(total_elements_partial_norm);
|
||||
|
||||
cutlass::reference::host::TensorFillRandomUniform(
|
||||
tensor_B.host_view(),
|
||||
options.seed + 19,
|
||||
ElementD(5),
|
||||
ElementD(-5),
|
||||
0
|
||||
);
|
||||
cutlass::reference::device::BlockFillRandomUniform(
|
||||
block_A.get(), total_elements_A, options.seed, ElementA(5), ElementA(-5), 0);
|
||||
|
||||
cutlass::reference::host::TensorFill(
|
||||
reference_D.host_view(),
|
||||
ElementD()
|
||||
);
|
||||
cutlass::reference::device::BlockFillRandomUniform(
|
||||
block_B.get(), total_elements_B, options.seed + 1, ElementB(5), ElementB(-5), 0);
|
||||
|
||||
cutlass::reference::device::BlockFillRandomUniform(
|
||||
block_C.get(), total_elements_C, options.seed + 2, ElementC(5), ElementC(-5), 0);
|
||||
|
||||
cutlass::reference::device::BlockFillRandomUniform(
|
||||
block_D.get(), total_elements_D, options.seed + 3, ElementD(5), ElementD(-5), 0);
|
||||
|
||||
cutlass::reference::device::BlockFillRandomUniform(
|
||||
block_Ref.get(), total_elements_D_per_batch, options.seed + 3, ElementD(5), ElementD(-5), 0);
|
||||
|
||||
cutlass::reference::device::BlockFillRandomUniform(
|
||||
block_Softmax.get(), total_elements_D, options.seed + 3, ElementSoftmax(5), ElementSoftmax(-5), 0);
|
||||
|
||||
cutlass::reference::host::TensorFill(
|
||||
reference_N.host_view(),
|
||||
ElementNorm()
|
||||
);
|
||||
|
||||
cutlass::reference::host::TensorFill(
|
||||
reference_Softmax.host_view(),
|
||||
ElementSoftmax()
|
||||
);
|
||||
|
||||
tensor_A.sync_device();
|
||||
tensor_B.sync_device();
|
||||
tensor_D.sync_device();
|
||||
tensor_N.sync_device();
|
||||
tensor_S.sync_device();
|
||||
tensor_Softmax.sync_device();
|
||||
}
|
||||
|
||||
cutlass::Status execute_device_kernel() {
|
||||
@@ -384,17 +407,24 @@ struct Testbed {
|
||||
GemmSoftmax::Arguments args(
|
||||
options.problem_size,
|
||||
options.batch_count,
|
||||
tensor_A.device_ref(),
|
||||
tensor_B.device_ref(),
|
||||
tensor_C.device_ref(),
|
||||
tensor_D.device_ref(),
|
||||
{block_A.get(), lda},
|
||||
{block_B.get(), ldb},
|
||||
{block_C.get(), ldc},
|
||||
{block_D.get(), ldc},
|
||||
{
|
||||
ElementCompute(options.alpha),
|
||||
ElementCompute(options.beta)
|
||||
},
|
||||
tensor_N.device_ref(),
|
||||
tensor_S.device_ref(),
|
||||
tensor_Softmax.device_ref()
|
||||
{block_Norm.get(), ldn},
|
||||
{block_Sum.get(), lds},
|
||||
{block_Softmax.get(), ldc},
|
||||
total_elements_A_per_batch,
|
||||
total_elements_B_per_batch,
|
||||
total_elements_C_per_batch,
|
||||
total_elements_D_per_batch,
|
||||
total_elements_partial_norm_per_batch,
|
||||
total_elements_partial_norm_per_batch,
|
||||
total_elements_D_per_batch
|
||||
);
|
||||
|
||||
//
|
||||
@@ -415,68 +445,21 @@ struct Testbed {
|
||||
return status;
|
||||
}
|
||||
|
||||
/// Reference calculation
|
||||
void compute_reference() {
|
||||
template<typename Element>
|
||||
bool verify_tensor(std::vector<Element> vector_Input, \
|
||||
std::vector<Element> vector_Input_Ref) {
|
||||
|
||||
// Compute GEMM
|
||||
|
||||
cutlass::reference::host::GemmComplex(
|
||||
options.problem_size,
|
||||
options.alpha,
|
||||
tensor_A.host_ref(),
|
||||
cutlass::ComplexTransform::kNone,
|
||||
tensor_B.host_ref(),
|
||||
cutlass::ComplexTransform::kNone,
|
||||
options.beta,
|
||||
tensor_C.host_ref(),
|
||||
reference_D.host_ref(),
|
||||
double()
|
||||
);
|
||||
|
||||
// Compute the norm
|
||||
for (int m = 0; m < options.problem_size.m(); ++m) {
|
||||
reference_N.at({m, 0}) = reference_D.at({m, 0});
|
||||
for (int n = 1; n < options.problem_size.n(); ++n) {
|
||||
reference_N.at({m, 0}) = std::max(reference_N.at({m, 0}), ElementNorm(reference_D.at({m, n})));
|
||||
}
|
||||
}
|
||||
|
||||
// Compute softmax
|
||||
for (int m = 0; m < options.problem_size.m(); ++m) {
|
||||
|
||||
float sum = float();
|
||||
|
||||
for (int n = 0; n < options.problem_size.n(); ++n) {
|
||||
sum += std::exp( float(reference_D.at({m, n})) - float(reference_N.at({m, 0})) );
|
||||
}
|
||||
|
||||
float inv_sum = float(1.0f / sum);
|
||||
|
||||
for (int n = 0; n < options.problem_size.n(); ++n) {
|
||||
|
||||
reference_Softmax.at({m, n}) = ElementSoftmax(
|
||||
std::exp( float(reference_D.at({m, n})) - float(reference_N.at({m, 0})) ) * inv_sum
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Emits all tensor values
|
||||
void emit_results() {
|
||||
std::cout << "D = \n" << tensor_D.host_view() << "\n\n";
|
||||
std::cout << "N = \n" << tensor_N.host_view() << "\n\n";
|
||||
std::cout << "Softmax = \n" << tensor_Softmax.host_view() << "\n\n";
|
||||
std::cout << "Reference N = \n" << reference_N.host_view() << "\n\n";
|
||||
std::cout << "Reference D = \n" << reference_D.host_view() << "\n\n";
|
||||
std::cout << "Reference Softmax = \n" << reference_Softmax.host_view() << "\n\n";
|
||||
}
|
||||
|
||||
bool verify_tensor_N(cutlass::HostTensor<ElementNorm, LayoutC> tensor_N, \
|
||||
cutlass::HostTensor<ElementNorm, LayoutC> reference_N) {
|
||||
|
||||
for (int m = 0; m < options.problem_size.m(); ++m) {
|
||||
float diff = (float)(tensor_N.at({0, m}) - reference_N.at({m, 0}));
|
||||
if (fabs(diff) > options.tolerance) {
|
||||
int64_t size = (vector_Input.size() < vector_Input_Ref.size()) ? vector_Input.size() : vector_Input_Ref.size();
|
||||
float abs_tol = options.tolerance;
|
||||
float rel_tol = options.tolerance;
|
||||
|
||||
for (int64_t i = 0; i < size; ++i) {
|
||||
float diff = (float)(vector_Input.at(i) - vector_Input_Ref.at(i));
|
||||
float abs_diff = fabs(diff);
|
||||
float abs_ref = fabs((float)vector_Input_Ref.at(i));
|
||||
float relative_diff = abs_ref > abs_tol ? abs_diff / abs_ref : 0;
|
||||
if ( (isnan(abs_diff) || isinf(abs_diff)) || (abs_diff > rel_tol && relative_diff > rel_tol)) {
|
||||
printf("diff = %f, {%f, %f}.\n", abs_diff, (float)(vector_Input.at(i)), (float)(vector_Input_Ref.at(i)));
|
||||
return false;
|
||||
}
|
||||
|
||||
@@ -488,80 +471,112 @@ struct Testbed {
|
||||
/// Verifies the reference matches
|
||||
bool verify() {
|
||||
|
||||
tensor_D.sync_host();
|
||||
tensor_N.sync_host();
|
||||
tensor_Softmax.sync_host();
|
||||
LayoutA layout_A(lda);
|
||||
LayoutB layout_B(ldb);
|
||||
LayoutC layout_C(ldc);
|
||||
LayoutN Layout_N(ldn);
|
||||
LayoutS Layout_S(lds);
|
||||
|
||||
double const kThreshold = options.tolerance;
|
||||
MatrixCoord extent_A{problem.m(), problem.k()};
|
||||
MatrixCoord extent_B{problem.k(), problem.n()};
|
||||
MatrixCoord extent_C{problem.m(), problem.n()};
|
||||
|
||||
// Verification checks - set any of these to 'true' to override the verification checks.
|
||||
bool verified_D = false;
|
||||
bool verified_N = false;
|
||||
bool verified_Softmax = false;
|
||||
for (int batch_idx = 0; batch_idx < options.batch_count; batch_idx++) {
|
||||
|
||||
// Verify softmax output
|
||||
if (!verified_D) {
|
||||
cutlass::TensorView<ElementA, LayoutA> view_A(block_A.get() + total_elements_A_per_batch * batch_idx, layout_A, extent_A);
|
||||
cutlass::TensorView<ElementB, LayoutB> view_B(block_B.get() + total_elements_B_per_batch * batch_idx, layout_B, extent_B);
|
||||
cutlass::TensorView<ElementC, LayoutC> view_C(block_C.get() + total_elements_C_per_batch * batch_idx, layout_C, extent_C);
|
||||
cutlass::TensorView<ElementC, LayoutC> view_Ref_device(block_Ref.get(), layout_C, extent_C);
|
||||
|
||||
double norm_diff = cutlass::reference::host::TensorNormDiff(
|
||||
tensor_D.host_view(),
|
||||
reference_D.host_view());
|
||||
cutlass::reference::device::GemmComplex<
|
||||
ElementA, LayoutA,
|
||||
ElementB, LayoutB,
|
||||
ElementC, LayoutC,
|
||||
ElementCompute, ElementCompute
|
||||
>(
|
||||
problem,
|
||||
options.alpha,
|
||||
view_A,
|
||||
cutlass::ComplexTransform::kNone,
|
||||
view_B,
|
||||
cutlass::ComplexTransform::kNone,
|
||||
options.beta,
|
||||
view_C,
|
||||
view_Ref_device,
|
||||
ElementCompute(0)
|
||||
);
|
||||
|
||||
double norm_reference = cutlass::reference::host::TensorNorm(
|
||||
reference_D.host_view());
|
||||
// Copy reference results to host memory for verification
|
||||
std::vector<ElementD> matrix_D_Ref(layout_C.capacity(extent_C));
|
||||
cutlass::device_memory::copy_to_host(matrix_D_Ref.data(), block_Ref.get(), matrix_D_Ref.size());
|
||||
cutlass::TensorView<ElementD, LayoutC> view_Ref(matrix_D_Ref.data(), layout_C, extent_C);
|
||||
|
||||
double rel_error = norm_diff / norm_reference;
|
||||
std::vector<ElementSoftmax> matrix_Softmax_Ref(layout_C.capacity(extent_C));
|
||||
cutlass::TensorView<ElementSoftmax, LayoutC> view_Softmax_Ref(matrix_Softmax_Ref.data(), layout_C, extent_C);
|
||||
|
||||
if (rel_error > kThreshold) {
|
||||
std::cerr << "\n\nTensor D Relative error: " << rel_error << std::endl;
|
||||
// Copy computed results to host memory
|
||||
std::vector<ElementD> matrix_D(layout_C.capacity(extent_C));
|
||||
cutlass::device_memory::copy_to_host(matrix_D.data(), block_D.get() + total_elements_D_per_batch * batch_idx, matrix_D.size());
|
||||
|
||||
std::vector<ElementD> matrix_Softmax(layout_C.capacity(extent_C));
|
||||
cutlass::device_memory::copy_to_host(matrix_Softmax.data(), block_Softmax.get() + total_elements_D_per_batch * batch_idx, matrix_Softmax.size());
|
||||
|
||||
// Compute the norm
|
||||
for (int m = 0; m < options.problem_size.m(); ++m) {
|
||||
reference_N.at({m, 0}) = view_Ref.ref().at({m, 0});
|
||||
for (int n = 1; n < options.problem_size.n(); ++n) {
|
||||
reference_N.at({m, 0}) = std::max(reference_N.at({m, 0}), ElementNorm(view_Ref.ref().at({m, n})));
|
||||
}
|
||||
}
|
||||
else {
|
||||
verified_D = true;
|
||||
|
||||
// Compute softmax
|
||||
for (int m = 0; m < options.problem_size.m(); ++m) {
|
||||
|
||||
float sum = float();
|
||||
|
||||
for (int n = 0; n < options.problem_size.n(); ++n) {
|
||||
sum += std::exp( float(view_Ref.ref().at({m, n})) - float(reference_N.at({m, 0})) );
|
||||
}
|
||||
|
||||
float inv_sum = float(1.0f / sum);
|
||||
|
||||
for (int n = 0; n < options.problem_size.n(); ++n) {
|
||||
|
||||
view_Softmax_Ref.ref().at({m, n}) = ElementSoftmax(
|
||||
std::exp( float(view_Ref.ref().at({m, n})) - float(reference_N.at({m, 0})) ) * inv_sum
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (!verified_N) {
|
||||
verified_N = verify_tensor_N(tensor_N, reference_N);
|
||||
}
|
||||
// Verification checks - set any of these to 'true' to override the verification checks.
|
||||
bool verified_D = false;
|
||||
bool verified_Softmax = false;
|
||||
|
||||
if (!verified_Softmax) {
|
||||
|
||||
double norm_diff = cutlass::reference::host::TensorNormDiff(
|
||||
tensor_Softmax.host_view(),
|
||||
reference_Softmax.host_view());
|
||||
|
||||
double norm_reference = cutlass::reference::host::TensorNorm(
|
||||
reference_Softmax.host_view());
|
||||
|
||||
double rel_error = norm_diff / norm_reference;
|
||||
|
||||
if (rel_error > kThreshold) {
|
||||
std::cerr << "\n\nSoftmax Relative error: " << rel_error << std::endl;
|
||||
}
|
||||
else {
|
||||
verified_Softmax = true;
|
||||
}
|
||||
}
|
||||
|
||||
if (!verified_D || !verified_N || !verified_Softmax) {
|
||||
|
||||
std::cerr << "Verification check failed for tensor Softmax" << std::endl;
|
||||
|
||||
emit_results();
|
||||
|
||||
// Summarize which checks failed
|
||||
// Verify softmax output
|
||||
if (!verified_D) {
|
||||
std::cerr << "Verification of D tensor failed\n";
|
||||
}
|
||||
|
||||
if (!verified_N) {
|
||||
std::cerr << "Verification of N tensor failed\n";
|
||||
verified_D = verify_tensor<ElementC>(matrix_D, matrix_D_Ref);
|
||||
}
|
||||
|
||||
if (!verified_Softmax) {
|
||||
std::cerr << "Verification of Softmax tensor failed\n";
|
||||
verified_Softmax = verify_tensor<ElementSoftmax>(matrix_Softmax, matrix_Softmax_Ref);
|
||||
}
|
||||
|
||||
if (!verified_D || !verified_Softmax) {
|
||||
|
||||
std::cerr << "Verification check failed for tensor Softmax at batch " << batch_idx << "\n";
|
||||
|
||||
// Summarize which checks failed
|
||||
if (!verified_D) {
|
||||
std::cerr << "Verification of D tensor failed\n";
|
||||
}
|
||||
|
||||
if (!verified_Softmax) {
|
||||
std::cerr << "Verification of Softmax tensor failed\n";
|
||||
}
|
||||
|
||||
return false;
|
||||
}
|
||||
|
||||
return false;
|
||||
}
|
||||
|
||||
return true;
|
||||
@@ -637,14 +652,17 @@ struct Testbed {
|
||||
int64_t flops = int64_t(options.problem_size.m()) * options.problem_size.n() * options.problem_size.k() * 2;
|
||||
int64_t bytes = (sizeof(ElementD) * 2 + sizeof(ElementSoftmax)) * options.problem_size.m() * options.problem_size.n();
|
||||
|
||||
double gflops_per_second = double(flops) * kIterations / double(elapsed_ms / 1000.0f) / double(1.0e9);
|
||||
double gbytes_per_second = double(bytes) * kIterations / double(elapsed_ms / 1000.0f) / double(1 << 30);
|
||||
double gflops_per_second = double(flops) * kIterations * options.batch_count / double(elapsed_ms / 1000.0f) / double(1.0e9);
|
||||
double gbytes_per_second = double(bytes) * kIterations * options.batch_count / double(elapsed_ms / 1000.0f) / double(1 << 30);
|
||||
|
||||
double elapsed_ms_per_iter = double(elapsed_ms) / kIterations;
|
||||
|
||||
std::cout << " Problem: "
|
||||
<< options.problem_size.m() << "-by-" << options.problem_size.n() << "-by-" << options.problem_size.k()
|
||||
<< ", batch size: " << options.batch_count
|
||||
<< std::endl;
|
||||
|
||||
std::cout << " Runtime: " << elapsed_ms << " ms\n" << std::endl;
|
||||
std::cout << " Runtime: " << elapsed_ms_per_iter << " ms\n" << std::endl;
|
||||
|
||||
std::cout << " GFLOPs: " << gflops_per_second << " GFLOPs" << std::endl;
|
||||
std::cout << "Memory bandwidth: " << gbytes_per_second << " GiB/s" << std::endl;
|
||||
|
||||
@@ -29,7 +29,8 @@
|
||||
*
|
||||
**************************************************************************************************/
|
||||
/*! \file
|
||||
\brief GEMM kernel to support the 'epilogue visitor' model for fusion.
|
||||
\brief GEMM kernel to support the epilogue visitor model
|
||||
for customized softmax partial reduction epilogue fusion.
|
||||
|
||||
This source file will likely be moved to `include/cutlass/gemm/kernel/` in the future once
|
||||
its usage has been stabilized. For now, it is included in this example to demonstrate
|
||||
@@ -78,6 +79,7 @@ public:
|
||||
|
||||
using ElementC = typename EpilogueVisitor::ElementOutput;
|
||||
using LayoutC = typename Epilogue::Layout;
|
||||
using TensorRefC = TensorRef<ElementC, LayoutC>;
|
||||
|
||||
static ComplexTransform const kTransformA = Mma::kTransformA;
|
||||
static ComplexTransform const kTransformB = Mma::kTransformB;
|
||||
@@ -89,6 +91,9 @@ public:
|
||||
using InstructionShape = typename Mma::Policy::Operator::InstructionShape;
|
||||
using ArchTag = typename Mma::ArchTag;
|
||||
|
||||
using ElementNorm = typename EpilogueVisitor::ElementNorm;
|
||||
using ElementSum = typename EpilogueVisitor::ElementSum;
|
||||
|
||||
static int const kStages = Mma::kStages;
|
||||
static int const kAlignmentA = Mma::IteratorA::AccessType::kElements;
|
||||
static int const kAlignmentB = Mma::IteratorB::AccessType::kElements;
|
||||
@@ -121,6 +126,11 @@ public:
|
||||
|
||||
TensorRefA ref_A;
|
||||
TensorRefB ref_B;
|
||||
TensorRefC ref_C;
|
||||
TensorRefC ref_D;
|
||||
|
||||
ElementNorm *ptr_Max;
|
||||
ElementSum *ptr_Sum;
|
||||
|
||||
int64_t batch_stride_A;
|
||||
int64_t batch_stride_B;
|
||||
@@ -144,6 +154,10 @@ public:
|
||||
int batch_count_,
|
||||
TensorRefA ref_A_,
|
||||
TensorRefB ref_B_,
|
||||
TensorRefC ref_C_,
|
||||
TensorRefC ref_D_,
|
||||
ElementNorm *ptr_Max_,
|
||||
ElementSum *ptr_Sum_,
|
||||
int64_t batch_stride_A_,
|
||||
int64_t batch_stride_B_,
|
||||
typename EpilogueVisitor::Arguments epilogue_visitor_
|
||||
@@ -153,6 +167,10 @@ public:
|
||||
batch_count(batch_count_),
|
||||
ref_A(ref_A_),
|
||||
ref_B(ref_B_),
|
||||
ref_C(ref_C_),
|
||||
ref_D(ref_D_),
|
||||
ptr_Max(ptr_Max_),
|
||||
ptr_Sum(ptr_Sum_),
|
||||
batch_stride_A(batch_stride_A_),
|
||||
batch_stride_B(batch_stride_B_),
|
||||
epilogue_visitor(epilogue_visitor_)
|
||||
@@ -174,6 +192,8 @@ public:
|
||||
|
||||
typename Mma::IteratorA::Params params_A;
|
||||
typename Mma::IteratorB::Params params_B;
|
||||
typename EpilogueVisitor::OutputTileIterator::Params params_C;
|
||||
typename EpilogueVisitor::OutputTileIterator::Params params_D;
|
||||
|
||||
GemmUniversalMode mode;
|
||||
int batch_count;
|
||||
@@ -181,6 +201,11 @@ public:
|
||||
|
||||
void * ptr_A;
|
||||
void * ptr_B;
|
||||
ElementC * ptr_C;
|
||||
ElementC * ptr_D;
|
||||
|
||||
ElementNorm * ptr_Max;
|
||||
ElementSum * ptr_Sum;
|
||||
|
||||
int64_t batch_stride_A;
|
||||
int64_t batch_stride_B;
|
||||
@@ -196,11 +221,17 @@ public:
|
||||
swizzle_log_tile(0),
|
||||
params_A(0),
|
||||
params_B(0),
|
||||
params_C(0),
|
||||
params_D(0),
|
||||
batch_count(0),
|
||||
gemm_k_size(0),
|
||||
mode(cutlass::gemm::GemmUniversalMode::kGemm),
|
||||
ptr_A(nullptr),
|
||||
ptr_B(nullptr),
|
||||
ptr_C(nullptr),
|
||||
ptr_D(nullptr),
|
||||
ptr_Max(nullptr),
|
||||
ptr_Sum(nullptr),
|
||||
batch_stride_A(0),
|
||||
batch_stride_B(0)
|
||||
{ }
|
||||
@@ -213,11 +244,17 @@ public:
|
||||
swizzle_log_tile(0),
|
||||
params_A(args.ref_A.layout()),
|
||||
params_B(args.ref_B.layout()),
|
||||
params_C(args.ref_C.layout()),
|
||||
params_D(args.ref_D.layout()),
|
||||
mode(args.mode),
|
||||
batch_count(args.batch_count),
|
||||
gemm_k_size(args.problem_size.k()),
|
||||
ptr_A(args.ref_A.data()),
|
||||
ptr_B(args.ref_B.data()),
|
||||
ptr_C(args.ref_C.data()),
|
||||
ptr_D(args.ref_D.data()),
|
||||
ptr_Max(args.ptr_Max),
|
||||
ptr_Sum(args.ptr_Sum),
|
||||
batch_stride_A(args.batch_stride_A),
|
||||
batch_stride_B(args.batch_stride_B),
|
||||
epilogue_visitor(args.epilogue_visitor)
|
||||
@@ -467,7 +504,14 @@ public:
|
||||
thread_idx,
|
||||
warp_idx,
|
||||
lane_idx,
|
||||
threadblock_offset);
|
||||
params.params_C,
|
||||
params.params_D,
|
||||
params.ptr_C,
|
||||
params.ptr_D,
|
||||
params.ptr_Max,
|
||||
params.ptr_Sum,
|
||||
threadblock_offset,
|
||||
blockIdx.y *params.problem_size.m() );
|
||||
|
||||
if (params.mode == GemmUniversalMode::kGemm) {
|
||||
// Indicate which position in a serial reduction the output operator is currently updating
|
||||
|
||||
@@ -49,10 +49,12 @@
|
||||
#include "cutlass/gemm/kernel/default_gemm.h"
|
||||
#include "cutlass/gemm/kernel/default_gemm_complex.h"
|
||||
#include "cutlass/gemm/device/default_gemm_configuration.h"
|
||||
#include "cutlass/epilogue/threadblock/epilogue_visitor_with_softmax.h"
|
||||
#include "cutlass/epilogue/threadblock/epilogue_with_visitor.h"
|
||||
#include "cutlass/reduction/kernel/reduce_softmax_final.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
#include "epilogue_with_visitor.h"
|
||||
#include "gemm_with_epilogue_visitor.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -209,6 +211,9 @@ private:
|
||||
int idx_m = block_m + thread_m;
|
||||
int idx_n = block_n + thread_n;
|
||||
|
||||
int batch_offset_norm = block_batch * params.args.batch_stride_N;
|
||||
int batch_offset_sum = block_batch * params.args.batch_stride_S;
|
||||
|
||||
// Kill off thread if it is outside the row boundary
|
||||
if (params.args.extent.row() <= idx_m) {
|
||||
return;
|
||||
@@ -251,8 +256,8 @@ private:
|
||||
params.args.batch_stride_Soft * block_batch +
|
||||
params.args.ref_Soft.layout()({idx_m, idx_n}));
|
||||
|
||||
ElementSum inv_sum = (params.args.ref_S.data())[block_m];
|
||||
ElementNorm norm = (params.args.ref_N.data())[block_m];
|
||||
ElementSum inv_sum = (params.args.ref_S.data())[block_m + batch_offset_sum];
|
||||
ElementNorm norm = (params.args.ref_N.data())[block_m + batch_offset_norm];
|
||||
|
||||
//
|
||||
// Loop
|
||||
@@ -281,556 +286,6 @@ private:
|
||||
}
|
||||
};
|
||||
|
||||
template <
|
||||
typename ElementNorm_,
|
||||
typename ElementSum_,
|
||||
typename ElementSoftmaxCompute_,
|
||||
typename ThreadblockShape_
|
||||
>
|
||||
class ApplyFinalReduction {
|
||||
public:
|
||||
|
||||
using ElementNorm = ElementNorm_;
|
||||
using ElementSum = ElementSum_;
|
||||
using ElementSoftmaxCompute = ElementSoftmaxCompute_;
|
||||
using ThreadblockShape = ThreadblockShape_;
|
||||
|
||||
using Layout = cutlass::layout::RowMajor;
|
||||
|
||||
using TensorRefN = TensorRef<ElementNorm, Layout>;
|
||||
using TensorRefSum = TensorRef<ElementSum, Layout>;
|
||||
|
||||
//
|
||||
// Arguments
|
||||
//
|
||||
|
||||
struct Arguments {
|
||||
|
||||
MatrixCoord extent; ///< Extent of D and Softmax matrices
|
||||
int batch_count; ///< Batch count
|
||||
TensorRefN ref_N; ///< Norm tensor (input / output)
|
||||
TensorRefSum ref_Sum; ///< Sum tensor (input / output)
|
||||
int64_t batch_stride_N; ///< Batch stride for N tensor
|
||||
int64_t batch_stride_Sum; ///< Batch stride for softmax tensor
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
Arguments():
|
||||
batch_count(1),
|
||||
batch_stride_N(0),
|
||||
batch_stride_Sum(0)
|
||||
{ }
|
||||
|
||||
Arguments(
|
||||
MatrixCoord extent_, ///< Extent of D and Softmax matrices
|
||||
int batch_count_, ///< Batch count
|
||||
TensorRefN ref_N_, ///< Output parameter for N
|
||||
TensorRefSum ref_Sum_ , ///< Sum
|
||||
int64_t batch_stride_N_ = 0,
|
||||
int64_t batch_stride_Sum_ = 0
|
||||
):
|
||||
extent(extent_),
|
||||
batch_count(batch_count_),
|
||||
ref_N(ref_N_),
|
||||
ref_Sum(ref_Sum_),
|
||||
batch_stride_N(batch_stride_N_),
|
||||
batch_stride_Sum(batch_stride_Sum_)
|
||||
{
|
||||
|
||||
}
|
||||
};
|
||||
|
||||
struct SharedStorage {
|
||||
|
||||
|
||||
};
|
||||
|
||||
//
|
||||
// Params struct
|
||||
//
|
||||
|
||||
struct Params {
|
||||
Arguments args;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
Params() { }
|
||||
|
||||
Params(Arguments const &args_): args(args_) { }
|
||||
};
|
||||
|
||||
private:
|
||||
|
||||
public:
|
||||
|
||||
CUTLASS_DEVICE
|
||||
ApplyFinalReduction() { }
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void operator()(Params const ¶ms, SharedStorage &shared_storage) {
|
||||
|
||||
apply(params, shared_storage);
|
||||
}
|
||||
|
||||
private:
|
||||
|
||||
/// Partial reduction
|
||||
CUTLASS_DEVICE
|
||||
void apply(Params const ¶ms, SharedStorage &shared_storage) {
|
||||
|
||||
int threadblock_num = (params.args.extent.column() + ThreadblockShape::kN - 1) / ThreadblockShape::kN;
|
||||
|
||||
int block_batch = blockIdx.z;
|
||||
|
||||
int block_n = blockIdx.x * blockDim.x;
|
||||
|
||||
int thread_n = threadIdx.x;
|
||||
|
||||
int idx_n = block_n + thread_n;
|
||||
|
||||
if (idx_n >= params.args.extent.row()) {
|
||||
return;
|
||||
}
|
||||
|
||||
|
||||
using ConvertSumOutput = cutlass::NumericConverter<ElementSum, ElementSoftmaxCompute>;
|
||||
using ConvertNormOutput = cutlass::NumericConverter<ElementNorm, ElementSoftmaxCompute>;
|
||||
|
||||
using ConvertSum = cutlass::NumericConverter<ElementSoftmaxCompute, ElementSum>;
|
||||
using ConvertNorm = cutlass::NumericConverter<ElementSoftmaxCompute, ElementNorm>;
|
||||
|
||||
ConvertSum convert_sum;
|
||||
ConvertNorm convert_norm;
|
||||
|
||||
ConvertSumOutput convert_sum_output;
|
||||
ConvertNormOutput convert_norm_output;
|
||||
|
||||
ElementNorm *access_n = params.args.ref_N.data() + params.args.batch_stride_N * block_batch + idx_n;
|
||||
ElementSum *access_s = params.args.ref_Sum.data() + params.args.batch_stride_Sum * block_batch + idx_n;
|
||||
|
||||
ElementNorm *access_n_bak = access_n;
|
||||
ElementSum *access_s_bak = access_s;
|
||||
|
||||
uint32_t float_max_bits = 0xff7fffff;
|
||||
float min_float = reinterpret_cast<float const &>(float_max_bits);
|
||||
|
||||
ElementSoftmaxCompute max_val = ElementSoftmaxCompute(min_float);
|
||||
ElementSoftmaxCompute sum_val = ElementSoftmaxCompute(0);
|
||||
ElementNorm fetch_n;
|
||||
ElementSum fetch_s;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int idx_m = 0; idx_m < threadblock_num; idx_m++) {
|
||||
arch::global_load<ElementNorm, sizeof(ElementNorm)>(fetch_n, access_n, true);
|
||||
max_val = fast_max(max_val, convert_norm(fetch_n));
|
||||
access_n += params.args.extent.row();
|
||||
}
|
||||
|
||||
access_n = access_n_bak;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int idx_m = 0; idx_m < threadblock_num; idx_m++) {
|
||||
arch::global_load<ElementNorm, sizeof(ElementNorm)>(fetch_n, access_n, true);
|
||||
arch::global_load<ElementSum, sizeof(ElementSum)>(fetch_s, access_s, true);
|
||||
sum_val += convert_sum(fetch_s) * fast_exp(convert_norm(fetch_n) - max_val);
|
||||
access_n += params.args.extent.row();
|
||||
access_s += params.args.extent.row();
|
||||
}
|
||||
|
||||
ElementSoftmaxCompute inv_sum = cutlass::constants::one<ElementSoftmaxCompute>() / sum_val;
|
||||
|
||||
access_n = access_n_bak;
|
||||
access_s = access_s_bak;
|
||||
|
||||
access_n[0] = convert_norm_output(max_val);
|
||||
access_s[0] = convert_sum_output(inv_sum);
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename ThreadblockShape_,
|
||||
int ThreadCount,
|
||||
typename OutputTileIterator_,
|
||||
typename ElementAccumulator_,
|
||||
typename ElementNorm_,
|
||||
typename ElementSum_,
|
||||
typename ElementSoftmaxCompute_,
|
||||
typename ElementwiseFunctor_
|
||||
>
|
||||
class EpilogueVisitorBiasMax {
|
||||
public:
|
||||
|
||||
using ThreadblockShape = ThreadblockShape_;
|
||||
static int const kThreadCount = ThreadCount;
|
||||
|
||||
using OutputTileIterator = OutputTileIterator_;
|
||||
using ElementwiseFunctor = ElementwiseFunctor_;
|
||||
|
||||
static int const kIterations = OutputTileIterator::kIterations;
|
||||
static int const kElementsPerAccess = OutputTileIterator::kElementsPerAccess;
|
||||
|
||||
using ElementOutput = typename OutputTileIterator::Element;
|
||||
using LayoutOutput = cutlass::layout::RowMajor;
|
||||
using ElementAccumulator = ElementAccumulator_;
|
||||
|
||||
using ElementNorm = ElementNorm_;
|
||||
using ElementSum = ElementSum_;
|
||||
using ElementSoftmaxCompute = ElementSoftmaxCompute_;
|
||||
|
||||
using AccumulatorFragment = Array<ElementAccumulator, kElementsPerAccess>;
|
||||
using SoftmaxFragment = Array<ElementSoftmaxCompute, kElementsPerAccess>;
|
||||
using OutputVector = Array<ElementOutput, kElementsPerAccess>;
|
||||
using TensorRefD = TensorRef<ElementOutput, LayoutOutput>;
|
||||
|
||||
/// Argument structure
|
||||
struct Arguments {
|
||||
|
||||
typename ElementwiseFunctor::Params elementwise;
|
||||
TensorRefD ref_C;
|
||||
TensorRefD ref_D;
|
||||
ElementNorm *ptr_Max;
|
||||
ElementSum *ptr_Sum;
|
||||
int64_t batch_stride_C;
|
||||
int64_t batch_stride_D;
|
||||
int64_t batch_stride_Max;
|
||||
int64_t batch_stride_Sum;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
Arguments():
|
||||
ptr_Max(nullptr),
|
||||
ptr_Sum(nullptr),
|
||||
batch_stride_C(0),
|
||||
batch_stride_D(0),
|
||||
batch_stride_Max(0),
|
||||
batch_stride_Sum(0)
|
||||
{
|
||||
|
||||
}
|
||||
|
||||
Arguments(
|
||||
typename ElementwiseFunctor::Params elementwise_,
|
||||
TensorRefD ref_C_,
|
||||
TensorRefD ref_D_,
|
||||
ElementNorm *ptr_Max_,
|
||||
ElementSum *ptr_Sum_,
|
||||
int64_t batch_stride_C_,
|
||||
int64_t batch_stride_D_,
|
||||
int64_t batch_stride_Max_,
|
||||
int64_t batch_stride_Sum_
|
||||
):
|
||||
elementwise(elementwise_),
|
||||
ref_C(ref_C_),
|
||||
ref_D(ref_D_),
|
||||
ptr_Max(ptr_Max_),
|
||||
ptr_Sum(ptr_Sum_),
|
||||
batch_stride_C(batch_stride_C_),
|
||||
batch_stride_D(batch_stride_D_),
|
||||
batch_stride_Max(batch_stride_Max_),
|
||||
batch_stride_Sum(batch_stride_Sum_)
|
||||
{
|
||||
|
||||
}
|
||||
};
|
||||
|
||||
struct Params {
|
||||
|
||||
typename ElementwiseFunctor::Params elementwise;
|
||||
typename OutputTileIterator::Params params_C;
|
||||
typename OutputTileIterator::Params params_D;
|
||||
typename OutputTileIterator::Element *ptr_C;
|
||||
typename OutputTileIterator::Element *ptr_D;
|
||||
ElementNorm *ptr_Max;
|
||||
ElementSum *ptr_Sum;
|
||||
int64_t batch_stride_C;
|
||||
int64_t batch_stride_D;
|
||||
int64_t batch_stride_Max;
|
||||
int64_t batch_stride_Sum;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params():
|
||||
ptr_D(nullptr),
|
||||
ptr_Max(nullptr),
|
||||
ptr_Sum(nullptr)
|
||||
{
|
||||
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(Arguments const &args):
|
||||
elementwise(args.elementwise),
|
||||
params_C(args.ref_C.layout()),
|
||||
params_D(args.ref_D.layout()),
|
||||
ptr_C(args.ref_C.data()),
|
||||
ptr_D(args.ref_D.data()),
|
||||
ptr_Max(args.ptr_Max),
|
||||
ptr_Sum(args.ptr_Sum),
|
||||
batch_stride_C(args.batch_stride_C),
|
||||
batch_stride_D(args.batch_stride_D),
|
||||
batch_stride_Max(args.batch_stride_Max),
|
||||
batch_stride_Sum(args.batch_stride_Sum)
|
||||
{
|
||||
|
||||
}
|
||||
};
|
||||
|
||||
/// Shared storage
|
||||
struct SharedStorage {
|
||||
|
||||
};
|
||||
|
||||
private:
|
||||
|
||||
Params const & params_;
|
||||
SharedStorage & shared_storage_;
|
||||
MatrixCoord extent_;
|
||||
ElementwiseFunctor elementwise_;
|
||||
|
||||
OutputTileIterator iterator_C_;
|
||||
OutputTileIterator iterator_D_;
|
||||
typename OutputTileIterator::Fragment fragment_C_;
|
||||
typename OutputTileIterator::Fragment fragment_D_;
|
||||
|
||||
ElementAccumulator alpha_;
|
||||
ElementAccumulator beta_;
|
||||
|
||||
ElementSoftmaxCompute accum_max_;
|
||||
int threadblock_row_;
|
||||
|
||||
public:
|
||||
|
||||
CUTLASS_DEVICE
|
||||
EpilogueVisitorBiasMax(
|
||||
Params const ¶ms, ///< Parameters routed to the epilogue
|
||||
SharedStorage &shared_storage, ///< Shared storage needed by the functors here
|
||||
MatrixCoord const &problem_size, ///< Problem size of the output
|
||||
int thread_idx, ///< Thread index within the threadblock
|
||||
int warp_idx, ///< Warp index within the threadblock
|
||||
int lane_idx, ///< Lane index within the warp
|
||||
MatrixCoord const &threadblock_offset = MatrixCoord(0, 0)
|
||||
):
|
||||
params_(params),
|
||||
shared_storage_(shared_storage),
|
||||
extent_(problem_size),
|
||||
elementwise_(params.elementwise),
|
||||
iterator_C_(params.params_C, params.ptr_C, problem_size, thread_idx, threadblock_offset),
|
||||
iterator_D_(params.params_D, params.ptr_D, problem_size, thread_idx, threadblock_offset),
|
||||
threadblock_row_(threadblock_offset.row())
|
||||
{
|
||||
alpha_ = (params.elementwise.alpha_ptr ? *params.elementwise.alpha_ptr : params.elementwise.alpha);
|
||||
beta_ = (params.elementwise.beta_ptr ? *params.elementwise.beta_ptr : params.elementwise.beta);
|
||||
|
||||
if (beta_ == ElementAccumulator()) {
|
||||
iterator_C_.clear_mask();
|
||||
}
|
||||
}
|
||||
|
||||
/// Helper to indicate split-K behavior
|
||||
CUTLASS_DEVICE
|
||||
void set_k_partition(
|
||||
int split_k_index, ///< Index of this threadblock within split-K partitioned scheme
|
||||
int split_k_slices) { ///< Total number of split-K slices
|
||||
|
||||
}
|
||||
|
||||
/// Called to set the batch index
|
||||
CUTLASS_DEVICE
|
||||
void set_batch_index(int batch_idx) {
|
||||
iterator_C_.add_pointer_offset(batch_idx * params_.batch_stride_C);
|
||||
iterator_D_.add_pointer_offset(batch_idx * params_.batch_stride_D);
|
||||
}
|
||||
|
||||
/// Called at the start of the epilogue just before iterating over accumulator slices
|
||||
CUTLASS_DEVICE
|
||||
void begin_epilogue() {
|
||||
|
||||
}
|
||||
|
||||
/// Called at the start of one step before starting accumulator exchange
|
||||
CUTLASS_DEVICE
|
||||
void begin_step(int step_idx) {
|
||||
fragment_D_.clear();
|
||||
fragment_C_.clear();
|
||||
|
||||
if (elementwise_.kScale != cutlass::epilogue::thread::ScaleType::OnlyAlphaScaling) {
|
||||
iterator_C_.load(fragment_C_);
|
||||
++iterator_C_;
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
/// Called at the start of a row
|
||||
CUTLASS_DEVICE
|
||||
void begin_row(int row_idx) {
|
||||
|
||||
}
|
||||
|
||||
/// Called after accumulators have been exchanged for each accumulator vector
|
||||
CUTLASS_DEVICE
|
||||
void visit(
|
||||
int row_idx,
|
||||
int column_idx,
|
||||
int frag_idx,
|
||||
AccumulatorFragment const &accum) {
|
||||
|
||||
using Mul = cutlass::multiplies<SoftmaxFragment>;
|
||||
using Minus = cutlass::minus<SoftmaxFragment>;
|
||||
using Exp = cutlass::fast_exp_op<SoftmaxFragment>;
|
||||
|
||||
Minus minus;
|
||||
Exp exponential;
|
||||
|
||||
SoftmaxFragment result;
|
||||
|
||||
using ConvertSumOutput = cutlass::NumericConverter<ElementSoftmaxCompute, ElementSum>;
|
||||
using ConvertNormOutput = cutlass::NumericConverter<ElementSoftmaxCompute, ElementNorm>;
|
||||
|
||||
ConvertSumOutput convert_sum_output;
|
||||
ConvertNormOutput convert_norm_output;
|
||||
|
||||
NumericArrayConverter<ElementSoftmaxCompute, ElementOutput, kElementsPerAccess> source_converter;
|
||||
OutputVector &source_vector = reinterpret_cast<OutputVector *>(&fragment_C_)[frag_idx];
|
||||
|
||||
if (elementwise_.kScale == cutlass::epilogue::thread::ScaleType::OnlyAlphaScaling) {
|
||||
result = source_converter(elementwise_(accum));
|
||||
}else{
|
||||
result = source_converter(elementwise_(accum, source_vector));
|
||||
}
|
||||
|
||||
MatrixCoord thread_offset =
|
||||
iterator_D_.thread_start() +
|
||||
OutputTileIterator::ThreadMap::iteration_offset(frag_idx);
|
||||
|
||||
int thread_in_row = OutputTileIterator::ThreadMap::Detail::RowArrangement::Detail::kShapeWidth;
|
||||
int half_thread_in_row = (thread_in_row >> 1);
|
||||
|
||||
bool column_guard = (thread_offset.column() < extent_.column());
|
||||
|
||||
// Compute the maximum within one row
|
||||
if (!column_idx) {
|
||||
// This is the first fragment in a new row
|
||||
if (column_guard) {
|
||||
accum_max_ = maximum_accumulator_(result);
|
||||
}
|
||||
}
|
||||
else {
|
||||
// This is an additional fragment in the same row
|
||||
if (column_guard) {
|
||||
accum_max_ = maximum_accumulator_(result, accum_max_);
|
||||
}
|
||||
}
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = half_thread_in_row; i > 0; i >>= 1) {
|
||||
ElementSoftmaxCompute tmp = __shfl_xor_sync(0xFFFFFFFF, accum_max_, i);
|
||||
accum_max_ = fast_max(accum_max_, tmp);
|
||||
}
|
||||
|
||||
SoftmaxFragment sum_frag = exponential(minus(result, accum_max_));
|
||||
|
||||
ElementSoftmaxCompute reduction_sum = sum_accumulator_(sum_frag);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = half_thread_in_row; i > 0; i >>= 1) {
|
||||
ElementSoftmaxCompute tmp = __shfl_xor_sync(0xFFFFFFFF, reduction_sum, i);
|
||||
reduction_sum += tmp;
|
||||
}
|
||||
|
||||
bool is_write_thread = (thread_offset.row() < extent_.row() && (threadIdx.x % thread_in_row) == 0);
|
||||
ElementNorm *curr_ptr_max = params_.ptr_Max + thread_offset.row() + blockIdx.y * extent_.row();
|
||||
ElementSum *curr_ptr_sum = params_.ptr_Sum + thread_offset.row() + blockIdx.y * extent_.row();
|
||||
|
||||
arch::global_store<ElementNorm, sizeof(ElementNorm)>(
|
||||
convert_norm_output(accum_max_),
|
||||
(void *)curr_ptr_max,
|
||||
is_write_thread);
|
||||
|
||||
arch::global_store<ElementSum, sizeof(ElementSum)>(
|
||||
convert_sum_output(reduction_sum),
|
||||
(void *)curr_ptr_sum,
|
||||
is_write_thread);
|
||||
|
||||
clear_accum_max_();
|
||||
|
||||
// Convert to the output
|
||||
NumericArrayConverter<ElementOutput, ElementSoftmaxCompute, kElementsPerAccess> output_converter;
|
||||
OutputVector &output = reinterpret_cast<OutputVector *>(&fragment_D_)[frag_idx];
|
||||
output = output_converter(result);
|
||||
}
|
||||
|
||||
/// Called at the start of a row
|
||||
CUTLASS_DEVICE
|
||||
void end_row(int row_idx) {
|
||||
|
||||
}
|
||||
|
||||
/// Called after all accumulator elements have been visited
|
||||
CUTLASS_DEVICE
|
||||
void end_step(int step_idx) {
|
||||
|
||||
iterator_D_.store(fragment_D_);
|
||||
++iterator_D_;
|
||||
}
|
||||
|
||||
/// Called after all steps have been completed
|
||||
CUTLASS_DEVICE
|
||||
void end_epilogue() {
|
||||
|
||||
}
|
||||
|
||||
private:
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void clear_accum_max_() {
|
||||
|
||||
uint32_t float_max_bits = 0xff7fffff; // -FLT_MAX
|
||||
float min_float = reinterpret_cast<float const &>(float_max_bits);
|
||||
accum_max_ = ElementSoftmaxCompute(min_float);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
ElementSoftmaxCompute sum_accumulator_(SoftmaxFragment const &accum) {
|
||||
ElementSoftmaxCompute sum_ = ElementSoftmaxCompute(0);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < SoftmaxFragment::kElements; ++i) {
|
||||
sum_ += ElementSoftmaxCompute(accum[i]);
|
||||
}
|
||||
|
||||
return sum_;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
ElementSoftmaxCompute maximum_accumulator_(SoftmaxFragment const &accum) {
|
||||
ElementSoftmaxCompute max_ = accum[0];
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 1; i < SoftmaxFragment::kElements; ++i) {
|
||||
max_ = fast_max(max_, ElementSoftmaxCompute(accum[i]));
|
||||
}
|
||||
|
||||
return max_;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
ElementSoftmaxCompute maximum_accumulator_(SoftmaxFragment const &accum, ElementSoftmaxCompute max_) {
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < SoftmaxFragment::kElements; ++i) {
|
||||
max_ = fast_max(max_, ElementSoftmaxCompute(accum[i]));
|
||||
}
|
||||
|
||||
return max_;
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -846,10 +301,18 @@ template <
|
||||
typename LayoutB_,
|
||||
typename ElementC_,
|
||||
typename ElementCompute_,
|
||||
typename OperatorClass_,
|
||||
typename ArchTag_,
|
||||
typename ThreadblockShape_,
|
||||
typename WarpShape_,
|
||||
typename InstructionShape_,
|
||||
typename EpilogueFunctorOp_,
|
||||
int kStages_,
|
||||
int AlignmentA_ = 128 / cutlass::sizeof_bits<ElementA_>::value,
|
||||
int AlignmentB_ = 128 / cutlass::sizeof_bits<ElementB_>::value,
|
||||
int AlignmentSoftmax_ = 128 / cutlass::sizeof_bits<ElementC_>::value,
|
||||
typename ElementNorm_ = float,
|
||||
typename ElementSum_ = float,
|
||||
int Alignment = 128 / cutlass::sizeof_bits<ElementA_>::value,
|
||||
typename ElementSoftmax_ = ElementC_
|
||||
>
|
||||
class GemmSoftmax {
|
||||
@@ -872,8 +335,6 @@ public:
|
||||
using LayoutA = LayoutA_;
|
||||
using LayoutB = LayoutB_;
|
||||
|
||||
static int const kAlignment = Alignment;
|
||||
|
||||
using EpilogueFunctorOp = EpilogueFunctorOp_;
|
||||
using ElementNorm = ElementNorm_;
|
||||
|
||||
@@ -890,13 +351,17 @@ public:
|
||||
using TensorRefSum = TensorRef<ElementSum, LayoutS>;
|
||||
using TensorRefSoft = TensorRef<ElementSoft, LayoutSoft>;
|
||||
|
||||
using ThreadblockShape = cutlass::gemm::GemmShape<128, 128, 32>;
|
||||
using WarpShape = cutlass::gemm::GemmShape<64, 64, 32>;
|
||||
using InstructionShape = cutlass::gemm::GemmShape<16, 8, 16>;
|
||||
using ThreadblockShape = ThreadblockShape_;
|
||||
using WarpShape = WarpShape_;
|
||||
using InstructionShape = InstructionShape_;
|
||||
|
||||
using OperatorClass = cutlass::arch::OpClassTensorOp;
|
||||
using ArchTag = cutlass::arch::Sm80;
|
||||
static int const kStages = 3;
|
||||
using OperatorClass = OperatorClass_;
|
||||
using ArchTag = ArchTag_;
|
||||
|
||||
static int const kStages = kStages_;
|
||||
static int const AlignmentA = AlignmentA_;
|
||||
static int const AlignmentB = AlignmentB_;
|
||||
static int const AlignmentSoftmax = AlignmentSoftmax_;
|
||||
|
||||
using ThreadblockSwizzle = cutlass::gemm::threadblock::GemmBatchedIdentityThreadblockSwizzle;
|
||||
|
||||
@@ -906,10 +371,10 @@ public:
|
||||
using DefaultGemmKernel = typename cutlass::gemm::kernel::DefaultGemm<
|
||||
ElementA,
|
||||
LayoutA,
|
||||
kAlignment,
|
||||
AlignmentA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
kAlignment,
|
||||
AlignmentB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
ElementCompute,
|
||||
@@ -930,7 +395,7 @@ public:
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// Epilogue visitor
|
||||
using EpilogueVisitor = kernel::EpilogueVisitorBiasMax<
|
||||
using EpilogueVisitor = typename cutlass::epilogue::threadblock::EpilogueVisitorSoftmax<
|
||||
ThreadblockShape,
|
||||
DefaultGemmKernel::kThreadCount,
|
||||
typename DefaultGemmKernel::Epilogue::OutputTileIterator,
|
||||
@@ -961,13 +426,13 @@ public:
|
||||
ElementSum,
|
||||
ElementSoft,
|
||||
ElementSoftmaxCompute,
|
||||
kAlignment,
|
||||
AlignmentSoftmax,
|
||||
MatrixShape<
|
||||
1, 1024
|
||||
>
|
||||
>;
|
||||
|
||||
using ApplyFinalReductionKernel = kernel::ApplyFinalReduction<
|
||||
using ApplyFinalReductionKernel = cutlass::reduction::kernel::ApplySoftmaxFinalReduction<
|
||||
ElementNorm,
|
||||
ElementSum,
|
||||
ElementSoftmaxCompute,
|
||||
@@ -983,6 +448,7 @@ public:
|
||||
typename SoftmaxApplyKernel::Arguments softmax;
|
||||
typename ApplyFinalReductionKernel::Arguments reduction;
|
||||
cutlass::gemm::GemmCoord extend;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
@@ -1013,14 +479,14 @@ public:
|
||||
batch_count_,
|
||||
ref_A_,
|
||||
ref_B_,
|
||||
ref_C_,
|
||||
ref_D_,
|
||||
ref_N_.data(),
|
||||
ref_S_.data(),
|
||||
batch_stride_A_,
|
||||
batch_stride_B_,
|
||||
typename EpilogueVisitor::Arguments(
|
||||
linear_scaling,
|
||||
ref_C_,
|
||||
ref_D_,
|
||||
ref_N_.data(),
|
||||
ref_S_.data(),
|
||||
batch_stride_C_,
|
||||
batch_stride_D_,
|
||||
batch_stride_Max_,
|
||||
@@ -1028,10 +494,9 @@ public:
|
||||
)
|
||||
),
|
||||
reduction(
|
||||
MatrixCoord(problem_size.m(), problem_size.n()),
|
||||
batch_count_,
|
||||
ref_N_,
|
||||
ref_S_,
|
||||
problem_size,
|
||||
ref_N_.data(),
|
||||
ref_S_.data(),
|
||||
batch_stride_Max_,
|
||||
batch_stride_Sum_
|
||||
),
|
||||
@@ -1127,28 +592,24 @@ public:
|
||||
// Launch the ApplyFinalReductionKernel
|
||||
//
|
||||
|
||||
int threadblock_num_in_column = (params_.extend.column() + ThreadblockShape::kN - 1) / ThreadblockShape::kN;
|
||||
int thread_per_block = 128;
|
||||
int block_per_row = (params_.extend.row() + thread_per_block - 1) / thread_per_block;
|
||||
if (block_per_row < 4) {
|
||||
thread_per_block = 32;
|
||||
block_per_row = (params_.extend.row() + thread_per_block - 1) / thread_per_block;
|
||||
}
|
||||
|
||||
if (threadblock_num_in_column > 1) {
|
||||
int thread_per_block = 128;
|
||||
int block_per_row = (params_.extend.row() + thread_per_block - 1) / thread_per_block;
|
||||
if (block_per_row < 4) {
|
||||
thread_per_block = 32;
|
||||
block_per_row = (params_.extend.row() + thread_per_block - 1) / thread_per_block;
|
||||
}
|
||||
dim3 final_reduction_grid(block_per_row, 1, params_.softmax.args.batch_count);
|
||||
dim3 final_reduction_block(thread_per_block);
|
||||
|
||||
dim3 final_reduction_grid(block_per_row);
|
||||
dim3 final_reduction_block(thread_per_block);
|
||||
Kernel<ApplyFinalReductionKernel><<<
|
||||
final_reduction_grid, final_reduction_block, sizeof(typename ApplyFinalReductionKernel::SharedStorage), stream
|
||||
>>>(params_.reduction);
|
||||
|
||||
Kernel<ApplyFinalReductionKernel><<<
|
||||
final_reduction_grid, final_reduction_block, sizeof(typename ApplyFinalReductionKernel::SharedStorage), stream
|
||||
>>>(params_.reduction);
|
||||
result = cudaGetLastError();
|
||||
|
||||
result = cudaGetLastError();
|
||||
|
||||
if (result != cudaSuccess) {
|
||||
return cutlass::Status::kErrorInternal;
|
||||
}
|
||||
if (result != cudaSuccess) {
|
||||
return cutlass::Status::kErrorInternal;
|
||||
}
|
||||
|
||||
//
|
||||
|
||||
@@ -40,18 +40,17 @@
|
||||
// for (int j = 0; j < options.index_size; ++j) {
|
||||
// int b_c_d_col = tensor_indices.at({j, 0});
|
||||
//
|
||||
// for (int k = 0; k < problem_size.k(); ++k) {
|
||||
// for (int k = 0; k < options.index_size; ++k) {
|
||||
// tensor_d_ref.at({i, b_c_d_col}) +=
|
||||
// alpha * tensor_a.at({i, k}) * tensor_b.at({k, b_c_d_col});
|
||||
// }
|
||||
// }
|
||||
// }
|
||||
//
|
||||
// Note that the index vector contains unique random integers with max to be N - 1
|
||||
//
|
||||
// The gather/scatter operation works best when we can still keep the biggest
|
||||
// alignment. For example, when the matrix is row major, we select rows. When
|
||||
// the matrix is column major, we selct columns.
|
||||
// the matrix is column major, we select columns.
|
||||
//
|
||||
// Not all the combination of gather and scatter are legal. For example, if A is
|
||||
// row major and C/D is column major, we cannot gather A and scatter C/D at the
|
||||
@@ -257,7 +256,7 @@ using Gemm = cutlass::gemm::device::GemmUniversal<ElementInputA,
|
||||
cutlass::arch::OpMultiplyAdd,
|
||||
cutlass::ComplexTransform::kNone,
|
||||
cutlass::ComplexTransform::kNone,
|
||||
false, /*GatherA*/
|
||||
false, /*GatherA*/
|
||||
true, /*GatherB*/
|
||||
true /*ScatterD*/
|
||||
>;
|
||||
@@ -353,7 +352,7 @@ int run(Options &options) {
|
||||
tensor_b.layout().stride(),
|
||||
tensor_c.layout().stride(),
|
||||
tensor_d_scattered.layout().stride(),
|
||||
nullptr, // <- pointer to index vector to gather A on device
|
||||
nullptr, // <- pointer to index vector to gather A on device
|
||||
tensor_indices.device_data(), // <- pointer to index vector to gather B on device
|
||||
tensor_indices.device_data()}; // <- pointer to index vector to scatter D on device
|
||||
|
||||
@@ -392,7 +391,7 @@ int run(Options &options) {
|
||||
tensor_d_ref.at({i, b_c_d_col}) +=
|
||||
alpha * tensor_a.at({i, k}) * tensor_b.at({k, b_c_d_col});
|
||||
}
|
||||
|
||||
|
||||
tensor_d_ref.at({i, b_c_d_col}) += (beta * tensor_c.at({i, b_c_d_col}));
|
||||
}
|
||||
}
|
||||
@@ -515,7 +514,7 @@ int main(int argc, const char ** argv) {
|
||||
cudaDeviceProp props;
|
||||
CUDA_CHECK(cudaGetDeviceProperties(&props, 0));
|
||||
|
||||
if (!(props.major > 8 || (props.major == 8 && props.minor >= 0))) {
|
||||
if (!(props.major >= 8)) {
|
||||
std::cerr << "Ampere Tensor Ops must be run on a machine with compute capability at least 80."
|
||||
<< std::endl;
|
||||
notSupported = true;
|
||||
|
||||
36
examples/37_gemm_layernorm_gemm_fusion/CMakeLists.txt
Normal file
36
examples/37_gemm_layernorm_gemm_fusion/CMakeLists.txt
Normal file
@@ -0,0 +1,36 @@
|
||||
|
||||
# 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.
|
||||
|
||||
|
||||
|
||||
cutlass_example_add_executable(
|
||||
37_gemm_layernorm_gemm_fusion
|
||||
gemm_layernorm.cu
|
||||
)
|
||||
|
||||
937
examples/37_gemm_layernorm_gemm_fusion/gemm_layernorm.cu
Normal file
937
examples/37_gemm_layernorm_gemm_fusion/gemm_layernorm.cu
Normal file
@@ -0,0 +1,937 @@
|
||||
/***************************************************************************************************
|
||||
* 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 CUTLASS Layernorm Example.
|
||||
|
||||
This workload provides a layer normalization example using a one-pass, square-sum-based
|
||||
variance calculation. Specifically, we fuse the reduction operation to find
|
||||
local mean and local square sum mean in the epilogue of 1st GEMM. After a light
|
||||
full reduction kernel, the mean / variance values are readily calculated for element-wise
|
||||
operations which are fused into the 2nd GEMM.
|
||||
|
||||
As stated in https://en.wikipedia.org/wiki/Algorithms_for_calculating_variance#Computing_shifted_data,
|
||||
the square-sum based one-pass implementation may raise concerns on numerical stability issues.
|
||||
That being said, though this fully fused layernorm example almost perfectly hides all the memory cost to
|
||||
access the intermediate matrix for layernorm computation, the numerical issue might hinder a persuasive
|
||||
usage in real-world scenarios. If that is the case, a user may turn to the stand-alone CUTLASS layernorm
|
||||
example in tools/util/include/cutlass/util/device_layernorm.h
|
||||
|
||||
Examples:
|
||||
|
||||
# Run a CUTLASS layernorm example with default setup ,
|
||||
# using the language of the transformer model as an example,
|
||||
(Column Major output matrix, hidden dimension = 768, valid word number = 4096, intermediate_scale = 4)
|
||||
$ ./examples/37_gemm_layernorm_gemm_fusion/37_gemm_layernorm_gemm_fusion
|
||||
|
||||
# Run an attention example with hidden dimension = 512
|
||||
$ ./examples/37_gemm_layernorm_gemm_fusion/37_gemm_layernorm_gemm_fusion --hidden_dim=512
|
||||
|
||||
*/
|
||||
|
||||
#include <cmath>
|
||||
#include <iostream>
|
||||
#include <vector>
|
||||
#include <limits>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/arch/memory.h"
|
||||
#include "cutlass/arch/memory_sm75.h"
|
||||
#include "cutlass/gemm/device/gemm_complex.h"
|
||||
#include "cutlass/epilogue/thread/scale_type.h"
|
||||
|
||||
#include "cutlass/util/command_line.h"
|
||||
#include "cutlass/util/host_tensor.h"
|
||||
#include "cutlass/util/reference/device/gemm.h"
|
||||
#include "cutlass/util/reference/host/gemm_complex.h"
|
||||
#include "cutlass/util/reference/host/tensor_reduce.h"
|
||||
#include "cutlass/util/reference/host/tensor_compare.h"
|
||||
#include "cutlass/util/reference/host/tensor_norm.h"
|
||||
#include "cutlass/util/reference/host/tensor_copy.h"
|
||||
#include "cutlass/util/reference/host/tensor_fill.h"
|
||||
#include "cutlass/util/reference/host/error_metrics.h"
|
||||
#include "cutlass/util/tensor_view_io.h"
|
||||
#include "cutlass/fast_math.h"
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
#include "gemm_with_layernorm.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
enum class Disposition {
|
||||
kPassed,
|
||||
kIncorrect,
|
||||
kNotVerified
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// Command line options parsing
|
||||
template<typename LayoutOutput_>
|
||||
struct Options {
|
||||
|
||||
using LayoutOutput = LayoutOutput_;
|
||||
|
||||
static bool const kIsColumnMajorOutput = cutlass::platform::is_same<LayoutOutput, cutlass::layout::ColumnMajor>::value;
|
||||
|
||||
bool help;
|
||||
cutlass::gemm::GemmCoord problem_size0;
|
||||
cutlass::gemm::GemmCoord problem_size1;
|
||||
int hidden_dim;
|
||||
int valid_word_num;
|
||||
int intermediate_scale;
|
||||
int iterations;
|
||||
unsigned seed;
|
||||
float alpha;
|
||||
float beta;
|
||||
bool verification_enabled;
|
||||
double tolerance;
|
||||
|
||||
Options():
|
||||
help(false),
|
||||
iterations(20),
|
||||
seed(2022),
|
||||
hidden_dim(768),
|
||||
valid_word_num(4096),
|
||||
intermediate_scale(4),
|
||||
alpha(1),
|
||||
beta(0),
|
||||
verification_enabled(true),
|
||||
tolerance(0.01),
|
||||
problem_size1(problem_size0.m() * 4, problem_size0.n(), problem_size0.m())
|
||||
{ }
|
||||
|
||||
bool valid() {
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
// Parses the command line
|
||||
void parse(int argc, char const **args) {
|
||||
cutlass::CommandLine cmd(argc, args);
|
||||
|
||||
if (cmd.check_cmd_line_flag("help")) {
|
||||
help = true;
|
||||
}
|
||||
|
||||
cmd.get_cmd_line_argument("hidden_dim", hidden_dim, 768);
|
||||
cmd.get_cmd_line_argument("valid_word_num", valid_word_num, 4096);
|
||||
cmd.get_cmd_line_argument("iterations", iterations);
|
||||
cmd.get_cmd_line_argument("verify", verification_enabled);
|
||||
cmd.get_cmd_line_argument("seed", seed);
|
||||
cmd.get_cmd_line_argument("tolerance", tolerance);
|
||||
|
||||
if (kIsColumnMajorOutput) {
|
||||
// column major output setup
|
||||
problem_size0.m() = hidden_dim;
|
||||
problem_size0.n() = valid_word_num;
|
||||
problem_size0.k() = hidden_dim;
|
||||
|
||||
problem_size1.m() = hidden_dim * intermediate_scale;
|
||||
problem_size1.n() = valid_word_num;
|
||||
problem_size1.k() = hidden_dim;
|
||||
}else{
|
||||
// row major output setup
|
||||
problem_size0.m() = valid_word_num;
|
||||
problem_size0.n() = hidden_dim;
|
||||
problem_size0.k() = hidden_dim;
|
||||
|
||||
problem_size1.m() = valid_word_num;
|
||||
problem_size1.n() = hidden_dim * intermediate_scale;
|
||||
problem_size1.k() = hidden_dim;
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
/// Prints the usage statement.
|
||||
std::ostream & print_usage(std::ostream &out) const {
|
||||
|
||||
out << "37_gemm_layernorm_gemm_fusion example\n\n"
|
||||
<< " This example uses the CUTLASS Library to compute GEMM + Layernorm for arbitrary problem sizes.\n\n"
|
||||
<< "Options:\n\n"
|
||||
<< " --help If specified, displays this usage statement.\n\n"
|
||||
<< " --hidden_dim=<int> Hidden dimension\n"
|
||||
<< " --valid_word_num=<int> Valid word number\n"
|
||||
<< " --seed=<int> Random number seed (1*)\n\n"
|
||||
<< " --iterations=<int> Number of profiling iterations to perform (0 to disable profiling).\n\n"
|
||||
<< " --verify=<bool> If true, performs reference calculation.\n\n"
|
||||
<< " --tolerance <float> Error tolerance\n"
|
||||
;
|
||||
|
||||
out << "\n\nExamples:\n\n"
|
||||
<< "$ ./examples/37_gemm_layernorm_gemm_fusion/37_gemm_layernorm_gemm_fusion \\\n"
|
||||
<< " --hidden_dim=768 --valid_word_num=1024 \n\n";
|
||||
|
||||
return out;
|
||||
}
|
||||
|
||||
/// Returns true if the environment and Toolkit support this
|
||||
bool supported(bool verbose = true) const {
|
||||
|
||||
// Ampere Tensor Core operations exposed with mma.sync and ldmatrix are first available
|
||||
// in CUDA 11.0.
|
||||
//
|
||||
// CUTLASS must be compiled with CUDA 11.0 Toolkit to run these examples.
|
||||
if (!(__CUDACC_VER_MAJOR__ >= 11)) {
|
||||
if (verbose) {
|
||||
std::cerr << "Ampere Tensor Core operations must be compiled with CUDA 11.0 Toolkit or later." << std::endl;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
cudaDeviceProp props;
|
||||
|
||||
cudaError_t error = cudaGetDeviceProperties(&props, 0);
|
||||
if (error != cudaSuccess) {
|
||||
if (verbose) {
|
||||
std::cerr << "cudaGetDeviceProperties() returned an error: " << cudaGetErrorString(error) << std::endl;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
if (!((props.major * 10 + props.minor) >= 80)) {
|
||||
if (verbose) {
|
||||
std::cerr << "Ampere Tensor Core operations must be run on a machine with compute capability at least 80."
|
||||
<< std::endl;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
//
|
||||
// CUTLASS attempts to load 128b vectors of cutlass::half_t (F16) elements. Consequently,
|
||||
// all pointers, strides, and tensor extents must be divisible by 8 elements.
|
||||
//
|
||||
int const kAlignment = 8;
|
||||
|
||||
if ((problem_size0.m() % kAlignment) ||
|
||||
(problem_size0.n() % kAlignment) ||
|
||||
(problem_size0.k() % kAlignment)) {
|
||||
if (verbose) {
|
||||
std::cerr << "Misaligned input in 1st GEMM." << std::endl;
|
||||
}
|
||||
// misaligned tensors for Gemm1
|
||||
return false;
|
||||
}
|
||||
|
||||
if ((problem_size1.m() % kAlignment) ||
|
||||
(problem_size1.n() % kAlignment) ||
|
||||
(problem_size1.k() % kAlignment)) {
|
||||
if (verbose) {
|
||||
std::cerr << "Misaligned input in 2nd GEMM." << std::endl;
|
||||
}
|
||||
// misaligned tensors for Gemm2
|
||||
return false;
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template<
|
||||
typename LayoutOutput_>
|
||||
struct Testbed {
|
||||
|
||||
//
|
||||
// Type definitions
|
||||
//
|
||||
|
||||
// User-defined data types
|
||||
using ElementInputA0 = cutlass::half_t;
|
||||
using ElementInputB0 = cutlass::half_t;
|
||||
using ElementOutput = cutlass::half_t;
|
||||
using ElementCompute = cutlass::half_t;
|
||||
|
||||
using LayoutInputA0 = cutlass::layout::RowMajor;
|
||||
using LayoutInputB0 = cutlass::layout::ColumnMajor;
|
||||
using LayoutOutput = LayoutOutput_;
|
||||
|
||||
static bool const kIsColumnMajorOutput = cutlass::platform::is_same<LayoutOutput, cutlass::layout::ColumnMajor>::value;
|
||||
// turn of shifted K by default
|
||||
static bool const kIsShiftedVariance = false;
|
||||
|
||||
/// Linear scaling operator
|
||||
using EpilogueFunctorOp = cutlass::epilogue::thread::LinearCombination<
|
||||
ElementOutput,
|
||||
128 / cutlass::sizeof_bits<ElementOutput>::value,
|
||||
ElementCompute,
|
||||
ElementCompute
|
||||
>;
|
||||
|
||||
using ThreadblockShape = cutlass::gemm::GemmShape<128, 128, 32>;
|
||||
using WarpShape = cutlass::gemm::GemmShape<64, 64, 32>;
|
||||
using InstructionShape = cutlass::gemm::GemmShape<16, 8, 16>;
|
||||
|
||||
static int const kStages0 = 3;
|
||||
static int const kStages1 = 4;
|
||||
|
||||
using GemmLayernorm = cutlass::GemmLayernorm<
|
||||
ElementInputA0,
|
||||
LayoutInputA0,
|
||||
ElementInputB0,
|
||||
LayoutInputB0,
|
||||
ElementOutput,
|
||||
LayoutOutput,
|
||||
ElementCompute,
|
||||
EpilogueFunctorOp,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
kStages0,
|
||||
kStages1,
|
||||
kIsShiftedVariance
|
||||
>;
|
||||
|
||||
using ElementInputA1 = typename GemmLayernorm::ElementInputA1;
|
||||
using ElementOutputC1 = typename GemmLayernorm::ElementOutputC1;
|
||||
using ElementInputScaleBias = typename GemmLayernorm::ElementInputScaleBias;
|
||||
using ElementLayernormCompute = typename GemmLayernorm::ElementLayernormCompute;
|
||||
|
||||
using LayoutInputA1 = typename GemmLayernorm::LayoutInputA1;
|
||||
using LayoutOutputC0 = typename GemmLayernorm::LayoutOutputC0;
|
||||
using LayoutOutputC1 = typename GemmLayernorm::LayoutOutputC1;
|
||||
using LayoutInputScaleBias = typename GemmLayernorm::LayoutInputScaleBias;
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
Options<LayoutOutput> const &options;
|
||||
|
||||
cutlass::HostTensor<ElementInputA0, LayoutInputA0> tensor_A0;
|
||||
cutlass::HostTensor<ElementInputB0, LayoutInputB0> tensor_B0;
|
||||
cutlass::HostTensor<ElementOutput, LayoutOutputC0> tensor_C0;
|
||||
cutlass::HostTensor<ElementInputA1, LayoutInputA1> tensor_A1;
|
||||
cutlass::HostTensor<ElementOutputC1, LayoutOutputC1> tensor_C1;
|
||||
|
||||
cutlass::HostTensor<ElementOutput, LayoutOutputC0> reference_C0;
|
||||
cutlass::HostTensor<ElementOutputC1, LayoutOutputC1> reference_C1;
|
||||
|
||||
cutlass::HostTensor<ElementInputScaleBias, LayoutInputScaleBias> tensor_Variance;
|
||||
cutlass::HostTensor<ElementInputScaleBias, LayoutInputScaleBias> tensor_Mean;
|
||||
cutlass::HostTensor<ElementInputScaleBias, LayoutInputScaleBias> tensor_Beta;
|
||||
cutlass::HostTensor<ElementInputScaleBias, LayoutInputScaleBias> tensor_Gamma;
|
||||
|
||||
cutlass::HostTensor<ElementInputScaleBias, LayoutInputScaleBias> reference_Mean;
|
||||
cutlass::HostTensor<ElementInputScaleBias, LayoutInputScaleBias> reference_Variance;
|
||||
|
||||
// shifted K tensor to better ensure the numerical stability
|
||||
// According to https://en.wikipedia.org/wiki/Algorithms_for_calculating_variance
|
||||
// the closer shifted K to the actual mean, the better numerical stability we'll observe
|
||||
cutlass::HostTensor<ElementOutput, LayoutOutputC0> tensor_Shifted_K;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
Testbed(
|
||||
Options<LayoutOutput> const &options_
|
||||
):
|
||||
options(options_)
|
||||
{
|
||||
|
||||
tensor_A0.reset({options.problem_size0.m(), options.problem_size0.k()});
|
||||
tensor_B0.reset({options.problem_size0.k(), options.problem_size0.n()});
|
||||
|
||||
tensor_C0.reset({options.problem_size0.m(), options.problem_size0.n()});
|
||||
|
||||
tensor_A1.reset({options.problem_size1.m(), options.problem_size1.k()});
|
||||
tensor_C1.reset({options.problem_size1.m(), options.problem_size1.n()});
|
||||
|
||||
reference_C0.reset({options.problem_size0.m(), options.problem_size0.n()});
|
||||
reference_C1.reset({options.problem_size1.m(), options.problem_size1.n()});
|
||||
|
||||
int leading_dim_0 = kIsColumnMajorOutput ? options.problem_size0.n() : options.problem_size0.m();
|
||||
int leading_dim_1 = kIsColumnMajorOutput ? options.problem_size0.m() : options.problem_size0.n();
|
||||
|
||||
int block_num = (leading_dim_1 + GemmLayernorm::ThreadblockShape::kM - 1) / GemmLayernorm::ThreadblockShape::kM;
|
||||
|
||||
tensor_Variance.reset({block_num, leading_dim_0});
|
||||
tensor_Mean.reset({block_num, leading_dim_0});
|
||||
tensor_Shifted_K.reset({1, leading_dim_0});
|
||||
|
||||
tensor_Beta.reset({1, leading_dim_1});
|
||||
tensor_Gamma.reset({1, leading_dim_1});
|
||||
|
||||
reference_Mean.reset({1, leading_dim_0}, false);
|
||||
reference_Variance.reset({1, leading_dim_0}, false);
|
||||
|
||||
}
|
||||
|
||||
/// Run
|
||||
Disposition run() {
|
||||
|
||||
Disposition disposition = Disposition::kNotVerified;
|
||||
|
||||
//
|
||||
// Initialize the workspace
|
||||
//
|
||||
|
||||
initialize();
|
||||
|
||||
//
|
||||
// Launch device kernel
|
||||
//
|
||||
cutlass::Status status = cutlass::Status::kSuccess;
|
||||
|
||||
status = execute_device_kernel();
|
||||
|
||||
if (status != cutlass::Status::kSuccess) {
|
||||
std::cerr << "Device execution failed." << std::endl;
|
||||
return disposition;
|
||||
}
|
||||
|
||||
cudaError_t result = cudaDeviceSynchronize();
|
||||
if (result != cudaSuccess) {
|
||||
std::cerr << "Device synchronize failed with error "
|
||||
<< cudaGetErrorString(result) << std::endl;
|
||||
return disposition;
|
||||
}
|
||||
|
||||
//
|
||||
// Compute the reference
|
||||
//
|
||||
compute_reference();
|
||||
|
||||
//
|
||||
// Verify
|
||||
//
|
||||
|
||||
if (options.verification_enabled) {
|
||||
|
||||
bool passed = verify();
|
||||
|
||||
if (passed) {
|
||||
disposition = Disposition::kPassed;
|
||||
}
|
||||
else {
|
||||
disposition = Disposition::kIncorrect;
|
||||
}
|
||||
}
|
||||
|
||||
//
|
||||
// Profiling
|
||||
//
|
||||
if (options.iterations) {
|
||||
profile();
|
||||
}
|
||||
|
||||
return disposition;
|
||||
}
|
||||
|
||||
/// Random initialization
|
||||
void initialize() {
|
||||
|
||||
cutlass::reference::host::TensorFillRandomUniform(
|
||||
tensor_A0.host_view(),
|
||||
options.seed,
|
||||
ElementInputA0(5),
|
||||
ElementInputA0(-5),
|
||||
0
|
||||
);
|
||||
|
||||
cutlass::reference::host::TensorFillRandomUniform(
|
||||
tensor_B0.host_view(),
|
||||
options.seed + 1,
|
||||
ElementInputB0(5),
|
||||
ElementInputB0(-5),
|
||||
0
|
||||
);
|
||||
|
||||
cutlass::reference::host::TensorFillRandomUniform(
|
||||
tensor_A1.host_view(),
|
||||
options.seed + 2,
|
||||
ElementInputA1(5),
|
||||
ElementInputA1(-5),
|
||||
0
|
||||
);
|
||||
|
||||
cutlass::reference::host::TensorFillRandomUniform(
|
||||
tensor_Beta.host_view(),
|
||||
options.seed + 3,
|
||||
ElementInputScaleBias(5),
|
||||
ElementInputScaleBias(-5),
|
||||
0
|
||||
);
|
||||
|
||||
cutlass::reference::host::TensorFillRandomUniform(
|
||||
tensor_Gamma.host_view(),
|
||||
options.seed + 4,
|
||||
ElementInputScaleBias(5),
|
||||
ElementInputScaleBias(-5),
|
||||
0
|
||||
);
|
||||
|
||||
cutlass::reference::host::TensorFillRandomUniform(
|
||||
tensor_Shifted_K.host_view(),
|
||||
options.seed + 5,
|
||||
ElementOutput(5),
|
||||
ElementOutput(-6),
|
||||
0
|
||||
);
|
||||
|
||||
tensor_A0.sync_device();
|
||||
tensor_B0.sync_device();
|
||||
tensor_A1.sync_device();
|
||||
tensor_Beta.sync_device();
|
||||
tensor_Gamma.sync_device();
|
||||
|
||||
}
|
||||
|
||||
|
||||
|
||||
cutlass::Status execute_device_kernel() {
|
||||
|
||||
cutlass::Status status = cutlass::Status::kSuccess;
|
||||
|
||||
//
|
||||
// Setup arguments
|
||||
//
|
||||
|
||||
typename GemmLayernorm::Arguments args(
|
||||
options.problem_size0,
|
||||
options.problem_size1,
|
||||
tensor_A0.device_ref().data(),
|
||||
tensor_B0.device_ref().data(),
|
||||
tensor_C0.device_ref().data(),
|
||||
tensor_C0.device_ref().data(),
|
||||
tensor_A1.device_ref().data(),
|
||||
tensor_C1.device_ref().data(),
|
||||
tensor_A0.device_ref().stride(0),
|
||||
tensor_B0.device_ref().stride(0),
|
||||
tensor_C0.device_ref().stride(0),
|
||||
tensor_C0.device_ref().stride(0),
|
||||
tensor_A1.device_ref().stride(0),
|
||||
tensor_C1.device_ref().stride(0),
|
||||
{
|
||||
ElementCompute(options.alpha),
|
||||
ElementCompute(options.beta)
|
||||
},
|
||||
tensor_Variance.device_ref(),
|
||||
tensor_Mean.device_ref(),
|
||||
tensor_Gamma.device_ref(),
|
||||
tensor_Beta.device_ref(),
|
||||
tensor_Shifted_K.device_ref().data()
|
||||
);
|
||||
|
||||
//
|
||||
// Launch
|
||||
//
|
||||
|
||||
GemmLayernorm gemm_layernorm;
|
||||
|
||||
// Initialize
|
||||
status = gemm_layernorm.initialize(args);
|
||||
if (status != cutlass::Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
// Run
|
||||
status = gemm_layernorm();
|
||||
|
||||
return status;
|
||||
}
|
||||
|
||||
/// Reference calculation
|
||||
void compute_reference() {
|
||||
|
||||
cutlass::reference::device::Gemm<
|
||||
ElementInputA0,
|
||||
LayoutInputA0,
|
||||
ElementInputB0,
|
||||
LayoutInputB0,
|
||||
ElementOutput,
|
||||
LayoutOutputC0,
|
||||
ElementCompute,
|
||||
ElementCompute
|
||||
> gemm_device0;
|
||||
|
||||
cutlass::reference::device::Gemm<
|
||||
ElementInputA1,
|
||||
LayoutInputA1,
|
||||
ElementOutput,
|
||||
LayoutOutputC0,
|
||||
ElementOutputC1,
|
||||
LayoutOutputC1,
|
||||
ElementCompute,
|
||||
ElementCompute
|
||||
> gemm_device1;
|
||||
|
||||
// Compute 1st GEMM
|
||||
gemm_device0(
|
||||
options.problem_size0,
|
||||
ElementCompute(options.alpha),
|
||||
tensor_A0.device_ref(),
|
||||
tensor_B0.device_ref(),
|
||||
ElementCompute(options.beta),
|
||||
tensor_C0.device_ref(),
|
||||
reference_C0.device_ref()
|
||||
);
|
||||
|
||||
reference_C0.sync_host();
|
||||
|
||||
tensor_Mean.sync_host();
|
||||
tensor_Variance.sync_host();
|
||||
tensor_Gamma.sync_host();
|
||||
tensor_Beta.sync_host();
|
||||
tensor_Shifted_K.sync_host();
|
||||
|
||||
// Compute the sum and square sum for verification purpose
|
||||
if (kIsColumnMajorOutput) {
|
||||
for (int n = 0; n < options.problem_size0.n(); ++n) {
|
||||
|
||||
ElementLayernormCompute sum = ElementLayernormCompute(0);
|
||||
ElementLayernormCompute square_sum = ElementLayernormCompute(0);
|
||||
for (int m = 0; m < options.problem_size0.m(); ++m) {
|
||||
sum += ElementLayernormCompute(reference_C0.at({m, n}));
|
||||
square_sum += ElementLayernormCompute(reference_C0.at({m, n})) * ElementLayernormCompute(reference_C0.at({m, n}));
|
||||
}
|
||||
|
||||
ElementLayernormCompute mean = sum / ElementLayernormCompute(options.problem_size0.m());
|
||||
ElementLayernormCompute square_mean = square_sum / ElementLayernormCompute(options.problem_size0.m());
|
||||
ElementLayernormCompute variance = cutlass::constants::one<ElementLayernormCompute>() / cutlass::fast_sqrt(square_mean - mean * mean + ElementLayernormCompute(1e-6) ) ;
|
||||
|
||||
mean = -mean * variance;
|
||||
|
||||
reference_Mean.at({0, n}) = ElementInputScaleBias(mean);
|
||||
reference_Variance.at({0, n}) = ElementInputScaleBias(variance);
|
||||
}
|
||||
}else{
|
||||
for (int m = 0; m < options.problem_size0.m(); ++m) {
|
||||
|
||||
ElementLayernormCompute sum = ElementLayernormCompute(0);
|
||||
ElementLayernormCompute square_sum = ElementLayernormCompute(0);
|
||||
for (int n = 0; n < options.problem_size0.n(); ++n) {
|
||||
sum += ElementLayernormCompute(reference_C0.at({m, n})) ;
|
||||
square_sum += ElementLayernormCompute(reference_C0.at({m, n})) * ElementLayernormCompute(reference_C0.at({m, n})) ;
|
||||
}
|
||||
|
||||
ElementLayernormCompute mean = sum / ElementLayernormCompute(options.problem_size0.n());
|
||||
ElementLayernormCompute square_mean = square_sum / ElementLayernormCompute(options.problem_size0.n());
|
||||
ElementLayernormCompute variance = cutlass::constants::one<ElementLayernormCompute>() / cutlass::fast_sqrt(square_mean - mean * mean + ElementLayernormCompute(1e-6)) ;
|
||||
|
||||
mean = -mean * variance;
|
||||
|
||||
reference_Mean.at({0, m}) = ElementInputScaleBias(mean);
|
||||
reference_Variance.at({0, m}) = ElementInputScaleBias(variance);
|
||||
}
|
||||
}
|
||||
|
||||
// Element-wise transform for OutputC0 using 1-pass layernorm algo
|
||||
if (kIsColumnMajorOutput) {
|
||||
for (int n = 0; n < options.problem_size0.n(); ++n) {
|
||||
|
||||
ElementLayernormCompute sum = ElementLayernormCompute(0);
|
||||
for (int m = 0; m < options.problem_size0.m(); ++m) {
|
||||
sum += ElementLayernormCompute(reference_C0.at({m, n})) ;
|
||||
}
|
||||
|
||||
ElementInputScaleBias mean = ElementInputScaleBias(sum / ElementLayernormCompute(options.problem_size0.m()));
|
||||
sum = ElementLayernormCompute(0);
|
||||
for (int m = 0; m < options.problem_size0.m(); ++m) {
|
||||
sum += ElementLayernormCompute(reference_C0.at({m, n}) - ElementLayernormCompute(mean)) * ElementLayernormCompute(reference_C0.at({m, n}) - ElementLayernormCompute(mean)) ;
|
||||
}
|
||||
|
||||
ElementLayernormCompute square_mean = sum / ElementLayernormCompute(options.problem_size0.m());
|
||||
ElementInputScaleBias variance = ElementInputScaleBias(cutlass::constants::one<ElementLayernormCompute>()
|
||||
/ cutlass::fast_sqrt(square_mean + ElementLayernormCompute(1e-6))) ;
|
||||
|
||||
for (int m = 0; m < options.problem_size0.m(); ++m) {
|
||||
reference_C0.at({m, n}) =
|
||||
ElementOutput( ( (ElementInputScaleBias(reference_C0.at({m, n})) - mean) * variance )
|
||||
* tensor_Gamma.at({0, m}) + tensor_Beta.at({0, m}));
|
||||
|
||||
}
|
||||
|
||||
}
|
||||
}else{
|
||||
|
||||
for (int m = 0; m < options.problem_size0.m(); ++m) {
|
||||
|
||||
float sum = float(0);
|
||||
for (int n = 0; n < options.problem_size0.n(); ++n) {
|
||||
sum += float(reference_C0.at({m, n})) ;
|
||||
}
|
||||
|
||||
float mean = sum / float(options.problem_size0.n());
|
||||
sum = float(0);
|
||||
for (int n = 0; n < options.problem_size0.n(); ++n) {
|
||||
sum += float(reference_C0.at({m, n}) - mean) * float(reference_C0.at({m, n}) - mean) ;
|
||||
}
|
||||
|
||||
float square_mean = sum / float(options.problem_size0.n());
|
||||
float variance = cutlass::constants::one<float>() / cutlass::fast_sqrt(square_mean + ElementLayernormCompute(1e-6)) ;
|
||||
|
||||
for (int n = 0; n < options.problem_size0.n(); ++n) {
|
||||
reference_C0.at({m, n}) =
|
||||
ElementOutput( ( (float(reference_C0.at({m, n})) - mean) * variance )
|
||||
* float(tensor_Gamma.at({0, n})) + float(tensor_Beta.at({0, n})));
|
||||
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
|
||||
// Sync host data with device after element-wise transform
|
||||
reference_C0.sync_device();
|
||||
|
||||
// Compute 2nd GEMM
|
||||
gemm_device1(
|
||||
options.problem_size1,
|
||||
ElementCompute(options.alpha),
|
||||
kIsColumnMajorOutput ? tensor_A1.device_ref() : reference_C0.device_ref(),
|
||||
kIsColumnMajorOutput ? reference_C0.device_ref() :tensor_A1.device_ref(),
|
||||
ElementCompute(options.beta),
|
||||
reference_C1.device_ref(),
|
||||
reference_C1.device_ref()
|
||||
);
|
||||
|
||||
}
|
||||
|
||||
/// Emits all tensor values
|
||||
void emit_results() {
|
||||
std::cout << "tensor_C1 = \n" << tensor_C1.host_view() << "\n\n";
|
||||
std::cout << "Reference C1 = \n" << reference_C1.host_view() << "\n\n";
|
||||
std::cout << "Mean = \n" << tensor_Mean.host_view() << "\n\n";
|
||||
std::cout << "rsqrt(Variance) = \n" << tensor_Variance.host_view() << "\n\n";
|
||||
std::cout << "Reference Mean = \n" << reference_Mean.host_view() << "\n\n";
|
||||
std::cout << "Reference rsqrt(Variance) = \n" << reference_Variance.host_view() << "\n\n";
|
||||
}
|
||||
|
||||
template<typename Element, typename Layout>
|
||||
bool verify_tensor(cutlass::HostTensor<Element, Layout> tensor, \
|
||||
cutlass::HostTensor<Element, Layout> reference,
|
||||
int leading_dim0, int leading_dim1, bool is_print = false) {
|
||||
float const kThreshold = float(options.tolerance);
|
||||
float const kAbsThreshold = 0.5f;
|
||||
float const kRelativeThreshold = 0.1f;
|
||||
// Adds a constant bias to avoid being divided by '0'
|
||||
float const kBias = 1e-5f;
|
||||
int counter = 0;
|
||||
for (int m = 0; m < leading_dim0; m++) {
|
||||
for (int n = 0; n < leading_dim1; ++n) {
|
||||
float diff = (float)(tensor.at({m, n}) - reference.at({m, n}));
|
||||
float rel_diff = fabs(diff) / fabs(reference.at({m, n}) + kBias);
|
||||
if (fabs(diff) > kAbsThreshold && rel_diff > kRelativeThreshold) {
|
||||
counter++;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
float err_rate = float(counter) / (float(leading_dim0) * float(leading_dim1));
|
||||
return (err_rate < kThreshold);
|
||||
}
|
||||
|
||||
/// Verifies the reference matches
|
||||
bool verify() {
|
||||
|
||||
tensor_Variance.sync_host();
|
||||
tensor_Mean.sync_host();
|
||||
tensor_C1.sync_host();
|
||||
reference_C1.sync_host();
|
||||
|
||||
// Verification checks - set any of these to 'true' to override the verification checks.
|
||||
bool verified_C1 = false;
|
||||
bool verified_Mean = false;
|
||||
bool verified_Variance = false;
|
||||
|
||||
// Verify layernorm output
|
||||
if (!verified_C1) {
|
||||
verified_C1 = verify_tensor<ElementOutputC1, LayoutOutputC1>(tensor_C1, reference_C1, options.problem_size1.m(), options.problem_size1.n());
|
||||
}
|
||||
|
||||
if (!verified_Variance) {
|
||||
verified_Variance = verify_tensor<ElementInputScaleBias, LayoutInputScaleBias>(tensor_Variance, reference_Variance, 1, options.problem_size0.n());
|
||||
}
|
||||
|
||||
if (!verified_Mean) {
|
||||
verified_Mean = verify_tensor<ElementInputScaleBias, LayoutInputScaleBias>(tensor_Mean, reference_Mean, 1, options.problem_size0.n());
|
||||
}
|
||||
|
||||
if (!verified_C1 || !verified_Mean || !verified_Variance) {
|
||||
|
||||
// emit_results();
|
||||
|
||||
std::cerr << "Verification check failed for tensor Layernorm" << std::endl;
|
||||
|
||||
// Summarize which checks failed
|
||||
if (!verified_C1) {
|
||||
std::cerr << "Verification of O tensor failed\n";
|
||||
}
|
||||
|
||||
if (!verified_Mean) {
|
||||
std::cerr << "Verification of Mean tensor failed\n";
|
||||
}
|
||||
|
||||
if (!verified_Variance) {
|
||||
std::cerr << "Verification of Variance tensor failed\n";
|
||||
}
|
||||
|
||||
return false;
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
/// Profiles
|
||||
bool profile() {
|
||||
|
||||
//
|
||||
// Profile
|
||||
//
|
||||
|
||||
cutlass::Status status = cutlass::Status::kSuccess;
|
||||
cudaError_t result;
|
||||
cudaEvent_t events[2];
|
||||
int const kIterations = options.iterations;
|
||||
|
||||
for (cudaEvent_t &evt : events) {
|
||||
result = cudaEventCreate(&evt);
|
||||
if (result != cudaSuccess) {
|
||||
std::cerr << "cudaEventCreate failed with error " << cudaGetErrorString(result) << std::endl;
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
result = cudaEventRecord(events[0]);
|
||||
|
||||
if (result != cudaSuccess) {
|
||||
std::cerr << "cudaEventRecord() failed with error " << cudaGetErrorString(result) << std::endl;
|
||||
return false;
|
||||
}
|
||||
|
||||
for (int iter = 0; iter < kIterations; ++iter) {
|
||||
|
||||
status = execute_device_kernel();
|
||||
|
||||
if (status != cutlass::Status::kSuccess) {
|
||||
std::cerr << "Device execution failed." << std::endl;
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
result = cudaEventRecord(events[1]);
|
||||
|
||||
if (result != cudaSuccess) {
|
||||
std::cerr << "cudaEventRecord() failed with error " << cudaGetErrorString(result) << std::endl;
|
||||
return false;
|
||||
}
|
||||
|
||||
result = cudaDeviceSynchronize();
|
||||
|
||||
if (result != cudaSuccess) {
|
||||
std::cerr << "cudaDeviceSynchronize() failed with error " << cudaGetErrorString(result) << std::endl;
|
||||
return false;
|
||||
}
|
||||
|
||||
float elapsed_ms = 0;
|
||||
result = cudaEventElapsedTime(&elapsed_ms, events[0], events[1]);
|
||||
|
||||
float elapsed_ms_per_iter = elapsed_ms / float(kIterations);
|
||||
|
||||
if (result != cudaSuccess) {
|
||||
std::cerr << "cudaEventElapsedTime() failed with error " << cudaGetErrorString(result) << std::endl;
|
||||
return false;
|
||||
}
|
||||
|
||||
for (cudaEvent_t &evt : events) {
|
||||
result = cudaEventDestroy(evt);
|
||||
if (result != cudaSuccess) {
|
||||
std::cerr << "cudaEventDestroy() failed with error " << cudaGetErrorString(result) << std::endl;
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
int64_t flops = int64_t(options.problem_size0.m()) * options.problem_size0.n() * options.problem_size0.k() * 2 \
|
||||
+ int64_t(options.problem_size1.m()) * options.problem_size1.n() * options.problem_size1.k() * 2;
|
||||
|
||||
double gflops_per_second = double(flops) * kIterations / double(elapsed_ms / 1000.0f) / double(1.0e9);
|
||||
|
||||
std::cout << " 1st GEMM: "
|
||||
<< options.problem_size0.m() << "-by-" << options.problem_size0.n() << "-by-" << options.problem_size0.k() << "\n"
|
||||
<< " 2nd GEMM: "
|
||||
<< options.problem_size1.m() << "-by-" << options.problem_size1.n() << "-by-" << options.problem_size1.k()
|
||||
<< std::endl;
|
||||
|
||||
std::cout << " Runtime / iteration: " << elapsed_ms_per_iter << " ms\n" << std::endl;
|
||||
std::cout << " GFLOPs: " << gflops_per_second << " GFLOPs" << std::endl;
|
||||
|
||||
return true;
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
int main(int argc, const char **argv) {
|
||||
|
||||
// Define final layout
|
||||
using LayoutOutput = cutlass::layout::ColumnMajor;
|
||||
|
||||
// Options parsing
|
||||
Options<LayoutOutput> options;
|
||||
options.parse(argc, argv);
|
||||
|
||||
if (options.help) {
|
||||
options.print_usage(std::cout) << std::endl;
|
||||
return 0;
|
||||
}
|
||||
|
||||
if (!options.supported()) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
// Run
|
||||
Testbed<LayoutOutput> testbed(options);
|
||||
|
||||
Disposition disposition = testbed.run();
|
||||
|
||||
std::cout << std::endl;
|
||||
|
||||
switch (disposition) {
|
||||
case Disposition::kPassed:
|
||||
std::cout << "Passed" << std::endl;
|
||||
break;
|
||||
case Disposition::kIncorrect:
|
||||
std::cout << "Incorrect" << std::endl;
|
||||
break;
|
||||
case Disposition::kNotVerified:
|
||||
std::cout << "Not verified" << std::endl;
|
||||
break;
|
||||
}
|
||||
|
||||
return (disposition == Disposition::kPassed ? 0 : -1);
|
||||
}
|
||||
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,450 @@
|
||||
/***************************************************************************************************
|
||||
* 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 GEMM kernel to support the epilogue visitor model
|
||||
for customized layernorm partial reduction epilogue fusion.
|
||||
|
||||
This source file will likely be moved to `include/cutlass/gemm/kernel/` in the future once
|
||||
its usage has been stabilized. For now, it is included in this example to demonstrate
|
||||
some basic output fusion options.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/fast_math.h"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/matrix_coord.h"
|
||||
#include "cutlass/complex.h"
|
||||
#include "cutlass/semaphore.h"
|
||||
|
||||
#include "cutlass/trace.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace kernel {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename Mma_, ///! Threadblock-scoped matrix multiply-accumulate
|
||||
typename Epilogue_, ///! Epilogue
|
||||
typename ThreadblockSwizzle_ ///! Threadblock swizzling function
|
||||
>
|
||||
struct GemmWithEpilogueVisitor {
|
||||
public:
|
||||
|
||||
using Mma = Mma_;
|
||||
using Epilogue = Epilogue_;
|
||||
using EpilogueVisitor = typename Epilogue::Visitor;
|
||||
using ThreadblockSwizzle = ThreadblockSwizzle_;
|
||||
|
||||
using ElementA = typename Mma::IteratorA::Element;
|
||||
using LayoutA = typename Mma::IteratorA::Layout;
|
||||
using TensorRefA = TensorRef<ElementA, LayoutA>;
|
||||
|
||||
using ElementB = typename Mma::IteratorB::Element;
|
||||
using LayoutB = typename Mma::IteratorB::Layout;
|
||||
using TensorRefB = TensorRef<ElementB, LayoutB>;
|
||||
|
||||
using ElementC = typename EpilogueVisitor::ElementOutput;
|
||||
using LayoutC = typename Epilogue::Layout;
|
||||
|
||||
static ComplexTransform const kTransformA = Mma::kTransformA;
|
||||
static ComplexTransform const kTransformB = Mma::kTransformB;
|
||||
using Operator = typename Mma::Operator;
|
||||
|
||||
using OperatorClass = typename Mma::Operator::OperatorClass;
|
||||
using ThreadblockShape = typename Mma::Shape;
|
||||
using WarpShape = typename Mma::Operator::Shape;
|
||||
using InstructionShape = typename Mma::Policy::Operator::InstructionShape;
|
||||
using ArchTag = typename Mma::ArchTag;
|
||||
|
||||
static int const kStages = Mma::kStages;
|
||||
static int const kAlignmentA = Mma::IteratorA::AccessType::kElements;
|
||||
static int const kAlignmentB = Mma::IteratorB::AccessType::kElements;
|
||||
static int const kAlignmentC = EpilogueVisitor::kElementsPerAccess;
|
||||
|
||||
/// Warp count (concept: GemmShape)
|
||||
using WarpCount = typename Mma::WarpCount;
|
||||
static int const kThreadCount = 32 * WarpCount::kCount;
|
||||
|
||||
/// Split-K preserves splits that are 128b aligned
|
||||
static int const kSplitKAlignment = const_max(
|
||||
128 / sizeof_bits<ElementA>::value,
|
||||
128 / sizeof_bits<ElementB>::value
|
||||
);
|
||||
|
||||
//
|
||||
// Structures
|
||||
//
|
||||
|
||||
/// Argument structure
|
||||
struct Arguments {
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
GemmUniversalMode mode;
|
||||
GemmCoord problem_size;
|
||||
|
||||
TensorRefA ref_A;
|
||||
TensorRefB ref_B;
|
||||
|
||||
typename EpilogueVisitor::Arguments epilogue_visitor;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
Arguments():
|
||||
mode(GemmUniversalMode::kGemm)
|
||||
{ }
|
||||
|
||||
|
||||
/// constructs an arguments structure
|
||||
Arguments(
|
||||
GemmUniversalMode mode_,
|
||||
GemmCoord problem_size_,
|
||||
TensorRefA ref_A_,
|
||||
TensorRefB ref_B_,
|
||||
typename EpilogueVisitor::Arguments epilogue_visitor_
|
||||
):
|
||||
mode(mode_),
|
||||
problem_size(problem_size_),
|
||||
ref_A(ref_A_),
|
||||
ref_B(ref_B_),
|
||||
epilogue_visitor(epilogue_visitor_)
|
||||
{
|
||||
|
||||
}
|
||||
};
|
||||
|
||||
//
|
||||
// Structure for precomputing values in host memory and passing to kernels
|
||||
//
|
||||
|
||||
/// Parameters structure
|
||||
struct Params {
|
||||
|
||||
cutlass::gemm::GemmCoord problem_size;
|
||||
cutlass::gemm::GemmCoord grid_tiled_shape;
|
||||
int swizzle_log_tile;
|
||||
|
||||
typename Mma::IteratorA::Params params_A;
|
||||
typename Mma::IteratorB::Params params_B;
|
||||
|
||||
GemmUniversalMode mode;
|
||||
int gemm_k_size;
|
||||
|
||||
void * ptr_A;
|
||||
void * ptr_B;
|
||||
|
||||
typename EpilogueVisitor::Params epilogue_visitor;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params():
|
||||
swizzle_log_tile(0),
|
||||
params_A(0),
|
||||
params_B(0),
|
||||
gemm_k_size(0),
|
||||
mode(cutlass::gemm::GemmUniversalMode::kGemm),
|
||||
ptr_A(nullptr),
|
||||
ptr_B(nullptr)
|
||||
{ }
|
||||
|
||||
|
||||
Params(
|
||||
Arguments const &args
|
||||
):
|
||||
problem_size(args.problem_size),
|
||||
swizzle_log_tile(0),
|
||||
params_A(args.ref_A.layout()),
|
||||
params_B(args.ref_B.layout()),
|
||||
mode(args.mode),
|
||||
gemm_k_size(args.problem_size.k()),
|
||||
ptr_A(args.ref_A.data()),
|
||||
ptr_B(args.ref_B.data()),
|
||||
epilogue_visitor(args.epilogue_visitor)
|
||||
{
|
||||
|
||||
ThreadblockSwizzle threadblock_swizzle;
|
||||
|
||||
grid_tiled_shape = threadblock_swizzle.get_tiled_shape(
|
||||
args.problem_size,
|
||||
{ThreadblockShape::kM, ThreadblockShape::kN, ThreadblockShape::kK}, 1);
|
||||
|
||||
if (args.mode == GemmUniversalMode::kGemm || args.mode == GemmUniversalMode::kGemmSplitKParallel) {
|
||||
|
||||
int const kAlignK = const_max(const_max(128 / sizeof_bits<ElementA>::value, 128 / sizeof_bits<ElementB>::value), 1);
|
||||
|
||||
gemm_k_size = round_up(args.problem_size.k(), kAlignK);
|
||||
|
||||
if (gemm_k_size) {
|
||||
grid_tiled_shape.k() = ceil_div(args.problem_size.k(), gemm_k_size);
|
||||
}
|
||||
}
|
||||
|
||||
swizzle_log_tile = threadblock_swizzle.get_log_tile(grid_tiled_shape);
|
||||
}
|
||||
};
|
||||
|
||||
/// Shared memory storage structure
|
||||
union SharedStorage {
|
||||
|
||||
typename Mma::SharedStorage main_loop;
|
||||
|
||||
struct {
|
||||
typename Epilogue::SharedStorage epilogue;
|
||||
typename EpilogueVisitor::SharedStorage visitor;
|
||||
} epilogue;
|
||||
};
|
||||
|
||||
public:
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
CUTLASS_DEVICE
|
||||
GemmWithEpilogueVisitor() { }
|
||||
|
||||
/// Determines whether kernel satisfies alignment
|
||||
static Status can_implement(
|
||||
cutlass::gemm::GemmCoord const & problem_size) {
|
||||
|
||||
CUTLASS_TRACE_HOST("GemmWithEpilogueVisitor::can_implement()");
|
||||
|
||||
static int const kAlignmentA = Mma::IteratorA::AccessType::kElements;
|
||||
static int const kAlignmentB = Mma::IteratorB::AccessType::kElements;
|
||||
static int const kAlignmentC = Epilogue::OutputTileIterator::kElementsPerAccess;
|
||||
|
||||
bool isAMisaligned = false;
|
||||
bool isBMisaligned = false;
|
||||
bool isCMisaligned = false;
|
||||
|
||||
if (platform::is_same<LayoutA, layout::RowMajor>::value) {
|
||||
isAMisaligned = problem_size.k() % kAlignmentA;
|
||||
} else if (platform::is_same<LayoutA, layout::ColumnMajor>::value) {
|
||||
isAMisaligned = problem_size.m() % kAlignmentA;
|
||||
} else if (platform::is_same<LayoutA, layout::ColumnMajorInterleaved<32>>::value
|
||||
|| platform::is_same<LayoutA, layout::ColumnMajorInterleaved<64>>::value) {
|
||||
isAMisaligned = problem_size.k() % kAlignmentA;
|
||||
}
|
||||
|
||||
if (platform::is_same<LayoutB, layout::RowMajor>::value) {
|
||||
isBMisaligned = problem_size.n() % kAlignmentB;
|
||||
} else if (platform::is_same<LayoutB, layout::ColumnMajor>::value) {
|
||||
isBMisaligned = problem_size.k() % kAlignmentB;
|
||||
} else if (platform::is_same<LayoutB, layout::RowMajorInterleaved<32>>::value
|
||||
|| platform::is_same<LayoutB, layout::RowMajorInterleaved<64>>::value) {
|
||||
isBMisaligned = problem_size.k() % kAlignmentB;
|
||||
}
|
||||
|
||||
if (platform::is_same<LayoutC, layout::RowMajor>::value) {
|
||||
isCMisaligned = problem_size.n() % kAlignmentC;
|
||||
} else if (platform::is_same<LayoutC, layout::ColumnMajor>::value) {
|
||||
isCMisaligned = problem_size.m() % kAlignmentC;
|
||||
} else if (platform::is_same<LayoutC, layout::ColumnMajorInterleaved<32>>::value
|
||||
|| platform::is_same<LayoutC, layout::ColumnMajorInterleaved<64>>::value) {
|
||||
isCMisaligned = problem_size.n() % kAlignmentC;
|
||||
}
|
||||
|
||||
if (isAMisaligned) {
|
||||
CUTLASS_TRACE_HOST(" returning kErrorMisalignedOperand for A operand");
|
||||
return Status::kErrorMisalignedOperand;
|
||||
}
|
||||
|
||||
if (isBMisaligned) {
|
||||
CUTLASS_TRACE_HOST(" returning kErrorMisalignedOperand for B operand");
|
||||
return Status::kErrorMisalignedOperand;
|
||||
}
|
||||
|
||||
if (isCMisaligned) {
|
||||
CUTLASS_TRACE_HOST(" returning kErrorMisalignedOperand for C operand");
|
||||
return Status::kErrorMisalignedOperand;
|
||||
}
|
||||
|
||||
CUTLASS_TRACE_HOST(" returning kSuccess");
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
static Status can_implement(Arguments const &args) {
|
||||
return can_implement(args.problem_size);
|
||||
}
|
||||
|
||||
static size_t get_extra_workspace_size(Arguments const &args,
|
||||
cutlass::gemm::GemmCoord const &grid_tiled_shape) {
|
||||
|
||||
return 0;
|
||||
}
|
||||
|
||||
/// Executes one GEMM
|
||||
CUTLASS_DEVICE
|
||||
void operator()(Params const ¶ms, SharedStorage &shared_storage) {
|
||||
|
||||
// Compute threadblock location
|
||||
ThreadblockSwizzle threadblock_swizzle;
|
||||
|
||||
cutlass::gemm::GemmCoord threadblock_tile_offset = threadblock_swizzle.get_tile_offset(params.swizzle_log_tile);
|
||||
|
||||
// Early exit if CTA is out of range
|
||||
if (params.grid_tiled_shape.m() <= threadblock_tile_offset.m() ||
|
||||
params.grid_tiled_shape.n() <= threadblock_tile_offset.n()) {
|
||||
|
||||
return;
|
||||
}
|
||||
|
||||
int offset_k = 0;
|
||||
int problem_size_k = params.problem_size.k();
|
||||
|
||||
ElementA *ptr_A = static_cast<ElementA *>(params.ptr_A);
|
||||
ElementB *ptr_B = static_cast<ElementB *>(params.ptr_B);
|
||||
|
||||
// Compute initial location in logical coordinates
|
||||
cutlass::MatrixCoord tb_offset_A{
|
||||
threadblock_tile_offset.m() * Mma::Shape::kM,
|
||||
offset_k,
|
||||
};
|
||||
|
||||
cutlass::MatrixCoord tb_offset_B{
|
||||
offset_k,
|
||||
threadblock_tile_offset.n() * Mma::Shape::kN
|
||||
};
|
||||
|
||||
// Compute position within threadblock
|
||||
int thread_idx = threadIdx.x;
|
||||
|
||||
// Construct iterators to A and B operands
|
||||
typename Mma::IteratorA iterator_A(
|
||||
params.params_A,
|
||||
ptr_A,
|
||||
{params.problem_size.m(), problem_size_k},
|
||||
thread_idx,
|
||||
tb_offset_A);
|
||||
|
||||
typename Mma::IteratorB iterator_B(
|
||||
params.params_B,
|
||||
ptr_B,
|
||||
{problem_size_k, params.problem_size.n()},
|
||||
thread_idx,
|
||||
tb_offset_B);
|
||||
|
||||
// Broadcast the warp_id computed by lane 0 to ensure dependent code
|
||||
// is compiled as warp-uniform.
|
||||
int warp_idx = __shfl_sync(0xffffffff, threadIdx.x / 32, 0);
|
||||
|
||||
int lane_idx = threadIdx.x % 32;
|
||||
|
||||
//
|
||||
// Main loop
|
||||
//
|
||||
|
||||
// Construct thread-scoped matrix multiply
|
||||
Mma mma(shared_storage.main_loop, thread_idx, warp_idx, lane_idx);
|
||||
|
||||
typename Mma::FragmentC accumulators;
|
||||
|
||||
accumulators.clear();
|
||||
|
||||
// Compute threadblock-scoped matrix multiply-add
|
||||
int gemm_k_iterations = (problem_size_k - offset_k + Mma::Shape::kK - 1) / Mma::Shape::kK;
|
||||
|
||||
// Compute threadblock-scoped matrix multiply-add
|
||||
mma(
|
||||
gemm_k_iterations,
|
||||
accumulators,
|
||||
iterator_A,
|
||||
iterator_B,
|
||||
accumulators);
|
||||
|
||||
//
|
||||
// Masked tile iterators constructed from members
|
||||
//
|
||||
|
||||
threadblock_tile_offset = threadblock_swizzle.get_tile_offset(params.swizzle_log_tile);
|
||||
|
||||
//assume identity swizzle
|
||||
MatrixCoord threadblock_offset(
|
||||
threadblock_tile_offset.m() * Mma::Shape::kM,
|
||||
threadblock_tile_offset.n() * Mma::Shape::kN
|
||||
);
|
||||
|
||||
int block_idx = threadblock_tile_offset.m() + threadblock_tile_offset.n() * params.grid_tiled_shape.m();
|
||||
|
||||
//
|
||||
// Construct the epilogue visitor
|
||||
//
|
||||
|
||||
EpilogueVisitor epilogue_visitor(
|
||||
params.epilogue_visitor,
|
||||
shared_storage.epilogue.visitor,
|
||||
params.problem_size.mn(),
|
||||
thread_idx,
|
||||
warp_idx,
|
||||
lane_idx,
|
||||
threadblock_offset);
|
||||
|
||||
if (params.mode == GemmUniversalMode::kGemm) {
|
||||
// Indicate which position in a serial reduction the output operator is currently updating
|
||||
epilogue_visitor.set_k_partition(threadblock_tile_offset.k(), params.grid_tiled_shape.k());
|
||||
}
|
||||
else if (params.mode == GemmUniversalMode::kBatched || params.mode == GemmUniversalMode::kArray) {
|
||||
epilogue_visitor.set_batch_index(threadblock_tile_offset.k());
|
||||
}
|
||||
|
||||
// Construct the epilogue
|
||||
Epilogue epilogue(
|
||||
shared_storage.epilogue.epilogue,
|
||||
thread_idx,
|
||||
warp_idx,
|
||||
lane_idx);
|
||||
|
||||
// Execute the epilogue operator to update the destination tensor.
|
||||
epilogue(epilogue_visitor, accumulators);
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace kernel
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
1066
examples/37_gemm_layernorm_gemm_fusion/gemm_with_layernorm.h
Normal file
1066
examples/37_gemm_layernorm_gemm_fusion/gemm_with_layernorm.h
Normal file
File diff suppressed because it is too large
Load Diff
36
examples/38_syr2k_grouped/CMakeLists.txt
Normal file
36
examples/38_syr2k_grouped/CMakeLists.txt
Normal file
@@ -0,0 +1,36 @@
|
||||
|
||||
# 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.
|
||||
|
||||
|
||||
|
||||
cutlass_example_add_executable(
|
||||
38_syr2k_grouped
|
||||
syr2k_grouped.cu
|
||||
)
|
||||
|
||||
1461
examples/38_syr2k_grouped/syr2k_grouped.cu
Normal file
1461
examples/38_syr2k_grouped/syr2k_grouped.cu
Normal file
File diff suppressed because it is too large
Load Diff
36
examples/39_gemm_permute/CMakeLists.txt
Normal file
36
examples/39_gemm_permute/CMakeLists.txt
Normal file
@@ -0,0 +1,36 @@
|
||||
|
||||
# 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.
|
||||
|
||||
|
||||
|
||||
cutlass_example_add_executable(
|
||||
39_gemm_permute
|
||||
gemm_permute.cu
|
||||
)
|
||||
|
||||
1126
examples/39_gemm_permute/gemm_permute.cu
Normal file
1126
examples/39_gemm_permute/gemm_permute.cu
Normal file
File diff suppressed because it is too large
Load Diff
162
examples/40_cutlass_py/README.md
Normal file
162
examples/40_cutlass_py/README.md
Normal file
@@ -0,0 +1,162 @@
|
||||
# CUTLASS Python Interface Example
|
||||
|
||||
## Using Docker
|
||||
You can run the PyCUTLASS on NGC pytorch container.
|
||||
```shell
|
||||
docker run --gpus all -it --rm nvcr.io/nvidia/pytorch:22.08-py3
|
||||
```
|
||||
PyCUTLASS requires additional dependency Boost C++ library, which can be installed with
|
||||
```bash
|
||||
apt-get update
|
||||
apt-get -y install libboost-all-dev
|
||||
```
|
||||
|
||||
|
||||
## Install the Python Interface
|
||||
The source code for python interface is allocated at `tools/library/script/pycutlass`. It requires two environment variables:
|
||||
* `CUTLASS_PATH`: the root directory of CUTLASS
|
||||
* `CUDA_INSTALL_PATH`: the directory where cuda toolkit is installed
|
||||
|
||||
After setting these two environment variables, PyCUTLASS can be installed with
|
||||
```shell
|
||||
cd $CUTLASS_PATH/tools/library/scripts/pycutlass && bash build.sh
|
||||
```
|
||||
***
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
### Issue 1: permission denied
|
||||
Building PyCUTLASS requires installing dependencies to python. So conda could an option if you don't have permission.
|
||||
|
||||
### Issue 2: rmm: module not found
|
||||
PyCUTLASS manages the device memory with [RMM](https://github.com/rapidsai/rmm). Our `build.sh` automatically pull the [rmm branch-22.08](https://github.com/rapidsai/rmm/tree/branch-22.08) from github and build it from source. The rmm is allocated at `$CUTLASS_PATH/tools/library/scripts/pycutlass/rmm`. It requires `cmake > 3.20.1`. If the build fails, it can be manually fixed with the following steps:
|
||||
```shell
|
||||
cd $CUTLASS_PATH/tools/library/scripts/pycutlass/rmm && ./build.sh librmm rmm
|
||||
|
||||
cd $CUTLASS_PATH/tools/library/scripts/pycutlass/rmm/python
|
||||
python setup.py build_ext --inplace
|
||||
python setup.py install
|
||||
```
|
||||
To test whether rmm is successfully installed, try `import rmm`. For other issues related to rmm, please check https://github.com/rapidsai/rmm/issues.
|
||||
|
||||
***
|
||||
For all the tests, add `--print_cuda` to print the underlying CUDA kernel. Use `-h` or `--help` to display the help message.
|
||||
## GEMM Examples
|
||||
The GEMM examples use numpy to create input tensors and verify the results.
|
||||
### GEMM F64 Example
|
||||
Example 1: SM80_Device_Gemm_f64t_f64n_f64n_tensor_op_f64_32x32x16_16x16x16
|
||||
```python
|
||||
python gemm.py -i 8 8 4 -ta float64 -tb float64 -tc float64 -tacc float64 -m multiply_add -op TensorOp -b 32 32 16 -s 4 -w 2 2 1 -cc 80 -la ColumnMajor -aa 1 -lb RowMajor -ab 1 -lc RowMajor -ac 1 -te float64 -ep LinearCombination -sw IdentitySwizzle1 -p 512 256 128 -alpha 1.0 -beta 0.5 -gm Gemm -k 1
|
||||
```
|
||||
Example 2: SM80_Device_Gemm_f64n_f64t_f64n_tensor_op_f64_64x64x16_32x32x16, split_k(2)_serial
|
||||
```python
|
||||
python gemm.py -i 8 8 4 -ta float64 -tb float64 -tc float64 -tacc float64 -m multiply_add -op TensorOp -b 64 64 16 -s 4 -w 2 2 1 -cc 80 -la RowMajor -aa 1 -lb ColumnMajor -ab 1 -lc RowMajor -ac 1 -te float64 -ep LinearCombination -sw IdentitySwizzle1 -p 512 256 128 -alpha 1.0 -beta 0.5 -gm Gemm -k 2
|
||||
```
|
||||
|
||||
### GEMM F32 Example
|
||||
Example 1: SM80_Device_Gemm_f32n_f32t_f32n_tensor_op_bf16_f32_128x128x32_64x64x32
|
||||
```python
|
||||
python gemm.py -i 16 8 8 -ta float32 -tb float32 -tc float32 -tacc float32 -m multiply_add_fast_bf16 -op TensorOp -b 128 128 32 -s 3 -w 2 2 1 -cc 80 -la RowMajor -aa 4 -lb ColumnMajor -ab 4 -lc RowMajor -ac 4 -te float32 -ep LinearCombination -sw IdentitySwizzle1 -p 512 256 128 -alpha 1.0 -beta 0.5 -gm Gemm -k 1
|
||||
```
|
||||
Example 2: SM80_Device_Gemm_f32t_f32t_f32n_tensor_op_f32_128x128x32_64x64x32, split_k(2)_parallel
|
||||
```python
|
||||
python gemm.py -i 16 8 8 -ta float32 -tb float32 -tc float32 -tacc float32 -m multiply_add -op TensorOp -b 128 128 32 -s 3 -w 2 2 1 -cc 80 -la ColumnMajor -aa 4 -lb ColumnMajor -ab 4 -lc RowMajor -ac 4 -te float32 -ep LinearCombination -sw IdentitySwizzle1 -p 512 256 128 -alpha 1.0 -beta 0.5 -gm GemmSplitKParallel -k 2
|
||||
```
|
||||
Example 3: SM80_Device_Gemm_f32t_f32t_f32n_tensor_op_fast_accurate_f32_64x64x32_32x32x32, split_k(4)_serial
|
||||
```python
|
||||
python gemm.py -i 16 8 8 -ta float32 -tb float32 -tc float32 -tacc float32 -m multiply_add_fast_f32 -op TensorOp -b 64 64 32 -s 3 -w 2 2 1 -cc 80 -la ColumnMajor -aa 4 -lb ColumnMajor -ab 4 -lc RowMajor -ac 4 -te float32 -ep LinearCombination -sw IdentitySwizzle1 -p 512 256 128 -alpha 1.0 -beta 0.5 -gm Gemm -k 4
|
||||
```
|
||||
|
||||
### GEMM F16 Example
|
||||
Example 1: SM80_Device_Gemm_f32t_f32n_f32t_tensor_op_bf16_f32_128x128x32_64x64x32
|
||||
```python
|
||||
python gemm.py -i 16 8 16 -ta float16 -tb float16 -tc float32 -tacc float32 -m multiply_add -op TensorOp -b 128 128 32 -s 3 -w 2 2 1 -cc 80 -la ColumnMajor -aa 8 -lb RowMajor -ab 8 -lc ColumnMajor -ac 4 -te float32 -ep LinearCombination -sw IdentitySwizzle4 -p 512 256 128 -alpha 1.0 -beta 0.5 -gm Gemm -k 1
|
||||
```
|
||||
Example 2: SM80_Device_Gemm_f16t_f16t_f16n_tensor_op_f32_128x128x64_64x64x64, split_k(2)_serial
|
||||
```python
|
||||
python gemm.py -i 16 8 16 -ta float16 -tb float16 -tc float16 -tacc float32 -m multiply_add -op TensorOp -b 128 128 64 -s 3 -w 2 2 1 -cc 80 -la ColumnMajor -aa 8 -lb ColumnMajor -ab 8 -lc RowMajor -ac 8 -te float32 -ep LinearCombination -sw IdentitySwizzle2 -p 512 256 128 -alpha 1.0 -beta 0.5 -gm Gemm -k 2
|
||||
```
|
||||
Example 3: SM80_Device_Gemm_f16t_f16t_f32n_tensor_op_f32_256x128x64_64x64x64, split_k(3)_serial
|
||||
```python
|
||||
python gemm.py -i 16 8 16 -ta float16 -tb float16 -tc float32 -tacc float32 -m multiply_add -op TensorOp -b 256 128 64 -s 3 -w 4 2 1 -cc 80 -la ColumnMajor -aa 8 -lb ColumnMajor -ab 8 -lc RowMajor -ac 4 -te float32 -ep LinearCombination -sw IdentitySwizzle1 -p 512 256 128 -alpha 1.0 -beta 0.5 -gm GemmSplitKParallel -k 3
|
||||
```
|
||||
|
||||
### GEMM BF16 Example
|
||||
Example 1: Device_Gemm_bf16t_bf16t_f32n_tensor_op_f32_64x128x64_32x64x64, split_k(5)_parallel
|
||||
```python
|
||||
python gemm.py -i 16 8 16 -ta bfloat16 -tb bfloat16 -tc float32 -tacc float32 -m multiply_add -op TensorOp -b 64 128 64 -s 3 -w 2 2 1 -cc 80 -la ColumnMajor -aa 8 -lb ColumnMajor -ab 8 -lc RowMajor -ac 4 -te float32 -ep LinearCombination -sw IdentitySwizzle2 -p 512 256 128 -alpha 1.0 -beta 0.5 -gm GemmSplitKParallel -k 5
|
||||
```
|
||||
|
||||
### GEMM Int8 Example
|
||||
Example 1: SM80_Device_Gemm_s8n_s8t_s8n_tensor_op_s32_256x128x128_64x64x128
|
||||
```python
|
||||
python gemm.py -i 16 8 32 -ta int8 -tb int8 -tc int8 -tacc int32 -m multiply_add -op TensorOp -b 128 128 128 -s 3 -w 2 2 1 -cc 80 -la RowMajor -aa 16 -lb ColumnMajor -ab 16 -lc RowMajor -ac 16 -te float32 -ep FastLinearCombinationClamp -sw IdentitySwizzle2 -p 512 512 512 -alpha 1.0 -beta 0.0 -gm Gemm -k 1
|
||||
```
|
||||
***
|
||||
## GEMM Grouped Examples
|
||||
The GEMM Grouped examples use numpy to create input tensors and verify the results.
|
||||
|
||||
Example 1: SM80_Device_GemmGrouped_f16t_f16t_f32t_tensor_op_f32_128x128x32_64x64x32, device schedule
|
||||
```python
|
||||
python gemm_grouped.py -i 16 8 16 -ta float16 -tb float16 -tc float32 -tacc float32 -m multiply_add -op TensorOp -b 128 128 32 -s 3 -w 2 2 1 -cc 80 -la ColumnMajor -aa 8 -lb ColumnMajor -ab 8 -lc ColumnMajor -ac 4 -te float32 -ep LinearCombination -sw IdentitySwizzle1 -p ./grouped_gemm_problem_size.csv -alpha 1.0 -beta 0.0 -pm Device
|
||||
```
|
||||
Example 2: SM80_Device_GemmGrouped_f64n_f64n_f64t_tensor_op_f64_64x64x16_32x32x16, host schedule
|
||||
```python
|
||||
python gemm_grouped.py -i 8 8 4 -ta float64 -tb float64 -tc float64 -tacc float64 -m multiply_add -op TensorOp -b 64 64 16 -s 4 -w 2 2 1 -cc 80 -la RowMajor -aa 1 -lb RowMajor -ab 1 -lc ColumnMajor -ac 1 -te float64 -ep LinearCombination -sw IdentitySwizzle2 -p ./grouped_gemm_problem_size.csv -alpha 1.0 -beta 1.0 -pm Host
|
||||
```
|
||||
Example 3: SM80_Device_GemmGrouped_f32n_f32n_f32n_simt_f32_128x64x8_64x32x1, device schedule
|
||||
```python
|
||||
python gemm_grouped.py -i 1 1 1 -ta float32 -tb float32 -tc float32 -tacc float32 -m multiply_add -op Simt -b 128 64 8 -s 4 -w 2 2 1 -cc 80 -la RowMajor -aa 1 -lb RowMajor -ab 1 -lc RowMajor -ac 1 -te float32 -ep LinearCombination -sw IdentitySwizzle4 -p ./grouped_gemm_problem_size.csv -alpha 2.0 -beta 1.0 -pm Device
|
||||
```
|
||||
Example 4: SM80_Device_GemmGrouped_f16t_f16t_f32t_tensor_op_f32_128x128x32_64x64x32, device schedule
|
||||
```python
|
||||
python gemm_grouped.py -i 16 8 16 -ta float16 -tb float16 -tc float32 -tacc float32 -m multiply_add -op TensorOp -b 128 128 32 -s 3 -w 2 2 1 -cc 80 -la ColumnMajor -aa 8 -lb ColumnMajor -ab 8 -lc ColumnMajor -ac 4 -te float32 -ep LinearCombination -sw IdentitySwizzle8 -p ./grouped_gemm_problem_size.csv -alpha 2.0 -beta 1.0 -pm Device
|
||||
```
|
||||
***
|
||||
## Conv2d Example
|
||||
The Conv2d examples use pytorch to create input tensors and verify the results. Pytorch can be installed following the [official website](https://pytorch.org/#:~:text=Aid%20to%20Ukraine.-,INSTALL%20PYTORCH,-Select%20your%20preferences).
|
||||
### Conv2d F32 Fprop
|
||||
Example 1: SM80_Device_Conv2d_Fprop_Analytic_ImplicitGemm_tf32nhwc_tf32nhwc_f32nhwc_tensor_op_f32
|
||||
```python
|
||||
python conv2d.py -i 16 8 8 -ta float32 -tb float32 -tc float32 -tacc float32 -m multiply_add -op TensorOp -b 128 128 16 -s 3 -w 2 2 1 -cc 80 -la TensorNHWC -aa 4 -lb TensorNHWC -ab 4 -lc TensorNHWC -ac 4 -te float32 -ep LinearCombination -sw IdentitySwizzle1 -co fprop -st Strided -ia optimized -sm Serial -k 1 -nhwc 1 13 17 8 -krsc 24 3 3 8 -pad 0 0 0 0 -stride 2 2 -dilation 1 1 -alpha 1.0 -beta 0.0
|
||||
```
|
||||
Example 2: SM80_Device_Conv2d_Fprop_Optimized_ImplicitGemm_tf32nhwc_tf32nhwc_f32nhwc_tensor_op_f32_align2
|
||||
```python
|
||||
python conv2d.py -i 16 8 8 -ta float32 -tb float32 -tc float32 -tacc float32 -m multiply_add -op TensorOp -b 128 128 16 -s 3 -w 2 2 1 -cc 80 -la TensorNHWC -aa 2 -lb TensorNHWC -ab 2 -lc TensorNHWC -ac 2 -te float32 -ep LinearCombination -sw IdentitySwizzle2 -co fprop -st Strided -ia optimized -sm Serial -k 2 -nhwc 1 4 4 12 -krsc 8 3 3 12 -pad 0 0 0 0 -stride 3 3 -dilation 1 1 -alpha 1.0 -beta 1.0
|
||||
```
|
||||
Example 3: SM80_Device_Conv2d_Fprop_Analytic_ImplicitGemm_f32nhwc_f32nhwc_f32nhwc_simt_f32
|
||||
```python
|
||||
python conv2d.py -i 1 1 1 -ta float32 -tb float32 -tc float32 -tacc float32 -m multiply_add -op Simt -b 128 128 8 -s 4 -w 4 2 1 -cc 80 -la TensorNHWC -aa 4 -lb TensorNHWC -ab 4 -lc TensorNHWC -ac 1 -te float32 -ep LinearCombination -sw IdentitySwizzle4 -co fprop -st Strided -ia analytic -sm Parallel -k 3 -nhwc 1 71 80 32 -krsc 64 5 5 32 -pad 2 2 2 2 -stride 2 2 -dilation 1 1 -alpha 1.0 -beta 1.0
|
||||
```
|
||||
### Conv2d F32 Wgrad
|
||||
Example 1: Device_Conv2d_Wgrad_Optimized_ImplicitGemm_tf32nhwc_tf32nhwc_f32nhwc_tensor_op_f32_align1
|
||||
```python
|
||||
python conv2d.py -i 16 8 8 -ta float32 -tb float32 -tc float32 -tacc float32 -m multiply_add -op TensorOp -b 128 128 32 -s 3 -w 2 2 1 -cc 80 -la TensorNHWC -aa 1 -lb TensorNHWC -ab 1 -lc TensorNHWC -ac 4 -te float32 -ep LinearCombination -sw IdentitySwizzle1 -co wgrad -st Strided -ia optimized -sm Serial -k 1 -nhwc 1 8 8 1 -krsc 1 3 3 1 -pad 1 1 1 1 -stride 1 1 -dilation 1 1 -alpha 1.0 -beta 0.0
|
||||
```
|
||||
Example 2: Device_Conv2d_Wgrad_Analytic_ImplicitGemm_f32nhwc_f32nhwc_f32nhwc_simt_f32
|
||||
```python
|
||||
python conv2d.py -i 1 1 1 -ta float32 -tb float32 -tc float32 -tacc float32 -m multiply_add -op Simt -b 128 128 8 -s 4 -w 2 4 1 -cc 80 -la TensorNHWC -aa 4 -lb TensorNHWC -ab 4 -lc TensorNHWC -ac 1 -te float32 -ep LinearCombination -sw IdentitySwizzle1 -co wgrad -st Strided -ia optimized -sm Serial -k 2 -nhwc 1 27 27 256 -krsc 512 3 3 256 -pad 1 1 1 1 -stride 2 1 -dilation 1 1 -alpha 1.0 -beta 0.0
|
||||
```
|
||||
### Conv2d F32 Dgrad
|
||||
Example 1: Device_Conv2d_Dgrad_Analytic_ImplicitGemm_tf32nhwc_tf32nhwc_f32nhwc_tensor_op_f32
|
||||
```python
|
||||
python conv2d.py -i 16 8 8 -ta float32 -tb float32 -tc float32 -tacc float32 -m multiply_add -op TensorOp -b 128 128 16 -s 3 -w 2 2 1 -cc 80 -la TensorNHWC -aa 4 -lb TensorNHWC -ab 4 -lc TensorNHWC -ac 4 -te float32 -ep LinearCombination -sw StridedDgradIdentitySwizzle1 -co dgrad -st Strided -ia optimized -sm Serial -k 2 -nhwc 1 27 27 256 -krsc 512 3 3 256 -pad 1 1 1 1 -stride 2 1 -dilation 1 1 -alpha 1.0 -beta 0.0
|
||||
```
|
||||
|
||||
### Conv2d F16 Fprop
|
||||
Example 1: SM80_Device_Conv2d_Fprop_Analytic_ImplicitGemm_f16nhwc_f16nhwc_f16nhwc_tensor_op_f32
|
||||
```python
|
||||
python conv2d.py -i 16 8 16 -ta float16 -tb float16 -tc float16 -tacc float32 -m multiply_add -op TensorOp -b 128 128 64 -s 3 -w 2 2 1 -cc 80 -la TensorNHWC -aa 8 -lb TensorNHWC -ab 8 -lc TensorNHWC -ac 8 -te float32 -ep LinearCombination -sw IdentitySwizzle1 -co fprop -st Strided -ia optimized -sm Serial -k 1 -nhwc 1 27 27 256 -krsc 512 3 3 256 -pad 1 1 1 1 -stride 2 1 -dilation 1 1 -alpha 1.0 -beta 0.0
|
||||
```
|
||||
Example 2: SM80_Device_Conv2d_Fprop_Few_Channels_ImplicitGemm_f16nhwc_f16nhwc_f16nhwc_tensor_op_f32_channels_2
|
||||
```python
|
||||
python conv2d.py -i 16 8 16 -ta float16 -tb float16 -tc float16 -tacc float32 -m multiply_add -op TensorOp -b 128 128 64 -s 3 -w 2 2 1 -cc 80 -la TensorNHWC -aa 2 -lb TensorNHWC -ab 2 -lc TensorNHWC -ac 8 -te float32 -ep LinearCombination -sw IdentitySwizzle1 -co fprop -st Strided -ia few_channels -sm Serial -k 1 -nhwc 1 16 16 2 -krsc 16 3 3 2 -pad 1 1 1 1 -stride 2 2 -dilation 1 1 -alpha 1.0 -beta 0.0
|
||||
```
|
||||
Example 3: SM80_Device_Conv2d_Fprop_Fixed_Channels_ImplicitGemm_f16nhwc_f16nhwc_f16nhwc_tensor_op_f32_channels_8
|
||||
```python
|
||||
python conv2d.py -i 16 8 16 -ta float16 -tb float16 -tc float16 -tacc float32 -m multiply_add -op TensorOp -b 128 128 64 -s 3 -w 2 2 1 -cc 80 -la TensorNHWC -aa 8 -lb TensorNHWC -ab 8 -lc TensorNHWC -ac 8 -te float32 -ep LinearCombination -sw IdentitySwizzle2 -co fprop -st Strided -ia fixed_channels -sm Serial -k 1 -nhwc 1 8 8 8 -krsc 16 3 3 8 -pad 1 1 1 1 -stride 2 2 -dilation 1 1 -alpha 1.0 -beta 0.0
|
||||
```
|
||||
Example 4: SM80_Device_Conv2d_Strided_Dgrad_Optimized_ImplicitGemm_f16nhwc_f16nhwc_f32nhwc_tensor_op_f32_128x128_32x3_64x64x32_align4
|
||||
```python
|
||||
python conv2d.py -i 16 8 16 -ta float16 -tb float16 -tc float16 -tacc float32 -m multiply_add -op TensorOp -b 128 128 32 -s 3 -w 2 2 1 -cc 80 -la TensorNHWC -aa 4 -lb TensorNHWC -ab 4 -lc TensorNHWC -ac 4 -te float32 -ep LinearCombination -sw StridedDgradIdentitySwizzle1 -co dgrad -st Strided -ia optimized -sm Serial -k 1 -nhwc 1 56 56 12 -krsc 8 1 1 12 -pad 0 0 0 0 -stride 2 2 -dilation 1 1 -alpha 1.0 -beta 0.0
|
||||
```
|
||||
277
examples/40_cutlass_py/conv2d.py
Normal file
277
examples/40_cutlass_py/conv2d.py
Normal file
@@ -0,0 +1,277 @@
|
||||
################################################################################
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
################################################################################
|
||||
import pycutlass
|
||||
from pycutlass import *
|
||||
from pycutlass.conv2d_operation import *
|
||||
from pycutlass.utils import reference_model
|
||||
|
||||
import argparse
|
||||
|
||||
# parse the arguments
|
||||
parser = argparse.ArgumentParser(description="Launch CUTLASS convolution 2d kernels from python")
|
||||
|
||||
# Operation description
|
||||
# math instruction description
|
||||
parser.add_argument("-i", "--instruction_shape",
|
||||
default=[1, 1, 1], nargs=3, type=int,
|
||||
help="This option describes the size of MMA op")
|
||||
parser.add_argument("-ta", "--element_a", default="float32", type=str,
|
||||
choices=['float64', 'float32', 'float16', 'bfloat16', 'int32', 'int8'],
|
||||
help='Data type of elements in input tensor A')
|
||||
parser.add_argument("-tb", "--element_b", default="float32", type=str,
|
||||
choices=['float64', 'float32', 'float16', 'bfloat16', 'int32', 'int8'],
|
||||
help='Data type of elements in input tensor B')
|
||||
parser.add_argument("-tc", "--element_c", default="float32", type=str,
|
||||
choices=['float64', 'float32', 'float16', 'bfloat16', 'int32', 'int8'],
|
||||
help='Data type of elements in input tensor C and output tensor D')
|
||||
parser.add_argument("-tacc", "--element_acc", default="float32", type=str,
|
||||
choices=['float64', 'float32', 'float16', 'bfloat16', 'int32', 'int8'],
|
||||
help='Data type of accumulator')
|
||||
parser.add_argument('-m', "--math", default="multiply_add",
|
||||
type=str, choices=["multiply_add", "multiply_add_fast_bf16", "multiply_add_fast_f32"], help="math instruction")
|
||||
parser.add_argument('-op', "--opcode", default="simt", type=str,
|
||||
choices=["Simt", 'TensorOp'],
|
||||
help='This option describes whether you want to use tensor \
|
||||
cores (TensorOp) or regular SIMT cores (Simt) on GPU SM')
|
||||
# tile description
|
||||
parser.add_argument("-b", "--threadblock_shape",
|
||||
default=[128, 128, 8], nargs=3, type=int,
|
||||
help="This option describes the tile size a thread block with compute")
|
||||
parser.add_argument("-s", "--stages", default=4,
|
||||
type=int, help="Number of pipelines you want to use")
|
||||
parser.add_argument("-w", "--warp_count", default=[
|
||||
4, 2, 1], nargs=3, type=int,
|
||||
help="This option describes the number of warps along M, N, and K of the threadblock")
|
||||
parser.add_argument("-cc", "--compute_capability", default=80,
|
||||
type=int, help="This option describes CUDA SM architecture number")
|
||||
# A
|
||||
parser.add_argument('-la', "--layout_a", default="TensorNHWC", type=str, choices=[
|
||||
"TensorNHWC", "TensorNC32HW32"],
|
||||
help="Memory layout of input tensor A")
|
||||
parser.add_argument('-aa', '--alignment_a', default=1,
|
||||
type=int, help="Memory alignement of input tensor A")
|
||||
# B
|
||||
parser.add_argument('-lb', "--layout_b", default="TensorNHWC", type=str, choices=[
|
||||
"TensorNHWC", "TensorC32RSK32"],
|
||||
help="Memory layout of input tensor B")
|
||||
parser.add_argument('-ab', '--alignment_b', default=1,
|
||||
type=int, help="Memory alignment of input tensor B")
|
||||
# C
|
||||
parser.add_argument('-lc', "--layout_c", default="TensorNHWC", type=str, choices=[
|
||||
"TensorNHWC", "TensorNC32HW32"],
|
||||
help="Memory layout of input tensor C and output tensor D")
|
||||
parser.add_argument('-ac', '--alignment_c', default=1,
|
||||
type=int, help="Memory alignment of input tensor C and output tensor D")
|
||||
# epilogue
|
||||
parser.add_argument("-te", "--element_epilogue", default="float32", type=str,
|
||||
choices=['float64', 'float32', 'float16', 'bfloat16'],
|
||||
help='Data type of computation in the epilogue')
|
||||
parser.add_argument("-ep", "--epilogue_functor", default="LinearCombination",
|
||||
type=str, choices=['LinearCombination', 'FastLinearCombinationClamp', 'LinearCombinationClamp'],
|
||||
help="This option describes the epilogue part of the kernel")
|
||||
# swizzling
|
||||
parser.add_argument("-sw", "--swizzling_functor", default="IdentitySwizzle1", type=str, choices=[
|
||||
"IdentitySwizzle1", "IdentitySwizzle2", "IdentitySwizzle4", "IdentitySwizzle8",
|
||||
"HorizontalSwizzle", "StridedDgradIdentitySwizzle1", "StridedDgradIdentitySwizzle4",
|
||||
"StridedDgradHorizontalSwizzle"],
|
||||
help="This option describes how thread blocks are scheduled on GPU")
|
||||
# conv related
|
||||
parser.add_argument("-co", "--conv_kind", default="fprop", type=str, choices=['fprop', 'dgrad', 'wgrad'],
|
||||
help="The type of convolution: forward propagation (fprop), \
|
||||
gradient of activation (dgrad), gradient of weight (wgrad)")
|
||||
parser.add_argument("-st", "--stride_support", default="Strided", type=str, choices=["Strided", "Unity"],
|
||||
)
|
||||
parser.add_argument("-ia", "--iterator_algorithm", default="analytic", type=str,
|
||||
choices=["analytic", "optimized", "fixed_channels", "few_channels"],
|
||||
help="This option describes iterator algorithm")
|
||||
|
||||
# arguments
|
||||
parser.add_argument("-sm", "--split_k_mode", default="Serial", type=str, choices=["Serial", "Parallel"],
|
||||
help="Split K Mode. Serial is used for non-splitK or serial-splitK.\
|
||||
Parallel is used for parallel splitK.")
|
||||
parser.add_argument('-k', '--split_k_slices', default=1,
|
||||
type=int, help="Number of split-k partitions. (default 1)")
|
||||
parser.add_argument("-nhwc", "--nhwc", nargs=4, type=int, help="input size (NHWC)")
|
||||
parser.add_argument("-krsc", "--krsc", nargs=4, type=int, help="filter size (KRSC)")
|
||||
parser.add_argument("-pad", "--pad", nargs=4, type=int, help="padding (pad_h, _, pad_w, _)")
|
||||
parser.add_argument("-stride", "--stride", nargs=2, type=int, help="stride (stride_h, stride_w)")
|
||||
parser.add_argument("-dilation", "--dilation", nargs=2, type=int, help="dilation (dilation_h, dilation_w)")
|
||||
parser.add_argument("-alpha", "--alpha", default=1.0, type=float, help="alpha")
|
||||
parser.add_argument("-beta", "--beta", default=0.0, type=float, help="beta")
|
||||
|
||||
parser.add_argument('--print_cuda', action="store_true",
|
||||
help="print the underlying CUDA kernel")
|
||||
|
||||
try:
|
||||
args = parser.parse_args()
|
||||
except:
|
||||
sys.exit(0)
|
||||
|
||||
pycutlass.get_memory_pool(init_pool_size=2**30, max_pool_size=2**32)
|
||||
|
||||
element_a = getattr(cutlass, args.element_a)
|
||||
element_b = getattr(cutlass, args.element_b)
|
||||
element_c = getattr(cutlass, args.element_c)
|
||||
element_acc = getattr(cutlass, args.element_acc)
|
||||
math_operation = getattr(MathOperation, args.math)
|
||||
opclass = getattr(cutlass.OpClass, args.opcode)
|
||||
|
||||
math_inst = MathInstruction(
|
||||
args.instruction_shape, element_a, element_b,
|
||||
element_acc, opclass, math_operation
|
||||
)
|
||||
|
||||
tile_description = TileDescription(
|
||||
args.threadblock_shape, args.stages, args.warp_count,
|
||||
math_inst, args.compute_capability, args.compute_capability
|
||||
)
|
||||
|
||||
layout_a = getattr(cutlass, args.layout_a)
|
||||
layout_b = getattr(cutlass, args.layout_b)
|
||||
layout_c = getattr(cutlass, args.layout_c)
|
||||
|
||||
A = TensorDescription(
|
||||
element_a, layout_a, args.alignment_a
|
||||
)
|
||||
|
||||
B = TensorDescription(
|
||||
element_b, layout_b, args.alignment_b
|
||||
)
|
||||
|
||||
C = TensorDescription(
|
||||
element_c, layout_c, args.alignment_c
|
||||
)
|
||||
|
||||
element_epilogue = getattr(cutlass, args.element_epilogue)
|
||||
epilogue_functor = getattr(EpilogueFunctor, args.epilogue_functor)
|
||||
iterator_algorithm = getattr(cutlass.conv.IteratorAlgorithm, args.iterator_algorithm)
|
||||
swizzling_functor = getattr(cutlass, args.swizzling_functor)
|
||||
stride_support = getattr(StrideSupport, args.stride_support)
|
||||
conv_kind = getattr(cutlass.conv.Operator, args.conv_kind)
|
||||
|
||||
operation = Conv2dOperation(
|
||||
conv_kind=conv_kind, iterator_algorithm=iterator_algorithm,
|
||||
arch=args.compute_capability, tile_description=tile_description,
|
||||
A=A, B=B, C=C, element_epilogue=element_epilogue, stride_support=stride_support,
|
||||
epilogue_functor=epilogue_functor, swizzling_functor=swizzling_functor
|
||||
)
|
||||
|
||||
if args.print_cuda:
|
||||
print(operation.rt_module.emit())
|
||||
|
||||
operations = [operation,]
|
||||
|
||||
if args.split_k_mode == "Parallel" and args.split_k_slices > 1:
|
||||
reduction_operation = ReductionOperation(
|
||||
shape=cutlass.MatrixCoord(4, 32 * C.alignment),
|
||||
C=C, element_accumulator=element_acc,
|
||||
element_compute=element_epilogue,
|
||||
count=C.alignment
|
||||
)
|
||||
operations.append(reduction_operation)
|
||||
|
||||
pycutlass.compiler.add_module(operations)
|
||||
|
||||
problem_size = cutlass.conv.Conv2dProblemSize(
|
||||
cutlass.Tensor4DCoord(args.nhwc[0], args.nhwc[1], args.nhwc[2], args.nhwc[3]),
|
||||
cutlass.Tensor4DCoord(args.krsc[0], args.krsc[1], args.krsc[2], args.krsc[3]),
|
||||
cutlass.Tensor4DCoord(args.pad[0], args.pad[1], args.pad[2], args.pad[3]),
|
||||
cutlass.MatrixCoord(args.stride[0], args.stride[1]),
|
||||
cutlass.MatrixCoord(args.dilation[0], args.dilation[1]),
|
||||
cutlass.conv.Mode.cross_correlation,
|
||||
args.split_k_slices, 1
|
||||
)
|
||||
|
||||
|
||||
# User-provide inputs
|
||||
tensor_A_size = cutlass.conv.implicit_gemm_tensor_a_size(
|
||||
conv_kind, problem_size
|
||||
)
|
||||
tensor_B_size = cutlass.conv.implicit_gemm_tensor_b_size(
|
||||
conv_kind, problem_size
|
||||
)
|
||||
tensor_C_size = cutlass.conv.implicit_gemm_tensor_c_size(
|
||||
conv_kind, problem_size
|
||||
)
|
||||
|
||||
if args.element_a != "int8":
|
||||
tensor_A = torch.ceil(torch.empty(size=(tensor_A_size,), dtype=getattr(torch, args.element_a), device="cuda").uniform_(-8.5, 7.5))
|
||||
else:
|
||||
tensor_A = torch.empty(size=(tensor_A_size,), dtype=getattr(torch, args.element_a), device="cuda").uniform_(-2, 2)
|
||||
|
||||
if args.element_b != "int8":
|
||||
tensor_B = torch.ceil(torch.empty(size=(tensor_B_size,), dtype=getattr(torch, args.element_b), device="cuda").uniform_(-8.5, 7.5))
|
||||
else:
|
||||
tensor_B = torch.empty(size=(tensor_B_size,), dtype=getattr(torch, args.element_b), device="cuda").uniform_(-2, 2)
|
||||
|
||||
if args.element_c != "int8":
|
||||
tensor_C = torch.ceil(torch.empty(size=(tensor_C_size,), dtype=getattr(torch, args.element_c), device="cuda").uniform_(-8.5, 7.5))
|
||||
else:
|
||||
tensor_C = torch.empty(size=(tensor_C_size,), dtype=getattr(torch, args.element_c), device="cuda").uniform_(-2, 2)
|
||||
|
||||
tensor_D = torch.ones_like(tensor_C)
|
||||
|
||||
arguments = Conv2dArguments(
|
||||
operation=operation, problem_size=problem_size, A=tensor_A,
|
||||
B=tensor_B, C=tensor_C, D=tensor_D,
|
||||
output_op = LinearCombinationFunctorArguments(args.alpha, args.beta),
|
||||
split_k_mode=getattr(cutlass.conv.SplitKMode, args.split_k_mode),
|
||||
split_k_slices=problem_size.split_k_slices
|
||||
)
|
||||
|
||||
if args.split_k_mode == "Parallel" and args.split_k_slices > 1:
|
||||
implicit_gemm_size = cutlass.conv.implicit_gemm_problem_size(conv_kind, arguments.problem_size)
|
||||
reduction_arguments = ReductionArguments(
|
||||
reduction_operation,
|
||||
problem_size=[implicit_gemm_size.m(), implicit_gemm_size.n()],
|
||||
partitions=problem_size.split_k_slices,
|
||||
workspace=arguments.ptr_D,
|
||||
destination=tensor_D,
|
||||
source=tensor_C,
|
||||
output_op = LinearCombinationFunctorArguments(args.alpha, args.beta)
|
||||
)
|
||||
|
||||
operation.run(arguments)
|
||||
|
||||
if args.split_k_mode == "Parallel" and args.split_k_slices > 1:
|
||||
reduction_operation.run(reduction_arguments)
|
||||
reduction_arguments.sync()
|
||||
else:
|
||||
arguments.sync()
|
||||
|
||||
reference_model = Conv2dReferenceModule(A, B, C, conv_kind)
|
||||
|
||||
tensor_D_ref = reference_model.run(tensor_A, tensor_B, tensor_C, arguments.problem_size, args.alpha, args.beta)
|
||||
|
||||
assert torch.equal(tensor_D, tensor_D_ref)
|
||||
|
||||
print("Passed.")
|
||||
266
examples/40_cutlass_py/gemm.py
Normal file
266
examples/40_cutlass_py/gemm.py
Normal file
@@ -0,0 +1,266 @@
|
||||
################################################################################
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
################################################################################
|
||||
import numpy as np
|
||||
import pycutlass
|
||||
from pycutlass import *
|
||||
import cutlass
|
||||
from bfloat16 import bfloat16
|
||||
|
||||
import argparse
|
||||
|
||||
|
||||
# parse the arguments
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Launch CUTLASS GEMM kernels from python: 'D = alpha * A * B + beta * C'")
|
||||
|
||||
# Operation description
|
||||
# math instruction description
|
||||
parser.add_argument("-i", "--instruction_shape",
|
||||
default=[1, 1, 1], nargs=3, type=int,
|
||||
help="This option describes the size of MMA op")
|
||||
parser.add_argument("-ta", "--element_a", default="float32", type=str,
|
||||
choices=['float64', 'float32', 'float16', 'bfloat16', 'int32', 'int8'],
|
||||
help='Data type of elements in input tensor A')
|
||||
parser.add_argument("-tb", "--element_b", default="float32", type=str,
|
||||
choices=['float64', 'float32', 'float16', 'bfloat16', 'int32', 'int8'],
|
||||
help='Data type of elements in input tensor B')
|
||||
parser.add_argument("-tc", "--element_c", default="float32", type=str,
|
||||
choices=['float64', 'float32', 'float16', 'bfloat16', 'int32', 'int8'],
|
||||
help='Data type of elements in input tensor C and output tensor D')
|
||||
parser.add_argument("-tacc", "--element_acc", default="float32", type=str,
|
||||
choices=['float64', 'float32', 'float16', 'bfloat16', 'int32', 'int8'],
|
||||
help='Data type of accumulator')
|
||||
parser.add_argument('-m', "--math", default="multiply_add",
|
||||
type=str, choices=["multiply_add", "multiply_add_fast_bf16", "multiply_add_fast_f32"], help="math instruction")
|
||||
parser.add_argument('-op', "--opcode", default="simt", type=str,
|
||||
choices=["Simt", 'TensorOp'],
|
||||
help="This option describes whether you want to use tensor \
|
||||
cores (TensorOp) or regular SIMT cores (Simt) on GPU SM")
|
||||
# tile description
|
||||
parser.add_argument("-b", "--threadblock_shape",
|
||||
default=[128, 128, 8], nargs=3, type=int,
|
||||
help="This option describes the tile size a thread block with compute")
|
||||
parser.add_argument("-s", "--stages", default=4,
|
||||
type=int, help="Number of pipelines you want to use")
|
||||
parser.add_argument("-w", "--warp_count", default=[4, 2, 1], nargs=3, type=int,
|
||||
help="This option describes the number of warps along M, N, and K of the threadblock")
|
||||
parser.add_argument("-cc", "--compute_capability", default=80,
|
||||
type=int, help="This option describes CUDA SM architecture number")
|
||||
# A
|
||||
parser.add_argument('-la', "--layout_a", default="RowMajor", type=str, choices=[
|
||||
"RowMajor", "ColumnMajor", "RowMajorInterleaved32", "ColumnMajorInterleaved32"],
|
||||
help="Memory layout of input tensor A")
|
||||
parser.add_argument('-aa', '--alignment_a', default=1,
|
||||
type=int, help="Memory alignement of input tensor A")
|
||||
# B
|
||||
parser.add_argument('-lb', "--layout_b", default="RowMajor", type=str, choices=[
|
||||
"RowMajor", "ColumnMajor", "RowMajorInterleaved32", "ColumnMajorInterleaved32"],
|
||||
help="Memory layout of input tensor B")
|
||||
parser.add_argument('-ab', '--alignment_b', default=1,
|
||||
type=int, help="Memory alignment of input tensor B")
|
||||
# C
|
||||
parser.add_argument('-lc', "--layout_c", default="RowMajor", type=str, choices=[
|
||||
"RowMajor", "ColumnMajor", "RowMajorInterleaved32", "ColumnMajorInterleaved32"],
|
||||
help="Memory layout of input tensor C and output tensor D")
|
||||
parser.add_argument('-ac', '--alignment_c', default=1,
|
||||
type=int, help="Memory alignment of input tensor C and output tensor D")
|
||||
# epilogue
|
||||
parser.add_argument("-te", "--element_epilogue", default="float32", type=str,
|
||||
choices=['float64', 'float32', 'float16', 'bfloat16'], help='Epilogue datatype')
|
||||
parser.add_argument("-ep", "--epilogue_functor", default="LinearCombination",
|
||||
type=str, choices=['LinearCombination', 'FastLinearCombinationClamp', 'LinearCombinationClamp'],
|
||||
help="This option describes the epilogue part of the kernel")
|
||||
# swizzling
|
||||
parser.add_argument("-sw", "--swizzling_functor", default="IdentitySwizzle1", type=str, choices=[
|
||||
"IdentitySwizzle1", "IdentitySwizzle2", "IdentitySwizzle4", "IdentitySwizzle8", "HorizontalSwizzle"],
|
||||
help="This option describes how thread blocks are scheduled on GPU")
|
||||
|
||||
# Argument
|
||||
parser.add_argument("-p", "--problem_size",
|
||||
default=[128, 128, 128], nargs=3, type=int,
|
||||
help="GEMM problem size M, N, K")
|
||||
parser.add_argument("-alpha", "--alpha", default=1.0, type=float,
|
||||
help="Scaling factor of A * B")
|
||||
parser.add_argument("-beta", "--beta", default=0.0, type=float,
|
||||
help="Scaling factor of C")
|
||||
parser.add_argument("-gm", "--gemm_mode", default="Gemm", type=str,
|
||||
choices=["Gemm", "GemmSplitKParallel"],
|
||||
help="GEMM mode. Gemm is used for non-splitK or serial-splitK. \
|
||||
GemmSplitKParallel is used for parallel splitK")
|
||||
parser.add_argument('-k', '--split_k_slices', default=1,
|
||||
type=int, help="Number of split-k partitions. (default 1)")
|
||||
|
||||
parser.add_argument('--print_cuda', action="store_true",
|
||||
help="print the underlying CUDA kernel")
|
||||
|
||||
# parser.add_argument('-h', '--help', action="store_true",
|
||||
# help="print help information")
|
||||
|
||||
try:
|
||||
args = parser.parse_args()
|
||||
except:
|
||||
sys.exit(0)
|
||||
|
||||
pycutlass.get_memory_pool(init_pool_size=2**30, max_pool_size=2**32)
|
||||
|
||||
element_a = getattr(cutlass, args.element_a)
|
||||
element_b = getattr(cutlass, args.element_b)
|
||||
element_c = getattr(cutlass, args.element_c)
|
||||
element_acc = getattr(cutlass, args.element_acc)
|
||||
math_operation = getattr(MathOperation, args.math)
|
||||
opclass = getattr(cutlass.OpClass, args.opcode)
|
||||
|
||||
math_inst = MathInstruction(
|
||||
args.instruction_shape, element_a, element_b,
|
||||
element_acc, opclass, math_operation
|
||||
)
|
||||
|
||||
tile_description = TileDescription(
|
||||
args.threadblock_shape, args.stages, args.warp_count,
|
||||
math_inst, args.compute_capability, args.compute_capability
|
||||
)
|
||||
|
||||
layout_a = getattr(cutlass, args.layout_a)
|
||||
layout_b = getattr(cutlass, args.layout_b)
|
||||
layout_c = getattr(cutlass, args.layout_c)
|
||||
|
||||
A = TensorDescription(
|
||||
element_a, layout_a, args.alignment_a
|
||||
)
|
||||
|
||||
B = TensorDescription(
|
||||
element_b, layout_b, args.alignment_b
|
||||
)
|
||||
|
||||
C = TensorDescription(
|
||||
element_c, layout_c, args.alignment_c
|
||||
)
|
||||
|
||||
element_epilogue = getattr(cutlass, args.element_epilogue)
|
||||
epilogue_functor = getattr(EpilogueFunctor, args.epilogue_functor)
|
||||
swizzling_functor = getattr(cutlass, args.swizzling_functor)
|
||||
|
||||
operation = GemmOperationUniversal(
|
||||
arch=args.compute_capability, tile_description=tile_description,
|
||||
A=A, B=B, C=C, element_epilogue=element_epilogue,
|
||||
epilogue_functor=epilogue_functor, swizzling_functor=swizzling_functor
|
||||
)
|
||||
|
||||
if args.print_cuda:
|
||||
print(operation.rt_module.emit())
|
||||
|
||||
operations = [operation, ]
|
||||
|
||||
if args.gemm_mode == "GemmSplitKParallel":
|
||||
reduction_operation = ReductionOperation(
|
||||
shape=cutlass.MatrixCoord(4, 32 * C.alignment),
|
||||
C=C, element_accumulator=element_acc,
|
||||
element_compute=element_epilogue,
|
||||
count=C.alignment
|
||||
)
|
||||
operations.append(reduction_operation)
|
||||
|
||||
pycutlass.compiler.add_module(operations)
|
||||
|
||||
# User-provide inputs
|
||||
|
||||
problem_size = cutlass.gemm.GemmCoord(
|
||||
args.problem_size[0], args.problem_size[1], args.problem_size[2])
|
||||
|
||||
if args.element_a != "int8":
|
||||
if args.element_a == "bfloat16":
|
||||
tensor_A = np.ceil(np.random.uniform(low=-8.5, high=7.5, size=(problem_size.m()
|
||||
* problem_size.k(),))).astype(bfloat16)
|
||||
else:
|
||||
tensor_A = np.ceil(np.random.uniform(low=-8.5, high=7.5, size=(problem_size.m()
|
||||
* problem_size.k(),))).astype(getattr(np, args.element_a))
|
||||
else:
|
||||
tensor_A = np.random.uniform(low=-2, high=2, size=(problem_size.m()
|
||||
* problem_size.k(),)).astype(getattr(np, args.element_a))
|
||||
|
||||
if args.element_b != "int8":
|
||||
if args.element_b == "bfloat16":
|
||||
tensor_B = np.ceil(np.random.uniform(low=-8.5, high=7.5, size=(problem_size.k()
|
||||
* problem_size.n(),))).astype(bfloat16)
|
||||
else:
|
||||
tensor_B = np.ceil(np.random.uniform(low=-8.5, high=7.5, size=(problem_size.k()
|
||||
* problem_size.n(),))).astype(getattr(np, args.element_b))
|
||||
else:
|
||||
tensor_B = np.random.uniform(low=-2, high=2, size=(problem_size.k()
|
||||
* problem_size.n(),)).astype(getattr(np, args.element_b))
|
||||
|
||||
if args.element_c != "int8":
|
||||
if args.element_c == "bfloat16":
|
||||
tensor_C = np.ceil(np.random.uniform(low=-8.5, high=7.5, size=(problem_size.m()
|
||||
* problem_size.n(),))).astype(bfloat16)
|
||||
else:
|
||||
tensor_C = np.ceil(np.random.uniform(low=-8.5, high=7.5, size=(problem_size.m()
|
||||
* problem_size.n(),))).astype(getattr(np, args.element_c))
|
||||
else:
|
||||
tensor_C = np.random.uniform(low=-2, high=2, size=(problem_size.m()
|
||||
* problem_size.n(),)).astype(getattr(np, args.element_c))
|
||||
|
||||
tensor_D = np.ones_like(tensor_C)
|
||||
|
||||
arguments = GemmArguments(
|
||||
operation=operation, problem_size=problem_size,
|
||||
A=tensor_A, B=tensor_B, C=tensor_C, D=tensor_D,
|
||||
output_op=LinearCombinationFunctorArguments(args.alpha, args.beta),
|
||||
gemm_mode=getattr(cutlass.gemm.Mode, args.gemm_mode),
|
||||
split_k_slices=args.split_k_slices
|
||||
)
|
||||
|
||||
if args.gemm_mode == "GemmSplitKParallel":
|
||||
reduction_arguments = ReductionArguments(
|
||||
operation=reduction_operation,
|
||||
problem_size=[problem_size.m(), problem_size.n()],
|
||||
partitions=args.split_k_slices, workspace=arguments.ptr_D,
|
||||
destination=tensor_D, source=tensor_C,
|
||||
output_op=LinearCombinationFunctorArguments(args.alpha, args.beta)
|
||||
)
|
||||
|
||||
operation.run(arguments)
|
||||
|
||||
if args.gemm_mode == "GemmSplitKParallel":
|
||||
reduction_operation.run(reduction_arguments)
|
||||
reduction_arguments.sync()
|
||||
else:
|
||||
arguments.sync()
|
||||
|
||||
# run the host reference module
|
||||
reference = ReferenceModule(A, B, C)
|
||||
tensor_D_ref = reference.run(
|
||||
tensor_A, tensor_B, tensor_C, problem_size, args.alpha, args.beta)
|
||||
|
||||
assert np.array_equal(tensor_D, tensor_D_ref)
|
||||
|
||||
print("Passed.")
|
||||
248
examples/40_cutlass_py/gemm_grouped.py
Normal file
248
examples/40_cutlass_py/gemm_grouped.py
Normal file
@@ -0,0 +1,248 @@
|
||||
################################################################################
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
################################################################################
|
||||
import pycutlass
|
||||
from pycutlass import *
|
||||
import csv
|
||||
|
||||
import argparse
|
||||
|
||||
# parse the arguments
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Launch CUTLASS GEMM Grouped kernels from python")
|
||||
|
||||
# Operation description
|
||||
# math instruction description
|
||||
parser.add_argument("-i", "--instruction_shape",
|
||||
default=[1, 1, 1], nargs=3, type=int,
|
||||
help="This option describes the size of MMA op")
|
||||
parser.add_argument("-ta", "--element_a", default="float32", type=str,
|
||||
choices=['float64', 'float32', 'float16', 'bfloat16', 'int32', 'int8'],
|
||||
help='Data type of elements in input tensor A')
|
||||
parser.add_argument("-tb", "--element_b", default="float32", type=str,
|
||||
choices=['float64', 'float32', 'float16', 'bfloat16', 'int32', 'int8'],
|
||||
help='Data type of elements in input tensor B')
|
||||
parser.add_argument("-tc", "--element_c", default="float32", type=str,
|
||||
choices=['float64', 'float32', 'float16', 'bfloat16', 'int32', 'int8'],
|
||||
help='Data type of elements in input tensor C and output tensor D')
|
||||
parser.add_argument("-tacc", "--element_acc", default="float32", type=str,
|
||||
choices=['float64', 'float32', 'float16', 'bfloat16', 'int32', 'int8'],
|
||||
help='Data type of accumulator')
|
||||
parser.add_argument('-m', "--math", default="multiply_add",
|
||||
type=str, choices=["multiply_add", "multiply_add_fast_bf16", "multiply_add_fast_f32"], help="math instruction")
|
||||
parser.add_argument('-op', "--opcode", default="simt", type=str,
|
||||
choices=["Simt", 'TensorOp'], help='This option describes whether you want to use tensor \
|
||||
cores (TensorOp) or regular SIMT cores (Simt) on GPU SM')
|
||||
# tile description
|
||||
parser.add_argument("-b", "--threadblock_shape",
|
||||
default=[128, 128, 8], nargs=3, type=int,
|
||||
help="This option describes the tile size a thread block with compute")
|
||||
parser.add_argument("-s", "--stages", default=4,
|
||||
type=int, help="Number of pipelines you want to use")
|
||||
parser.add_argument("-w", "--warp_count", default=[
|
||||
4, 2, 1], nargs=3, type=int,
|
||||
help="This option describes the number of warps along M, N, and K of the threadblock")
|
||||
parser.add_argument("-cc", "--compute_capability", default=80,
|
||||
type=int, help="This option describes CUDA SM architecture number")
|
||||
# A
|
||||
parser.add_argument('-la', "--layout_a", default="RowMajor", type=str, choices=[
|
||||
"RowMajor", "ColumnMajor", "RowMajorInterleaved32", "ColumnMajorInterleaved32"],
|
||||
help="Memory layout of input tensor A")
|
||||
parser.add_argument('-aa', '--alignment_a', default=1,
|
||||
type=int, help="Memory alignment of input tensor A")
|
||||
# B
|
||||
parser.add_argument('-lb', "--layout_b", default="RowMajor", type=str, choices=[
|
||||
"RowMajor", "ColumnMajor", "RowMajorInterleaved32", "ColumnMajorInterleaved32"],
|
||||
help="Memory layout of input tensor B")
|
||||
parser.add_argument('-ab', '--alignment_b', default=1,
|
||||
type=int, help="Memory alignment of input tensor B")
|
||||
# C
|
||||
parser.add_argument('-lc', "--layout_c", default="RowMajor", type=str, choices=[
|
||||
"RowMajor", "ColumnMajor", "RowMajorInterleaved32", "ColumnMajorInterleaved32"],
|
||||
help="Memory layout of input tensor C and output tensor D")
|
||||
parser.add_argument('-ac', '--alignment_c', default=1,
|
||||
type=int, help="Memory alignment of input tensor C and output tensor D")
|
||||
# epilogue
|
||||
parser.add_argument("-te", "--element_epilogue", default="float32", type=str,
|
||||
choices=['float64', 'float32', 'float16', 'bfloat16'], help='Epilogue datatype')
|
||||
parser.add_argument("-ep", "--epilogue_functor", default="LinearCombination",
|
||||
type=str, choices=['LinearCombination', 'FastLinearCombinationClamp', 'LinearCombinationClamp'],
|
||||
help="This option describes the epilogue part of the kernel")
|
||||
# swizzling
|
||||
parser.add_argument("-sw", "--swizzling_functor", default="IdentitySwizzle1", type=str, choices=[
|
||||
"IdentitySwizzle1", "IdentitySwizzle2", "IdentitySwizzle4", "IdentitySwizzle8", "HorizontalSwizzle"],
|
||||
help="This option describes how thread blocks are scheduled on GPU")
|
||||
# precompute mode
|
||||
parser.add_argument("-pm", "--precompute_mode",
|
||||
default="Device", type=str, choices=["Host", "Device"],
|
||||
help="Grouped Gemm Scheduing on device only (Device) or using host precompute (Host)")
|
||||
# arguments
|
||||
parser.add_argument("-p", "--problem_size_dir", type=str,
|
||||
help="path to the csv file contains the problem sizes")
|
||||
parser.add_argument("-alpha", "--alpha", default=1.0, type=float, help="alpha")
|
||||
parser.add_argument("-beta", "--beta", default=0.0, type=float, help="beta")
|
||||
|
||||
parser.add_argument('--print_cuda', action="store_true",
|
||||
help="print the underlying CUDA kernel")
|
||||
|
||||
try:
|
||||
args = parser.parse_args()
|
||||
except:
|
||||
sys.exit(0)
|
||||
|
||||
pycutlass.get_memory_pool(init_pool_size=2**30, max_pool_size=2**32)
|
||||
|
||||
element_a = getattr(cutlass, args.element_a)
|
||||
element_b = getattr(cutlass, args.element_b)
|
||||
element_c = getattr(cutlass, args.element_c)
|
||||
element_acc = getattr(cutlass, args.element_acc)
|
||||
math_operation = getattr(MathOperation, args.math)
|
||||
opclass = getattr(cutlass.OpClass, args.opcode)
|
||||
|
||||
math_inst = MathInstruction(
|
||||
args.instruction_shape, element_a, element_b,
|
||||
element_acc, opclass, math_operation
|
||||
)
|
||||
|
||||
tile_description = TileDescription(
|
||||
args.threadblock_shape, args.stages, args.warp_count,
|
||||
math_inst, args.compute_capability, args.compute_capability
|
||||
)
|
||||
|
||||
layout_a = getattr(cutlass, args.layout_a)
|
||||
layout_b = getattr(cutlass, args.layout_b)
|
||||
layout_c = getattr(cutlass, args.layout_c)
|
||||
|
||||
A = TensorDescription(
|
||||
element_a, layout_a, args.alignment_a
|
||||
)
|
||||
|
||||
B = TensorDescription(
|
||||
element_b, layout_b, args.alignment_b
|
||||
)
|
||||
|
||||
C = TensorDescription(
|
||||
element_c, layout_c, args.alignment_c
|
||||
)
|
||||
|
||||
element_epilogue = getattr(cutlass, args.element_epilogue)
|
||||
epilogue_functor = getattr(EpilogueFunctor, args.epilogue_functor)
|
||||
swizzling_functor = getattr(cutlass, args.swizzling_functor)
|
||||
precompute_mode = getattr(SchedulerMode, args.precompute_mode)
|
||||
|
||||
operation = GemmOperationGrouped(
|
||||
arch=args.compute_capability, tile_description=tile_description,
|
||||
A=A, B=B, C=C, element_epilogue=element_epilogue,
|
||||
epilogue_functor=epilogue_functor, swizzling_functor=swizzling_functor,
|
||||
precompute_mode=precompute_mode
|
||||
)
|
||||
|
||||
if args.print_cuda:
|
||||
print(operation.rt_module.emit())
|
||||
|
||||
pycutlass.compiler.add_module([operation, ])
|
||||
|
||||
reference_module = ReferenceModule(A, B, C)
|
||||
|
||||
# get problems
|
||||
problem_sizes = []
|
||||
with open(args.problem_size_dir) as csv_file:
|
||||
reader = csv.reader(csv_file)
|
||||
for row in reader:
|
||||
problem_sizes.append(
|
||||
cutlass.gemm.GemmCoord(int(row[0]), int(row[1]), int(row[2]))
|
||||
)
|
||||
|
||||
problem_count = len(problem_sizes)
|
||||
|
||||
tensor_As = []
|
||||
tensor_Bs = []
|
||||
tensor_Cs = []
|
||||
tensor_Ds = []
|
||||
problem_sizes_coord = []
|
||||
tensor_D_refs = []
|
||||
|
||||
for problem_size in problem_sizes:
|
||||
if args.element_a != "int8":
|
||||
if args.element_a == "bfloat16":
|
||||
tensor_A = np.ceil(np.random.uniform(low=-8.5, high=7.5, size=(problem_size.m()
|
||||
* problem_size.k(),))).astype(bfloat16)
|
||||
else:
|
||||
tensor_A = np.ceil(np.random.uniform(low=-8.5, high=7.5, size=(problem_size.m()
|
||||
* problem_size.k(),))).astype(getattr(np, args.element_a))
|
||||
else:
|
||||
tensor_A = np.random.uniform(low=-2, high=2, size=(problem_size.m()
|
||||
* problem_size.k(),)).astype(getattr(np, args.element_a))
|
||||
|
||||
if args.element_b != "int8":
|
||||
if args.element_b == "bfloat16":
|
||||
tensor_B = np.ceil(np.random.uniform(low=-8.5, high=7.5, size=(problem_size.k()
|
||||
* problem_size.n(),))).astype(bfloat16)
|
||||
else:
|
||||
tensor_B = np.ceil(np.random.uniform(low=-8.5, high=7.5, size=(problem_size.k()
|
||||
* problem_size.n(),))).astype(getattr(np, args.element_b))
|
||||
else:
|
||||
tensor_B = np.random.uniform(low=-2, high=2, size=(problem_size.k()
|
||||
* problem_size.n(),)).astype(getattr(np, args.element_b))
|
||||
|
||||
if args.element_c != "int8":
|
||||
if args.element_c == "bfloat16":
|
||||
tensor_C = np.ceil(np.random.uniform(low=-8.5, high=7.5, size=(problem_size.m()
|
||||
* problem_size.n(),))).astype(bfloat16)
|
||||
else:
|
||||
tensor_C = np.ceil(np.random.uniform(low=-8.5, high=7.5, size=(problem_size.m()
|
||||
* problem_size.n(),))).astype(getattr(np, args.element_c))
|
||||
else:
|
||||
tensor_C = np.random.uniform(low=-2, high=2, size=(problem_size.m()
|
||||
* problem_size.n(),)).astype(getattr(np, args.element_c))
|
||||
tensor_D = np.zeros_like(tensor_C)
|
||||
|
||||
tensor_As.append(tensor_A)
|
||||
tensor_Bs.append(tensor_B)
|
||||
tensor_Cs.append(tensor_C)
|
||||
tensor_Ds.append(tensor_D)
|
||||
tensor_D_refs.append(reference_module.run(
|
||||
tensor_A, tensor_B, tensor_C, problem_size, args.alpha, args.beta))
|
||||
problem_sizes_coord.append(problem_size)
|
||||
|
||||
arguments = GemmGroupedArguments(
|
||||
operation, problem_sizes_coord, tensor_As, tensor_Bs, tensor_Cs, tensor_Ds,
|
||||
output_op=LinearCombinationFunctorArguments(args.alpha, args.beta)
|
||||
)
|
||||
|
||||
operation.run(arguments)
|
||||
|
||||
arguments.sync()
|
||||
|
||||
for tensor_d, tensor_d_ref in zip(tensor_Ds, tensor_D_refs):
|
||||
assert np.array_equal(tensor_d, tensor_d_ref)
|
||||
|
||||
print("Passed.")
|
||||
3
examples/40_cutlass_py/grouped_gemm_problem_size.csv
Normal file
3
examples/40_cutlass_py/grouped_gemm_problem_size.csv
Normal file
@@ -0,0 +1,3 @@
|
||||
128,128,128
|
||||
128,128,256
|
||||
512,128,384
|
||||
|
@@ -1,169 +0,0 @@
|
||||
|
||||
# System modules
|
||||
import numpy as np
|
||||
import os.path
|
||||
import sys
|
||||
import ctypes
|
||||
|
||||
# CUDA Python modules
|
||||
from cuda import cuda
|
||||
from cuda import nvrtc
|
||||
|
||||
# CUTLASS modules
|
||||
import library
|
||||
import manifest as cutlass_manifest
|
||||
import generator
|
||||
import rt
|
||||
|
||||
|
||||
#
|
||||
# Construct an SGEMM
|
||||
#
|
||||
|
||||
manifest = cutlass_manifest.Manifest()
|
||||
|
||||
generator.GenerateSM50_Simt(manifest, "11.5.0")
|
||||
|
||||
#
|
||||
# Construct a GEMM operation
|
||||
#
|
||||
|
||||
operation = manifest.operations_by_name['cutlass_simt_sgemm_128x128_8x2_nt_align1']
|
||||
|
||||
#
|
||||
# Construct a runtime GEMM operation
|
||||
#
|
||||
gemm = rt.Gemm(operation)
|
||||
|
||||
#
|
||||
# Initialize context
|
||||
#
|
||||
err, = cuda.cuInit(0)
|
||||
|
||||
if err != cuda.CUresult.CUDA_SUCCESS:
|
||||
raise RuntimeError("CUDA Error %s" % str(err))
|
||||
|
||||
err, device = cuda.cuDeviceGet(0)
|
||||
|
||||
if err != cuda.CUresult.CUDA_SUCCESS:
|
||||
raise RuntimeError("CUDA Error %s" % str(err))
|
||||
|
||||
err, context = cuda.cuCtxCreate(0, device)
|
||||
|
||||
if err != cuda.CUresult.CUDA_SUCCESS:
|
||||
raise RuntimeError("CUDA Error %s" % str(err))
|
||||
|
||||
#
|
||||
# Construct a module
|
||||
#
|
||||
|
||||
architectures = [80,]
|
||||
include_paths = [
|
||||
'../../include',
|
||||
'../../tools/util/include',
|
||||
]
|
||||
|
||||
compilation_options = rt.CompilationOptions(architectures, include_paths)
|
||||
|
||||
module = rt.Module('module.cu', [gemm], compilation_options)
|
||||
|
||||
#
|
||||
# Setup a workspace
|
||||
#
|
||||
|
||||
M, N, K = (128, 128, 128)
|
||||
|
||||
tensor_A = np.ndarray(M * K, dtype=np.float32)
|
||||
tensor_B = np.ndarray(N * K, dtype=np.float32)
|
||||
tensor_C = np.ndarray(M * N, dtype=np.float32)
|
||||
tensor_D = np.ndarray(M * N, dtype=np.float32)
|
||||
|
||||
err, tensor_A_d = cuda.cuMemAlloc(tensor_A.size * tensor_A.itemsize)
|
||||
if err != cuda.CUresult.CUDA_SUCCESS:
|
||||
raise RuntimeError("CUDA Error %s" % str(err))
|
||||
|
||||
err, tensor_B_d = cuda.cuMemAlloc(tensor_B.size * tensor_B.itemsize)
|
||||
if err != cuda.CUresult.CUDA_SUCCESS:
|
||||
raise RuntimeError("CUDA Error %s" % str(err))
|
||||
|
||||
err, tensor_C_d = cuda.cuMemAlloc(tensor_C.size * tensor_C.itemsize)
|
||||
if err != cuda.CUresult.CUDA_SUCCESS:
|
||||
raise RuntimeError("CUDA Error %s" % str(err))
|
||||
|
||||
err, tensor_D_d = cuda.cuMemAlloc(tensor_D.size * tensor_D.itemsize)
|
||||
if err != cuda.CUresult.CUDA_SUCCESS:
|
||||
raise RuntimeError("CUDA Error %s" % str(err))
|
||||
|
||||
err, stream = cuda.cuStreamCreate(0)
|
||||
if err != cuda.CUresult.CUDA_SUCCESS:
|
||||
raise RuntimeError("CUDA Error %s" % str(err))
|
||||
|
||||
tensors = [
|
||||
(tensor_A_d, tensor_A),
|
||||
(tensor_B_d, tensor_B),
|
||||
(tensor_C_d, tensor_C),
|
||||
(tensor_D_d, tensor_D)
|
||||
]
|
||||
|
||||
for tensor_device, tensor_host in tensors:
|
||||
bytes = tensor_host.size * tensor_host.itemsize
|
||||
print("Tensor has dimensions: %s (%d bytes)" % (str(tensor_host.size), tensor_host.itemsize))
|
||||
err, = cuda.cuMemcpyHtoDAsync(tensor_device, tensor_host, bytes, stream)
|
||||
print("updating tensor in device memory ", hex(int(tensor_device)))
|
||||
if err != cuda.CUresult.CUDA_SUCCESS:
|
||||
raise RuntimeError('CUDA Error %s' % str(err))
|
||||
|
||||
#
|
||||
# Initialize a host buffer
|
||||
#
|
||||
|
||||
arguments = rt.GemmArguments()
|
||||
|
||||
arguments.problem_size = rt.GemmCoord(M, N, K)
|
||||
|
||||
arguments.A = rt.TensorRef(tensor_A_d, M)
|
||||
arguments.B = rt.TensorRef(tensor_B_d, N)
|
||||
arguments.C = rt.TensorRef(tensor_C_d, M)
|
||||
arguments.D = rt.TensorRef(tensor_D_d, M)
|
||||
|
||||
host_workspace = bytearray(gemm.get_host_workspace_size(arguments))
|
||||
device_workspace = None
|
||||
|
||||
launch_config = gemm.plan(arguments)
|
||||
|
||||
byte_count = gemm.initialize(host_workspace, device_workspace, launch_config, arguments)
|
||||
|
||||
#
|
||||
# Launch the kernel
|
||||
#
|
||||
|
||||
err = gemm.run(host_workspace, device_workspace, launch_config)
|
||||
|
||||
if err != cuda.CUresult.CUDA_SUCCESS:
|
||||
raise RuntimeError('CUDA Error %s' % str(err))
|
||||
|
||||
#
|
||||
# Verify results
|
||||
#
|
||||
err, = cuda.cuStreamSynchronize(stream)
|
||||
|
||||
if err != cuda.CUresult.CUDA_SUCCESS:
|
||||
raise RuntimeError("CUDA Error %s" % str(err))
|
||||
|
||||
|
||||
#
|
||||
# Debug reporting of byte array contents
|
||||
#
|
||||
|
||||
def PrintBytearray(host_workspace):
|
||||
uint_str = None
|
||||
prefix = None
|
||||
print("uint32_t host_workspace[] = {")
|
||||
for idx, byte in enumerate(host_workspace):
|
||||
if not (idx % 4):
|
||||
if uint_str is not None:
|
||||
print(prefix, uint_str, ",")
|
||||
prefix = "/* offset: %d B */ 0x" % idx
|
||||
uint_str = ""
|
||||
uint_str = "{:02x}".format(byte) + uint_str
|
||||
print("};")
|
||||
36
examples/41_multi_head_attention/CMakeLists.txt
Normal file
36
examples/41_multi_head_attention/CMakeLists.txt
Normal file
@@ -0,0 +1,36 @@
|
||||
|
||||
# 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.
|
||||
|
||||
|
||||
|
||||
cutlass_example_add_executable(
|
||||
41_multi_head_attention
|
||||
fused_multihead_attention.cu
|
||||
)
|
||||
|
||||
1145
examples/41_multi_head_attention/fused_multihead_attention.cu
Normal file
1145
examples/41_multi_head_attention/fused_multihead_attention.cu
Normal file
File diff suppressed because it is too large
Load Diff
626
examples/41_multi_head_attention/gemm_attention.h
Normal file
626
examples/41_multi_head_attention/gemm_attention.h
Normal file
@@ -0,0 +1,626 @@
|
||||
/***************************************************************************************************
|
||||
* 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 holdvr 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 the FusedMultiHeadAttention Class
|
||||
|
||||
The class contains the following:
|
||||
1) GEMM0 with epilogue fusion,
|
||||
2) GEMM1 with mainloop fusion, and
|
||||
3) A lightweight full softmax reduction kernel.
|
||||
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
#include <cmath>
|
||||
#include <iostream>
|
||||
#include <vector>
|
||||
#include <limits>
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/arch/memory.h"
|
||||
#include "cutlass/arch/memory_sm75.h"
|
||||
#include "cutlass/epilogue/threadblock/epilogue_visitor_with_softmax.h"
|
||||
#include "cutlass/epilogue/thread/scale_type.h"
|
||||
#include "cutlass/gemm/kernel/default_gemm_grouped_softmax_mainloop_fusion.h"
|
||||
#include "cutlass/reduction/kernel/reduce_softmax_final.h"
|
||||
#include "gemm_grouped_with_softmax_visitor.h"
|
||||
|
||||
namespace cutlass {
|
||||
|
||||
template <
|
||||
typename ElementQ_,
|
||||
typename LayoutQ_,
|
||||
typename ElementK_,
|
||||
typename LayoutK_,
|
||||
typename ElementP_,
|
||||
typename LayoutP_,
|
||||
typename ElementCompute_,
|
||||
typename OperatorClass_,
|
||||
typename ArchTag_,
|
||||
typename ThreadblockShape0_,
|
||||
typename ThreadblockShape1_,
|
||||
typename WarpShape0_,
|
||||
typename WarpShape1_,
|
||||
typename InstructionShape_,
|
||||
int kStages0_,
|
||||
int kStages1_,
|
||||
bool UseMasking_ = false,
|
||||
cutlass::gemm::kernel::GroupScheduleMode GroupScheduleMode0_ = cutlass::gemm::kernel::GroupScheduleMode::kHostPrecompute,
|
||||
cutlass::gemm::kernel::GroupScheduleMode GroupScheduleMode1_ = cutlass::gemm::kernel::GroupScheduleMode::kHostPrecompute,
|
||||
int Alignment = 128 / cutlass::sizeof_bits<ElementQ_>::value,
|
||||
typename ElementSoftmax_ = ElementP_
|
||||
>
|
||||
class FusedMultiHeadAttention {
|
||||
public:
|
||||
|
||||
using ElementQ = ElementQ_;
|
||||
using ElementK = ElementK_;
|
||||
using ElementP = ElementP_;
|
||||
using ElementV = ElementK;
|
||||
using ElementOutput = ElementP;
|
||||
using ElementAccumulator = ElementCompute_;
|
||||
|
||||
using LayoutQ = LayoutQ_;
|
||||
using LayoutK = LayoutK_;
|
||||
using LayoutP = LayoutP_;
|
||||
using LayoutV = LayoutK;
|
||||
using LayoutO = LayoutP;
|
||||
|
||||
using ElementNorm = cutlass::half_t;
|
||||
using ElementSum = cutlass::half_t;
|
||||
using ElementSoftmaxCompute = float;
|
||||
using LayoutNorm = cutlass::layout::RowMajor;
|
||||
|
||||
using ThreadblockSwizzle = cutlass::gemm::threadblock::GemmBatchedIdentityThreadblockSwizzle;
|
||||
|
||||
using OperatorClass = OperatorClass_;
|
||||
using ArchTag = ArchTag_;
|
||||
|
||||
using ThreadblockShape0 = ThreadblockShape0_;
|
||||
using WarpShape0 = WarpShape0_;
|
||||
|
||||
using ThreadblockShape1 = ThreadblockShape1_;
|
||||
using WarpShape1 = WarpShape1_;
|
||||
|
||||
static int const Stages0 = kStages0_;
|
||||
static int const Stages1 = kStages1_;
|
||||
|
||||
using InstructionShape = InstructionShape_;
|
||||
|
||||
using EpilogueOutputOp0 = cutlass::epilogue::thread::LinearCombination<
|
||||
ElementOutput, 128 / cutlass::sizeof_bits<ElementOutput>::value,
|
||||
ElementAccumulator, ElementAccumulator, cutlass::epilogue::thread::ScaleType::OnlyAlphaScaling>;
|
||||
|
||||
using EpilogueOutputOp1 = cutlass::epilogue::thread::LinearCombination<
|
||||
ElementOutput, 128 / cutlass::sizeof_bits<ElementOutput>::value,
|
||||
ElementAccumulator, ElementAccumulator, cutlass::epilogue::thread::ScaleType::Nothing>;
|
||||
|
||||
using Operator = typename cutlass::gemm::device::DefaultGemmConfiguration<
|
||||
OperatorClass, ArchTag, ElementQ, ElementK, ElementP,
|
||||
ElementAccumulator>::Operator;
|
||||
static bool const kInternalTranspose = cutlass::platform::is_same<LayoutP, cutlass::layout::ColumnMajor>::value;
|
||||
|
||||
static bool const kUseMasking = UseMasking_;
|
||||
|
||||
static cutlass::gemm::kernel::GroupScheduleMode const kGroupScheduleMode0 = GroupScheduleMode0_;
|
||||
static cutlass::gemm::kernel::GroupScheduleMode const kGroupScheduleMode1 = GroupScheduleMode1_;
|
||||
|
||||
using MapArguments = cutlass::gemm::kernel::detail::MapArguments<
|
||||
ElementQ,
|
||||
LayoutQ,
|
||||
cutlass::ComplexTransform::kNone,
|
||||
8,
|
||||
ElementK,
|
||||
LayoutK,
|
||||
cutlass::ComplexTransform::kNone,
|
||||
8,
|
||||
LayoutP,
|
||||
kInternalTranspose
|
||||
>;
|
||||
|
||||
using DefaultGemmKernel = typename cutlass::gemm::kernel::DefaultGemm<
|
||||
typename MapArguments::ElementA,
|
||||
typename MapArguments::LayoutA,
|
||||
MapArguments::kAlignmentA,
|
||||
typename MapArguments::ElementB,
|
||||
typename MapArguments::LayoutB,
|
||||
MapArguments::kAlignmentB,
|
||||
ElementP,
|
||||
typename MapArguments::LayoutC,
|
||||
ElementAccumulator,
|
||||
OperatorClass,
|
||||
ArchTag,
|
||||
ThreadblockShape0,
|
||||
WarpShape0,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp0,
|
||||
ThreadblockSwizzle,
|
||||
Stages0,
|
||||
true,
|
||||
Operator,
|
||||
cutlass::gemm::SharedMemoryClearOption::kNone
|
||||
>::GemmKernel;
|
||||
|
||||
using EpilogueVisitor = typename cutlass::epilogue::threadblock::EpilogueVisitorSoftmax<
|
||||
ThreadblockShape0,
|
||||
DefaultGemmKernel::kThreadCount,
|
||||
typename DefaultGemmKernel::Epilogue::OutputTileIterator,
|
||||
typename EpilogueOutputOp0::ElementCompute,
|
||||
ElementNorm,
|
||||
ElementSum,
|
||||
ElementSoftmaxCompute,
|
||||
EpilogueOutputOp0,
|
||||
kUseMasking
|
||||
>;
|
||||
|
||||
using Epilogue = typename cutlass::epilogue::threadblock::EpilogueWithVisitorFromExistingEpilogue<
|
||||
EpilogueVisitor,
|
||||
typename DefaultGemmKernel::Epilogue
|
||||
>::Epilogue;
|
||||
|
||||
using GemmKernel0 = cutlass::gemm::kernel::GemmGroupedWithEpilogueVistor<
|
||||
typename DefaultGemmKernel::Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
kGroupScheduleMode0,
|
||||
kInternalTranspose,
|
||||
kUseMasking
|
||||
>;
|
||||
|
||||
using GemmGrouped0 = cutlass::gemm::device::GemmGrouped<GemmKernel0>;
|
||||
|
||||
using ApplyFinalReductionDevice = cutlass::reduction::kernel::ApplySoftmaxFinalReduction<
|
||||
ElementNorm,
|
||||
ElementSum,
|
||||
typename GemmGrouped0::GemmKernel::EpilogueVisitor::ElementSoftmaxCompute,
|
||||
typename GemmGrouped0::GemmKernel::EpilogueVisitor::ThreadblockShape,
|
||||
true
|
||||
>;
|
||||
|
||||
using GemmKernel1 = typename cutlass::gemm::kernel::DefaultGemmGroupedSoftmaxMainloopFusion<
|
||||
ElementP,
|
||||
LayoutP,
|
||||
cutlass::ComplexTransform::kNone,
|
||||
128 / cutlass::sizeof_bits<ElementQ>::value,
|
||||
ElementV,
|
||||
LayoutV,
|
||||
cutlass::ComplexTransform::kNone,
|
||||
128 / cutlass::sizeof_bits<ElementK>::value,
|
||||
ElementNorm,
|
||||
LayoutNorm,
|
||||
ElementOutput,
|
||||
LayoutO,
|
||||
ElementAccumulator,
|
||||
OperatorClass,
|
||||
ArchTag,
|
||||
ThreadblockShape1,
|
||||
WarpShape1,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp1,
|
||||
ThreadblockSwizzle,
|
||||
Stages1,
|
||||
kGroupScheduleMode1
|
||||
>::GemmKernel;
|
||||
|
||||
using GemmGrouped1 = cutlass::gemm::device::GemmGrouped<GemmKernel1>;
|
||||
|
||||
public:
|
||||
|
||||
/// Arguments class
|
||||
struct Arguments {
|
||||
cutlass::gemm::GemmCoord *problem_sizes0;
|
||||
cutlass::gemm::GemmCoord *problem_sizes0_real;
|
||||
cutlass::gemm::GemmCoord *problem_sizes1;
|
||||
int problem_count;
|
||||
int threadblock_count;
|
||||
|
||||
ElementQ ** ptr_Q;
|
||||
ElementK ** ptr_K;
|
||||
ElementP ** ptr_P;
|
||||
ElementP ** ptr_V;
|
||||
ElementP ** ptr_O;
|
||||
|
||||
ElementNorm **ptr_Max;
|
||||
ElementSum **ptr_Sum;
|
||||
|
||||
ElementP *block_P;
|
||||
ElementNorm *block_Norm;
|
||||
ElementSum *block_Sum;
|
||||
int64_t *offset_P;
|
||||
int64_t *offset_Norm_Device;
|
||||
int64_t *offset_Sum_Device;
|
||||
|
||||
typename LayoutQ::Stride::LongIndex *ldq;
|
||||
typename LayoutK::Stride::LongIndex *ldk;
|
||||
typename LayoutP::Stride::LongIndex *ldp;
|
||||
typename LayoutP::Stride::LongIndex *ldv;
|
||||
typename LayoutP::Stride::LongIndex *ldo;
|
||||
|
||||
cutlass::gemm::GemmCoord *problem_sizes0_host;
|
||||
cutlass::gemm::GemmCoord *problem_sizes1_host;
|
||||
|
||||
ElementAccumulator alpha0;
|
||||
ElementAccumulator alpha1;
|
||||
ElementAccumulator beta;
|
||||
|
||||
int head_number;
|
||||
int batch_size;
|
||||
int seq_length;
|
||||
|
||||
typename ApplyFinalReductionDevice::Arguments reduction;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
Arguments():
|
||||
problem_count(0),
|
||||
threadblock_count(0),
|
||||
ptr_Q(nullptr),
|
||||
ptr_K(nullptr),
|
||||
ptr_P(nullptr),
|
||||
ptr_V(nullptr),
|
||||
ptr_O(nullptr),
|
||||
ptr_Max(nullptr),
|
||||
ptr_Sum(nullptr),
|
||||
block_P(nullptr),
|
||||
block_Norm(nullptr),
|
||||
block_Sum(nullptr),
|
||||
offset_P(nullptr),
|
||||
offset_Norm_Device(nullptr),
|
||||
offset_Sum_Device(nullptr),
|
||||
ldq(nullptr),
|
||||
ldk(nullptr),
|
||||
ldp(nullptr),
|
||||
ldv(nullptr),
|
||||
ldo(nullptr),
|
||||
head_number(0),
|
||||
batch_size(0),
|
||||
seq_length(0)
|
||||
{
|
||||
|
||||
}
|
||||
|
||||
Arguments(
|
||||
cutlass::gemm::GemmCoord *problem_sizes0,
|
||||
cutlass::gemm::GemmCoord *problem_sizes1,
|
||||
int problem_count,
|
||||
int threadblock_count,
|
||||
ElementQ ** ptr_Q,
|
||||
ElementK ** ptr_K,
|
||||
ElementP ** ptr_P,
|
||||
ElementP ** ptr_V,
|
||||
ElementP ** ptr_O,
|
||||
ElementNorm **ptr_Max,
|
||||
ElementSum **ptr_Sum,
|
||||
ElementP *block_P,
|
||||
ElementNorm *block_Norm,
|
||||
ElementSum *block_Sum,
|
||||
int64_t *offset_P,
|
||||
int64_t *offset_Norm_Device,
|
||||
int64_t *offset_Sum_Device,
|
||||
typename LayoutQ::Stride::LongIndex *ldq,
|
||||
typename LayoutK::Stride::LongIndex *ldk,
|
||||
typename LayoutP::Stride::LongIndex *ldp,
|
||||
typename LayoutP::Stride::LongIndex *ldv,
|
||||
typename LayoutP::Stride::LongIndex *ldo,
|
||||
ElementAccumulator alpha0,
|
||||
ElementAccumulator alpha1,
|
||||
ElementAccumulator beta,
|
||||
int head_number,
|
||||
int batch_size,
|
||||
int seq_length,
|
||||
cutlass::gemm::GemmCoord *problem_sizes0_host = nullptr,
|
||||
cutlass::gemm::GemmCoord *problem_sizes1_host = nullptr,
|
||||
cutlass::gemm::GemmCoord *problem_sizes0_real = nullptr
|
||||
):
|
||||
problem_sizes0(problem_sizes0),
|
||||
problem_sizes1(problem_sizes1),
|
||||
problem_count(problem_count),
|
||||
threadblock_count(threadblock_count),
|
||||
ptr_Q(ptr_Q),
|
||||
ptr_K(ptr_K),
|
||||
ptr_P(ptr_P),
|
||||
ptr_V(ptr_V),
|
||||
ptr_O(ptr_O),
|
||||
ptr_Max(ptr_Max),
|
||||
ptr_Sum(ptr_Sum),
|
||||
block_P(block_P),
|
||||
block_Norm(block_Norm),
|
||||
block_Sum(block_Sum),
|
||||
offset_P(offset_P),
|
||||
offset_Norm_Device(offset_Norm_Device),
|
||||
offset_Sum_Device(offset_Sum_Device),
|
||||
ldq(ldq),
|
||||
ldk(ldk),
|
||||
ldp(ldp),
|
||||
ldv(ldv),
|
||||
ldo(ldo),
|
||||
alpha0(alpha0),
|
||||
alpha1(alpha1),
|
||||
beta(beta),
|
||||
head_number(head_number),
|
||||
batch_size(batch_size),
|
||||
seq_length(seq_length),
|
||||
problem_sizes0_host(problem_sizes0_host),
|
||||
problem_sizes1_host(problem_sizes1_host),
|
||||
problem_sizes0_real(problem_sizes0_real),
|
||||
reduction(
|
||||
problem_sizes0,
|
||||
block_Norm,
|
||||
block_Sum,
|
||||
offset_Norm_Device,
|
||||
offset_Sum_Device
|
||||
)
|
||||
{
|
||||
|
||||
}
|
||||
|
||||
|
||||
};
|
||||
|
||||
struct Params {
|
||||
cutlass::gemm::GemmCoord *problem_sizes0;
|
||||
cutlass::gemm::GemmCoord *problem_sizes0_real;
|
||||
cutlass::gemm::GemmCoord *problem_sizes1;
|
||||
int problem_count;
|
||||
int threadblock_count;
|
||||
|
||||
ElementQ ** ptr_Q;
|
||||
ElementK ** ptr_K;
|
||||
ElementP ** ptr_P;
|
||||
ElementP ** ptr_V;
|
||||
ElementP ** ptr_O;
|
||||
|
||||
ElementNorm **ptr_Max;
|
||||
ElementSum **ptr_Sum;
|
||||
|
||||
ElementP *block_P;
|
||||
ElementNorm *block_Norm;
|
||||
ElementSum *block_Sum;
|
||||
int64_t *offset_P;
|
||||
int64_t *offset_Norm_Device;
|
||||
int64_t *offset_Sum_Device;
|
||||
|
||||
typename LayoutQ::Stride::LongIndex *ldq;
|
||||
typename LayoutK::Stride::LongIndex *ldk;
|
||||
typename LayoutP::Stride::LongIndex *ldp;
|
||||
typename LayoutP::Stride::LongIndex *ldv;
|
||||
typename LayoutP::Stride::LongIndex *ldo;
|
||||
|
||||
cutlass::gemm::GemmCoord *problem_sizes0_host;
|
||||
cutlass::gemm::GemmCoord *problem_sizes1_host;
|
||||
|
||||
ElementAccumulator alpha0;
|
||||
ElementAccumulator alpha1;
|
||||
ElementAccumulator beta;
|
||||
|
||||
int head_number;
|
||||
int batch_size;
|
||||
int seq_length;
|
||||
|
||||
typename ApplyFinalReductionDevice::Params reduction;
|
||||
|
||||
Params():
|
||||
problem_count(0),
|
||||
threadblock_count(0),
|
||||
ptr_Q(nullptr),
|
||||
ptr_K(nullptr),
|
||||
ptr_P(nullptr),
|
||||
ptr_V(nullptr),
|
||||
ptr_O(nullptr),
|
||||
ptr_Max(nullptr),
|
||||
ptr_Sum(nullptr),
|
||||
block_P(nullptr),
|
||||
block_Norm(nullptr),
|
||||
block_Sum(nullptr),
|
||||
offset_P(nullptr),
|
||||
offset_Norm_Device(nullptr),
|
||||
offset_Sum_Device(nullptr),
|
||||
ldq(nullptr),
|
||||
ldk(nullptr),
|
||||
ldp(nullptr),
|
||||
ldv(nullptr),
|
||||
ldo(nullptr),
|
||||
problem_sizes0(nullptr),
|
||||
problem_sizes1(nullptr),
|
||||
problem_sizes0_real(nullptr),
|
||||
head_number(0),
|
||||
batch_size(0),
|
||||
seq_length(0)
|
||||
{
|
||||
|
||||
}
|
||||
|
||||
Params(Arguments const &args, void *workspace = nullptr):
|
||||
problem_sizes0(args.problem_sizes0),
|
||||
problem_sizes1(args.problem_sizes1),
|
||||
problem_count(args.problem_count),
|
||||
threadblock_count(args.threadblock_count),
|
||||
ptr_Q(args.ptr_Q),
|
||||
ptr_K(args.ptr_K),
|
||||
ptr_P(args.ptr_P),
|
||||
ptr_V(args.ptr_V),
|
||||
ptr_O(args.ptr_O),
|
||||
ptr_Max(args.ptr_Max),
|
||||
ptr_Sum(args.ptr_Sum),
|
||||
block_P(args.block_P),
|
||||
block_Norm(args.block_Norm),
|
||||
block_Sum(args.block_Sum),
|
||||
offset_P(args.offset_P),
|
||||
offset_Norm_Device(args.offset_Norm_Device),
|
||||
offset_Sum_Device(args.offset_Sum_Device),
|
||||
ldq(args.ldq),
|
||||
ldk(args.ldk),
|
||||
ldp(args.ldp),
|
||||
ldv(args.ldv),
|
||||
ldo(args.ldo),
|
||||
problem_sizes0_host(args.problem_sizes0_host),
|
||||
problem_sizes1_host(args.problem_sizes1_host),
|
||||
problem_sizes0_real(args.problem_sizes0_real),
|
||||
alpha0(args.alpha0),
|
||||
alpha1(args.alpha1),
|
||||
beta(args.beta),
|
||||
head_number(args.head_number),
|
||||
batch_size(args.batch_size),
|
||||
seq_length(args.seq_length),
|
||||
reduction(args.reduction)
|
||||
{
|
||||
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
private:
|
||||
|
||||
Params params_;
|
||||
GemmGrouped0 gemm_grouped0;
|
||||
GemmGrouped1 gemm_grouped1;
|
||||
|
||||
|
||||
public:
|
||||
|
||||
/// Ctor
|
||||
FusedMultiHeadAttention() {
|
||||
|
||||
}
|
||||
|
||||
/// Initialize
|
||||
Status initialize(Arguments const &args,
|
||||
void *workspace0 = nullptr,
|
||||
void *workspace1 = nullptr) {
|
||||
|
||||
params_ = Params(args);
|
||||
|
||||
typename GemmGrouped0::Arguments args_gemm0(
|
||||
params_.problem_sizes0,
|
||||
params_.problem_count,
|
||||
params_.threadblock_count,
|
||||
params_.ptr_Q,
|
||||
params_.ptr_K,
|
||||
params_.ptr_P,
|
||||
params_.ptr_P,
|
||||
params_.ptr_Max,
|
||||
params_.ptr_Sum,
|
||||
params_.ldq,
|
||||
params_.ldk,
|
||||
params_.ldp,
|
||||
params_.ldp,
|
||||
typename GemmGrouped0::GemmKernel::EpilogueVisitor::Arguments(
|
||||
{
|
||||
params_.alpha0,
|
||||
params_.beta
|
||||
}
|
||||
),
|
||||
params_.problem_sizes0_host,
|
||||
params_.problem_sizes0_real
|
||||
);
|
||||
|
||||
|
||||
Status result0 = gemm_grouped0.initialize(args_gemm0, workspace0);
|
||||
|
||||
typename EpilogueOutputOp1::Params epilogue_op1(params_.alpha1, params_.beta);
|
||||
|
||||
typename GemmGrouped1::Arguments args_gemm1(
|
||||
params_.problem_sizes1,
|
||||
params_.problem_count,
|
||||
params_.threadblock_count,
|
||||
epilogue_op1,
|
||||
params_.ptr_P,
|
||||
params_.ptr_V,
|
||||
params_.ptr_O,
|
||||
params_.ptr_O,
|
||||
(void**)params_.ptr_Max,
|
||||
(void**)params_.ptr_Sum,
|
||||
params_.ldp,
|
||||
params_.ldv,
|
||||
params_.ldo,
|
||||
params_.ldo,
|
||||
params_.problem_sizes1_host
|
||||
);
|
||||
|
||||
Status result1 = gemm_grouped1.initialize(args_gemm1, workspace1);
|
||||
|
||||
if ((result0 == cutlass::Status::kSuccess) && (result1 == cutlass::Status::kSuccess) ) {
|
||||
return cutlass::Status::kSuccess;
|
||||
}else{
|
||||
if (result0 != cutlass::Status::kSuccess) {
|
||||
return result0;
|
||||
}else{
|
||||
return result1;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Run
|
||||
Status run(cudaStream_t stream = nullptr) {
|
||||
|
||||
Status result = gemm_grouped0.run();
|
||||
cudaError_t error_info;
|
||||
|
||||
if (result != cutlass::Status::kSuccess) {
|
||||
return cutlass::Status::kErrorInternal;
|
||||
}
|
||||
|
||||
int thread_per_block = 1024;
|
||||
|
||||
dim3 final_reduction_grid(params_.head_number * params_.batch_size);
|
||||
dim3 final_reduction_block(thread_per_block);
|
||||
|
||||
cutlass::Kernel<ApplyFinalReductionDevice><<<
|
||||
final_reduction_grid, final_reduction_block, sizeof(typename ApplyFinalReductionDevice::SharedStorage), stream
|
||||
>>>(params_.reduction);
|
||||
|
||||
error_info = cudaGetLastError();
|
||||
|
||||
if (error_info != cudaSuccess) {
|
||||
return cutlass::Status::kErrorInternal;
|
||||
}
|
||||
|
||||
result = gemm_grouped1.run();
|
||||
|
||||
if (result != cutlass::Status::kSuccess) {
|
||||
return cutlass::Status::kErrorInternal;
|
||||
}
|
||||
|
||||
return cutlass::Status::kSuccess;
|
||||
}
|
||||
|
||||
/// Function call operator
|
||||
Status operator()(cudaStream_t stream = nullptr) {
|
||||
return run(stream);
|
||||
}
|
||||
};
|
||||
|
||||
}
|
||||
@@ -0,0 +1,522 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017 - 2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
/*! \file
|
||||
\brief Grouped GEMM kernel with epilogue visitor customized for softmax
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/fast_math.h"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/matrix_coord.h"
|
||||
#include "cutlass/complex.h"
|
||||
#include "cutlass/semaphore.h"
|
||||
#include "cutlass/util/device_memory.h"
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/trace.h"
|
||||
#include "cutlass/gemm/kernel/gemm_transpose_operands.h"
|
||||
#include "cutlass/gemm/kernel/gemm_grouped_problem_visitor.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace kernel {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename Mma_, ///! Threadblock-scoped matrix multiply-accumulate
|
||||
typename Epilogue_, ///! Epilogue
|
||||
typename ThreadblockSwizzle_, ///! Threadblock swizzling function
|
||||
GroupScheduleMode GroupScheduleMode_, ///! Type of scheduling to perform
|
||||
bool Transposed_ = false,
|
||||
bool UseMask_ = false
|
||||
>
|
||||
struct GemmGroupedWithEpilogueVistor {
|
||||
public:
|
||||
|
||||
using Mma = Mma_;
|
||||
using Epilogue = Epilogue_;
|
||||
using ThreadblockSwizzle = ThreadblockSwizzle_;
|
||||
|
||||
static GroupScheduleMode const kGroupScheduleMode = GroupScheduleMode_;
|
||||
|
||||
using EpilogueVisitor = typename Epilogue::Visitor;
|
||||
using EpilogueOutputOp = typename EpilogueVisitor::ElementwiseFunctor;
|
||||
static bool const kTransposed = Transposed_;
|
||||
|
||||
// Optional transpose
|
||||
using MapArguments = kernel::detail::MapArguments<
|
||||
typename Mma::IteratorA::Element,
|
||||
typename Mma::IteratorA::Layout,
|
||||
Mma::kTransformA,
|
||||
Mma::IteratorA::AccessType::kElements,
|
||||
typename Mma::IteratorB::Element,
|
||||
typename Mma::IteratorB::Layout,
|
||||
Mma::kTransformB,
|
||||
Mma::IteratorB::AccessType::kElements,
|
||||
typename Mma::LayoutC,
|
||||
kTransposed
|
||||
>;
|
||||
|
||||
// Public-facing type definitions related to operand element type, layout, and complex conjugate
|
||||
// operation. Must interact with the 'kTransposed' notion.
|
||||
using ElementA = typename MapArguments::ElementA;
|
||||
using LayoutA = typename MapArguments::LayoutA;
|
||||
using ElementB = typename MapArguments::ElementB;
|
||||
using LayoutB = typename MapArguments::LayoutB;
|
||||
using ElementC = typename EpilogueVisitor::ElementOutput;
|
||||
using LayoutC = typename MapArguments::LayoutC;
|
||||
|
||||
using ElementNorm = typename EpilogueVisitor::ElementNorm;
|
||||
using ElementSum = typename EpilogueVisitor::ElementSum;
|
||||
|
||||
static ComplexTransform const kTransformA = MapArguments::kTransformA;
|
||||
static ComplexTransform const kTransformB = MapArguments::kTransformB;
|
||||
|
||||
// Type definitions about the mainloop.
|
||||
using Operator = typename Mma::Operator;
|
||||
using OperatorClass = typename Mma::Operator::OperatorClass;
|
||||
using ThreadblockShape = typename Mma::Shape;
|
||||
using WarpShape = typename Mma::Operator::Shape;
|
||||
using InstructionShape = typename Mma::Policy::Operator::InstructionShape;
|
||||
using ArchTag = typename Mma::ArchTag;
|
||||
|
||||
static int const kStages = Mma::kStages;
|
||||
static int const kAlignmentA = MapArguments::kAlignmentA;
|
||||
static int const kAlignmentB = MapArguments::kAlignmentB;
|
||||
static int const kAlignmentC = EpilogueVisitor::kElementsPerAccess;
|
||||
|
||||
/// Warp count (concept: GemmShape)
|
||||
using WarpCount = typename Mma::WarpCount;
|
||||
static int const kThreadCount = 32 * WarpCount::kCount;
|
||||
|
||||
using ProblemVisitor = GemmGroupedProblemVisitor<
|
||||
ThreadblockShape,
|
||||
kGroupScheduleMode,
|
||||
kThreadCount,
|
||||
kThreadCount,
|
||||
kTransposed>;
|
||||
|
||||
//
|
||||
// Structures
|
||||
//
|
||||
|
||||
/// Argument structure
|
||||
struct Arguments {
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
GemmCoord *problem_sizes;
|
||||
// when using mask, real problem sizes may not be aligned
|
||||
// then we need to mask out unpadded elements in softmax
|
||||
GemmCoord *problem_sizes_real;
|
||||
int problem_count;
|
||||
int threadblock_count;
|
||||
|
||||
ElementA ** ptr_A;
|
||||
ElementB ** ptr_B;
|
||||
ElementC ** ptr_C;
|
||||
ElementC ** ptr_D;
|
||||
|
||||
ElementNorm **ptr_Max;
|
||||
ElementSum **ptr_Sum;
|
||||
|
||||
typename LayoutA::Stride::LongIndex *lda;
|
||||
typename LayoutB::Stride::LongIndex *ldb;
|
||||
typename LayoutC::Stride::LongIndex *ldc;
|
||||
typename LayoutC::Stride::LongIndex *ldd;
|
||||
|
||||
typename EpilogueVisitor::Arguments epilogue_visitor;
|
||||
|
||||
// Only used by device-level operator
|
||||
GemmCoord *host_problem_sizes;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
/// Default ctor
|
||||
CUTLASS_HOST_DEVICE
|
||||
Arguments():
|
||||
problem_count(0),
|
||||
threadblock_count(0),
|
||||
ptr_A(nullptr),
|
||||
ptr_B(nullptr),
|
||||
ptr_C(nullptr),
|
||||
ptr_D(nullptr),
|
||||
ptr_Max(nullptr),
|
||||
ptr_Sum(nullptr),
|
||||
lda(nullptr),
|
||||
ldb(nullptr),
|
||||
ldc(nullptr),
|
||||
ldd(nullptr),
|
||||
host_problem_sizes(nullptr)
|
||||
{
|
||||
|
||||
}
|
||||
|
||||
/// Ctor
|
||||
CUTLASS_HOST_DEVICE
|
||||
Arguments(
|
||||
GemmCoord *problem_sizes,
|
||||
int problem_count,
|
||||
int threadblock_count,
|
||||
ElementA ** ptr_A,
|
||||
ElementB ** ptr_B,
|
||||
ElementC ** ptr_C,
|
||||
ElementC ** ptr_D,
|
||||
ElementNorm **ptr_Max,
|
||||
ElementSum **ptr_Sum,
|
||||
typename LayoutA::Stride::LongIndex *lda,
|
||||
typename LayoutB::Stride::LongIndex *ldb,
|
||||
typename LayoutC::Stride::LongIndex *ldc,
|
||||
typename LayoutC::Stride::LongIndex *ldd,
|
||||
typename EpilogueVisitor::Arguments epilogue_visitor_,
|
||||
GemmCoord *host_problem_sizes=nullptr,
|
||||
GemmCoord *problem_sizes_real=nullptr
|
||||
):
|
||||
problem_sizes(problem_sizes),
|
||||
problem_count(problem_count),
|
||||
threadblock_count(threadblock_count),
|
||||
ptr_A(ptr_A),
|
||||
ptr_B(ptr_B),
|
||||
ptr_C(ptr_C),
|
||||
ptr_D(ptr_D),
|
||||
ptr_Max(ptr_Max),
|
||||
ptr_Sum(ptr_Sum),
|
||||
lda(lda),
|
||||
ldb(ldb),
|
||||
ldc(ldc),
|
||||
ldd(ldd),
|
||||
epilogue_visitor(epilogue_visitor_),
|
||||
host_problem_sizes(host_problem_sizes),
|
||||
problem_sizes_real(problem_sizes_real)
|
||||
{
|
||||
|
||||
}
|
||||
};
|
||||
|
||||
//
|
||||
// Structure for precomputing values in host memory and passing to kernels
|
||||
//
|
||||
|
||||
/// Parameters structure
|
||||
struct Params {
|
||||
|
||||
typename ProblemVisitor::Params problem_visitor;
|
||||
GemmCoord *problem_sizes_real;
|
||||
int threadblock_count;
|
||||
|
||||
ElementA ** ptr_A;
|
||||
ElementB ** ptr_B;
|
||||
ElementC ** ptr_C;
|
||||
ElementC ** ptr_D;
|
||||
|
||||
ElementNorm **ptr_Max;
|
||||
ElementSum **ptr_Sum;
|
||||
|
||||
typename LayoutA::Stride::LongIndex *lda;
|
||||
typename LayoutB::Stride::LongIndex *ldb;
|
||||
typename LayoutC::Stride::LongIndex *ldc;
|
||||
typename LayoutC::Stride::LongIndex *ldd;
|
||||
|
||||
typename EpilogueVisitor::Params epilogue_visitor;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params():
|
||||
ptr_A(nullptr),
|
||||
ptr_B(nullptr),
|
||||
ptr_C(nullptr),
|
||||
ptr_D(nullptr),
|
||||
ptr_Max(nullptr),
|
||||
ptr_Sum(nullptr),
|
||||
lda(nullptr),
|
||||
ldb(nullptr),
|
||||
ldc(nullptr),
|
||||
ldd(nullptr),
|
||||
problem_sizes_real(problem_sizes_real)
|
||||
{ }
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(Arguments const &args, void *workspace = nullptr, int32_t tile_count = 0):
|
||||
problem_visitor(args.problem_sizes, args.problem_count, workspace, tile_count),
|
||||
threadblock_count(args.threadblock_count),
|
||||
ptr_A(args.ptr_A),
|
||||
ptr_B(args.ptr_B),
|
||||
ptr_C(args.ptr_C),
|
||||
ptr_D(args.ptr_D),
|
||||
ptr_Max(args.ptr_Max),
|
||||
ptr_Sum(args.ptr_Sum),
|
||||
lda(args.lda),
|
||||
ldb(args.ldb),
|
||||
ldc(args.ldc),
|
||||
ldd(args.ldd),
|
||||
epilogue_visitor(args.epilogue_visitor),
|
||||
problem_sizes_real(args.problem_sizes_real)
|
||||
{
|
||||
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
void update(
|
||||
Arguments const &args,
|
||||
void *workspace = nullptr,
|
||||
int32_t tile_count = -1) {
|
||||
|
||||
problem_visitor = typename ProblemVisitor::Params(args.problem_sizes, args.problem_count, workspace, tile_count);
|
||||
threadblock_count = args.threadblock_count;
|
||||
ptr_A = args.ptr_A;
|
||||
ptr_B = args.ptr_B;
|
||||
ptr_C = args.ptr_C;
|
||||
ptr_D = args.ptr_D;
|
||||
ptr_Max = args.ptr_Max;
|
||||
ptr_Sum = args.ptr_Sum;
|
||||
lda = args.lda;
|
||||
ldb = args.ldb;
|
||||
ldc = args.ldc;
|
||||
ldd = args.ldd;
|
||||
problem_sizes_real = args.problem_sizes_real;
|
||||
}
|
||||
};
|
||||
|
||||
/// Shared memory storage structure
|
||||
struct SharedStorage {
|
||||
union {
|
||||
typename Mma::SharedStorage main_loop;
|
||||
struct {
|
||||
typename Epilogue::SharedStorage epilogue;
|
||||
typename EpilogueVisitor::SharedStorage visitor;
|
||||
} epilogue;
|
||||
} kernel;
|
||||
|
||||
// ProblemVisitor shared storage can't be overlapped with others
|
||||
typename ProblemVisitor::SharedStorage problem_visitor;
|
||||
};
|
||||
|
||||
|
||||
public:
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
CUTLASS_DEVICE
|
||||
GemmGroupedWithEpilogueVistor() { }
|
||||
|
||||
/// Determines whether kernel satisfies alignment
|
||||
static Status can_implement(cutlass::gemm::GemmCoord const & problem_size) {
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
static Status can_implement(Arguments const &args) {
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
static size_t get_extra_workspace_size(
|
||||
Arguments const &args,
|
||||
cutlass::gemm::GemmCoord const &grid_tiled_shape) {
|
||||
|
||||
return 0;
|
||||
}
|
||||
|
||||
/// Executes one GEMM
|
||||
CUTLASS_DEVICE
|
||||
void operator()(Params const ¶ms, SharedStorage &shared_storage) {
|
||||
|
||||
//
|
||||
// These types shadow the type-level definitions and support the ability to implement
|
||||
// a 'transposed' GEMM that computes the transposed problems.
|
||||
//
|
||||
using ElementA = typename Mma::IteratorA::Element;
|
||||
using LayoutA = typename Mma::IteratorA::Layout;
|
||||
using ElementB = typename Mma::IteratorB::Element;
|
||||
using LayoutB = typename Mma::IteratorB::Layout;
|
||||
using ElementC = typename EpilogueVisitor::ElementOutput;
|
||||
using LayoutC = typename Mma::LayoutC;
|
||||
|
||||
//
|
||||
// Problem visitor.
|
||||
//
|
||||
ProblemVisitor problem_visitor(
|
||||
params.problem_visitor,
|
||||
shared_storage.problem_visitor,
|
||||
blockIdx.x);
|
||||
|
||||
// Outer 'persistent' loop to iterate over tiles
|
||||
while (problem_visitor.next_tile()) {
|
||||
|
||||
GemmCoord problem_size = problem_visitor.problem_size();
|
||||
int32_t problem_idx = problem_visitor.problem_index();
|
||||
int32_t threadblock_idx = int32_t(problem_visitor.threadblock_idx());
|
||||
|
||||
GemmCoord grid_shape = problem_visitor.grid_shape(problem_size);
|
||||
|
||||
cutlass::gemm::GemmCoord threadblock_offset(
|
||||
int(threadblock_idx / grid_shape.n()) * Mma::Shape::kM,
|
||||
int(threadblock_idx % grid_shape.n()) * Mma::Shape::kN,
|
||||
0);
|
||||
|
||||
// Load element pointers. Exchange pointers and strides if working on the transpose
|
||||
ElementA *ptr_A = reinterpret_cast<ElementA *>((kTransposed ? params.ptr_B[problem_idx] : params.ptr_A[problem_idx]));
|
||||
typename LayoutA::LongIndex ldm_A = (kTransposed ? params.ldb[problem_idx] : params.lda[problem_idx]);
|
||||
|
||||
ElementB *ptr_B = reinterpret_cast<ElementB *>((kTransposed ? params.ptr_A[problem_idx] : params.ptr_B[problem_idx]));
|
||||
typename LayoutB::LongIndex ldm_B = (kTransposed ? params.lda[problem_idx] : params.ldb[problem_idx]);
|
||||
|
||||
// Compute initial location in logical coordinates
|
||||
cutlass::MatrixCoord tb_offset_A{
|
||||
threadblock_offset.m(),
|
||||
0,
|
||||
};
|
||||
|
||||
cutlass::MatrixCoord tb_offset_B{
|
||||
0,
|
||||
threadblock_offset.n()
|
||||
};
|
||||
|
||||
// Compute position within threadblock
|
||||
int thread_idx = threadIdx.x;
|
||||
|
||||
// Construct iterators to A and B operands
|
||||
typename Mma::IteratorA iterator_A(
|
||||
LayoutA(ldm_A),
|
||||
ptr_A,
|
||||
{problem_size.m(), problem_size.k()},
|
||||
thread_idx,
|
||||
tb_offset_A);
|
||||
|
||||
typename Mma::IteratorB iterator_B(
|
||||
LayoutB(ldm_B),
|
||||
ptr_B,
|
||||
{problem_size.k(), problem_size.n()},
|
||||
thread_idx,
|
||||
tb_offset_B);
|
||||
|
||||
typename Mma::FragmentC accumulators;
|
||||
|
||||
accumulators.clear();
|
||||
|
||||
// Broadcast the warp_id computed by lane 0 to ensure dependent code
|
||||
// is compiled as warp-uniform.
|
||||
int warp_idx = __shfl_sync(0xffffffff, threadIdx.x / 32, 0);
|
||||
|
||||
int lane_idx = threadIdx.x % 32;
|
||||
|
||||
//
|
||||
// Matrix multiply phase
|
||||
//
|
||||
|
||||
// Construct thread-scoped matrix multiply
|
||||
Mma mma(shared_storage.kernel.main_loop, thread_idx, warp_idx, lane_idx);
|
||||
|
||||
// Compute threadblock-scoped matrix multiply-add
|
||||
int gemm_k_iterations = (problem_size.k() + Mma::Shape::kK - 1) / Mma::Shape::kK;
|
||||
|
||||
// Wait for all threads to finish their epilogue phases from the previous tile.
|
||||
__syncthreads();
|
||||
|
||||
// Compute threadblock-scoped matrix multiply-add
|
||||
mma(
|
||||
gemm_k_iterations,
|
||||
accumulators,
|
||||
iterator_A,
|
||||
iterator_B,
|
||||
accumulators);
|
||||
|
||||
ElementC *ptr_C = params.ptr_C[problem_idx];
|
||||
ElementC *ptr_D = params.ptr_D[problem_idx];
|
||||
|
||||
ElementNorm *ptr_Max = params.ptr_Max[problem_idx];
|
||||
ElementSum *ptr_Sum = params.ptr_Sum[problem_idx];
|
||||
|
||||
LayoutC layout_C(params.ldc[problem_idx]);
|
||||
LayoutC layout_D(params.ldd[problem_idx]);
|
||||
|
||||
int column_offset = (threadblock_offset.n() / ThreadblockShape::kN) * problem_size.m();
|
||||
|
||||
typename EpilogueVisitor::OutputTileIterator::Params params_C(layout_C);
|
||||
typename EpilogueVisitor::OutputTileIterator::Params params_D(layout_D);
|
||||
|
||||
//
|
||||
// Construct the epilogue visitor
|
||||
//
|
||||
|
||||
EpilogueVisitor epilogue_visitor(
|
||||
params.epilogue_visitor,
|
||||
shared_storage.kernel.epilogue.visitor,
|
||||
problem_size.mn(),
|
||||
thread_idx,
|
||||
warp_idx,
|
||||
lane_idx,
|
||||
params_C,
|
||||
params_D,
|
||||
ptr_C,
|
||||
ptr_D,
|
||||
ptr_Max,
|
||||
ptr_Sum,
|
||||
threadblock_offset.mn(),
|
||||
column_offset,
|
||||
params.problem_sizes_real[problem_idx].mn()
|
||||
);
|
||||
|
||||
// Construct the epilogue
|
||||
Epilogue epilogue(
|
||||
shared_storage.kernel.epilogue.epilogue,
|
||||
thread_idx,
|
||||
warp_idx,
|
||||
lane_idx);
|
||||
|
||||
// Execute the epilogue operator to update the destination tensor
|
||||
epilogue(epilogue_visitor, accumulators);
|
||||
|
||||
// Next tile
|
||||
problem_visitor.advance(gridDim.x);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace kernel
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -116,6 +116,10 @@ foreach(EXAMPLE
|
||||
34_transposed_conv2d
|
||||
35_gemm_softmax
|
||||
36_gather_scatter_fusion
|
||||
37_gemm_layernorm_gemm_fusion
|
||||
38_syr2k_grouped
|
||||
39_gemm_permute
|
||||
41_multi_head_attention
|
||||
)
|
||||
|
||||
add_subdirectory(${EXAMPLE})
|
||||
|
||||
Reference in New Issue
Block a user