CUTLASS 2.10 (#615)

Co-authored-by: Aniket Shivam <ashivam@nvidia.com>
This commit is contained in:
ANIKET SHIVAM
2022-09-03 18:48:46 -04:00
committed by GitHub
co-authored by Aniket Shivam
parent ca23ff7924
commit b72cbf957d
289 changed files with 43708 additions and 2513 deletions
+16 -31
View File
@@ -417,46 +417,27 @@ struct TestbedGrouped {
return passed;
}
/// Returns the number of threadblocks to launch if the kernel can run on the target
/// device. Otherwise, returns zero.
int sufficient() const {
cudaDeviceProp properties;
int device_idx;
cudaError_t result = cudaGetDevice(&device_idx);
if (result != cudaSuccess) {
throw std::runtime_error("cudaGetDevice() API call failed.");
}
result = cudaGetDeviceProperties(&properties, device_idx);
if (result != cudaSuccess) {
throw std::runtime_error("cudaGetDeviceProperties() failed");
}
int occupancy = Gemm::maximum_active_blocks();
return properties.multiProcessorCount * occupancy;
}
/// Executes one test
bool run(
int problem_count,
ElementCompute alpha = ElementCompute(1),
ElementCompute beta = ElementCompute(0)) {
int threadblock_count = sufficient();
// Early exit
if (!threadblock_count) {
return false;
}
this->problem_count = problem_count;
// Initialize the problem
initialize();
int threadblock_count = Gemm::sufficient(problem_sizes_host.data(), problem_count);
// Early exit
if (!threadblock_count) {
if (CUTLASS_TEST_UNIT_ENABLE_WARNINGS) {
std::cerr << "Test waived due to insufficient CUDA device resources." << std::endl;
}
return true;
}
// Configure the GEMM arguments
typename EpilogueOutputOp::Params epilogue_op(alpha, beta);
@@ -473,13 +454,17 @@ struct TestbedGrouped {
lda.get(),
ldb.get(),
ldc.get(),
ldd.get()
ldd.get(),
problem_sizes_host.data()
);
// Initialize the GEMM object
Gemm gemm;
cutlass::Status status = gemm.initialize(args);
size_t workspace_size = gemm.get_workspace_size(args);
cutlass::DeviceAllocation<uint8_t> workspace(workspace_size);
cutlass::Status status = gemm.initialize(args, workspace.get());
if (status != cutlass::Status::kSuccess) {
return false;