CUTLASS 2.0 (#62)

CUTLASS 2.0

Substantially refactored for

- Better performance, particularly for native Turing Tensor Cores
- Robust and durable templates spanning the design space
- Encapsulated functionality embodying modern C++11 programming techniques
- Optimized containers and data types for efficient, generic, portable device code

Updates to:
- Quick start guide
- Documentation
- Utilities
- CUTLASS Profiler

Native Turing Tensor Cores
- Efficient GEMM kernels targeting Turing Tensor Cores
- Mixed-precision floating point, 8-bit integer, 4-bit integer, and binarized operands

Coverage of existing CUTLASS functionality:
- GEMM kernels targeting CUDA and Tensor Cores in NVIDIA GPUs
- Volta Tensor Cores through native mma.sync and through WMMA API
- Optimizations such as parallel reductions, threadblock rasterization, and intra-threadblock reductions
- Batched GEMM operations
- Complex-valued GEMMs

Note: this commit and all that follow require a host compiler supporting C++11 or greater.
This commit is contained in:
Andrew Kerr
2019-11-19 16:55:34 -08:00
committed by GitHub
parent b5cab177a9
commit fb335f6a5f
5434 changed files with 599799 additions and 250176 deletions
+33
View File
@@ -0,0 +1,33 @@
# Copyright (c) 2017-2019, 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.
cutlass_test_unit_add_executable(
cutlass_test_unit_gemm_thread
gemm_sm50.cu
gemm_sm60.cu
gemm_sm61.cu
testbed.h
)
add_subdirectory(host)
add_dependencies(test_unit_gemm_thread test_unit_gemm_thread_host)
+169
View File
@@ -0,0 +1,169 @@
/***************************************************************************************************
* Copyright (c) 2017-2019, 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 Unit tests for thread-level GEMM
*/
#include "../../common/cutlass_unit_test.h"
#include "cutlass/gemm/thread/mma.h"
#include "testbed.h"
/////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SM50_Sgemm_thread, col_row_3x4x2) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<3, 4, 2>,
float,
cutlass::layout::ColumnMajor,
float,
cutlass::layout::RowMajor,
float,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM50_Sgemm_thread, col_row_4x4x2) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<4, 4, 2>,
float,
cutlass::layout::ColumnMajor,
float,
cutlass::layout::RowMajor,
float,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM50_Sgemm_thread, row_col_4x4x2) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<4, 4, 2>,
float,
cutlass::layout::RowMajor,
float,
cutlass::layout::ColumnMajor,
float,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM50_Sgemm_thread, col_row_4x5x3) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<4, 5, 3>,
float,
cutlass::layout::ColumnMajor,
float,
cutlass::layout::RowMajor,
float,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM50_Sgemm_thread, col_row) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<8, 8, 1>,
float,
cutlass::layout::ColumnMajor,
float,
cutlass::layout::RowMajor,
float,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM50_Sgemm_thread, row_col) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<8, 8, 1>,
float,
cutlass::layout::RowMajor,
float,
cutlass::layout::ColumnMajor,
float,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM50_Sgemm_thread, col_col) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<8, 8, 1>,
float,
cutlass::layout::ColumnMajor,
float,
cutlass::layout::ColumnMajor,
float,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM50_Sgemm_thread, row_row) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<8, 8, 1>,
float,
cutlass::layout::RowMajor,
float,
cutlass::layout::RowMajor,
float,
cutlass::layout::ColumnMajor
>().run();
}
/////////////////////////////////////////////////////////////////////////////////////////////////
TEST(SM50_Dgemm_thread, col_row) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<8, 8, 1>,
double,
cutlass::layout::ColumnMajor,
double,
cutlass::layout::RowMajor,
double,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM50_Dgemm_thread, row_col) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<8, 8, 1>,
double,
cutlass::layout::RowMajor,
double,
cutlass::layout::ColumnMajor,
double,
cutlass::layout::ColumnMajor
>().run();
}
/////////////////////////////////////////////////////////////////////////////////////////////////
+493
View File
@@ -0,0 +1,493 @@
/***************************************************************************************************
* Copyright (c) 2017-2019, 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 Unit tests for thread-level GEMM
*/
#include "../../common/cutlass_unit_test.h"
#include "cutlass/gemm/thread/mma.h"
#include "testbed.h"
/////////////////////////////////////////////////////////////////////////////////////////////////
//
// Compute capability SM60
//
TEST(SM60_Hgemm_thread, col_row_col_1x1x16) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<1, 1, 16>,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM60_Hgemm_thread, row_col_row_1x1x16) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<1, 1, 16>,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::RowMajor
>().run();
}
TEST(SM60_Hgemm_thread, row_col_col_1x3x8) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<1, 3, 8>,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM60_Hgemm_thread, row_row_row_7x8x3) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<7, 8, 3>,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::RowMajor
>().run();
}
TEST(SM60_Hgemm_thread, row_col_row_7x8x3) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<7, 8, 3>,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::RowMajor
>().run();
}
TEST(SM60_Hgemm_thread, col_row_row_7x8x3) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<7, 8, 3>,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::RowMajor
>().run();
}
TEST(SM60_Hgemm_thread, col_col_row_7x8x3) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<7, 8, 3>,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::RowMajor
>().run();
}
TEST(SM60_Hgemm_thread, row_row_row_7x8x4) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<7, 8, 4>,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::RowMajor
>().run();
}
TEST(SM60_Hgemm_thread, row_col_row_7x8x4) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<7, 8, 4>,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::RowMajor
>().run();
}
TEST(SM60_Hgemm_thread, col_row_row_7x8x4) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<7, 8, 4>,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::RowMajor
>().run();
}
TEST(SM60_Hgemm_thread, col_col_row_7x8x4) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<7, 8, 4>,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::RowMajor
>().run();
}
TEST(SM60_Hgemm_thread, row_row_col_16x3x3) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<16, 3, 3>,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM60_Hgemm_thread, row_col_col_16x3x3) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<16, 3, 3>,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM60_Hgemm_thread, col_row_col_16x3x3) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<16, 3, 3>,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM60_Hgemm_thread, col_col_col_16x3x3) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<16, 3, 3>,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM60_Hgemm_thread, row_row_col_16x3x4) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<16, 3, 4>,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM60_Hgemm_thread, row_col_col_16x3x4) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<16, 3, 4>,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM60_Hgemm_thread, col_row_col_16x3x4) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<16, 3, 4>,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM60_Hgemm_thread, col_col_col_16x3x4) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<16, 3, 4>,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM60_Hgemm_thread, row_row_row_16x8x3) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<16, 8, 3>,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::RowMajor
>().run();
}
TEST(SM60_Hgemm_thread, row_row_col_16x8x3) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<16, 8, 3>,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM60_Hgemm_thread, row_col_row_16x8x3) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<16, 8, 3>,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::RowMajor
>().run();
}TEST(SM60_Hgemm_thread, row_col_col_16x8x3) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<16, 8, 3>,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM60_Hgemm_thread, col_row_row_16x8x3) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<16, 8, 3>,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::RowMajor
>().run();
}
TEST(SM60_Hgemm_thread, col_row_col_16x8x3) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<16, 8, 3>,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM60_Hgemm_thread, col_col_row_16x8x3) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<16, 8, 3>,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::RowMajor
>().run();
}
TEST(SM60_Hgemm_thread, col_col_col_16x8x3) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<16, 8, 3>,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM60_Hgemm_thread, row_row_row_16x8x4) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<16, 8, 4>,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::RowMajor
>().run();
}
TEST(SM60_Hgemm_thread, row_row_col_16x8x4) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<16, 8, 4>,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM60_Hgemm_thread, row_col_row_16x8x4) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<16, 8, 4>,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::RowMajor
>().run();
}
TEST(SM60_Hgemm_thread, row_col_col_16x8x4) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<16, 8, 4>,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM60_Hgemm_thread, col_row_row_16x8x4) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<16, 8, 4>,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::RowMajor
>().run();
}
TEST(SM60_Hgemm_thread, col_row_col_16x8x4) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<16, 8, 4>,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM60_Hgemm_thread, col_col_row_16x8x4) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<16, 8, 4>,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::RowMajor
>().run();
}
TEST(SM60_Hgemm_thread, col_col_col_16x8x4) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<16, 8, 4>,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor
>().run();
}
/////////////////////////////////////////////////////////////////////////////////////////////////
+81
View File
@@ -0,0 +1,81 @@
/***************************************************************************************************
* Copyright (c) 2017-2019, 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 Unit tests for thread-level GEMM
*/
#include "../../common/cutlass_unit_test.h"
#include "cutlass/gemm/thread/mma.h"
#include "testbed.h"
/////////////////////////////////////////////////////////////////////////////////////////////////
/////////////////////////////////////////////////////////////////////////////////////////////////
//
// Compute capability SM61
//
TEST(SM61_Igemm_thread, col_row_1x1x4) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<1, 1, 4>,
int8_t,
cutlass::layout::RowMajor,
int8_t,
cutlass::layout::ColumnMajor,
int32_t,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM61_Igemm_thread, col_row_2x3x4) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<2, 3, 4>,
int8_t,
cutlass::layout::RowMajor,
int8_t,
cutlass::layout::ColumnMajor,
int32_t,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM61_Igemm_thread, col_row_8x8x4) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<8, 8, 4>,
int8_t,
cutlass::layout::RowMajor,
int8_t,
cutlass::layout::ColumnMajor,
int32_t,
cutlass::layout::ColumnMajor
>().run();
}
/////////////////////////////////////////////////////////////////////////////////////////////////
+27
View File
@@ -0,0 +1,27 @@
# Copyright (c) 2017-2019, 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.
cutlass_test_unit_add_executable(
cutlass_test_unit_gemm_thread_host
gemm_sm60_host.cu
testbed_host.h
)
@@ -0,0 +1,170 @@
/***************************************************************************************************
* Copyright (c) 2017-2019, 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 Unit tests for thread-level GEMM
*/
#include "../../../common/cutlass_unit_test.h"
#include "cutlass/gemm/thread/mma.h"
#include "testbed_host.h"
/////////////////////////////////////////////////////////////////////////////////////////////////
//
// Compute capability SM60
//
TEST(SM60_host_Hgemm_thread, col_row_col_1x1x16) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<1, 1, 16>,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM60_host_Hgemm_thread, row_col_row_1x1x16) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<1, 1, 16>,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::RowMajor
>().run();
}
TEST(SM60_host_Hgemm_thread, row_row_row_2x2x2) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<2, 2, 2>,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::RowMajor
>().run();
}
TEST(SM60_host_Hgemm_thread, row_row_col_2x2x2) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<2, 2, 2>,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM60_host_Hgemm_thread, row_col_row_2x2x2) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<2, 2, 2>,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::RowMajor
>().run();
}
TEST(SM60_host_Hgemm_thread, row_col_col_2x2x2) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<2, 2, 2>,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM60_host_Hgemm_thread, col_row_row_2x2x2) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<2, 2, 2>,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::RowMajor
>().run();
}
TEST(SM60_host_Hgemm_thread, col_row_col_2x2x2) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<2, 2, 2>,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::RowMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor
>().run();
}
TEST(SM60_host_Hgemm_thread, col_col_row_2x2x2) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<2, 2, 2>,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::RowMajor
>().run();
}
TEST(SM60_host_Hgemm_thread, col_col_col_2x2x2) {
test::gemm::thread::Testbed<
cutlass::gemm::GemmShape<2, 2, 2>,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor,
cutlass::half_t,
cutlass::layout::ColumnMajor
>().run();
}
/////////////////////////////////////////////////////////////////////////////////////////////////
+226
View File
@@ -0,0 +1,226 @@
/***************************************************************************************************
* Copyright (c) 2017-2019, 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 Unit tests for thread-level GEMM
*/
#pragma once
#include "cutlass/gemm/thread/mma.h"
#include "cutlass/layout/vector.h"
#include "cutlass/util/host_tensor.h"
#include "cutlass/util/tensor_view_io.h"
#include "cutlass/util/reference/host/tensor_copy.h"
#include "cutlass/util/reference/host/tensor_fill.h"
#include "cutlass/util/reference/host/tensor_compare.h"
#include "cutlass/util/reference/host/gemm.h"
namespace test {
namespace gemm {
namespace thread {
/////////////////////////////////////////////////////////////////////////////////////////////////
/// Thread-level matrix multiply-accumulate
template <typename Mma>
void kernel(
typename Mma::ElementC *D,
typename Mma::ElementA const *A,
typename Mma::ElementB const *B,
typename Mma::ElementC const *C) {
auto ptr_D = reinterpret_cast<cutlass::Array<typename Mma::ElementC, Mma::Shape::kMN> *>(D);
auto ptr_A = reinterpret_cast<cutlass::Array<typename Mma::ElementA, Mma::Shape::kMK> const *>(A);
auto ptr_B = reinterpret_cast<cutlass::Array<typename Mma::ElementB, Mma::Shape::kKN> const *>(B);
auto ptr_C = reinterpret_cast<cutlass::Array<typename Mma::ElementC, Mma::Shape::kMN> const *>(C);
Mma mma;
auto a = *ptr_A;
auto b = *ptr_B;
auto c = *ptr_C;
using Btype = typename Mma::ElementB;
cutlass::Array<typename Mma::ElementC, Mma::Shape::kMN> d;
mma(d, a, b, c);
*ptr_D = d;
}
/////////////////////////////////////////////////////////////////////////////////////////////////
/// Structure to compute the matrix product
template <
/// Size of the Gemm problem - concept: gemm::GemmShape<>
typename Shape,
/// Data type of A elements
typename ElementA,
/// Layout of A matrix (concept: MatrixLayout)
typename LayoutA,
/// Data type of B elements
typename ElementB,
/// Layout of B matrix (concept: MatrixLayout)
typename LayoutB,
/// Element type of C matrix
typename ElementC,
/// Layout of C matrix (concept: MatrixLayout)
typename LayoutC
>
struct Testbed {
/// Thread-level matrix multiply-accumulate operator
using Mma = cutlass::gemm::thread::Mma<
Shape,
ElementA,
LayoutA,
ElementB,
LayoutB,
ElementC,
LayoutC
>;
//
// Data members
//
cutlass::HostTensor<ElementA, LayoutA> tensor_A;
cutlass::HostTensor<ElementB, LayoutB> tensor_B;
cutlass::HostTensor<ElementC, LayoutC> tensor_C;
cutlass::HostTensor<ElementC, LayoutC> tensor_D_computed;
cutlass::HostTensor<ElementC, LayoutC> tensor_D_reference;
//
// Methods
//
/// Allocates workspace in device memory
Testbed() {
tensor_A.reset(cutlass::make_Coord(Shape::kM, Shape::kK), false);
tensor_B.reset(cutlass::make_Coord(Shape::kK, Shape::kN), false);
tensor_C.reset(cutlass::make_Coord(Shape::kM, Shape::kN), false);
tensor_D_computed.reset(cutlass::make_Coord(Shape::kM, Shape::kN), false);
tensor_D_reference.reset(cutlass::make_Coord(Shape::kM, Shape::kN), false);
}
/// Runs the test
bool run() {
//
// initialize device memory
//
cutlass::reference::host::detail::RandomUniformFunc< ElementA > tfill_rand_func(
0, // seed
10, // max
0, // min
0); // bits after decimal
cutlass::reference::host::detail::TensorFillRandomUniformFunc< ElementA, LayoutA > tfill_rand(
tensor_A.host_view(),
tfill_rand_func);
for (auto i=0; i< Shape::kM; i++)
for (auto j=0; j< Shape::kK; j++)
tfill_rand(cutlass::make_Coord(i,j));
cutlass::reference::host::BlockFillSequential(
tensor_B.host_data(),
tensor_B.capacity(),
ElementB(1),
ElementB(2)
);
cutlass::reference::host::TensorFill(
tensor_C.host_view(),
ElementC(0)
);
cutlass::reference::host::TensorFill(
tensor_D_computed.host_view(),
ElementC(0)
);
cutlass::reference::host::TensorFill(
tensor_D_reference.host_view(),
ElementC(0)
);
// Host side call
kernel<Mma>(
tensor_D_computed.host_data(),
tensor_A.host_data(),
tensor_B.host_data(),
tensor_C.host_data());
//
// Reference implementation
//
cutlass::reference::host::Gemm<ElementA, LayoutA, ElementB, LayoutB,
ElementC, LayoutC, ElementC, ElementC>
reference_gemm;
reference_gemm(
{Shape::kM, Shape::kN, Shape::kK},
ElementC(1),
tensor_A.host_ref(),
tensor_B.host_ref(),
ElementC(0),
tensor_D_reference.host_ref()
);
//
// Verify equivalence
//
// compare
bool passed = cutlass::reference::host::TensorEquals(
tensor_D_computed.host_view(),
tensor_D_reference.host_view()
);
EXPECT_TRUE(passed)
<< "A:\n" << tensor_A.host_view() << "\n\n"
<< "B:\n" << tensor_B.host_view() << "\n\n"
<< "C:\n" << tensor_C.host_view() << "\n\n"
<< "Reference:\n" << tensor_D_reference.host_view() << "\n\n"
<< "Computed:\n" << tensor_D_computed.host_view() << std::endl;
return passed;
}
};
/////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace thread
} // namespace gemm
} // namespace test
+230
View File
@@ -0,0 +1,230 @@
/***************************************************************************************************
* Copyright (c) 2017-2019, 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 Unit tests for thread-level GEMM
*/
#pragma once
#include "cutlass/gemm/thread/mma.h"
#include "cutlass/util/host_tensor.h"
#include "cutlass/util/tensor_view_io.h"
#include "cutlass/util/reference/host/tensor_copy.h"
#include "cutlass/util/reference/host/tensor_fill.h"
#include "cutlass/util/reference/host/tensor_compare.h"
#include "cutlass/util/reference/host/gemm.h"
namespace test {
namespace gemm {
namespace thread {
/////////////////////////////////////////////////////////////////////////////////////////////////
/// Thread-level matrix multiply-accumulate
template <typename Mma>
__global__ void kernel(
typename Mma::ElementC *D,
typename Mma::ElementA const *A,
typename Mma::ElementB const *B,
typename Mma::ElementC const *C) {
auto ptr_D = reinterpret_cast<cutlass::Array<typename Mma::ElementC, Mma::Shape::kMN> *>(D);
auto ptr_A = reinterpret_cast<cutlass::Array<typename Mma::ElementA, Mma::Shape::kMK> const *>(A);
auto ptr_B = reinterpret_cast<cutlass::Array<typename Mma::ElementB, Mma::Shape::kKN> const *>(B);
auto ptr_C = reinterpret_cast<cutlass::Array<typename Mma::ElementC, Mma::Shape::kMN> const *>(C);
Mma mma;
auto a = *ptr_A;
auto b = *ptr_B;
auto c = *ptr_C;
cutlass::Array<typename Mma::ElementC, Mma::Shape::kMN> d;
mma(d, a, b, c);
*ptr_D = d;
}
/////////////////////////////////////////////////////////////////////////////////////////////////
/// Structure to compute the matrix product
template <
/// Size of the Gemm problem - concept: gemm::GemmShape<>
typename Shape,
/// Data type of A elements
typename ElementA,
/// Layout of A matrix (concept: MatrixLayout)
typename LayoutA,
/// Data type of B elements
typename ElementB,
/// Layout of B matrix (concept: MatrixLayout)
typename LayoutB,
/// Element type of C matrix
typename ElementC,
/// Layout of C matrix (concept: MatrixLayout)
typename LayoutC
>
struct Testbed {
/// Thread-level matrix multiply-accumulate operator
using Mma = cutlass::gemm::thread::Mma<
Shape,
ElementA,
LayoutA,
ElementB,
LayoutB,
ElementC,
LayoutC
>;
//
// Data members
//
cutlass::HostTensor<ElementA, LayoutA> tensor_A;
cutlass::HostTensor<ElementB, LayoutB> tensor_B;
cutlass::HostTensor<ElementC, LayoutC> tensor_C;
cutlass::HostTensor<ElementC, LayoutC> tensor_D_computed;
cutlass::HostTensor<ElementC, LayoutC> tensor_D_reference;
//
// Methods
//
/// Allocates workspace in device memory
Testbed() {
tensor_A.reset(cutlass::make_Coord(Shape::kM, Shape::kK));
tensor_B.reset(cutlass::make_Coord(Shape::kK, Shape::kN));
tensor_C.reset(cutlass::make_Coord(Shape::kM, Shape::kN));
tensor_D_computed.reset(cutlass::make_Coord(Shape::kM, Shape::kN));
tensor_D_reference.reset(cutlass::make_Coord(Shape::kM, Shape::kN), false);
}
/// Runs the test
bool run() {
//
// initialize device memory
//
cutlass::reference::host::BlockFillSequential(
tensor_A.host_data(),
tensor_A.capacity()
);
cutlass::reference::host::BlockFillSequential(
tensor_B.host_data(),
tensor_B.capacity(),
ElementB(1),
ElementB(2)
);
cutlass::reference::host::TensorFill(
tensor_C.host_view(),
ElementC(0)
);
cutlass::reference::host::TensorFill(
tensor_D_computed.host_view(),
ElementC(0)
);
cutlass::reference::host::TensorFill(
tensor_D_reference.host_view(),
ElementC(0)
);
tensor_A.sync_device();
tensor_B.sync_device();
tensor_C.sync_device();
tensor_D_computed.sync_device();
// launch kernel
kernel<Mma><<< dim3(1, 1), dim3(1, 1, 1) >>>(
tensor_D_computed.device_data(),
tensor_A.device_data(),
tensor_B.device_data(),
tensor_C.device_data());
// verify no errors
cudaError_t result = cudaDeviceSynchronize();
EXPECT_EQ(result, cudaSuccess) << "CUDA ERROR: " << cudaGetErrorString(result);
if (result != cudaSuccess) {
return false;
}
tensor_D_computed.sync_host();
//
// Reference implementation
//
//tensor_D_reference.fill(tensor_C.host_view());
cutlass::reference::host::Gemm<ElementA, LayoutA, ElementB, LayoutB,
ElementC, LayoutC, ElementC, ElementC>
reference_gemm;
reference_gemm(
{Shape::kM, Shape::kN, Shape::kK},
ElementC(1),
tensor_A.host_ref(),
tensor_B.host_ref(),
ElementC(0),
tensor_D_reference.host_ref()
);
//
// Verify equivalence
//
// compare
bool passed = cutlass::reference::host::TensorEquals(
tensor_D_computed.host_view(),
tensor_D_reference.host_view()
);
EXPECT_TRUE(passed)
<< "A:\n" << tensor_A.host_view() << "\n\n"
<< "B:\n" << tensor_B.host_view() << "\n\n"
<< "C:\n" << tensor_C.host_view() << "\n\n"
<< "Reference:\n" << tensor_D_reference.host_view() << "\n\n"
<< "Computed:\n" << tensor_D_computed.host_view() << std::endl;
return passed;
}
};
/////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace thread
} // namespace gemm
} // namespace test