#include #include #include #include #include #define CHECK_CUDA(call) do { if ((call) != cudaSuccess) std::exit(2); } while (0) #define CHECK_BLAS(call) do { if ((call) != CUBLAS_STATUS_SUCCESS) std::exit(3); } while (0) int main() { constexpr int m = 128, n = 128, k = 128; constexpr size_t a_bytes = m * k * sizeof(float); constexpr size_t b_bytes = k * n * sizeof(float); constexpr size_t c_bytes = m * n * sizeof(float); float *a = nullptr, *b = nullptr, *c = nullptr; CHECK_CUDA(cudaMalloc(&a, a_bytes)); CHECK_CUDA(cudaMalloc(&b, b_bytes)); CHECK_CUDA(cudaMalloc(&c, c_bytes)); CHECK_CUDA(cudaMemset(a, 0, a_bytes)); CHECK_CUDA(cudaMemset(b, 0, b_bytes)); const float alpha = 1.0F, beta = 0.0F; cublasHandle_t blas{}; CHECK_BLAS(cublasCreate(&blas)); CHECK_BLAS(cublasSgemm(blas, CUBLAS_OP_N, CUBLAS_OP_N, m, n, k, &alpha, a, m, b, k, &beta, c, m)); CHECK_CUDA(cudaDeviceSynchronize()); CHECK_BLAS(cublasDestroy(blas)); cublasLtHandle_t lt{}; cublasLtMatmulDesc_t operation{}; cublasLtMatrixLayout_t a_layout{}, b_layout{}, c_layout{}; cublasLtMatmulPreference_t preference{}; CHECK_BLAS(cublasLtCreate(<)); CHECK_BLAS(cublasLtMatmulDescCreate(&operation, CUBLAS_COMPUTE_32F, CUDA_R_32F)); CHECK_BLAS(cublasLtMatrixLayoutCreate(&a_layout, CUDA_R_32F, m, k, m)); CHECK_BLAS(cublasLtMatrixLayoutCreate(&b_layout, CUDA_R_32F, k, n, k)); CHECK_BLAS(cublasLtMatrixLayoutCreate(&c_layout, CUDA_R_32F, m, n, m)); CHECK_BLAS(cublasLtMatmulPreferenceCreate(&preference)); constexpr size_t workspace_bytes = 4 << 20; void* workspace = nullptr; CHECK_CUDA(cudaMalloc(&workspace, workspace_bytes)); CHECK_BLAS(cublasLtMatmulPreferenceSetAttribute( preference, CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES, &workspace_bytes, sizeof(workspace_bytes))); cublasLtMatmulHeuristicResult_t heuristic{}; int returned = 0; CHECK_BLAS(cublasLtMatmulAlgoGetHeuristic( lt, operation, a_layout, b_layout, c_layout, c_layout, preference, 1, &heuristic, &returned)); if (returned == 0) { std::fprintf(stderr, "no cuBLASLt heuristic\n"); return 4; } CHECK_BLAS(cublasLtMatmul( lt, operation, &alpha, a, a_layout, b, b_layout, &beta, c, c_layout, c, c_layout, &heuristic.algo, workspace, workspace_bytes, nullptr)); CHECK_CUDA(cudaDeviceSynchronize()); std::printf("cublas_baseline=ok cublaslt_heuristic=ok workspace_bytes=%zu\n", workspace_bytes); CHECK_CUDA(cudaFree(workspace)); CHECK_BLAS(cublasLtMatmulPreferenceDestroy(preference)); CHECK_BLAS(cublasLtMatrixLayoutDestroy(c_layout)); CHECK_BLAS(cublasLtMatrixLayoutDestroy(b_layout)); CHECK_BLAS(cublasLtMatrixLayoutDestroy(a_layout)); CHECK_BLAS(cublasLtMatmulDescDestroy(operation)); CHECK_BLAS(cublasLtDestroy(lt)); CHECK_CUDA(cudaFree(c)); CHECK_CUDA(cudaFree(b)); CHECK_CUDA(cudaFree(a)); }