releaase 2.11 (#703)
This commit is contained in:
@@ -50,9 +50,11 @@
|
||||
#include "cutlass/util/reference/host/gemm.h"
|
||||
|
||||
#include "testbed_utils.h"
|
||||
#include "testbed_universal.h"
|
||||
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/matrix_coord.h"
|
||||
#include "cutlass/gemm/device/gemm_universal_adapter.h"
|
||||
|
||||
namespace test {
|
||||
namespace gemm {
|
||||
@@ -309,7 +311,7 @@ struct Testbed {
|
||||
throw std::runtime_error("cudaGetDeviceProperties() failed");
|
||||
}
|
||||
|
||||
if (properties.sharedMemPerMultiprocessor < smem_size) {
|
||||
if (properties.sharedMemPerBlockOptin < smem_size) {
|
||||
return false;
|
||||
}
|
||||
|
||||
@@ -319,10 +321,19 @@ struct Testbed {
|
||||
|
||||
/// Executes one test
|
||||
bool run(
|
||||
cutlass::gemm::GemmCoord problem_size,
|
||||
cutlass::gemm::GemmCoord problem_size,
|
||||
int split_k_slices = 1,
|
||||
ElementCompute alpha = ElementCompute(1),
|
||||
ElementCompute beta = ElementCompute(0)) {
|
||||
ElementCompute alpha = ElementCompute(1),
|
||||
ElementCompute beta = ElementCompute(0))
|
||||
{
|
||||
/*
|
||||
std::cout << "\n-----------------------\n";
|
||||
std::cout << "problem size: " << problem_size << "\n";
|
||||
std::cout << "split_k_slices: " << split_k_slices << "\n";
|
||||
std::cout << "alpha: " << alpha << "\n";
|
||||
std::cout << "beta: " << beta << "\n";
|
||||
std::cout << "-----------------------\n\n";
|
||||
*/
|
||||
|
||||
// Waive test if insufficient CUDA device
|
||||
if (!sufficient()) {
|
||||
@@ -387,7 +398,7 @@ struct Testbed {
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename Gemm, bool Relu=false>
|
||||
bool TestAllGemm(
|
||||
bool TestAllGemmBasic(
|
||||
const typename Gemm::LayoutA::Stride& stride_factor_A = typename Gemm::LayoutA::Stride(),
|
||||
const typename Gemm::LayoutB::Stride& stride_factor_B = typename Gemm::LayoutB::Stride(),
|
||||
const typename Gemm::LayoutC::Stride& stride_factor_C = typename Gemm::LayoutC::Stride()) {
|
||||
@@ -477,6 +488,52 @@ bool TestAllGemm(
|
||||
return passed;
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename Gemm, bool Relu=false>
|
||||
bool TestAllGemm(
|
||||
const typename Gemm::LayoutA::Stride& stride_factor_A,
|
||||
const typename Gemm::LayoutB::Stride& stride_factor_B = typename Gemm::LayoutB::Stride(),
|
||||
const typename Gemm::LayoutC::Stride& stride_factor_C = typename Gemm::LayoutC::Stride())
|
||||
{
|
||||
// Test basic GEMM with non-default stride factors
|
||||
return TestAllGemmBasic<Gemm, Relu>(stride_factor_A, stride_factor_B, stride_factor_C);
|
||||
}
|
||||
|
||||
template <typename Gemm, bool Relu=false>
|
||||
bool TestAllGemm()
|
||||
{
|
||||
#ifdef NDEBUG
|
||||
// Non-debug builds also test basic GEMM with default stride factors
|
||||
if (!TestAllGemmBasic<Gemm, Relu>()) {
|
||||
return false;
|
||||
}
|
||||
#endif // NDEBUG
|
||||
|
||||
// Test universal GEMM
|
||||
#if 0
|
||||
// Define the universal kernel
|
||||
using UniversalKernel = cutlass::gemm::kernel::GemmUniversal<
|
||||
typename Gemm::GemmKernel::Mma, // Mma
|
||||
typename Gemm::GemmKernel::Epilogue, // Epilogue
|
||||
cutlass::gemm::threadblock::GemmIdentityThreadblockSwizzle<> // ThreadblockSwizzle
|
||||
>;
|
||||
#else
|
||||
// Define the streamk universal kernel
|
||||
using UniversalKernel = cutlass::gemm::kernel::GemmUniversalStreamk<
|
||||
typename Gemm::GemmKernel::Mma, // Mma
|
||||
typename Gemm::GemmKernel::Epilogue, // Epilogue
|
||||
cutlass::gemm::threadblock::ThreadblockSwizzleStreamK // ThreadblockSwizzle
|
||||
>;
|
||||
#endif
|
||||
|
||||
// Define the universal adaptor
|
||||
using UniversalGemm = cutlass::gemm::device::GemmUniversalAdapter<UniversalKernel>;
|
||||
|
||||
// Test universal GEMM
|
||||
return TestAllGemmUniversal<UniversalGemm, Relu>();
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
template <typename Gemm>
|
||||
bool TestGemmPerf(int iterations = 1) {
|
||||
|
||||
Reference in New Issue
Block a user