22#pragma once
33
44#include < cuda_fp16.h>
5+ #ifndef __HIP_PLATFORM_HCC__
56#include < cuda_profiler_api.h>
7+ #endif
68#include < array>
79#include < cstdio>
810#include < cstdlib>
@@ -58,7 +60,11 @@ class GemmTest {
5860 B,
5961 A,
6062 C,
63+ #ifdef __HIP_PLATFORM_HCC__
64+ static_cast <rocblas_gemm_algo>(algo));
65+ #else
6166 static_cast <cublasGemmAlgo_t>(algo));
67+ #endif
6268 });
6369
6470 int algo_bw1 = Run (loops, [=](int algo) {
@@ -73,7 +79,11 @@ class GemmTest {
7379 A,
7480 C,
7581 B,
82+ #ifdef __HIP_PLATFORM_HCC__
83+ static_cast <rocblas_gemm_algo>(algo));
84+ #else
7685 static_cast <cublasGemmAlgo_t>(algo));
86+ #endif
7787 });
7888
7989 int algo_bw2 = Run (loops, [=](int algo) {
@@ -88,7 +98,11 @@ class GemmTest {
8898 B,
8999 C,
90100 A,
101+ #ifdef __HIP_PLATFORM_HCC__
102+ static_cast <rocblas_gemm_algo>(algo));
103+ #else
91104 static_cast <cublasGemmAlgo_t>(algo));
105+ #endif
92106 });
93107
94108 return std::array<int , 3 >({algo_fw, algo_bw1, algo_bw2});
@@ -100,8 +114,12 @@ class GemmTest {
100114 float fast_latency = (std::numeric_limits<float >::max)();
101115 int fast_algo = 0 ;
102116
117+ #ifdef __HIP_PLATFORM_HCC__
118+ for (int algo = (int )rocblas_gemm_algo_standard; algo <= (int )rocblas_gemm_algo_standard;
119+ #else
103120 for (int algo = (int )CUBLAS_GEMM_DEFAULT_TENSOR_OP ;
104121 algo <= (int )CUBLAS_GEMM_ALGO15_TENSOR_OP ;
122+ #endif
105123 algo++) {
106124 int warm_up = 5 ;
107125 for (int i = 0 ; i < warm_up; ++i) f (algo);
@@ -186,7 +204,11 @@ class StridedGemmTest {
186204 stride_b,
187205 stride_c,
188206 bsz,
207+ #ifdef __HIP_PLATFORM_HCC__
208+ static_cast <rocblas_gemm_algo>(algo));
209+ #else
189210 static_cast <cublasGemmAlgo_t>(algo));
211+ #endif
190212 });
191213
192214 int algo_bw1 = Run (loops, [=](int algo) {
@@ -216,7 +238,11 @@ class StridedGemmTest {
216238 stride_b,
217239 stride_c,
218240 bsz,
241+ #ifdef __HIP_PLATFORM_HCC__
242+ static_cast <rocblas_gemm_algo>(algo));
243+ #else
219244 static_cast <cublasGemmAlgo_t>(algo));
245+ #endif
220246 });
221247
222248 int algo_bw2 = Run (loops, [=](int algo) {
@@ -243,7 +269,11 @@ class StridedGemmTest {
243269 stride_b,
244270 stride_c,
245271 bsz,
272+ #ifdef __HIP_PLATFORM_HCC__
273+ static_cast <rocblas_gemm_algo>(algo));
274+ #else
246275 static_cast <cublasGemmAlgo_t>(algo));
276+ #endif
247277 });
248278
249279 return std::array<int , 3 >({algo_fw, algo_bw1, algo_bw2});
@@ -255,8 +285,12 @@ class StridedGemmTest {
255285 float fast_latency = (std::numeric_limits<float >::max)();
256286 int fast_algo = 0 ;
257287
288+ #ifdef __HIP_PLATFORM_HCC__
289+ for (int algo = (int )rocblas_gemm_algo_standard; algo <= (int )rocblas_gemm_algo_standard;
290+ #else
258291 for (int algo = (int )CUBLAS_GEMM_DEFAULT_TENSOR_OP ;
259292 algo <= (int )CUBLAS_GEMM_ALGO15_TENSOR_OP ;
293+ #endif
260294 algo++) {
261295 int warm_up = 5 ;
262296 for (int i = 0 ; i < warm_up; ++i) f (algo);
0 commit comments