CUTLASS 1.2

This commit is contained in:
akerr
2018-10-26 14:38:46 -07:00
parent 2332df492e
commit 74df0331f2
97 changed files with 11301 additions and 632 deletions

View File

@@ -21,6 +21,8 @@
# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
include_directories("external/googletest/googletest/include")
add_subdirectory(external/googletest/googletest)
add_subdirectory(test)
add_subdirectory(nvrtc)

View File

@@ -29,6 +29,7 @@ set(CUTLASS_PERF_TEST_HEADERS
performance_result.h
gemm/cublas_dispatch.h
gemm/cutlass_dispatch.h
gemm/cutlass_dispatch_splitK_PI.h
gemm/gemm_perf_testbed.h
gemm/gemm_profiler.h
)
@@ -36,9 +37,11 @@ set(CUTLASS_PERF_TEST_HEADERS
set(CUTLASS_PERF_TEST_SOURCES
cutlass_perf_test.cu
gemm/sgemm.cu
gemm/sgemm_splitK.cu
gemm/dgemm.cu
gemm/hgemm.cu
gemm/igemm.cu
gemm/igemm_splitK.cu
gemm/wmma_gemm.cu
gemm/wmma_binary_gemm.cu
gemm/wmma_integer_gemm.cu

View File

@@ -1,121 +0,0 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (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 {nv-internal-release}
#if (defined(__CUDACC__) && (!defined(__CUDA_ARCH__) || __CUDA_ARCH__ >= 750))
#pragma warning( disable : 4503)
////////////////////////////////////////////////////////////////////////////////////////////////////
#include "cutlass/gemm/gemm.h"
#include "cutlass/gemm/bmma_gemm_traits.h"
#include "tools/test/perf/cutlass_perf_test.h"
#include "tools/test/perf/gemm/gemm_profiler.h"
#include "tools/test/perf/gemm/cutlass_dispatch.h"
#include "tools/test/perf/gemm/gemm_perf_testbed.h"
////////////////////////////////////////////////////////////////////////////////////////////////////
template<typename Traits>
struct BmmaGemmDispatch {
typedef cutlass::gemm::Gemm<Traits> Gemm;
typedef typename Gemm::Params Params;
/// Indicate warp-level GEMM
static bool const kThreadMultiplyAdd = false;
static bool const kRunCuBLAS = false;
static cutlass::MatrixLayout::Kind const kLayoutA = Traits::kLayoutA;
static cutlass::MatrixLayout::Kind const kLayoutB = Traits::kLayoutB;
//
// Data members
//
/// Params argument
Params params;
//
// Methods
//
BmmaGemmDispatch() {}
/// Initializes params object
BmmaGemmDispatch(int m, int n, int k, int alpha,
cutlass::Vector<cutlass::bin1_t, 32> const* d_a, int lda,
cutlass::Vector<cutlass::bin1_t, 32> const* d_b, int ldb, int beta,
int const* d_c, int ldc, int* d_d, int ldd) {
params.initialize(m, n, k * 32, alpha, d_a, lda, d_b, ldb, beta, d_c, ldc, d_d, ldd);
}
/// Initializes params object
BmmaGemmDispatch(Params const& _params) : params(_params) {}
/// Launches kernel
cudaError_t operator()() { return Gemm::launch(params); }
};
////////////////////////////////////////////////////////////////////////////////////////////////////
namespace perf {
////////////////////////////////////////////////////////////////////////////////////////////////////
int profile_bmma_gemm(TestbenchOutput<GemmProblem> &output, TestbenchOptions const &options, Config const &config) {
typedef perf::GemmProfiler<cutlass::Vector<cutlass::bin1_t, 32>, cutlass::Vector<cutlass::bin1_t, 32>, int, int, int> GemmProfiler;
int results = 0;
{
typedef cutlass::gemm::BmmaGemmTraits<cutlass::Shape<1024, 128, 128>,
cutlass::Shape<1024, 32, 32>,
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor>
BmmaGemmTraits;
typedef BmmaGemmDispatch<BmmaGemmTraits> Dispatch;
results |= profile_gemm<Dispatch, GemmProfiler>(output, "bmma_gemm_tn", options, config);
}
return results;
}
////////////////////////////////////////////////////////////////////////////////////////////////////
struct BmmaGemmRegistrar {
BmmaGemmRegistrar() { RegisterGemmProfileFunc(profile_bmma_gemm); }
};
volatile BmmaGemmRegistrar _BmmaGemmRegistrar;
} // namespace perf
#endif // if (defined(__CUDACC__) && (!defined(__CUDA_ARCH__) || __CUDA_ARCH__ >= 750)

View File

@@ -89,4 +89,76 @@ struct CublasGemmDispatch {
}
};
/// Dispatcher for batched strided cuBLAS kernels
template <typename AType, typename BType, typename CType, typename Accumulator, typename Scalar>
struct CublasBatchedStridedGemmDispatch {
/// Type used for device-side allocations
typedef typename cutlass::TypeTraits<AType>::device_type ADeviceType;
typedef typename cutlass::TypeTraits<BType>::device_type BDeviceType;
typedef typename cutlass::TypeTraits<CType>::device_type CDeviceType;
typedef typename cutlass::TypeTraits<Accumulator>::device_type AccumulatorDeviceType;
typedef typename cutlass::TypeTraits<Scalar>::device_type ScalarDeviceType;
static cublasOperation_t convert(cutlass::MatrixLayout::Kind layout) {
switch (layout) {
case cutlass::MatrixLayout::kRowMajor:
return CUBLAS_OP_T;
case cutlass::MatrixLayout::kColumnMajor:
return CUBLAS_OP_N;
default:
break;
}
return CUBLAS_OP_N;
}
/// Launches a cuBLAS GEMM kernel
cublasStatus_t operator()(cublasHandle_t handle,
cutlass::MatrixLayout::Kind layout_a,
cutlass::MatrixLayout::Kind layout_b,
int m,
int n,
int k,
Scalar alpha,
const ADeviceType *A,
int lda,
long long int batch_stride_A,
const BDeviceType *B,
int ldb,
long long int batch_stride_B,
Scalar beta,
CDeviceType *C,
int ldc,
long long int batch_stride_C,
int batch_count,
cublasGemmAlgo_t algorithm) {
#if defined(CUDA_VERSION) && CUDA_VERSION >= 9010
return cublasGemmStridedBatchedEx(handle,
convert(layout_a),
convert(layout_b),
m,
n,
k,
reinterpret_cast<ScalarDeviceType const *>(&alpha),
A,
cutlass::TypeTraits<ADeviceType>::cublas_type,
lda,
batch_stride_A,
B,
cutlass::TypeTraits<BDeviceType>::cublas_type,
ldb,
batch_stride_B,
reinterpret_cast<ScalarDeviceType const *>(&beta),
C,
cutlass::TypeTraits<CDeviceType>::cublas_type,
ldc,
batch_stride_C,
batch_count,
cutlass::TypeTraits<AccumulatorDeviceType>::cublas_type,
algorithm);
#else
return CUBLAS_STATUS_NOT_SUPPORTED;
#endif
}
};
} // namespace perf

View File

@@ -81,6 +81,32 @@ struct CutlassDispatch {
params.initialize(m, n, k, alpha, d_a, lda, d_b, ldb, beta, d_c, ldc, d_d, ldd);
}
/// Initializes batched strided params object
CutlassDispatch(Index m,
Index n,
Index k,
ScalarEpilogue alpha,
ScalarA const* d_a,
Index lda,
long long int batch_stride_A,
ScalarB const* d_b,
Index ldb,
long long int batch_stride_B,
ScalarEpilogue beta,
ScalarC const* d_c,
Index ldc,
long long int batch_stride_C,
ScalarD* d_d,
Index ldd,
long long int batch_stride_D,
Index batch_count) {
params.initialize(m, n, k, alpha, d_a, lda, batch_stride_A,
d_b, ldb, batch_stride_B,
beta, d_c, ldc, batch_stride_C,
d_d, ldd, batch_stride_D,
batch_count);
}
/// Initializes params object
CutlassDispatch(Params const& _params) : params(_params) {}

View File

@@ -0,0 +1,172 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#pragma once
#include "cutlass/matrix_traits.h"
#include "tools/util/type_traits.h"
#include <cuda_runtime_api.h>
#include <assert.h>
namespace perf {
template <typename KernelClass_,
typename Index_,
typename ScalarA_,
typename ScalarB_,
typename ScalarC_,
typename ScalarD_,
typename Compute_,
typename ScalarEpilogue_,
bool ThreadMultiplyAdd_,
bool RunCuBLAS_ = true>
struct CutlassDispatchSplitKPIGemm {
typedef typename KernelClass_::Params Params;
typedef KernelClass_ KernelClass;
typedef Index_ Index;
typedef ScalarA_ ScalarA;
typedef ScalarB_ ScalarB;
typedef ScalarC_ ScalarC;
typedef ScalarD_ ScalarD;
typedef Compute_ Compute;
typedef ScalarEpilogue_ ScalarEpilogue;
static bool const kThreadMultiplyAdd = ThreadMultiplyAdd_;
static bool const kRunCuBLAS = RunCuBLAS_;
static cutlass::MatrixLayout::Kind const kLayoutA = KernelClass::Traits::kLayoutA;
static cutlass::MatrixLayout::Kind const kLayoutB = KernelClass::Traits::kLayoutB;
//
// Data members
//
/// Params argument
Params params;
/// splitK PI require workspace
typename cutlass::TypeTraits<Compute>::device_type *workspace_ptr;
//
// Methods
//
/// Ctor Initializes params object
CutlassDispatchSplitKPIGemm(Index m,
Index n,
Index k,
ScalarEpilogue alpha,
ScalarA const* d_a,
Index lda,
ScalarB const* d_b,
Index ldb,
ScalarEpilogue beta,
ScalarC const* d_c,
Index ldc,
ScalarD* d_d,
Index ldd) {
params.init_problem(m, n, k);
int workspace_size_in_byte = params.required_workspace_memory_in_byte();
cudaError_t workspace_err = cudaMalloc(&workspace_ptr, workspace_size_in_byte);
if (workspace_err != cudaSuccess) {
std::cout << "\nCUDA workspace malloc error: " << cudaGetErrorString(workspace_err)
<< "\n";
}
params.initialize(alpha, d_a, lda, d_b, ldb, beta, d_c, ldc, d_d, ldd, workspace_ptr);
}
/// Initializes batched strided params object
CutlassDispatchSplitKPIGemm(Index m,
Index n,
Index k,
ScalarEpilogue alpha,
ScalarA const* d_a,
Index lda,
long long int batch_stride_A,
ScalarB const* d_b,
Index ldb,
long long int batch_stride_B,
ScalarEpilogue beta,
ScalarC const* d_c,
Index ldc,
long long int batch_stride_C,
ScalarD* d_d,
Index ldd,
long long int batch_stride_D,
Index batch_count) {
assert(0);//batched strided splitK should never be called
}
/// Launches kernel
cudaError_t operator()() { return KernelClass::launch(params); }
~CutlassDispatchSplitKPIGemm() {
cudaError_t workspace_err = cudaFree(workspace_ptr);
if (workspace_err != cudaSuccess) {
std::cout << "\nCUDA workspace malloc error: " << cudaGetErrorString(workspace_err)
<< "\n";
}
}
};
template<
typename SplitKPIGemmTraits_
>
struct CutlassDispatchSplitKPIGemmBasic {
///
typedef SplitKPIGemmTraits_ Traits;
///
typedef typename Traits::KernelClass KernelClass;
/// Index type
typedef typename Traits::Index Index;
/// The scalar for A.
typedef typename Traits::ScalarA ScalarA;
/// The scalar for B.
typedef typename Traits::ScalarB ScalarB;
/// The scalar for C.
typedef typename Traits::ScalarC ScalarC;
/// The scalar for D.
typedef typename Traits::ScalarD ScalarD;
// TODO - support alternative accumulator and scalar types
typedef ScalarD Compute;
typedef Compute ScalarEpilogue;
typedef CutlassDispatchSplitKPIGemm<KernelClass,
Index,
ScalarA,
ScalarB,
ScalarC,
ScalarD,
Compute,
ScalarEpilogue,
true>
Dispatch;
};
} //namespace perf

View File

@@ -78,6 +78,7 @@ class GemmTestbed {
/// Dispatch object to cuBLAS GEMM
typedef CublasGemmDispatch<AType, BType, CType, Accumulator, Scalar> CublasDispatch;
typedef CublasBatchedStridedGemmDispatch<AType, BType, CType, Accumulator, Scalar> CublasBatchedStridedGemmDispatch;
//
// Type definitions
@@ -160,18 +161,20 @@ class GemmTestbed {
/// Resizes each tensor
void resize_helper(GemmProblem const &problem) {
resize_device_allocation(A,
initial_distribution.dist_A,
initial_distribution.seed,
problem.m,
problem.k,
problem.layout_A);
resize_device_allocation(A,
initial_distribution.dist_A,
initial_distribution.seed,
problem.m,
problem.k * problem.batch_count,
problem.layout_A);
resize_device_allocation(
B,
initial_distribution.dist_B,
initial_distribution.seed + 17, // compute distinct value from initial seed
problem.k,
problem.k * problem.batch_count,
problem.n,
problem.layout_B);
@@ -180,21 +183,21 @@ class GemmTestbed {
initial_distribution.dist_C,
initial_distribution.seed + 101, // compute distinct value from initial seed
problem.m,
problem.n,
problem.n * problem.batch_count,
cutlass::MatrixLayout::kColumnMajor);
resize_device_allocation(reference,
cutlass::Distribution(),
0,
problem.m,
problem.n,
problem.n * problem.batch_count,
cutlass::MatrixLayout::kColumnMajor);
resize_device_allocation(experimental,
cutlass::Distribution(),
0,
problem.m,
problem.n,
problem.n * problem.batch_count,
cutlass::MatrixLayout::kColumnMajor);
}
@@ -315,24 +318,36 @@ class GemmTestbed {
/// Inner dimension of GEMM problem
int K() const { return problem.k; }
/// batch count
int batch_count() const { return problem.batch_count; }
/// Returns a pointer to the A operand
ADeviceType *ptr_A() const { return A.get(); }
/// Leading dimension of A
int lda() const { return problem.lda(); }
///
long long int batch_stride_a() const{ return problem.batch_stride_a(); }
/// Returns a pointer to the B operand
BDeviceType *ptr_B() const { return B.get(); }
/// Leading dimension of B
int ldb() const { return problem.ldb(); }
///
long long int batch_stride_b() const{ return problem.batch_stride_b(); }
/// Returns a pointer to the initial state of the result tensor in device memory
CDeviceType *ptr_C_initial() const { return C_initial.get(); }
/// Leading dimension of C
int ldc() const { return problem.ldc(); }
///
long long int batch_stride_c() const { return problem.batch_stride_c(); }
/// Returns a pointer to the result tensor in device memory
CDeviceType *ptr_experimental() const { return experimental.get(); }
@@ -341,7 +356,7 @@ class GemmTestbed {
/// Returns the number of flops implied by the computation (1 multiply-accumulate = 2 flops)
uint64_t flops() const {
return uint64_t(problem.m) * uint64_t(problem.n) * uint64_t(problem.k) * detail::ElementCount<AType>::kValue * 2ULL;
return uint64_t(problem.batch_count) * uint64_t(problem.m) * uint64_t(problem.n) * uint64_t(problem.k) * detail::ElementCount<AType>::kValue * 2ULL;
}
/// Computes the speed of the computation in GFLOPs/s
@@ -373,28 +388,59 @@ class GemmTestbed {
/// Launches the cuBLAS GEMM - does not initialize output matrix
cublasStatus_t launch_cublas(cublasGemmAlgo_t algo) {
CublasDispatch dispatch;
if (problem.batch_count == 1) {
CublasDispatch dispatch;
Scalar alpha(Scalar(problem.alpha));
Scalar beta(Scalar(problem.beta));
Scalar alpha(Scalar(problem.alpha));
Scalar beta(Scalar(problem.beta));
status = dispatch(handle,
problem.layout_A,
problem.layout_B,
problem.m,
problem.n,
problem.k,
alpha,
ptr_A(),
lda(),
ptr_B(),
ldb(),
beta,
ptr_reference(),
ldc(),
algo);
status = dispatch(handle,
problem.layout_A,
problem.layout_B,
problem.m,
problem.n,
problem.k,
alpha,
ptr_A(),
lda(),
ptr_B(),
ldb(),
beta,
ptr_reference(),
ldc(),
algo);
return status;
return status;
}
else {
// call batched strided cublas
CublasBatchedStridedGemmDispatch dispatch;
Scalar alpha(Scalar(problem.alpha));
Scalar beta(Scalar(problem.beta));
status = dispatch(handle,
problem.layout_A,
problem.layout_B,
problem.m,
problem.n,
problem.k,
alpha,
ptr_A(),
lda(),
batch_stride_a(),
ptr_B(),
ldb(),
batch_stride_b(),
beta,
ptr_reference(),
ldc(),
batch_stride_c(),
batch_count(),
algo);
return status;
}
}
/// Verifies the 'test' tensor with 'ref'

View File

@@ -164,24 +164,52 @@ class GemmProfiler {
result.disposition = Disposition::Passed;
}
CutlassDispatch dispatch(testbed.M(),
testbed.N(),
testbed.K(),
testbed.alpha(),
testbed.ptr_A(),
testbed.lda(),
testbed.ptr_B(),
testbed.ldb(),
testbed.beta(),
testbed.ptr_C_initial(),
testbed.ldc(),
testbed.ptr_experimental(),
testbed.ldc());
CutlassDispatch *dispatch_ptr;
dispatch();
// check to see if we need to launch batched strided gemm
if (testbed.batch_count() == 1) {
dispatch_ptr = new CutlassDispatch(testbed.M(),
testbed.N(),
testbed.K(),
testbed.alpha(),
testbed.ptr_A(),
testbed.lda(),
testbed.ptr_B(),
testbed.ldb(),
testbed.beta(),
testbed.ptr_C_initial(),
testbed.ldc(),
testbed.ptr_experimental(),
testbed.ldc());
dispatch_ptr->operator()();
}
else {
dispatch_ptr = new CutlassDispatch(testbed.M(),
testbed.N(),
testbed.K(),
testbed.alpha(),
testbed.ptr_A(),
testbed.lda(),
testbed.batch_stride_a(),
testbed.ptr_B(),
testbed.ldb(),
testbed.batch_stride_b(),
testbed.beta(),
testbed.ptr_C_initial(),
testbed.ldc(),
testbed.batch_stride_c(),
testbed.ptr_experimental(),
testbed.ldc(),
testbed.batch_stride_c(),
testbed.batch_count());
dispatch_ptr->operator()();
}
if (cudaDeviceSynchronize() != cudaSuccess) {
result.disposition = Disposition::Failed;
delete dispatch_ptr;
return result;
}
@@ -202,35 +230,40 @@ class GemmProfiler {
}
// warmup launch
dispatch();
dispatch_ptr->operator()();
if (cudaDeviceSynchronize() != cudaSuccess) {
result.disposition = Disposition::Failed;
delete dispatch_ptr;
return result;
}
if (cudaEventRecord(events[0]) != cudaSuccess) {
result.disposition = Disposition::Failed;
delete dispatch_ptr;
return result;
}
for (int iter = 0; iter < options.iterations; ++iter) {
dispatch();
dispatch_ptr->operator()();
}
if (cudaEventRecord(events[1]) != cudaSuccess) {
result.disposition = Disposition::Failed;
delete dispatch_ptr;
return result;
}
if (cudaEventSynchronize(events[1]) != cudaSuccess) {
result.disposition = Disposition::Failed;
delete dispatch_ptr;
return result;
}
float average_ms = 0;
if (cudaEventElapsedTime(&average_ms, events[0], events[1]) != cudaSuccess) {
result.disposition = Disposition::Failed;
delete dispatch_ptr;
return result;
}
@@ -242,6 +275,7 @@ class GemmProfiler {
<< " failed with disposition: " << result.disposition << "\n";
}
delete dispatch_ptr;
return result;
}
@@ -265,7 +299,7 @@ class GemmProfiler {
std::vector<PerformanceResult<GemmProblem> > results;
results.push_back(execute_cutlass<CutlassDispatch>(problem, algorithm));
results.push_back(execute_cutlass<CutlassDispatch>(problem, algorithm));
// cool-down period
if (!options.dry_run) {
pause(options.sleep_time);
@@ -276,28 +310,30 @@ class GemmProfiler {
/// Runs the test and collects performance for all results
template <typename CutlassDispatch>
void schmoo(Range const &M, Range const &N, Range const &K) {
for (int m = M.start; m <= M.end; m = M.next(m)) {
for (int n = N.start; n <= N.end; n = N.next(n)) {
for (int k = K.start; k <= K.end; k = K.next(k)) {
std::vector<PerformanceResult<GemmProblem> > results =
void schmoo(Range const &M, Range const &N, Range const &K, Range const &batch_count) {
for (int b = batch_count.start; b <= batch_count.end; b = batch_count.next(b)) {
for (int m = M.start; m <= M.end; m = M.next(m)) {
for (int n = N.start; n <= N.end; n = N.next(n)) {
for (int k = K.start; k <= K.end; k = K.next(k)) {
std::vector<PerformanceResult<GemmProblem> > results =
execute<CutlassDispatch>(GemmProblem(m,
n,
k,
CutlassDispatch::kLayoutA,
CutlassDispatch::kLayoutB,
config.alpha,
config.beta));
n,
k,
CutlassDispatch::kLayoutA,
CutlassDispatch::kLayoutB,
config.alpha,
config.beta,
b));
for (std::vector<PerformanceResult<GemmProblem> >::const_iterator it = results.begin();
it != results.end();
++it) {
output.append(*it);
}
}
}
}
for (std::vector<PerformanceResult<GemmProblem> >::const_iterator it = results.begin();
it != results.end();
++it) {
output.append(*it);
}
}//k
}//n
}//m
}//batch_count
}
/// Runs the test over the problem space and reports only the best performance
@@ -369,7 +405,7 @@ int profile_gemm(TestbenchOutput<GemmProblem> &output,
config.problem_range.M, config.problem_range.N, config.problem_range.K);
} else {
perf.template schmoo<Dispatch>(
config.problem_range.M, config.problem_range.N, config.problem_range.K);
config.problem_range.M, config.problem_range.N, config.problem_range.K, config.problem_range.batch_count);
}
}

View File

@@ -0,0 +1,202 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#include "cutlass/gemm/gemm.h"
#include "cutlass/gemm/igemm_traits.h"
#include "cutlass/reduction/batched_reduction_traits.h"
#include "cutlass/gemm/device_gemm_traits.h"
#include "tools/test/perf/cutlass_perf_test.h"
#include "tools/test/perf/gemm/gemm_perf_testbed.h"
#include "tools/test/perf/gemm/gemm_profiler.h"
#include "tools/test/perf/gemm/cutlass_dispatch.h"
#include "tools/test/perf/gemm/cutlass_dispatch_splitK_PI.h"
#pragma warning( disable : 4503)
namespace perf {
////////////////////////////////////////////////////////////////////////////////////////////////////
template <typename OutputTile, int splits_count>
int profile_igemm_splitkpi_kernel(
TestbenchOutput<GemmProblem> &output,
TestbenchOptions const &options,
Config const &config,
std::string const &name,
std::string const &algo) {
typedef perf::GemmProfiler<int8_t, int8_t, int, int, int> GemmProfiler;
int results = 0;
{
/*batched igemm traits*/
typedef cutlass::gemm::IgemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor,
OutputTile
> IgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<int,
int,
int,
int,
int, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
// create a device gemm
typedef typename cutlass::gemm::SplitkPIGemmTraits<IgemmTraits, BatchedReductionTraits> deviceGemmTraits;
typedef typename CutlassDispatchSplitKPIGemmBasic<deviceGemmTraits>::Dispatch Dispatch;
results |= profile_gemm<Dispatch, GemmProfiler>(output, name + "_nn", options, config, algo + "_splitk_pi");
}
{
/*batched igemm traits*/
typedef cutlass::gemm::IgemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor,
OutputTile
> IgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<int,
int,
int,
int,
int, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
// create a device gemm
typedef typename cutlass::gemm::SplitkPIGemmTraits<IgemmTraits, BatchedReductionTraits> deviceGemmTraits;
typedef typename CutlassDispatchSplitKPIGemmBasic<deviceGemmTraits>::Dispatch Dispatch;
results |= profile_gemm<Dispatch, GemmProfiler>(output, name + "_nt", options, config, algo + "_splitk_pi");
}
{
/*batched igemm traits*/
typedef cutlass::gemm::IgemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor,
OutputTile
> IgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<int,
int,
int,
int,
int, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
// create a device gemm
typedef typename cutlass::gemm::SplitkPIGemmTraits<IgemmTraits, BatchedReductionTraits> deviceGemmTraits;
typedef typename CutlassDispatchSplitKPIGemmBasic<deviceGemmTraits>::Dispatch Dispatch;
results |= profile_gemm<Dispatch, GemmProfiler>(output, name + "_tn", options, config, algo + "_splitk_pi");
}
{
/*batched igemm traits*/
typedef cutlass::gemm::IgemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor,
OutputTile
> IgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<int,
int,
int,
int,
int, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
// create a device gemm
typedef typename cutlass::gemm::SplitkPIGemmTraits<IgemmTraits, BatchedReductionTraits> deviceGemmTraits;
typedef typename CutlassDispatchSplitKPIGemmBasic<deviceGemmTraits>::Dispatch Dispatch;
results |= profile_gemm<Dispatch, GemmProfiler>(output, name + "_tt", options, config, algo + "_splitk_pi");
}
return results;
}
/// Profiles all SGEMM tile sizes
int profile_igemm_splitkpi(TestbenchOutput<GemmProblem> &output, TestbenchOptions const &options, Config const &config) {
int results = 0;
/*128x128x32*/
results |= profile_igemm_splitkpi_kernel<cutlass::Shape<32, 128, 128>, 8 >(output, options, config, "igemm_128x128x32_splitk_pi_split8", "128x128");
results |= profile_igemm_splitkpi_kernel<cutlass::Shape<32, 128, 128>, 16 >(output, options, config, "igemm_128x128x32_splitk_pi_split16", "128x128");
results |= profile_igemm_splitkpi_kernel<cutlass::Shape<32, 128, 128>, 32 >(output, options, config, "igemm_128x128x32_splitk_pi_split32", "128x128");
results |= profile_igemm_splitkpi_kernel<cutlass::Shape<32, 128, 128>, 64 >(output, options, config, "igemm_128x128x32_splitk_pi_split64", "128x128");
/*128x64x32*/
results |= profile_igemm_splitkpi_kernel<cutlass::Shape<32, 64, 128>, 8 >(output, options, config, "igemm_128x64x32_splitk_pi_split8", "128x64");
results |= profile_igemm_splitkpi_kernel<cutlass::Shape<32, 64, 128>, 16 >(output, options, config, "igemm_128x64x32_splitk_pi_split16", "128x64");
results |= profile_igemm_splitkpi_kernel<cutlass::Shape<32, 64, 128>, 20 >(output, options, config, "igemm_128x64x32_splitk_pi_split20", "128x64");
results |= profile_igemm_splitkpi_kernel<cutlass::Shape<32, 64, 128>, 32 >(output, options, config, "igemm_128x64x32_splitk_pi_split32", "128x64");
results |= profile_igemm_splitkpi_kernel<cutlass::Shape<32, 64, 128>, 64 >(output, options, config, "igemm_128x64x32_splitk_pi_split64", "128x64");
/*128x32x32*/
results |= profile_igemm_splitkpi_kernel<cutlass::Shape<32, 32, 128>, 8 >(output, options, config, "igemm_128x32x32_splitk_pi_split8", "128x32");
results |= profile_igemm_splitkpi_kernel<cutlass::Shape<32, 32, 128>, 16 >(output, options, config, "igemm_128x32x32_splitk_pi_split16", "128x32");
results |= profile_igemm_splitkpi_kernel<cutlass::Shape<32, 32, 128>, 20 >(output, options, config, "igemm_128x32x32_splitk_pi_split20", "128x32");
results |= profile_igemm_splitkpi_kernel<cutlass::Shape<32, 32, 128>, 32 >(output, options, config, "igemm_128x32x32_splitk_pi_split32", "128x32");
results |= profile_igemm_splitkpi_kernel<cutlass::Shape<32, 32, 128>, 64 >(output, options, config, "igemm_128x32x32_splitk_pi_split64", "128x32");
/*64x64x32*/
results |= profile_igemm_splitkpi_kernel<cutlass::Shape<32, 64, 64>, 8 >(output, options, config, "igemm_64x64x32_splitk_pi_split8", "64x64");
results |= profile_igemm_splitkpi_kernel<cutlass::Shape<32, 64, 64>, 16 >(output, options, config, "igemm_64x64x32_splitk_pi_split16", "64x64");
results |= profile_igemm_splitkpi_kernel<cutlass::Shape<32, 64, 64>, 20 >(output, options, config, "igemm_64x64x32_splitk_pi_split20", "64x64");
results |= profile_igemm_splitkpi_kernel<cutlass::Shape<32, 64, 64>, 32 >(output, options, config, "igemm_64x64x32_splitk_pi_split32", "64x64");
results |= profile_igemm_splitkpi_kernel<cutlass::Shape<32, 64, 64>, 64 >(output, options, config, "igemm_64x64x32_splitk_pi_split64", "64x64");
return results;
}
struct IgemmSplitKPIRegistrar {
IgemmSplitKPIRegistrar() { RegisterGemmProfileFunc(profile_igemm_splitkpi); }
};
volatile IgemmSplitKPIRegistrar _IgemmSplitKPIRegistrar;
////////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace perf

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,187 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#include "cutlass/gemm/gemm.h"
#include "cutlass/gemm/sgemm_traits.h"
#include "cutlass/reduction/batched_reduction_traits.h"
#include "cutlass/gemm/device_gemm_traits.h"
#include "tools/test/perf/cutlass_perf_test.h"
#include "tools/test/perf/gemm/gemm_perf_testbed.h"
#include "tools/test/perf/gemm/gemm_profiler.h"
#include "tools/test/perf/gemm/cutlass_dispatch.h"
#include "tools/test/perf/gemm/cutlass_dispatch_splitK_PI.h"
#pragma warning( disable : 4503)
namespace perf {
////////////////////////////////////////////////////////////////////////////////////////////////////
template <typename OutputTile, int splits_count>
int profile_sgemm_splitkpi_kernel(
TestbenchOutput<GemmProblem> &output,
TestbenchOptions const &options,
Config const &config,
std::string const &name,
std::string const &algo) {
typedef perf::GemmProfiler<float, float, float, float, float> SGemmProfiler;
int results = 0;
{
/*batched sgemm traits*/
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor, OutputTile>
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float,
float,
float,
float,
float,
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
// create a device gemm
typedef typename cutlass::gemm::SplitkPIGemmTraits<SgemmTraits, BatchedReductionTraits> deviceGemmTraits;
typedef typename CutlassDispatchSplitKPIGemmBasic<deviceGemmTraits>::Dispatch Dispatch;
results |= profile_gemm<Dispatch, SGemmProfiler>(output, name + "_nn", options, config, algo + "_splitk_pi");
}
{
/*batched sgemm traits*/
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor, OutputTile>
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float,
float,
float,
float,
float,
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
// create a device gemm
typedef typename cutlass::gemm::SplitkPIGemmTraits<SgemmTraits, BatchedReductionTraits> deviceGemmTraits;
typedef typename CutlassDispatchSplitKPIGemmBasic<deviceGemmTraits>::Dispatch Dispatch;
results |= profile_gemm<Dispatch, SGemmProfiler>(output, name + "_nt", options, config, algo + "_splitk_pi");
}
{
/*batched sgemm traits*/
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor, OutputTile>
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float,
float,
float,
float,
float,
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
// create a device gemm
typedef typename cutlass::gemm::SplitkPIGemmTraits<SgemmTraits, BatchedReductionTraits> deviceGemmTraits;
typedef typename CutlassDispatchSplitKPIGemmBasic<deviceGemmTraits>::Dispatch Dispatch;
results |= profile_gemm<Dispatch, SGemmProfiler>(output, name + "_tn", options, config, algo + "_splitk_pi");
}
{
/*batched sgemm traits*/
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor, OutputTile>
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float,
float,
float,
float,
float,
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
// create a device gemm
typedef typename cutlass::gemm::SplitkPIGemmTraits<SgemmTraits, BatchedReductionTraits> deviceGemmTraits;
typedef typename CutlassDispatchSplitKPIGemmBasic<deviceGemmTraits>::Dispatch Dispatch;
results |= profile_gemm<Dispatch, SGemmProfiler>(output, name + "_tt", options, config, algo + "_splitk_pi");
}
return results;
}
/// Profiles all SGEMM tile sizes
int profile_sgemm_splitkpi(TestbenchOutput<GemmProblem> &output, TestbenchOptions const &options, Config const &config) {
int results = 0;
results |= profile_sgemm_splitkpi_kernel<cutlass::Shape<8, 128, 128>, 32 >(output, options, config, "sgemm_128x128x8_splitk_pi_split32", "128x128");
/*128x64x8*/
results |= profile_sgemm_splitkpi_kernel<cutlass::Shape<8, 64, 128>, 8 >(output, options, config, "sgemm_128x64x8_splitk_pi_split8", "128x64");
results |= profile_sgemm_splitkpi_kernel<cutlass::Shape<8, 64, 128>, 16 >(output, options, config, "sgemm_128x64x8_splitk_pi_split16", "128x64");
results |= profile_sgemm_splitkpi_kernel<cutlass::Shape<8, 64, 128>, 20 >(output, options, config, "sgemm_128x64x8_splitk_pi_split20", "128x64");
results |= profile_sgemm_splitkpi_kernel<cutlass::Shape<8, 64, 128>, 24 >(output, options, config, "sgemm_128x64x8_splitk_pi_split24", "128x64");
results |= profile_sgemm_splitkpi_kernel<cutlass::Shape<8, 64, 128>, 28 >(output, options, config, "sgemm_128x64x8_splitk_pi_split28", "128x64");
results |= profile_sgemm_splitkpi_kernel<cutlass::Shape<8, 64, 128>, 32 >(output, options, config, "sgemm_128x64x8_splitk_pi_split32", "128x64");
results |= profile_sgemm_splitkpi_kernel<cutlass::Shape<8, 64, 128>, 64 >(output, options, config, "sgemm_128x64x8_splitk_pi_split64", "128x64");
/*128x32x8*/
results |= profile_sgemm_splitkpi_kernel<cutlass::Shape<8, 32, 128>, 8 >(output, options, config, "sgemm_128x32x8_splitk_pi_split8", "128x32");
results |= profile_sgemm_splitkpi_kernel<cutlass::Shape<8, 32, 128>, 16 >(output, options, config, "sgemm_128x32x8_splitk_pi_split16", "128x32");
results |= profile_sgemm_splitkpi_kernel<cutlass::Shape<8, 32, 128>, 20 >(output, options, config, "sgemm_128x32x8_splitk_pi_split20", "128x32");
results |= profile_sgemm_splitkpi_kernel<cutlass::Shape<8, 32, 128>, 24 >(output, options, config, "sgemm_128x32x8_splitk_pi_split24", "128x32");
results |= profile_sgemm_splitkpi_kernel<cutlass::Shape<8, 32, 128>, 28 >(output, options, config, "sgemm_128x32x8_splitk_pi_split28", "128x32");
results |= profile_sgemm_splitkpi_kernel<cutlass::Shape<8, 32, 128>, 32 >(output, options, config, "sgemm_128x32x8_splitk_pi_split32", "128x32");
results |= profile_sgemm_splitkpi_kernel<cutlass::Shape<8, 32, 128>, 64 >(output, options, config, "sgemm_128x32x8_splitk_pi_split64", "128x32");
return results;
}
struct SgemmSplitKPIRegistrar {
SgemmSplitKPIRegistrar() { RegisterGemmProfileFunc(profile_sgemm_splitkpi); }
};
volatile SgemmSplitKPIRegistrar _SgemmSplitKPIRegistrar;
////////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace perf

View File

@@ -76,6 +76,15 @@ struct WmmaBinaryGemmDispatch {
params.initialize(m, n, k * 32, alpha, d_a, lda, d_b, ldb, beta, d_c, ldc, d_d, ldd);
}
/// batched strided bmma
WmmaBinaryGemmDispatch(int m, int n, int k, int alpha,
cutlass::Vector<cutlass::bin1_t, 32> const* d_a, int lda, long long int batch_stride_a,
cutlass::Vector<cutlass::bin1_t, 32> const* d_b, int ldb, long long int batch_stride_b, int beta,
int const* d_c, int ldc, long long int batch_stride_c, int* d_d, int ldd, long long int batch_stride_d,
int batch_count) {
assert(0);
}
/// Initializes params object
WmmaBinaryGemmDispatch(Params const& _params) : params(_params) {}

View File

@@ -1,27 +1,27 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#include "cutlass/wmma_matrix.h"
#ifdef CUTLASS_USE_WMMA_API
@@ -92,6 +92,31 @@ struct WmmaGemmDispatch {
params.initialize(m, n, k, alpha, d_a, lda, d_b, ldb, beta, d_c, ldc, d_d, ldd);
}
WmmaGemmDispatch(int m,
int n,
int k,
Scalar alpha,
ScalarA const* d_a,
int lda,
long long int batch_stride_A,
ScalarB const* d_b,
int ldb,
long long int batch_stride_B,
Scalar beta,
ScalarC const* d_c,
int ldc,
long long int batch_stride_C,
ScalarD* d_d,
int ldd,
long long int batch_stride_D,
int batch_count) {
params.initialize(m, n, k, alpha, d_a, lda, batch_stride_A,
d_b, ldb, batch_stride_B,
beta, d_c, ldc, batch_stride_C,
d_d, ldd, batch_stride_D,
batch_count);
}
/// Initializes params object
WmmaGemmDispatch(Params const& _params) : params(_params) {}
@@ -105,6 +130,7 @@ namespace perf {
////////////////////////////////////////////////////////////////////////////////////////////////////
template <typename DummyT>
int profile_wmma_gemm_f32(TestbenchOutput<GemmProblem> &output, TestbenchOptions const &options, Config const &config) {
typedef perf::GemmProfiler<cutlass::half_t, cutlass::half_t, float, float, float> GemmProfiler;
@@ -112,8 +138,8 @@ int profile_wmma_gemm_f32(TestbenchOutput<GemmProblem> &output, TestbenchOptions
{
typedef cutlass::gemm::WmmaGemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor>
WmmaGemmTraits;
cutlass::MatrixLayout::kRowMajor>
WmmaGemmTraits;
typedef WmmaGemmDispatch<WmmaGemmTraits> Dispatch;
@@ -122,8 +148,8 @@ int profile_wmma_gemm_f32(TestbenchOutput<GemmProblem> &output, TestbenchOptions
{
typedef cutlass::gemm::WmmaGemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor>
WmmaGemmTraits;
cutlass::MatrixLayout::kColumnMajor>
WmmaGemmTraits;
typedef WmmaGemmDispatch<WmmaGemmTraits> Dispatch;
@@ -132,7 +158,7 @@ int profile_wmma_gemm_f32(TestbenchOutput<GemmProblem> &output, TestbenchOptions
{
typedef cutlass::gemm::WmmaGemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor>
cutlass::MatrixLayout::kColumnMajor>
WmmaGemmTraits;
typedef WmmaGemmDispatch<WmmaGemmTraits> Dispatch;
@@ -142,7 +168,7 @@ int profile_wmma_gemm_f32(TestbenchOutput<GemmProblem> &output, TestbenchOptions
{
typedef cutlass::gemm::WmmaGemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor>
cutlass::MatrixLayout::kRowMajor>
WmmaGemmTraits;
typedef WmmaGemmDispatch<WmmaGemmTraits> Dispatch;
@@ -155,10 +181,11 @@ int profile_wmma_gemm_f32(TestbenchOutput<GemmProblem> &output, TestbenchOptions
////////////////////////////////////////////////////////////////////////////////////////////////////
template <typename DummyT>
int profile_wmma_gemm_f16(
TestbenchOutput<GemmProblem> &output,
TestbenchOptions const &options,
Config const &config) {
TestbenchOutput<GemmProblem> &output,
TestbenchOptions const &options,
Config const &config) {
typedef perf::GemmProfiler<
cutlass::half_t,
@@ -173,7 +200,7 @@ int profile_wmma_gemm_f16(
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 128, 128>,
cutlass::Shape<32, 256, 128>,
half,
half,
half,
@@ -192,7 +219,7 @@ int profile_wmma_gemm_f16(
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 128, 128>,
cutlass::Shape<32, 256, 128>,
half,
half,
half,
@@ -211,7 +238,7 @@ int profile_wmma_gemm_f16(
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 128, 128>,
cutlass::Shape<32, 256, 128>,
half,
half,
half,
@@ -230,7 +257,7 @@ int profile_wmma_gemm_f16(
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 128, 128>,
cutlass::Shape<32, 256, 128>,
half,
half,
half,
@@ -248,12 +275,283 @@ int profile_wmma_gemm_f16(
return results;
}
template <typename DummyT>
int profile_wmma_4_gemm_f16(
TestbenchOutput<GemmProblem> &output,
TestbenchOptions const &options,
Config const &config) {
typedef perf::GemmProfiler<
cutlass::half_t,
cutlass::half_t,
cutlass::half_t,
cutlass::half_t,
cutlass::half_t> GemmProfiler;
int results = 0;
// a set of test requires leading dim to be multiple of 4 instead of 8
{
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
half,
cutlass::gemm::LinearScaling<half>,
half,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
4, /*kScalarsPerLdgA_*/
4, /*kScalarsPerLdgB_*/
4, /*KScalarsPerLdsA_*/
4, /*KScalarsPerLdsB_*/
4 / sizeof(half), /*kScalarsPerLdgCAndStgD_*/
4 / sizeof(half), /*kScalarsPerStsD_*/
4 / sizeof(half) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
typedef WmmaGemmDispatch<WmmaGemmTraits> Dispatch;
results |= profile_gemm<Dispatch, GemmProfiler>(output, "wmma_4_gemm_f16_nt", options, config);
}
{
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
half,
cutlass::gemm::LinearScaling<half>,
half,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
4, /*kScalarsPerLdgA_*/
4, /*kScalarsPerLdgB_*/
4, /*KScalarsPerLdsA_*/
4, /*KScalarsPerLdsB_*/
4 / sizeof(half), /*kScalarsPerLdgCAndStgD_*/
4 / sizeof(half), /*kScalarsPerStsD_*/
4 / sizeof(half) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
typedef WmmaGemmDispatch<WmmaGemmTraits> Dispatch;
results |= profile_gemm<Dispatch, GemmProfiler>(output, "wmma_4_gemm_f16_nn", options, config);
}
{
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
half,
cutlass::gemm::LinearScaling<half>,
half,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
4, /*kScalarsPerLdgA_*/
4, /*kScalarsPerLdgB_*/
4, /*KScalarsPerLdsA_*/
4, /*KScalarsPerLdsB_*/
4 / sizeof(half), /*kScalarsPerLdgCAndStgD_*/
4 / sizeof(half), /*kScalarsPerStsD_*/
4 / sizeof(half) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
typedef WmmaGemmDispatch<WmmaGemmTraits> Dispatch;
results |= profile_gemm<Dispatch, GemmProfiler>(output, "wmma_4_gemm_f16_tn", options, config);
}
{
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
half,
cutlass::gemm::LinearScaling<half>,
half,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
4, /*kScalarsPerLdgA_*/
4, /*kScalarsPerLdgB_*/
4, /*KScalarsPerLdsA_*/
4, /*KScalarsPerLdsB_*/
4 / sizeof(half), /*kScalarsPerLdgCAndStgD_*/
4 / sizeof(half), /*kScalarsPerStsD_*/
4 / sizeof(half) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
typedef WmmaGemmDispatch<WmmaGemmTraits> Dispatch;
results |= profile_gemm<Dispatch, GemmProfiler>(output, "wmma_4_gemm_f16_tt", options, config);
}
return results;
}
template <typename DummyT>
int profile_wmma_4_fp16_sgemm_fp16(
TestbenchOutput<GemmProblem> &output,
TestbenchOptions const &options,
Config const &config) {
typedef perf::GemmProfiler<
cutlass::half_t,
cutlass::half_t,
cutlass::half_t,
float,
float> GemmProfiler;
int results = 0;
// a set of test requires leading dim to be multiple of 4 instead of 8
{
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
4, /*kScalarsPerLdgA_*/
4, /*kScalarsPerLdgB_*/
4, /*KScalarsPerLdsA_*/
4, /*KScalarsPerLdsB_*/
8 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
8 / sizeof(accumu_type), /*kScalarsPerStsD_*/
8 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
typedef WmmaGemmDispatch<WmmaGemmTraits> Dispatch;
results |= profile_gemm<Dispatch, GemmProfiler>(output, "wmma_4_fp16_sgemm_fp16_nt", options, config);
}
{
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
4, /*kScalarsPerLdgA_*/
4, /*kScalarsPerLdgB_*/
4, /*KScalarsPerLdsA_*/
4, /*KScalarsPerLdsB_*/
8 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
8 / sizeof(accumu_type), /*kScalarsPerStsD_*/
8 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
typedef WmmaGemmDispatch<WmmaGemmTraits> Dispatch;
results |= profile_gemm<Dispatch, GemmProfiler>(output, "wmma_4_fp16_sgemm_fp16_nn", options, config);
}
{
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
4, /*kScalarsPerLdgA_*/
4, /*kScalarsPerLdgB_*/
4, /*KScalarsPerLdsA_*/
4, /*KScalarsPerLdsB_*/
8 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
8 / sizeof(accumu_type), /*kScalarsPerStsD_*/
8 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
typedef WmmaGemmDispatch<WmmaGemmTraits> Dispatch;
results |= profile_gemm<Dispatch, GemmProfiler>(output, "wmma_4_fp16_sgemm_fp16_tn", options, config);
}
{
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
4, /*kScalarsPerLdgA_*/
4, /*kScalarsPerLdgB_*/
4, /*KScalarsPerLdsA_*/
4, /*KScalarsPerLdsB_*/
8 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
8 / sizeof(accumu_type), /*kScalarsPerStsD_*/
8 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
typedef WmmaGemmDispatch<WmmaGemmTraits> Dispatch;
results |= profile_gemm<Dispatch, GemmProfiler>(output, "wmma_4_fp16_sgemm_fp16_tt", options, config);
}
return results;
}
////////////////////////////////////////////////////////////////////////////////////////////////////
struct WmmaGemmRegistrar {
WmmaGemmRegistrar() {
RegisterGemmProfileFunc(profile_wmma_gemm_f32);
RegisterGemmProfileFunc(profile_wmma_gemm_f16);
RegisterGemmProfileFunc(profile_wmma_gemm_f32<void>);
RegisterGemmProfileFunc(profile_wmma_gemm_f16<void>);
//#ifdef EXHAUSTIVE_PROF
RegisterGemmProfileFunc(profile_wmma_4_gemm_f16<void>);
//fp32 accum with fp16 input and output
RegisterGemmProfileFunc(profile_wmma_4_fp16_sgemm_fp16<void>);
//#endif // defined EXHAUSTIVE_PROF
}
};
@@ -266,3 +564,4 @@ volatile WmmaGemmRegistrar _WmmaGemmRegistrar;
////////////////////////////////////////////////////////////////////////////////////////////////////
#endif // defined CUTLASS_USE_WMMA_API

View File

@@ -74,6 +74,15 @@ struct WmmaIntegerGemmDispatch {
params.initialize(m, n, k, alpha, d_a, lda, d_b, ldb, beta, d_c, ldc, d_d, ldd);
}
///
WmmaIntegerGemmDispatch(int m, int n, int k, int alpha,
ScalarA const* d_a, int lda, long long int batch_stride_a,
ScalarB const* d_b, int ldb, long long int batch_stride_b, int beta,
int const* d_c, int ldc, long long int batch_stride_c, int* d_d, int ldd, long long int batch_stride_d,
int batch_count) {
assert(0);
}
/// Initializes params object
WmmaIntegerGemmDispatch(Params const& _params) : params(_params) {}
@@ -125,6 +134,15 @@ struct WmmaIntegerGemmDispatch<Traits,
params.initialize(m, n, k * 8, alpha, d_a, lda, d_b, ldb, beta, d_c, ldc, d_d, ldd);
}
///
WmmaIntegerGemmDispatch(int m, int n, int k, int alpha,
ScalarA const* d_a, int lda, long long int batch_stride_a,
ScalarB const* d_b, int ldb, long long int batch_stride_b, int beta,
int const* d_c, int ldc, long long int batch_stride_c, int* d_d, int ldd, long long int batch_stride_d,
int batch_count) {
assert(0);
}
/// Initializes params object
WmmaIntegerGemmDispatch(Params const& _params) : params(_params) {}
@@ -176,6 +194,15 @@ struct WmmaIntegerGemmDispatch<Traits,
params.initialize(m, n, k * 8, alpha, d_a, lda, d_b, ldb, beta, d_c, ldc, d_d, ldd);
}
///
WmmaIntegerGemmDispatch(int m, int n, int k, int alpha,
ScalarA const* d_a, int lda, long long int batch_stride_a,
ScalarB const* d_b, int ldb, long long int batch_stride_b, int beta,
int const* d_c, int ldc, long long int batch_stride_c, int* d_d, int ldd, long long int batch_stride_d,
int batch_count) {
assert(0);
}
/// Initializes params object
WmmaIntegerGemmDispatch(Params const& _params) : params(_params) {}

View File

@@ -24,7 +24,7 @@
**************************************************************************************************/
#pragma once
#include <assert.h>
#include "cutlass/matrix_traits.h"
#include "tools/util/command_line.h"
#include "tools/test/perf/provider.h"
@@ -85,6 +85,7 @@ struct GemmProblem {
int m;
int n;
int k;
int batch_count;
cutlass::MatrixLayout::Kind layout_A;
cutlass::MatrixLayout::Kind layout_B;
@@ -96,7 +97,7 @@ struct GemmProblem {
//
/// Static method to print GemmProblem headers
static std::string header() { return "M,N,K,Layout_A,Layout_B,Beta"; }
static std::string header() { return "M,N,K,Layout_A,Layout_B,Beta,batch_count"; }
//
// Methods
@@ -108,21 +109,24 @@ struct GemmProblem {
cutlass::MatrixLayout::Kind _layout_A = cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::Kind _layout_B = cutlass::MatrixLayout::kRowMajor,
double _alpha = 1,
double _beta = 0)
: m(_m), n(_n), k(_k), layout_A(_layout_A), layout_B(_layout_B), alpha(_alpha), beta(_beta) {}
double _beta = 0,
int _batch_count = 1)
: m(_m), n(_n), k(_k), layout_A(_layout_A), layout_B(_layout_B), alpha(_alpha), beta(_beta), batch_count(_batch_count) {
assert(batch_count >= 1);
}
/// leading dimension of A
int lda() const {
if (layout_A == cutlass::MatrixLayout::kColumnMajor) {
return m;
}
return k;
return k * batch_count;
}
/// leading dimension of B
int ldb() const {
if (layout_B == cutlass::MatrixLayout::kColumnMajor) {
return k;
return k * batch_count;
}
return n;
}
@@ -130,10 +134,35 @@ struct GemmProblem {
/// leading dimension of C
int ldc() const { return m; }
/// batch_stride_a. only makes sense when batch_count > 1
long long int batch_stride_a() const {
assert(batch_count > 1);
if (layout_A == cutlass::MatrixLayout::kColumnMajor) {
return static_cast<long long int>(k) * static_cast<long long int>(lda());
}
return static_cast<long long int>(k);
}
/// batch_stride_b. only makes sense when batch_count > 1
long long int batch_stride_b() const {
assert(batch_count > 1);
if (layout_B == cutlass::MatrixLayout::kColumnMajor) {
return static_cast<long long int>(k);
}
return static_cast<long long int>(k) * static_cast<long long int>(ldb());
}
/// batch_stride_c. only makes sense when batch_count > 1
long long int batch_stride_c() const {
assert(batch_count > 1);
return static_cast<long long int>(n) * static_cast<long long int>(ldc());
}
/// Pretty prints output
std::ostream &pretty_print(std::ostream &out) const {
out << m << "-by-" << n << "-by-" << k << ", A: " << layout_A << "-major, B: " << layout_B
<< "-major, beta: " << beta;
<< "-major, beta: " << beta << ", batch: " << batch_count;
return out;
}
@@ -142,7 +171,7 @@ struct GemmProblem {
/// Prints a problem to an output stream
inline std::ostream &operator<<(std::ostream &out, GemmProblem const &problem) {
out << problem.m << "," << problem.n << "," << problem.k << "," << problem.layout_A << ","
<< problem.layout_B << "," << problem.beta;
<< problem.layout_B << "," << problem.beta << "," << problem.batch_count;
return out;
}

View File

@@ -125,13 +125,16 @@ struct GemmProblemRange {
/// Range of sizes in GEMM K dimension
Range K;
/// Range of sizes in batch dimeion
Range batch_count;
//
// Methods
//
/// Constructor to define a space of probelm sizes
GemmProblemRange(Range _M = Range(256), Range _N = Range(256), Range _K = Range(256))
: M(_M), N(_N), K(_K) {}
GemmProblemRange(Range _M = Range(256), Range _N = Range(256), Range _K = Range(256), Range _batch_count = Range(1))
: M(_M), N(_N), K(_K), batch_count(_batch_count) {}
/// Parses a command line argument as a Range object
static void get_range(Range &range,
@@ -155,6 +158,7 @@ struct GemmProblemRange {
get_range(M, args, "m", Range(10240));
get_range(N, args, "n", Range(4096));
get_range(K, args, "k", Range(4096));
get_range(batch_count, args, "batch", Range(1));
}
};
@@ -368,7 +372,7 @@ struct TestbenchOptions {
/// Number of iterations
int iterations;
/// Defines how to run the benchmark
ExecutionMode::Kind execution_mode;
@@ -599,6 +603,9 @@ struct TestbenchOptions {
<< " --k=<depth>[:max depth[:step]] "
<< " Size of inner dimension of A and B. May specify a range with optional step size.\n"
<< " --batch=<batch> "
<< " Number of batches for a bached gemm. "
<< " --kernels=<{s|d|h|i|wmma_|wmma_binary_|wmma_integer_}gemm_{nn,nt,tn,tt}>\n"
<< " "
<< " Select GEMM datatype and layout to use for tests\n"

View File

@@ -39,10 +39,18 @@ set(CUTLASS_UNIT_TEST_HEADERS
core/layout_verification.h
gemm/run_gemm.h
gemm/gemm_testbed.h
reduction/batched_reduction_testbed.h
reduction/test_batched_reduction.h
)
set(CUTLASS_UNIT_TEST_SOURCES_BACKUP
cutlass_unit_test.cpp
gemm/batched_strided_sgemm_128x128x8.cu
)
set(CUTLASS_UNIT_TEST_SOURCES
cutlass_unit_test.cpp
tile_iterator_test.cu
core/tensor_ref.cu
core/tensor_view.cu
util/unique_ptr.cu
@@ -80,6 +88,9 @@ set(CUTLASS_UNIT_TEST_SOURCES
gemm/fp16_sgemm_fp32_128x128x16.cu
gemm/fp16_sgemm_fp16_128x128x16.cu
gemm/wmma_gemm.cu
gemm/fp16_wmma_gemm_fp16.cu
gemm/wmma_gemm_non_multiple16.cu
gemm/fp16_wmma_gemm_fp16_non_multiple16.cu
gemm/wmma_binary_gemm.cu
gemm/wmma_integer_gemm.cu
gemm/sgemm_threadblock_swizzle_nn.cu
@@ -89,7 +100,18 @@ set(CUTLASS_UNIT_TEST_SOURCES
gemm/batched_strided_sgemm_128x128x8.cu
gemm/batched_strided_dgemm_128x128x8.cu
gemm/batched_strided_hgemm_128x128x8.cu
gemm/batched_strided_wmma_gemm.cu
gemm/batched_strided_fp16_wmma_gemm_fp16.cu
gemm/epilogue_functor.cu
reduction/batched_reduction.cu
reduction/mixed_batched_reduction.cu
gemm/splitK_sgemm.cu
gemm/splitK_igemm.cu
gemm/splitK_fp16_sgemm_fp16.cu
gemm/splitK_dgemm.cu
gemm/splitK_hgemm.cu
gemm/splitK_wmma_gemm.cu
gemm/partitionedK_sgemm_128x128x8.cu
)
if (CUTLASS_NVRTC_ENABLE)

View File

@@ -124,120 +124,120 @@ TEST(PredicateVector, Count) {
{
typedef cutlass::PredicateVector<4, 8> PredicateVector;
EXPECT_EQ(int(PredicateVector::kWordCount), 1)
<< "PredicateVector<4, 8> word count: " << PredicateVector::kWordCount;
<< "PredicateVector<4, 8> word count: " << int(PredicateVector::kWordCount);
}
{
typedef cutlass::PredicateVector<4, 4> PredicateVector;
EXPECT_EQ(int(PredicateVector::kWordCount), 1)
<< "PredicateVector<4, 4> word count: " << PredicateVector::kWordCount;
<< "PredicateVector<4, 4> word count: " << int(PredicateVector::kWordCount);
}
{
typedef cutlass::PredicateVector<4, 2> PredicateVector;
EXPECT_EQ(int(PredicateVector::kWordCount), 1)
<< "PredicateVector<4, 2> word count: " << PredicateVector::kWordCount;
<< "PredicateVector<4, 2> word count: " << int(PredicateVector::kWordCount);
}
{
typedef cutlass::PredicateVector<4, 1> PredicateVector;
EXPECT_EQ(int(PredicateVector::kWordCount), 1)
<< "PredicateVector<4, 1> word count: " << PredicateVector::kWordCount;
<< "PredicateVector<4, 1> word count: " << int(PredicateVector::kWordCount);
}
{
typedef cutlass::PredicateVector<8, 8> PredicateVector;
EXPECT_EQ(int(PredicateVector::kWordCount), 1)
<< "PredicateVector<8, 8> word count: " << PredicateVector::kWordCount;
<< "PredicateVector<8, 8> word count: " << int(PredicateVector::kWordCount);
}
{
typedef cutlass::PredicateVector<8, 4> PredicateVector;
EXPECT_EQ(int(PredicateVector::kWordCount), 1)
<< "PredicateVector<8, 4> word count: " << PredicateVector::kWordCount;
<< "PredicateVector<8, 4> word count: " << int(PredicateVector::kWordCount);
}
{
typedef cutlass::PredicateVector<8, 2> PredicateVector;
EXPECT_EQ(int(PredicateVector::kWordCount), 1)
<< "PredicateVector<8, 2> word count: " << PredicateVector::kWordCount;
<< "PredicateVector<8, 2> word count: " << int(PredicateVector::kWordCount);
}
{
typedef cutlass::PredicateVector<8, 1> PredicateVector;
EXPECT_EQ(int(PredicateVector::kWordCount), 2)
<< "PredicateVector<8, 1> word count: " << PredicateVector::kWordCount;
<< "PredicateVector<8, 1> word count: " << int(PredicateVector::kWordCount);
}
{
typedef cutlass::PredicateVector<16, 8> PredicateVector;
EXPECT_EQ(int(PredicateVector::kWordCount), 1)
<< "PredicateVector<16, 8> word count: " << PredicateVector::kWordCount;
<< "PredicateVector<16, 8> word count: " << int(PredicateVector::kWordCount);
}
{
typedef cutlass::PredicateVector<16, 4> PredicateVector;
EXPECT_EQ(int(PredicateVector::kWordCount), 1)
<< "PredicateVector<16, 4> word count: " << PredicateVector::kWordCount;
<< "PredicateVector<16, 4> word count: " << int(PredicateVector::kWordCount);
}
{
typedef cutlass::PredicateVector<16, 2> PredicateVector;
EXPECT_EQ(int(PredicateVector::kWordCount), 2)
<< "PredicateVector<16, 2> word count: " << PredicateVector::kWordCount;
<< "PredicateVector<16, 2> word count: " << int(PredicateVector::kWordCount);
}
{
typedef cutlass::PredicateVector<16, 1> PredicateVector;
EXPECT_EQ(int(PredicateVector::kWordCount), 4)
<< "PredicateVector<16, 1> word count: " << PredicateVector::kWordCount;
<< "PredicateVector<16, 1> word count: " << int(PredicateVector::kWordCount);
}
{
typedef cutlass::PredicateVector<32, 8> PredicateVector;
EXPECT_EQ(int(PredicateVector::kWordCount), 1)
<< "PredicateVector<32, 8> word count: " << PredicateVector::kWordCount;
<< "PredicateVector<32, 8> word count: " << int(PredicateVector::kWordCount);
}
{
typedef cutlass::PredicateVector<32, 4> PredicateVector;
EXPECT_EQ(int(PredicateVector::kWordCount), 2)
<< "PredicateVector<32, 4> word count: " << PredicateVector::kWordCount;
<< "PredicateVector<32, 4> word count: " << int(PredicateVector::kWordCount);
}
{
typedef cutlass::PredicateVector<32, 2> PredicateVector;
EXPECT_EQ(int(PredicateVector::kWordCount), 4)
<< "PredicateVector<32, 2> word count: " << PredicateVector::kWordCount;
<< "PredicateVector<32, 2> word count: " << int(PredicateVector::kWordCount);
}
{
typedef cutlass::PredicateVector<32, 1> PredicateVector;
EXPECT_EQ(int(PredicateVector::kWordCount), 8)
<< "PredicateVector<32, 1> word count: " << PredicateVector::kWordCount;
<< "PredicateVector<32, 1> word count: " << int(PredicateVector::kWordCount);
}
{
typedef cutlass::PredicateVector<64, 8> PredicateVector;
EXPECT_EQ(int(PredicateVector::kWordCount), 2)
<< "PredicateVector<64, 8> word count: " << PredicateVector::kWordCount;
<< "PredicateVector<64, 8> word count: " << int(PredicateVector::kWordCount);
}
{
typedef cutlass::PredicateVector<64, 4> PredicateVector;
EXPECT_EQ(int(PredicateVector::kWordCount), 4)
<< "PredicateVector<64, 4> word count: " << PredicateVector::kWordCount;
<< "PredicateVector<64, 4> word count: " << int(PredicateVector::kWordCount);
}
{
typedef cutlass::PredicateVector<64, 2> PredicateVector;
EXPECT_EQ(int(PredicateVector::kWordCount), 8)
<< "PredicateVector<64, 2> word count: " << PredicateVector::kWordCount;
<< "PredicateVector<64, 2> word count: " << int(PredicateVector::kWordCount);
}
{
typedef cutlass::PredicateVector<64, 1> PredicateVector;
EXPECT_EQ(int(PredicateVector::kWordCount), 16)
<< "PredicateVector<64, 1> word count: " << PredicateVector::kWordCount;
<< "PredicateVector<64, 1> word count: " << int(PredicateVector::kWordCount);
}
}

View File

@@ -64,15 +64,30 @@ void set_gtest_flag() {
/// If true, the tests are enabled strictly for one compute capability
bool experimental;
} test_filters[] = {
{ "Sgemm*", 50, false },
{ "Dgemm*", 60, false },
{ "Fp16_sgemm*", 60, false },
{ "Hgemm*", 60, false },
{ "Igemm*", 61, false },
{ "WmmaGemm*", 70, false },
{ "WmmaInt8*", 72, false },
{ "WmmaInt4*", 75, true },
{ "WmmaBinary*", 75, true },
{ "Sgemm*", 50, false },
{ "*sgemm*", 50, false },
{ "Dgemm*", 60, false },
{ "*dgemm*", 60, false },
{ "Fp16_sgemm*", 60, false },
{ "*fp16_sgemm*", 60, false },
{ "Batched_reduction*", 60, false },
{ "*batched_reduction*", 60, false },
{ "Float_batched_reduction*", 60, false },
{ "*float_batched_reduction*", 60, false },
{ "SplitK*", 60, false },
{ "*splitK*", 60, false },
{ "Hgemm*", 60, false },
{ "*hgemm*", 60, false },
{ "Igemm*", 61, false },
{ "*igemm*", 61, false },
{ "WmmaGemm*", 70, false },
{ "*wmma*", 70, false },
{ "WmmaInt8*", 72, false },
{ "*wmmaInt8*", 72, false },
{ "WmmaInt4*", 75, true },
{ "*wmmaInt4*", 75, true },
{ "WmmaBinary*", 75, true },
{ "*wmmaBinary*", 75, true },
{ 0, 0, false }
};

View File

@@ -0,0 +1,385 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#include "cutlass/wmma_matrix.h"
#if defined(CUTLASS_USE_WMMA_API)
#include "cutlass_unit_test.h"
#include "cutlass/gemm/gemm.h"
#include "cutlass/gemm/wmma_gemm_traits.h"
#include "tools/test/unit/gemm/gemm_testbed.h"
#include "tools/test/unit/gemm/run_gemm.h"
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_strided_batched_16x16x32_f32, fp16_wmma_gemm_fp16_32x32x16_nn) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
8, /*kScalarsPerLdgA_*/
8, /*kScalarsPerLdgB_*/
8, /*KScalarsPerLdsA_*/
8, /*KScalarsPerLdsB_*/
16 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
16 / sizeof(accumu_type), /*kScalarsPerStsD_*/
16 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_batched_strided_gemm<WmmaGemmTraits>(32, 32, 64, 3);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_strided_batched_16x16x32_f32, fp16_wmma_gemm_fp16_32x32x16_nt) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
8, /*kScalarsPerLdgA_*/
8, /*kScalarsPerLdgB_*/
8, /*KScalarsPerLdsA_*/
8, /*KScalarsPerLdsB_*/
16 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
16 / sizeof(accumu_type), /*kScalarsPerStsD_*/
16 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_batched_strided_gemm<WmmaGemmTraits>(32, 32, 64, 3);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_strided_batched_16x16x32_f32, fp16_wmma_gemm_fp16_32x32x16_tn) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
8, /*kScalarsPerLdgA_*/
8, /*kScalarsPerLdgB_*/
8, /*KScalarsPerLdsA_*/
8, /*KScalarsPerLdsB_*/
16 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
16 / sizeof(accumu_type), /*kScalarsPerStsD_*/
16 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_batched_strided_gemm<WmmaGemmTraits>(32, 32, 64, 3);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_strided_batched_16x16x32_f32, fp16_wmma_gemm_fp16_32x32x16_tt) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
8, /*kScalarsPerLdgA_*/
8, /*kScalarsPerLdgB_*/
8, /*KScalarsPerLdsA_*/
8, /*KScalarsPerLdsB_*/
16 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
16 / sizeof(accumu_type), /*kScalarsPerStsD_*/
16 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_batched_strided_gemm<WmmaGemmTraits>(32, 32, 64, 3);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
//mulitple of 4
TEST(WmmaGemm_strided_batched_16x16x32_f32, fp16_wmma_gemm_fp16_36x36x16_nn) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
4, /*kScalarsPerLdgA_*/
4, /*kScalarsPerLdgB_*/
4, /*KScalarsPerLdsA_*/
4, /*KScalarsPerLdsB_*/
8 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
8 / sizeof(accumu_type), /*kScalarsPerStsD_*/
8 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_batched_strided_gemm<WmmaGemmTraits>(36, 36, 64, 3);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_strided_batched_16x16x32_f32, fp16_wmma_gemm_fp16_36x36x16_nt) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
4, /*kScalarsPerLdgA_*/
4, /*kScalarsPerLdgB_*/
4, /*KScalarsPerLdsA_*/
4, /*KScalarsPerLdsB_*/
8 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
8 / sizeof(accumu_type), /*kScalarsPerStsD_*/
8 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_batched_strided_gemm<WmmaGemmTraits>(36, 36, 64, 3);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_strided_batched_16x16x32_f32, fp16_wmma_gemm_fp16_36x36x16_tn) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
4, /*kScalarsPerLdgA_*/
4, /*kScalarsPerLdgB_*/
4, /*KScalarsPerLdsA_*/
4, /*KScalarsPerLdsB_*/
8 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
8 / sizeof(accumu_type), /*kScalarsPerStsD_*/
8 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_batched_strided_gemm<WmmaGemmTraits>(36, 36, 64, 3);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_strided_batched_16x16x32_f32, fp16_wmma_gemm_fp16_36x36x16_tt) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
4, /*kScalarsPerLdgA_*/
4, /*kScalarsPerLdgB_*/
4, /*KScalarsPerLdsA_*/
4, /*KScalarsPerLdsB_*/
8 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
8 / sizeof(accumu_type), /*kScalarsPerStsD_*/
8 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_batched_strided_gemm<WmmaGemmTraits>(36, 36, 64, 3);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
//mulitple of 2
TEST(WmmaGemm_strided_batched_16x16x32_f32, fp16_wmma_gemm_fp16_34x34x16_nn) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
2, /*kScalarsPerLdgA_*/
2, /*kScalarsPerLdgB_*/
2, /*KScalarsPerLdsA_*/
2, /*KScalarsPerLdsB_*/
4 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
4 / sizeof(accumu_type), /*kScalarsPerStsD_*/
4 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_batched_strided_gemm<WmmaGemmTraits>(34, 34, 64, 3);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_strided_batched_16x16x32_f32, fp16_wmma_gemm_fp16_34x34x16_nt) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
2, /*kScalarsPerLdgA_*/
2, /*kScalarsPerLdgB_*/
2, /*KScalarsPerLdsA_*/
2, /*KScalarsPerLdsB_*/
4 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
4 / sizeof(accumu_type), /*kScalarsPerStsD_*/
4 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_batched_strided_gemm<WmmaGemmTraits>(34, 34, 64, 3);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_strided_batched_16x16x32_f32, fp16_wmma_gemm_fp16_34x34x16_tn) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
2, /*kScalarsPerLdgA_*/
2, /*kScalarsPerLdgB_*/
2, /*KScalarsPerLdsA_*/
2, /*KScalarsPerLdsB_*/
4 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
4 / sizeof(accumu_type), /*kScalarsPerStsD_*/
4 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_batched_strided_gemm<WmmaGemmTraits>(34, 34, 64, 3);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_strided_batched_16x16x32_f32, fp16_wmma_gemm_fp16_34x34x16_tt) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
2, /*kScalarsPerLdgA_*/
2, /*kScalarsPerLdgB_*/
2, /*KScalarsPerLdsA_*/
2, /*KScalarsPerLdsB_*/
4 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
4 / sizeof(accumu_type), /*kScalarsPerStsD_*/
4 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_batched_strided_gemm<WmmaGemmTraits>(34, 34, 64, 3);
}
#endif

View File

@@ -34,6 +34,7 @@ TEST(Sgemm_strided_batched_128x128x8, sgemm_256x384x64x3_nn) {
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
//think about using run_gemm directly
run_batched_strided_gemm<SgemmTraits>(256/*m*/, 384/*n*/, 64/*k*/, 3 /*batch_size*/);
}
@@ -43,6 +44,7 @@ TEST(Sgemm_strided_batched_128x128x8, sgemm_128x384x192x2_nn) {
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
//think about using run_gemm directly
run_batched_strided_gemm<SgemmTraits>(128/*m*/, 384/*n*/, 192/*k*/, 2 /*batch_size*/);
}
@@ -52,6 +54,7 @@ TEST(Sgemm_strided_batched_128x128x8, sgemm_127x384x192x2_nn) {
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
//think about using run_gemm directly
run_batched_strided_gemm<SgemmTraits>(127/*m*/, 384/*n*/, 192/*k*/, 2 /*batch_size*/);
}
@@ -61,6 +64,7 @@ TEST(Sgemm_strided_batched_128x128x8, sgemm_127x388x190x2_nn) {
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
//think about using run_gemm directly
run_batched_strided_gemm<SgemmTraits>(127/*m*/, 388/*n*/, 190/*k*/, 2 /*batch_size*/);
}
@@ -70,6 +74,7 @@ TEST(Sgemm_strided_batched_128x128x8, sgemm_256x384x64x3_nt) {
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
//think about using run_gemm directly
run_batched_strided_gemm<SgemmTraits>(256/*m*/, 384/*n*/, 64/*k*/, 3 /*batch_size*/);
}
@@ -79,6 +84,7 @@ TEST(Sgemm_strided_batched_128x128x8, sgemm_128x384x192x2_nt) {
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
//think about using run_gemm directly
run_batched_strided_gemm<SgemmTraits>(128/*m*/, 384/*n*/, 192/*k*/, 2 /*batch_size*/);
}
@@ -90,6 +96,7 @@ TEST(Sgemm_strided_batched_128x128x8, sgemm_256x384x64x3_tn) {
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
//think about using run_gemm directly
run_batched_strided_gemm<SgemmTraits>(256/*m*/, 384/*n*/, 64/*k*/, 3 /*batch_size*/);
}
@@ -99,6 +106,7 @@ TEST(Sgemm_strided_batched_128x128x8, sgemm_128x384x192x2_tn) {
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
//think about using run_gemm directly
run_batched_strided_gemm<SgemmTraits>(128/*m*/, 384/*n*/, 192/*k*/, 2 /*batch_size*/);
}
@@ -110,6 +118,7 @@ TEST(Sgemm_strided_batched_128x128x8, sgemm_256x384x64x3_tt) {
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
//think about using run_gemm directly
run_batched_strided_gemm<SgemmTraits>(256/*m*/, 384/*n*/, 64/*k*/, 3 /*batch_size*/);
}
@@ -119,8 +128,8 @@ TEST(Sgemm_strided_batched_128x128x8, sgemm_128x384x192x2_tt) {
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
//think about using run_gemm directly
run_batched_strided_gemm<SgemmTraits>(128/*m*/, 384/*n*/, 192/*k*/, 2 /*batch_size*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////

View File

@@ -0,0 +1,240 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#include "cutlass/wmma_matrix.h"
#if defined(CUTLASS_USE_WMMA_API)
#include "cutlass_unit_test.h"
#include "cutlass/gemm/gemm.h"
#include "cutlass/gemm/wmma_gemm_traits.h"
#include "tools/test/unit/gemm/gemm_testbed.h"
#include "tools/test/unit/gemm/run_gemm.h"
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_strided_batched_16x16x32_f16, wmma_gemm_32x32x16_nn) {
/*
this wmmaTraits requires leading dim to be divisible by 4
*/
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
half,
cutlass::gemm::LinearScaling<half>,
half
>
WmmaGemmTraits;
run_batched_strided_gemm<WmmaGemmTraits>(32, 32, 64, 3);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_strided_batched_16x16x32_f16, wmma_gemm_32x32x16_nt) {
/*
this wmmaTraits requires leading dim to be divisible by 4
*/
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
half,
cutlass::gemm::LinearScaling<half>,
half
>
WmmaGemmTraits;
run_batched_strided_gemm<WmmaGemmTraits>(32, 32, 64, 3);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_strided_batched_16x16x32_f16, wmma_gemm_32x32x16_tn) {
/*
this wmmaTraits requires leading dim to be divisible by 4
*/
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
half,
cutlass::gemm::LinearScaling<half>,
half
>
WmmaGemmTraits;
run_batched_strided_gemm<WmmaGemmTraits>(32, 32, 64, 3);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_strided_batched_16x16x32_f16, wmma_gemm_32x32x16_tt) {
/*
this wmmaTraits requires leading dim to be divisible by 4
*/
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
half,
cutlass::gemm::LinearScaling<half>,
half
>
WmmaGemmTraits;
run_batched_strided_gemm<WmmaGemmTraits>(32, 32, 64, 3);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
//non multiple of 16
TEST(WmmaGemm_strided_batched_16x16x32_f16, wmma_gemm_36x36x16_nn) {
/*
this wmmaTraits requires leading dim to be divisible by 4
*/
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
half,
cutlass::gemm::LinearScaling<half>,
half,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
4, /*kScalarsPerLdgA_*/
4, /*kScalarsPerLdgB_*/
4, /*KScalarsPerLdsA_*/
4, /*KScalarsPerLdsB_*/
4 / sizeof(half), /*kScalarsPerLdgCAndStgD_*/
4 / sizeof(half), /*kScalarsPerStsD_*/
4 / sizeof(half) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_batched_strided_gemm<WmmaGemmTraits>(36, 36, 64, 3);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_strided_batched_16x16x32_f16, wmma_gemm_36x36x16_nt) {
/*
this wmmaTraits requires leading dim to be divisible by 4
*/
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
half,
cutlass::gemm::LinearScaling<half>,
half,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
4, /*kScalarsPerLdgA_*/
4, /*kScalarsPerLdgB_*/
4, /*KScalarsPerLdsA_*/
4, /*KScalarsPerLdsB_*/
4 / sizeof(half), /*kScalarsPerLdgCAndStgD_*/
4 / sizeof(half), /*kScalarsPerStsD_*/
4 / sizeof(half) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_batched_strided_gemm<WmmaGemmTraits>(36, 36, 64, 3);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_strided_batched_16x16x32_f16, wmma_gemm_36x36x16_tn) {
/*
this wmmaTraits requires leading dim to be divisible by 4
*/
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
half,
cutlass::gemm::LinearScaling<half>,
half,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
4, /*kScalarsPerLdgA_*/
4, /*kScalarsPerLdgB_*/
4, /*KScalarsPerLdsA_*/
4, /*KScalarsPerLdsB_*/
4 / sizeof(half), /*kScalarsPerLdgCAndStgD_*/
4 / sizeof(half), /*kScalarsPerStsD_*/
4 / sizeof(half) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_batched_strided_gemm<WmmaGemmTraits>(36, 36, 64, 3);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_strided_batched_16x16x32_f16, wmma_gemm_36x36x16_tt) {
/*
this wmmaTraits requires leading dim to be divisible by 4
*/
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
half,
cutlass::gemm::LinearScaling<half>,
half,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
4, /*kScalarsPerLdgA_*/
4, /*kScalarsPerLdgB_*/
4, /*KScalarsPerLdsA_*/
4, /*KScalarsPerLdsB_*/
4 / sizeof(half), /*kScalarsPerLdgCAndStgD_*/
4 / sizeof(half), /*kScalarsPerStsD_*/
4 / sizeof(half) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_batched_strided_gemm<WmmaGemmTraits>(36, 36, 64, 3);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
#endif

View File

@@ -28,7 +28,7 @@
#include "tools/test/unit/gemm/gemm_testbed.h"
#include "tools/test/unit/gemm/run_gemm.h"
#if !defined(__CUDA_ARCH__) || __CUDA_ARCH__ >= 530
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Fp16_sgemm_alphaFp16_fp16_128x128x16, fp16_sgemm_fp16_128x128x16_nn) {
@@ -319,3 +319,5 @@ TEST(Fp16_sgemm_alphaFp32_fp16_128x128x16, fp16_sgemm_fp16_128x112x17_tt) {
run_gemm<SgemmTraits>(128, 112, 17);
}
#endif

View File

@@ -28,7 +28,7 @@
#include "tools/test/unit/gemm/gemm_testbed.h"
#include "tools/test/unit/gemm/run_gemm.h"
#if !defined(__CUDA_ARCH__) || __CUDA_ARCH__ >= 530
TEST(Fp16_sgemm_alphaFp32_fp32_128x128x16, fp16_sgemm_fp32_128x128x16_nn) {
@@ -172,3 +172,6 @@ TEST(Fp16_sgemm_alphaFp32_fp32_128x128x16, fp16_sgemm_fp32_128x112x17_tt) {
SgemmTraits;
run_gemm<SgemmTraits>(128, 112, 17);
}
#endif

View File

@@ -0,0 +1,381 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#include "cutlass/wmma_matrix.h"
#if defined(CUTLASS_USE_WMMA_API)
#include "cutlass_unit_test.h"
#include "cutlass/gemm/gemm.h"
#include "cutlass/gemm/wmma_gemm_traits.h"
#include "tools/test/unit/gemm/gemm_testbed.h"
#include "tools/test/unit/gemm/run_gemm.h"
TEST(WmmaGemm_16x16x32_fp32, fp16_wmma_gemm_fp16_16x16x16_nn) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
8, /*kScalarsPerLdgA_*/
8, /*kScalarsPerLdgB_*/
8, /*KScalarsPerLdsA_*/
8, /*KScalarsPerLdsB_*/
16 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
16 / sizeof(accumu_type), /*kScalarsPerStsD_*/
16 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_gemm<WmmaGemmTraits>(16, 16, 16);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_16x16x32_fp32, fp16_wmma_gemm_fp16_8x8x16_nn) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
8, /*kScalarsPerLdgA_*/
8, /*kScalarsPerLdgB_*/
8, /*KScalarsPerLdsA_*/
8, /*KScalarsPerLdsB_*/
16 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
16 / sizeof(accumu_type), /*kScalarsPerStsD_*/
16 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_gemm<WmmaGemmTraits>(8, 8, 16);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_16x16x32_fp32, fp16_wmma_gemm_fp16_256x256x64_nn) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
8, /*kScalarsPerLdgA_*/
8, /*kScalarsPerLdgB_*/
8, /*KScalarsPerLdsA_*/
8, /*KScalarsPerLdsB_*/
16 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
16 / sizeof(accumu_type), /*kScalarsPerStsD_*/
16 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_gemm<WmmaGemmTraits>(256, 256, 64);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_16x16x32_fp32, fp16_wmma_gemm_fp16_16x16x16_nt) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
8, /*kScalarsPerLdgA_*/
8, /*kScalarsPerLdgB_*/
8, /*KScalarsPerLdsA_*/
8, /*KScalarsPerLdsB_*/
16 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
16 / sizeof(accumu_type), /*kScalarsPerStsD_*/
16 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_gemm<WmmaGemmTraits>(16, 16, 16);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_16x16x32_fp32, fp16_wmma_gemm_fp16_8x8x16_nt) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
8, /*kScalarsPerLdgA_*/
8, /*kScalarsPerLdgB_*/
8, /*KScalarsPerLdsA_*/
8, /*KScalarsPerLdsB_*/
16 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
16 / sizeof(accumu_type), /*kScalarsPerStsD_*/
16 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_gemm<WmmaGemmTraits>(8, 8, 16);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_16x16x32_fp32, fp16_wmma_gemm_fp16_256x256x64_nt) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
8, /*kScalarsPerLdgA_*/
8, /*kScalarsPerLdgB_*/
8, /*KScalarsPerLdsA_*/
8, /*KScalarsPerLdsB_*/
16 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
16 / sizeof(accumu_type), /*kScalarsPerStsD_*/
16 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_gemm<WmmaGemmTraits>(256, 256, 64);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_16x16x32_fp32, fp16_wmma_gemm_fp16_16x16x16_tn) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
8, /*kScalarsPerLdgA_*/
8, /*kScalarsPerLdgB_*/
8, /*KScalarsPerLdsA_*/
8, /*KScalarsPerLdsB_*/
16 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
16 / sizeof(accumu_type), /*kScalarsPerStsD_*/
16 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_gemm<WmmaGemmTraits>(16, 16, 16);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_16x16x32_fp32, fp16_wmma_gemm_fp16_8x8x16_tn) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
8, /*kScalarsPerLdgA_*/
8, /*kScalarsPerLdgB_*/
8, /*KScalarsPerLdsA_*/
8, /*KScalarsPerLdsB_*/
16 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
16 / sizeof(accumu_type), /*kScalarsPerStsD_*/
16 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_gemm<WmmaGemmTraits>(8, 8, 16);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_16x16x32_fp32, fp16_wmma_gemm_fp16_256x256x64_tn) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
8, /*kScalarsPerLdgA_*/
8, /*kScalarsPerLdgB_*/
8, /*KScalarsPerLdsA_*/
8, /*KScalarsPerLdsB_*/
16 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
16 / sizeof(accumu_type), /*kScalarsPerStsD_*/
16 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_gemm<WmmaGemmTraits>(256, 256, 64);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_16x16x32_fp32, fp16_wmma_gemm_fp16_16x16x16_tt) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
8, /*kScalarsPerLdgA_*/
8, /*kScalarsPerLdgB_*/
8, /*KScalarsPerLdsA_*/
8, /*KScalarsPerLdsB_*/
16 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
16 / sizeof(accumu_type), /*kScalarsPerStsD_*/
16 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_gemm<WmmaGemmTraits>(16, 16, 16);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_16x16x32_fp32, fp16_wmma_gemm_fp16_8x8x16_tt) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
8, /*kScalarsPerLdgA_*/
8, /*kScalarsPerLdgB_*/
8, /*KScalarsPerLdsA_*/
8, /*KScalarsPerLdsB_*/
16 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
16 / sizeof(accumu_type), /*kScalarsPerStsD_*/
16 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_gemm<WmmaGemmTraits>(8, 8, 16);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_16x16x32_fp32, fp16_wmma_gemm_fp16_256x256x64_tt) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
8, /*kScalarsPerLdgA_*/
8, /*kScalarsPerLdgB_*/
8, /*KScalarsPerLdsA_*/
8, /*KScalarsPerLdsB_*/
16 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
16 / sizeof(accumu_type), /*kScalarsPerStsD_*/
16 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_gemm<WmmaGemmTraits>(256, 256, 64);
}
#endif //#if defined(CUTLASS_USE_WMMA_API)

View File

@@ -0,0 +1,273 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#include "cutlass/wmma_matrix.h"
#if defined(CUTLASS_USE_WMMA_API)
#include "cutlass_unit_test.h"
#include "cutlass/gemm/gemm.h"
#include "cutlass/gemm/wmma_gemm_traits.h"
#include "tools/test/unit/gemm/gemm_testbed.h"
#include "tools/test/unit/gemm/run_gemm.h"
////////////////////////////////////////////////////////////////////////////////////////////////////
/* mulitple of 4*/
TEST(WmmaGemm_16x16x32_fp32, fp16_wmma_gemm_fp16_36x36x64_nn) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
4, /*kScalarsPerLdgA_*/
4, /*kScalarsPerLdgB_*/
4, /*KScalarsPerLdsA_*/
4, /*KScalarsPerLdsB_*/
8 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
8 / sizeof(accumu_type), /*kScalarsPerStsD_*/
8 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_gemm<WmmaGemmTraits>(36, 36, 64);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_16x16x32_fp32, fp16_wmma_gemm_fp16_36x36x64_nt) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
4, /*kScalarsPerLdgA_*/
4, /*kScalarsPerLdgB_*/
4, /*KScalarsPerLdsA_*/
4, /*KScalarsPerLdsB_*/
8 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
8 / sizeof(accumu_type), /*kScalarsPerStsD_*/
8 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_gemm<WmmaGemmTraits>(36, 36, 64);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_16x16x32_fp32, fp16_wmma_gemm_fp16_36x36x64_tn) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
4, /*kScalarsPerLdgA_*/
4, /*kScalarsPerLdgB_*/
4, /*KScalarsPerLdsA_*/
4, /*KScalarsPerLdsB_*/
8 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
8 / sizeof(accumu_type), /*kScalarsPerStsD_*/
8 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_gemm<WmmaGemmTraits>(36, 36, 64);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_16x16x32_fp32, fp16_wmma_gemm_fp16_36x36x64_tt) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
4, /*kScalarsPerLdgA_*/
4, /*kScalarsPerLdgB_*/
4, /*KScalarsPerLdsA_*/
4, /*KScalarsPerLdsB_*/
8 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
8 / sizeof(accumu_type), /*kScalarsPerStsD_*/
8 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_gemm<WmmaGemmTraits>(36, 36, 64);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
/* mulitple of 2*/
TEST(WmmaGemm_16x16x32_fp32, fp16_wmma_gemm_fp16_34x34x64_nn) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
2, /*kScalarsPerLdgA_*/
2, /*kScalarsPerLdgB_*/
2, /*KScalarsPerLdsA_*/
2, /*KScalarsPerLdsB_*/
4 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
4 / sizeof(accumu_type), /*kScalarsPerStsD_*/
4 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_gemm<WmmaGemmTraits>(34, 34, 64);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
/* mulitple of 2*/
TEST(WmmaGemm_16x16x32_fp32, fp16_wmma_gemm_fp16_34x34x64_nt) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
2, /*kScalarsPerLdgA_*/
2, /*kScalarsPerLdgB_*/
2, /*KScalarsPerLdsA_*/
2, /*KScalarsPerLdsB_*/
4 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
4 / sizeof(accumu_type), /*kScalarsPerStsD_*/
4 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_gemm<WmmaGemmTraits>(34, 34, 64);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
/* mulitple of 2*/
TEST(WmmaGemm_16x16x32_fp32, fp16_wmma_gemm_fp16_34x34x64_tn) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
2, /*kScalarsPerLdgA_*/
2, /*kScalarsPerLdgB_*/
2, /*KScalarsPerLdsA_*/
2, /*KScalarsPerLdsB_*/
4 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
4 / sizeof(accumu_type), /*kScalarsPerStsD_*/
4 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_gemm<WmmaGemmTraits>(34, 34, 64);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
/* mulitple of 2*/
TEST(WmmaGemm_16x16x32_fp32, fp16_wmma_gemm_fp16_34x34x64_tt) {
typedef float accumu_type;
typedef half c_type;
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
c_type,
cutlass::gemm::LinearScaling<accumu_type>,
accumu_type,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
2, /*kScalarsPerLdgA_*/
2, /*kScalarsPerLdgB_*/
2, /*KScalarsPerLdsA_*/
2, /*KScalarsPerLdsB_*/
4 / sizeof(c_type), /*kScalarsPerLdgCAndStgD_*/
4 / sizeof(accumu_type), /*kScalarsPerStsD_*/
4 / sizeof(accumu_type) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_gemm<WmmaGemmTraits>(34, 34, 64);
}
#endif

View File

@@ -0,0 +1,382 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#include "cutlass_unit_tests.h"
#include "tools/util/host_tensor.h"
#include "tools/test/unit/core/layout_verification.h"
#include "tools/util/tensor_view_io.h"
#include "cutlass/gemm/gemm.h"
#include "cutlass/shape.h"
#include "cutlass/gemm/sgemm_traits.h"
#include "cutlass/gemm/dgemm_traits.h"
#include "cutlass/gemm/hgemm_traits.h"
////////////////////////////////////////////////////////////////////////////////////////////////////
namespace test {
// M/N/K struct.
struct GemmDesc {
int m, n, k;
CUTLASS_HOST_DEVICE GemmDesc(int m_, int n_, int k_) : m(m_), n(n_), k(k_) {}
};
/// Simple test to load from global memory and store to shared memory
// Loading from global memory and storing to shared memory for A
template <typename Traits>
__global__ void Gemm_load_global_store_shared_a(
typename Traits::GlobalLoadStreamA::Scalar *output,
typename Traits::GlobalLoadStreamA::Scalar const *input,
int M,
int N,
int K,
int ldm) {
//Create shared memory.
__shared__ typename Traits::SharedStorage shared_storage;
// Create those iterators.
typedef typename Traits::GlobalLoadStreamA GlobalLoadStreamA;
typename GlobalLoadStreamA::Params global_load_params;
GemmDesc desc(M, N, K);
global_load_params.initialize(desc, input, ldm);
GlobalLoadStreamA stream_a(global_load_params, shared_storage.main_loop.stream_a.global, M, N, K, cutlass::make_Coord(0, 0, 0));
stream_a.copy();
stream_a.commit();
// store barrier
__syncthreads();
// one thread writes everything out
if (threadIdx.x == 0) {
for (int i = 0; i < M*K; ++i) {
output[i] = shared_storage.main_loop.stream_a.shared[i];
}
}
}
// Loading from global memory and storing to shared memory for B
template <typename Traits>
__global__ void Gemm_load_global_store_shared_b(
typename Traits::GlobalLoadStreamB::Scalar *output,
typename Traits::GlobalLoadStreamB::Scalar const *input,
int M,
int N,
int K,
int ldm) {
//Create shared memory.
__shared__ typename Traits::SharedStorage shared_storage;
// Create those iterators.
typedef typename Traits::GlobalLoadStreamB GlobalLoadStreamB;
typename GlobalLoadStreamB::Params global_load_params;
GemmDesc desc(M, N, K);
global_load_params.initialize(desc, input, ldm);
GlobalLoadStreamB stream_b(global_load_params, shared_storage.main_loop.stream_b.global, M, N, K, cutlass::make_Coord(0, 0, 0));
stream_b.copy();
stream_b.commit();
// store barrier
__syncthreads();
// one thread writes everything out
if (threadIdx.x == 0) {
for (int i = 0; i < M*K; ++i) {
output[i] = shared_storage.main_loop.stream_b.shared[i];
}
}
}
////////////////////////////////////////////////////////////////////////////////////////////////////
template <
typename CtaTile, // concept: Shape
typename DestType, // raw data type
typename SourceType // raw data type
>
class VerifyDataMovement {
public:
/// Tensor to store the destination data
cutlass::HostTensor<DestType> destination;
/// Tensor to store the source data
cutlass::HostTensor<SourceType> source;
/// Verification utility
typedef test::VerifyLayout<
DestType,
test::CoordinatePack<DestType>,
SourceType,
test::CoordinatePack<SourceType> > VerifyLayout;
/// Verification object
VerifyLayout verify_layout;
public:
VerifyDataMovement() { }
VerifyDataMovement(test::Layout const &source_layout) {
// Actual layout here doesn't matter here, just the number of elements
destination.resize_matrix(CtaTile::kH, CtaTile::kW, cutlass::MatrixLayout::kRowMajor);
source.resize_matrix(CtaTile::kH, CtaTile::kW, cutlass::MatrixLayout::kRowMajor);
verify_layout.initialize(source, source_layout);
destination.fill(0);
destination.sync_device();
source.sync_device();
}
/// Verifies resulting layout
bool verify(test::Layout const & destination_layout) {
destination.sync_host();
typename VerifyLayout::VisitorVerbose visitor(std::cout);
bool passed = verify_layout.verify(
destination,
destination_layout,
visitor);
return passed;
}
};
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Gemm_shared_tile, A_float_contiguous) {
static int const M = 64;
static int const N = 64;
static int const K = 8;
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<K, N, M> >
SgemmTraits;
typedef test::Layout::Span Span;
test::Layout::SpanVector dst_layout;
test::Layout::SpanVector src_layout;
// define the source layout
src_layout.push_back(Span(0, K));
src_layout.push_back(Span(1, M));
typedef VerifyDataMovement<
cutlass::Shape<1, M, K, 1>,
float,
float
> VerifyDataMovement_t;
VerifyDataMovement_t testbed(src_layout);
test::Gemm_load_global_store_shared_a< SgemmTraits ><<<
dim3(1,1,1),
dim3(SgemmTraits::kThreads, 1)
>>>(
testbed.destination.device_data(),
testbed.source.device_data(),
M,
N,
K,
M
);
cudaError_t result = cudaDeviceSynchronize();
ASSERT_EQ(result, cudaSuccess) << "\nCUDA kernel launch error: " << cudaGetErrorString(result)
<< "\n";
// define the destination layout
dst_layout.push_back(Span(0, K));
dst_layout.push_back(Span(1, M));
EXPECT_TRUE(testbed.verify(dst_layout));
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Gemm_shared_tile, A_double_contiguous) {
static int const M = 64;
static int const N = 64;
static int const K = 8;
typedef cutlass::gemm::DgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<K, N, M> >
DgemmTraits;
typedef test::Layout::Span Span;
test::Layout::SpanVector dst_layout;
test::Layout::SpanVector src_layout;
// define the source layout
src_layout.push_back(Span(0, K));
src_layout.push_back(Span(1, M));
typedef VerifyDataMovement<
cutlass::Shape<1, M, K, 1>,
double,
double
> VerifyDataMovement_t;
VerifyDataMovement_t testbed(src_layout);
test::Gemm_load_global_store_shared_a< DgemmTraits ><<<
dim3(1,1,1),
dim3(DgemmTraits::kThreads, 1)
>>>(
testbed.destination.device_data(),
testbed.source.device_data(),
M,
N,
K,
M
);
cudaError_t result = cudaDeviceSynchronize();
ASSERT_EQ(result, cudaSuccess) << "\nCUDA kernel launch error: " << cudaGetErrorString(result)
<< "\n";
// define the destination layout
dst_layout.push_back(Span(0, K));
dst_layout.push_back(Span(1, M));
EXPECT_TRUE(testbed.verify(dst_layout));
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Gemm_shared_tile, B_float_contiguous) {
static int const M = 64;
static int const N = 64;
static int const K = 8;
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<K, N, M> >
SgemmTraits;
typedef test::Layout::Span Span;
test::Layout::SpanVector dst_layout;
test::Layout::SpanVector src_layout;
// define the source layout
src_layout.push_back(Span(0, K));
src_layout.push_back(Span(1, M));
typedef VerifyDataMovement<
cutlass::Shape<1, M, K, 1>,
float,
float
> VerifyDataMovement_t;
VerifyDataMovement_t testbed(src_layout);
test::Gemm_load_global_store_shared_b< SgemmTraits ><<<
dim3(1,1,1),
dim3(SgemmTraits::kThreads, 1)
>>>(
testbed.destination.device_data(),
testbed.source.device_data(),
M,
N,
K,
M
);
cudaError_t result = cudaDeviceSynchronize();
ASSERT_EQ(result, cudaSuccess) << "\nCUDA kernel launch error: " << cudaGetErrorString(result)
<< "\n";
// define the destination layout
dst_layout.push_back(Span(0, K));
dst_layout.push_back(Span(1, M));
EXPECT_TRUE(testbed.verify(dst_layout));
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Gemm_shared_tile, B_double_contiguous) {
static int const M = 64;
static int const N = 64;
static int const K = 8;
typedef cutlass::gemm::DgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<K, N, M> >
DgemmTraits;
typedef test::Layout::Span Span;
test::Layout::SpanVector dst_layout;
test::Layout::SpanVector src_layout;
// define the source layout
src_layout.push_back(Span(0, K));
src_layout.push_back(Span(1, M));
typedef VerifyDataMovement<
cutlass::Shape<1, M, K, 1>,
double,
double
> VerifyDataMovement_t;
VerifyDataMovement_t testbed(src_layout);
test::Gemm_load_global_store_shared_b< DgemmTraits ><<<
dim3(1,1,1),
dim3(DgemmTraits::kThreads, 1)
>>>(
testbed.destination.device_data(),
testbed.source.device_data(),
M,
N,
K,
M
);
cudaError_t result = cudaDeviceSynchronize();
ASSERT_EQ(result, cudaSuccess) << "\nCUDA kernel launch error: " << cudaGetErrorString(result)
<< "\n";
// define the destination layout
dst_layout.push_back(Span(0, K));
dst_layout.push_back(Span(1, M));
EXPECT_TRUE(testbed.verify(dst_layout));
}
////////////////////////////////////////////////////////////////////////////////////////////////////
}

View File

@@ -46,6 +46,7 @@
#include "tools/util/type_traits.h"
#include "tools/util/reference/host/gemm.h"
#include "tools/util/reference/device/gemm.h"
#include "tools/util/reference/host/tensor_elementwise.h"
//////////////////////////////////////////////////////////////////////////////////////////
@@ -247,6 +248,9 @@ struct GemmTestbed {
/// Reference result computed on the host
HostMatrixC ref_host;
/// Reference result computed on the device
HostMatrixC ref_device;
/// Reference result computed with cublas
HostMatrixC ref_cublas;
@@ -262,6 +266,9 @@ struct GemmTestbed {
/// batch count
int batch_count;
/// partitionK count
int partitionK_count;
/// distance between A[i] and A[i+1] for strided batched gemm
long long int batch_stride_A;
@@ -308,6 +315,7 @@ struct GemmTestbed {
beta(beta_),
algorithm(algorithm_),
batch_count(1),
partitionK_count(1),
batch_stride_A(static_cast<long long int>(0)),
batch_stride_B(static_cast<long long int>(0)),
batch_stride_C(static_cast<long long int>(0)) {
@@ -320,6 +328,7 @@ struct GemmTestbed {
resize(B, K_, N_, layout_b);
resize(C_initial, M_, N_, layout_c);
resize(ref_host, M_, N_, layout_c);
resize(ref_device, M_, N_, layout_c);
resize(ref_cublas, M_, N_, layout_c);
resize(computed, M_, N_, layout_c);
}
@@ -345,6 +354,7 @@ struct GemmTestbed {
beta(beta_),
algorithm(algorithm_),
batch_count(1),
partitionK_count(1),
batch_stride_A(static_cast<long long int>(0)),
batch_stride_B(static_cast<long long int>(0)),
batch_stride_C(static_cast<long long int>(0)) {
@@ -353,6 +363,7 @@ struct GemmTestbed {
resize(B, K_ * batch_count, N_, layout_b);
resize(C_initial, M_, N_ * batch_count, layout_c);
resize(ref_host, M_, N_ * batch_count, layout_c);
resize(ref_device, M_, N_ * batch_count, layout_c);
resize(ref_cublas, M_, N_ * batch_count, layout_c);
resize(computed, M_, N_ * batch_count, layout_c);
}
@@ -377,6 +388,7 @@ struct GemmTestbed {
beta(beta_),
algorithm(algorithm_),
batch_count(1),
partitionK_count(1),
batch_stride_A(static_cast<long long int>(0)),
batch_stride_B(static_cast<long long int>(0)),
batch_stride_C(static_cast<long long int>(0)) {
@@ -389,6 +401,7 @@ struct GemmTestbed {
resize(B, K_, N_, layout_b, ldb);
resize(C_initial, M_, N_, layout_c, ldc);
resize(ref_host, M_, N_, layout_c, ldc);
resize(ref_device, M_, N_, layout_c, ldc);
resize(ref_cublas, M_, N_, layout_c, ldc);
resize(computed, M_, N_, layout_c, ldc);
}
@@ -414,6 +427,7 @@ struct GemmTestbed {
beta(beta_),
algorithm(algorithm_),
batch_count(1),
partitionK_count(1),
batch_stride_A(static_cast<long long int>(0)),
batch_stride_B(static_cast<long long int>(0)),
batch_stride_C(static_cast<long long int>(0)) {
@@ -422,6 +436,7 @@ struct GemmTestbed {
resize(B, K_ * batch_count, N_, layout_b);
resize(C_initial, M_, N_ * batch_count, layout_c);
resize(ref_host, M_, N_ * batch_count, layout_c);
resize(ref_device, M_, N_ * batch_count, layout_c);
resize(ref_cublas, M_, N_ * batch_count, layout_c);
resize(computed, M_, N_ * batch_count, layout_c);
}
@@ -446,7 +461,8 @@ struct GemmTestbed {
alpha(alpha_),
beta(beta_),
algorithm(algorithm_),
batch_count(batch_count_) {
batch_count(batch_count_),
partitionK_count(1) {
status = cublasCreate(&handle);
if (status != CUBLAS_STATUS_SUCCESS) {
@@ -457,6 +473,7 @@ struct GemmTestbed {
resize(B, K_ * batch_count, N_, layout_b);
resize(C_initial, M_, N_ * batch_count, layout_c);
resize(ref_host, M_, N_ * batch_count, layout_c);
resize(ref_device, M_, N_ * batch_count, layout_c);
resize(ref_cublas, M_, N_ * batch_count, layout_c);
resize(computed, M_, N_ * batch_count, layout_c);
@@ -465,6 +482,50 @@ struct GemmTestbed {
batch_stride_C = M_ * N_;
}
/// Constructs a workspace for verifying partitionedK GEMM, assumes
/// dense packing.
/// in partitionedK GEMM, the K is partitioned by partitionK_size
/// each partition is of the same size, except for the last partition
/// each partition, except for the last one, is of size K / partitionK_count
/// if K is not divisible by partitionK_size, the last partitionK = K % partitionK_count + K / partitionK_count
GemmTestbed(int M_,
int N_,
std::pair<int, int> K_pair_, /*(k, partitionK_count)*/
cublasOperation_t layout_a,
cublasOperation_t layout_b,
Scalar alpha_ = Scalar(1),
Scalar beta_ = Scalar(0),
cublasGemmAlgo_t algorithm_ = CUBLAS_GEMM_DEFAULT,
cublasOperation_t layout_c = CUBLAS_OP_N)
: problem_size(K_pair_.first, N_, M_, 1),
layout_A(layout_a),
layout_B(layout_b),
alpha(alpha_),
beta(beta_),
algorithm(algorithm_),
batch_count(1),
partitionK_count(K_pair_.second) {
status = cublasCreate(&handle);
if (status != CUBLAS_STATUS_SUCCESS) {
throw cutlass::cuda_exception("Failed to create CUBLAS handle");
}
resize(A, M_, K_pair_.first, layout_a);
resize(B, K_pair_.first, N_, layout_b);
resize(C_initial, M_, N_ * partitionK_count, layout_c);
resize(ref_host, M_, N_ * partitionK_count, layout_c);
resize(ref_device, M_, N_ * partitionK_count, layout_c);
resize(ref_cublas, M_, N_ * partitionK_count, layout_c);
resize(computed, M_, N_ * partitionK_count, layout_c);
// we can use a combination of batched stried gemm and regular gemm
// to simulation partitionedK, which is what we will do for reference code
int partitionK_size = K() / partitionK_count;
batch_stride_A = (layout_a == CUBLAS_OP_N) ? M_ * partitionK_size : partitionK_size;
batch_stride_B = (layout_b == CUBLAS_OP_N) ? partitionK_size : partitionK_size * N_;
batch_stride_C = M_ * N_;
}
/// Destructs the GEMM testbed
~GemmTestbed() {
if (status != CUBLAS_STATUS_NOT_INITIALIZED) {
@@ -504,7 +565,14 @@ struct GemmTestbed {
/// Returns the number of flops implied by the computation (1 multiply-accumulate = 2 flops)
uint64_t flops() const {
return uint64_t(batch_count) * uint64_t(M()) * uint64_t(N()) * uint64_t(K()) * 2ULL;
if (partitionK_count == 1) {
return uint64_t(batch_count) * uint64_t(M()) * uint64_t(N()) * uint64_t(K()) * 2ULL;
}
else {
int partitionK_size = K() / partitionK_count;
return (uint64_t(partitionK_count - 1) * uint64_t(batch_count) * uint64_t(M()) * uint64_t(N()) * uint64_t(partitionK_size) * 2ULL)
+ (uint64_t(batch_count) * uint64_t(M()) * uint64_t(N()) * uint64_t(K() - partitionK_size * (partitionK_count - 1)) * 2ULL);
}
}
/// Computes the speed of the computation in GFLOPs/s
@@ -555,14 +623,15 @@ struct GemmTestbed {
// Initialize the source matrix with a uniform distribution
cutlass::Distribution dist;
dist.set_uniform(-8, 8);
cutlass::reference::host::TensorInitialize(A.host_view(), seed, dist);
cutlass::reference::host::TensorInitialize(B.host_view(), seed + 11, dist);
cutlass::reference::host::TensorInitialize(C_initial.host_view(), seed + 13, dist);
A.sync_device();
B.sync_device();
C_initial.sync_device();
}
/// Initializes binary data
@@ -585,56 +654,121 @@ struct GemmTestbed {
/// Computes the matrix product on the host
void compute_host() {
ref_host.fill(C_initial);
cutlass::reference::host::Gemm(problem_size, alpha, A.host_ref(), B.host_ref(), beta, ref_host.host_ref(), Accumulator(0));
}
/// Compute the matrix product using the device-side reference
void compute_device_reference() {
ref_device.fill(C_initial);
cutlass::reference::device::Gemm(
problem_size,
cutlass::TypeTraits<Scalar>::to_device(alpha),
A.device_ref(),
B.device_ref(),
cutlass::TypeTraits<Scalar>::to_device(beta),
ref_device.device_ref(),
cutlass::TypeTraits<Accumulator>::to_device(0)
);
}
/// Excutes an equivalent GEMM using cuBLAS
bool execute_cublas() {
if (batch_count == 1) {
status = cublasGemmEx(handle,
layout_a(),
layout_b(),
M(),
N(),
K(),
&alpha,
ptr_A(),
cutlass::TypeTraits<AType>::cublas_type,
lda(),
ptr_B(),
cutlass::TypeTraits<BType>::cublas_type,
ldb(),
&beta,
ref_cublas.device_data(),
cutlass::TypeTraits<CType>::cublas_type,
ldc(),
cutlass::TypeTraits<Accumulator>::cublas_type,
algorithm);
if (partitionK_count == 1) {
if (batch_count == 1) {
status = cublasGemmEx(handle,
layout_a(),
layout_b(),
M(),
N(),
K(),
&alpha,
ptr_A(),
cutlass::TypeTraits<AType>::cublas_type,
lda(),
ptr_B(),
cutlass::TypeTraits<BType>::cublas_type,
ldb(),
&beta,
ref_cublas.device_data(),
cutlass::TypeTraits<CType>::cublas_type,
ldc(),
cutlass::TypeTraits<Accumulator>::cublas_type,
algorithm);
return status == CUBLAS_STATUS_SUCCESS;
} else {
// call strided batched gemm
return status == CUBLAS_STATUS_SUCCESS;
}
else {
// call strided batched gemm
status = cublasGemmStridedBatchedTemplate(handle,
layout_a(),
layout_b(),
M(),
N(),
K(),
&alpha,
ptr_A(),
lda(),
batch_stride_A,
ptr_B(),
ldb(),
batch_stride_B,
&beta,
ref_cublas.device_data(),
ldc(),
batch_stride_C,
batch_count);
return status == CUBLAS_STATUS_SUCCESS;
}
}
else {
assert(batch_count == 1);
//the last batch is of a different K
//first call strided batched gemm
int partitionK_size = K() / partitionK_count;
//int lastK_size = (K() % partitionK_size) + partitionK_size;
int lastK_size = K() - partitionK_size * (partitionK_count - 1);
status = cublasGemmStridedBatchedTemplate(handle,
layout_a(),
layout_b(),
M(),
N(),
K(),
&alpha,
ptr_A(),
lda(),
batch_stride_A,
ptr_B(),
ldb(),
batch_stride_B,
&beta,
ref_cublas.device_data(),
ldc(),
batch_stride_C,
batch_count);
layout_a(),
layout_b(),
M(),
N(),
partitionK_size,
&alpha,
ptr_A(),
lda(),
batch_stride_A,
ptr_B(),
ldb(),
batch_stride_B,
&beta,
ref_cublas.device_data(),
ldc(),
batch_stride_C,
partitionK_count - 1);
//then call gemm for the last batch
status = cublasGemmEx(handle,
layout_a(),
layout_b(),
M(),
N(),
lastK_size,
&alpha,
ptr_A() + (partitionK_count - 1) * batch_stride_A,
cutlass::TypeTraits<AType>::cublas_type,
lda(),
ptr_B() + (partitionK_count - 1) * batch_stride_B,
cutlass::TypeTraits<BType>::cublas_type,
ldb(),
&beta,
ref_cublas.device_data() + (partitionK_count - 1) * batch_stride_C,
cutlass::TypeTraits<CType>::cublas_type,
ldc(),
cutlass::TypeTraits<Accumulator>::cublas_type,
algorithm);
return status == CUBLAS_STATUS_SUCCESS;
}
}
@@ -787,6 +921,24 @@ struct GemmTestbed {
return passed;
}
/// Verifies the reference implementation with cuBLAS
bool verify_reference_with_cublas(bool save_on_error = true, bool always_print = false) {
compute_device_reference();
ref_device.sync_host();
compute_cublas();
ref_cublas.sync_host();
bool passed = ref_device.bit_equals(ref_cublas);
if ((!passed && save_on_error) || always_print) {
save_workspace(ref_device, ref_cublas);
}
return passed;
}
/// Verifies with host-side and device-side computations
bool verify_with_all() {
bool passed = true;
@@ -917,4 +1069,44 @@ template<> inline cublasStatus_t GemmTestbed<cutlass::half_t, cutlass::half_t, c
batchCount);
}
template<> inline cublasStatus_t GemmTestbed<cutlass::half_t, cutlass::half_t, cutlass::half_t, float, float>::cublasGemmStridedBatchedTemplate(cublasHandle_t handle,
cublasOperation_t transa,
cublasOperation_t transb,
int M,
int N,
int K,
const float *alpha,
const half *ptr_A,
int lda,
long long int stride_A,
const half *ptr_B,
int ldb,
long long int stride_B,
const float *beta,
half *ptr_C,
int ldc,
long long int stride_C,
int batchCount) {
return cublasGemmStridedBatchedEx(handle,
transa,
transb,
M, N, K,
alpha,
ptr_A,
cutlass::TypeTraits<cutlass::half_t>::cublas_type,
lda,
stride_A,
ptr_B,
cutlass::TypeTraits<cutlass::half_t>::cublas_type,
ldb,
stride_B,
beta,
ptr_C,
cutlass::TypeTraits<cutlass::half_t>::cublas_type,
ldc,
stride_C,
batchCount,
cutlass::TypeTraits<float>::cublas_type,
CUBLAS_GEMM_DEFAULT);
}
} // namespace test

View File

@@ -29,6 +29,8 @@
#include "tools/test/unit/gemm/gemm_testbed.h"
#include "tools/test/unit/gemm/run_gemm.h"
#if !defined(__CUDA_ARCH__) || __CUDA_ARCH__ >= 530
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Hgemm_128x128x16, hgemm_128x128x16_nt) {
@@ -326,4 +328,5 @@ TEST(Hgemm_128x128x16, hgemm_124x126x32_ragged_alpha2_beta1_nt) {
run_gemm<HgemmTraits>(124, 126, 32, cutlass::half_t(2), cutlass::half_t(1));
}
////////////////////////////////////////////////////////////////////////////////////////////////////
#endif

View File

@@ -29,6 +29,8 @@
#include "tools/test/unit/gemm/gemm_testbed.h"
#include "tools/test/unit/gemm/run_gemm.h"
#if !defined(__CUDA_ARCH__) || __CUDA_ARCH__ >= 530
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Hgemm_128x128x8, hgemm_128x128x1_nt) {
@@ -384,5 +386,5 @@ TEST(Hgemm_128x128x8, hgemm_124x126x32_ragged_alpha2_beta1_nt) {
}
////////////////////////////////////////////////////////////////////////////////////////////////////
#endif

View File

@@ -28,6 +28,7 @@
#include "tools/test/unit/gemm/gemm_testbed.h"
#include "tools/test/unit/gemm/run_gemm.h"
#if !defined(__CUDA_ARCH__) || __CUDA_ARCH__ >= 530
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Hgemm_128x32x8, hgemm_128x32x1_nt) {
@@ -312,3 +313,5 @@ TEST(Hgemm_128x32x8, hgemm_256x64x16_tt) {
run_gemm<HgemmTraits>(256, 64, 16);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
#endif

View File

@@ -28,6 +28,7 @@
#include "tools/test/unit/gemm/gemm_testbed.h"
#include "tools/test/unit/gemm/run_gemm.h"
#if !defined(__CUDA_ARCH__) || __CUDA_ARCH__ >= 530
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Hgemm_128x64x8, hgemm_128x64x1_nt) {
@@ -312,3 +313,5 @@ TEST(Hgemm_128x64x8, hgemm_256x128x16_tt) {
run_gemm<HgemmTraits>(256, 128, 16);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
#endif

View File

@@ -0,0 +1,378 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#include "cutlass_unit_test.h"
#include "cutlass/gemm/gemm.h"
#include "cutlass/gemm/sgemm_traits.h"
#include "tools/test/unit/gemm/gemm_testbed.h"
#include "tools/test/unit/gemm/run_gemm.h"
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Sgemm_partitionedK_128x128x8, sgemm_128x256x100x8_nn) {
/*
for example
partitionedK sgemm, m = 128, n = 256, overall_K = 100, partitionK_count = 8
for the first 7 partition k = overall_k / partitionK_count = 12
for the last partition last_k = overall_k - (partitionK_count - 1) * k = 16
*/
int m = 128;
int n = 256;
int overall_k = 100;
int partitionK_count = 8;
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
run_partitioned_k_gemm<SgemmTraits>(m, n, overall_k, partitionK_count);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Sgemm_partitionedK_128x128x8, sgemm_128x256x175x8_nn) {
int m = 128;
int n = 256;
int overall_k = 175;
int partitionK_count = 8;
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
run_partitioned_k_gemm<SgemmTraits>(m, n, overall_k, partitionK_count);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Sgemm_partitionedK_128x128x8, sgemm_10x12x20x3_nn) {
int m = 10;
int n = 12;
int overall_k = 20;
int partitionK_count = 3;
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
run_partitioned_k_gemm<SgemmTraits>(m, n, overall_k, partitionK_count);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Sgemm_partitionedK_128x128x8, sgemm_10x12x60x8_nn) {
int m = 10;
int n = 12;
int overall_k = 60;
int partitionK_count = 8;
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
run_partitioned_k_gemm<SgemmTraits>(m, n, overall_k, partitionK_count);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Sgemm_partitionedK_128x128x8, sgemm_128x256x100x4_nn) {
int m = 128;
int n = 256;
int overall_k = 100;
int partitionK_count = 4;
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
run_partitioned_k_gemm<SgemmTraits>(m, n, overall_k, partitionK_count);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Sgemm_partitionedK_128x128x8, sgemm_128x256x100x8_nt) {
/*
for example
partitionedK sgemm, m = 128, n = 256, overall_K = 100, partitionK_count = 8
for the first 7 partition k = overall_k / partitionK_count = 12
for the last partition last_k = overall_k - (partitionK_count - 1) * k = 16
*/
int m = 128;
int n = 256;
int overall_k = 100;
int partitionK_count = 8;
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
run_partitioned_k_gemm<SgemmTraits>(m, n, overall_k, partitionK_count);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Sgemm_partitionedK_128x128x8, sgemm_128x256x175x8_nt) {
int m = 128;
int n = 256;
int overall_k = 175;
int partitionK_count = 8;
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
run_partitioned_k_gemm<SgemmTraits>(m, n, overall_k, partitionK_count);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Sgemm_partitionedK_128x128x8, sgemm_10x12x20x3_nt) {
int m = 10;
int n = 12;
int overall_k = 20;
int partitionK_count = 3;
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
run_partitioned_k_gemm<SgemmTraits>(m, n, overall_k, partitionK_count);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Sgemm_partitionedK_128x128x8, sgemm_10x12x60x8_nt) {
int m = 10;
int n = 12;
int overall_k = 60;
int partitionK_count = 8;
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
run_partitioned_k_gemm<SgemmTraits>(m, n, overall_k, partitionK_count);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Sgemm_partitionedK_128x128x8, sgemm_128x256x100x4_nt) {
int m = 128;
int n = 256;
int overall_k = 100;
int partitionK_count = 4;
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
run_partitioned_k_gemm<SgemmTraits>(m, n, overall_k, partitionK_count);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Sgemm_partitionedK_128x128x8, sgemm_128x256x100x8_tn) {
/*
for example
partitionedK sgemm, m = 128, n = 256, overall_K = 100, partitionK_count = 8
for the first 7 partition k = overall_k / partitionK_count = 12
for the last partition last_k = overall_k - (partitionK_count - 1) * k = 16
*/
int m = 128;
int n = 256;
int overall_k = 100;
int partitionK_count = 8;
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
run_partitioned_k_gemm<SgemmTraits>(m, n, overall_k, partitionK_count);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Sgemm_partitionedK_128x128x8, sgemm_128x256x175x8_tn) {
int m = 128;
int n = 256;
int overall_k = 175;
int partitionK_count = 8;
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
run_partitioned_k_gemm<SgemmTraits>(m, n, overall_k, partitionK_count);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Sgemm_partitionedK_128x128x8, sgemm_10x12x20x3_tn) {
int m = 10;
int n = 12;
int overall_k = 20;
int partitionK_count = 3;
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
run_partitioned_k_gemm<SgemmTraits>(m, n, overall_k, partitionK_count);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Sgemm_partitionedK_128x128x8, sgemm_10x12x60x8_tn) {
int m = 10;
int n = 12;
int overall_k = 60;
int partitionK_count = 8;
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
run_partitioned_k_gemm<SgemmTraits>(m, n, overall_k, partitionK_count);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Sgemm_partitionedK_128x128x8, sgemm_128x256x100x4_tn) {
int m = 128;
int n = 256;
int overall_k = 100;
int partitionK_count = 4;
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
run_partitioned_k_gemm<SgemmTraits>(m, n, overall_k, partitionK_count);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Sgemm_partitionedK_128x128x8, sgemm_128x256x100x8_tt) {
/*
for example
partitionedK sgemm, m = 128, n = 256, overall_K = 100, partitionK_count = 8
for the first 7 partition k = overall_k / partitionK_count = 12
for the last partition last_k = overall_k - (partitionK_count - 1) * k = 16
*/
int m = 128;
int n = 256;
int overall_k = 100;
int partitionK_count = 8;
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
run_partitioned_k_gemm<SgemmTraits>(m, n, overall_k, partitionK_count);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Sgemm_partitionedK_128x128x8, sgemm_128x256x175x8_tt) {
int m = 128;
int n = 256;
int overall_k = 175;
int partitionK_count = 8;
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
run_partitioned_k_gemm<SgemmTraits>(m, n, overall_k, partitionK_count);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Sgemm_partitionedK_128x128x8, sgemm_10x12x20x3_tt) {
int m = 10;
int n = 12;
int overall_k = 20;
int partitionK_count = 3;
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
run_partitioned_k_gemm<SgemmTraits>(m, n, overall_k, partitionK_count);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Sgemm_partitionedK_128x128x8, sgemm_10x12x60x8_tt) {
int m = 10;
int n = 12;
int overall_k = 60;
int partitionK_count = 8;
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
run_partitioned_k_gemm<SgemmTraits>(m, n, overall_k, partitionK_count);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Sgemm_partitionedK_128x128x8, sgemm_128x256x100x4_tt) {
int m = 128;
int n = 256;
int overall_k = 100;
int partitionK_count = 4;
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
run_partitioned_k_gemm<SgemmTraits>(m, n, overall_k, partitionK_count);
}
////////////////////////////////////////////////////////////////////////////////////////////////////

View File

@@ -25,8 +25,12 @@
#pragma once
#include <utility>
#include "cutlass/cutlass.h"
#include "tools/test/unit/gemm/gemm_testbed.h"
#include "cutlass/gemm/device_gemm.h"
#include "cutlass/gemm/device_gemm_traits.h"
template <typename GemmTraits_>
static void run_gemm(
int m,
@@ -36,9 +40,9 @@ static void run_gemm(
int ldb,
int ldc,
typename test::GemmTestbedTraits<typename GemmTraits_::Epilogue::Scalar>::host_type alpha =
typename test::GemmTestbedTraits<typename GemmTraits_::Epilogue::Scalar>::host_type(1),
typename test::GemmTestbedTraits<typename GemmTraits_::Epilogue::Scalar>::host_type(1.0f),
typename test::GemmTestbedTraits<typename GemmTraits_::Epilogue::Scalar>::host_type beta =
typename test::GemmTestbedTraits<typename GemmTraits_::Epilogue::Scalar>::host_type(0)) {
typename test::GemmTestbedTraits<typename GemmTraits_::Epilogue::Scalar>::host_type(0.0f)) {
typedef typename GemmTraits_::KernelClass Gemm;
typename Gemm::Params params;
@@ -69,8 +73,10 @@ static void run_gemm(
if (testbed.has_cublas_support()) {
EXPECT_TRUE(testbed.verify_host_with_cublas());
EXPECT_TRUE(testbed.verify_reference_with_cublas());
}
params.initialize(testbed.M(),
testbed.N(),
testbed.K(),
@@ -137,6 +143,7 @@ static void run_gemm(
if (testbed.has_cublas_support()) {
EXPECT_TRUE(testbed.verify_host_with_cublas());
EXPECT_TRUE(testbed.verify_reference_with_cublas());
}
params.initialize(testbed.M(),
@@ -175,9 +182,9 @@ static void run_batched_strided_gemm(
int k,
int batch_count,
typename test::GemmTestbedTraits<typename GemmTraits_::Epilogue::Scalar>::host_type alpha =
typename test::GemmTestbedTraits<typename GemmTraits_::Epilogue::Scalar>::host_type(1),
typename test::GemmTestbedTraits<typename GemmTraits_::Epilogue::Scalar>::host_type(1.0f),
typename test::GemmTestbedTraits<typename GemmTraits_::Epilogue::Scalar>::host_type beta =
typename test::GemmTestbedTraits<typename GemmTraits_::Epilogue::Scalar>::host_type(0)) {
typename test::GemmTestbedTraits<typename GemmTraits_::Epilogue::Scalar>::host_type(0.0f)) {
//typedef cutlass::gemm::Gemm<GemmTraits_> Gemm;
typedef typename GemmTraits_::KernelClass Gemm;
typename Gemm::Params params;
@@ -242,3 +249,153 @@ static void run_batched_strided_gemm(
}
////////////////////////////////////////////////////////////////////////////////////////////////////
template <typename GemmTraits_, typename ReductionTraits_>
static void run_splitK_gemm(int m,
int n,
int k,
typename test::GemmTestbedTraits<typename ReductionTraits_::ScalarAlphaBeta>::host_type alpha =
typename test::GemmTestbedTraits<typename ReductionTraits_::ScalarAlphaBeta>::host_type(1.0f),
typename test::GemmTestbedTraits<typename ReductionTraits_::ScalarAlphaBeta>::host_type beta =
typename test::GemmTestbedTraits<typename ReductionTraits_::ScalarAlphaBeta>::host_type(0.0f),
bool use_host_reference = false){
test::GemmTestbed<
typename test::GemmTestbedTraits<
typename GemmTraits_::GemmConfig::ScalarA>::host_type, // AType
typename test::GemmTestbedTraits<
typename GemmTraits_::GemmConfig::ScalarB>::host_type, // BType
typename test::GemmTestbedTraits<
typename ReductionTraits_::ScalarC>::host_type, // CType
typename test::GemmTestbedTraits<
typename GemmTraits_::GemmConfig::ScalarD>::host_type, // Workspace Accumulator
typename test::GemmTestbedTraits<typename ReductionTraits_::ScalarAlphaBeta>::host_type // Scalar
>
testbed(m,
n,
k,
test::convert(GemmTraits_::kLayoutA),
test::convert(GemmTraits_::kLayoutB),
alpha,
beta);
testbed.initialize();
// create a device gemm
typedef cutlass::gemm::SplitkPIGemmTraits<GemmTraits_, ReductionTraits_> deviceGemmTraits;
typedef typename deviceGemmTraits::KernelClass deviceGemm;
typename deviceGemm::Params deviceGemmParams(testbed.M(), testbed.N(), testbed.K());
// query if workspace is needed
int workspace_size = deviceGemmParams.required_workspace_memory_in_byte();
typename test::GemmTestbedTraits<typename GemmTraits_::GemmConfig::ScalarD>::device_type
*workspace_ptr = 0;
if (workspace_size != 0) {
cudaError_t workspace_err = cudaMalloc(&workspace_ptr, workspace_size);
ASSERT_EQ(workspace_err, cudaSuccess) << "\nCUDA workspace malloc error: " << cudaGetErrorString(workspace_err)
<< "\n";
}
deviceGemmParams.initialize(testbed.alpha,
testbed.ptr_A(),
testbed.lda(),
testbed.ptr_B(),
testbed.ldb(),
testbed.beta,
testbed.ptr_C_initial(),
testbed.ldc(),
testbed.ptr_computed(),
testbed.ldc(),
workspace_ptr);
deviceGemm::launch(deviceGemmParams);
cudaError_t result = cudaDeviceSynchronize();
ASSERT_EQ(result, cudaSuccess) << "\nCUDA kernel launch error: " << cudaGetErrorString(result)
<< "\n";
if (workspace_size != 0) {
cudaError_t workspace_err = cudaFree(workspace_ptr);
ASSERT_EQ(workspace_err, cudaSuccess) << "\nCUDA workspace free error: " << cudaGetErrorString(workspace_err)
<< "\n";
}
if (use_host_reference == true || testbed.has_cublas_support() == false) {
ASSERT_TRUE(testbed.verify_with_host());
}
else {
ASSERT_TRUE(testbed.verify_with_cublas());
}
}
////////////////////////////////////////////////////////////////////////////////////////////////////
template <typename GemmTraits_>
static void run_partitioned_k_gemm(
int m,
int n,
int k,
int partitionK_count,
typename test::GemmTestbedTraits<typename GemmTraits_::Epilogue::Scalar>::host_type alpha =
typename test::GemmTestbedTraits<typename GemmTraits_::Epilogue::Scalar>::host_type(1.0f),
typename test::GemmTestbedTraits<typename GemmTraits_::Epilogue::Scalar>::host_type beta =
typename test::GemmTestbedTraits<typename GemmTraits_::Epilogue::Scalar>::host_type(0.0f)) {
//typedef cutlass::gemm::Gemm<GemmTraits_> Gemm;
typedef typename GemmTraits_::KernelClass Gemm;
typename Gemm::Params params;
test::GemmTestbed<
typename test::GemmTestbedTraits<
typename GemmTraits_::GemmConfig::ScalarA>::host_type, // AType
typename test::GemmTestbedTraits<
typename GemmTraits_::GemmConfig::ScalarB>::host_type, // BType
typename test::GemmTestbedTraits<
typename GemmTraits_::Epilogue::ScalarC>::host_type, // CType
typename test::GemmTestbedTraits<
typename GemmTraits_::Epilogue::Accumulators::Element>::host_type, // Accumulator
typename test::GemmTestbedTraits<typename GemmTraits_::Epilogue::Scalar>::host_type // Scalar
>
testbed(m,
n,
std::make_pair(k, partitionK_count),
test::convert(GemmTraits_::kLayoutA),
test::convert(GemmTraits_::kLayoutB),
alpha,
beta);
testbed.initialize();
// host support is not implemented for strided batched gemm
// if (testbed.has_cublas_support()) {
// EXPECT_TRUE(testbed.verify_host_with_cublas());
//}
params.initialize(testbed.M(),
testbed.N(),
testbed.K(),
testbed.alpha,
testbed.ptr_A(),
testbed.lda(),
testbed.ptr_B(),
testbed.ldb(),
testbed.beta,
testbed.ptr_C_initial(),
testbed.ldc(),
testbed.ptr_computed(),
testbed.ldc(),
partitionK_count);
Gemm::launch(params);
cudaError_t result = cudaDeviceSynchronize();
ASSERT_EQ(result, cudaSuccess) << "\nCUDA kernel launch error: " << cudaGetErrorString(result)
<< "\n";
if (testbed.has_cublas_support()) {
ASSERT_TRUE(testbed.verify_with_cublas());
}
else {
// ASSERT_TRUE(testbed.verify_with_host());
ASSERT_TRUE(false) << "host support is not implemented for strided batched gemm" << std::endl;
}
}

View File

@@ -0,0 +1,247 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#include "cutlass_unit_test.h"
#include "cutlass/gemm/gemm.h"
#include "cutlass/gemm/dgemm_traits.h"
#include "cutlass/reduction/batched_reduction_traits.h"
#include "tools/test/unit/gemm/gemm_testbed.h"
#include "tools/test/unit/gemm/run_gemm.h"
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_dgemm_128x128x8_splits16, dgemm_128x256x512_nn) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 512;
/*batched sgemm traits*/
typedef cutlass::gemm::DgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
DgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<double,
double,
double,
double,
double, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<DgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_dgemm_128x128x8_splits16, dgemm_128x256x512_nt) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 512;
/*batched sgemm traits*/
typedef cutlass::gemm::DgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
DgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<double,
double,
double,
double,
double, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<DgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_dgemm_128x128x8_splits16, dgemm_128x256x512_tn) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 512;
/*batched sgemm traits*/
typedef cutlass::gemm::DgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
DgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<double,
double,
double,
double,
double, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<DgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_dgemm_128x128x8_splits16, dgemm_128x256x512_tt) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 512;
/*batched sgemm traits*/
typedef cutlass::gemm::DgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
DgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<double,
double,
double,
double,
double, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<DgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_dgemm_128x128x8_splits16, dgemm_128x256x500_nn) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 500;
/*batched sgemm traits*/
typedef cutlass::gemm::DgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
DgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<double,
double,
double,
double,
double, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<DgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_dgemm_128x128x8_splits16, dgemm_128x256x500_nt) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 500;
/*batched sgemm traits*/
typedef cutlass::gemm::DgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
DgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<double,
double,
double,
double,
double, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<DgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_dgemm_128x128x8_splits16, dgemm_128x256x500_tn) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 500;
/*batched sgemm traits*/
typedef cutlass::gemm::DgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
DgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<double,
double,
double,
double,
double, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<DgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_dgemm_128x128x8_splits16, dgemm_128x256x500_tt) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 500;
/*batched sgemm traits*/
typedef cutlass::gemm::DgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
DgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<double,
double,
double,
double,
double, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<DgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}

View File

@@ -0,0 +1,579 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#include "cutlass_unit_test.h"
#include "cutlass/gemm/gemm.h"
#include "cutlass/gemm/fp16_sgemm_traits.h"
#include "cutlass/reduction/batched_reduction_traits.h"
#include "tools/test/unit/gemm/gemm_testbed.h"
#include "tools/test/unit/gemm/run_gemm.h"
////////////////////////////////////////////////////////////////////////////////////////////////////
/*
for fp16_sgemm_fp16 A, B, C and D are half typed. alpha and beta can be half or float typed.
Accumulation is float typed.
1. in batched gemm kernel, Ab and Bb are half typed, and pointing to A and B.
Cb and Db are float typed, since Db is actually pointing to the workspace memory
thus is of the same type with accumulation. Cb is generally ignored since beta is zero. alpha is one.
2. in the reduction kernel. Dr = alpha * Reduction(Ar) + beta * Cr. Ar is float typed and pointing to the same
workspace memory with Db. Cr is half typed and pointing to C. Dr is half typed and pointing to D.
ALPHAr is the same with alpha, BETAr is the same with beta.
*/
TEST(SplitK_fp16_sgemm_fp16_alphabetaFloat_128x128x8_splits16, sgemm_128x256x512_nn) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 512;
/*batched sgemm traits*/
typedef cutlass::gemm::Fp16SgemmSgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<16, 128, 128>,
half, /*Ab type*/
half, /*Bb type*/
float, /*Cb type*/
float, /*Db type*/
float /*alpha, beta type*/
>
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float, /*Ar type*/
half, /*Cr type*/
half, /*Dr type*/
float, /*alpha, beta type*/
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f, true/*use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_fp16_sgemm_fp16_alphabetaFloat_128x128x8_splits16, sgemm_128x256x512_nt) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 512;
/*batched sgemm traits*/
typedef cutlass::gemm::Fp16SgemmSgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<16, 128, 128>,
half, /*Ab type*/
half, /*Bb type*/
float, /*Cb type*/
float, /*Db type*/
float /*alpha, beta type*/
>
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float, /*Ar type*/
half, /*Cr type*/
half, /*Dr type*/
float, /*alpha, beta type*/
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f, true/*use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_fp16_sgemm_fp16_alphabetaFloat_128x128x8_splits16, sgemm_128x256x512_tn) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 512;
/*batched sgemm traits*/
typedef cutlass::gemm::Fp16SgemmSgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<16, 128, 128>,
half, /*Ab type*/
half, /*Bb type*/
float, /*Cb type*/
float, /*Db type*/
float /*alpha, beta type*/
>
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float, /*Ar type*/
half, /*Cr type*/
half, /*Dr type*/
float, /*alpha, beta type*/
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f, true/*use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_fp16_sgemm_fp16_alphabetaFloat_128x128x8_splits16, sgemm_128x256x512_tt) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 512;
/*batched sgemm traits*/
typedef cutlass::gemm::Fp16SgemmSgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<16, 128, 128>,
half, /*Ab type*/
half, /*Bb type*/
float, /*Cb type*/
float, /*Db type*/
float /*alpha, beta type*/
>
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float, /*Ar type*/
half, /*Cr type*/
half, /*Dr type*/
float, /*alpha, beta type*/
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f, true/*use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_fp16_sgemm_fp16_alphabetaFloat_128x128x8_splits16, sgemm_128x256x500_nn) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 500;
/*batched sgemm traits*/
typedef cutlass::gemm::Fp16SgemmSgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<16, 128, 128>,
half, /*Ab type*/
half, /*Bb type*/
float, /*Cb type*/
float, /*Db type*/
float /*alpha, beta type*/
>
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float, /*Ar type*/
half, /*Cr type*/
half, /*Dr type*/
float, /*alpha, beta type*/
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f, true/*use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_fp16_sgemm_fp16_alphabetaFloat_128x128x8_splits16, sgemm_128x256x500_nt) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 500;
/*batched sgemm traits*/
typedef cutlass::gemm::Fp16SgemmSgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<16, 128, 128>,
half, /*Ab type*/
half, /*Bb type*/
float, /*Cb type*/
float, /*Db type*/
float /*alpha, beta type*/
>
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float, /*Ar type*/
half, /*Cr type*/
half, /*Dr type*/
float, /*alpha, beta type*/
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f, true/*use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_fp16_sgemm_fp16_alphabetaFloat_128x128x8_splits16, sgemm_128x256x500_tn) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 500;
/*batched sgemm traits*/
typedef cutlass::gemm::Fp16SgemmSgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<16, 128, 128>,
half, /*Ab type*/
half, /*Bb type*/
float, /*Cb type*/
float, /*Db type*/
float /*alpha, beta type*/
>
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float, /*Ar type*/
half, /*Cr type*/
half, /*Dr type*/
float, /*alpha, beta type*/
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f, true/*use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_fp16_sgemm_fp16_alphabetaFloat_128x128x8_splits16, sgemm_128x256x500_tt) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 500;
/*batched sgemm traits*/
typedef cutlass::gemm::Fp16SgemmSgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<16, 128, 128>,
half, /*Ab type*/
half, /*Bb type*/
float, /*Cb type*/
float, /*Db type*/
float /*alpha, beta type*/
>
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float, /*Ar type*/
half, /*Cr type*/
half, /*Dr type*/
float, /*alpha, beta type*/
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f, true/*use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_fp16_sgemm_fp16_alphabetaFp16_128x128x8_splits16, sgemm_128x256x512_nn) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 512;
/*batched sgemm traits*/
typedef cutlass::gemm::Fp16SgemmSgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<16, 128, 128>,
half, /*Ab type*/
half, /*Bb type*/
float, /*Cb type*/
float, /*Db type*/
float /*alpha, beta type*/
>
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float, /*Ar type*/
half, /*Cr type*/
half, /*Dr type*/
half, /*alpha, beta type*/
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f, true/*use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_fp16_sgemm_fp16_alphabetaFp16_128x128x8_splits16, sgemm_128x256x512_nt) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 512;
/*batched sgemm traits*/
typedef cutlass::gemm::Fp16SgemmSgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<16, 128, 128>,
half, /*Ab type*/
half, /*Bb type*/
float, /*Cb type*/
float, /*Db type*/
float /*alpha, beta type*/
>
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float, /*Ar type*/
half, /*Cr type*/
half, /*Dr type*/
half, /*alpha, beta type*/
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f, true/*use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_fp16_sgemm_fp16_alphabetaFp16_128x128x8_splits16, sgemm_128x256x512_tn) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 512;
/*batched sgemm traits*/
typedef cutlass::gemm::Fp16SgemmSgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<16, 128, 128>,
half, /*Ab type*/
half, /*Bb type*/
float, /*Cb type*/
float, /*Db type*/
float /*alpha, beta type*/
>
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float, /*Ar type*/
half, /*Cr type*/
half, /*Dr type*/
half, /*alpha, beta type*/
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f, true/*use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_fp16_sgemm_fp16_alphabetaFp16_128x128x8_splits16, sgemm_128x256x512_tt) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 512;
/*batched sgemm traits*/
typedef cutlass::gemm::Fp16SgemmSgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<16, 128, 128>,
half, /*Ab type*/
half, /*Bb type*/
float, /*Cb type*/
float, /*Db type*/
float /*alpha, beta type*/
>
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float, /*Ar type*/
half, /*Cr type*/
half, /*Dr type*/
half, /*alpha, beta type*/
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f, true/*use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
//k = 500
TEST(SplitK_fp16_sgemm_fp16_alphabetaFp16_128x128x8_splits16, sgemm_128x256x500_nn) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 500;
/*batched sgemm traits*/
typedef cutlass::gemm::Fp16SgemmSgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<16, 128, 128>,
half, /*Ab type*/
half, /*Bb type*/
float, /*Cb type*/
float, /*Db type*/
float /*alpha, beta type*/
>
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float, /*Ar type*/
half, /*Cr type*/
half, /*Dr type*/
half, /*alpha, beta type*/
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f, true/*use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_fp16_sgemm_fp16_alphabetaFp16_128x128x8_splits16, sgemm_128x256x500_nt) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 500;
/*batched sgemm traits*/
typedef cutlass::gemm::Fp16SgemmSgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<16, 128, 128>,
half, /*Ab type*/
half, /*Bb type*/
float, /*Cb type*/
float, /*Db type*/
float /*alpha, beta type*/
>
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float, /*Ar type*/
half, /*Cr type*/
half, /*Dr type*/
half, /*alpha, beta type*/
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f, true/*use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_fp16_sgemm_fp16_alphabetaFp16_128x128x8_splits16, sgemm_128x256x500_tn) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 500;
/*batched sgemm traits*/
typedef cutlass::gemm::Fp16SgemmSgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<16, 128, 128>,
half, /*Ab type*/
half, /*Bb type*/
float, /*Cb type*/
float, /*Db type*/
float /*alpha, beta type*/
>
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float, /*Ar type*/
half, /*Cr type*/
half, /*Dr type*/
half, /*alpha, beta type*/
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f, true/*use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_fp16_sgemm_fp16_alphabetaFp16_128x128x8_splits16, sgemm_128x256x500_tt) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 500;
/*batched sgemm traits*/
typedef cutlass::gemm::Fp16SgemmSgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<16, 128, 128>,
half, /*Ab type*/
half, /*Bb type*/
float, /*Cb type*/
float, /*Db type*/
float /*alpha, beta type*/
>
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float, /*Ar type*/
half, /*Cr type*/
half, /*Dr type*/
half, /*alpha, beta type*/
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f, true/*use host reference*/);
}

View File

@@ -0,0 +1,248 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#include "cutlass_unit_test.h"
#include "cutlass/gemm/gemm.h"
#include "cutlass/gemm/hgemm_traits.h"
#include "cutlass/reduction/batched_reduction_traits.h"
#include "tools/test/unit/gemm/gemm_testbed.h"
#include "tools/test/unit/gemm/run_gemm.h"
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_hgemm_128x128x8_splits16, hgemm_128x256x64_nn) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 64;
/*batched sgemm traits*/
typedef cutlass::gemm::HgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
HgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<half,
half,
half,
half,
half, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<HgemmTraits, BatchedReductionTraits>(m, n, k, 1.0f, 0.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_hgemm_128x128x8_splits16, hgemm_128x256x64_nt) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 64;
/*batched sgemm traits*/
typedef cutlass::gemm::HgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
HgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<half,
half,
half,
half,
half, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<HgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_hgemm_128x128x8_splits16, hgemm_128x256x64_tn) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 64;
/*batched sgemm traits*/
typedef cutlass::gemm::HgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
HgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<half,
half,
half,
half,
half, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<HgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_hgemm_128x128x8_splits16, hgemm_128x256x64_tt) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 64;
/*batched sgemm traits*/
typedef cutlass::gemm::HgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
HgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<half,
half,
half,
half,
half, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<HgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_hgemm_128x128x8_splits16, hgemm_128x256x66_nn) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 66;
/*batched sgemm traits*/
typedef cutlass::gemm::HgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
HgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<half,
half,
half,
half,
half, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<HgemmTraits, BatchedReductionTraits>(m, n, k, 1.0f, 0.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_hgemm_128x128x8_splits16, hgemm_128x256x66_nt) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 66;
/*batched sgemm traits*/
typedef cutlass::gemm::HgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
HgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<half,
half,
half,
half,
half, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<HgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_hgemm_128x128x8_splits16, hgemm_128x256x66_tn) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 66;
/*batched sgemm traits*/
typedef cutlass::gemm::HgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
HgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<half,
half,
half,
half,
half, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<HgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_hgemm_128x128x8_splits16, hgemm_128x256x66_tt) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 66;
/*batched sgemm traits*/
typedef cutlass::gemm::HgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
HgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<half,
half,
half,
half,
half, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<HgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}

View File

@@ -0,0 +1,367 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#include "cutlass_unit_test.h"
#include "cutlass/gemm/gemm.h"
#include "cutlass/gemm/igemm_traits.h"
#include "cutlass/reduction/batched_reduction_traits.h"
#include "tools/test/unit/gemm/gemm_testbed.h"
#include "tools/test/unit/gemm/run_gemm.h"
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_igemm_128x128x32_splits16, igemm_128x256x512_nn) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 512;
/*batched igemm traits*/
typedef cutlass::gemm::IgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 128, 128>, int, cutlass::gemm::LinearScaling<int> >
IgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<int,
int,
int,
int,
int, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<IgemmTraits, BatchedReductionTraits>(m, n, k, 2, 1, true /*use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_igemm_128x128x32_splits16, igemm_128x256x512_nt) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 512;
/*batched igemm traits*/
typedef cutlass::gemm::IgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 128, 128>, int, cutlass::gemm::LinearScaling<int> >
IgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<int,
int,
int,
int,
int, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<IgemmTraits, BatchedReductionTraits>(m, n, k, 2, 1, true /*use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_igemm_128x128x32_splits16, igemm_128x256x512_tn) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 512;
/*batched igemm traits*/
typedef cutlass::gemm::IgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 128, 128>, int, cutlass::gemm::LinearScaling<int> >
IgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<int,
int,
int,
int,
int, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<IgemmTraits, BatchedReductionTraits>(m, n, k, 2, 1, true /*use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_igemm_128x128x32_splits16, igemm_128x256x512_tt) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 512;
/*batched igemm traits*/
typedef cutlass::gemm::IgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 128, 128>, int, cutlass::gemm::LinearScaling<int> >
IgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<int,
int,
int,
int,
int, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<IgemmTraits, BatchedReductionTraits>(m, n, k, 2, 1, true /*use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_igemm_128x128x32_splits16, igemm_1024x64x4096_nn) {
const int splits_count = 16;
const int m = 1024;
const int n = 64;
const int k = 4096;
/*batched igemm traits*/
typedef cutlass::gemm::IgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 128, 128>, int, cutlass::gemm::LinearScaling<int> >
IgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<int,
int,
int,
int,
int, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<IgemmTraits, BatchedReductionTraits>(m, n, k, 1, 0, false /*not use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_igemm_128x128x32_splits16, igemm_1024x64x4096_nt) {
const int splits_count = 16;
const int m = 1024;
const int n = 64;
const int k = 4096;
/*batched igemm traits*/
typedef cutlass::gemm::IgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 128, 128>, int, cutlass::gemm::LinearScaling<int> >
IgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<int,
int,
int,
int,
int, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<IgemmTraits, BatchedReductionTraits>(m, n, k, 1, 0, false /*not use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_igemm_128x128x32_splits16, igemm_1024x64x4096_tn) {
const int splits_count = 16;
const int m = 1024;
const int n = 64;
const int k = 4096;
/*batched igemm traits*/
typedef cutlass::gemm::IgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 128, 128>, int, cutlass::gemm::LinearScaling<int> >
IgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<int,
int,
int,
int,
int, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<IgemmTraits, BatchedReductionTraits>(m, n, k, 1, 0, false /*not use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_igemm_128x128x32_splits16, igemm_1024x64x4096_tt) {
const int splits_count = 16;
const int m = 1024;
const int n = 64;
const int k = 4096;
/*batched igemm traits*/
typedef cutlass::gemm::IgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 128, 128>, int, cutlass::gemm::LinearScaling<int> >
IgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<int,
int,
int,
int,
int, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<IgemmTraits, BatchedReductionTraits>(m, n, k, 1, 0, false /*not use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_igemm_128x32x32_splits16, igemm_1024x64x4096_nn) {
const int splits_count = 16;
const int m = 1024;
const int n = 64;
const int k = 4096;
/*batched igemm traits*/
typedef cutlass::gemm::IgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 32, 128>, int, cutlass::gemm::LinearScaling<int> >
IgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<int,
int,
int,
int,
int, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<IgemmTraits, BatchedReductionTraits>(m, n, k, 1, 0, false /*not use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_igemm_128x32x32_splits16, igemm_1024x64x4096_nt) {
const int splits_count = 16;
const int m = 1024;
const int n = 64;
const int k = 4096;
/*batched igemm traits*/
typedef cutlass::gemm::IgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 32, 128>, int, cutlass::gemm::LinearScaling<int> >
IgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<int,
int,
int,
int,
int, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<IgemmTraits, BatchedReductionTraits>(m, n, k, 1, 0, false /*not use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_igemm_128x32x32_splits16, igemm_1024x64x4096_tn) {
const int splits_count = 16;
const int m = 1024;
const int n = 64;
const int k = 4096;
/*batched igemm traits*/
typedef cutlass::gemm::IgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 32, 128>, int, cutlass::gemm::LinearScaling<int> >
IgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<int,
int,
int,
int,
int, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<IgemmTraits, BatchedReductionTraits>(m, n, k, 1, 0, false /*not use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_igemm_128x32x32_splits16, igemm_1024x64x4096_tt) {
const int splits_count = 16;
const int m = 1024;
const int n = 64;
const int k = 4096;
/*batched igemm traits*/
typedef cutlass::gemm::IgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 32, 128>, int, cutlass::gemm::LinearScaling<int> >
IgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<int,
int,
int,
int,
int, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<IgemmTraits, BatchedReductionTraits>(m, n, k, 1, 0, false /*not use host reference*/);
}

View File

@@ -0,0 +1,355 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#include "cutlass_unit_test.h"
#include "cutlass/gemm/gemm.h"
#include "cutlass/gemm/sgemm_traits.h"
#include "cutlass/reduction/batched_reduction_traits.h"
#include "tools/test/unit/gemm/gemm_testbed.h"
#include "tools/test/unit/gemm/run_gemm.h"
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_sgemm_128x128x8_splits16, sgemm_128x256x512_nn) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 512;
/*batched sgemm traits*/
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float,
float,
float,
float,
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_sgemm_128x128x8_splits16, sgemm_128x256x512_nt) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 512;
/*batched sgemm traits*/
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float,
float,
float,
float,
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_sgemm_128x128x8_splits16, sgemm_128x256x512_tn) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 512;
/*batched sgemm traits*/
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float,
float,
float,
float,
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_sgemm_128x128x8_splits16, sgemm_128x256x512_tt) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 512;
/*batched sgemm traits*/
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float,
float,
float,
float,
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_sgemm_128x128x8_splits16, sgemm_128x256x500_nn) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 500;
/*batched sgemm traits*/
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float,
float,
float,
float,
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_sgemm_128x128x8_splits16, sgemm_128x256x500_nt) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 500;
/*batched sgemm traits*/
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float,
float,
float,
float,
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_sgemm_128x128x8_splits16, sgemm_128x256x500_tn) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 500;
/*batched sgemm traits*/
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float,
float,
float,
float,
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_sgemm_128x128x8_splits16, sgemm_128x256x500_tt) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 500;
/*batched sgemm traits*/
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float,
float,
float,
float,
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_sgemm_128x128x8_splits16, sgemm_1024x64x4096_nn) {
const int splits_count = 16;
const int m = 1024;
const int n = 64;
const int k = 4096;
/*batched sgemm traits*/
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float,
float,
float,
float,
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_sgemm_128x128x8_splits16, sgemm_1024x64x4096_nt) {
const int splits_count = 16;
const int m = 1024;
const int n = 64;
const int k = 4096;
/*batched sgemm traits*/
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float,
float,
float,
float,
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_sgemm_128x128x8_splits16, sgemm_1024x64x4096_tn) {
const int splits_count = 16;
const int m = 1024;
const int n = 64;
const int k = 4096;
/*batched sgemm traits*/
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float,
float,
float,
float,
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_sgemm_128x128x8_splits16, sgemm_1024x64x4096_tt) {
const int splits_count = 16;
const int m = 1024;
const int n = 64;
const int k = 4096;
/*batched sgemm traits*/
typedef cutlass::gemm::SgemmTraits<cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor, cutlass::Shape<8, 128, 128> >
SgemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<float,
float,
float,
float,
float, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<SgemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f);
}

View File

@@ -0,0 +1,175 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#include "cutlass/wmma_matrix.h"
#if defined(CUTLASS_USE_WMMA_API)
#include "cutlass_unit_test.h"
#include "cutlass/gemm/gemm.h"
#include "cutlass/gemm/wmma_gemm_traits.h"
#include "cutlass/reduction/batched_reduction_traits.h"
#include "tools/test/unit/gemm/gemm_testbed.h"
#include "tools/test/unit/gemm/run_gemm.h"
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_wmma_gemm_16x16x32_splits16, wmma_gemm_128x256x512_nn) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 512;
/*batched wmma gemm traits*/
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
half,
cutlass::gemm::LinearScaling<half>,
half
>
WmmaGemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<half,
half,
half,
half,
half, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<WmmaGemmTraits, BatchedReductionTraits>(m, n, k, 2.0f, 1.0f, true/*use host reference*/);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_wmma_gemm_16x16x32_splits16, wmma_gemm_128x256x512_nt) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 512;
/*batched wmma gemm traits*/
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
half,
cutlass::gemm::LinearScaling<half>,
half
>
WmmaGemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<half,
half,
half,
half,
half, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<WmmaGemmTraits, BatchedReductionTraits>(m, n, k, 1.0f, 0.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_wmma_gemm_16x16x32_splits16, wmma_gemm_128x256x512_tn) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 512;
/*batched wmma gemm traits*/
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
half,
cutlass::gemm::LinearScaling<half>,
half
>
WmmaGemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<half,
half,
half,
half,
half, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<WmmaGemmTraits, BatchedReductionTraits>(m, n, k, 1.0f, 0.0f);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SplitK_wmma_gemm_16x16x32_splits16, wmma_gemm_128x256x512_tt) {
const int splits_count = 16;
const int m = 128;
const int n = 256;
const int k = 512;
/*batched wmma gemm traits*/
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
half,
cutlass::gemm::LinearScaling<half>,
half
>
WmmaGemmTraits;
/*batched reduction traits*/
typedef cutlass::reduction::BatchedReductionTraits<half,
half,
half,
half,
half, /*accumulation type*/
splits_count,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits;
run_splitK_gemm<WmmaGemmTraits, BatchedReductionTraits>(m, n, k, 1.0f, 0.0f);
}
#endif

View File

@@ -53,6 +53,7 @@ TEST(WmmaGemm_16x16x32_f16, wmma_gemm_16x16x16_nn) {
run_gemm<WmmaGemmTraits>(16, 16, 16);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_16x16x32_f16, wmma_gemm_16x16x32_nn) {
@@ -367,7 +368,5 @@ TEST(WmmaGemm_128x128x32, wmma_32x8x16_gemm_256x256x128_tn) {
run_gemm<WmmaGemmTraits>(256, 256, 128);
}
#endif
////////////////////////////////////////////////////////////////////////////////////////////////////
#endif // defined CUTLASS_USE_WMMA_API

View File

@@ -0,0 +1,155 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#include "cutlass/wmma_matrix.h"
#if defined(CUTLASS_USE_WMMA_API)
#include "cutlass_unit_test.h"
#include "cutlass/gemm/gemm.h"
#include "cutlass/gemm/wmma_gemm_traits.h"
#include "tools/test/unit/gemm/gemm_testbed.h"
#include "tools/test/unit/gemm/run_gemm.h"
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_16x16x32_f16, wmma_gemm_36x36x16_nn) {
/*
this wmmaTraits requires leading dim to be divisible by 4
*/
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
half,
cutlass::gemm::LinearScaling<half>,
half,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
4, /*kScalarsPerLdgA_*/
4, /*kScalarsPerLdgB_*/
4, /*KScalarsPerLdsA_*/
4, /*KScalarsPerLdsB_*/
4 / sizeof(half), /*kScalarsPerLdgCAndStgD_*/
4 / sizeof(half), /*kScalarsPerStsD_*/
4 / sizeof(half) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_gemm<WmmaGemmTraits>(36, 36, 64);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_16x16x32_f16, wmma_gemm_36x36x16_nt) {
/*
this wmmaTraits requires leading dim to be divisible by 4
*/
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kColumnMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
half,
cutlass::gemm::LinearScaling<half>,
half,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
4, /*kScalarsPerLdgA_*/
4, /*kScalarsPerLdgB_*/
4, /*KScalarsPerLdsA_*/
4, /*KScalarsPerLdsB_*/
4 / sizeof(half), /*kScalarsPerLdgCAndStgD_*/
4 / sizeof(half), /*kScalarsPerStsD_*/
4 / sizeof(half) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_gemm<WmmaGemmTraits>(36, 36, 64);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_16x16x32_f16, wmma_gemm_36x36x16_tn) {
/*
this wmmaTraits requires leading dim to be divisible by 4
*/
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kColumnMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
half,
cutlass::gemm::LinearScaling<half>,
half,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
4, /*kScalarsPerLdgA_*/
4, /*kScalarsPerLdgB_*/
4, /*KScalarsPerLdsA_*/
4, /*KScalarsPerLdsB_*/
4 / sizeof(half), /*kScalarsPerLdgCAndStgD_*/
4 / sizeof(half), /*kScalarsPerStsD_*/
4 / sizeof(half) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_gemm<WmmaGemmTraits>(36, 36, 64);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(WmmaGemm_16x16x32_f16, wmma_gemm_36x36x16_tt) {
/*
this wmmaTraits requires leading dim to be divisible by 4
*/
typedef cutlass::gemm::WmmaGemmTraits<
cutlass::MatrixLayout::kRowMajor,
cutlass::MatrixLayout::kRowMajor,
cutlass::Shape<32, 16, 16>,
half,
half,
half,
cutlass::gemm::LinearScaling<half>,
half,
typename cutlass::gemm::WmmaGemmAccumulatorsPerWarp<typename cutlass::Shape<32, 16, 16> >::Shape,
typename cutlass::Shape<16, 16, 16>,
4, /*kScalarsPerLdgA_*/
4, /*kScalarsPerLdgB_*/
4, /*KScalarsPerLdsA_*/
4, /*KScalarsPerLdsB_*/
4 / sizeof(half), /*kScalarsPerLdgCAndStgD_*/
4 / sizeof(half), /*kScalarsPerStsD_*/
4 / sizeof(half) /*kScalarsPerLdsD_*/
>
WmmaGemmTraits;
run_gemm<WmmaGemmTraits>(36, 36, 64);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
#endif

View File

@@ -0,0 +1,307 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#include "cutlass_unit_test.h"
#include "cutlass/shape.h"
#include "tools/util/host_tensor.h"
#include "cutlass/reduction/batched_reduction.h"
#include "cutlass/reduction/batched_reduction_traits.h"
#include "tools/test/unit/reduction/test_batched_reduction.h"
#include "tools/test/unit/reduction/batched_reduction_testbed.h"
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Batched_reduction_float, batched_reduction_128x256x16) {
/*
The output matrix is 128x256
The input matrix is 128x256x16
The reduction will be applied at the third dim of input matrix
*/
const int m = 128;
const int n = 256;
const int lda = 128;
const int ldc = 128;
const int ldd = 128;
const int reduction_size = 16;
typedef cutlass::reduction::BatchedReductionTraits<float, /*A*/
float, /*C*/
float, /*D*/
float, /*alpha and beta*/
float, /*accumulation type*/
reduction_size,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits_16;
test_batched_reduction<BatchedReductionTraits_16>(m, n, lda, ldc, ldd);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Batched_reduction_double, batched_reduction_128x256x16) {
/*
D = alpha * Reduction(A) + beta * C
The output matrix D is 128x256
The input matrix A is 128x256x16
The input matrix C is 128x256
The reduction will be applied at the third dim of input matrix
*/
const int m = 128;
const int n = 256;
const int lda = 128;
const int ldc = 128;
const int ldd = 128;
const int reduction_size = 16;
typedef cutlass::reduction::BatchedReductionTraits<double,
double,
double,
double,
double, /*accumulation type*/
reduction_size,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits_16;
test_batched_reduction<BatchedReductionTraits_16>(m, n, lda, ldc, ldd);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Batched_reduction_half, batched_reduction_128x256x16) {
/*
The output matrix is 128x256
The input matrix is 128x256x16
The reduction will be applied at the third dim of input matrix
*/
const int m = 128;
const int n = 256;
const int lda = 128;
const int ldc = 128;
const int ldd = 128;
const int reduction_size = 16;
typedef cutlass::reduction::BatchedReductionTraits<half,
half,
half,
half,
half, /*accumulation type*/
reduction_size,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits_16;
test_batched_reduction<BatchedReductionTraits_16>(m, n, lda, ldc, ldd);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Batched_reduction_float, batched_reduction_128x64x80) {
/*
The output matrix is 128x64
The input matrix is 128x64x80
The reduction will be applied at the third dim of input matrix
*/
const int m = 128;
const int n = 64;
const int lda = 128;
const int ldc = 128;
const int ldd = 128;
const int reduction_size = 80;
typedef cutlass::reduction::BatchedReductionTraits<float,
float,
float,
float,
float, /*accumulation type*/
reduction_size,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits_80;
test_batched_reduction<BatchedReductionTraits_80>(m, n, lda, ldc, ldd);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Batched_reduction_double, batched_reduction_128x64x80) {
/*
The output matrix is 128x64
The input matrix is 128x64x80
The reduction will be applied at the third dim of input matrix
*/
const int m = 128;
const int n = 64;
const int lda = 128;
const int ldc = 128;
const int ldd = 128;
const int reduction_size = 80;
typedef cutlass::reduction::BatchedReductionTraits<double,
double,
double,
double,
double, /*accumulation type*/
reduction_size,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits_80;
test_batched_reduction<BatchedReductionTraits_80>(m, n, lda, ldc, ldd);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Batched_reduction_half, batched_reduction_128x64x80) {
/*
The output matrix is 128x64
The input matrix is 128x64x80
The reduction will be applied at the third dim of input matrix
*/
const int m = 128;
const int n = 64;
const int lda = 128;
const int ldc = 128;
const int ldd = 128;
const int reduction_size = 80;
typedef cutlass::reduction::BatchedReductionTraits<half,
half,
half,
half,
half, /*accumulation type*/
reduction_size,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits_80;
test_batched_reduction<BatchedReductionTraits_80>(m, n, lda, ldc, ldd);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Batched_reduction_float_threadShape1, batched_reduction_128x256x90) {
/*
The output matrix is 128x256
The input matrix is 128x256x90
The reduction will be applied at the third dim of input matrix
*/
const int m = 128;
const int n = 256;
const int lda = 128;
const int ldc = 128;
const int ldd = 128;
const int reduction_size = 90;
typedef cutlass::reduction::BatchedReductionTraits<float, /*A*/
float, /*C*/
float, /*D*/
float, /*alpha and beta*/
float, /*accumulation type*/
reduction_size,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 1> >
BatchedReductionTraits_16;
test_batched_reduction<BatchedReductionTraits_16>(m, n, lda, ldc, ldd);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Batched_reduction_double_threadShape1, batched_reduction_128x256x90) {
/*
The output matrix is 128x256
The input matrix is 128x256x90
The reduction will be applied at the third dim of input matrix
*/
const int m = 128;
const int n = 256;
const int lda = 128;
const int ldc = 128;
const int ldd = 128;
const int reduction_size = 90;
typedef cutlass::reduction::BatchedReductionTraits<double, /*A*/
double, /*C*/
double, /*D*/
double, /*alpha and beta*/
double, /*accumulation type*/
reduction_size,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 1> >
BatchedReductionTraits_16;
test_batched_reduction<BatchedReductionTraits_16>(m, n, lda, ldc, ldd);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Batched_reduction_half_threadShape1, batched_reduction_128x256x90) {
/*
The output matrix is 128x256
The input matrix is 128x256x90
The reduction will be applied at the third dim of input matrix
*/
const int m = 128;
const int n = 256;
const int lda = 128;
const int ldc = 128;
const int ldd = 128;
const int reduction_size = 90;
typedef cutlass::reduction::BatchedReductionTraits<half, /*A*/
half, /*C*/
half, /*D*/
half, /*alpha and beta*/
half, /*accumulation type*/
reduction_size,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 1> >
BatchedReductionTraits_16;
test_batched_reduction<BatchedReductionTraits_16>(m, n, lda, ldc, ldd);
}

View File

@@ -0,0 +1,301 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (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 Test environment for batched reduction
*/
#pragma once
#include "cutlass/matrix_traits.h"
#include "cutlass/util/platform.h"
#include "tools/util/host_matrix.h"
#include "tools/util/host_matrix_view.h"
#include "tools/util/host_tensor.h"
#include "tools/util/tensor_view_io.h"
#include "tools/util/type_traits.h"
#include <assert.h>
namespace test {
inline cublasOperation_t convert(cutlass::MatrixLayout::Kind layout) {
switch (layout) {
case cutlass::MatrixLayout::kRowMajor:
return CUBLAS_OP_T;
case cutlass::MatrixLayout::kColumnMajor:
return CUBLAS_OP_N;
default:
break;
}
return CUBLAS_OP_N;
}
inline cutlass::MatrixLayout::Kind convert(cublasOperation_t transform) {
switch (transform) {
case CUBLAS_OP_T:
return cutlass::MatrixLayout::kRowMajor;
case CUBLAS_OP_N:
return cutlass::MatrixLayout::kColumnMajor;
default:
break;
}
return cutlass::MatrixLayout::kColumnMajor;
}
/// Testbed for evaluating batched reduction
template <
typename AType,
typename CType,
typename DType,
typename ScalarAlpha,
typename ScalarBeta,
typename ScalarAccum,
// input matrix depth size to be sumed
int ReductionSize
>
struct BatchedReductionTestbed {
//
// Type definitions
//
/// Host tensor for operand C
typedef cutlass::HostTensor<AType, 3> HostTensorA;
/// Host tensor for operand C
typedef cutlass::HostMatrix<CType> HostMatrixC;
/// Host tensor for operand D
typedef cutlass::HostMatrix<DType> HostMatrixD;
/// Generates random elements
template <typename T>
struct RandomGenerator {
RandomGenerator(int seed = -1, bool only_ones_ = false) : only_ones(only_ones_) { srand(seed); }
T operator()() {
if (only_ones) {
return T(1);
}
else {
int val = (rand() % 16) - 8;
return T(val);
}
}
bool only_ones;
};
template <typename T>
struct RandomBitGenerator {
RandomBitGenerator(int seed = -1) { srand(seed); }
T operator()() {
uint32_t val = 0;
for (int i = 0; i < 32; i++) {
val |= rand() % 2;
val <<= 1;
}
return T(val);
}
};
//
// Data members
//
/// input/output number of rows
int m;
/// input/output number of columns
int n;
/// A matrix operand, always column major, no trans
HostTensorA A;
/// C matrix operand, always column major, no trans
HostMatrixC C;
/// D matrix operand, always column major, no trans
HostMatrixD D;
/// Reference
cutlass::HostTensor<AType, 3> ref_A;
///
cutlass::HostMatrix<CType> ref_C;
/// Reference result computed on the host
cutlass::HostMatrix<DType> ref_D;
/// lda
int lda;
/// ldc
int ldc;
/// ldd
int ldd;
/// Linear scalaring factor
ScalarAlpha alpha;
/// Linear scaling factor
ScalarBeta beta;
/// stride between two element that will be sumed
long long int reduction_stride;
//
// Static helpers
//
/// Helper to resize a matrix with a given size and layout
template <typename T>
static void resize(cutlass::HostMatrix<T>& tensor,
int rows,
int columns,
cublasOperation_t layout,
int ldm = 0,
bool device_backed = true) {
tensor.resize(cutlass::make_Coord(rows, columns), convert(layout), ldm, device_backed);
}
template <typename T>
static void resize(cutlass::HostTensor<T, 3>& tensor,
int rows,
int columns,
int batches,
cublasOperation_t layout,
int ldm,
long long int batch_stride,
bool device_backed = true) {
assert(CUBLAS_OP_N == layout);
//tensor.resize(cutlass::make_Coord(rows, columns), convert(layout), ldm, device_backed);
tensor.reset(cutlass::make_Coord(static_cast<int>(batch_stride), ldm, 1), /*stride, slowest moving dim on the left*/
cutlass::make_Coord(batches, columns, rows), /*size, slowest moving dim on the left*/
device_backed);
}
//
// Methods
//
/// Ctor.
BatchedReductionTestbed(int m_,
int n_,
int lda_,
int ldc_,
int ldd_,
typename cutlass::TypeTraits<ScalarAlpha>::host_type alpha_ =
typename cutlass::TypeTraits<ScalarAlpha>::host_type(2),
typename cutlass::TypeTraits<ScalarAlpha>::host_type beta_ =
typename cutlass::TypeTraits<ScalarAlpha>::host_type(3))
: m(m_),
n(n_),
lda(lda_),
ldc(ldc_),
ldd(ldd_),
alpha(alpha_),
beta(beta_),
reduction_stride(ldc_ * n_) {
/// column major, batch along rows
resize(A, m_, n_, ReductionSize, CUBLAS_OP_N, lda_, reduction_stride, true);
resize(C, m_, n_, CUBLAS_OP_N, ldc_, true);
resize(D, m_, n_, CUBLAS_OP_N, ldd_, true);
resize(ref_A, m_, n_, ReductionSize, CUBLAS_OP_N, lda_, reduction_stride, false);
resize(ref_C, m_, n_, CUBLAS_OP_N, ldc_, false);
resize(ref_D, m_, n_, CUBLAS_OP_N, ldd_, false);
}
/// Dtor
~BatchedReductionTestbed() { }
/// Getters
/// Returns a pointer to the C operand
typename HostTensorA::DeviceType* ptr_A() const { return A.device_data(); }
/// Returns a pointer to the C operand
typename HostMatrixC::DeviceType* ptr_C() const { return C.device_data(); }
/// Returns a pointer to the D operand
typename HostMatrixD::DeviceType* ptr_D() const { return D.device_data(); }
///
int M() const { return m; }
///
int N() const { return n; }
///
int get_lda() const { return lda; }
///
int get_ldc() const { return ldc; }
///
int get_ldd() const { return ldd; }
///
ScalarAlpha get_alpha() const { return alpha; }
///
ScalarBeta get_beta() const { return beta; }
///
long long int get_reduction_stride() const { return reduction_stride; }
/// Initializes data, randomly
void initialize(int seed = -1) {
A.fill_random(RandomGenerator<AType>(seed + 7));
//A.fill(3);
C.fill_random(RandomGenerator<CType>(seed));
//C.fill(1);
D.fill_random(RandomGenerator<DType>(seed + 11));
//D.fill(2);
}
/// compute_host
void compute_host() {
ref_A.fill(A);
ref_C.fill(C);
ref_D.fill(D);
/// D = alpha * reduction(A) + beta * C
for (int m_idx = 0; m_idx < m; m_idx++) {
for (int n_idx = 0; n_idx < n; n_idx++) {
ScalarAccum accum = static_cast<ScalarAccum>(0.0);
for (int r_idx = 0; r_idx < static_cast<int>(ReductionSize); r_idx++) {
accum += static_cast<ScalarAccum>(ref_A.at(cutlass::make_Coord(r_idx, n_idx, m_idx)));
}
ref_D.at(cutlass::make_Coord(m_idx, n_idx)) = static_cast<DType>(
alpha * static_cast<ScalarAlpha>(accum) +
beta * static_cast<ScalarBeta>(ref_C.at(cutlass::make_Coord(m_idx, n_idx)))
);
}
}
}
/// Verifies the contents of C equal the host-side reference
bool verify_with_host() {
compute_host();
D.sync_host();
bool passed = D.bit_equals(ref_D);
return passed;
}
};
} //namespace test

View File

@@ -0,0 +1,161 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#include "cutlass_unit_test.h"
#include "cutlass/shape.h"
#include "tools/util/host_tensor.h"
#include "cutlass/reduction/batched_reduction.h"
#include "cutlass/reduction/batched_reduction_traits.h"
#include "tools/test/unit/reduction/test_batched_reduction.h"
#include "tools/test/unit/reduction/batched_reduction_testbed.h"
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Float_batched_reduction_half_alphabeta_float, batched_reduction_128x256x16) {
/*
The output matrix is 128x256
The input matrix is 128x256x16
The reduction will be applied at the third dim of input matrix
A is float, Accumulation is float
alpha and beta are float
C and D are half
*/
const int m = 128;
const int n = 256;
const int lda = 128;
const int ldc = 128;
const int ldd = 128;
const int reduction_size = 16;
typedef cutlass::reduction::BatchedReductionTraits<float, /*A*/
half, /*C*/
half, /*D*/
float, /*alpha and beta*/
float, /*accumulation type*/
reduction_size,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits_16;
test_batched_reduction<BatchedReductionTraits_16>(m, n, lda, ldc, ldd);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Float_batched_reduction_half_alphabeta_half, batched_reduction_128x256x16) {
/*
The output matrix is 128x256
The input matrix is 128x256x16
The reduction will be applied at the third dim of input matrix
A is float, Accumulation is float
alpha and beta are float
C and D are half
*/
const int m = 128;
const int n = 256;
const int lda = 128;
const int ldc = 128;
const int ldd = 128;
const int reduction_size = 16;
typedef cutlass::reduction::BatchedReductionTraits<float, /*A*/
half, /*C*/
half, /*D*/
half, /*alpha and beta*/
float, /*accumulation type*/
reduction_size,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits_16;
test_batched_reduction<BatchedReductionTraits_16>(m, n, lda, ldc, ldd);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Float_batched_reduction_half_alphabeta_float, batched_reduction_128x64x80) {
/*
The output matrix is 128x64
The input matrix is 128x64x80
The reduction will be applied at the third dim of input matrix
*/
const int m = 128;
const int n = 64;
const int lda = 128;
const int ldc = 128;
const int ldd = 128;
const int reduction_size = 80;
typedef cutlass::reduction::BatchedReductionTraits<float, /*A*/
half, /*C*/
half, /*D*/
float, /*alpha and beta*/
float, /*accumulation type*/
reduction_size,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits_80;
test_batched_reduction<BatchedReductionTraits_80>(m, n, lda, ldc, ldd);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
TEST(Float_batched_reduction_half_alphabeta_half, batched_reduction_128x64x80) {
/*
The output matrix is 128x64
The input matrix is 128x64x80
The reduction will be applied at the third dim of input matrix
*/
const int m = 128;
const int n = 64;
const int lda = 128;
const int ldc = 128;
const int ldd = 128;
const int reduction_size = 80;
typedef cutlass::reduction::BatchedReductionTraits<float, /*A*/
half, /*C*/
half, /*D*/
half, /*alpha and beta*/
float, /*accumulation type*/
reduction_size,
cutlass::Shape<1, 1, 128>,
cutlass::Shape<1, 1, 64>,
cutlass::Shape<1, 1, 2> >
BatchedReductionTraits_80;
test_batched_reduction<BatchedReductionTraits_80>(m, n, lda, ldc, ldd);
}

View File

@@ -0,0 +1,73 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (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 Test environment for batched reduction
*/
#pragma once
#include "cutlass/cutlass.h"
#include "tools/test/unit/reduction/batched_reduction_testbed.h"
template <typename BatchedRecutionTraits_>
static void test_batched_reduction(
int m,
int n,
int lda,
int ldc,
int ldd) {
typedef BatchedRecutionTraits_ Traits;
typedef cutlass::reduction::BatchedReduction<Traits> batched_reduction;
typename batched_reduction::Params params;
test::BatchedReductionTestbed<typename cutlass::TypeTraits<typename Traits::ScalarA>::host_type,
typename cutlass::TypeTraits<typename Traits::ScalarC>::host_type,
typename cutlass::TypeTraits<typename Traits::ScalarD>::host_type,
typename cutlass::TypeTraits<typename Traits::ScalarAlphaBeta>::host_type,
typename cutlass::TypeTraits<typename Traits::ScalarAlphaBeta>::host_type,
typename cutlass::TypeTraits<typename Traits::ScalarAccum>::host_type,
Traits::ReductionSize>
testbed(m, n, lda, ldc, ldd);
testbed.initialize();
params.initialize(testbed.M(),
testbed.N(),
testbed.get_alpha(),
testbed.get_beta(),
testbed.get_reduction_stride(),
testbed.ptr_A(),
testbed.get_lda(),
testbed.ptr_C(),
testbed.get_ldc(),
testbed.ptr_D(),
testbed.get_ldd());
batched_reduction::launch(params);
cudaError_t result = cudaDeviceSynchronize();
ASSERT_EQ(result, cudaSuccess) << "\nCUDA kernel launch error: " << cudaGetErrorString(result)
<< "\n";
ASSERT_TRUE(testbed.verify_with_host());
}

View File

@@ -0,0 +1,125 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
#include "cutlass_unit_test.h"
#include "cutlass/shape.h"
#include "cutlass/tile_iterator.h"
#include "gtest/gtest.h"
using ::cutlass::Coord;
using ::cutlass::Fragment;
using ::cutlass::IteratorAdvance;
using ::cutlass::make_Coord;
using ::cutlass::MemorySpace;
using ::cutlass::Shape;
using ::cutlass::TileLoadIterator;
using ::cutlass::TileTraits;
using ::testing::Test;
// TODO: Move the following to standard test helper infrastructure
// Returns randomly initialized array
//
// Caller is responsible for deallocation.
float* malloc_randomly_initialized_array(int elements) {
float* matrix = (float*)calloc(sizeof(float), elements);
for (int i = 0; i < elements; i++) {
matrix[i] = float((rand() - RAND_MAX/2) % 10);
}
return matrix;
}
#define kWarpSize 32
#define kCtaWarpCnt 6
#define kDimXPerWarp 16
#define kDimYPerWarp 2
#define kWarpTileWidth kDimXPerWarp
#define kDimYPerThread (kWarpSize / kDimYPerWarp)
#define kDimX 2400
#define kDimY 800
struct TileThreadOffset {
public:
TileThreadOffset() : xidx(0), yidx(0) {}
TileThreadOffset(int x, int y) : xidx(x), yidx(y) {}
__host__ __device__ Coord<4> operator()() const {
int column = (yidx / kDimYPerWarp) * kDimXPerWarp +
(yidx & (kDimYPerWarp - 1)) * kDimYPerThread;
return make_Coord(0, column, xidx, 0);
}
private:
int xidx, yidx;
};
TEST(TileIteratorTest, BasicCpuSideIterateTile) {
// Basic test demonstrating CPU-side tile iteration mimicking a 16x16 tile load/warp with 6 warp
// CTAs iterating over the Y.
float* matrix = malloc_randomly_initialized_array(kDimX*kDimY);
typedef Shape</*kD=*/1, /*kH=*/kCtaWarpCnt * kDimXPerWarp, /*kW=*/kDimXPerWarp> TileShape;
typedef TileLoadIterator<
TileTraits<TileShape,
/* Delta = */ Shape</*kD=*/1, /*kH=*/1, /*kW=*/1>,
/* Iter = */ Shape</*kD=*/1, /*kH=*/kDimYPerThread, /*kW=*/1>,
TileThreadOffset, /*AccessSize=*/1>,
float, IteratorAdvance::kH, MemorySpace::kGlobal> GlobalTileLoader;
typedef GlobalTileLoader::Fragment BufferType;
//
// TODO: The following loop should probably be refactored out into standard test helper code for
// tile iteration.
//
// Iterate: gridDim(1, 1, kDimX / kDimXPerWarp), blockDim(1, kDimXPerWarp, kDimYPerWarp)
for (int blockIdx_x = 0; blockIdx_x < kDimX / kDimXPerWarp; blockIdx_x++) {
for (int threadIdx_x = 0; threadIdx_x < kDimXPerWarp; threadIdx_x++) {
for (int threadIdx_y = 0; threadIdx_y < kCtaWarpCnt * kDimYPerWarp; threadIdx_y++) {
GlobalTileLoader loader(
GlobalTileLoader::Params(matrix,
/* stride_d=*/1, /*stride_h=*/kDimX, /*stride_w=*/1),
make_Coord(/*d=*/0, /*h=*/0, /*w=*/blockIdx_x * kDimXPerWarp),
TileThreadOffset(threadIdx_x, threadIdx_y));
BufferType b;
for (int yidx = 0; (yidx + threadIdx_y * kWarpTileWidth) < kDimY;
yidx += kCtaWarpCnt*kWarpTileWidth) {
loader.load_post_increment(b);
for (int i = 0; i < BufferType::kElements; i++) {
int matrix_idx = blockIdx_x * kDimXPerWarp + threadIdx_x + // row offset
kDimX * ((threadIdx_y & (kDimYPerWarp - 1)) * kDimYPerThread +
(threadIdx_y / kDimYPerWarp) * kWarpTileWidth + i + yidx);
ASSERT_EQ(b[i], matrix[matrix_idx])
<< "blockIdx.x = " << blockIdx_x << " threadIdx.x = " << threadIdx_x
<< " threadIdx.y = " << threadIdx_y << " yidx = " << yidx
<< " tile_idx = " << i << " matrix_idx = " << matrix_idx;
}
}
}
}
}
free(matrix);
}

View File

@@ -0,0 +1,127 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (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 Reference implementation for GEMM in host-side code.
*/
#pragma once
#include "cutlass/cutlass.h"
#include "cutlass/vector.h"
namespace cutlass {
namespace reference {
namespace detail {
////////////////////////////////////////////////////////////////////////////////////////////////////
/// Template function to compute an inner product.
#pragma hd_warning_disable // Suppresses warnings when attempting to instantiate with a
// host-only type
template <typename Atype, typename Btype, typename Ctype>
CUTLASS_HOST_DEVICE
Ctype inner_product(Atype a, Btype b, Ctype c) {
return Ctype(a) * Ctype(b) + c;
}
/// Specialization for matrix multiplication with binary operands
template <>
CUTLASS_HOST_DEVICE
int inner_product<Vector<bin1_t, 32>, Vector<bin1_t, 32>, int>(
Vector<bin1_t, 32> a,
Vector<bin1_t, 32> b,
int c) {
int accum = 0;
for (int bit = 0; bit < 32; bit++) {
accum += a[bit] ^ b[bit];
}
return accum + c;
}
/// Specialization for matrix multiplication with signed 4-bit integer operands
template <>
CUTLASS_HOST_DEVICE
int inner_product<Vector<int4_t, 8>, Vector<int4_t, 8>, int>(
Vector<int4_t, 8> a,
Vector<int4_t, 8> b,
int c) {
int accum = 0;
for (int k = 0; k < 8; k++) {
accum += a[k] * b[k];
}
return accum + c;
}
/// Specialization for matrix multiplication with unsigned 4-bit integer operands
template <>
CUTLASS_HOST_DEVICE
int inner_product<Vector<uint4_t, 8>, Vector<uint4_t, 8>, int>(
Vector<uint4_t, 8> a,
Vector<uint4_t, 8> b,
int c) {
int accum = 0;
for (int k = 0; k < 8; k++) {
accum += a[k] * b[k];
}
return accum + c;
}
////////////////////////////////////////////////////////////////////////////////////////////////////
template <typename SrcType, typename DstType>
struct Cast {
// Default behavior: convert to the destination type
#pragma hd_warning_disable // Suppresses warnings when attempting to instantiate complex<T> with a
// host-only type
CUTLASS_HOST_DEVICE
static DstType apply(SrcType src) { return static_cast<DstType>(src); };
};
template <>
struct Cast<float, int8_t> {
CUTLASS_HOST_DEVICE
static int8_t apply(float src) {
// Clamp to the range of signed 8-bit integers.
return static_cast<int8_t>(fmaxf(-128.f, fminf(127.f, src)));
};
};
template <>
struct Cast<float, uint8_t> {
CUTLASS_HOST_DEVICE
static uint8_t apply(float src) {
// Clamp to the range of signed 8-bit integers.
return static_cast<uint8_t>(fmaxf(0.f, fminf(255.f, src)));
};
};
////////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace detail
} // namespace reference
} // namespace cutlass

View File

@@ -0,0 +1,224 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (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 Reference implementation for GEMM in device-side code.
*/
#pragma once
#include "cutlass/coord.h"
#include "cutlass/matrix_traits.h"
#include "cutlass/tensor_view.h"
#include "cutlass/gemm/gemm_coord.h"
#include "tools/util/reference/device/kernel/gemm.h"
namespace cutlass {
namespace reference {
namespace device {
////////////////////////////////////////////////////////////////////////////////////////////////////
/// Computes a general matrix product among matrices (tensors of rank=2) pointed to by TensorRef
/// objects.
///
/// Explicitly naming types needed by this template can be cumbersome, particularly for the
/// accumulator type, so a function argument 'initial_accum' is exposed. Passing
/// AccumulatorType(0) as the last function argument can be easier than naming all template
/// arguments explicitly.
template <
typename TensorRefA,
typename TensorRefB,
typename TensorRefC,
typename ScalarType,
typename AccumulatorType
>
void Gemm(
gemm::GemmCoord problem_size,
ScalarType alpha,
TensorRefA tensor_a,
TensorRefB tensor_b,
ScalarType beta,
TensorRefC tensor_c,
AccumulatorType initial_accum) {
typedef typename TensorRefA::Storage AType;
typedef typename TensorRefB::Storage BType;
typedef typename TensorRefC::Storage CType;
static_assert(
TensorRefA::kRank == 2 &&
TensorRefB::kRank == 2 &&
TensorRefC::kRank == 2, "Tensors must be of rank 2");
// Blocking structure potentially improves performance of reference implementation
// with a minor increase in complexity.
//
// Note, this reference implementation is NOT expected to approach peak performance.
typedef Shape<1, 4, 4> OutputTile;
dim3 block(16, 8);
dim3 grid(
(problem_size.m() + block.x * OutputTile::kW - 1) / (block.x * OutputTile::kW),
(problem_size.n() + block.y * OutputTile::kH - 1) / (block.y * OutputTile::kH)
);
// Launch a GEMM kernel
kernel::Gemm<
TensorRefA,
TensorRefB,
TensorRefC,
ScalarType,
AccumulatorType,
OutputTile
><<< grid, block >>>(
problem_size,
alpha,
tensor_a,
tensor_b,
beta,
tensor_c,
initial_accum
);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
/// Computes a general matrix product among matrices (tensors of rank=2) pointed to by TensorRef
/// objects.
///
/// This assumes the accumulator type is the same type as the scalars.
template <
typename TensorRefA,
typename TensorRefB,
typename TensorRefC,
typename ScalarType
>
void Gemm(
gemm::GemmCoord problem_size,
ScalarType alpha,
TensorRefA tensor_a,
TensorRefB tensor_b,
ScalarType beta,
TensorRefC tensor_c) {
Gemm(problem_size, alpha, tensor_a, tensor_b, beta, tensor_c, ScalarType(0));
}
////////////////////////////////////////////////////////////////////////////////////////////////////
//
// Batched GEMM
//
////////////////////////////////////////////////////////////////////////////////////////////////////
/// Computes a batch of GEMMs over a set of matrices of common dimension.
//
// TensorRefCollection* is a type satisfying the TensorRefCollection concept.
//
template <
typename TensorRefCollectionA,
typename TensorRefCollectionB,
typename TensorRefCollectionC,
typename ScalarType,
typename AccumulatorType
>
void BatchedGemm(
gemm::GemmCoord problem_size,
ScalarType alpha,
TensorRefCollectionA tensor_a,
TensorRefCollectionB tensor_b,
ScalarType beta,
TensorRefCollectionC tensor_c,
AccumulatorType initial_accum) {
typedef typename TensorRefCollectionA::Storage AType;
typedef typename TensorRefCollectionB::Storage BType;
typedef typename TensorRefCollectionC::Storage CType;
static_assert(
TensorRefCollectionA::kRank == 2 &&
TensorRefCollectionB::kRank == 2 &&
TensorRefCollectionC::kRank == 2, "Tensors must be of rank 2");
// Blocking structure potentially improves performance of reference implementation
// with a minor increase in complexity.
//
// Note, this reference implementation is NOT expected to approach peak performance.
typedef Shape<1, 4, 4> OutputTile;
dim3 block(16, 8);
dim3 grid(
(problem_size.m() + block.x * OutputTile::kW - 1) / (block.x * OutputTile::kW),
(problem_size.n() + block.y * OutputTile::kH - 1) / (block.y * OutputTile::kH),
problem_size.batch()
);
// Launch a GEMM kernel
kernel::BatchedGemm<
TensorRefCollectionA,
TensorRefCollectionB,
TensorRefCollectionC,
ScalarType,
AccumulatorType,
OutputTile
><<< grid, block >>>(
problem_size,
alpha,
tensor_a,
tensor_b,
beta,
tensor_c,
initial_accum
);
}
/// Computes a general matrix product among matrices (tensors of rank=2) pointed to by TensorRef
/// objects.
//
// TensorRefCollection* is a type satisfying the TensorRefCollection concept.
//
template <
typename TensorRefCollectionA,
typename TensorRefCollectionB,
typename TensorRefCollectionC,
typename ScalarType,
typename AccumulatorType
>
void BatchedGemm(
gemm::GemmCoord problem_size,
ScalarType alpha,
TensorRefCollectionA tensor_a,
TensorRefCollectionB tensor_b,
ScalarType beta,
TensorRefCollectionC tensor_c) {
BatchedGemm(problem_size, alpha, tensor_a, tensor_b, beta, tensor_c, ScalarType(0));
}
////////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace host
} // namespace reference
} // namespace cutlass

View File

@@ -0,0 +1,148 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (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 Reference implementation for GEMM in host-side code.
*/
#pragma once
#include "cutlass/coord.h"
#include "cutlass/matrix_traits.h"
#include "cutlass/tensor_view.h"
#include "cutlass/gemm/gemm_coord.h"
#include "tools/util/reference/device/thread/gemm.h"
namespace cutlass {
namespace reference {
namespace device {
namespace kernel {
////////////////////////////////////////////////////////////////////////////////////////////////////
/// Computes a general matrix product among matrices (tensors of rank=2) pointed to by TensorRef
/// objects.
template <
typename TensorRefA,
typename TensorRefB,
typename TensorRefC,
typename ScalarType,
typename AccumulatorType,
typename OutputTile
>
__global__ void Gemm(
gemm::GemmCoord problem_size,
ScalarType alpha,
TensorRefA tensor_a,
TensorRefB tensor_b,
ScalarType beta,
TensorRefC tensor_c,
AccumulatorType initial_accum) {
// Map each thread to a unique tile of the output matrix
MatrixCoord output_coord(
(threadIdx.x + blockIdx.x * blockDim.x) * OutputTile::kW,
(threadIdx.y + blockIdx.y * blockDim.y) * OutputTile::kH
);
// Compute the general matrix product
thread::Gemm<
TensorRefA,
TensorRefB,
TensorRefC,
ScalarType,
AccumulatorType,
OutputTile
> gemm(initial_accum);
gemm.multiply_add(
problem_size,
tensor_a,
tensor_b,
output_coord);
gemm.epilogue(problem_size, alpha, beta, tensor_c, output_coord);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
/// Computes a general matrix product among matrices (tensors of rank=2) pointed to by TensorRef
/// objects.
template <
typename TensorRefCollectionA,
typename TensorRefCollectionB,
typename TensorRefCollectionC,
typename ScalarType,
typename AccumulatorType,
typename OutputTile
>
__global__ void BatchedGemm(
gemm::GemmCoord problem_size,
ScalarType alpha,
TensorRefCollectionA tensor_collection_a,
TensorRefCollectionB tensor_collection_b,
ScalarType beta,
TensorRefCollectionC tensor_collection_c,
AccumulatorType initial_accum) {
// Obtain batch ID
int batch_id = blockIdx.z;
// Dereference based on batch_id
typename TensorRefCollectionA::TensorRef tensor_a = tensor_collection_a.at(batch_id);
typename TensorRefCollectionB::TensorRef tensor_b = tensor_collection_b.at(batch_id);
typename TensorRefCollectionC::TensorRef tensor_c = tensor_collection_c.at(batch_id);
// Map each thread to a unique tile of the output matrix
MatrixCoord output_coord(
(threadIdx.x + blockIdx.x * blockDim.x) * OutputTile::kW,
(threadIdx.y + blockIdx.y * blockDim.y) * OutputTile::kH
);
// Compute the general matrix product
thread::Gemm<
typename TensorRefCollectionA::TensorRef,
typename TensorRefCollectionB::TensorRef,
typename TensorRefCollectionC::TensorRef,
ScalarType,
AccumulatorType,
OutputTile
> gemm(initial_accum);
gemm.multiply_add(
problem_size,
tensor_a,
tensor_b,
output_coord);
gemm.epilogue(problem_size, alpha, beta, tensor_c, output_coord);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace kernel
} // namespace device
} // namespace reference
} // namespace cutlass

View File

@@ -0,0 +1,95 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (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 Reference implementation for GEMM in host-side code.
*/
#pragma once
#include "cutlass/coord.h"
#include "cutlass/matrix_traits.h"
#include "cutlass/tensor_view.h"
#include "cutlass/gemm/gemm_coord.h"
#include "cutlass/util/complex.h"
#include "tools/util/reference/device/thread/split_complex_gemm.h"
namespace cutlass {
namespace reference {
namespace device {
namespace kernel {
////////////////////////////////////////////////////////////////////////////////////////////////////
/// Computes a general matrix product among matrices (tensors of rank=2) pointed to by TensorRef
/// objects.
template <
typename TensorRefA, /// concept: ZipTensorRef
typename TensorRefB, /// concept: ZipTensorRef
typename TensorRefC, /// concept: ZipTensorRef
typename ScalarType, /// real-valued type underlying complex scalars
typename AccumulatorType, /// real-valued type underlying complex accumulators
typename OutputTile /// concept: Shape
>
__global__ void SplitComplexGemm(
gemm::GemmCoord problem_size,
platform::complex<ScalarType> alpha,
TensorRefA tensor_a,
TensorRefB tensor_b,
platform::complex<ScalarType> beta,
TensorRefC tensor_c,
platform::complex<AccumulatorType> initial_accum) {
// Map each thread to a unique tile of the output matrix
MatrixCoord output_coord(
(threadIdx.x + blockIdx.x * blockDim.x) * OutputTile::kW,
(threadIdx.y + blockIdx.y * blockDim.y) * OutputTile::kH
);
// Compute the general matrix product
thread::Gemm<
TensorRefA,
TensorRefB,
TensorRefC,
ScalarType,
AccumulatorType,
OutputTile
> gemm(initial_accum);
gemm.multiply_add(
problem_size,
tensor_a,
tensor_b,
output_coord);
gemm.epilogue(problem_size, alpha, beta, tensor_c, output_coord);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace kernel
} // namespace device
} // namespace reference
} // namespace cutlass

View File

@@ -0,0 +1,103 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (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 Reference implementation for GEMM in device-side code.
*/
#pragma once
#include "cutlass/coord.h"
#include "cutlass/matrix_traits.h"
#include "cutlass/tensor_view.h"
#include "cutlass/gemm/gemm_coord.h"
#include "cutlass/util/complex.h"
#include "tools/util/reference/device/kernel/gemm.h"
namespace cutlass {
namespace reference {
namespace device {
////////////////////////////////////////////////////////////////////////////////////////////////////
/// Computes a complex-valued GEMM whose operands are in the split-complex format.
template <
typename TensorRefA, /// concept: ZipTensorRef
typename TensorRefB, /// concept: ZipTensorRef
typename TensorRefC, /// concept: ZipTensorRef
typename ScalarType, /// real-valued type underlying complex scalars
typename AccumulatorType /// real-valued type underlying complex accumulators
>
void SplitComplexGemm(
gemm::GemmCoord problem_size,
platform::complex<ScalarType> alpha,
TensorRefA tensor_a,
TensorRefB tensor_b,
platform::complex<ScalarType> beta,
TensorRefC tensor_c,
platform::complex<ScalarType> initial_accum) {
static_assert(
TensorRefA::First::kRank == 2 && TensorRefA::Second::kRank == 2 &&
TensorRefB::First::kRank == 2 && TensorRefB::Second::kRank == 2 &&
TensorRefC::First::kRank == 2 && TensorRefC::Second::kRank == 2,
"Tensors must be of rank 2");
// Blocking structure potentially improves performance of reference implementation
// with a minor increase in complexity.
//
// Note, this reference implementation is NOT expected to approach peak performance.
typedef Shape<1, 4, 4> OutputTile;
dim3 block(16, 8);
dim3 grid(
(problem_size.m() + block.x * OutputTile::kW - 1) / (block.x * OutputTile::kW),
(problem_size.n() + block.y * OutputTile::kH - 1) / (block.y * OutputTile::kH)
);
// Launch a GEMM kernel
kernel::SplitComplexGemm<
TensorRefA,
TensorRefB,
TensorRefC,
ScalarType,
AccumulatorType,
OutputTile
><<< grid, block >>>(
problem_size,
alpha,
tensor_a,
tensor_b,
beta,
tensor_c,
initial_accum
);
}
////////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace device
} // namespace reference
} // namespace cutlass

View File

@@ -0,0 +1,176 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (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 Reference implementation for GEMM in host-side code.
*/
#pragma once
#include "cutlass/coord.h"
#include "cutlass/matrix_traits.h"
#include "cutlass/tensor_view.h"
#include "cutlass/gemm/gemm_coord.h"
#include "tools/util/reference/detail/inner_product.h"
namespace cutlass {
namespace reference {
namespace device {
namespace thread {
////////////////////////////////////////////////////////////////////////////////////////////////////
/// Thread-level blocked general matrix product.
//
// Note, this is a reference implementation. Performance is not expected to approach peak.
//
template <
typename TensorRefA,
typename TensorRefB,
typename TensorRefC,
typename ScalarType,
typename AccumulatorType,
typename OutputTile
>
struct Gemm {
typedef typename TensorRefA::Storage ScalarA;
typedef typename TensorRefB::Storage ScalarB;
typedef typename TensorRefC::Storage ScalarC;
//
// Data members
//
/// Tile for A operand
ScalarA A_tile[OutputTile::kW];
/// Tile for B operand
ScalarB B_tile[OutputTile::kH];
/// Tile for Accumulator
AccumulatorType accum[OutputTile::kH][OutputTile::kW];
//
// Methods
//
/// Constructor
CUTLASS_HOST_DEVICE
Gemm(AccumulatorType initial_accum = AccumulatorType(0)) {
// Clear fetch registers
for (int i = 0; i < OutputTile::kW; ++i) {
A_tile[i] = ScalarA(0);
}
for (int j = 0; j < OutputTile::kW; ++j) {
B_tile[j] = ScalarB(0);
}
// Clear accumulators
CUTLASS_PRAGMA_UNROLL
for (int j = 0; j < OutputTile::kH; ++j) {
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < OutputTile::kW; ++i) {
accum[j][i] = initial_accum;
}
}
}
/// Computes a matrix product
CUTLASS_HOST_DEVICE
Gemm & multiply_add(
gemm::GemmCoord problem_size,
TensorRefA tensor_a,
TensorRefB tensor_b,
MatrixCoord output_coord = MatrixCoord()) {
// Loop over the GEMM K dimension
CUTLASS_PRAGMA_NO_UNROLL
for (int k = 0; k < problem_size.k(); ++k) {
// Fetch a slice of the A matrix
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < OutputTile::kW; ++i) {
if (output_coord.row() + i < problem_size.m()) {
A_tile[i] = tensor_a.at(make_Coord(output_coord.row() + i, k));
}
}
// Fetch a slice of the B matrix
CUTLASS_PRAGMA_UNROLL
for (int j = 0; j < OutputTile::kH; ++j) {
if (output_coord.column() + j < problem_size.n()) {
B_tile[j] = tensor_b.at(make_Coord(k, output_coord.column() + j));
}
}
// Compute an accumulated matrix product
CUTLASS_PRAGMA_UNROLL
for (int j = 0; j < OutputTile::kH; ++j) {
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < OutputTile::kW; ++i) {
accum[j][i] = detail::inner_product(A_tile[i], B_tile[j], accum[j][i]);
}
}
}
return *this;
}
/// Performs linear scaling of matrix product and updates output tensor
CUTLASS_HOST_DEVICE
Gemm & epilogue(
gemm::GemmCoord problem_size,
ScalarType alpha,
ScalarType beta,
TensorRefC tensor_c,
MatrixCoord output_coord = MatrixCoord()) {
// Update the output tensor
for (int j = 0; j < OutputTile::kH; ++j) {
for (int i = 0; i < OutputTile::kW; ++i) {
MatrixCoord coord = output_coord + MatrixCoord(i, j);
if (coord.row() < problem_size.m() && coord.column() < problem_size.n()) {
tensor_c.at(coord) = detail::Cast<ScalarType, ScalarC>::apply(
alpha * ScalarType(accum[j][i]) +
beta * ScalarType(tensor_c.at(coord))
);
}
}
}
return *this;
}
};
////////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace thread
} // namespace device
} // namespace reference
} // namespace cutlass

View File

@@ -0,0 +1,192 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (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 Reference implementation for GEMM in host-side code.
*/
#pragma once
#include "cutlass/coord.h"
#include "cutlass/matrix_traits.h"
#include "cutlass/tensor_view.h"
#include "cutlass/gemm/gemm_coord.h"
#include "tools/util/reference/detail/inner_product.h"
namespace cutlass {
namespace reference {
namespace device {
namespace thread {
////////////////////////////////////////////////////////////////////////////////////////////////////
/// Thread-level blocked general matrix product.
//
// Note, this is a reference implementation. Performance is not expected to approach peak.
//
template <
typename TensorRefA, /// concept: ZipTensorRef
typename TensorRefB, /// concept: ZipTensorRef
typename TensorRefC, /// concept: ZipTensorRef
typename ScalarType, /// real-valued type underlying complex scalars
typename AccumulatorType, /// real-valued type underlying complex accumulators
typename OutputTile /// concept: Shape
>
struct SplitComplexGemm {
typedef typename TensorRefA::First::Storage RealScalarA;
typedef typename TensorRefB::First::Storage RealScalarB;
typedef typename TensorRefC::First::Storage RealScalarC;
typedef platform::complex<RealScalarA> ScalarA;
typedef platform::complex<RealScalarB> ScalarB;
typedef platform::complex<AccumulatorType> ComplexAccumulator;
typedef platform::complex<ScalarType> ComplexScalar;
//
// Data members
//
/// Tile for A operand
ScalarA A_tile[OutputTile::kW];
/// Tile for B operand
ScalarB B_tile[OutputTile::kH];
/// Tile for Accumulator
ComplexAccumulator accum[OutputTile::kH][OutputTile::kW];
//
// Methods
//
/// Constructor
CUTLASS_HOST_DEVICE
Gemm(ComplexAccumulator initial_accum = AccumulatorType(0)) {
// Clear fetch registers
for (int i = 0; i < OutputTile::kW; ++i) {
A_tile[i] = ScalarA(0);
}
for (int j = 0; j < OutputTile::kW; ++j) {
B_tile[j] = ScalarB(0);
}
// Clear accumulators
CUTLASS_PRAGMA_UNROLL
for (int j = 0; j < OutputTile::kH; ++j) {
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < OutputTile::kW; ++i) {
accum[j][i] = initial_accum;
}
}
}
/// Computes a matrix product
CUTLASS_HOST_DEVICE
Gemm & multiply_add(
gemm::GemmCoord problem_size,
TensorRefA tensor_a,
TensorRefB tensor_b,
MatrixCoord output_coord = MatrixCoord()) {
// Loop over the GEMM K dimension
CUTLASS_PRAGMA_NO_UNROLL
for (int k = 0; k < problem_size.k(); ++k) {
// Fetch a slice of the A matrix - zip into complex values
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < OutputTile::kW; ++i) {
if (output_coord.row() + i < problem_size.m()) {
MatrixCoord coord(output_coord.row() + i, k);
A_tile[i].real() = tensor_a.first.at(coord);
A_tile[i].imag() = tensor_a.second.at(coord);
}
}
// Fetch a slice of the B matrix - zip into complex values
CUTLASS_PRAGMA_UNROLL
for (int j = 0; j < OutputTile::kH; ++j) {
if (output_coord.column() + j < problem_size.n()) {
MatrixCoord coord(k, output_coord.column() + j);
B_tile[j].real() = tensor_b.first.at(coord);
B_tile[j].imag() = tensor_b.second.at(coord);
}
}
// Compute an accumulated matrix product on complex values
CUTLASS_PRAGMA_UNROLL
for (int j = 0; j < OutputTile::kH; ++j) {
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < OutputTile::kW; ++i) {
accum[j][i] = detail::inner_product(A_tile[i], B_tile[j], accum[j][i]);
}
}
}
return *this;
}
/// Performs linear scaling of matrix product and updates output tensor
CUTLASS_HOST_DEVICE
Gemm & epilogue(
gemm::GemmCoord problem_size,
ComplexScalar alpha,
ComplexScalar beta,
TensorRefC tensor_c,
MatrixCoord output_coord = MatrixCoord()) {
// Update the output tensor
for (int j = 0; j < OutputTile::kH; ++j) {
for (int i = 0; i < OutputTile::kW; ++i) {
MatrixCoord coord = output_coord + MatrixCoord(i, j);
if (coord < problem_size.mn()) {
ComplexScalar source(
tensor_c.first.at(coord),
tensor_c.second.at(coord)
);
// Final calculation is performed in data type of scalars
ComplexScalar result = alpha * ComplexScalar(accum[j][i].real(), accum[j][i].imag()) + beta * source;
// Unzip and convert into output tensor data type
tensor_c.first.at(coord) = detail::Cast<ScalarType, RealScalarC>::apply(result.real());
tensor_c.second.at(coord) = detail::Cast<ScalarType, RealScalarC>::apply(result.imag());
}
}
}
return *this;
}
};
////////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace thread
} // namespace device
} // namespace reference
} // namespace cutlass

View File

@@ -33,90 +33,14 @@
#include "cutlass/tensor_view.h"
#include "cutlass/gemm/gemm_coord.h"
#include "tools/util/reference/detail/inner_product.h"
namespace cutlass {
namespace reference {
namespace host {
////////////////////////////////////////////////////////////////////////////////////////////////////
namespace detail {
/// Template function to compute an inner product.
template <typename Atype, typename Btype, typename Ctype>
Ctype inner_product(Atype a, Btype b, Ctype c) {
return Ctype(a) * Ctype(b) + c;
}
/// Specialization for matrix multiplication with binary operands
template <>
inline int inner_product<Vector<bin1_t, 32>, Vector<bin1_t, 32>, int>(
Vector<bin1_t, 32> a,
Vector<bin1_t, 32> b,
int c) {
int accum = 0;
for (int bit = 0; bit < 32; bit++) {
accum += a[bit] ^ b[bit];
}
return accum + c;
}
/// Specialization for matrix multiplication with signed 4-bit integer operands
template <> inline
int inner_product<Vector<int4_t, 8>, Vector<int4_t, 8>, int>(
Vector<int4_t, 8> a,
Vector<int4_t, 8> b,
int c) {
int accum = 0;
for (int k = 0; k < 8; k++) {
accum += a[k] * b[k];
}
return accum + c;
}
/// Specialization for matrix multiplication with unsigned 4-bit integer operands
template <> inline
int inner_product<Vector<uint4_t, 8>, Vector<uint4_t, 8>, int>(
Vector<uint4_t, 8> a,
Vector<uint4_t, 8> b,
int c) {
int accum = 0;
for (int k = 0; k < 8; k++) {
accum += a[k] * b[k];
}
return accum + c;
}
////////////////////////////////////////////////////////////////////////////////////////////////////
template <typename SrcType, typename DstType>
struct Cast {
// Default behavior: convert to the destination type
static inline DstType apply(SrcType src) { return static_cast<DstType>(src); };
};
template <>
struct Cast<float, int8_t> {
static inline int8_t apply(float src) {
// Clamp to the range of signed 8-bit integers.
return static_cast<int8_t>(fmaxf(-128.f, fminf(127.f, src)));
};
};
template <>
struct Cast<float, uint8_t> {
static inline uint8_t apply(float src) {
// Clamp to the range of signed 8-bit integers.
return static_cast<uint8_t>(fmaxf(0.f, fminf(255.f, src)));
};
};
} // namespace detail
////////////////////////////////////////////////////////////////////////////////////////////////////
/// Computes a general matrix product among matrices (tensors of rank=2) pointed to by TensorRef
/// objects.
///
@@ -178,7 +102,7 @@ void Gemm(
AType a = tensor_a.at(MatrixCoord(row, k_block));
BType b = tensor_b.at(MatrixCoord(k_block, col));
accum[i][j] = detail::inner_product(a, b, accum[i][j]);
accum[i][j] = cutlass::reference::detail::inner_product(a, b, accum[i][j]);
}
}
}
@@ -192,7 +116,7 @@ void Gemm(
MatrixCoord coord = MatrixCoord(row, col);
if (row < M && col < N) {
tensor_c.at(coord) = detail::Cast<ScalarType, CType>::apply(
tensor_c.at(coord) = cutlass::reference::detail::Cast<ScalarType, CType>::apply(
alpha * ScalarType(accum[i][j]) +
beta * ScalarType(tensor_c.at(coord)));
}
@@ -225,9 +149,16 @@ void Gemm(
Gemm(problem_size, alpha, tensor_a, tensor_b, beta, tensor_c, ScalarType(0));
}
////////////////////////////////////////////////////////////////////////////////////////////////////
//
// Batched GEMM
//
////////////////////////////////////////////////////////////////////////////////////////////////////
/// Computes a batch of GEMMs over a set of matrices of common dimension.
//
// TensorRefCollection* is a type satisfying the TensorRefCollection concept.
//
template <
typename TensorRefCollectionA,
typename TensorRefCollectionB,
@@ -235,14 +166,14 @@ template <
typename ScalarType,
typename AccumulatorType
>
void BatchGemm(
void BatchedGemm(
gemm::GemmCoord problem_size,
ScalarType alpha,
TensorRefCollectionA const& tensor_a,
TensorRefCollectionB const& tensor_b,
ScalarType beta,
TensorRefCollectionC &tensor_c,
AccumulatorType initial_accum = AccumulatorType(0)) {
AccumulatorType initial_accum) {
typename TensorRefCollectionA::ConstIterator tensor_a_it = tensor_a.begin();
typename TensorRefCollectionB::ConstIterator tensor_b_it = tensor_b.begin();
@@ -263,6 +194,29 @@ void BatchGemm(
}
}
/// Computes a general matrix product among matrices (tensors of rank=2) pointed to by TensorRef
/// objects.
//
// TensorRefCollection* is a type satisfying the TensorRefCollection concept.
//
template <
typename TensorRefCollectionA,
typename TensorRefCollectionB,
typename TensorRefCollectionC,
typename ScalarType,
typename AccumulatorType
>
void BatchedGemm(
gemm::GemmCoord problem_size,
ScalarType alpha,
TensorRefCollectionA const& tensor_a,
TensorRefCollectionB const& tensor_b,
ScalarType beta,
TensorRefCollectionC &tensor_c) {
BatchedGemm(problem_size, alpha, tensor_a, tensor_b, beta, tensor_c, ScalarType(0));
}
////////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace host

View File

@@ -0,0 +1,254 @@
/***************************************************************************************************
* Copyright (c) 2017-2018, NVIDIA CORPORATION. All rights reserved.
*
* Redistribution and use in source and binary forms, with or without modification, are permitted
* provided that the following conditions are met:
* * Redistributions of source code must retain the above copyright notice, this list of
* conditions and the following disclaimer.
* * 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.
* * Neither the name of the NVIDIA CORPORATION 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 NVIDIA CORPORATION 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 TOR (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 Reference implementation for split-complex GEMM in device-side code.
*/
#pragma once
#include "cutlass/coord.h"
#include "cutlass/matrix_traits.h"
#include "cutlass/tensor_view.h"
#include "cutlass/gemm/gemm_coord.h"
#include "cutlass/util/complex.h"
namespace cutlass {
namespace reference {
namespace host {
////////////////////////////////////////////////////////////////////////////////////////////////////
/// Computes a complex-valued GEMM whose operands are in the split-complex format.
template <
typename TensorRefA, /// concept: ZipTensorRef
typename TensorRefB, /// concept: ZipTensorRef
typename TensorRefC, /// concept: ZipTensorRef
typename ScalarType, /// real-valued type underlying complex scalars
typename AccumulatorType /// real-valued type underlying complex accumulators
>
void SplitComplexGemm(
gemm::GemmCoord problem_size,
platform::complex<ScalarType> alpha,
TensorRefA tensor_a,
TensorRefB tensor_b,
platform::complex<ScalarType> beta,
TensorRefC tensor_c,
platform::complex<AccumulatorType> initial_accum) {
typedef typename TensorRefA::First::Storage AType;
typedef typename TensorRefB::First::Storage BType;
typedef typename TensorRefC::First::Storage CType;
typedef platform::complex<AType> ComplexAType;
typedef platform::complex<BType> ComplexBType;
typedef platform::complex<CType> ComplexCType;
typedef platform::complex<ScalarType> ComplexScalarType;
typedef platform::complex<AccumulatorType> ComplexAccumulatorType;
static_assert(
TensorRefA::First::kRank == 2 && TensorRefA::Second::kRank == 2 &&
TensorRefB::First::kRank == 2 && TensorRefB::Second::kRank == 2 &&
TensorRefC::First::kRank == 2 && TensorRefC::Second::kRank == 2,
"Tensors must be of rank 2");
// Note: batch is ignored.
int const M = problem_size.m();
int const N = problem_size.n();
int const K = problem_size.k();
// Blocking necessary to speedup reference implementation
int const Mblock = 32;
int const Nblock = 32;
for (int row_block = 0; row_block < M; row_block += Mblock) {
for (int col_block = 0; col_block < N; col_block += Nblock) {
ComplexAccumulatorType accum[Mblock][Nblock];
for (int j = 0; j < Nblock; j++) {
for (int i = 0; i < Mblock; i++) {
accum[i][j] = initial_accum;
}
}
for (int k_block = 0; k_block < K; ++k_block) {
for (int j = 0; j < Nblock; j++) {
for (int i = 0; i < Mblock; i++) {
int row = row_block + i;
int col = col_block + j;
if (row < M && col < N) {
ComplexAType a(
tensor_a.first.at(MatrixCoord(row, k_block)),
tensor_a.second.at(MatrixCoord(row, k_block))
);
ComplexBType b(
tensor_b.first.at(MatrixCoord(k_block, col)),
tensor_b.second.at(MatrixCoord(k_block, col))
);
accum[i][j] = detail::inner_product(a, b, accum[i][j]);
}
}
}
}
for (int j = 0; j < Nblock; j++) {
for (int i = 0; i < Mblock; i++) {
int row = row_block + i;
int col = col_block + j;
MatrixCoord coord = MatrixCoord(row, col);
if (row < M && col < N) {
ComplexScalarType product(
detail::Cast<AccumulatorType, ScalarType>::apply(accum[i][j].real()),
detail::Cast<AccumulatorType, ScalarType>::apply(accum[i][j].imag())
);
ComplexScalarType source(
detail::Cast<CType, ScalarType>::apply(tensor_c.first.at(coord)),
detail::Cast<CType, ScalarType>::apply(tensor_c.second.at(coord))
);
ComplexScalarType result = alpha * product + beta * source;
tensor_c.first.at(coord) = detail::Cast<ScalarType, CType>::apply(result.real());
tensor_c.second.at(coord) = detail::Cast<ScalarType, CType>::apply(result.imag());
}
}
}
}
}
}
////////////////////////////////////////////////////////////////////////////////////////////////////
/// Computes a complex-valued GEMM whose operands are in the split-complex format.
template <
typename TensorRefA, /// concept: ZipTensorRef
typename TensorRefB, /// concept: ZipTensorRef
typename TensorRefC, /// concept: ZipTensorRef
typename ScalarType, /// real-valued type underlying complex scalars
typename AccumulatorType /// real-valued type underlying complex accumulators
>
void SplitComplexGemm(
gemm::GemmCoord problem_size,
platform::complex<ScalarType> alpha,
TensorRefA tensor_a,
TensorRefB tensor_b,
platform::complex<ScalarType> beta,
TensorRefC tensor_c) {
return SplitComplexGemm(problem_size, alpha, tensor_a, tensor_b,beta, tensor_c, ScalarType(0));
}
////////////////////////////////////////////////////////////////////////////////////////////////////
//
// Batched Split-Complex GEMM
//
////////////////////////////////////////////////////////////////////////////////////////////////////
/// Computes a complex-valued GEMM whose operands are in the split-complex format.
template <
typename TensorRefCollectionA, /// concept: Pair<TensorRefCollection, TensorRefCollection>
typename TensorRefCollectionB, /// concept: Pair<TensorRefCollection, TensorRefCollection>
typename TensorRefCollectionC, /// concept: Pair<TensorRefCollection, TensorRefCollection>
typename ScalarType, /// real-valued type underlying complex scalars
typename AccumulatorType /// real-valued type underlying complex accumulators
>
void BatchedSplitComplexGemm(
gemm::GemmCoord problem_size,
platform::complex<ScalarType> alpha,
TensorRefCollectionA tensor_a,
TensorRefCollectionB tensor_b,
platform::complex<ScalarType> beta,
TensorRefCollectionC tensor_c,
platform::complex<AccumulatorType> initial_accum) {
typename TensorRefCollectionA::ConstIterator tensor_a_real = tensor_a.first.begin();
typename TensorRefCollectionA::ConstIterator tensor_a_imag = tensor_a.second.begin();
typename TensorRefCollectionB::ConstIterator tensor_b_real = tensor_b.first.begin();
typename TensorRefCollectionB::ConstIterator tensor_b_imag = tensor_b.second.begin();
typename TensorRefCollectionC::ConstIterator tensor_c_real = tensor_c.first.begin();
typename TensorRefCollectionC::ConstIterator tensor_c_imag = tensor_c.second.begin();
for (int batch = 0; batch < problem_size.batch(); ++batch) {
SplitComplexGemm(
problem_size,
alpha,
make_ZipTensorRef(*tensor_a_real, *tensor_a_imag),
make_ZipTensorRef(*tensor_b_real, *tensor_b_imag),
beta,
make_ZipTensorRef(*tensor_c_real, *tensor_c_imag),
initial_accum);
++tensor_a_real;
++tensor_a_imag;
++tensor_b_real;
++tensor_b_imag;
++tensor_c_real;
++tensor_c_imag;
}
}
////////////////////////////////////////////////////////////////////////////////////////////////////
/// Computes a complex-valued GEMM whose operands are in the split-complex format.
template <
typename TensorRefCollectionA, /// concept: pair<TensorRefCollection, TensorRefCollection>
typename TensorRefCollectionB, /// concept: pair<TensorRefCollection, TensorRefCollection>
typename TensorRefCollectionC, /// concept: pair<TensorRefCollection, TensorRefCollection>
typename ScalarType, /// real-valued type underlying complex scalars
typename AccumulatorType /// real-valued type underlying complex accumulators
>
void BatchedSplitComplexGemm(
gemm::GemmCoord problem_size,
platform::complex<ScalarType> alpha,
TensorRefCollectionA tensor_a,
TensorRefCollectionB tensor_b,
platform::complex<ScalarType> beta,
TensorRefCollectionC tensor_c) {
BatchedSplitComplexGemm(
problem_size,
alpha,
tensor_a,
tensor_b,
beta,
tensor_c,
platform::complex<ScalarType>(0, 0));
}
////////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace host
} // namespace reference
} // namespace cutlass

View File

@@ -45,6 +45,7 @@ struct TypeTraits {
typedef T device_type;
static inline T remove_negative_zero(T x) { return x; }
static inline T to_print(T x) { return x; }
static inline device_type to_device(host_type x) { return x; }
};
template <>
@@ -56,6 +57,7 @@ struct TypeTraits<Vector<bin1_t, 32> > {
typedef uint32_t unsigned_type;
static inline uint32_t remove_negative_zero(uint32_t x) { return x; }
static inline uint32_t to_print(uint32_t x) { return x; }
static inline device_type to_device(host_type x) { return x; }
};
template <>
@@ -67,6 +69,7 @@ struct TypeTraits< Vector<int4_t, 8> > {
typedef uint32_t unsigned_type;
static inline uint32_t remove_negative_zero(uint32_t x) { return x; }
static inline uint32_t to_print(uint32_t x) { return x; }
static inline device_type to_device(host_type x) { return x; }
};
template <>
@@ -78,6 +81,7 @@ struct TypeTraits< Vector<uint4_t, 8> > {
typedef uint32_t unsigned_type;
static inline uint32_t remove_negative_zero(uint32_t x) { return x; }
static inline uint32_t to_print(uint32_t x) { return x; }
static inline device_type to_device(host_type x) { return x; }
};
template <>
@@ -89,6 +93,7 @@ struct TypeTraits<int8_t> {
typedef uint8_t unsigned_type;
static inline int8_t remove_negative_zero(int8_t x) { return x; }
static inline int to_print(int8_t x) { return (int)x; }
static inline device_type to_device(host_type x) { return x; }
};
template <>
@@ -100,6 +105,7 @@ struct TypeTraits<uint8_t> {
typedef uint8_t unsigned_type;
static inline uint8_t remove_negative_zero(uint8_t x) { return x; }
static inline uint32_t to_print(uint8_t x) { return (uint32_t)x; }
static inline device_type to_device(host_type x) { return x; }
};
template <>
@@ -111,6 +117,7 @@ struct TypeTraits<int> {
typedef uint32_t unsigned_type;
static inline int32_t remove_negative_zero(int32_t x) { return x; }
static inline int to_print(int x) { return x; }
static inline device_type to_device(host_type x) { return x; }
};
template <>
@@ -122,6 +129,7 @@ struct TypeTraits<unsigned> {
typedef uint32_t unsigned_type;
static inline uint32_t remove_negative_zero(uint32_t x) { return x; }
static inline uint32_t to_print(uint32_t x) { return x; }
static inline device_type to_device(host_type x) { return x; }
};
template <>
@@ -140,6 +148,7 @@ struct TypeTraits<half> {
return x;
}
static inline half to_print(half x) { return x; }
static inline device_type to_device(half x) { return reinterpret_cast<device_type const &>(x); }
};
template <>
@@ -151,6 +160,7 @@ struct TypeTraits<int64_t> {
typedef uint64_t unsigned_type;
static inline int64_t remove_negative_zero(int64_t x) { return x; }
static inline int64_t to_print(int64_t x) { return x; }
static inline device_type to_device(host_type x) { return x; }
};
template <>
@@ -162,6 +172,7 @@ struct TypeTraits<uint64_t> {
typedef uint64_t unsigned_type;
static inline uint64_t remove_negative_zero(uint64_t x) { return x; }
static inline uint64_t to_print(uint64_t x) { return x; }
static inline device_type to_device(host_type x) { return x; }
};
template <>
@@ -175,6 +186,7 @@ struct TypeTraits<cutlass::half_t> {
return (x.raw() == 0x8000 ? half_t::bitcast(0) : x);
}
static inline half_t to_print(half_t x) { return x; }
static inline device_type to_device(cutlass::half_t x) { return reinterpret_cast<device_type const &>(x); }
};
template <>
@@ -186,6 +198,7 @@ struct TypeTraits<float> {
typedef uint32_t unsigned_type;
static inline float remove_negative_zero(float x) { return x == -0.f ? 0.f : x; }
static inline float to_print(float x) { return x; }
static inline device_type to_device(host_type x) { return x; }
};
template <>
@@ -197,6 +210,7 @@ struct TypeTraits<double> {
typedef uint64_t unsigned_type;
static inline double remove_negative_zero(double x) { return x == -0.0 ? 0.0 : x; }
static inline double to_print(double x) { return x; }
static inline device_type to_device(host_type x) { return x; }
};
///////////////////////////////////////////////////////////////////////////////////////////////////
@@ -212,6 +226,7 @@ struct TypeTraits<platform::complex<half> > {
typedef platform::complex<half> device_type;
typedef int16_t integer_type;
typedef uint16_t unsigned_type;
static inline device_type to_device(platform::complex<half> x) { return reinterpret_cast<device_type const &>(x); }
};
template <>
@@ -228,6 +243,7 @@ struct TypeTraits<platform::complex<half_t> > {
);
}
static inline platform::complex<half_t> to_print(platform::complex<half_t> x) { return x; }
static inline device_type to_device(platform::complex<half_t> x) { return reinterpret_cast<device_type const &>(x); }
};
template <>
@@ -247,6 +263,7 @@ struct TypeTraits<platform::complex<float> > {
}
static inline platform::complex<float> to_print(platform::complex<float> x) { return x; }
static inline device_type to_device(platform::complex<float> x) { return reinterpret_cast<device_type const &>(x); }
};
template <>
@@ -263,6 +280,7 @@ struct TypeTraits<platform::complex<double> > {
);
}
static inline platform::complex<double> to_print(platform::complex<double> x) { return x; }
static inline device_type to_device(platform::complex<double> x) { return reinterpret_cast<device_type const &>(x); }
};
///////////////////////////////////////////////////////////////////////////////////////////////////